Skip to content

cuda: BATCH_INVARIANT: occupancy-independent FA KV split, PTQ1_0 mat-vec up to 8 columns - #322

Open
professorpalmer wants to merge 1 commit into
PrismML-Eng:prismfrom
professorpalmer:claude/batch-invariant-fa
Open

professorpalmer wants to merge 1 commit into
PrismML-Eng:prismfrom
professorpalmer:claude/batch-invariant-fa

Conversation

@professorpalmer

Copy link
Copy Markdown

One of the five pieces of #285, split per bri-prism's request on #317. Single commit on current prism (eaecb50c7), applies without conflicts. Only active under GGML_CUDA_BATCH_INVARIANT=1.

What it does

  • The flash-attention vector path's KV split no longer depends on the occupancy of the kernel instance that runs, so the same query produces the same split whether it is decoded alone or inside a verify batch.
  • PTQ1_0 batches of 5-8 columns stay on the PT mat-vec path instead of crossing to MMQ: a 5-column verify on MMQ did not match a single decode. GGML_CUDA_PTQ1_MMVQ_MAX overrides the crossover; the default of 4 is confirmed by measurement (pp5: 172 tok/s on MMQ vs 162 on the mat-vec, so 4 is where MMQ starts to win and where the verify columns of a draft of 4 still match).

Receipts

…vec up to 8 columns

Under GGML_CUDA_BATCH_INVARIANT:
- the non-stream-k flash-attention KV split is sized from a fixed blocks-per-SM value instead of the
  occupancy of the template instance that runs; the 1-query and multi-query instances differ in registers
  and shared memory, so their splits (and combine order) differed.
- PTQ1_0 batches up to MMVQ_MAX_BATCH_SIZE stay on the PT mat-vec, whose per-column arithmetic does not
  depend on the column count; a 5-column speculative verify on MMQ did not match the same tokens decoded
  alone. Measured cost: pp5 172 -> 162 tok/s, only on 5-8 column batches.
GGML_CUDA_PTQ1_MMVQ_MAX overrides the mat-vec / MMQ crossover (default 4, confirmed: MMQ wins from 5).

Not covered: the stream-k split of the MMA kernel follows the padded KV length (and the tile instance
follows the query count), so attention in a verify batch is not bit-identical to single-token decode past
~32k, or at any depth on the MMA decode route. Measured at 4k/20k/40k: the weight path matches, a few
continuations diverge at the rounding level.

Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
(cherry picked from commit dc0cd6c)
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant