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 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 when rows_per_sg * tokens_per_pass is large.

Was this page helpful?