首页/新闻资讯/正文详情

Neural Additive Models(NAM)实战指南:用神经网络实现可解释的加性机器学习

发布时间:2026/9/23 4:09:14 来源:云帆数科 栏目:资讯中心
Neural Additive Models(NAM)实战指南:用神经网络实现可解释的加性机器学习
人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载导读本文围绕neural_additive_models开源项目系统讲解 Neural Additive ModelsNAM这一将深度网络与广义加性模型GAM相结合的模型家族它以tf.keras.Model形式提供可直接嵌入任意神经网络训练流程的 NAM 实现并附带基于tf.compat.v1的分类/回归计算图构建工具与单数据集训练脚本。读完本文你将掌握 NAM 的 ExU 隐藏单元原理、FeatureNN 与 NAM 的 Keras 实现细节、9 个公开数据集的加载与预处理管线、nam_train.py的全部命令行超参数以及如何用run.sh验证环境并复现论文实验。该仓库对应论文Neural Additive Models: Interpretable Machine Learning with Neural NetsNeurIPS 2021开源代码覆盖模型定义、图构建、数据加载、训练脚本与单元测试是学习和二次开发可解释深度模型的完整参考实现。NAM 核心思想可加性 × 深度网络广义加性模型将预测建模为每个特征独立形状函数的和每个特征由一个专属的神经网络FeatureNN建模网络输出仅依赖该单一特征模型总输出为所有 FeatureNN 输出之和再加上一个可训练的偏置项bias由于每个特征的贡献被解耦训练结束后可以单独绘制每个特征的 shape function从而获得与树模型、线性模型同等级别的可解释性同时保留深度网络在表格数据上的表达能力。从源码结构看这种每个特征一个网络、输出求和的架构在 models.py 的NAM.call中体现得非常直接calc_outputs将输入按特征维度切分tf.split后分别送入各 FeatureNN再将各子网络输出tf.stack、做feature_dropout后reduce_sum最后叠加bias。ExU 隐藏单元NAM 的关键创新NAM 的第一隐藏层使用论文提出的ExUExponential Unit单元而不是普通 ReLU。在 models.py 中其定义为def exu(x, weight, bias): ExU hidden unit modification. return tf.exp(weight) * (x - bias)ExU 的设计动机是让网络具备学习跳变jump形状函数的能力参数以指数形式出现通过梯度可以更有效地移动特征值的分割点从而在特征空间内实现尖锐的阈值切换。实际使用时ExU 输出会进一步通过relu_nReLU 裁剪到[0, n]得到最终激活self._activation lambda x, weight, bias: relu_n(exu(x, weight, bias))与 ExU 配套的beta初始化为均值为 4.0、标准差 0.5 的截断正态分布见 ActivationLayer.build而标准 ReLU 分支则使用glorot_uniform初始化。--activation参数可在exu与relu之间切换这也是复现论文实验时最主要的模型配置开关之一。仓库结构与核心模块neural_additive_models目录结构如下文件职责models.py纯 Keras 模型ExU、ActivationLayer、FeatureNN、NAM、DNNgraph_builder.py基于tf.compat.v1构建训练/评估计算图、损失函数与正则化data_utils.py数据集加载、预处理、训练/验证/测试切分nam_train.py单数据集 split 的训练脚本absl flags 入口run.sh环境自检脚本建虚拟环境、装依赖、跑测试requirements.txt / setup.py依赖声明与 pip 打包tests/models_test.py、graph_builder_test.py、data_utils_test.pynam_train_test.py端到端训练流水线测试三种可加载架构models.py 提供了三类模型与测试 tests/models_test.py 中的三种架构一一对应exu_nam浅层 NAMshallowTrue第一层为 ExUActivationLayer随后直接接无偏置线性输出层是论文主推配置relu_nam深度 NAMshallowFalse第一层为 ExU 或 ReLU 激活层其后追加 64、32 个单元的 ReLU Dense 层最后接线性输出层dnn10 层 × 100 单元的 ReLU 深度网络基线he_normal初始化用于与 NAM 做可解释性 vs 性能的对照。FeatureNN的核心属性见 models.pynum_units第一隐藏层神经元基函数数量dropout每个 FeatureNN 内部的 dropout 比率训练时生效、评估时为 0通过tf.cond(training, ...)控制shallow是否只使用单隐藏层activationexu或relu。NAM的顶层参数见 models.py还包括feature_dropout——以给定概率整条丢弃某个特征的 FeatureNN论文中的 feature dropout 正则化以及支持num_units传入整数所有特征统一或与特征数等长的列表按特征定制基函数数量。环境依赖与安装README 注明代码在 Ubuntu 16 下测试通过依赖包见 requirements.txttensorflow1.15numpy1.15.2sklearn0.23pandas0.24absl-py注意两点适用前提代码大量使用tensorflow.compat.v1graph_builder 与 nam_train 均以tf.compat.v1导入因此需要能加载 TF1.x 语义的 TensorFlow 环境setup.py 声明支持 Python 3.5–3.8建议使用对应版本的 Python 3 虚拟环境。仓库自带一键环境验证脚本 run.sh流程为创建 python3 虚拟环境 → 激活 → 安装 requirements → 运行python -m neural_additive_models.nam_train_test.py冒烟测试virtualenv -p python3 . source ./bin/activate pip install -r neural_additive_models/requirements.txt python -m neural_additive_models.nam_train_test.py也可以直接通过 pip 安装包pip install .setup.py 已配置find_packages与依赖声明。数据集加载、预处理与切分数据集清单与下载方式论文实验所用数据集除 MIMIC-II 外托管在公开 GCP 存储桶gs://nam_datasets/data可借助 gsutil 下载gsutil 的安装步骤参见官方 storage 文档。data_utils.py 的load_dataset支持 9 个数据集其中 7 个分类、2 个回归名称任务类型数据说明Telco分类Telco 客户流失预测KaggleBreastCancer分类威斯康星乳腺癌数据集sklearn 内置Adult分类成人收入预测Census IncomeCredit分类信用卡欺诈检测高度不平衡正样本占比约 0.172%Heart分类Cleveland 心脏病数据Mimic2分类MIMIC-II ICU 死亡率预测需签署数据使用协议Recidivism分类ProPublica COMPAS 再犯风险数据Fico回归FICO 信用评分Housing回归California Housing 房价中位数预测MIMIC-II 的预处理版本因涉及临床数据只能在你向 PhysioNet 提交 MIMIC-III 临床数据库的签署版数据使用协议后共享——README 对此有明确说明请遵循数据使用合规要求。数据加载测试 tests/data_utils_test.py 验证了各数据集的样本规模如 BreastCancer 569 条、Adult 32561 条、Credit 284807 条、Housing 20640 条可作为下载完整性的校验参考。预处理管线transform_data见 data_utils.py执行固定三步转换类别特征 One-Hot 编码通过ColumnTransformer对dtype.kind O的列做OneHotEncoderhandle_unknownignore数值特征恒等变换数值列走FunctionTransformerMin-Max 缩放将所有特征统一缩放到(-1, 1)区间MinMaxScaler((-1, 1))。值得注意的工程细节CustomPipeline.apply_transformation只执行steps[:-1]不含最后的 estimator从而在不拟合任何模型的前提下完成数据变换。训练/验证/测试切分get_train_test_fold基于StratifiedKFold分类或KFold回归做 5 折交叉验证划分由fold_num指定测试折split_training_dataset基于StratifiedShuffleSplit/ShuffleSplit生成训练-验证切分默认验证占比test_size0.125返回可迭代生成器供多 split 实验使用。分类场景下训练数据通过create_balanced_dataset见 graph_builder.py按正负类别分别建流并sample_from_datasets等概率采样以缓解类别不平衡。训练脚本nam_train.py 与全部超参数nam_train.py 提供单数据集 split 的训练示例入口为main先按fold_num切出训练测试折再按data_split从训练集中取 (train, validation) split 调用training。logdir与training_epochs被标记为必填 flag。完整命令行参数如下Flag默认值说明--training_epochs必填训练轮数epochs--learning_rate1e-2初始学习率--output_regularization0.0特征输出惩罚系数feature reg--l2_regularization0.0L2 权重衰减系数--batch_size1024批大小--logdir必填存放 checkpoint 与摘要的目录--dataset_nameTeleco数据集名见上表--decay_rate0.995优化器学习率衰减率每个 epoch 衰减一次--dropout0.5FeatureNN 内部 dropout 率--data_split1使用的 split 索引取值 1 到num_splits--tf_seed1TensorFlow 随机种子--feature_dropout0.0整条丢弃特征的 dropout 概率--num_basis_functions1000实数特征在 FeatureNN 中的基函数神经元数量上限--units_multiplier2类别特征基函数数量的乘子--cross_valFalse是否执行交叉验证--max_checkpoints_to_keep1保留的最近 checkpoint 数量--save_checkpoint_every_n_epochs10每隔多少 epoch 保存一次 checkpoint--n_models1并行训练多少个模型可多模型投票--num_splits3数据 split 数量--fold_num15 折交叉验证中使用的折索引--activationexu激活函数relu或exu--regressionFalse是否为回归任务否则为二分类--debugFalse调试模式额外记录 TensorBoard 摘要--shallowFalse使用浅层单隐藏层还是深层 FeatureNN--use_dnnFalse使用 10 层 DNN 基线替代 NAM--early_stopping_epochs60早停耐心值典型分类训练命令以 Telco 为例python -m neural_additive_models.nam_train \ --training_epochs100 \ --dataset_nameTelco \ --logdir/tmp/nam_telco \ --learning_rate1e-2 \ --batch_size1024 \ --activationexu \ --shallowTrue \ --dropout0.5 \ --regressionFalse回归任务则加--regressionTrue如--dataset_nameHousing评估指标随之切换。data_split与cross_val互斥——nam_train.py 中通过multi_flags_validator保证二者不能同时使用。训练循环与模型选择机制训练采用tf.train.MonitoredSessionCheckpointSaverHook见 nam_train.py每个 epoch 结束后按decay_rate衰减学习率lr_decay_op每save_checkpoint_every_n_epochs轮在验证集上计算指标回归用 RMSE越小越好、分类用 AUROC越大越好达到更优即把当前 checkpoint 拷贝到best_checkpoint目录若最优指标对应的 epoch 距今超过early_stopping_epochs轮则该模型提前停止支持n_models个模型各自的独立早停支持多模型--n_models并行训练最终返回多个模型训练/验证指标的平均值。计算图构建graph_builder.py 的内部原理graph_builder.py 的build_graph将数据管线、模型、优化器、评估指标组织成一张可sess.run的 TF1 计算图核心要素包括模型创建create_nam_model依据每个特征取值的唯一数量num_unique_vals自适应设置各特征网络的基函数数——num_units min(num_basis_functions, num_unique_vals * units_multiplier)即类别特征按取值数乘子分配、实数特征封顶 1000损失函数分类用penalized_cross_entropy_losssoftmax 交叉熵实现见回归用penalized_mse_loss双重正则化penalized_loss见 graph_builder.pyoutput_regularization惩罚各 FeatureNN 输出的 L2 范数均值feature_output_regularizationl2_regularization对模型可训练变量做权重衰减weight_decay按网络数量归一化优化器tf.train.AdamOptimizer 全局步数配tf.metrics.mean维护训练损失运行均值评估指标回归返回 RMSErmse_loss分类返回 ROC AUCroc_auc_score预测值经 sigmoid 后与sklearn.metrics.roc_auc_score比较分别绑定训练集与验证集迭代器。图构建测试 tests/graph_builder_test.py 展示了最小可运行用法build_graph返回的graph_tensors_and_ops中iterator_initializer、running_vars_initializer、train_op需要先初始化再迭代执行metric_scorestrain返回 float 指标。运行与测试验证仓库提供两层验证手段端到端训练测试nam_train_test.py以BreastCancer分类与Housing回归为例用num_basis_functions16的小配置跑 4 个 epoch 训练验证整个流水线无错误运行单元测试tests/覆盖三种模型架构前向计算models_test.py、计算图构建与指标计算graph_builder_test.py、9 个数据集加载规模data_utils_test.py。执行全部测试可运行python -m neural_additive_models.nam_train_test.py或按模块单独运行python -m neural_additive_models.tests.models_test等测试文件。论文引用如果你在研究中使用了本仓库代码请引用以下论文NeurIPS 2021Agarwal, R., Melnick, L., Frosst, N., Zhang, X., Lengerich, B., Caruana, R., Hinton, G. E. (2021). Neural additive models: Interpretable machine learning with neural nets. Advances in Neural Information Processing Systems, 34.article{agarwal2021neural, title{Neural additive models: Interpretable machine learning with neural nets}, author{Agarwal, Rishabh and Melnick, Levi and Frosst, Nicholas and Zhang, Xuezhou and Lengerich, Ben and Caruana, Rich and Hinton, Geoffrey E}, journal{Advances in Neural Information Processing Systems}, volume{34}, year{2021} }关于多任务 NAM 与 COMPAS 数据的伦理提示多任务 NAM多任务版本multi-task NAMs的代码独立维护在作者的另一仓库lemeln/nam本仓库聚焦单任务 NAMCOMPAS 数据使用说明README 特别强调用机器学习模型预测审前羁押存在重要的伦理考量——可参考 Partnership on AI 发布的《美国刑事司法系统中的算法风险评估工具报告》Google 为该多利益相关方组织的成员。COMPAS 数据集在此仅用作如何识别与修复数据公平性问题的示例它是算法公平性文献中的经典数据集本仓库为非官方 Google 产品not an official Google product以 Apache 2.0 协议开源发布见 setup.py。小结neural_additive_models提供了一个完整、可直接复现论文实验的 NAM 实现models.py给出可插拔的 Keras 模型ExU 单元 每特征 FeatureNN 加性求和graph_builder.py封装了分类/回归的训练图与双正则化机制nam_train.py以 20 个 absl flag 覆盖论文中的全部超参数旋钮data_utils.py打通了 9 个标准数据集的下载、预处理与切分。无论你是要在业务中落地可解释模型还是复现/扩展 NeurIPS 2021 论文都可以从本仓库直接起步。赞分享人工智能深度学习NLP计算机视觉强化学习【免费下载链接】google-researchGoogle Research项目地址https://gitcode.com/gh_mirrors/go/google-research点击查看免费下载相关推荐ACE-Step UI终极指南如何在本地免费生成专业AI音乐完全替代SunoACE Step UI终极指南如何在本地免费生成专业AI音乐完全替代Suno 还在为每月支付高昂的Suno或Udio订阅费而烦恼吗想要一个完全免费、本地运人工智能AI 应用音频媒体生成本地部署前端后端DGL 实现 GNNExplainer从训练到可视化的图神经网络可解释性实战指南DGL 实现 GNNExplainer从训练到可视化的图神经网络可解释性实战指南 GNNExplainerGenerating Explanations f人工智能机器学习深度学习图计算多智能体强化学习新突破on-policy项目全面解析与入门指南多智能体强化学习新突破on policy项目全面解析与入门指南 多智能体强化学习 Multi Agent Reinforcement Learning, M人工智能强化学习多智能体创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

