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