lucidrains/perceiver-pytorch

Implementation of Perceiver, General Perception with Iterative Attention, in Pytorch

perceiver‑pytorch – Perceiverファミリーの注意力モデルのPyTorch実装

概要 – 論文『Perceiver: General Perception with Iterative Attention』および『Perceiver IO: A General Architecture for Structured Inputs & Outputs』で提案された Perceiver および Perceiver IO アーキテクチャ(および小さな実験的バリアント)を再現した純粋なPythonライブラリです。これらのモデルは、非常に大きく高次元な入力(画像、動画、音声、点群など)を、まず学習された小さな潜在ベクトルセットに投影し、その後クロスアテンションとセルフアテンションを反復的に適用することで処理するように設計された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より)

  • 1行インストール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でこのアーキテクチャを実験したいすべての人に適しています。

関連

  • プロジェクト
  • プロジェクト
  • プロジェクト
  • プロジェクト