相关推荐

用Grok Bot三天冲刺从零推进公司:Day1方向验证与市场扫描实操
用Grok Bot三天冲刺从零推进公司:Day1方向验证与市场扫描实操

不用急着给公司起名、注册、装修办公室,真正的从零推进一家公司,指的是在最短时间内把一个模糊的想法,变成一套有客户、有产品方向、有验证结论的推进方案。我最近在做的这个三天冲刺项目里,把Grok Bot作为唯一的AI主引擎&#xf… · 2026/9/23 4:09:14

和飞信是什么?搞懂这1个高频面试题,配置不再卡半天
和飞信是什么?搞懂这1个高频面试题,配置不再卡半天

和飞信是什么?搞懂这1个高频面试题,配置不再卡半天 配置环境就卡半天,这是很多初入职场的开发者最真实的写照。你明明照着教程敲命令,结果终端里全是红字报错,重启电脑也没用。这时候,如果你能把“和飞信是什么”这个看似与代码无关的概念讲清楚,往往… · 2026/9/23 4:09:14

垃圾分类图像分类实战:从数据集到模型部署的避坑指南
垃圾分类图像分类实战:从数据集到模型部署的避坑指南

简介:这份深度学习图像分类数据集面向从事计算机视觉入门与垃圾分类识别实践的开发者、学生及算法爱好者,围绕塑料瓶、玻璃瓶、金属瓶等可回收物类别构建,可直接用于训练与评估卷积神经网络分类模型。资源包共约2000个文件,以1998… · 2026/9/23 4:09:14

