如何用 annotated_deep_learning_paper_implementations 实现 PonderNet 自适应计算并运行 parity 实验

发布时间:2026/9/9 15:11:57
如何用 annotated_deep_learning_paper_implementations 实现 PonderNet 自适应计算并运行 parity 实验 如何用 annotated_deep_learning_paper_implementations 实现 PonderNet 自适应计算并运行 parity 实验【免费下载链接】annotated_deep_learning_paper_implementations‍ 60 Implementations/tutorials of deep learning papers with side-by-side notes ; including transformers (original, xl, switch, feedback, vit, ...), optimizers (adam, adabelief, sophia, ...), gans(cyclegan, stylegan2, ...), reinforcement learning (ppo, dqn), capsnet, distillation, ... 项目地址: https://gitcode.com/gh_mirrors/an/annotated_deep_learning_paper_implementations这篇文章面向想动手验证 PonderNet 自适应计算机制的读者在 annotated_deep_learning_paper_implementations 仓库中PonderNet 已经以 PyTorch 实现好配套了一个 parity奇偶校验实验脚本。你只需要安装依赖、运行一个脚本就能观察网络如何根据输入动态决定循环计算步数并在屏幕上看到 loss、accuracy、steps 等训练指标。前提是你的环境能安装 PyTorch 与 labml 相关包仓库根目录 requirements.txt 给出的最低依赖包括torch1.10、labml0.4.147、labml-helpers0.4.84等。parity 任务与实验目标实验基于论文 PonderNet: Learning to Ponder 的实现模块说明见 labml_nn/adaptive_computation/ponder_net/readme.md。任务定义在 parity.py 的ParityDataset中输入是一个只含0、1、-1的向量其中1/-1的个数是 1 到总长度之间的随机数并随机打乱位置标签是1的个数的奇偶性奇数个1输出1偶数个输出0。PonderNet 的核心机制见init.py 的注释网络的每一步由 step function 输出当前步预测 $\hat{y}n$ 和停步概率 $\lambda_n$在最多 $N$ 步内按 $p_n \lambda_n \prod{j1}^{n-1}(1-\lambda_j)$ 的分布决定在哪一步停步。训练时对每一步的预测按 $p_n$ 加权计算重构损失 $L_{Rec}$再加一个把停步分布拉向几何分布 $p_G(\lambda_p)$ 的正则化项 $L_{Reg}$总损失为 $L L_{Rec} \beta L_{Reg}$。准备环境并安装获取仓库代码clone 到你自己的工作目录例如git clone https://gitcode.com/gh_mirrors/an/annotated_deep_learning_paper_implementations。在仓库根目录安装本仓库的labml_nn包。Makefile 提供install目标对应命令pip install -e .该命令会把labml_nn以可编辑模式装进当前 Python 环境。README 中给出的另一种方式是直接安装 PyPI 上的发布包pip install labml-nn但那样你运行的是 PyPI 版本代码本文的实验脚本必须从本仓库目录执行因此在仓库内用pip install -e .更直接。运行 PonderNet parity 实验实验入口是 labml_nn/adaptive_computation/ponder_net/experiment.py文件末尾有if __name__ __main__: main()在仓库根目录执行python labml_nn/adaptive_computation/ponder_net/experiment.py脚本会调用experiment.create(nameponder_net)创建名为ponder_net的 labml 实验然后按 Configs 中的默认配置训练。关键默认值如下均来自Configs均有注释说明配置项默认值文档中的说明epochs100训练轮数n_batches500每轮批次数batch_size128批大小n_elems8输入向量元素个数注释明确说 We keep it low for demonstration; otherwise, training takes a lot of timen_hidden64GRU 隐层状态单元数max_steps20最大步数 $N$lambda_p0.2几何分布 $p_G(\lambda_p)$ 的参数与停步概率 $\lambda_n$ 无关beta0.01正则化损失 $L_{Reg}$ 的系数 $\beta$grad_norm_clip1.0按范数裁剪梯度优化器在 main() 中固定为 Adam学习率0.0003。训练/验证数据在Configs.init中构建训练集ParityDataset(batch_size * n_batches, n_elems)即默认 128×50064000 个样本验证集ParityDataset(batch_size * 32, n_elems)即默认 4096 个样本两者都包在DataLoader中、批大小为 128。注意 parity 数据是__getitem__里按随机规则现生成的每次取到的是新样本不是静态数据集。如何观察与判断训练结果实验没有额外的评估脚本验证手段就是运行过程中的屏幕输出。Configs.init里通过 labml 的 tracker 打开四类标量的屏幕打印tracker.set_scalar(loss.*, True) # 重构损失 L_Rec tracker.set_scalar(loss_reg.*, True) # 正则化损失 L_Reg tracker.set_scalar(accuracy.*, True) # 训练/验证 accuracy tracker.set_scalar(steps.*, True) # 期望停步数每一批的处理逻辑step 方法前向返回四个张量各步停步概率p、各步预测y_hat、采样停步处的p_sampled、y_hat_sampledloss.记录 $L_{Rec}$用nn.BCEWithLogitsLoss(reductionnone)逐样本计算后按 $p_n$ 加权求和见 ReconstructionLossloss_reg.记录 $L_{Reg}$即 $KL(p_n \Vert p_G(\lambda_p))$RegularizationLosssteps.记录期望步数按expected_steps (p * steps[:, None]).sum(dim0)计算其中steps为1..Naccuracy 指标用AccuracyDirect比较对象是采样停步预测y_hat_sampled 0与真实标签训练与验证都会计算 epoch 级 accuracy。因此运行脚本后你应该在终端看到持续刷新的loss.、loss_reg.、accuracy.、steps.数值。文档未给出固定的达标阈值只说明了正则化项的作用是biases the network towards taking $1/\lambda_p$ steps所以steps.数值围绕 $1/\lambda_p$默认配置下即 5波动是文档描述的预期行为其余指标变化只能依据你本次运行打印的数值本身判断。可调参数与推理时停步以下调整都有源码注释依据按需修改Configs即可仓库只读时请在你自己的工作副本中修改n_elems控制 parity 向量长度。注释同时提醒Although the parity task seems simple, figuring out the pattern by looking at samples is quite hard默认取 8 只是演示考虑调大后文档明确说 training takes a lot of time。max_steps步数上限 $N$。模型最后一步n max_steps会强制 $\lambda_N 1$ 保证一定停步见 forward 实现。lambda_p/beta分别控制正则化分布的形状与正则项权重。is_halt模型上有一个self.is_halt False的选项源码注释写明它是An option to set during inference so that computation is actually halted at inference time。设为True后batch 内所有样本都停步时循环会提前break用于推理阶段的实际省计算训练时保持False每步预测都要算完以计算加权损失。限制与边界文档说明n_elems8是刻意调低以缩短演示训练时间不要把它当作 parity 任务的标准规模。实验脚本一次运行就是完整的 100 epoch 训练流程没有提供提前中断或断点恢复的文档化参数中途 CtrlC 终止即可labml 的 tracker 每步调用tracker.save()记录数据。该实现面向 parity 演示模型结构是ParityPonderGRUGRU Cell 作 step function 线性输出层不是通用 PonderNet 封装换任务需要自行按init.py 中的结构改写。更多带注释的文档版页面可参考 docs/adaptive_computation/ponder_net/experiment.html 和 docs/adaptive_computation/parity.html。【免费下载链接】annotated_deep_learning_paper_implementations‍ 60 Implementations/tutorials of deep learning papers with side-by-side notes ; including transformers (original, xl, switch, feedback, vit, ...), optimizers (adam, adabelief, sophia, ...), gans(cyclegan, stylegan2, ...), reinforcement learning (ppo, dqn), capsnet, distillation, ... 项目地址: https://gitcode.com/gh_mirrors/an/annotated_deep_learning_paper_implementations创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考