ixaxaar/pytorch-dnc
Differentiable Neural Computers, Sparse Access Memory and Sparse Differentiable Neural Computers, for Pytorch
📚 什么是 pytorch‑dnc?
pytorch‑dnc 是一个 纯 Python 库,为 PyTorch 实现了多种 带外部记忆的神经网络 架构:
| 架构 | 论文(原始) | 添加功能 |
|---|---|---|
| DNC(可微神经计算机) | Graves et al., Nature 2016 | 一个循环控制器加上一个可微读写矩阵值外部记忆。 |
| SDNC(稀疏 DNC) | Rae et al., NeurIPS 2018 | 与 DNC 相同,但使用 稀疏读写 以扩展到数千个记忆槽。 |
| SAM(稀疏访问记忆) | Rae et al., NeurIPS 2018 | 一个独立的稀疏记忆模块,可插入任意 RNN 控制器。 |
该库允许你将这些模块之一插入 PyTorch 模型中,并在需要算法推理的任务(如复制序列、加法、arg-max 等)上进行端到端训练。
🚀 快速开始
# 从 PyPI 安装包
pip install dnc
# (可选)GPU 加速稀疏操作需要 FAISS
conda install -c pytorch faiss-gpu
或从源码安装:
git clone https://github.com/ixaxaar/pytorch-dnc
cd pytorch-dnc
pip install -r requirements.txt
pip install -e .
🛠️ 如何使用模块
这三个类具有相同的构造函数签名(大多数参数都有合理的默认值)。以下是经典 DNC 的最小示例:
import torch
from dnc import DNC
# 创建一个输入维度为 64、隐藏状态维度为 128 的 DNC
model = DNC(
input_size=64,
hidden_size=128,
nr_cells=100, # 100 个记忆槽
cell_size=32, # 每个槽存储一个 32 维向量
read_heads=4,
batch_first=True,
device=torch.device('cuda:0')
)
# 初始隐藏状态(控制器、记忆、读向量)—— 让模型懒加载创建
h = (None, None, None)
# 前向传播一个随机批次(seq_len=10, batch=4)
output, (h_ctrl, h_mem, h_read) = model(
torch.randn(10, 4, 64), # (seq_len, batch, input_dim)
h,
reset_experience=True # 在新任务开始时清空记忆
)
SDNC 和 SAM 也使用相同模式;唯一额外的参数是与稀疏性相关的 sparse_reads 和 temporal_reads。
调试模式
在构造函数中传入 debug=True。此时前向调用将返回第三个值——一个包含 NumPy 数组的字典,其中包含原始记忆矩阵(memory, link_matrix, read_weights, …)。这些数据可使用 Visdom 等工具可视化,以检查网络如何使用其外部记忆。
📊 仓库自带的示例任务
仓库包含可直接运行的脚本,复现经典的 DNC 实验:
| 任务 | 测试内容 | 如何运行 |
|---|---|---|
| 复制任务 | 存储并重现任意长度输入序列的能力。 | python ./tasks/copy_task.py -cuda 0 -optim adam -sequence_max_length 8 |
| 加法任务 | 学习在长序列中对两个数字求和(原始的“学习加法”基准)。 | python ./tasks/add_task.py …(查看脚本帮助) |
| argmax 任务 | 找到序列中最大元素的索引。 | python ./tasks/argmax_task.py … |
所有任务都接受丰富的命令行选项(学习率、优化器、记忆大小、课程学习等)。对于复制任务,你还可以启动一个 Visdom 服务器(pip install visdom && python -m visdom.server),在训练期间观察记忆矩阵的热力图。
🏗️ 代码组织
├─ dnc/ # 核心实现(DNC、SDNC、SAM)
├─ tasks/ # 复制 / 加法 / argmax 的训练脚本
├─ docs/ # 架构图和截图
├─ requirements.txt # Python 依赖项
└─ tests/ # pytest 套件
核心模块是纯 PyTorch 实现,仅在 SDNC 和 SAM 中使用 GPU 加速稀疏读写时依赖 FAISS。
📦 谁会使用这个?
- 实验 神经图灵机 或其他可微数据结构的研究人员。
- 需要为序列到序列模型提供 即插即用外部记忆 的实践者,例如程序合成或长上下文推理。
- 寻找具体、可运行实现以演示 DNC 论文概念的教育者。
📚 进一步阅读
- 原始 DNC 论文 – 使用动态外部记忆的神经网络进行混合计算(Graves et al., Nature 2016)。
- 稀疏记忆论文 – 使用稀疏读写扩展记忆增强神经网络(Rae et al., NeurIPS 2018)。
✅ 总结
pytorch‑dnc 提供了现成的、文档齐全的 PyTorch 模块,用于 DNC、SDNC 和 SAM,包含示例训练脚本和可窥探外部记忆的调试模式。通过 pip install dnc 安装,将 DNC/SDNC/SAM 插入模型,即可开始训练需要长程记忆的算法任务。
相关
- 项目
- 项目
- 项目
- 项目
- Dispatch