AI眼镜与可控核聚变:技术路线争议与商业化前景
AI眼镜与可控核聚变:技术路线争议与商业化前景

1. 为什么AI眼镜与可控核聚变会成为技术路线的争议焦点?最近科技圈有个特别有意思的现象:一边是各大科技公司扎堆研发AI眼镜,另一边则是少数硬核团队在可控核聚变领域默默耕耘。这两种看似毫不相干的技术路线,实际上代表着完全不同… · 2026/9/23 6:35:25

大模型推理优化框架对比与选型指南
大模型推理优化框架对比与选型指南

1. 大模型推理部署的现状与挑战当前大语言模型(LLM)在实际业务落地过程中面临的核心矛盾是:模型规模持续增长与推理效率难以提升之间的鸿沟。以Llama 3-70B为例,单次推理需要占用140GB以上的GPU显存,即使使用A100 80GB… · 2026/9/23 6:35:19

个人品牌建设:差异化定位与记忆点设计实战
个人品牌建设:差异化定位与记忆点设计实战

1. 项目背景与核心价值"大家好,我是The One"这个看似简单的自我介绍,背后蕴含着个人品牌建设的完整方法论。在当今注意力经济时代,如何用一句话让人记住你,已经成为职场人士、创业者、自由职业者的必备技能。这个标题实… · 2026/9/23 6:35:19

