
max.nn.kernels 内核封装模块全解析MAX 平台 KV 缓存与注意力算子的 Python 入口【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo导读max.nn.kernels是 MAXModular Platform包含 MAX 与 MojoPython 前端中面向 GPU 推理的内核封装层它以约 1.1 万行、170 余个函数的规模将 Mojo/C 侧编译好的mo.*自定义算子以类型安全、带严格输入校验的 Python API 暴露给上层流水线与模型架构代码。本文以仓库中 nn.kernels.rstSphinxautomodule文档入口为主题结合 kernels.py 源码逐族讲解其设计思想、核心 API、量化/精度约束与调用方式帮助你掌握如何在 MAX 平台上编写自定义注意力与 KV 缓存内核或理解现有架构代码Qwen、DeepSeek、MiniMax 等底层算子的行为。模块定位包装器而非内核实现max.nn.kernels的模块文档字符串只有一句话Helper functions for wrapping custom kv cache/attention related ops.这句话界定了它的角色它不实现注意力、RoPE、MoE 的数学逻辑而是作为 Python 侧的统一接线层通过三个底层机制把图构建期的算子描述映射到 Mojo 内核ops.inplace_custom(name, device, values, out_types, parameters)声明一个带字符串名字如mo.fused_qkv_matmul.ragged.paged的自定义算子可原地读写 KV 缓存ops.custom(...)非原地版本用于无缓存耦合的纯计算内核如mo.mha.no_cache、mo.grouped.matmul.raggedGraph.current._add_op_generated(mo.CompositeXxxOp, ...)生成一等复合算子composite op如CompositeGroupedMatmulBlockScaledOp、CompositeMaskedFlashAttentionGpuOp便于图编译器以类型而非字符串做匹配与融合kernels.py。从调用链看mo.*内核名遵循统一的命名模式mo.功能.ragged|padded.paged|no_cache[.量化后缀]例如内核字符串语义mo.fused_qkv_matmul.ragged.paged分块ragged输入的融合 QKV 投影 分页缓存写入mo.fused_qkv_matmul.padded.paged填充padded输入版本mo.mha.ragged.paged分块输入的 flash attention自注意力mo.mha.ragged.no_cache无缓存 flash attentionmo.kv_cache.store.paged.ragged分页 KV 缓存写入mo.rope.ragged.with_position_id带显式位置 ID 的 RoPEmo.grouped.matmul.raggedMoE 分组矩阵乘mo.mla.graph.decode.paged.fp8FP8 量化的 MLA 解码图内核上层架构代码如 multihead_attention.py 及 max/python/max/pipelines/architectures 下各模型的 attention 层直接调用这些包装函数测试则集中在 max/tests/integration/kv_cache 与 max/tests/integration/nn 目录。输入校验类型、秩、设备的三重防线kernels.py定义了五个通用校验辅助函数几乎每个公共 API 开头都会调用违规时抛出带参数名的ValueError而不是在 Mojo 编译期才失败_check_dtype(expected, **tensors)要求张量 dtype 等于期望值L71_check_rank(expected, **tensors)要求张量秩等于期望值L83_check_same_dtype(**tensors)/_check_same_device(**tensors)要求多个张量共享 dtype / 设备L95、L109_validate_argument_tensor(name, tensor, dtype, rank, device, device_type)更精细的联合校验可同时检查 dtype、秩、具体设备或设备类型如DeviceKind.GPUL2868。常见的硬性约束包括layer_idx必须是uint32且位于 CPUDeviceRef.CPU()上因为它在内核里用于索引input_row_offsets/valid_lengths/position_ids均为uint32scale 常量总是以ops.constant(scale, DType.float32, deviceDeviceRef.CPU())传入。此外还有一处值得注意的工程细节rope_ragged在发现freqs_cis与输入不在同一设备时会显式插入一次设备转移freqs_cis.to(input.device)原因注释说明若把 CPU 上的切片视图直接交给ops.custom图编译器会把mo.slice融合进 GPU 消费者融合视图与隐式传输的生命周期存在竞态可能间歇性越界读L2581-L2591。通用工具函数ceildiv(n, d) - Dim天花板除法返回(n d - 1) // d用于计算量化 scale 的块数L140模块常量KEY_CACHE_INDEX 0、VALUE_CACHE_INDEX 1L67-L68以及_MX_SF_VECTOR_SIZE 32与 Mojo 侧MXFP8_SF_VECTOR_SIZE/MXFP4_SF_VECTOR_SIZE对齐L48-L51_MHA_MASK_VARIANT_TO_ATTENTION_MASK映射表把MHAMaskVariant定义于 attention/mask_config.py映射为注意力内核的mask_str字符串支持CAUSAL_MASK、NULL_MASK、CHUNKED_CAUSAL_MASK、SLIDING_WINDOW_CAUSAL_MASK、SLIDING_WINDOW_NONCAUSAL_MASK五种变体L53-L65。QKV 投影内核族把投影 写缓存熔进一次内核启动分块ragged与填充padded两条路径fused_qkv_ragged_matmulL213是解码/预填充主路径输入为[total_seq_len, hidden_dim]配合input_row_offsets: [batch_size 1]前缀和实现 ragged 语义输出只返回 Q 投影[total_seq_len, n_heads * head_dim]K/V 投影由内核直接写入分页 KV 缓存。可选bias按[q, k, v]拼接与_output_dim覆盖。对应的fused_qkv_padded_matmulL153面向已填充为统一形状[batch, seq_len, hidden_dim]的批量输入多一个valid_lengths: [batch]uint32缓冲K/V 只对有效位置写缓存。两者的共同输入约定可视为本模块的签名规范参数类型/形状约束inputrank 3padded/ rank 2ragged与wqkv同 dtypewqkv[N, hidden_dim]N (n_heads 2*n_kv_heads)*head_dim与input同 dtypelayer_idx标量uint32必须在 CPUinput_row_offsets/valid_lengthsuint32ragged/padded 分别要求 rank 1n_headsint决定输出 Q 的列数量化变体矩阵模块针对不同量化格式提供了专属 QKV 融合内核均以_fused_qkv_ragged_matmul_scaled_*命名Python 侧以下划线开头表示半私有由上层按量化配置选择调用FP8float8_e4m3fn动态缩放L582mo.fused_qkv_matmul.ragged.paged.scale通过input_scale/weight_scale的形状推断 per-tensor(-1,-1,-1)或 per-channel(1,1,-1)粒度也可显式传quant_config来自 max/nn/quant_config.py 的QuantConfig含scales_granularity_mnk输出恒为bfloat16FP4NVFP4 风格L729mo.fused_qkv_matmul.ragged.paged.scale.float4tensor_sf为 buffer 级缩放因子weight_scale_2 * input_scale必须为 float 或 CPU 上的 rank-0 float32 张量MXFP8L851float8_e4m3fn数据 float8_e8m0fnuE8M032 元素 K 块缩放SM100 上 scale 走 rank-5 SF-atom 交错布局CDNA4AMD gfx950走 rank-2[M, K//32]内核还会乘一个恒等 per-tensor scale因此 Python 侧传入tensor_sf 1.0的 CPU 常量MXFP6L966uint8打包每 3 字节 4 个 FP6 码 E8M0 块缩放仅 CDNA4mo.fused_qkv_matmul.ragged.paged.scale.mxfp6.amd非 AMD 直接抛ValueErrorfp6_format支持e2m3/e3m2MiniMax-M3 五路融合L1080、L1226、L1370_fused_qkv_index_ragged_matmul_scaled_mxfp8/mxfp6把[Wq|Wk|Wv|Wiq|Wik]拼成一个 GEMM一次启动同时产出 Q、IndexQ 并原地写入主 KV 缓存与 indexMLA 单潜头缓存输出为二元组(q, index_q)。GGUF 与 GPTQ 量化unfused_qkv_ragged_matmul_gguf_quantizedL1486为 GGUF 检查点提供 Q/K/V 三个独立量化权重矩阵的路径要求quantization_encoding_q/k/v均为 GGUF 编码is_gguf为真且输入为float32权重先经repack_gguf_quantized_weights重排fused_qkv_ragged_matmul_quantizedL1561GPTQ 风格 group-wise 量化group_size参数has_zp_int恒为 0当传入perm_idx激活排序时走GPTQ_gpu_repack_b4_g128_desc_act重排路径否则走GPTQ_gpu_repack_b4_g128最终落入mo.fused_qkv_matmul.ragged.paged[.bias.]quantized。此外rope_split_store_raggedL284把从扁平 QKV 输出中读取、对 Q/K 施加 RoPE、K/V 写入分页缓存、输出已旋转 Q四步熔为一个算子支持interleaved、mrope_section多头 RoPE配合position_ids要求长度等于position_ids.shape[0]并换算成前缀和字符串参数、以及可选的 per-head Q/K/V RMSNorm 融合q_norm_weightk_norm_weightrms_norm_eps此时eps以其整数倒数eps_recip传入因为自定义算子参数不接受 floatfuseFalse时会退化为 splitropestore 三个独立算子用于测试图编译器融合能力_rope_split_store_ragged_unfusedL496。RoPE 内核族标准、显式位置、多段 mRoPErope_raggedL2557标准 ragged RoPEstart_pos取各序列的 cache 长度freqs_cis若窄于 head_dim只旋转每头前若干列默认旋转末尾列rope_firstTrue改为旋转前导列输出 dtype 可用output_dtype覆盖。该函数通过CompositeRopeRaggedOp生成rope_ragged_with_position_idsL2735无缓存耦合、显式position_ids的版本。无mrope_section时走内核快速路径mo.rope.ragged.with_position_id有 mRoPE 段时退化为图实现_freqs_cis_from_position_idsL2645按mrope_section用gather/scatter组装逐 token 频率对 Qwen2.5-VL 等 3D RoPE 模型很有用再交给_apply_rope_with_freqs_cisL2610纯图级复数旋转实现支持 interleaved 与非 interleavedfused_qk_ragged_ropeL1800与 KV 缓存耦合的 QK RoPE直接在缓存中的 K 上施转并原地写回Q 返回新张量cache_dtype取自kv_params.dtype支持position_idsmrope_section同 M3/Qwen2.5-VL 场景fused_qk_padded_ropeL2297padded 输入的对应版本输入 rank 4[batch, seq_len, n_heads, head_dim]用valid_lengths限定施转范围。KV 缓存写入与就地规约KV 缓存存储 API 由PagedCacheValues承载kv_blocksrank 6、cache_lengths、lookup_table、max_prompt_length、max_cache_length、kv_scales、attention_dispatch_metadata等字段定义于 max/nn/kv_cache.py 的KVCacheParams/MHAKVCacheParams/PagedCacheValueskv_cache_store_paged_raggedL2383与kv_cache_store_paged_paddedL2470是底层入口key_or_value取KEY_CACHE_INDEX/VALUE_CACHE_INDEX_validate_kv_cache_store_common会校验kv_blocks为 rank 6、layer_idx为 rank 0 uint32 等便捷封装store_k_cache_ragged/store_v_cache_ragged/store_k_cache_padded/store_v_cache_padded分别绑定 K/V 与 ragged/padded 两种输入形态kv_cache_ragged_raddL5165把张量加到缓存按 batch 切出的切片上用于 speculative decoding 等场景batch_offset指定起始批次内部会先slice_tensor截取input_row_offsets[batch_offset:]rms_norm_key_cache/rms_norm_value_cacheL5210、L5296对缓存中新增条目由input_row_offsets界定做就地 RMSNormper_head_normTrue时 gamma 为[head_dim]逐头归一化per_head_normFalse时 gamma 为[n_kv_heads*head_dim]按 token 跨所有头归一化gamma 尺寸与 head_dim 不一致时必须显式传rms_norm_colsstore_k_scale_cache_raggedL463把量化 K 的 scale 张量写入量化 KV 缓存kv_collection.kv_scalesquantization_granularity作为内核参数。Flash Attention 与 MHA 内核族flash_attention_gpuL3592无缓存的 GPU flash attentionq/k/v 均为 rank 4[batch, seq_len, n_heads, head_dim]传valid_length时走mo.mha.padded.no_cache变体mask_variantlocal_window_size组合出掩码语义滑窗注意力默认local_window_size-1关闭masked_flash_attention_gpuL3678显式 additive mask 版本mask 可为 rank 3[batch, q_seq, kv_seq]跨头广播或 rank 4[batch, n_heads, q_seq, kv_seq]per-head 偏置通过CompositeMaskedFlashAttentionGpuOp生成一等算子flash_attention_raggedL3786自注意力Q/KV 序列等长核心约束是input.dtype kv_params.dtypemask 在内核内物化。可选项丰富sink_weightsper-head 可学习 sink 权重mo.mha.ragged.paged.sink_weights与rel_logits相对位置偏置表[total_q, heads, extent]mo.mha.ragged.paged.rel_logits互斥后者仅支持CAUSAL_MASK要求local_window_size -1与SLIDING_WINDOW_CAUSAL_MASK要求正窗口flash_attention_padded_kv_cacheL2801padded 输入 分页缓存的 MHAvalid_lengths批量尺寸必须等于q的 batch 尺寸flash_attention_ragged_gpuL3958无缓存 ragged 版本max_seq_len必须是 CPU 上的 uint32cross_attention_raggedL5084交叉注意力Q 与 KV 序列长度可不同多出kv_input_row_offsets与q_max_seq_lenCPU uint32两个参数。MLAMulti-head Latent Attention内核族MLA 是 DeepSeek-V3.2、MiniMax-M3 等模型的核心注意力形态KV 缓存只存低秩 latent 表示注意力前再上投影还原。kernels.py提供了从规划、解压、预填充、解码、量化到图捕获友好的完整内核链flare_mla_prefill_planL4353mo.mla.prefill.ragged.plan用buffer_size把变长序列切成 chunk返回(buffer_row_offsets, cache_offsets, buffer_lengths)三张规划张量max_chunks默认 16flare_mla_decompress_k_cacheL5023把 latent K 从分页缓存拷入连续缓冲并用weight上投影还原为完整 Kk k_latent weight.T返回[buffer_size, weight.shape[0]]flare_mla_prefill_raggedL4256MLA 预填充需先解压 K/V避免 OOM总 cache 长于缓冲时按 chunk 迭代返回输出张量与 softmax 信息张量供下一迭代续算qk_rope_dim默认 64决定输出末维head_dim - qk_rope_dimflare_mla_decode_ragged/flare_mla_decode_ragged_scaledL4067、L4154MLA 解码。scaled 版本接收显式 per-token KV scaleskv_scales: [num_blocks,1,1,page_size,1,1]float32与 Q scalesper_token_scale_rope_awareTrue时使用 FP8BF16 交错布局输出末维为head_dim - 2*qk_rope_dimquantization_granularity默认 640rope-aware 下等于 KV cache head_dimmla_prefill_graphL4489、mla_decode_graphL4687、mla_prefill_decode_graphL4857手工融合的图内核把 RoPE含 cache 内就地、KV latent 的 RMSNorm、latent→KV 上投影、FP8 量化、MLA 注意力全部塞进一次启动mo.mla.graph.prefill.paged[.fp8]等。mla_decode_graph的 decode 分支支持 FP8 batched matmul 投影w_uk/w_uv与稀疏解码sparse_indices/sparse_topk_lengths/sparse_attn_sink均需同时提供sparse_indices_stride必填mla_prefill_decode_graph按批次最大序列长度在 prefill/decode 间切换。两者都接收scalar_args由compute_mla_dispatch_args_scalar产生int64[3]用于 CUDA graph 捕获与num_partitions_scalar由compute_mha_decode_num_partitions产生运行期动态计算分区数FP8 scale 粒度由_fp8_mla_scale_paramsL4463从quant_config推导scale_granularity_override用于 per-head 行数跨越磁盘 block 时覆盖如 64 vs 128。稀疏注意力MSA与 MLA 索引内核针对长上下文稀疏注意力模块提供两块能力mla_fp8_index_top_kL2898mo.mla.indexer.ragged.float8.paged对 FP8 Q 与量化 K 缓存做带 scale 的 FP8 打分、跨头聚合、掩码返回每 token 的 top-k 键索引int32无效位填 -1kpool 1时缓存每 k 个 token 只存一个池化键返回的是 pool idmsa_sparse_indexerL3020MiniMax-M3 的块稀疏索引器按 128-token 块对 index-K 缓存打分返回每 (index head, query token) 的 topk 块 idscore_scratch是持久 FP32 缓冲[num_index_heads, max_rows, MAX_NUM_BLOCKS]decode 在 graph-capture 区内不能临时分配缓冲过窄会使多 token 路由回退到 prefill 索引器init_blocks/local_blocks分别强制保留前导与尾部块msa_sparse_attention_raggedL3219及其 MXFP8/MXFP6 变体L3299、L3397按块 id gather 稀疏 KV 带后执行块稀疏 MHAhead_dim128sparse_block_size必须等于缓存页大小与内核BN。MXFP8/MXFP6 变体直接输出 o_proj 就绪的量化激活与 E8M0 块缩放要求n_heads*head_dim是 32 的倍数与quantize_dynamic_block_scaled的输出位级一致省去单独量化分发。MoE 内核族路由、重排与分组矩阵乘moe_create_indicesL5361根据 router 的topk_ids生成五个张量——按专家重排的 token 顺序、各专家起始下标、恢复原序的下标、活跃专家 id、专家使用统计[max_tokens_per_expert, num_active_experts]moe_router_group_limitedL5444DeepSeek-V3 风格 group-limited router。n_groups 1时走专用单组路径mo.moe.single.group.router此时topk_group被忽略、权重归一化恒开启n_groups 1时先选组再选专家mo.moe.router.group.limitedrouted_scaling_factor以 CPU float32 常量传入moe_sink_gate_routerL5544Inkling gate 公式的融合 sigmoid 门控 常开 sink 专家路由mo.moe.sink.gate.router。参数约束严格n_routed_experts必须是 warp 宽度AMD 64 / NVIDIA 32的正倍数且 ≤1024n_experts_per_tok n_shared_experts必须是不超过 warp 宽度的 2 的幂top-k 幸存者还需装进一个 warp——代码中给出了如 256 专家在 32 warp 下每 token 最多 10 个、512 专家最多 8 个的示例边界moe_eplb_remapL5803EPLB 逻辑→物理专家 id 重映射单内核完成共享内存缓存 logcnt/log2phy 切片hash_decorrelateTrue时用 Knuth 乘法哈希打散位置对齐moe_router_single_group_eplbL5907进一步把单组 router 与 EPLB remap 融合为一次启动grouped_matmul_raggedL5999BF16 分组矩阵乘mo.grouped.matmul.raggedweight为 rank 3[num_experts, N, K]expert_usage_stats需要与计算设备同侧host 常驻时自动转量化分组 GEMMgrouped_dynamic_scaled_fp8_matmulL6881FP8 InputScaleSpec/WeightScaleSpec块粒度支持 (1,128)grouped_dynamic_block_scaled_matmul_amdL6077MXFP4/MXFP8E8M0 缩放lane_bytes区分 16/32支持preshuffled_b与 A-scale 预打乱/融合grouped_dynamic_scaled_mxfp6_matmulL6313CDNA4-onlygrouped_matmul_block_scaledL6486支持 NVFP4/MXFP4/MXFP8/W4A8 四种组合以一等复合算子下发SF_VECTOR_SIZE由 scale dtype 推断 16/32grouped_matmul_blocked_swigluL6698SM100 的 GEMMSwiGLU 融合要求权重按sigma(2i)i, sigma(2i1)Di在 N 轴预排列。量化与 MX 格式工具内核FP8 量化quantize_static_scaled_float8/quantize_tensor_dynamic_scaled_float8/quantize_dynamic_scaled_float8L7307 起与matmul_static_scaled_float8L8879、dynamic_scaled_matmulL7514块缩放block-scale / MX 格式quantize_dynamic_block_scaledL8409输出 E8M0 缩放_MX_SF_VECTOR_SIZE32、grouped_quantize_dynamic_block_scaledL8540、quantize_dynamic_block_scaled_mxfp4L8637、quantize_dynamic_block_scaled_mxfp6L8171打包 uint8 FP6_FORMAT参数、mxfp4_dequant/mxfp6_dequantL8237、L8303、dynamic_block_scaled_matmul/dynamic_block_scaled_matmul_amd/dynamic_block_scaled_matmul_mxfp6L7635、L7978、L8079布局工具block_scales_interleaveL8699SM100 的 rank-5 SF-atom 交错、block_scaled_preshuffle_grouped_scale_4dL8769、block_scaled_preshuffle_b_5dL8840、needs_fp8_fnuz_conversion/normalize_e4m3fn_to_e4m3fnuz/convert_weights_to_fp8_fnuz_if_neededL8968 起兼容 e4m3fnuz 硬件平台探测_is_sm10x_gpu/_is_sm12x_gpu/_is_apple_gpu/_is_amd_gpuL8377 起。采样、LoRA/SGMV 与其他杂项内核模块还承载了一批与注意力/KV 无直接关系但同属图内核封装的算子可单独取用采样topk_fused_sampling/topk_fused_sampling_with_distL9587、L9743、topk_topp_masked_probsL9825、gumbel_argmax_from_probsL9879、apply_penalties_to_logitsL9283、update_frequency_dataL9390、scatter_set_constantL9425、apply_packed_bitmaskL9465、scatter_nd_skip_oob_indicesL9540LoRA/SGMVsgmv_kernelL9921、sgmv_lora_kernelL9982、sgmv_lora_qkv_shrinkL10055、sgmv_qkv_lora_kernel/sgmv_qkv_lora_fusedL10199、L10295序列工具merge_ragged_tensorsL9054、spatial_mergeL10418、learnable_2d_interp_pos_embL10481、sliced_addL10538、tpool_patch_mergerL10822宿主协作inplace_memcpyL10607、launch_host_funcL10656、wait_host_value/wait_host_value_with_depL10689、L10735、sleepL10788统计与规范row_mean_of_squaresL10913、row_mean_of_squares_qkL10966、apply_qk_rms_normL11022、mtp_eh_normL9135、eagle_prefill_shift_tokensL9253。使用方式与验证max.nn.kernels的典型消费方是模型架构层——例如 max/python/max/pipelines/architectures 下各模型的layers/attention.py、layers/moe.py以及 multihead_attention.py 中基于MHAMaskVariant/KVCacheParams组织 attention 前向采样算子被 sampling.py 等消费。测试与基准方面单元测试test_kernels.py集成测试注意力与 KV 缓存集中在 max/tests/integration/kv_cache如attention/test_mla_gpu.py、attention/test_ragged_attention_gpu.py、test_fused_qk_rms_norm_rope_fusion_gpu.py、test_kv_cache_store_gpu.py量化 GEMM 与内核杂项在 max/tests/integration/nn如test_fused_qkv_mxfp8_matmul_gpu.py、test_moe_gpu.py、test_sgmv_kernels_gpu.py、test_smallm_streaming_matmul_gpu.py内核基准max/kernels/benchmarks/misc/comparison如bench_blackwell_mla_decode.py、bench_amd_mla.py对比不同内核路径。小结max.nn.kernels是 MAX 平台 Python 推理栈的内核契约层一方面用ops.inplace_custom/ops.custom/Composite*Op三类机制把 Mojo 内核稳定地暴露给图构建器另一方面通过统一的参数签名uint32的层索引与偏移量、CPU 常量的 scale/epsilon、ragged 的前缀和表示与逐参数的运行时校验把绝大多数形状/dtype/设备错误提前到 Python 侧抛出。理解这层封装是阅读 MAX 上任意模型架构代码、或为自定义注意力/量化方案编写图内核的起点。【免费下载链接】mojoThe Modular Platform (includes MAX Mojo)项目地址: https://gitcode.com/GitHub_Trending/mo/mojo创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考