diff --git a/crates/kern-manifest/src/types.rs b/crates/kern-manifest/src/types.rs index 579bf98..2b7d776 100644 --- a/crates/kern-manifest/src/types.rs +++ b/crates/kern-manifest/src/types.rs @@ -944,11 +944,38 @@ pub struct ExternLaunch { /// Run this launch only while a var is in range, e.g. `{"var": "tokens", "max": 16}`: an op picks its implementation by shape with one launch per range. Outside it the launch does nothing. #[serde(default, skip_serializing_if = "Option::is_none")] pub when: Option, + /// For `cublaslt_bf16_tn` and `cublaslt_bf16_tn_acc`: the cuBLASLt algorithm to run instead of the heuristic's first answer. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub algo: Option, /// Where each launch param comes from (default: the op's params in order). #[serde(default, skip_serializing_if = "Option::is_none")] args: Option>, } +/// One cuBLASLt algorithm as its `cublasLtMatmulAlgoConfigAttributes_t` values, e.g. `{"id": 6, "tile": 24, "stages": 9, "split_k": 1, "reduction": 0, "swizzle": 0, "custom": 0, "inner_shape": 0, "cluster_shape": 0}`. Built at load and checked against every shape the launch runs at; one that does not fit is an error, never a fallback. +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)] +#[serde(deny_unknown_fields)] +pub struct GemmAlgo { + /// `CUBLASLT_ALGO_CONFIG_ID`. + pub id: i32, + /// `CUBLASLT_ALGO_CONFIG_TILE_ID`. + pub tile: u32, + /// `CUBLASLT_ALGO_CONFIG_STAGES_ID`. + pub stages: u32, + /// `CUBLASLT_ALGO_CONFIG_SPLITK_NUM`. + pub split_k: i32, + /// `CUBLASLT_ALGO_CONFIG_REDUCTION_SCHEME`. + pub reduction: u32, + /// `CUBLASLT_ALGO_CONFIG_CTA_SWIZZLING`. + pub swizzle: u32, + /// `CUBLASLT_ALGO_CONFIG_CUSTOM_OPTION`. + pub custom: u32, + /// `CUBLASLT_ALGO_CONFIG_INNER_SHAPE_ID`. + pub inner_shape: u16, + /// `CUBLASLT_ALGO_CONFIG_CLUSTER_SHAPE_ID`. + pub cluster_shape: u16, +} + /// The inclusive range of a var a launch runs in, e.g. `{"var": "tokens", "min": 17}`; an end left out is open. #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, schemars::JsonSchema)] #[serde(deny_unknown_fields)] diff --git a/crates/kern-manifest/src/verify.rs b/crates/kern-manifest/src/verify.rs index 11d03e7..95780aa 100644 --- a/crates/kern-manifest/src/verify.rs +++ b/crates/kern-manifest/src/verify.rs @@ -478,6 +478,11 @@ fn diagnostics(m: &Manifest) -> Vec { "{ctx}: a launch without a module must be a runtime built-in (`extern:`)" )); } + if e.algo.is_some() + && !matches!(e.entry.as_str(), "extern:cublaslt_bf16_tn" | "extern:cublaslt_bf16_tn_acc") + { + errs.push(format!("{ctx}: `algo` pins a cublaslt_bf16_tn[_acc] launch, not `{}`", e.entry)); + } } Launch::Kernel(k) => { if k.entry.is_empty() { diff --git a/crates/kern-runtime/src/compile.rs b/crates/kern-runtime/src/compile.rs index 6283c42..b6f2eee 100644 --- a/crates/kern-runtime/src/compile.rs +++ b/crates/kern-runtime/src/compile.rs @@ -10,13 +10,15 @@ use std::ffi::CString; use std::path::{Path, PathBuf}; use std::sync::Arc; +use cudarc::cublaslt::CudaBlasLT; use cudarc::driver::{result as cu, sys, CudaStream}; use kern_manifest::types::{ - Arg, Call, Dim, Expr, FieldSrc, LaunchArg, Manifest, Op, Pack, ParamType, TensorMap, TmaDType, Var, + Arg, Call, Dim, Expr, FieldSrc, LaunchArg, Manifest, Op, Pack, ParamType, TensorMap, TmaDType, Var, When, }; use std::os::raw::c_void; use crate::cubin::{param_sizes, LoadedModule, MulticastScan}; +use crate::cublas::{bf16_operands, Pinned}; use crate::device::{alloc, DeviceBuf}; use crate::error::{bail, cuda_check, Error, Result}; use crate::nccl::{Coll, Elem}; @@ -159,8 +161,9 @@ pub(crate) enum LaunchKind { pdl: bool, }, /// `extern:cublaslt_bf16_tn` / `..._acc` (beta 0.0 / 1.0); 6 args, or - /// 7 with C's row stride, or 8 with C's and A's row strides. - Gemm { beta: f32 }, + /// 7 with C's row stride, or 8 with C's and A's row strides. On the + /// launch's pinned algorithm when it names one. + Gemm { beta: f32, algo: Option }, /// `extern:cublas_bf16_tn_f32`: same operands, f32 result (cublasGemmEx). GemmF32, /// `extern:cublaslt_fp8_tn` / `..._f32`: e4m3 operands with per-tensor @@ -226,6 +229,7 @@ enum LaunchImpl { }, GemmBf16Tn { beta: f32, + algo: Option, }, /// cublasGemmEx with an f32 result (`extern:cublas_bf16_tn_f32`). GemmBf16TnF32, @@ -295,6 +299,7 @@ pub(crate) fn resolve_ops( modules: &[LoadedModule], kernels_dir: Option<&Path>, stream: &Arc, + blt: &CudaBlasLT, vars_max: &BTreeMap, ) -> Result> { let mut ops = BTreeMap::new(); @@ -305,8 +310,13 @@ pub(crate) fn resolve_ops( kern_manifest::types::Launch::Extern(e) => { let ext = e.entry.strip_prefix("extern:").unwrap_or(&e.entry); match ext { - "cublaslt_bf16_tn" => launches.push(LaunchImpl::GemmBf16Tn { beta: 0.0 }), - "cublaslt_bf16_tn_acc" => launches.push(LaunchImpl::GemmBf16Tn { beta: 1.0 }), + "cublaslt_bf16_tn" | "cublaslt_bf16_tn_acc" => { + let algo = e.algo.as_ref().map(|a| Pinned::new(blt, a)).transpose().map_err(|err| { + Error::Call { context: format!("op `{name}` launch #{li}"), source: Box::new(err) } + })?; + let beta = if ext == "cublaslt_bf16_tn" { 0.0 } else { 1.0 }; + launches.push(LaunchImpl::GemmBf16Tn { beta, algo }); + } "cublas_bf16_tn_f32" => launches.push(LaunchImpl::GemmBf16TnF32), "cublaslt_fp8_tn" => launches.push(LaunchImpl::GemmFp8Tn { f32_out: false }), "cublaslt_fp8_tn_f32" => launches.push(LaunchImpl::GemmFp8Tn { f32_out: true }), @@ -412,6 +422,40 @@ struct Ranks<'a> { peer_buffers: &'a BTreeSet, } +/// What checking a pinned GEMM needs: the Lt handle and every var's upper bound. +struct PinnedCheck<'a> { + blt: &'a CudaBlasLT, + maxima: &'a [u64], +} + +impl PinnedCheck<'_> { + /// `cublasLtMatmulAlgoCheck` the launch's algorithm at the shapes its + /// args take with every var at its lower bound and at its upper one, the + /// `when` var held to the launch's range. + fn check(&self, algo: &Pinned, slots: &[Slot], when: Option<&When>, vars: &BTreeMap<&str, usize>) -> Result<()> { + let range = when.map(|w| Ok::<_, Error>((var_index(vars, &w.var)?, w))).transpose()?; + for high in [false, true] { + let mut v = if high { self.maxima.to_vec() } else { vec![Var::MIN; self.maxima.len()] }; + if let Some((i, w)) = range { + v[i] = if high { v[i].min(w.max.unwrap_or(u64::MAX)) } else { v[i].max(w.min.unwrap_or(0)) }; + } + let at = Dense(v); + let vals = slots + .iter() + .map(|s| match s { + Slot::Const(rv) => Ok(*rv), + Slot::Expr(e) => Ok(RVal { val: e.eval(&at)?, bytes: 0 }), + Slot::Pack(_) => bail!(Manifest, "an extern gemm takes no pack"), + }) + .collect::>>()?; + if let Some((_, _, _, shape)) = bf16_operands(&vals)? { + algo.check(self.blt, shape)?; + } + } + Ok(()) + } +} + /// Lower every program's call list into a flat launch list. Every launch /// that receives a peer buffer is SASS-scanned for multicast TMA first. fn compile_programs( @@ -420,8 +464,11 @@ fn compile_programs( buffers: &BTreeMap, states: &BTreeMap, ranks: &Ranks, + blt: &CudaBlasLT, ) -> Result> { let vars: BTreeMap<&str, usize> = manifest.vars.keys().enumerate().map(|(i, s)| (s.as_str(), i)).collect(); + let maxima: Vec = manifest.vars.values().map(|v| v.max).collect(); + let pinned = PinnedCheck { blt, maxima: &maxima }; let mut scan = MulticastScan::new(); let mut programs = BTreeMap::new(); for (pname, p) in &manifest.programs { @@ -434,7 +481,7 @@ fn compile_programs( bail!(Manifest, "program `{pname}` {cctx}: unknown op"); }; let lo = launches.len(); - compile_call(c, op, rop, &cctx, buffers, states, &vars, ranks, &mut scan, &mut launches) + compile_call(c, op, rop, &cctx, buffers, states, &vars, ranks, &pinned, &mut scan, &mut launches) .map_err(|e| Error::Call { context: format!("program `{pname}` {cctx}"), source: Box::new(e) })?; call_ranges.push((lo, launches.len())); } @@ -454,9 +501,10 @@ pub(crate) fn compile( states: &BTreeMap, ranks: &BTreeMap, peers: &BTreeMap, + blt: &CudaBlasLT, ) -> Result> { let peer_buffers: BTreeSet = peers.keys().cloned().collect(); - compile_programs(manifest, ops, buffers, states, &Ranks { ranks, peer_buffers: &peer_buffers }) + compile_programs(manifest, ops, buffers, states, &Ranks { ranks, peer_buffers: &peer_buffers }, blt) } /// The impl-private scratch of every resolved op, whose addresses the @@ -479,6 +527,7 @@ fn compile_call( states: &BTreeMap, vars: &BTreeMap<&str, usize>, ranks: &Ranks, + pinned: &PinnedCheck, scan: &mut MulticastScan, launches: &mut Vec, ) -> Result<()> { @@ -553,7 +602,12 @@ fn compile_call( ); } match imp { - LaunchImpl::GemmBf16Tn { beta } => LaunchKind::Gemm { beta: *beta }, + LaunchImpl::GemmBf16Tn { beta, algo } => { + if let Some(algo) = algo { + pinned.check(algo, &slots, l.when(), vars)?; + } + LaunchKind::Gemm { beta: *beta, algo: *algo } + } _ => LaunchKind::GemmF32, } } diff --git a/crates/kern-runtime/src/cublas.rs b/crates/kern-runtime/src/cublas.rs index dfe9743..4fac49f 100644 --- a/crates/kern-runtime/src/cublas.rs +++ b/crates/kern-runtime/src/cublas.rs @@ -9,6 +9,7 @@ use cudarc::cublas; use cudarc::cublaslt::{self, CudaBlasLT, Matmul, MatmulConfig, MatmulShared}; use cudarc::driver::{sys, CudaSlice, CudaStream, DevicePtr, DevicePtrMut, DeviceSlice, SyncOnDrop}; use half::bf16; +use kern_manifest::types::GemmAlgo; use crate::compile::RVal; use crate::error::{bail, Error, Result}; @@ -42,15 +43,11 @@ impl DevicePtrMut for RawBf16 { } } -/// `extern:cublaslt_bf16_tn`: row-major `C[m,n] = A[m,k] @ W[n,k]^T`, -/// resolved args `[a, w, c, m, n, k]`. Column-major mapping: compute -/// `C_cm[n,m] = W_cm^T[n,k] x A_cm[k,m]` -> transa=T on W (lda=k), -/// transb=N on A (ldb=k), m'=n, n'=m, ldc=n. -/// `extern:cublaslt_bf16_tn_acc` is the same with beta=1: `C += A @ W^T`. -pub(crate) fn gemm_bf16_tn(blt: &CudaBlasLT, stream: &Arc, args: &[RVal], beta: f32) -> Result<()> { - // `c[m, n] (+)= a[m, k] @ w[n, k]^T`; an optional 7th arg is C's row - // stride in elements (default n), and an optional 8th arg is A's row - // stride (default k). This permits group views without transposing A. +/// The operands of a `cublaslt_bf16_tn[_acc]` call, `[a, w, c, m, n, k]` plus +/// an optional C row stride (default n) and A row stride (default k), in +/// elements, so a group view needs no transposed copy of A; checked against +/// their buffers, `None` for an empty product. +pub(crate) fn bf16_operands(args: &[RVal]) -> Result> { let (a, w, c, m, n, k, ldc, a_stride) = match args { [a, w, c, m, n, k] => (a, w, c, m.val, n.val, k.val, n.val, k.val), [a, w, c, m, n, k, ldc] => (a, w, c, m.val, n.val, k.val, ldc.val, k.val), @@ -64,11 +61,24 @@ pub(crate) fn gemm_bf16_tn(blt: &CudaBlasLT, stream: &Arc, args: &[R bail!(Manifest, "gemm: A row stride {a_stride} < k {k}"); } if m == 0 || n == 0 || k == 0 { - return Ok(()); + return Ok(None); } if a.bytes < ((m - 1) * a_stride + k) * 2 || w.bytes < n * k * 2 || c.bytes < ((m - 1) * ldc + n) * 2 { bail!(Manifest, "gemm: operands too small for m={m} n={n} k={k} ldc={ldc} A row stride={a_stride}"); } + Ok(Some((a, w, c, GemmShape { m, n, k, ldc, a_stride }))) +} + +/// `extern:cublaslt_bf16_tn`: row-major `C[m,n] = A[m,k] @ W[n,k]^T`, +/// resolved args `[a, w, c, m, n, k]`. Column-major mapping: compute +/// `C_cm[n,m] = W_cm^T[n,k] x A_cm[k,m]` -> transa=T on W (lda=k), +/// transb=N on A (ldb=k), m'=n, n'=m, ldc=n. +/// `extern:cublaslt_bf16_tn_acc` is the same with beta=1: `C += A @ W^T`. +/// Optional row strides as [`bf16_operands`] reads them. +pub(crate) fn gemm_bf16_tn(blt: &CudaBlasLT, stream: &Arc, args: &[RVal], beta: f32) -> Result<()> { + let Some((a, w, c, GemmShape { m, n, k, ldc, a_stride })) = bf16_operands(args)? else { + return Ok(()); + }; let view = |rv: &RVal| RawBf16 { ptr: rv.val, len: (rv.bytes / 2) as usize, stream: stream.clone() }; let cfg = MatmulConfig { transa: true, @@ -96,6 +106,176 @@ pub(crate) fn gemm_bf16_tn(blt: &CudaBlasLT, stream: &Arc, args: &[R Ok(()) } +/// A bf16 GEMM's shape: `C[m, n] = A[m, k] · W[n, k]ᵀ` with C's and A's row strides. +#[derive(Clone, Copy, Debug)] +pub(crate) struct GemmShape { + m: u64, + n: u64, + k: u64, + ldc: u64, + a_stride: u64, +} + +/// A matmul descriptor and the `w`, `a`, `c` layouts of a bf16 `C = A · Wᵀ`, +/// column-major as in [`gemm_bf16_tn`]; destroyed on drop. +struct Described { + desc: cublaslt::sys::cublasLtMatmulDesc_t, + layouts: [cublaslt::sys::cublasLtMatrixLayout_t; 3], +} + +impl Described { + fn new(s: GemmShape) -> Result { + use cublaslt::result; + use cublaslt::sys::{cublasComputeType_t, cublasLtMatmulDescAttributes_t as Attr, cudaDataType}; + let err = |e: cublaslt::result::CublasError| Error::Cuda(format!("cublasLt descriptor ({s:?}): {e:?}")); + let bf = cudaDataType::CUDA_R_16BF; + unsafe { + let desc = result::create_matmul_desc(cublasComputeType_t::CUBLAS_COMPUTE_32F, cudaDataType::CUDA_R_32F) + .map_err(err)?; + let mut d = Described { desc, layouts: [std::ptr::null_mut(); 3] }; + let t = cublas::sys::cublasOperation_t::CUBLAS_OP_T; + result::set_matmul_desc_attribute( + desc, + Attr::CUBLASLT_MATMUL_DESC_TRANSA, + &t as *const _ as *const c_void, + 4, + ) + .map_err(err)?; + d.layouts[0] = result::create_matrix_layout(bf, s.k, s.n, s.k as i64).map_err(err)?; + d.layouts[1] = result::create_matrix_layout(bf, s.k, s.m, s.a_stride as i64).map_err(err)?; + d.layouts[2] = result::create_matrix_layout(bf, s.n, s.m, s.ldc as i64).map_err(err)?; + Ok(d) + } + } +} + +impl Drop for Described { + fn drop(&mut self) { + use cublaslt::result; + unsafe { + for l in self.layouts { + if !l.is_null() { + let _ = result::destroy_matrix_layout(l); + } + } + let _ = result::destroy_matmul_desc(self.desc); + } + } +} + +/// A launch's pinned cuBLASLt algorithm, built from the manifest's config. +#[derive(Clone, Copy)] +pub(crate) struct Pinned(cublaslt::sys::cublasLtMatmulAlgo_t); + +impl Pinned { + /// `cublasLtMatmulAlgoInit` for the bf16 GEMM's types, then every config + /// attribute set as `algo` says. + pub(crate) fn new(blt: &CudaBlasLT, algo: &GemmAlgo) -> Result { + use cublaslt::sys::{ + self as lt, cublasComputeType_t, cublasLtMatmulAlgoConfigAttributes_t as Cfg, cudaDataType, + }; + let bf = cudaDataType::CUDA_R_16BF; + let mut raw = std::mem::MaybeUninit::::zeroed(); + let status = unsafe { + lt::cublasLtMatmulAlgoInit( + *blt.handle(), + cublasComputeType_t::CUBLAS_COMPUTE_32F, + cudaDataType::CUDA_R_32F, + bf, + bf, + bf, + bf, + algo.id, + raw.as_mut_ptr(), + ) + }; + if status != lt::cublasStatus_t::CUBLAS_STATUS_SUCCESS { + bail!(KernelArtifact, "cublasLt has no bf16 algorithm {} here ({status:?})", algo.id); + } + let mut raw = unsafe { raw.assume_init() }; + let mut set = |attr: Cfg, (v, size): (*const c_void, usize), what: &str| { + let status = unsafe { lt::cublasLtMatmulAlgoConfigSetAttribute(&mut raw, attr, v, size) }; + if status == lt::cublasStatus_t::CUBLAS_STATUS_SUCCESS { + Ok(()) + } else { + Err(Error::KernelArtifact(format!("cublasLt algorithm {}: {what} rejected ({status:?})", algo.id))) + } + }; + fn raw_of(v: &T) -> (*const c_void, usize) { + (v as *const T as *const c_void, std::mem::size_of::()) + } + set(Cfg::CUBLASLT_ALGO_CONFIG_TILE_ID, raw_of(&algo.tile), "tile")?; + set(Cfg::CUBLASLT_ALGO_CONFIG_STAGES_ID, raw_of(&algo.stages), "stages")?; + set(Cfg::CUBLASLT_ALGO_CONFIG_SPLITK_NUM, raw_of(&algo.split_k), "split_k")?; + set(Cfg::CUBLASLT_ALGO_CONFIG_REDUCTION_SCHEME, raw_of(&algo.reduction), "reduction")?; + set(Cfg::CUBLASLT_ALGO_CONFIG_CTA_SWIZZLING, raw_of(&algo.swizzle), "swizzle")?; + set(Cfg::CUBLASLT_ALGO_CONFIG_CUSTOM_OPTION, raw_of(&algo.custom), "custom")?; + set(Cfg::CUBLASLT_ALGO_CONFIG_INNER_SHAPE_ID, raw_of(&algo.inner_shape), "inner_shape")?; + set(Cfg::CUBLASLT_ALGO_CONFIG_CLUSTER_SHAPE_ID, raw_of(&algo.cluster_shape), "cluster_shape")?; + Ok(Pinned(raw)) + } + + /// `cublasLtMatmulAlgoCheck` at `shape`, the workspace within `Blas`'s. + pub(crate) fn check(&self, blt: &CudaBlasLT, shape: GemmShape) -> Result<()> { + use cublaslt::sys as lt; + let d = Described::new(shape)?; + let mut result = std::mem::MaybeUninit::::zeroed(); + let [w, a, c] = d.layouts; + let status = + unsafe { lt::cublasLtMatmulAlgoCheck(*blt.handle(), d.desc, w, a, c, c, &self.0, result.as_mut_ptr()) }; + if status != lt::cublasStatus_t::CUBLAS_STATUS_SUCCESS { + bail!(KernelArtifact, "the pinned cublasLt algorithm does not run at {shape:?} ({status:?})"); + } + let need = unsafe { result.assume_init() }.workspaceSize; + if need > Blas::WORKSPACE { + bail!( + KernelArtifact, + "the pinned cublasLt algorithm needs {need} bytes of workspace at {shape:?}, over {}", + Blas::WORKSPACE + ); + } + Ok(()) + } +} + +/// [`gemm_bf16_tn`] on a pinned algorithm: no heuristic, the workspace `Blas`'s. +pub(crate) fn gemm_bf16_tn_pinned( + blt: &CudaBlasLT, + blas: &Blas, + stream: &Arc, + args: &[RVal], + beta: f32, + algo: &Pinned, +) -> Result<()> { + let Some((a, w, c, shape)) = bf16_operands(args)? else { + return Ok(()); + }; + let d = Described::new(shape)?; + let [wl, al, cl] = d.layouts; + let alpha = 1.0f32; + unsafe { + cublaslt::result::matmul( + *blt.handle(), + d.desc, + &alpha as *const _ as *const c_void, + &beta as *const _ as *const c_void, + w.val as *const c_void, + wl, + a.val as *const c_void, + al, + c.val as *const c_void, + cl, + c.val as *mut c_void, + cl, + &algo.0, + blas.ws as *mut c_void, + Blas::WORKSPACE, + stream.cu_stream() as *mut _, + ) + } + .map_err(|e| Error::Cuda(format!("cublasLt matmul, pinned algorithm ({shape:?}): {e:?}"))) +} + /// A cuBLAS handle bound to the runtime's stream with its own workspace, for /// the f32-result GEMM built-in (`cublasGemmEx`; cublasLt's typed `Matmul` /// only lands in the operand type). Kept separate from the Lt handle so the diff --git a/crates/kern-runtime/src/exec.rs b/crates/kern-runtime/src/exec.rs index d7c0634..b0b3184 100644 --- a/crates/kern-runtime/src/exec.rs +++ b/crates/kern-runtime/src/exec.rs @@ -26,7 +26,7 @@ use std::os::raw::c_void; use cudarc::driver::sys; use crate::compile::{CompiledProgram, Dense, Launch, LaunchKind, RVal, Slot}; -use crate::cublas::{gemm_bf16_tn, gemm_bf16_tn_f32, gemm_fp8_tn}; +use crate::cublas::{gemm_bf16_tn, gemm_bf16_tn_f32, gemm_bf16_tn_pinned, gemm_fp8_tn}; use crate::error::{bail, cuda_check}; use crate::{Error, Result, Runtime}; @@ -224,7 +224,10 @@ impl Runtime { images.push(m); } match &l.kind { - LaunchKind::Gemm { beta } => gemm_bf16_tn(&self.blt, &self.stream, &vals, *beta), + LaunchKind::Gemm { beta, algo: None } => gemm_bf16_tn(&self.blt, &self.stream, &vals, *beta), + LaunchKind::Gemm { beta, algo: Some(algo) } => { + gemm_bf16_tn_pinned(&self.blt, &self.blas, &self.stream, &vals, *beta, algo) + } LaunchKind::GemmF32 => gemm_bf16_tn_f32(&self.blas, &vals), LaunchKind::GemmFp8 { f32_out } => gemm_fp8_tn(&self.blt, &self.blas, &self.stream, &vals, *f32_out), LaunchKind::Nccl { coll, elem, group } => self.collective(*coll, *elem, group, &vals), diff --git a/crates/kern-runtime/src/host.rs b/crates/kern-runtime/src/host.rs index 655c5fd..9b20194 100644 --- a/crates/kern-runtime/src/host.rs +++ b/crates/kern-runtime/src/host.rs @@ -76,7 +76,15 @@ impl Runtime { for (_, exec) in std::mem::take(&mut self.graphs) { cuda_check(unsafe { sys::cuGraphExecDestroy(exec) }, "cuGraphExecDestroy")?; } - match compile::compile(&self.manifest, &self.ops, &self.buffers, &self.states, &self.ranks, &self.peers) { + match compile::compile( + &self.manifest, + &self.ops, + &self.buffers, + &self.states, + &self.ranks, + &self.peers, + &self.blt, + ) { Ok(programs) => { self.programs = programs; Ok(()) diff --git a/crates/kern-runtime/src/load.rs b/crates/kern-runtime/src/load.rs index 06281e0..899f272 100644 --- a/crates/kern-runtime/src/load.rs +++ b/crates/kern-runtime/src/load.rs @@ -189,7 +189,7 @@ impl Runtime { } // Op scratch is allocated here, at var max, like the buffers. - let resolved = compile::resolve_ops(&manifest, &modules, kernels_dir, &stream, &vars_max)?; + let resolved = compile::resolve_ops(&manifest, &modules, kernels_dir, &stream, &blt, &vars_max)?; // Everything but the states is on the device now: what is left is // the states' to take (weights are bound into buffers already @@ -274,7 +274,7 @@ impl Runtime { let programs = if hosted { BTreeMap::new() } else { - compile::compile(&manifest, &resolved, &buffers, &states, &ranks, &peers)? + compile::compile(&manifest, &resolved, &buffers, &states, &ranks, &peers, &blt)? }; let provision = Provision { tokens: pool.pages_max() as u64 * page, seq_slots: pool.slots_max() as u64 }; diff --git a/crates/kern-runtime/tests/gemm_algo.rs b/crates/kern-runtime/tests/gemm_algo.rs new file mode 100644 index 0000000..0322a97 --- /dev/null +++ b/crates/kern-runtime/tests/gemm_algo.rs @@ -0,0 +1,159 @@ +//! A pinned cuBLASLt algorithm on `extern:cublaslt_bf16_tn`: it lands the +//! default op's exact product, per `when` range, in a run and in a captured +//! graph; one the shape cannot run fails the load. + +use std::collections::BTreeMap; +use std::os::raw::c_void; + +use cudarc::cublaslt::sys as lt; +use half::bf16; +use kern_manifest::Verified; +use kern_runtime::{Capacity, Runtime}; +use serde_json::{json, Value}; + +const MAX: usize = 256; +const N: usize = 512; +const K: usize = 1024; + +/// cuBLASLt's first heuristic answer for the bf16 `C[m, N] = A[m, K] · W[N, K]ᵀ` +/// as a manifest `algo`: what a caller that times candidates would write. +fn heuristic_algo(m: usize) -> Value { + use lt::cublasLtMatmulAlgoConfigAttributes_t as Cfg; + let ok = |s: lt::cublasStatus_t| assert_eq!(s, lt::cublasStatus_t::CUBLAS_STATUS_SUCCESS); + let bf = lt::cudaDataType::CUDA_R_16BF; + unsafe { + cudarc::driver::result::init().unwrap(); + let ctx = cudarc::driver::CudaContext::new(0).unwrap(); + ctx.bind_to_thread().unwrap(); + let mut handle = std::ptr::null_mut(); + ok(lt::cublasLtCreate(&mut handle)); + let mut desc = std::ptr::null_mut(); + ok(lt::cublasLtMatmulDescCreate( + &mut desc, + lt::cublasComputeType_t::CUBLAS_COMPUTE_32F, + lt::cudaDataType::CUDA_R_32F, + )); + let t = cudarc::cublas::sys::cublasOperation_t::CUBLAS_OP_T; + ok(lt::cublasLtMatmulDescSetAttribute( + desc, + lt::cublasLtMatmulDescAttributes_t::CUBLASLT_MATMUL_DESC_TRANSA, + &t as *const _ as *const c_void, + 4, + )); + let mut layouts = [std::ptr::null_mut(); 3]; + for (l, (rows, cols, ld)) in layouts.iter_mut().zip([(K, N, K), (K, m, K), (N, m, N)]) { + ok(lt::cublasLtMatrixLayoutCreate(l, bf, rows as u64, cols as u64, ld as i64)); + } + let mut pref = std::ptr::null_mut(); + ok(lt::cublasLtMatmulPreferenceCreate(&mut pref)); + let ws: usize = 32 << 20; + ok(lt::cublasLtMatmulPreferenceSetAttribute( + pref, + lt::cublasLtMatmulPreferenceAttributes_t::CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES, + &ws as *const _ as *const c_void, + 8, + )); + let mut found = std::mem::zeroed::(); + let mut count = 0; + let [w, a, c] = layouts; + ok(lt::cublasLtMatmulAlgoGetHeuristic(handle, desc, w, a, c, c, pref, 1, &mut found, &mut count)); + assert_eq!(count, 1); + let get = |attr: Cfg, size: usize| { + let mut v = 0u64; + let mut written = 0; + ok(lt::cublasLtMatmulAlgoConfigGetAttribute( + &found.algo, + attr, + &mut v as *mut _ as *mut c_void, + size, + &mut written, + )); + v + }; + let algo = json!({ + "id": get(Cfg::CUBLASLT_ALGO_CONFIG_ID, 4) as i32, + "tile": get(Cfg::CUBLASLT_ALGO_CONFIG_TILE_ID, 4), + "stages": get(Cfg::CUBLASLT_ALGO_CONFIG_STAGES_ID, 4), + "split_k": get(Cfg::CUBLASLT_ALGO_CONFIG_SPLITK_NUM, 4) as i32, + "reduction": get(Cfg::CUBLASLT_ALGO_CONFIG_REDUCTION_SCHEME, 4), + "swizzle": get(Cfg::CUBLASLT_ALGO_CONFIG_CTA_SWIZZLING, 4), + "custom": get(Cfg::CUBLASLT_ALGO_CONFIG_CUSTOM_OPTION, 4), + "inner_shape": get(Cfg::CUBLASLT_ALGO_CONFIG_INNER_SHAPE_ID, 2), + "cluster_shape": get(Cfg::CUBLASLT_ALGO_CONFIG_CLUSTER_SHAPE_ID, 2), + }); + lt::cublasLtMatmulPreferenceDestroy(pref); + for l in layouts { + lt::cublasLtMatrixLayoutDestroy(l); + } + lt::cublasLtMatmulDescDestroy(desc); + lt::cublasLtDestroy(handle); + algo + } +} + +/// The default op and a pinned one, one launch per range of `rows`. +fn manifest(small: Value, large: Value) -> Verified { + let params = json!(["in buffer", "in buffer", "out buffer", "i32", "i32", "i32"]); + let call = |op: &str, out: &str| json!({"op": op, "args": [{"buf": "a"}, {"buf": "w"}, {"buf": out}, {"var": "rows"}, {"i32": N}, {"i32": K}]}); + let m = json!({ + "schema_version": 5, "model": "gemm-algo-test", "vars": {"rows": {"max": MAX}}, "states": {}, + "buffers": { + "a": {"kind": "input", "dtype": "bf16", "shape": ["rows", K]}, + "w": {"kind": "input", "dtype": "bf16", "shape": [N, K]}, + "default": {"kind": "output", "dtype": "bf16", "shape": ["rows", N]}, + "pinned": {"kind": "output", "dtype": "bf16", "shape": ["rows", N]} + }, + "modules": {}, + "ops": { + "gemm": {"params": params, "impl": {"launches": [{"entry": "extern:cublaslt_bf16_tn"}]}}, + "gemm_pinned": {"params": params, "impl": {"launches": [ + {"entry": "extern:cublaslt_bf16_tn", "when": {"var": "rows", "max": 128}, "algo": small}, + {"entry": "extern:cublaslt_bf16_tn", "when": {"var": "rows", "min": 129}, "algo": large} + ]}} + }, + "programs": { + "both": {"batch": {"groups": 1, "rows": "rows"}, "graph": true, + "calls": [call("gemm", "default"), call("gemm_pinned", "pinned")]} + } + }); + Verified::from_json(&m.to_string()).unwrap() +} + +fn load(v: &Verified) -> kern_runtime::Result { + Runtime::load(v, None, 0, Some(Capacity { tokens: Some(1), seqs: 1 }), None) +} + +#[test] +#[ignore = "requires a CUDA GPU"] +fn a_pinned_algorithm_lands_the_default_product_in_each_range() { + let mut rt = load(&manifest(heuristic_algo(128), heuristic_algo(MAX))).unwrap(); + // a is a permutation (one 1.0 per row at column 7i mod K), w cycles through -2..=2: the product is exact. + let a: Vec = (0..MAX * K) + .flat_map(|i| bf16::from_f32(if i % K == (i / K * 7) % K { 1.0 } else { 0.0 }).to_le_bytes()) + .collect(); + let wv = |n: usize, k: usize| ((n + k) % 5) as f32 - 2.0; + let w: Vec = (0..N * K).flat_map(|i| bf16::from_f32(wv(i / K, i % K)).to_le_bytes()).collect(); + for rows in [100, MAX] { + let vars = BTreeMap::from([("rows".to_string(), rows as u64)]); + rt.write_input_at("a", &a[..rows * K * 2], &vars).unwrap(); + rt.write_input("w", &w).unwrap(); + let expected: Vec = + (0..rows * N).flat_map(|i| bf16::from_f32(wv(i % N, (i / N * 7) % K)).to_le_bytes()).collect(); + rt.run("both", &vars).unwrap(); + assert_eq!(rt.read_output("default").unwrap()[..rows * N * 2], expected[..], "default, {rows} rows"); + assert_eq!(rt.read_output("pinned").unwrap()[..rows * N * 2], expected[..], "pinned, {rows} rows"); + rt.capture("both", &vars).unwrap(); + rt.run_captured("both", &vars).unwrap(); + assert_eq!(rt.read_output("pinned").unwrap()[..rows * N * 2], expected[..], "pinned in a graph, {rows} rows"); + } +} + +#[test] +#[ignore = "requires a CUDA GPU"] +fn an_algorithm_the_shape_cannot_run_fails_the_load() { + let good = heuristic_algo(MAX); + let mut bad = good.clone(); + bad["tile"] = json!(9999); + let err = load(&manifest(good, bad)).err().expect("an unknown tile loaded"); + assert!(err.to_string().contains("kernel artifact"), "{err}"); +} diff --git a/docs/manifest.md b/docs/manifest.md index 189efa2..a5d1e61 100644 --- a/docs/manifest.md +++ b/docs/manifest.md @@ -97,7 +97,9 @@ topology.groups. group 多卡 SPMD 的 rank 组:只有名字和 `[min, max]` 内才发射,两端可省;区间外该 launch 什么都不做。manifest 没有控制流,一个 op 按形状选实现就写几个 launch 各管一段区间,例如小 batch 走自研 split-K、`min: 17` 起交回 `extern:cublaslt_bf16_tn`; - verifier 查 var 已声明、区间非空,不查各段是否覆盖全部取值)、`args` 连线:`{"param": i}` 转发接口第 i 参 / + verifier 查 var 已声明、区间非空,不查各段是否覆盖全部取值)、`algo` + (只限 `cublaslt_bf16_tn` / `_acc`:钉住的 cublasLt 算法配置,装载时对 + 该 launch 的形状检查,见 runtime.md)、`args` 连线:`{"param": i}` 转发接口第 i 参 / `{"scratch": name}` 接私有工作区 / 字面量标量(impl 私有常量)/ `{"rank": group}` / `{"pack": {...}}`;**不写 = 按序转发接口参数**。 **bytes / pack** 是 launch 私有的参数类型:核的 ABI 收 struct diff --git a/docs/qwen38-bringup.md b/docs/qwen38-bringup.md index d9caabe..7f32f18 100644 --- a/docs/qwen38-bringup.md +++ b/docs/qwen38-bringup.md @@ -402,8 +402,8 @@ forward,launch 开销被摊薄。 - **manifest 里唯一没被 sha256 钉住的东西就是残差的来源。** Stage 1 的逐 op 对比把 kern 与 vLLM 的差异收敛到 cuBLAS 在 M=43、N=96 GEMM 上的算法选择 (1 ulp / 4 个元素)——`extern:cublaslt_bf16_tn` 是 manifest 里唯一由 - runtime 自行挑算法的 dispatch。下一步自然是把 cublasLt 的 algo id 也写进 - manifest(`extern` 带 `algo` 字段),让 GEMM 和 Triton 核一样可钉、可 diff。 + runtime 自行挑算法的 dispatch。extern launch 现在可以带 `algo` 把 cublasLt + 算法写进 manifest(见 runtime.md),让 GEMM 和 Triton 核一样可钉、可 diff。 - **`kern test` 需要"外部参考"这一侧。** 现在它 diff 的是两份 manifest; 这次真正有用的是"manifest vs vLLM 的逐 op 中间量"(`KERN_PROBE_LAYER` + `qwen38_probe_vllm.py`)。把 vLLM 的 forward hook 输出当作一份"参考 tap" diff --git a/docs/runtime.md b/docs/runtime.md index 09d9640..422676d 100644 --- a/docs/runtime.md +++ b/docs/runtime.md @@ -137,6 +137,17 @@ park 31 ms、wake 28 ms,发起 0.9 / 0.5 ms(2500 页折成 ~50 次拷贝) (默认 n),再追加第 8 参 A 行步长(默认 k),单位都是元素。 配合 buffer 字节 offset,可直接读取交错存储的组并写入输出列带, 无需先转置或复制输入。 +这两个入口的 launch 可写 `algo` 钉住 cublasLt 算法(`cublasLtMatmulAlgoConfigAttributes_t` +的全部取值:`id` / `tile` / `stages` / `split_k` / `reduction` / `swizzle` / `custom` / +`inner_shape` / `cluster_shape`),不写时照旧取启发式第一名。装载时 `cublasLtMatmulAlgoInit` +建算法、逐项设属性,再在该 launch 会遇到的形状两端(所有 var 取下界、取上界,`when` 的 var +夹在自己的区间里)跑 `cublasLtMatmulAlgoCheck`,并要求 workspace 不超过 `Blas` 的 32 MiB; +设不上或查不过就是 kernel artifact 错误,绝不静默退回启发式。发射时不查启发式、不同步、 +不回读,workspace 用 `Blas` 的。按形状换算法用 `when`:一个 op 每段 m 区间一个 launch, +各带自己的 `algo`。算法怎么选不归 kern:调用方在自己的真实负载里计时,把赢家写进 +manifest,输出照常由 `kern test` 对参考把关。不切 K 不等于逐位相同:sm_89 上 `sliced` +的核(algo 30/31/16)在块内按 warp 切 k,`SPLITK_NUM` 却报 1,结果与 algo 5/6/21 不同; +cuBLAS 只承诺同一算法在同架构、同库版本上可复现,所以换 `algo` 前要对输出重新把关。 `extern:cublas_bf16_tn_f32` 同一映射但结果落 **f32**(cublasGemmEx, `CUBLAS_COMPUTE_32F` / `DEFAULT_TENSOR_OP`,独立 cuBLAS handle + 32 MiB workspace,可捕获):K3 的每条稠密投影都是 f32 partial 再由认证的 `k3_land` diff --git a/schema/manifest-v5.schema.json b/schema/manifest-v5.schema.json index 950c881..62b8d84 100644 --- a/schema/manifest-v5.schema.json +++ b/schema/manifest-v5.schema.json @@ -642,6 +642,17 @@ "additionalProperties": false, "description": "A runtime built-in launch, e.g. `{\"entry\": \"extern:cublaslt_bf16_tn\"}`.", "properties": { + "algo": { + "anyOf": [ + { + "$ref": "#/$defs/GemmAlgo" + }, + { + "type": "null" + } + ], + "description": "For `cublaslt_bf16_tn` and `cublaslt_bf16_tn_acc`: the cuBLASLt algorithm to run instead of the heuristic's first answer." + }, "args": { "description": "Where each launch param comes from (default: the op's params in order).", "items": { @@ -962,6 +973,168 @@ } ] }, + "GemmAlgo": { + "additionalProperties": false, + "description": "One cuBLASLt algorithm as its `cublasLtMatmulAlgoConfigAttributes_t` values, e.g. `{\"id\": 6, \"tile\": 24, \"stages\": 9, \"split_k\": 1, \"reduction\": 0, \"swizzle\": 0, \"custom\": 0, \"inner_shape\": 0, \"cluster_shape\": 0}`. Built at load and checked against every shape the launch runs at; one that does not fit is an error, never a fallback.", + "properties": { + "cluster_shape": { + "anyOf": [ + { + "description": "`CUBLASLT_ALGO_CONFIG_CLUSTER_SHAPE_ID`.", + "format": "uint16", + "maximum": 65535, + "minimum": 0, + "type": "integer" + }, + { + "description": "Name of a numeric constant declared in constants.", + "minLength": 1, + "type": "string" + } + ], + "description": "`CUBLASLT_ALGO_CONFIG_CLUSTER_SHAPE_ID`." + }, + "custom": { + "anyOf": [ + { + "description": "`CUBLASLT_ALGO_CONFIG_CUSTOM_OPTION`.", + "format": "uint32", + "minimum": 0, + "type": "integer" + }, + { + "description": "Name of a numeric constant declared in constants.", + "minLength": 1, + "type": "string" + } + ], + "description": "`CUBLASLT_ALGO_CONFIG_CUSTOM_OPTION`." + }, + "id": { + "anyOf": [ + { + "description": "`CUBLASLT_ALGO_CONFIG_ID`.", + "format": "int32", + "type": "integer" + }, + { + "description": "Name of a numeric constant declared in constants.", + "minLength": 1, + "type": "string" + } + ], + "description": "`CUBLASLT_ALGO_CONFIG_ID`." + }, + "inner_shape": { + "anyOf": [ + { + "description": "`CUBLASLT_ALGO_CONFIG_INNER_SHAPE_ID`.", + "format": "uint16", + "maximum": 65535, + "minimum": 0, + "type": "integer" + }, + { + "description": "Name of a numeric constant declared in constants.", + "minLength": 1, + "type": "string" + } + ], + "description": "`CUBLASLT_ALGO_CONFIG_INNER_SHAPE_ID`." + }, + "reduction": { + "anyOf": [ + { + "description": "`CUBLASLT_ALGO_CONFIG_REDUCTION_SCHEME`.", + "format": "uint32", + "minimum": 0, + "type": "integer" + }, + { + "description": "Name of a numeric constant declared in constants.", + "minLength": 1, + "type": "string" + } + ], + "description": "`CUBLASLT_ALGO_CONFIG_REDUCTION_SCHEME`." + }, + "split_k": { + "anyOf": [ + { + "description": "`CUBLASLT_ALGO_CONFIG_SPLITK_NUM`.", + "format": "int32", + "type": "integer" + }, + { + "description": "Name of a numeric constant declared in constants.", + "minLength": 1, + "type": "string" + } + ], + "description": "`CUBLASLT_ALGO_CONFIG_SPLITK_NUM`." + }, + "stages": { + "anyOf": [ + { + "description": "`CUBLASLT_ALGO_CONFIG_STAGES_ID`.", + "format": "uint32", + "minimum": 0, + "type": "integer" + }, + { + "description": "Name of a numeric constant declared in constants.", + "minLength": 1, + "type": "string" + } + ], + "description": "`CUBLASLT_ALGO_CONFIG_STAGES_ID`." + }, + "swizzle": { + "anyOf": [ + { + "description": "`CUBLASLT_ALGO_CONFIG_CTA_SWIZZLING`.", + "format": "uint32", + "minimum": 0, + "type": "integer" + }, + { + "description": "Name of a numeric constant declared in constants.", + "minLength": 1, + "type": "string" + } + ], + "description": "`CUBLASLT_ALGO_CONFIG_CTA_SWIZZLING`." + }, + "tile": { + "anyOf": [ + { + "description": "`CUBLASLT_ALGO_CONFIG_TILE_ID`.", + "format": "uint32", + "minimum": 0, + "type": "integer" + }, + { + "description": "Name of a numeric constant declared in constants.", + "minLength": 1, + "type": "string" + } + ], + "description": "`CUBLASLT_ALGO_CONFIG_TILE_ID`." + } + }, + "required": [ + "id", + "tile", + "stages", + "split_k", + "reduction", + "swizzle", + "custom", + "inner_shape", + "cluster_shape" + ], + "type": "object" + }, "HostTensor": { "additionalProperties": false, "description": "The layout of a host-allocated state: a strided tensor over the host's memory. Only the outermost extent may be 0, the host's to choose (how many blocks it allocated); every other extent and every stride is fixed.",