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