尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

如何加载neural-backed-decision-trees预训练模型?30行代码解析SoftNBDT推理全流程

如何加载neural-backed-decision-trees预训练模型?30行代码解析SoftNBDT推理全流程 如何加载neural-backed-decision-trees预训练模型30行代码解析SoftNBDT推理全流程【免费下载链接】neural-backed-decision-treesMaking decision trees competitive with neural networks on CIFAR10, CIFAR100, TinyImagenet200, Imagenet项目地址: https://gitcode.com/gh_mirrors/ne/neural-backed-decision-trees NBDTneural-backed-decision-trees神经网络支撑决策树是一个让决策树在精度上正面对抗深度神经网络的开源项目在 CIFAR10、CIFAR100、TinyImagenet200 乃至 ImageNet 上均达到或超过 SOTA 神经网络水平同时提供人类可读的推理路径。本文带你用30 行代码加载 NBDT 预训练模型并逐行拆解SoftNBDT的推理全流程小白也能一次跑通图像分类。一、为什么值得关注可解释的准神经网络NBDT 的核心思想是推理时走决策树图像先由神经网络骨干如 WideResNet28提取特征再由内嵌决策规则沿层级结构动物 → 哺乳动物 → 猫逐层判断精度不打折CIFAR10 达到 97.55%、CIFAR100 达到 82.97%、ImageNet 达到 76.60%与纯神经网络互有胜负泛化更强对训练时未见过的类别泛化能力最高提升 16%例如没见过熊仍能正确判断它是动物而非车辆。 想零配置体验可先运行 CLIpip install nbdt后直接执行nbdt 图片路径即可输出预测类别和每一步中间决策及其置信度。二、快速安装3 步完成环境准备第一步克隆仓库并安装依赖git clone https://gitcode.com/gh_mirrors/ne/neural-backed-decision-trees cd neural-backed-decision-trees python setup.py develop第二步确认已安装 PyTorch 与torchvisionsetup.py develop会自动安装requirements.txt中的其余依赖如pytorchcv、nltk等。第三步跑一下测试确认环境正常pytest tests✅ 安装完成后你就可以直接from nbdt.model import SoftNBDT使用全部模型与层级结构。三、30 行代码加载预训练 NBDT 并推理一张图片官方示例位于examples/load_pretrained_nbdts.ipynb下面按 4 步拆解总共约 30 行。第 1 步导入核心模块from nbdt.model import SoftNBDT from nbdt.models import wrn28_10_cifar10 from torchvision import transforms from nbdt.utils import DATASET_TO_CLASSES, load_image_from_pathSoftNBDT软推理决策树封装器源码见nbdt/model.pywrn28_10_cifar10在 CIFAR10 上预训练好的 WideResNet28x10 骨干网络来自nbdt/models/wideresnet.pyCIFAR100 用wrn28_10_cifar100TinyImagenet200 用wrn28_10。第 2 步加载预训练模型只有 5 行model wrn28_10_cifar10() model SoftNBDT( pretrainedTrue, datasetCIFAR10, archwrn28_10_cifar10, modelmodel)⚡ 关键在pretrainedTrue它会触发nbdt/model.py中的model_urls查找表按(arch, dataset)自动下载官方发布的 NBDT 检查点.pth并通过load_state_dict装入骨干网络。arch参数必须显式传入——项目对加载预训练 NBDT 的硬性要求用于定位正确的权重文件。第 3 步加载并预处理图像im load_image_from_path(cat.jpg) # 本地路径或图片URL均可 transforms transforms.Compose([ transforms.Resize(32), transforms.CenterCrop(32), transforms.ToTensor(), transforms.Normalize((0.4914, 0.4822, 0.4465), (0.2023, 0.1994, 0.2010)), ]) x transforms(im)[None] # 增加batch维度: (1,3,32,32)⚠️ 这里的Normalize均值/方差是 CIFAR10 的标准参数务必与训练时保持一致否则精度会明显下降。第 4 步执行推理并输出结果outputs model(x) # 输出logits: (1, 10) _, predicted outputs.max(1) # 取概率最大的类别索引 cls DATASET_TO_CLASSES[CIFAR10][predicted[0]] print(cls) # 例如输出: cat如果想看到决策树是怎么想的把model(x)换成model.forward_with_decisions(x)即可拿到从根节点到叶节点的每一步中间决策及置信度。四、内部机制SoftNBDT 一次 forward 到底发生了什么SoftNBDT的推理入口在nbdt/model.py的NBDT.forward中整个流程只有两行x self.model(x) # ① 神经网络骨干输出10维logits x self.rules(x) # ② SoftEmbeddedDecisionRules 做软决策树遍历① 特征提取wrn28_10_cifar10骨干照常输出 10 个类别的 logits与普通 CNN 分类器完全相同。② 软推理遍历决策树项目把 10 个类别组织成一棵树JSON 结构如nbdt/hierarchies/CIFAR10/graph-induced-wrn28_10_cifar10.json例如 CIFAR10 的结构是whole ┌─────┴─────┐ animal vehicle ┌───┴───┐ ┌──┴──┐ chordate vertebrate craft motor_vehicle └─┬─┘ └─┬─┘ carnivore ungulate ... ┌─┴─┐ ┌─┴─┐ cat dog deer horse airplane ship car truck ...硬推理HardNBDT每个内部节点只选概率最大的子节点沿一条路径走到叶子软推理SoftNBDT对每个内部节点把该节点下所有子组的类概率取出做 softmax然后把每个叶子到根路径上所有节点的子概率连乘得到全局 10 维分布。这样即使某一步走错分支最终结果仍可被修正且整条计算图可微分——这正是它精度能追平纯神经网络的原因。实现细节可参考nbdt/model.py中的SoftEmbeddedDecisionRules.traverse_tree概率连乘与HardEmbeddedDecisionRules.traverse_tree单路径遍历。 补充一点训练时对应的是树监督损失nbdt/loss.py的SoftTreeSupLoss要求骨干网络在每个内部节点上也学出正确判断从而让内嵌决策规则在推理时真正可用。五、可选预训练检查点清单nbdt/model.py的model_urls注册了以下官方预训练权重pretrainedTrue时自动下载骨干 arch数据集 dataset说明ResNet18CIFAR10ResNet18 骨干wrn28_10_cifar10CIFAR10论文主力模型另有hierarchywordnet变体wrn28_10_cifar100CIFAR100WideResNet28x10ResNet18TinyImagenet200ResNet18 骨干wrn28_10TinyImagenet200WideResNet28x10使用 ImageNet EfficientNet 等其他模型可参考examples/imagenet下的 ClassyVision 集成示例含examples/imagenet/losses/nbdt_losses.py。六、常见踩坑与排查清单 UserWarning: To load a pretrained NBDT, you need to specify the arch加载预训练权重时必须传arch且需与dataset组合匹配上表精度异常低检查预处理是否用了对应数据集的Normalize参数以及输入尺寸CIFAR 系列为 32×32权重与论文数字略有差异官方公开检查点为重新训练的版本可能与论文数字相差 0.1–0.2%属正常现象想看中间决策使用model.forward_with_decisions(x)若希望节点名更贴合词义可在构造SoftNBDT时指定hierarchywordnet。七、核心文件速查文件路径作用nbdt/model.pySoftNBDT/HardNBDT定义与预训练权重映射nbdt/loss.pySoftTreeSupLoss等树监督损失nbdt/hierarchy.py生成/加载层级决策树结构nbdt/hierarchies/各数据集的树结构 JSON 与 wnids 映射nbdt/models/wideresnet.pyWideResNet28 骨干工厂函数nbdt/utils.py类别表、图像加载等工具函数examples/load_pretrained_nbdts.ipynb官方 30 行推理示例 Notebookmain.py完整的训练/评估入口脚本 到这里你已经掌握了 NBDT 预训练模型的加载方式与 SoftNBDT 的完整推理链路——加载 5 行、推理 5 行剩下的时间可以用来探索它的可解释性分析了。【免费下载链接】neural-backed-decision-treesMaking decision trees competitive with neural networks on CIFAR10, CIFAR100, TinyImagenet200, Imagenet项目地址: https://gitcode.com/gh_mirrors/ne/neural-backed-decision-trees创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表