IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).

Python module

max.experimental.nn.common_layers.functional_kernels

Functional wrappers for MAX kernel operations used in attention layers.

Functions​

flash_attention_gpuComputes flash attention using a GPU-optimized kernel.
flash_attention_raggedComputes flash (self) attention provided the !mo.opaque KV Cache.
flash_attention_ragged_gpuComputes flash attention for ragged inputs using a GPU-optimized kernel, without a KV cache.
fused_siluPerforms the fused SILU operation for all the MLPs in the EP MoE module.
grouped_matmul_raggedPerforms the grouped matmul used in the MoE layer.
moe_create_indicesCreates indices for the MoE layer.
moe_router_group_limitedRoutes tokens with the group-limited MoE router.
rms_norm_key_cacheApplies RMSNorm to the new entries in the KV cache.
rope_split_store_raggedApplies RoPE to Q and K from a flat QKV buffer and stores K/V to the cache.
stack_device_shardsReassembles a per-device weight-shard bundle into one Sharded tensor.