如何用 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),仅供参考

相关新闻

最新新闻

网络安全——kali中的set工具

网络安全——kali中的set工具

一、kali中的set工具利用的是社会工程学攻击 二、常见的社会工程学: 环境渗透、引诱、伪装欺骗、说服、恐吓、恭维、反向社会工程学 三、SET工具的使用1、建立钓鱼网站收集目标凭证 (1)、打开set:有两种途径(2&#xf…

2026/9/9 15:52:01
达签是什么平台?主要提供哪些签证服务?

达签是什么平台?主要提供哪些签证服务?

达签平台概述达签是一家专业的线上签证服务平台,致力于为用户提供便捷、高效的签证申请服务。平台通过线上化的操作流程,帮助用户简化签证办理的繁琐步骤,让出国签证申请更加简单透明。企业背景达签平台由杭州达签信息科技有限公司运营&#…

2026/9/9 15:52:01
从hello world到可执行程序:C语言编译链接全流程指南

从hello world到可执行程序:C语言编译链接全流程指南

1. 一个hello world到底经历了什么:先看全流程地图我最早学C语言的时候,写了一个hello world,然后老师告诉我在命令行敲gcc hello.c -o hello,回车,程序就出来了。当时觉得这事儿没什么大不了,直到后来自己…

2026/9/9 15:52:01
基于Python Flask与Vue3的求职招聘资讯交流系统全栈开发实战

基于Python Flask与Vue3的求职招聘资讯交流系统全栈开发实战

1. 项目整体设计:为什么是Python Offer求职招聘资讯交流系统 做这个项目的起因很直接,我身边不少同学和朋友在找工作时,信息非常分散。今天在哪家公司笔试,明天又有一家发来Offer,薪资结构、面试进度、公司口碑这些信息…

2026/9/9 15:52:01
doc3D数据集解析:从合成数据到DewarpNet文档矫正实践

doc3D数据集解析:从合成数据到DewarpNet文档矫正实践

简介:面向文档图像矫正与三维重建研究者的doc3D数据集下载工具包,解决真实纸张复杂变形场景下缺乏高真实性3D训练数据的问题。借助Shell脚本,可批量获取10万张包含3D坐标、深度、UV映射、反照率、法线及棋盘格标注的混合图像,适用…

2026/9/9 15:52:01
Kotlin协程实战指南:从挂起函数到Flow与结构化并发

Kotlin协程实战指南:从挂起函数到Flow与结构化并发

1. 从线程到协程:为什么Kotlin敢把异步做得这么简单 过去几年我一直在用Java写后端服务,后来转到Kotlin开发,最直观的感受就是异步编程的姿势发生了巨大变化。Java里典型的异步写法是线程池配Future,或者用CompletableFuture串回调…

2026/9/9 15:47:00