网络热词“cua”走红:从CUBA到拟声词的流行密码
网络热词“cua”走红:从CUBA到拟声词的流行密码

“cua”这四个字母最近在各大平台的热搜榜上窜得很快,很多人第一次看到时一脸懵——是拟声词?是新游戏?还是什么缩写?我翻了一下各个讨论区,发现这个词的走红路径挺有意思的,它不是某一个人带火的&#xff… · 2026/9/23 6:35:12

AI工具PaperZZ:15分钟搞定专业学术PPT
AI工具PaperZZ:15分钟搞定专业学术PPT

1. 学术PPT制作的痛点与效率革命作为一名经历过无数次学术答辩的老手,我深知制作PPT这个看似简单的任务背后隐藏着多少时间黑洞。每次答辩前,我们总要在文献堆里反复筛选数据、调整版式、纠结配色,最后往往在Deadline前通宵赶工。直到遇到Pap… · 2026/9/23 6:35:06

专业降AIGC工具:提升AI生成内容质量的关键技术
专业降AIGC工具:提升AI生成内容质量的关键技术

1. 项目概述:专业降AIGC工具的诞生背景最近两年AI生成内容(AIGC)技术爆发式发展,从文字创作到图像生成,AI正在重塑内容生产流程。但随之而来的问题是:大量AI生成内容存在质量参差不齐、专业度不足、风格同质… · 2026/9/23 6:35:06

