深入解析 Transformers 中的 MRA 模型:多分辨率近似自注意力机制的架构、配置与实战
MRA(Multi Resolution Analysis,多分辨率分析)是 🤗 Transformers 中一类基于经典小波与多分辨率思想近似自注意力矩阵的高效 Transformer 模型,专门用于在保持精度的同时显著压缩自注意力的计算与显存开销。本文以 MRA 官方模型文档 为核心主线,结合 src/transformers/models/mra/ 下的配置源码、PyTorch 实现与转换脚本,完整讲解 MRA 的原理、MraConfig 关键参数、内部注意力计算流程、六大任务模型头的使用方法以及权重转换与测试验证方式,帮助你直接用 AutoModel 生态无缝复用 MRA 模型。
MRA 模型概述:为什么要做"多分辨率"自注意力
MRA 模型由 Zhanpeng Zeng、Sourav Pal、Jeffery Kline、Glenn M Fung 和 Vikas Singh 在论文 Multi Resolution Analysis (MRA) for Approximate Self-Attention 中提出(对应文档标注的论文发布时间为 2022-07-21),并于 2023-07-10 被合入 Hugging Face Transformers 仓库。该实现由社区贡献者 [novice03] 贡献,原始的论文代码库为 mlpen/mra-attention。
从论文摘要可以提炼出它的核心动机:Transformer 在 NLP 与视觉任务中已经成为主流模型,但自注意力矩阵是架构中的核心计算瓶颈。业界此前已经提出了许多近似手段,包括预定义的稀疏模式(prespecified sparsity patterns)、**低秩基展开(low-rank basis expansions)以及二者的组合。MRA 的独特之处在于,它回头重访了信号处理领域的经典多分辨率分析(Multiresolution Analysis, MRA)**概念——尤其是小波(Wavelets)——这些工具在自注意力近似场景中长期被低估。
MRA 的核心思想并不复杂:先在低分辨率下快速概览整条序列的"块与块"之间的相关性(类似把序列按块做平均池化后跑一遍小规模注意力),再据此选出少数真正重要的高分辨率块,只在选出的块上计算精细的自注意力。这种"低分辨率全局引导 + 高分辨率局部精算"的两级方案,是理解 MRA 一切代码细节的主线。
仓库中的 MRA 实现布局
在当前仓库中,MRA 的全部实现集中在一个目录下,结构非常清晰:
- 配置类源码:定义
MraConfig与全部默认超参; - PyTorch 模型实现:约 1300 行,包含从自注意力到六个下游任务头的完整实现;
- 权重转换脚本:用于把原论文仓库
mlpen/mra-attention的 checkpoint 转换为 Transformers 格式; - 模型测试文件:
MraModelTester负责覆盖全部模型头的正确性; - 模块入口 __init__.py:通过
_LazyModule惰性导出,只有实际使用时才加载依赖。
在架构层面,MRA 沿用了 BERT 风格的 encoder-only 骨架,因此它天然适合 mask 语言建模以及各种句子/词元级下游任务,而不是自回归生成。modeling_mra.py 中大量模块(如 MraSelfOutput、MraIntermediate、MraOutput、MLM 头等)都明确标注了 Copied from transformers.models.bert.modeling_bert.*,也就是说其前馈网络、残差连接与 LayerNorm 结构与 BERT 完全一致,真正的差异集中在自注意力模块 MraSelfAttention 与 mra2_attention。
MraConfig:四个专属参数与完整默认值
与近似策略直接相关的四个专属参数
MraConfig 在标准 BERT 式配置之外,专门定义了四个用于控制"多分辨率近似预算"的参数。理解这四个参数是调好 MRA 模型的前提:
| 参数 | 类型/默认值 | 作用 |
|---|---|---|
block_per_row |
int,默认 4 |
设置高分辨率尺度的预算,即序列每一行中参与高分辨率精算的块数上限,直接决定近似带来的计算节省幅度 |
approx_mode |
str,默认 "full" |
控制是否同时使用低分辨率与高分辨率近似:"full" 表示两者都参与(推荐,精度更佳);"sparse" 表示只用低分辨率来选择块、且最终只用高分辨率结果(仅做 top-k 稀疏注意力的近似) |
initial_prior_first_n_blocks |
int,默认 0 |
前 N 个块强制使用高分辨率(作为"初始先验"),用于保证序列开头位置始终被精细建模 |
initial_prior_diagonal_n_blocks |
int,默认 0 |
主对角线附近 N 个块强制使用高分辨率(对角线先验),用于保证相邻 token 之间的局部注意力不丢失 |
需要特别强调的是 block_per_row 的含义:它在模型初始化时会被换算成真正的"每行块预算"self.num_block,计算公式见 modeling_mra.py:
self.num_block = (config.max_position_embeddings // 32) * config.block_per_row
self.num_block = min(self.num_block, int((config.max_position_embeddings // 32) ** 2))
其中常量 32 是 CUDA kernel 依赖的 GPU warp 宽度(详见后文)。以 max_position_embeddings=512、block_per_row=4 为例:序列被切成 512 / 32 = 16 个块,每行最多只精算 16 * 4 = 64 个块条目,而完整稠密注意力每行要算 16 * 16 = 256 个条目——这正是 MRA 降低复杂度的来源。min(..., num_block²) 这层截断保证当块数很少(短序列)时预算不会超过全连接的上限。
完整的默认配置值
MraConfig 采用 model_type = "mra",其余与 BERT 同构的默认值如下(源码见 configuration_mra.py):
vocab_size=50265、hidden_size=768、num_hidden_layers=12、num_attention_heads=12intermediate_size=3072、hidden_act="gelu"、hidden_dropout_prob=0.1、attention_probs_dropout_prob=0.1max_position_embeddings=512、type_vocab_size=1、initializer_range=0.02、layer_norm_eps=1e-5pad_token_id=1、bos_token_id=0、eos_token_id=2tie_word_embeddings=True(MLM 解码器与输入词嵌入共享权重)、add_cross_attention=False
配置文档自带的最小使用示例(从 configuration_mra.py 原样继承)如下,uw-madison/mra-base-512-4 是默认文档串引用的参考 checkpoint 命名风格:
>>> from transformers import MraConfig, MraModel
>>> # 初始化一个 uw-madison/mra-base-512-4 风格的配置
>>> configuration = MraConfig()
>>> # 基于配置初始化一个(随机权重的)模型
>>> model = MraModel(configuration)
>>> # 访问模型配置
>>> configuration = model.config
源码级原理:MRA 两级注意力到底怎么算
MRA 的注意力实现集中在 modeling_mra.py 的 mra2_attention 函数与配套的工具函数中。整条计算链可以拆成四步,理解它能帮你判断改哪个超参会产生什么影响。
第一步:把序列切成块并构造低分辨率"概览"。 模型要求序列长度能被 block_size=32 整除。get_low_resolution_logit(L272-L309)先把 query/key(在 mask 存在时按有效 token 计数做平均)重塑为 (batch, seq_len/32, 32, head_dim) 并在块内取均值,得到每块一个代表向量的"低分辨率"Q/K;随后在块粒度上做一次完整的小矩阵乘法 query_hat @ key_hat^T,得到一个 (batch, num_block, num_block) 的块间相关矩阵,即低分辨率 logits。若提供了注意力 mask,还会用每个块内的有效 token 数把空块的行列压到极小值。
第二步:按低分辨率打分挑选高分辨率块。 get_block_idxes(L312-L347)在低分辨率 logits 上做全局 top-k,选出每行得分最高的若干"query 块 × key 块"组合,k 就是前文换算出的 num_block。若 initial_prior_first_n_blocks > 0,会给前 N 列/行块加上 5e3 的偏置强制入选;若 initial_prior_diagonal_n_blocks > 0,则会给以主对角线为中心的带状块同样加偏置。这里的逻辑会生成两类输出:indices(被选中的块坐标)以及(仅 "full" 模式下)high_resolution_mask(低分辨率 logits 中达到阈值的块掩码)。
第三步:只在选中的块上做高分辨率稠密计算。 这一步不生成完整的 seq_len × seq_len 矩阵,而是借助自定义 CUDA kernel 做采样稠密矩阵乘法(Sampled Dense Matrix Multiplication, SDDMM)。仓库中通过 load_cuda_kernels()(L52-L58)从 kernels-community/mra 拉取并缓存 kernel,然后依次调用:
mm_to_sparse:在选中的块上计算Q·K^T,得到稀疏高分辨率 logits;sparse_max:在稀疏 logits 上做按行 max,用于 softmax 数值稳定性;sparse_dense_mm:把稀疏注意力权重与稠密 V 相乘,得到稀疏上下文;MraReduceSum:累加稀疏注意力权重得到归一化因子。
这些算子被封装成可自定义反向传播的 torch.autograd.Function——MraSampledDenseMatMul 与 MraSparseDenseMatMul(L196-L240),其 backward 通过 transpose_indices 转置稀疏索引复用同一批 kernel,从而支持端到端训练,而不是只做推理近似。
第四步:按 approx_mode 融合两级结果。 这是 mra2_attention 中分支最复杂的部分(L384-L462):
- 在
"full"模式下,低分辨率分支把块均值池化后的 V 乘上低分辨率注意力权重再广播回每个 token(repeat(1, 1, block_size, 1)),并用log_correction修正高、低分辨率行最大值不一致带来的尺度偏移,最后按(高分辨率输出 + 低分辨率输出) / (两个归一化因子之和)融合; - 在
"sparse"模式下,低分辨率仅在torch.no_grad()中用于选块,最终上下文只来自高分辨率分支high_resolution_attn_out / high_resolution_normalizer。
值得注意的工程细节:由于 CUDA kernel 对 GPU warp 大小(32)最友好,MraSelfAttention.forward 在 attention_head_size < 32 时会把 Q/K/V 的最后一维用 0 填充到 32,计算完后再截回原维度;同时在 MraSelfAttention.__init__ 里,只要满足 torch CUDA 可用 + CUDA 平台 + ninja 可用 就会尝试加载 kernel(L527-L532),加载失败则降级为告警——这意味着没有 CUDA kernel 环境时 MRA 无法获得真正的加速,这是使用前必须确认的前提。另外,从 mra2_attention 的入口断言可以看出,序列长度必须能被 block_size=32 整除,这也是后续使用 MRA 时需要 pad 序列的硬性约束。
模型类族:MRA 提供的六个任务模型
MRA 沿 BERT 惯例把共享骨架 MraModel 封装成 6 个开箱即用的类,全部注册在文档与 modeling_mra.py 的 __all__ 中。它们的共有特性是:都支持 output_hidden_states 与 return_dict,并在 return_dict=True 时返回对应语义的 ModelOutput 对象。
MraModel:基础编码器
MraModel 由 MraEmbeddings(word + position + token_type 三份嵌入、LayerNorm 与 dropout)和 MraEncoder(堆叠若干 MraLayer)组成。输入参数包括 input_ids、attention_mask、token_type_ids、position_ids、inputs_embeds 等,返回 BaseModelOutputWithCrossAttentions。注意 MRA 是双向(bidirectional)注意力模型,其 mask 通过 create_bidirectional_mask 统一构造。MraPreTrainedModel 声明 base_model_prefix = "mra" 且 supports_gradient_checkpointing = True,因此配合 Trainer 时可直接启用梯度检查点以降低显存。
MraForMaskedLM:掩码语言建模
MraForMaskedLM 在 MraModel 之上叠加了与 BERT 相同的 MLM 头(MraLMPredictionHead:先过一个 hidden→hidden 的 dense + GELU + LayerNorm 变换,再投影回词表)。它声明了 _tied_weights_keys 把 cls.predictions.decoder.weight 与输入词嵌入 mra.embeddings.word_embeddings.weight 绑定(对应 tie_word_embeddings=True)。labels 中 -100 位置被忽略,仅对 [0, vocab_size) 的 token 计算交叉熵。
MraForSequenceClassification:句子级分类/回归
MraForSequenceClassification 使用 MraClassificationHead——取 <s>(等价 [CLS])位置的向量,过 dropout → dense → 激活 → dropout → out_proj。它的 loss 逻辑是自适应的:num_labels == 1 时走 MSE 回归;否则根据 labels 的 dtype 自动判定单标签/多标签分类并选择交叉熵或 BCEWithLogitsLoss。
MraForMultipleChoice:多项选择
MraForMultipleChoice 的输入形状是 (batch_size, num_choices, sequence_length),模型内部先把前三维拍平送入 MraModel,取每个序列的 [CLS] 向量过 pre_classifier + ReLU + classifier(→1),最终 reshape 回 (batch, num_choices) 计算交叉熵。
MraForTokenClassification:词元级分类
MraForTokenClassification 直接在 sequence_output 上加 dropout 和一个 Linear(hidden_size, num_labels)。计算 loss 时,它用 attention_mask 把非有效位置替换为 loss_fct.ignore_index,只对真实 token 计算交叉熵。
MraForQuestionAnswering:抽取式问答
MraForQuestionAnswering 构造时强制 config.num_labels = 2,通过 qa_outputs 输出 start/end logits,loss 为起点与终点交叉熵的平均值;越界的 start/end 位置会被 clamp 到序列长度并用 ignore_index 忽略。
快速上手:加载、推理与 AutoModel 生态
由于 MRA 已完整接入 Transformers 的注册表(model_type="mra"),你可以像使用 BERT 一样通过 AutoModel/AutoTokenizer 体系加载。测试文件中出现过的实际 checkpoint 标识如 uw-madison/mra-base-512-4(MaskedLM 任务,512 长度)与 uw-madison/mra-base-4096-8-d3(长序列变体,可参见 test_modeling_mra.py 的加载路径)可作为命名参考。一个完整的 MLM 使用片段如下:
>>> from transformers import AutoTokenizer, MraForMaskedLM
>>> import torch
>>> tokenizer = AutoTokenizer.from_pretrained("uw-madison/mra-base-512-4")
>>> model = MraForMaskedLM.from_pretrained("uw-madison/mra-base-512-4")
>>> inputs = tokenizer("The capital of France is [MASK].", return_tensors="pt")
>>> with torch.no_grad():
... logits = model(**inputs).logits
>>> mask_index = (inputs["input_ids"] == tokenizer.mask_token_id)[0].nonzero(as_tuple=True)[0]
>>> predicted_token_id = logits[0, mask_index].argmax(dim=-1)
>>> tokenizer.decode(predicted_token_id)
实际动手时请务必遵守本仓库代码验证过的三条约束:
- 序列长度必须是 32 的倍数(kernel 的块大小限制),长于
max_position_embeddings的输入需要自行做截断与 padding 策略; - 近似预算参数要对应长度配置:改
max_position_embeddings时建议同步评估block_per_row,可参考测试配置中max_position_embeddings=64(2 个块)与序列长度对齐的做法(见 test_modeling_mra.py); - 追求精度的训练任务使用默认
approx_mode="full",追求极致稀疏/速度时可尝试"sparse"。
权重转换:从原论文仓库复现 checkpoint
如果你需要把 mlpen/mra-attention 官方仓库训练的 checkpoint 迁移到 Transformers 生态,可以直接使用仓库自带的 convert_mra_pytorch_to_pytorch.py,其入口函数为 convert_mra_checkpoint(checkpoint_path, mra_config_file, pytorch_dump_path),命令行用法为:
python src/transformers/models/mra/convert_mra_pytorch_to_pytorch.py \
--pytorch_model_path /path/to/mra_checkpoint.pt \
--config_file /path/to/mra_config.json \
--pytorch_dump_path /path/to/output_dir
转换脚本的核心是 rename_key(L23-L60):它把原仓库形如 transformer_{i}.mha.attn.W_q、norm1/norm2、ff.0/ff.2、mlm_class、backbone.backbone.encoders 这类命名逐一映射到 Transformers 的 encoder.layer.{i}.attention.self.query、attention.output.LayerNorm、intermediate.dense、cls.predictions.decoder 等标准命名。此外,convert_checkpoint_helper 还会显式重建 position_ids 缓冲(从 2 开始,因为 MRA 的 position embedding 表在 MraEmbeddings 中被扩成了 max_position_embeddings + 2),并跳过原仓库中与 pooler、sentence 分类相关的无关权重。
测试验证:MRA 如何被守护
MRA 的回归测试位于 tests/models/mra/test_modeling_mra.py。MraModelTester 使用的微型配置充分体现了 MRA 的约束与特性:seq_length=64、max_position_embeddings=64(注释明确指出序列长度必须等于最大长度且是 block_size 32 的倍数)、hidden_size=16、num_attention_heads=2 等。测试类同时覆盖了梯度检查点训练的多个变体(use_reentrant 的 true/false 分支),说明模型前向与自定义 CUDA autograd 算子在反向传播与重计算场景下都经过了验证;另有专项测试会真实调用 MraModel.from_pretrained("uw-madison/mra-base-512-4") 与 MraForMaskedLM.from_pretrained(...) 来比对 checkpoint 输出。
使用边界与结论
汇总上文,MRA 在 Transformers 中的定位是面向中长序列、可训练的高效 encoder 自注意力方案,其工程实现强依赖 CUDA 自定义 kernel 与 warp 对齐的 32 尺寸假设,且只支持双向注意力(适合预训练/理解类任务,不适合因果解码生成)。在动手之前,请确认运行环境具备可加载 kernels-community/mra 的 GPU 与 ninja,并将序列长度规整为 32 的倍数。在这些前提满足的前提下,你可以通过 MraConfig 的 block_per_row、approx_mode 与两类 initial prior 参数,在精度与算力之间自由调节近似粒度,并在六个现成任务模型头上直接开展 MLM、GLUE 式分类、多项选择、NER 与抽取式问答等下游实验。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0625
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00