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
enqueue_apple_grouped_gemv
def enqueue_apple_grouped_gemv[*, rows_per_sg: Int = Int(2), tokens_per_pass: Int = Int(1), num_sg: Int = Int(8), acc_width: Int = Int(8)](c: TileTensor[Engine=c.Engine, address_space=c.address_space, linear_idx_type=c.linear_idx_type], a: TileTensor[Engine=a.Engine, address_space=a.address_space, linear_idx_type=a.linear_idx_type], b: TileTensor[Engine=b.Engine, address_space=b.address_space, linear_idx_type=b.linear_idx_type], a_offsets: TileTensor[.uint32, Engine=a_offsets.Engine, address_space=a_offsets.address_space, linear_idx_type=a_offsets.linear_idx_type], expert_ids: TileTensor[.int32, Engine=expert_ids.Engine, address_space=expert_ids.address_space, linear_idx_type=expert_ids.linear_idx_type], num_active_experts: Int, ctx: DeviceContext)
Enqueues the grouped GEMV.
b is the weight stack [num_experts, N, K] with static N and K;
a is [total_M, K] and c is [total_M, N], both token-major. a
and b are both bf16 or both fp16; c is fp16, bf16 or fp32, and
accumulation is fp32. Correct for any group size, but it re-reads the
weight every tokens_per_pass rows, so the dispatch only uses it for
small groups.
Parameters:
- βrows_per_sg (
Int): Weight rows (output columns) per simdgroup. - βtokens_per_pass (
Int): Tokens of a group accumulated per pass over the weight. - βnum_sg (
Int): Simdgroups per threadgroup. - βacc_width (
Int): Fp32 lanes kept per accumulator (1 to 8); narrower trades adds for registers whenrows_per_sg * tokens_per_passis large.