3招搞定手机怎么下载微信面试难题实战项目解析
3招搞定手机怎么下载微信面试难题实战项目解析

3招搞定手机怎么下载微信面试难题实战项目解析 面试被问“手机怎么下载微信”背后的原理,90%的人答不上来。别笑,这看似弱智的问题,实则是考察你对移动应用分发机制、安全校验及网络协议理解的试金石。我带过不少校招新人,他们背了八股文,却连一个A… · 2026/9/23 0:00:03

你有新短消息请注意查收:3个新手避坑指南搞定消息系统选型
你有新短消息请注意查收:3个新手避坑指南搞定消息系统选型

你有新短消息请注意查收:3个新手避坑指南搞定消息系统选型 面试被问“高并发下如何保证消息不丢失”,你张口就是“用Redis”,结果面试官追问“如果Redis宕机了怎么办”,你瞬间卡壳。这种场景太常见了,很多新手在背八股文时,只记住了技术名词… · 2026/9/23 0:00:29

Win7无线热点配置工具源码解析:解决API失效的3个实战技巧
Win7无线热点配置工具源码解析:解决API失效的3个实战技巧

Win7无线热点配置工具源码解析:解决API失效的3个实战技巧 Win7无线热点配置工具在Win10/11上跑不动?不是你的问题,是版本升级后 API 全变了。很多老项目里的 netsh wlan… · 2026/9/23 0:00:36

了解更多?预约专属演示

我们的顾问将为您一对一讲解产品与方案

企业微信二维码