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 中实验该架构,而无需深入研究原始研究代码库的用户。

相关

  • 项目
  • 项目
  • 项目
  • 项目