Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 27 additions & 0 deletions crates/kern-manifest/src/types.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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<When>,
/// 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<GemmAlgo>,
/// Where each launch param comes from (default: the op's params in order).
#[serde(default, skip_serializing_if = "Option::is_none")]
args: Option<Vec<LaunchArg>>,
}

/// 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)]
Expand Down
5 changes: 5 additions & 0 deletions crates/kern-manifest/src/verify.rs
Original file line number Diff line number Diff line change
Expand Up @@ -478,6 +478,11 @@ fn diagnostics(m: &Manifest) -> Vec<String> {
"{ctx}: a launch without a module must be a runtime built-in (`extern:<name>`)"
));
}
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() {
Expand Down
70 changes: 62 additions & 8 deletions crates/kern-runtime/src/compile.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};
Expand Down Expand Up @@ -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<Pinned> },
/// `extern:cublas_bf16_tn_f32`: same operands, f32 result (cublasGemmEx).
GemmF32,
/// `extern:cublaslt_fp8_tn` / `..._f32`: e4m3 operands with per-tensor
Expand Down Expand Up @@ -226,6 +229,7 @@ enum LaunchImpl {
},
GemmBf16Tn {
beta: f32,
algo: Option<Pinned>,
},
/// cublasGemmEx with an f32 result (`extern:cublas_bf16_tn_f32`).
GemmBf16TnF32,
Expand Down Expand Up @@ -295,6 +299,7 @@ pub(crate) fn resolve_ops(
modules: &[LoadedModule],
kernels_dir: Option<&Path>,
stream: &Arc<CudaStream>,
blt: &CudaBlasLT,
vars_max: &BTreeMap<String, u64>,
) -> Result<BTreeMap<String, ResolvedOp>> {
let mut ops = BTreeMap::new();
Expand All @@ -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 }),
Expand Down Expand Up @@ -412,6 +422,40 @@ struct Ranks<'a> {
peer_buffers: &'a BTreeSet<String>,
}

/// 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::<Result<Vec<_>>>()?;
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(
Expand All @@ -420,8 +464,11 @@ fn compile_programs(
buffers: &BTreeMap<String, DeviceBuf>,
states: &BTreeMap<String, DeviceBuf>,
ranks: &Ranks,
blt: &CudaBlasLT,
) -> Result<BTreeMap<String, CompiledProgram>> {
let vars: BTreeMap<&str, usize> = manifest.vars.keys().enumerate().map(|(i, s)| (s.as_str(), i)).collect();
let maxima: Vec<u64> = 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 {
Expand All @@ -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()));
}
Expand All @@ -454,9 +501,10 @@ pub(crate) fn compile(
states: &BTreeMap<String, DeviceBuf>,
ranks: &BTreeMap<String, u64>,
peers: &BTreeMap<String, crate::peers::PeerSlot>,
blt: &CudaBlasLT,
) -> Result<BTreeMap<String, CompiledProgram>> {
let peer_buffers: BTreeSet<String> = 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
Expand All @@ -479,6 +527,7 @@ fn compile_call(
states: &BTreeMap<String, DeviceBuf>,
vars: &BTreeMap<&str, usize>,
ranks: &Ranks,
pinned: &PinnedCheck,
scan: &mut MulticastScan,
launches: &mut Vec<Launch>,
) -> Result<()> {
Expand Down Expand Up @@ -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,
}
}
Expand Down
Loading
Loading