runtime: a cublaslt bf16 gemm launch may pin its algorithm - #25
Conversation
|
Thanks for this, and for the measurements. The sliced-kernel finding on sm_89 (algo 30/31/16 report We'd like to land the capability in a different shape, though. The runtime should not measure or decide at run time. It should only execute what the manifest says. Planning inside the runtime breaks things in practice:
We don't need to switch algorithms online, so the proposal is to pin the algorithm in the manifest:
How the |
Signed-off-by: Feathbow <feathbow@gmail.com>
8fd6ddf to
37fed33
Compare
|
Thanks, reworked along these lines. The runtime now only executes what the manifest pins: |
Signed-off-by: Feathbow <feathbow@gmail.com>
What
Reworked as the review asked: the runtime executes the algorithm the manifest names and decides nothing at run time.
extern:cublaslt_bf16_tnandextern:cublaslt_bf16_tn_acctake an optionalalgo, everycublasLtMatmulAlgoConfigAttributes_tvalue of one algorithm:id,tile,stages,split_k,reduction,swizzle,custom,inner_shape,cluster_shape. Withoutalgonothing changes.cublasLtMatmulAlgoInitand every attribute set, thencublasLtMatmulAlgoCheckfor each call of the launch at the shapes its args take with every var at its lower bound and at its upper one, thewhenvar held to the launch's range, with the workspace withinBlas's 32 MiB. An algorithm that cannot be built or does not fit is a kernel artifact error, never a fallback.cublasLtMatmulon the pinned algorithm withBlas's workspace: no heuristic, no sync, no readback, so eager multi-rank issue,kern benchcaptures andwhen-skipped launches need nothing from it.whenrange, each with its ownalgo.algoon any other extern.docs/runtime.mdrecords how a caller chooses one and the sm_89 finding thatSPLITK_NUM1 does not make two algorithms agree bitwise (the sliced kernels, algo 30, 31 and 16). Schema regenerated.Verification
tests/gemm_algo.rs(GPU, ignored likefp8_gemm.rs): an op with twowhenranges ofrows, each pinned to cuBLASLt's first heuristic answer at the end of its range, lands the default op's exact product (a permutation times small integers) at 100 and 256 rows, eagerly and through a captured graph; a tile the shape cannot run fails the load as a kernel artifact error.gemm_algoandfp8_gemmpass on Single GPU (sm_89, x86_64) and Single GPU (GH200, aarch64), both CUDA 13.1.cargo fmt --check,cargo clippy --all-targets -- -D warnings,cargo test -p kern-manifest -p kern-pool -p kern-runtime -p kern-testand the schema golden pass.