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   # 在新任务开始时清空记忆
)

SDNCSAM 也使用相同模式;唯一额外的参数是与稀疏性相关的 sparse_readstemporal_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