首页
/ 深入解析 Transformers 中的 MRA 模型:多分辨率近似自注意力机制的架构、配置与实战

深入解析 Transformers 中的 MRA 模型:多分辨率近似自注意力机制的架构、配置与实战

2026-09-07 14:37:12作者:宣利权Counsellor

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 中大量模块(如 MraSelfOutputMraIntermediateMraOutput、MLM 头等)都明确标注了 Copied from transformers.models.bert.modeling_bert.*,也就是说其前馈网络、残差连接与 LayerNorm 结构与 BERT 完全一致,真正的差异集中在自注意力模块 MraSelfAttentionmra2_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=512block_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=50265hidden_size=768num_hidden_layers=12num_attention_heads=12
  • intermediate_size=3072hidden_act="gelu"hidden_dropout_prob=0.1attention_probs_dropout_prob=0.1
  • max_position_embeddings=512type_vocab_size=1initializer_range=0.02layer_norm_eps=1e-5
  • pad_token_id=1bos_token_id=0eos_token_id=2
  • tie_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.pymra2_attention 函数与配套的工具函数中。整条计算链可以拆成四步,理解它能帮你判断改哪个超参会产生什么影响。

第一步:把序列切成块并构造低分辨率"概览"。 模型要求序列长度能被 block_size=32 整除。get_low_resolution_logitL272-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_idxesL312-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——MraSampledDenseMatMulMraSparseDenseMatMulL196-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.forwardattention_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_statesreturn_dict,并在 return_dict=True 时返回对应语义的 ModelOutput 对象。

MraModel:基础编码器

MraModelMraEmbeddings(word + position + token_type 三份嵌入、LayerNorm 与 dropout)和 MraEncoder(堆叠若干 MraLayer)组成。输入参数包括 input_idsattention_masktoken_type_idsposition_idsinputs_embeds 等,返回 BaseModelOutputWithCrossAttentions。注意 MRA 是双向(bidirectional)注意力模型,其 mask 通过 create_bidirectional_mask 统一构造。MraPreTrainedModel 声明 base_model_prefix = "mra"supports_gradient_checkpointing = True,因此配合 Trainer 时可直接启用梯度检查点以降低显存。

MraForMaskedLM:掩码语言建模

MraForMaskedLMMraModel 之上叠加了与 BERT 相同的 MLM 头(MraLMPredictionHead:先过一个 hidden→hidden 的 dense + GELU + LayerNorm 变换,再投影回词表)。它声明了 _tied_weights_keyscls.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)

实际动手时请务必遵守本仓库代码验证过的三条约束:

  1. 序列长度必须是 32 的倍数(kernel 的块大小限制),长于 max_position_embeddings 的输入需要自行做截断与 padding 策略;
  2. 近似预算参数要对应长度配置:改 max_position_embeddings 时建议同步评估 block_per_row,可参考测试配置中 max_position_embeddings=64(2 个块)与序列长度对齐的做法(见 test_modeling_mra.py);
  3. 追求精度的训练任务使用默认 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_keyL23-L60):它把原仓库形如 transformer_{i}.mha.attn.W_qnorm1/norm2ff.0/ff.2mlm_classbackbone.backbone.encoders 这类命名逐一映射到 Transformers 的 encoder.layer.{i}.attention.self.queryattention.output.LayerNormintermediate.densecls.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.pyMraModelTester 使用的微型配置充分体现了 MRA 的约束与特性:seq_length=64max_position_embeddings=64(注释明确指出序列长度必须等于最大长度且是 block_size 32 的倍数)、hidden_size=16num_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 的倍数。在这些前提满足的前提下,你可以通过 MraConfigblock_per_rowapprox_mode 与两类 initial prior 参数,在精度与算力之间自由调节近似粒度,并在六个现成任务模型头上直接开展 MLM、GLUE 式分类、多项选择、NER 与抽取式问答等下游实验。

登录后查看全文
热门项目推荐
相关项目推荐