For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /max/get-started.md).
Mojo module
mha_utils
Shared configuration types and dispatch helpers for MHA GPU kernels.
Provides MHAConfig (tile/pipeline configuration), FlashAttentionAlgorithm
(algorithm variant selector), mask-dispatch helpers, and partition-scheme
types used by both prefill and decode attention kernels.
comptime valuesβ
callback_fn_typeβ
comptime callback_fn_type = def[mask_t: MHAMask](mask: mask_t) raises capturing thin -> None
is_sm100β
comptime is_sm100 = (StringSlice("sm_100") in String(_accelerator_arch())) or (StringSlice("sm_103") in String(_accelerator_arch()))
is_sm90β
comptime is_sm90 = (StringSlice("sm_90") in String(_accelerator_arch()))
is_sm90or100β
comptime is_sm90or100 = is_sm90 or is_sm100
MHA_PDL_LEVELβ
comptime MHA_PDL_LEVEL = PDLLevel.OVERLAP_AT_END if get_defined_bool[StringSlice("MHA_PDL"), True]() else PDLLevel.OFF
Structsβ
- β
DynamicInt: A runtime integer value that satisfiesOptionallyStaticInt. - β
FlashAttentionAlgorithm: Identifies which flash-attention algorithm variant to use for a kernel launch. - β
MHAConfig: Compile-time and runtime tile-shape configuration for MHA GPU kernels. - β
NoPartition: A single-partition (non-split-K) scheme for MHA decoding. - β
SplitKPartition: A multi-partition split-K scheme for MHA decoding over long sequences. - β
StaticInt: A compile-time constant integer that satisfiesOptionallyStaticInt.
Traitsβ
- β
MHAPartitionScheme: Trait describing how the key-value sequence is partitioned for split-K decoding. - β
OptionallyStaticInt: Trait for integer values that may be statically known at compile time.
Functionsβ
- β
as_dynamic_row_major_1d: Reinterprets a generic-addressLayoutTensoras a 1-D dynamic row-major tensor. - β
dispatch_mask: Instantiate anMHAMaskby name and invoke a callback with it. - β
dispatch_materialized_mask: Wrap a dense mask tensor in aMaterializedMaskand invoke a callback. - β
get_start_and_end_for_partitions: Calculate start and end indices for a partition.
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!