lucidrains/perceiver-pytorch
Implementation of Perceiver, General Perception with Iterative Attention, in Pytorch
perceiver‑pytorch – Perceiver 系列注意力模型的 PyTorch 實作
它是什麼 – 一個純 Python 函式庫,重新實作了論文《Perceiver: General Perception with Iterative Attention》與《Perceiver IO: A General Architecture for Structured Inputs & Outputs》中的 Perceiver 和 Perceiver IO 架構(以及一個小型實驗性變體)。這些模型是一類 Transformer 風格的神經網路,旨在透過先將輸入投影到一組較小的學習潛在向量,然後迭代應用交叉注意力與自注意力,來處理非常大且高維度的輸入(圖像、影片、音訊、點雲等)。
為什麼重要 – 原版 Perceiver 論文展示了單一架構可以處理多種模態,而無需針對特定模態進行設計。此儲存庫讓研究人員與開發者能輕鬆在 PyTorch 中實驗此概念,並將模型應用於圖像分類、語言建模或任何自定義任務。
快速開始
pip install perceiver-pytorch
圖像分類範例 (vanilla Perceiver)
import torch
from perceiver_pytorch import Perceiver
model = Perceiver(
input_channels=3, # RGB 圖像
input_axis=2, # 2‑D 資料 (高度 × 寬度)
num_freq_bands=6,
max_freq=10.,
depth=6,
num_latents=256,
latent_dim=512,
cross_heads=1,
latent_heads=8,
cross_dim_head=64,
latent_dim_head=64,
num_classes=1000,
attn_dropout=0.,
ff_dropout=0.,
weight_tie_layers=False,
fourier_encode_data=True,
self_per_cross_attn=2,
)
img = torch.randn(1, 224, 224, 3) # batch‑size‑1 ImageNet‑size 圖像
logits = model(img) # → shape (1, 1000)
彈性輸出 Perceiver IO
from perceiver_pytorch import PerceiverIO
model = PerceiverIO(
dim=32,
queries_dim=32,
logits_dim=100,
depth=6,
num_latents=256,
latent_dim=512,
cross_heads=1,
latent_heads=8,
cross_dim_head=64,
latent_dim_head=64,
weight_tie_layers=False,
seq_dropout_prob=0.2,
)
seq = torch.randn(1, 512, 32) # 任意輸入序列
queries = torch.randn(128, 32) # 解碼器查詢
out = model(seq, queries=queries) # → (1, 128, 100)
語言建模變體 (PerceiverLM)
from perceiver_pytorch import PerceiverLM
model = PerceiverLM(
num_tokens=20000,
dim=32,
depth=6,
max_seq_len=2048,
num_latents=256,
latent_dim=512,
cross_heads=1,
latent_heads=8,
cross_dim_head=64,
latent_dim_head=64,
weight_tie_layers=False,
)
seq = torch.randint(0, 20000, (1, 512))
mask = torch.ones(1, 512).bool()
logits = model(seq, mask=mask) # → (1, 512, 20000)
主要功能 (如 README 所述)
- 單行安裝:透過
pip安裝。 - 模組化 API:
Perceiver、PerceiverIO和PerceiverLM類別分別涵蓋分類、彈性輸出任務與語言建模。 - 內建傅立葉位置編碼:可透過
fourier_encode_data切換。 - 可配置參數:深度、潛在維度與注意力頭數皆可調整,以符合原論文或進行小型模型實驗。
- 實驗性自下而上注意力:變體可在
perceiver_pytorch.experimental.Perceiver下使用(增加類似 Set Transformers 的誘導集注意力區塊)。 - 準備好引用:儲存庫包含可直接複製的 Perceiver、Perceiver IO 及相關工作的 BibTeX 條目。
使用時機
- 您需要一個模態無關的骨幹網路,能夠攝取非常大的張量(例如高解析度圖像、影片幀、點雲)而不會耗盡記憶體。
- 您想要原型化跨模態或多模態研究,其中相同的編碼器可以重複用於不同的資料類型。
- 您正在探索彈性輸出形狀(例如分割圖、語言生成) – Perceiver IO 的查詢式解碼器使這變得簡單直接。
限制與注意事項
- 本函式庫僅提供模型定義;訓練迴圈、資料管線與效能優化由使用者自行處理。
- 不提供預訓練權重;您必須從頭開始訓練或載入自己的檢查點。
- 實驗性自下而上變體是選用的,需要從
experimental子模組匯入。
參考文獻 (來自 README)
- Perceiver: General Perception with Iterative Attention – arXiv:2103.03206
- Perceiver IO: A General Architecture for Structured Inputs & Outputs – arXiv:2107.14795
- 亦列出了相關注意力機制的額外引用。
總結 – perceiver-pytorch 是一個忠實且輕量級的 Perceiver 系列重新實作,適合任何想在 PyTorch 中實驗該架構,而無需深入研究原始研究程式碼庫的使用者。
相關
- 專案
- 專案
- 專案
- 專案