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