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》中的 PerceiverPerceiver 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 安裝。
  • 模組化 APIPerceiverPerceiverIOPerceiverLM 類別分別涵蓋分類、彈性輸出任務與語言建模。
  • 內建傅立葉位置編碼:可透過 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 中實驗該架構,而無需深入研究原始研究程式碼庫的使用者。

相關

  • 專案
  • 專案
  • 專案
  • 專案