MAX 中 Step-3.5-Flash 架构解析:混合注意力、MoE 与 TP/EP/DP 并行实现指南
MAX 中 Step-3.5-Flash 架构解析混合注意力、MoE 与 TP/EP/DP 并行实现指南【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo本文围绕 MAX 平台max.pipelines.architectures.step3p5模块展开深入解读 MAX 如何将 Step-3.5-Flash 文本生成模型接入其推理管线从架构注册、模型配置到混合注意力、共享专家 MoE、三种并行模式与权重适配的完整实现。读者读完可掌握该架构模块的配置项语义、并行模式选择规则与源码级实现脉络便于在 MAX 中正确加载与调优该类模型。模块定位一条文档入口背后的完整架构文档 pipelines.architectures.step3p5.rst 是 Sphinx 的automodule文档入口其渲染内容即模块 step3p5/init.py 及其子模块的 API 文档。该模块对外导出四个核心对象Step3p5Config—— 模型配置类定义在 model_config.pyStep3p5Inputs/Step3p5Model—— 推理输入与管线模型定义在 model.pystep3p5_arch—— 架构注册实例定义在 arch.py模块目录中还包含注意力层layers/attention.py、MoE 门控layers/moe_gate.py、批处理器batch_processor.py、权重适配器weight_adapters.py与模型主体step3p5.py。该模块同时被 architectures/init.py、all_arches.bzl 引用并收录于文档索引 pipelines.architectures.rst 的文本生成架构列表中。架构注册如何让模型进入 MAX 管线MAX 的每个模型架构都通过一个SupportedArchitecture实例向管线系统声明其加载、配置与执行方式。arch.py 中的step3p5_arch注册信息如下字段值含义nameStep3p5ForCausalLM架构名与 HuggingFace 模型类型对应taskPipelineTask.TEXT_GENERATION文本生成任务example_repo_ids[stepfun-ai/Step-3.5-Flash]示例模型仓库default_weights_formatsafetensors默认权重格式default_encoding/supported_encodingsbfloat16默认及支持的量化编码pipeline_modelStep3p5Model管线模型类tokenizer/context_typeTextTokenizer/TextContext分词器与上下文类型weight_adapterssafetensors - convert_step3p5_state_dict权重转换适配器configStep3p5Config模型配置类batchingStep3p5BatchProcessor批处理类multi_gpu_supportedTrue支持多 GPUmemory_plannerPagedMemoryPlanner.with_activation_reservation(0, always_signal_buffersTrue)分页 KV 缓存内存规划supports_overlap_scheduler/supports_device_graph_captureFalse/False不支持重叠调度与设备图捕获其中Step3p5Model继承自AlwaysSignalBuffersMixin, LlamaModelBasemodel.py并固定norm_methodrms_norm、attention_biasFalse这决定了其基线结构复用 Llama3同时叠加 Step-3.5 特有的混合注意力与 MoE 层。Step3p5Config配置项全解Step3p5Config继承自Llama3Configmodel_config.py默认编码为bfloat16。其 Step-3.5 特有配置项如下配置项默认值说明num_attention_groups8KV 头组数全注意力层等价于 KV 头数head_dim128每个注意力头的维度sliding_window512滑动窗口注意力SWA的窗口大小layer_types[]逐层注意力类型full_attention或sliding_attentionsliding_num_attention_heads96SWA 层的注意力头数sliding_num_attention_groups8SWA 层的 KV 头组数per_layer_rope_theta[]逐层 RoPE theta为空时退化为单一rope_thetapartial_rotary_factors[]逐层部分旋转因子全注意力 0.5SWA 1.0yarn_only_types[]仅这些层类型应用rope_scaling如[full_attention]use_head_wise_attn_gateTrue是否启用逐头 sigmoid 注意力门控g_projmoe_num_experts288MoE 层路由专家数moe_top_k8每 token 激活的专家数moe_intermediate_size1280每个专家 MLP 的中间维度share_expert_dim1280共享专家 MLP 的中间维度moe_layersset()使用 MoE而非稠密 MLP的层索引集合moe_router_scaling_factor3.0路由专家权重缩放因子norm_expert_weightTrue是否将 top-k 专家权重归一化为和为 1swiglu_limits[]路由专家的逐层 SwiGLU 激活裁剪阈值0.0 表示不裁剪swiglu_limits_shared[]共享专家的逐层 SwiGLU 激活裁剪阈值配置初始化流程initialize()要求模型仓库包含合法config.json否则抛出明确错误model_config.py。initialize_from_config()则完成全部推导KV 参数调用construct_kv_params()构造混合 SWA 全注意力 KV 树内部委托hybrid_swa_full_kv_params同时读取attention_other_setting[num_attention_groups]作为滑动层的 KV 头数model_config.py注意力缩放calculate_attention_multiplier()返回sqrt(1.0 / head_dim)model_config.py逐层参数截断layer_types、per_layer_rope_theta、partial_rotary_factors、swiglu_limits均按num_hidden_layers截断config 中可能含有 MTP 层的多余条目MoE 层索引若moe_layers_enum非空则按逗号解析否则默认除第一层外全部为 MoEset(range(1, num_hidden_layers))HF 配置别名补齐_ensure_hf_config_aliases()为trust_remote_codeTrue场景补齐num_key_value_heads、rms_norm_eps、rope_scaling、hidden_act等标准字段并处理rope_theta为逐层列表时的折叠与恢复model_config.py。模型结构从模块 docstring 到实现step3p5.py 的模块 docstring 概括了 Step-3.5-Flash 实现的五个核心特性混合注意力全注意力 滑动窗口注意力逐层切换逐层 RoPE每层独立的 theta 与部分旋转因子MoE共享专家 sigmoid 路由 路由偏置零中心 RMSNormweight_offset1逐头注意力门控g_proj sigmoid。Transformer 块按层类型分派Step3p5TransformerBlock根据layer_types[layer_idx]决定注意力头数滑动层用sliding_num_attention_heads否则用num_attention_heads并根据layer_idx in config.moe_layers选择 MoE 或稠密 MLPstep3p5.py。每一层包含输入 LayerNorm、自注意力、后注意力 LayerNorm、MLP/MoE 与残差连接并根据并行模式决定 allreduce 的施加位置。逐层 RoPE 与部分旋转Step-3.5 的 RoPE 配置按层区分step3p5.py全注意力层partial_rotary_factor0.5旋转 64/128 维带rope_scalingSWA 层partial_rotary_factor1.0旋转 128/128 维无rope_scaling。_PartialRotaryEmbedding继承Llama3RotaryEmbedding只在前rotary_dim维产生真实旋转频率其余维补零为恒等cos1、sin0从而规避内核只支持交错式部分 RoPE 的限制step3p5.py。缩放YaRN必须在补零前施加否则会对零频做除零。所有 RoPE 对象按(theta, rotary_dim, use_scaling)键缓存freqs_cis的计算延迟到图上下文内执行。图构建与推理Step3p5Model._build_graph_for_compile()构建名为step3p5的 Graph输入顺序为tokens、input_row_offsets、return_n_logitsDP_EP 模式额外插入 CPU 上的host_input_row_offsets与data_parallel_splits随后是信号缓冲、扁平化 KV 缓存输入末尾追加 EP 通信缓冲model.py。input_types()中同样按此布局声明符号输入类型step3p5.py。注意力层QK 归一化、滑动窗口与逐头门控Step3p5Attention在标准注意力之上叠加四项特性attention.py逐头零中心 RMSNormQ/K、逐头 sigmoid 注意力门控、逐层全/滑动窗口注意力、逐层 RoPE。前向过程依次执行attention.pyfused_qkv_ragged_matmul融合 QKV 矩阵乘并写入 KV 缓存对 Q 施加q_norm零中心 RMSNormweight_offset1.0rms_norm_key_cache对 K 施加k_norm并归一化写入缓存per_head_normTruefused_qk_ragged_rope融合 QK 旋转嵌入flash_attention_ragged按层类型选择SLIDING_WINDOW_CAUSAL_MASK或CAUSAL_MASK滑动窗口大小由sliding_window传入若启用门控输出乘上sigmoid(g_proj(x))的逐头缩放最后经o_proj投影。切分策略上张量并行时 Q/K/V 投影采用rowwiseO 投影采用head_aware_columnwiseQK 归一化与门控投影分别采用 replicate / rowwiseattention.py。MoE共享专家、sigmoid 路由与激活裁剪Step3p5MoEWithSharedExpert实现output routed_moe(x) shared_expert(x)其中共享专家为始终开启的 MLP维度share_expert_dimstep3p5.py。路由专家与共享专家分别接受swiglu_limits[layer_idx]与swiglu_limits_shared[layer_idx]的激活裁剪通过make_concatenated_gated_activation_fn(ops.silu, limit)构造带裁剪的 SwiGLU 激活用于防止长上下文下的数值爆炸NaN。Step3p5MoEGate的路由算法moe_gate.pyscores sigmoid(gate(x))—— 门控计算强制在 FP32 中进行以匹配 HF 参考need_fp32_gateTruecorrected scores router_bias—— 加入可学习的逐专家路由偏置基于修正分数做 top-k 选择权重取自未修正的原始分数可选 top-k 权重归一化为和为 1norm_topk_prob乘以routed_scaling_factor。另外代码对 top-k 索引做了安全钳制clamp 到[0, num_experts-1]作为 NaN/Inf 分数导致越界索引的兜底防护moe_gate.py。router_bias以Weight(namerouter_bias, dtypefloat32)形式参与加载。并行模式TP_TP / TP_EP / DP_EPParallelismMode枚举定义了三种并行策略step3p5.py模式注意力MoE通信TP_TP张量并行张量并行无 EP注意力后与 MoE 后各一次 AllreduceTP_EP张量并行专家并行注意力后 AllreduceMoE 后由 EP combine 归约DP_EP数据并行每 rank 复制一份并处理各自的 batch 分片专家并行无注意力 AllreduceMoE 后由 EP combine 归约模式由_select_parallelism_mode()推导step3p5.py单 GPU 一律退化为TP_TP启用 EPep_size 1时data_parallel_degree 1选TP_EPdata_parallel_degree 设备数选DP_EP其他取值报错未启用 EP 而data_parallel_degree ! 1时报错提示 DP 注意力需要--ep-size 1。EP 配置由_create_ep_config()生成model.py要求ep_size能被 GPU 数整除dispatch_dtype由所选编码推导并通过calculate_ep_max_tokens_per_rank计算每 rank 最大 token 数。DP_EP 模式下return_logits仅支持LAST_TOKEN、且不允许返回 hidden states违反时在编译期直接报错model.py。DP_EP 的 logits 后处理会先在各设备上收集最后一个 token 的隐状态allgather聚合后再做归一化与词表并行 LM head 投影step3p5.py。此外模型支持按(is_sliding, is_moe)分组将同构层合并为子图step3p5_{sliding|full}_{moe|mlp}_block用于编译优化step3p5.py。权重适配从 HF checkpoint 到 MAX 布局weight_adapters.py 定义了convert_step3p5_state_dict的转换规则去除model.前缀堆叠专家权重拆分layers.{i}.moe.{proj}.weight [num_experts, ...]按专家逐一切片为layers.{i}.mlp.moe.experts.{j}.{proj}.weight门控重命名moe.gate.weight - mlp.moe.gate.gate_score.weight路由偏置moe.router_bias - mlp.moe.gate.router_bias并强制转为 float32共享专家share_expert. - mlp.share_expert.跳过 MTP 层层索引 ≥num_hidden_layers的权重及mtp.*前缀权重被跳过Step-3.5 的 config 可能包含 MTP 层部分 RoPE 权重置换对partial_rotary_factor 1.0的全注意力层Q/K 投影与 QK 归一化权重执行维度置换_build_partial_rope_perm使 MAX 内核的全头配对与 HF 参考的子空间配对等价该置换为自逆置换weight_adapters.py归一化权重原样传递零中心 RMSNorm 在运行时统一加 1。输入批处理与推理输入Step3p5BatchProcessor继承自Llama3EpBatchProcessor在三种模式间分派输入构造batch_processor.pyTP_TP直接复用 Llama3 批处理器TP_EP基础输入 EP 通信缓冲DP_EP额外构造 CPU 上的host_input_row_offsetsDP 模式下由_prepare_ep_moe_token_inputs生成与data_parallel_splits用于图内的 batch 分片。Step3p5Inputs的buffers属性按固定顺序拼装tokens、行偏移、返回 logits 数、DP 专属缓冲、信号缓冲、KV 缓存输入与 EP 输入model.py其中强制要求data_parallel_splits必须是 CPU Buffer否则断言失败以避免静默错放设备。相关文件索引文档入口pipelines.architectures.step3p5.rst / 架构索引pipelines.architectures.rst架构注册arch.py配置类model_config.py模型与输入model.py模型主体与并行模式step3p5.py注意力层layers/attention.pyMoE 门控layers/moe_gate.py批处理器batch_processor.py权重适配器weight_adapters.py构建目标BUILD.bazel以上内容均以当前仓库源码为准。需要说明的是本文描述的是该架构模块的实现事实具体的量化编码、并行规模等运行时能力请以实际部署环境与模型仓库配置为准。【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考