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).

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 = (StringSpan("sm_100") in String(_accelerator_arch())) or (StringSpan("sm_103") in String(_accelerator_arch()))

is_sm90

comptime is_sm90 = (StringSpan("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[StringSpan("MHA_PDL"), True]() else PDLLevel.OFF

Structs

  • DynamicInt: A runtime integer value that satisfies OptionallyStaticInt.
  • 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.
  • NonNullPointer: A pointer with a compile-time guarantee of being non-null.
  • NoPartition: A single-partition (non-split-K) scheme for MHA decoding.
  • NullPointer: A pointer known at compile time to be null, used when an optional pointer argument is absent.
  • SplitKPartition: A multi-partition split-K scheme for MHA decoding over long sequences.
  • StaticInt: A compile-time constant integer that satisfies OptionallyStaticInt.

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.
  • OptionalPointer: Abstracts over nullable pointers, providing a uniform interface for NonNullPointer and NullPointer.

Functions