Skip to content

runtime: bf16 gemm runs a measured cuBLASLt pick per shape, not the heuristic's first answer #26

Description

@xiaguan

Summary

extern:cublaslt_bf16_tn runs cuBLASLt's first heuristic answer. That answer is not the fastest one, and which algorithm is fastest depends on the GPU generation and the library version. #25 fixes the choice to one tile (256x128, no split-K), measured on sm_89. On GB300 that same policy makes a step 4.3% slower. A measured per-shape pick is 2.6% faster than the heuristic.

GEMM is a basic capability of the runtime, so a measured choice is worth having. It should be a tuned table outside the manifest, not a tile name inside it.

Measurement (GB300, sm_103, cuBLASLt 13.7.0, 1400 W cap)

Shapes: HiDream-O1 / Qwen3-VL-8B text tower at m = 4097: qkv 6144x4096, o 4096x4096, gate_up 24576x4096, down 4096x12288. Random bf16 inputs, same column-major mapping as gemm_bf16_tn. #25's selection logic is copied exactly.

  • The heuristic's candidates. It returns 8 at every shape, and all 8 are non-split-K cluster kernels (algo id 66). Top-1 is already a big tile: 256x240 with a 2x2 cluster for qkv/o/down. runtime: a cublaslt bf16 gemm launch may pin its algorithm #25's 256x128 is offered at qkv/o/gate_up and borrowed from k = n for down. On this arch it is the narrower tile.
  • Sustained load, 4 s per candidate. Every candidate sits at the power cap (~1365 W, throttle reason = SW power cap). Power is flat, so time tracks J/GEMM, which confirms runtime: a cublaslt bf16 gemm launch may pin its algorithm #25's mechanism. But 256x128 is 4–6% worse than top-1 on qkv/o, in both time and energy.
  • Cold timing picks the same winners as sustained load at all four shapes (qkv 128x192, o 256x224, gate_up 256x208, down 256x224). The sm_89 observation that isolated timing misranks does not reproduce here.
  • Whole-step run. Each step is 36 layers x 4 GEMMs. The three policies alternate, 4 s x 6 rounds, and the round-to-round spread is < 0.5%.
policy ms/step vs top-1
heuristic top-1 35.12 —
#25 wide (256x128) 36.62 +4.3%
best per shape (from top-8, sustained) 34.22 −2.6%

Constraints

  • The algorithm cannot live in the manifest. A cuBLASLt algorithm moves with the library version (the reason extern:cublaslt_bf16_tn_tile was removed in cd39a57), and the best one moves with the hardware (above). Static shapes in the manifest do not change this. n and k are already literals, and m is fixed per captured (program, vars). What is missing is a place to record the measured choice, not the shape.
  • The choice must be deterministic. Candidates are often within noise of each other (qkv: 0.129 / 0.130 / 0.130 ms), and different tiles are not guaranteed to produce the same bits. Timing at every startup would break "same input, same bytes out". The choice should be tuned once and pinned.
  • Timing cannot run inside graph capture, and the first call to a graph program happens during capture. Tuning belongs offline (or in an explicit warmup), not in the first call.

Proposal

Follow the registry's [[pick]] model: a measured choice keyed by shape, with a report behind it.

  • kern bench (or kern tune) writes the table. It takes the workload's m values (the vars the service actually captures) x every GEMM call site's (n, k, ldc, A row stride) in the manifest, times cuBLASLt's top-N at each shape, and writes the winners.
  • The table's key and value. Key: GPU name + cuBLASLt version + shape. Value: the opaque cublasLtMatmulAlgo_t, which is stable within one library version. The table ships with the deployment, not the manifest.
  • The runtime only looks up. A hit runs the tabled algorithm; a miss runs heuristic top-1. Both are deterministic, and the manifest keeps saying extern:cublaslt_bf16_tn.
  • runtime: a cublaslt bf16 gemm launch may pin its algorithm #25's _wide entry is then not needed. sm_89 and GB300 each get their own measured winner.

Open question: small m

Everything above is m = 4097, where the heuristic is already within 2.6% of the best. Decode-shaped GEMMs (m = 1..256) are where cuBLASLt heuristics are usually weakest, and they have not been measured yet. A sweep there decides how much this is worth before it goes into the roadmap.

Search space

cublasLtMatmulAlgoGetHeuristic returns only 8 candidates on GB300 even when asked for 64. cublasLtMatmulAlgoGetIds lists 20 algo ids for bf16 TN, and id 66 alone reports 645 tile ids via CUBLASLT_ALGO_CAP_TILE_IDS. Configurations beyond top-8 can be built with cublasLtMatmulAlgoConfigSetAttribute and validated with cublasLtMatmulAlgoCheck, if top-8 turns out not to be enough.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions