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 function
mla_decode_sm100_dispatch
def mla_decode_sm100_dispatch[q_type: DType, k_t: MHAOperand, output_type: DType, mask_t: MHAMask, config: MHAConfig[config.dtype], depth: Int, num_heads: Int, group: Int = Int(1), *, ragged: Bool = False, _is_cache_length_accurate: Bool = False, decoding_warp_split_k: Bool = False, per_token_scale_rope_aware: Bool = False, sparse: Bool = False, rope_aware_kv_sparse: Bool = False, fold_shared_index: Bool = False, has_extra_k: Bool = False, AttnSinkPtrType: OptionalPointer = NullPointer[.float32], TopkLengthsPtrType: OptionalPointer = NullPointer[.int32]](q: TileTensor[q_type, Engine=q.Engine, linear_idx_type=q.linear_idx_type], k: k_t, output: TileTensor[output_type, Engine=output.Engine, linear_idx_type=output.linear_idx_type], scale: Float32, valid_length: TileTensor[.uint32, Engine=valid_length.Engine, linear_idx_type=valid_length.linear_idx_type], mask: mask_t, scalar_args_buf: TileTensor[.int64, Engine=scalar_args_buf.Engine, linear_idx_type=scalar_args_buf.linear_idx_type], batch_size: Int, q_max_seq_len: Int, max_cache_valid_length: Int, ctx: DeviceContext, q_scale_ptr: OptionalReg[Pointer[Float32, MutAnyOrigin]] = None, d_indices: OptionalReg[Pointer[Int32, MutAnyOrigin]] = None, indices_stride: Int = Int(0), topk_lengths: TopkLengthsPtrType = null_pointer[TopkLengthsPtrType](), attn_sink_ptr: AttnSinkPtrType = null_pointer[AttnSinkPtrType](), extra_k: OptionalReg[k_t] = None, extra_d_indices: OptionalReg[Pointer[Int32, MutAnyOrigin]] = None, extra_indices_stride: Int = Int(0), extra_topk_lengths: TopkLengthsPtrType = unread_pointer[TopkLengthsPtrType](), extra_scales_ptr: OptionalReg[Pointer[Float32, MutAnyOrigin]] = None, num_partitions_in: Optional[Int] = None, logical_indices: OptionalReg[Pointer[Int32, MutAnyOrigin]] = None)