MAX Pipelines 的 LFM2 架构支持:混合全注意力 + 短卷积解码器实现解析
2026/9/12 20:36:49 网站建设 项目流程

MAX Pipelines 的 LFM2 架构支持:混合全注意力 + 短卷积解码器实现解析

【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo

导读

本文围绕 MAX 推理框架(Modular Platform 的一部分,即当前仓库max/目录所对应的 Python 服务端)中max.pipelines.architectures.lfm2模块展开,剖析它如何将 LiquidAI 的 LFM2 系列模型(如LiquidAI/LFM2.5-350M)接入 MAX Pipelines 推理管线。读者将了解 LFM2 的「全注意力 + 短卷积」混合架构在 MAX 中的图构建、权重适配、分页 KV Cache 与卷积状态缓存(conv-state cache)实现,以及如何在仓库中定位与验证这套支持代码。

文档定位:一个由 automodule 驱动的 API 模块页

关联文档 max/python/docs/pipelines.architectures.lfm2.rst 是 MAX Python 文档体系中max.pipelines.architectures命名空间下的一个 Sphinx 模块页。其正文只有一条automodule指令:

.. automodule:: max.pipelines.architectures.lfm2 :members: :imported-members: :show-inheritance:

也就是说,该文档的实质性内容完全来自被自动索引的 Python 模块max.pipelines.architectures.lfm2docs/目录下还包含pipelines.architectures.rstpipelines.rst等一系列同构页面(如pipelines.architectures.llama3.rstpipelines.architectures.qwen3.rst),共同构成架构层 API 文档。因此,理解本文主题的正确路径是直接阅读模块源码目录 max/python/max/pipelines/architectures/lfm2/,该目录包含:

  • arch.py— 架构注册(SupportedArchitecture
  • model.py— 管线模型LFM2Model与卷积状态缓存ConvStateCache
  • lfm2.py— 图级模型LFM2、解码层LFM2DecoderLayer、短卷积LFM2ShortConv、MLP
  • full_attention.py— 全注意力块LFM2FullAttention
  • model_config.py— 配置类LFM2Config
  • batch_processor.py— 批处理器LFM2BatchProcessor
  • weight_adapters.py— safetensors 权重名映射
  • __init__.py— 公共导出(ConvStateCacheLFM2ConfigLFM2InputsLFM2Modellfm2_arch

架构注册:LFM2 如何被 MAX Pipelines 识别

在 arch.py 中,模块通过SupportedArchitecture将 LFM2 注册为一种因果语言模型架构:

字段说明
name"Lfm2ForCausalLM"与 Hugging Face 模型 config 中的architectures字段对应
taskPipelineTask.TEXT_GENERATION文本生成任务
default_encoding/supported_encodings"float32"/{"float32", "bfloat16"}默认 FP32 推理,支持 BF16
example_repo_ids["LiquidAI/LFM2.5-350M", "LiquidAI/LFM2.5-350M-Base"]官方示例仓库 ID
tokenizerTextTokenizer文本 tokenizer
context_typeTextContext文本上下文
default_weights_formatWeightsFormat.safetensors默认 safetensors 权重
required_argumentsallow_safetensors_weights_fp32_bf16_bidirectional_cast=Truetrust_remote_code=True加载 LFM2 权重时必须开启的两个参数
multi_gpu_supportedFalse当前不支持多 GPU 数据并行
weight_adapterssafetensors →convert_lfm2_safetensor_state_dict权重名转换
configLFM2Config配置类
batchingLFM2BatchProcessor批处理
memory_plannerPagedMemoryPlanner分页内存规划
supports_overlap_scheduler/supports_device_graph_captureFalse不支持重叠调度与设备图捕获

从源码结构看,MAX Pipelines 通过arch_lookup(见 max/python/docs/pipelines.lib.arch_lookup.rst 对应的max.pipelines.lib模块)与各架构arch.py中的注册信息,根据 Hugging Faceconfig.jsonarchitectures字段自动选择Lfm2ForCausalLM实现。

__init__.py通过from .arch import lfm2_arch导出lfm2_arch,并公开LFM2ModelLFM2ConfigLFM2InputsConvStateCache,这正是 Sphinx 文档页:members:会索引的核心 API。

混合架构核心:全注意力层 + 短卷积层

LFM2 的关键设计是按层混合layer_types列表逐层标记每层是"full_attention"还是卷积层。这一结构在 lfm2.py 的LFM2LFM2DecoderLayer中体现:

if self.layer_type == "full_attention": self.self_attn = LFM2FullAttention(...) self.conv = None else: self.conv = LFM2ShortConv(config, linear_cls) self.self_attn = None
  • 全注意力层使用LFM2FullAttention(继承自 Qwen3.5 风格 GQA 实现,被 vendor 到 full_attention.py);
  • 非注意力层使用LFM2ShortConv短卷积,对局部窗口内的 token 做因果卷积(convolution over the sequence dimension),从而在保持全局建模能力的同时降低注意力层数量带来的开销。

LFM2在构造时统计layer_types中非full_attention的层数(num_conv_layers),KV 索引kv_idx只对全注意力层递增,卷积层不占用 KV cache 页。

短卷积实现(LFM2ShortConv)

LFM2ShortConv 的图级实现要点:

  • in_proj将隐藏状态投影为 3 份(bcxv),计算bx = b * xv(类似门控线性单元);
  • conv_weight形状为[hidden, 1, kernel_size]kernel_size来自配置conv_L_cache(默认 3);
  • 通过ops.while_loop逐 token 滑动窗口:每个 token 从对应请求的[1, hidden, K]状态中取出、滑入新 token(state_r_next = ops.concat((state_r[:, :, 1:], bx_t_k), axis=2)),再用scatter_nd写回状态栈;
  • 卷积输出conv_out = sum(state_r_next * conv_w, axis=2),最终y = c * conv_out经过out_proj输出,并返回更新后的new_state

代码注释特别指出,scatter_nd是 GPU 友好的(ops.scatter会经由 CPU round-trip,见 TODO(GEX-2197)),因此在循环内更新状态栈使用scatter_nd。这是从源码注释可确认的工程取舍。

全注意力块(LFM2FullAttention)

full_attention.py 实现了 Qwen3.5 风格的全注意力路径,与qwen3_5_moe.layers.attention逻辑一致(模块头部 docstring 明确说明该实现被 vendor 于此,避免 LFM2 依赖整个 qwen3_5_moe 架构包):

  • per-head Q/K RMSNormq_norm/k_norm,Qwen3 风格;
  • 可选输出门(attn_output_gate:开启时q_proj输出维翻倍,checkpoint 按「逐头交错」布局存放[Q_h0, Gate_h0, Q_h1, Gate_h1, ...],因此必须 reshape 为[seq_len, n_heads, 2*head_dim]后沿最后一维切分,而非扁平切分;门在注意力输出后以attn_output * sigmoid(gate)形式生效(非 silu);
  • 非融合 KV 路径:门开启时fused_qkv_ragged_matmul无法处理翻倍的 Q 维,故改用matmul_kv_cache_raggedconcat(k_proj.weight, v_proj.weight)写入分页 cache)并配合rms_norm_key_cache对缓存中的 K 做 per-head norm;
  • partial RoPEfreqs_cis.shape[-1] == head_dim * partial_rotary_factorfused_qk_ragged_rope只旋转前导维度;
  • 因果 Flash Attention:flash_attention_ragged,支持local_window_size(默认 512)与MHAMaskVariant.CAUSAL_MASK
  • 张量并行:实现Shardable,支持 replicate 与 tensor_parallel 两种分片策略;门开启时q_proj使用gate_up分片保证每个分片拿到[q_shard | gate_shard]o_proj使用head_aware_columnwise,KV 投影使用rowwise。需要注意:该支持是图内张量并行,与arch.pymulti_gpu_supported=False并不矛盾(后者针对数据并行多副本)。

配置类 LFM2Config:从 Hugging Face config 到图构建参数

model_config.py 中LFM2Config(Llama3Config)新增字段:

字段默认值来源(HF config 键)说明
layer_types[]layer_types每层类型,"full_attention"或卷积层
conv_L_cache3conv_L_cache卷积缓存窗口长度(kernel size)
conv_biasFalseconv_bias卷积是否带偏置
norm_eps1e-5norm_eps归一化 epsilon

initialize_from_config从 HFAutoConfig读取这些字段,并处理rope_parameters.rope_theta(默认DEFAULT_ROPE_THETA = 10000.0)与可选 rope 字段缺失的兼容(_ensure_optional_rope_fields)。

值得注意的两个兼容细节(均有源码注释佐证):

  1. norm_epsvsrms_norm_eps:LFM2 将归一化 epsilon 存在norm_eps而非 LLaMA 风格的rms_norm_epsfinalize()先以norm_method="layer_norm"调用父类以跳过对rms_norm_eps的属性读取,随后再改回rms_norm_eps并显式赋值self.rms_norm_eps = norm_eps
  2. tie_embeddingvstie_word_embeddings:transformers ≥ 5 将 LiquidAI 的自定义tie_embedding键并入标准tie_word_embeddings,代码先读标准键再回退到旧键。

_resolve_intermediate_size还处理block_auto_adjust_ff_dim:SwiGLU FFN 有两个门控投影,有效宽度按 2/3 折算(与 LLaMA 系列同约定),并依次应用block_ffn_dim_multiplierblock_multiple_of对齐。

LFM2图模型(lfm2.py)的组件包括:embed_tokens嵌入、rope(由create_rope_embedding创建,interleaved_rope_weights=False)、layers层列表、norm(RMSNorm)与lm_head;当tie_word_embeddings为真时lm_head共享embed_tokens.weight。前向流程为:嵌入 → 逐层(残差 + 注意力/卷积 + SwiGLU FFN)→logits_postprocess(含可选的 logits 缩放)→ 拼接各卷积层新状态。

推理状态管理:ConvStateCache 与卷积状态生命周期

卷积层没有 KV cache,但需要跨 token 维护滑动窗口状态,因此 MAX 在管线侧实现了专门的 ConvStateCache:

  • 每个 slot 为每层分配Buffer.zeros([1, hidden_size, conv_L_cache], dtype, device)
  • claim(request_id)为请求申请 slot(无空闲 slot 时抛RuntimeError("No free LFM2 conv-state slots.")),release(request_id)归还 slot;
  • get_states(request_ids):单请求时零拷贝直接返回 slot buffer;多请求(N > 1)时按层沿 batch 维拼接为[N, hidden, kernel]
  • update_states(request_ids, new_states):执行图输出新状态后按请求写回各自 slot。

拼接/切分通过 numpy round-trip 完成(_cat_buffers/_split_buffer_dim0),其中对 bfloat16 的处理非常细致:numpy 不支持 bfloat16,DLPack 转换会抛RuntimeError: Unsupported dtype in DLTensor,因此先将 bf16 视图为同字节宽的uint16走 numpy,再视图回原 dtype。由于卷积状态本身很小(每 slot 为hidden * kernel),这种 round-trip 开销可接受——这是源码 docstring 明确给出的设计权衡。

LFM2Model(model.py)继承Llama3ModelBase

  • _create_model_configLFM2Config.initializefinalize
  • _build_graph_for_compile构建名为"lfm2"Graph:输入为 tokens、input_row_offsetsreturn_n_logits、展开的 KV 输入、以及每个卷积层一个[conv_batch, hidden, conv_L_cache]状态输入(见LFM2.input_types),输出包含 logits(数量由_num_logit_outputs决定,依据ReturnLogits/ReturnHiddenStates)与各卷积层新状态;
  • execute将图输出切成 logits 与新状态两部分,新状态回写ConvStateCache
  • release(request_id)释放卷积状态 slot。

批处理与权重适配

LFM2BatchProcessor 继承Llama3BatchProcessor,在基类 ragged batching 基础上叠加卷积状态 slot:

  • bind_conv_cacheLFM2Model.__init__在构造后调用,将ConvStateCache注入批处理器;
  • prepare_initial_token_inputs先调父类准备 KV/输入,再对扁平化后的请求列表逐个claim,最后get_states组装LFM2Inputs
  • LFM2InputsLlama3Inputs基础上增加conv_states: list[Buffer]request_ids: list[RequestID]buffers属性把卷积状态追加到父类 buffer 列表之后,保证与图输入顺序一致。

weight_adapters.py 定义了 HF checkpoint → MAX 命名映射:

LFM2_SAFETENSOR_MAPPING = ( ("model.embed_tokens", "embed_tokens"), ("model.embedding_norm", "norm"), ("model.layers", "layers"), ("self_attn.out_proj", "self_attn.o_proj"), ("self_attn.q_layernorm", "self_attn.q_norm"), ("self_attn.k_layernorm", "self_attn.k_norm"), ("conv.conv.weight", "conv.conv_weight"), ("conv.conv.bias", "conv.conv_bias"), )

convert_lfm2_safetensor_state_dict对 state dict 中每个键按序做字符串替换(如self_attn.q_layernormself_attn.q_normconv.conv.weightconv.conv_weight),完成权重名对齐。该适配器与arch.pyweight_adapters={WeightsFormat.safetensors: ...}对应,是_build_graph_for_compilemodel.load_state_dict(state_dict, ...)之前必经的一步。

如何在仓库中验证与运行

  • 单元/集成证据:LFM2 架构测试位于 max/tests/ 下,模型集成测试可参考max/python/test/max/tests/integration/中的 pipelines 测试组织方式(当前仓库以 Bazel 管理,相关测试目标定义于对应 BUILD.bazel)。
  • 注册证据max/python/max/pipelines/architectures/__init__.pyall_arches.bzl中均列出lfm2,确认其参与架构自动发现。
  • 实际使用前提:加载 LFM2 权重必须满足arch.pyrequired_arguments的两项——trust_remote_code=True(LiquidAI 模型包含自定义建模代码)与 safetensors 的 FP32/BF16 双向转换开关;且当前注册信息显示不支持多 GPU 数据并行(multi_gpu_supported=False)、不支持重叠调度与设备图捕获。
  • 文档入口:完整的架构级 API 索引见 max/python/docs/pipelines.architectures.rst,模块页即本文开头的pipelines.architectures.lfm2.rst;服务部署与运行参数请参考 max/docs/serve/ 与 max/docs/get-started.mdx。

小结

max.pipelines.architectures.lfm2为 LFM2 系列混合架构模型提供了一套完整的 MAX Pipelines 支持:通过SupportedArchitecture注册让Lfm2ForCausalLM可被自动发现;LFM2/LFM2DecoderLayerlayer_types混合实例化 Qwen3.5 风格全注意力块与短卷积块;ConvStateCacheLFM2BatchProcessor负责卷积状态在分页推理与 ragged batching 下的生命周期管理;LFM2Configconvert_lfm2_safetensor_state_dict分别完成 HF config 与权重名的兼容对齐。这套实现既复用了 LLaMA 3 管线基座,又以 vendor 方式引入 Qwen3.5 注意力逻辑,是研究「如何在大型推理框架中接入混合注意力新架构」的典型样本。

【免费下载链接】mojoThe Modular Platform (includes MAX & Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询