ndif-team/nnsight
The nnsight package enables interpreting and manipulating the internals of deep learned models.
nnsight – 解釋與編輯 PyTorch 模型的內部結構
是什麼 – nnsight 是一個可透過 pip install nnsight 安裝的 Python 庫,讓您能檢視任何 PyTorch 模型(如 GPT‑2、LLaMA)的內部狀態,讀取任意層的隱藏狀態張量,即時修改它們,計算中間張量的梯度,甚至對模型進行永久性編輯。它可在您自己的 GPU/CPU 上本地運行,也可在 ND 研究所的遠端基礎設施上執行,適用於非常大的模型。
為何重要 – 現代基礎模型(如 GPT‑2、LLaMA 等)是黑箱。希望理解模型為何做出特定預測、測試因果假設或原型化模型編輯技術的研究人員,需要一種無需重寫模型程式碼即可存取和操作內部激活的清晰方法。nnsight 提供一個高階、Python 風格的 API,抽象掉了鉤子注入和追蹤的重複程式碼。
核心功能(如 README 所述)
| 功能 | 如何使用 | 你將獲得 |
|---|---|---|
| 激活存取 | with model.trace(prompt): hidden = model.transformer.h[5].output[0].save() |
包含給定提示下第 5 層隱藏狀態的實際張量。 |
| 就地干預 | 在追蹤區塊內使用 model.transformer.h[0].output[0][:] = 0 |
前向傳播繼續使用修改後的激活,可用於測試因果效應。 |
| 中間張量的梯度 | with loss.backward(): grad = hs.grad.save() |
任意張量(如隱藏狀態)相對於您定義的損失函數的梯度。 |
| 批量呼叫 | with tracer.invoke(prompt): … |
並行執行多個提示;每個呼叫順序執行,但透過 .save() 共享值。 |
| 分步控制生成 | with model.generate(prompt, max_new_tokens=5) as tracer: for step in tracer.iter[:]: … |
可在特定步驟進行干預的自回歸生成。 |
| 模型編輯 | with model.edit() as edited: edited.transformer.h[0].output[0][:] = 0 |
產生一個新 LanguageModel 實例,永久包含編輯內容,原始模型保持不變。 |
| 僅形狀掃描 | with model.scan(prompt): dim = nnsight.save(layer.output.shape[-1]) |
無需完整前向傳播即可取得張量形狀。 |
| 快取與會話 | cache = tracer.cache() 或 with model.session() as s: … |
在多個追蹤中重用先前捕獲的激活,提升效率。 |
| 遠端執行 | CONFIG.set_default_api_key(..); model = LanguageModel('meta-llama/Meta-Llama-3.1-8B'); with model.trace(..., remote=True): … |
在 NDIF 的雲端服務上執行追蹤程式碼,適用於本地無法容納的模型。 |
| vLLM 集成 | from nnsight.modeling.vllm import VLLM; model = VLLM('gpt2', ...) |
保持相同追蹤 API 的同時實現高效能推理。 |
| 任意 PyTorch 模型 | NNsight(net) 其中 net 是任意 torch.nn.Module |
相同的追蹤/干預工具適用於非 Transformer 模型(如簡單前饋網路)。 |
快速入門範例(來自 README)
from nnsight import LanguageModel
model = LanguageModel('openai-community/gpt2', device_map='auto', dispatch=True)
with model.trace('The Eiffel Tower is in the city of'):
# 將第一層的激活設為零
model.transformer.h[0].output[0][:] = 0
# 保存最終隱藏狀態和 logits
hidden = model.transformer.h[-1].output[0].save()
logits = model.output.save()
print(model.tokenizer.decode(logits.logits.argmax(dim=-1)[0]))
該程式碼片段展示了載入模型、開啟追蹤、干預某一層並取得最終預測的過程。
典型用例
- 機制可解釋性 – 檢查注意力頭或 MLP 塊如何貢獻於某個 token 的預測。
- 因果探測 – 將激活設為零、加入雜訊或替換,以測試資訊流的假設。
- 模型編輯研究 – 建立可逆編輯(例如,「在『艾菲爾鐵塔』之後總是輸出『巴黎』」)。
- 調試自訂架構 – 對任意
torch.nn.Module使用相同 API,驗證前向傳播行為。 - 高效批量實驗 – 使用
invoke/session機制在單次前向傳播中執行多個提示。
限制與注意事項(如 README 所述)
- 執行順序很重要 – 在追蹤中必須按模型實際執行順序存取模組;否則會觸發
OutOfOrderError死鎖。 - 基於執行緒的同步 – 此庫在獨立工作執行緒中執行您的追蹤程式碼;值僅在呼叫
.save()後才可用。 - 無界迭代器 –
tracer.iter[:]永不返回,因此後續程式碼不會執行,除非放在獨立的invoke區塊中。 - 遠端執行需要 API 金鑰,且模型必須在 NDIF 平台上可用。
- vLLM 集成僅限 vLLM 支援的模型(目前為 Transformer 風格語言模型)。
如何進一步了解
- 文件網站 – https://www.nnsight.net
- 論文 – NNsight and NDIF: Democratizing Access to Foundation Model Internals (arXiv 2407.14561)
- Discord 與論壇 – README 中連結的社群支援與討論頻道。
- Colab 入門指南 – 用於實際探索的互動式筆記本。
總結(TL;DR)
nnsight 為研究人員提供了一種簡潔、Python 風格的方法,用於 追蹤、讀取、修改 和 永久編輯 任何 PyTorch 模型的隱藏狀態,支援批量處理、梯度提取、遠端執行以及 vLLM 等高效能後端。它是現代基礎模型機制研究的真正、面向研究的工具。
相關
- 專案
- 專案
- 專案
- 專案
- 專案