use super::super::buffers::{
GpuBuffer, ManagedAllocationEpochStamp, managed_allocation_epoch_for_ranges,
};
use super::super::context::{GpuCtx, HalfTriadPolicy};
use super::super::kernels::MambaKernels as GpuKernels;
use super::contract::*;
use super::dispatch::*;
use super::sm89_tf32_joint_source::{
Sm89Tf32JointGemmParams, Sm89Tf32JointTransposeParams, TN_PRE_RNA_TRANSPOSE_SYMBOL,
};
use crate::mamba_ssm::gpu::kernel_identity::{
FramedSha256, GemmPolicy, GemmRouteIdentity, ModuleKind, NoPhysicalObserver,
PhysicalConversionArguments, PhysicalCudaLaunchError, PhysicalGemmBackend,
PhysicalLaunchObservation, PhysicalLaunchObserver, PolicyDtype, RecordingPhysicalObserver,
ResolvedGemmLaunchSet, ResolvedGemmLaunchSetBuilder, ResolvedGemmOp, ResolvedGemmRoute,
ResolvedInputTransform, ResolvedInstructionFamily, ResolvedInstructionShape,
ResolvedKernelLaunch, ResolvedNumericContract, ResolvedOperandConversion,
ResolvedOutputOwnership, ResolvedPhysicalKernelLaunch, ResolvedTransformOutputOwnership,
SCHEDULE_REVISION, SM89_EXACT_F32_D128_ROUTE_REVISION, SM89_EXACT_F32_TN_ROUTE_REVISION,
SM89_FIXED_COPYPLAN_ROUTE_REVISION, SM89_HALF_ROUTE_REVISION, Sha256Digest,
TUNING_TABLE_REVISION, build_resolved_gemm_launch_set, build_zero_reduction_route_identity,
enqueue_prepared_physical_launch, enqueue_with_physical_observation,
prepare_recording_physical_observer, resolve_physical_launch_observation,
};
use cudarc::driver::{
CudaFunction, CudaStream, DeviceRepr, LaunchArgs, LaunchConfig, PushKernelArg,
};
use std::collections::HashMap;
use std::sync::Arc;
#[derive(Clone, Copy)]
pub(in crate::mamba_ssm::gpu) struct PhysicalArgumentRange {
pub(crate) pointer: CUptr,
pub(crate) required_bytes: u64,
}
pub(in crate::mamba_ssm::gpu) fn prepare_physical_observer(
ctx: &GpuCtx,
capacity: usize,
ranges: &[PhysicalArgumentRange],
) -> Result<RecordingPhysicalObserver, String> {
let allocation_domain =
validated_allocation_domain(&ctx.stream, &ctx.kernels, "physical half launch")?;
let mut argument_allocations = Vec::new();
argument_allocations
.try_reserve_exact(ranges.len())
.map_err(|error| format!("reserve physical argument identities: {error}"))?;
for range in ranges {
argument_allocations.push(Sm90aAllocationIdentity::query(
range.pointer,
range.required_bytes,
allocation_domain,
"physical half launch",
"argument",
)?);
}
if argument_allocations.is_empty() {
return Err("physical launch observer requires prevalidated arguments".into());
}
let mut managed_ranges = Vec::new();
managed_ranges
.try_reserve_exact(ranges.len())
.map_err(|error| format!("reserve physical allocation liveness ranges: {error}"))?;
managed_ranges.extend(
ranges
.iter()
.map(|range| (range.pointer, range.required_bytes)),
);
let managed_epoch =
managed_allocation_epoch_for_ranges(allocation_domain.context_handle, &managed_ranges);
let argument_allocations = argument_allocations.into_boxed_slice();
let observer = prepare_recording_physical_observer(
ctx,
capacity,
managed_epoch,
move |pointer, required_bytes| {
argument_allocations
.iter()
.find_map(|identity| identity.physical_subrange(pointer, required_bytes))
.map(Sm90aAllocationIdentity::physical_digest)
.ok_or_else(|| {
format!(
"physical launch argument range is not covered by its prevalidated allocation: pointer={pointer:#x}, bytes={required_bytes}"
)
})
},
)?;
observer.validate_start(capacity)?;
Ok(observer)
}
#[derive(Clone, Copy, Debug, PartialEq)]
#[repr(C)]
struct Sm80Tf32KernelParams {
alpha: f32,
beta: f32,
m: i32,
k: i32,
n: i32,
lda: i32,
ldb: i32,
ldc: i32,
}
unsafe impl DeviceRepr for Sm80Tf32KernelParams {}
unsafe impl DeviceRepr for Sm89Tf32JointGemmParams {}
unsafe impl DeviceRepr for Sm89Tf32JointTransposeParams {}
#[derive(Clone, Copy, Debug, PartialEq)]
#[repr(C)]
struct SgbZeroReductionParams {
alpha: f32,
beta: f32,
m: i32,
k: i32,
n: i32,
lda: i32,
ldb: i32,
ldc: i32,
}
unsafe impl DeviceRepr for SgbZeroReductionParams {}
#[derive(Clone, Copy, Debug, PartialEq)]
#[repr(C)]
struct SgbNnM64N64Params {
alpha: f32,
beta: f32,
m: i32,
n: i32,
k: i32,
lda: i32,
ldb: i32,
ldc: i32,
}
unsafe impl DeviceRepr for SgbNnM64N64Params {}
#[derive(Clone, Copy, Debug, PartialEq)]
#[repr(C)]
struct Sm90aTf32KernelParams {
a_x: i32,
a_y: i32,
b_x: i32,
b_y: i32,
alpha: f32,
beta: f32,
m: i32,
k: i32,
n: i32,
ldc: i32,
}
unsafe impl DeviceRepr for Sm90aTf32KernelParams {}
#[derive(Clone, Copy, Debug, PartialEq)]
#[repr(C)]
struct Sm100KernelParams {
a_x: i32,
a_y: i32,
b_x: i32,
b_y: i32,
alpha: f32,
beta: f32,
m: i32,
k: i32,
n: i32,
ldc: i32,
}
unsafe impl DeviceRepr for Sm100KernelParams {}
#[derive(Clone, Copy, Debug, PartialEq)]
#[repr(C)]
struct Sm120KernelParams {
a_x: i32,
a_y: i32,
b_x: i32,
b_y: i32,
alpha: f32,
beta: f32,
m: i32,
k: i32,
n: i32,
ldc: i32,
}
unsafe impl DeviceRepr for Sm120KernelParams {}
unsafe impl DeviceRepr for Sm120FmaKernelParams {}
pub(super) const GEMM_BI_ZERO_REDUCTION_PARAMS_SIZE: usize =
std::mem::size_of::<SgbZeroReductionParams>();
pub(super) const fn tf32_kernel_params_size(module_kind: ModuleKind) -> Option<usize> {
match module_kind {
ModuleKind::TriadSm80 | ModuleKind::TriadSm89Finalist => {
Some(std::mem::size_of::<Sm80Tf32KernelParams>())
}
ModuleKind::TriadSm90a => Some(std::mem::size_of::<Sm90aTf32KernelParams>()),
ModuleKind::TriadSm100 => Some(std::mem::size_of::<Sm100KernelParams>()),
ModuleKind::TriadSm120 => Some(std::mem::size_of::<Sm120KernelParams>()),
_ => None,
}
}
impl Sm100KernelParams {
fn into_words(self) -> [u32; 10] {
[
self.a_x as u32,
self.a_y as u32,
self.b_x as u32,
self.b_y as u32,
self.alpha.to_bits(),
self.beta.to_bits(),
self.m as u32,
self.k as u32,
self.n as u32,
self.ldc as u32,
]
}
fn from_words(words: [u32; 10]) -> Self {
Self {
a_x: words[0] as i32,
a_y: words[1] as i32,
b_x: words[2] as i32,
b_y: words[3] as i32,
alpha: f32::from_bits(words[4]),
beta: f32::from_bits(words[5]),
m: words[6] as i32,
k: words[7] as i32,
n: words[8] as i32,
ldc: words[9] as i32,
}
}
}
impl Sm120KernelParams {
fn into_words(self) -> [u32; 10] {
[
self.a_x as u32,
self.a_y as u32,
self.b_x as u32,
self.b_y as u32,
self.alpha.to_bits(),
self.beta.to_bits(),
self.m as u32,
self.k as u32,
self.n as u32,
self.ldc as u32,
]
}
fn from_words(words: [u32; 10]) -> Self {
Self {
a_x: words[0] as i32,
a_y: words[1] as i32,
b_x: words[2] as i32,
b_y: words[3] as i32,
alpha: f32::from_bits(words[4]),
beta: f32::from_bits(words[5]),
m: words[6] as i32,
k: words[7] as i32,
n: words[8] as i32,
ldc: words[9] as i32,
}
}
}
fn prepare_f32_maps_with<Plan, Encode>(
request: F32TriadRequest,
operands: F32TriadOperands,
route: Tf32PhysicalRoute,
binding: Tf32MapBinding,
plan: Plan,
encode: Encode,
) -> Result<F32PreparedTensorMaps, String>
where
Plan: FnOnce(
F32TriadRequest,
F32TriadOperands,
Tf32PhysicalRoute,
) -> Result<Tf32TensorMapPlan, String>,
Encode: FnOnce(Tf32TensorMapPlan) -> Result<F32PreparedTensorMaps, String>,
{
if request.shape.reduction(request.op) == 0 {
let format = match route {
Tf32PhysicalRoute::Sm120TmaMmaTf32Rna(_)
| Tf32PhysicalRoute::Sm120TmaMmaTf32RnaStreamKV1(_) => Tf32TensorMapFormat::Uint32,
_ => Tf32TensorMapFormat::Tfloat32,
};
return Ok(F32PreparedTensorMaps::zero_reduction(
request,
Some(route),
Some(binding),
format,
));
}
encode(plan(request, operands, route)?)
}
macro_rules! abi_const_assert_eq {
($left:expr, $right:expr) => {
const _: [(); $right] = [(); $left];
};
}
abi_const_assert_eq!(std::mem::size_of::<Sm80Tf32KernelParams>(), 32);
abi_const_assert_eq!(std::mem::align_of::<Sm80Tf32KernelParams>(), 4);
abi_const_assert_eq!(std::mem::offset_of!(Sm80Tf32KernelParams, alpha), 0);
abi_const_assert_eq!(std::mem::offset_of!(Sm80Tf32KernelParams, beta), 4);
abi_const_assert_eq!(std::mem::offset_of!(Sm80Tf32KernelParams, m), 8);
abi_const_assert_eq!(std::mem::offset_of!(Sm80Tf32KernelParams, k), 12);
abi_const_assert_eq!(std::mem::offset_of!(Sm80Tf32KernelParams, n), 16);
abi_const_assert_eq!(std::mem::offset_of!(Sm80Tf32KernelParams, lda), 20);
abi_const_assert_eq!(std::mem::offset_of!(Sm80Tf32KernelParams, ldb), 24);
abi_const_assert_eq!(std::mem::offset_of!(Sm80Tf32KernelParams, ldc), 28);
abi_const_assert_eq!(std::mem::size_of::<SgbZeroReductionParams>(), 32);
abi_const_assert_eq!(std::mem::align_of::<SgbZeroReductionParams>(), 4);
abi_const_assert_eq!(std::mem::offset_of!(SgbZeroReductionParams, alpha), 0);
abi_const_assert_eq!(std::mem::offset_of!(SgbZeroReductionParams, beta), 4);
abi_const_assert_eq!(std::mem::offset_of!(SgbZeroReductionParams, m), 8);
abi_const_assert_eq!(std::mem::offset_of!(SgbZeroReductionParams, k), 12);
abi_const_assert_eq!(std::mem::offset_of!(SgbZeroReductionParams, n), 16);
abi_const_assert_eq!(std::mem::offset_of!(SgbZeroReductionParams, lda), 20);
abi_const_assert_eq!(std::mem::offset_of!(SgbZeroReductionParams, ldb), 24);
abi_const_assert_eq!(std::mem::offset_of!(SgbZeroReductionParams, ldc), 28);
abi_const_assert_eq!(std::mem::size_of::<SgbNnM64N64Params>(), 32);
abi_const_assert_eq!(std::mem::align_of::<SgbNnM64N64Params>(), 4);
abi_const_assert_eq!(std::mem::offset_of!(SgbNnM64N64Params, alpha), 0);
abi_const_assert_eq!(std::mem::offset_of!(SgbNnM64N64Params, beta), 4);
abi_const_assert_eq!(std::mem::offset_of!(SgbNnM64N64Params, m), 8);
abi_const_assert_eq!(std::mem::offset_of!(SgbNnM64N64Params, n), 12);
abi_const_assert_eq!(std::mem::offset_of!(SgbNnM64N64Params, k), 16);
abi_const_assert_eq!(std::mem::offset_of!(SgbNnM64N64Params, lda), 20);
abi_const_assert_eq!(std::mem::offset_of!(SgbNnM64N64Params, ldb), 24);
abi_const_assert_eq!(std::mem::offset_of!(SgbNnM64N64Params, ldc), 28);
abi_const_assert_eq!(std::mem::size_of::<Sm90aTf32KernelParams>(), 40);
abi_const_assert_eq!(std::mem::align_of::<Sm90aTf32KernelParams>(), 4);
abi_const_assert_eq!(std::mem::offset_of!(Sm90aTf32KernelParams, a_x), 0);
abi_const_assert_eq!(std::mem::offset_of!(Sm90aTf32KernelParams, a_y), 4);
abi_const_assert_eq!(std::mem::offset_of!(Sm90aTf32KernelParams, b_x), 8);
abi_const_assert_eq!(std::mem::offset_of!(Sm90aTf32KernelParams, b_y), 12);
abi_const_assert_eq!(std::mem::offset_of!(Sm90aTf32KernelParams, alpha), 16);
abi_const_assert_eq!(std::mem::offset_of!(Sm90aTf32KernelParams, beta), 20);
abi_const_assert_eq!(std::mem::offset_of!(Sm90aTf32KernelParams, m), 24);
abi_const_assert_eq!(std::mem::offset_of!(Sm90aTf32KernelParams, k), 28);
abi_const_assert_eq!(std::mem::offset_of!(Sm90aTf32KernelParams, n), 32);
abi_const_assert_eq!(std::mem::offset_of!(Sm90aTf32KernelParams, ldc), 36);
abi_const_assert_eq!(std::mem::size_of::<Sm100KernelParams>(), 40);
abi_const_assert_eq!(std::mem::align_of::<Sm100KernelParams>(), 4);
abi_const_assert_eq!(std::mem::offset_of!(Sm100KernelParams, a_x), 0);
abi_const_assert_eq!(std::mem::offset_of!(Sm100KernelParams, a_y), 4);
abi_const_assert_eq!(std::mem::offset_of!(Sm100KernelParams, b_x), 8);
abi_const_assert_eq!(std::mem::offset_of!(Sm100KernelParams, b_y), 12);
abi_const_assert_eq!(std::mem::offset_of!(Sm100KernelParams, alpha), 16);
abi_const_assert_eq!(std::mem::offset_of!(Sm100KernelParams, beta), 20);
abi_const_assert_eq!(std::mem::offset_of!(Sm100KernelParams, m), 24);
abi_const_assert_eq!(std::mem::offset_of!(Sm100KernelParams, k), 28);
abi_const_assert_eq!(std::mem::offset_of!(Sm100KernelParams, n), 32);
abi_const_assert_eq!(std::mem::offset_of!(Sm100KernelParams, ldc), 36);
abi_const_assert_eq!(std::mem::size_of::<Sm120KernelParams>(), 40);
abi_const_assert_eq!(std::mem::align_of::<Sm120KernelParams>(), 4);
abi_const_assert_eq!(std::mem::offset_of!(Sm120KernelParams, a_x), 0);
abi_const_assert_eq!(std::mem::offset_of!(Sm120KernelParams, a_y), 4);
abi_const_assert_eq!(std::mem::offset_of!(Sm120KernelParams, b_x), 8);
abi_const_assert_eq!(std::mem::offset_of!(Sm120KernelParams, b_y), 12);
abi_const_assert_eq!(std::mem::offset_of!(Sm120KernelParams, alpha), 16);
abi_const_assert_eq!(std::mem::offset_of!(Sm120KernelParams, beta), 20);
abi_const_assert_eq!(std::mem::offset_of!(Sm120KernelParams, m), 24);
abi_const_assert_eq!(std::mem::offset_of!(Sm120KernelParams, k), 28);
abi_const_assert_eq!(std::mem::offset_of!(Sm120KernelParams, n), 32);
abi_const_assert_eq!(std::mem::offset_of!(Sm120KernelParams, ldc), 36);
#[derive(Clone, Copy)]
enum PreparedTf32Params {
Sm80(Sm80Tf32KernelParams),
Sm90a(Sm90aTf32KernelParams),
Sm100(Sm100KernelParams),
Sm120(Sm120KernelParams),
Sm120Fma(Sm120FmaKernelParams),
}
#[derive(Clone, Copy)]
struct Tf32RawLaunch<'a> {
operands: F32TriadOperands,
route: Tf32PhysicalRoute,
maps: Option<&'a F32PreparedTensorMaps>,
params: PreparedTf32Params,
config: LaunchConfig,
zero_reduction: bool,
symbol: &'static str,
observation: Option<PhysicalLaunchObservation>,
streamk: Option<Tf32StreamKWorkspace>,
}
#[derive(Clone, Copy, Debug)]
struct Tf32StreamKWorkspace {
partial: CUptr,
flags: CUptr,
}
#[derive(Clone, Copy, Debug)]
struct Tf32StreamKLaunchPlan {
grid: u32,
partial_elements: usize,
flag_elements: usize,
}
fn tf32_streamk_launch_plan(
request: F32TriadRequest,
spec: &Tf32KernelSpec,
multiprocessor_count: u32,
) -> Result<Tf32StreamKLaunchPlan, String> {
request.shape.validate(request.op)?;
if request.op != spec.op {
return Err(format!(
"stream-K TF32 specification {:?} does not match {:?}",
spec.op, request.op
));
}
if multiprocessor_count == 0 {
return Err("stream-K TF32 launch requires at least one multiprocessor".into());
}
let rows = request.shape.output_rows(request.op);
let columns = request.shape.output_columns(request.op);
let tiles = rows
.div_ceil(spec.tile.0 as usize)
.checked_mul(columns.div_ceil(spec.tile.1 as usize))
.ok_or_else(|| invalid_gemm_dimensions("stream-K TF32 tile count overflows usize"))?;
let k_tiles = request
.shape
.reduction(request.op)
.div_ceil(spec.bk as usize)
.max(1);
let units = tiles
.checked_mul(k_tiles)
.ok_or_else(|| invalid_gemm_dimensions("stream-K TF32 unit count overflows usize"))?;
if units > i32::MAX as usize {
return Err(invalid_gemm_dimensions(
"stream-K TF32 unit count exceeds the kernel's 32-bit range",
));
}
let grid = multiprocessor_count;
let partial_elements = (grid as usize)
.checked_mul(SM120_TF32_STREAMK_SLOTS_PER_CTA)
.and_then(|slots| slots.checked_mul(SM120_TF32_STREAMK_SLAB_FLOATS))
.ok_or_else(|| invalid_gemm_dimensions("stream-K TF32 slab extent overflows usize"))?;
if partial_elements > SPLITK_SCRATCH_CAP {
return Err(invalid_gemm_dimensions(
"stream-K TF32 slabs exceed the fixed workspace",
));
}
let flag_elements = (grid as usize) * SM120_TF32_STREAMK_SLOTS_PER_CTA;
if flag_elements > TF32_SPLITK_COUNTER_CAP {
return Err(invalid_gemm_dimensions(
"stream-K TF32 flags exceed the fixed counter workspace",
));
}
Ok(Tf32StreamKLaunchPlan {
grid,
partial_elements,
flag_elements,
})
}
enum PreparedF32Kind {
Scalar(ScalarDispatchPlan),
ScalarZero {
maps: F32PreparedTensorMaps,
params: SgbZeroReductionParams,
config: cudarc::driver::LaunchConfig,
},
Tf32 {
route: Tf32PhysicalRoute,
maps: Option<F32PreparedTensorMaps>,
params: PreparedTf32Params,
config: cudarc::driver::LaunchConfig,
},
Tf32SplitK {
route: Tf32PhysicalRoute,
params: Sm80Tf32KernelParams,
plan: Tf32SplitKLaunchPlan,
workspace: Tf32SplitKWorkspace,
},
Tf32StreamK {
route: Tf32PhysicalRoute,
maps: F32PreparedTensorMaps,
params: PreparedTf32Params,
config: cudarc::driver::LaunchConfig,
plan: Tf32StreamKLaunchPlan,
workspace: Tf32StreamKWorkspace,
},
Tf32TnPreRna {
route: Tf32PhysicalRoute,
transpose_params: Sm89Tf32JointTransposeParams,
gemm_params: Sm89Tf32JointGemmParams,
transpose_config: cudarc::driver::LaunchConfig,
gemm_config: cudarc::driver::LaunchConfig,
scratch: CUptr,
scratch_elements: usize,
transform: Box<ResolvedInputTransform>,
},
}
pub(in crate::mamba_ssm::gpu) struct PreparedF32TriadLaunch {
context_token: u64,
stream_token: usize,
request: F32TriadRequest,
operands: F32TriadOperands,
resources: F32LaunchResourceSnapshot,
managed_epoch: Option<ManagedAllocationEpochStamp>,
routes: Box<[ResolvedGemmRoute]>,
resolved_launch_set: ResolvedGemmLaunchSet,
kind: PreparedF32Kind,
}
impl PreparedF32TriadLaunch {
pub(in crate::mamba_ssm::gpu) fn physical_graph_request(&self) -> F32TriadRequest {
self.request
}
pub(in crate::mamba_ssm::gpu) fn physical_graph_operands(&self) -> F32TriadOperands {
self.operands
}
pub(in crate::mamba_ssm::gpu) fn physical_graph_is_direct(&self) -> bool {
matches!(
self.kind,
PreparedF32Kind::ScalarZero { .. }
| PreparedF32Kind::Tf32 { .. }
| PreparedF32Kind::Tf32SplitK { .. }
| PreparedF32Kind::Tf32StreamK { .. }
| PreparedF32Kind::Tf32TnPreRna { .. }
)
}
pub(in crate::mamba_ssm::gpu) fn physical_graph_launch_count(&self) -> usize {
match self.kind {
PreparedF32Kind::Tf32TnPreRna { .. } => 2,
_ => self.routes.len(),
}
}
pub(in crate::mamba_ssm::gpu) fn physical_graph_scratch_range(
&self,
) -> Option<PhysicalArgumentRange> {
match self.kind {
PreparedF32Kind::Tf32TnPreRna {
scratch,
scratch_elements,
..
} => Some(PhysicalArgumentRange {
pointer: scratch,
required_bytes: u64::try_from(scratch_elements).ok()?.checked_mul(4)?,
}),
_ => None,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
enum F32PreparedSelection {
Automatic,
ExactScalar,
Forced(Tf32PhysicalRoute),
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
struct PreparedF32Key {
context_token: u64,
policy: GemmPolicy,
selection: F32PreparedSelection,
request: F32TriadRequest,
output: CUptr,
a: CUptr,
b: CUptr,
bias: Option<CUptr>,
alpha_bits: u32,
beta_bits: u32,
}
impl PreparedF32Key {
fn new(
context_token: u64,
policy: GemmPolicy,
selection: F32PreparedSelection,
request: F32TriadRequest,
operands: F32TriadOperands,
) -> Self {
Self {
context_token,
policy,
selection,
request,
output: operands.output,
a: operands.a,
b: operands.b,
bias: operands.bias,
alpha_bits: operands.alpha.to_bits(),
beta_bits: operands.beta.to_bits(),
}
}
}
#[derive(Default)]
pub(crate) struct F32PreparedLaunchCache {
entries: HashMap<PreparedF32Key, Box<PreparedF32TriadLaunch>>,
}
const F32_PREPARED_CACHE_LIMIT: usize = 1024;
const SM120_PREPARED_CACHE_LIMIT: usize = 1024;
fn make_room_in_bounded_cache<K, V>(
entries: &mut HashMap<K, V>,
incoming: &K,
limit: usize,
mut keep: impl FnMut(&V) -> bool,
) where
K: Eq + std::hash::Hash,
{
entries.retain(|_, value| keep(value));
if !entries.contains_key(incoming) && entries.len() >= limit {
entries.clear();
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct Sm120PreparedKey {
context: GemmRouteIdentity,
route: Sm120ForcedRoute,
a: CUptr,
b: CUptr,
output: CUptr,
bias: CUptr,
alpha_bits: u32,
beta_bits: u32,
}
impl std::hash::Hash for Sm120PreparedKey {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.context.hash(state);
self.route.op.hash(state);
match self.route.dtype {
WeightDtype::F32 | WeightDtype::Tf32 => 0_u8,
WeightDtype::F16 => 1,
WeightDtype::Bf16 => 2,
}
.hash(state);
self.route.physical.hash(state);
self.route.shape.hash(state);
self.a.hash(state);
self.b.hash(state);
self.output.hash(state);
self.bias.hash(state);
self.alpha_bits.hash(state);
self.beta_bits.hash(state);
}
}
impl Sm120PreparedKey {
fn new(context: GemmRouteIdentity, route: Sm120ForcedRoute, request: Sm120AutoRequest) -> Self {
Self {
context,
route,
a: request.a_ptr,
b: request.b_ptr,
output: request.operands.output_ptr,
bias: request.operands.bias_ptr,
alpha_bits: request.operands.alpha.to_bits(),
beta_bits: request.operands.beta.to_bits(),
}
}
}
struct Sm120PreparedCacheEntry {
prepared: Box<Sm120PreparedLaunch>,
managed_epoch: Option<ManagedAllocationEpochStamp>,
}
#[derive(Default)]
pub(crate) struct Sm120PreparedLaunchCache {
entries: HashMap<Sm120PreparedKey, Sm120PreparedCacheEntry>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Sm120ManagedEpochState {
Missing,
Current,
Stale,
Untracked,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum Sm120CacheAction {
UsePrepared,
Validate,
Prepare,
CaptureMissing,
CaptureStale,
CaptureUntracked,
}
fn sm120_cache_action(
capturing: bool,
hit: bool,
epoch: Sm120ManagedEpochState,
) -> Sm120CacheAction {
if !hit {
return if capturing {
Sm120CacheAction::CaptureMissing
} else {
Sm120CacheAction::Prepare
};
}
if !capturing {
return Sm120CacheAction::Validate;
}
match epoch {
Sm120ManagedEpochState::Current => Sm120CacheAction::UsePrepared,
Sm120ManagedEpochState::Stale => Sm120CacheAction::CaptureStale,
Sm120ManagedEpochState::Untracked | Sm120ManagedEpochState::Missing => {
Sm120CacheAction::CaptureUntracked
}
}
}
fn sm120_capture_cache_error(action: Sm120CacheAction) -> &'static str {
match action {
Sm120CacheAction::CaptureMissing => {
"prepared SM120 Triad cache entry is missing during graph capture; run eager warmup again"
}
Sm120CacheAction::CaptureStale => {
"prepared SM120 Triad allocation epoch changed during graph capture; run eager warmup again"
}
Sm120CacheAction::CaptureUntracked => {
"prepared SM120 Triad automatic capture requires managed allocations; run eager warmup again"
}
Sm120CacheAction::UsePrepared | Sm120CacheAction::Validate | Sm120CacheAction::Prepare => {
"SM120 cache action is not a capture error"
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(in crate::mamba_ssm::gpu) struct Sm120AutoBranchSeal {
pub(in crate::mamba_ssm::gpu) route: Sm120ForcedRoute,
}
impl Sm120PreparedLaunchCache {
fn managed_epoch_state(entry: Option<&Sm120PreparedCacheEntry>) -> Sm120ManagedEpochState {
match entry {
None => Sm120ManagedEpochState::Missing,
Some(entry) => match entry.managed_epoch.as_ref() {
Some(epoch) if epoch.is_current() => Sm120ManagedEpochState::Current,
Some(_) => Sm120ManagedEpochState::Stale,
None => Sm120ManagedEpochState::Untracked,
},
}
}
fn ensure_sm120_prepared(
&mut self,
ctx: &GpuCtx,
key: Sm120PreparedKey,
route: Sm120ForcedRoute,
request: Sm120AutoRequest,
) -> Result<&Sm120PreparedLaunch, String> {
let capturing = ctx
.stream
.capture_status()
.map_err(|error| format!("query SM120 TMA capture status: {error:?}"))?
!= cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE;
let epoch = Self::managed_epoch_state(self.entries.get(&key));
let action = sm120_cache_action(capturing, self.entries.contains_key(&key), epoch);
match action {
Sm120CacheAction::UsePrepared => {
return Ok(&self.entries.get(&key).expect("cache hit above").prepared);
}
Sm120CacheAction::CaptureMissing => {
return Err(sm120_capture_cache_error(action).into());
}
Sm120CacheAction::CaptureStale | Sm120CacheAction::CaptureUntracked => {
return Err(sm120_capture_cache_error(action).into());
}
Sm120CacheAction::Validate => {
let validation = validate_sm120_graph_replay(
&ctx.stream,
&ctx.kernels,
&self.entries.get(&key).expect("cache hit above").prepared,
);
if validation.is_ok() {
return Ok(&self
.entries
.get(&key)
.expect("validated cache hit")
.prepared);
}
self.entries.remove(&key);
}
Sm120CacheAction::Prepare => {}
}
let maps = prepare_sm120_tensor_maps(
&ctx.stream,
&ctx.kernels,
Sm120MapRequest {
op: route.op,
dtype: route.dtype,
tile: route.physical.tile,
bk: route.physical.bk,
a_ptr: request.a_ptr,
b_ptr: request.b_ptr,
shape: route.shape,
},
)?;
let prepared =
prepare_sm120_tma_forced(&ctx.stream, &ctx.kernels, route, &maps, request.operands)?;
let entry = Sm120PreparedCacheEntry {
managed_epoch: prepared.managed_epoch(),
prepared: Box::new(prepared),
};
make_room_in_bounded_cache(
&mut self.entries,
&key,
SM120_PREPARED_CACHE_LIMIT,
|cached| {
cached
.managed_epoch
.as_ref()
.is_none_or(ManagedAllocationEpochStamp::is_current)
},
);
if !self.entries.contains_key(&key) {
self.entries
.try_reserve(1)
.map_err(|error| format!("reserve prepared SM120 TMA cache: {error}"))?;
}
self.entries.insert(key, entry);
Ok(&self
.entries
.get(&key)
.expect("prepared SM120 TMA cache entry was inserted above")
.prepared)
}
}
const SM100_PREPARED_CACHE_LIMIT: usize = 1024;
#[derive(Clone, Copy, Debug, PartialEq)]
struct Sm100PreparedKey {
context: GemmRouteIdentity,
route: Sm100ForcedRoute,
a: CUptr,
b: CUptr,
output: CUptr,
bias: CUptr,
alpha_bits: u32,
beta_bits: u32,
}
impl Eq for Sm100PreparedKey {}
impl std::hash::Hash for Sm100PreparedKey {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.context.hash(state);
self.route.op.hash(state);
match self.route.dtype {
WeightDtype::F32 | WeightDtype::Tf32 => 0_u8,
WeightDtype::F16 => 1,
WeightDtype::Bf16 => 2,
}
.hash(state);
self.route.physical.hash(state);
self.route.shape.hash(state);
self.a.hash(state);
self.b.hash(state);
self.output.hash(state);
self.bias.hash(state);
self.alpha_bits.hash(state);
self.beta_bits.hash(state);
}
}
impl Sm100PreparedKey {
fn new(context: GemmRouteIdentity, route: Sm100ForcedRoute, request: Sm100AutoRequest) -> Self {
Self {
context,
route,
a: request.a_ptr,
b: request.b_ptr,
output: request.operands.output_ptr,
bias: request.operands.bias_ptr,
alpha_bits: request.operands.alpha.to_bits(),
beta_bits: request.operands.beta.to_bits(),
}
}
}
struct Sm100PreparedCacheEntry {
prepared: Box<Sm100PreparedLaunch>,
managed_epoch: Option<ManagedAllocationEpochStamp>,
}
#[derive(Default)]
pub(crate) struct Sm100PreparedLaunchCache {
entries: HashMap<Sm100PreparedKey, Sm100PreparedCacheEntry>,
}
fn sm100_capture_cache_error(action: Sm120CacheAction) -> &'static str {
match action {
Sm120CacheAction::CaptureMissing => {
"prepared SM100 Triad cache entry is missing during graph capture; run eager warmup again"
}
Sm120CacheAction::CaptureStale => {
"prepared SM100 Triad allocation epoch changed during graph capture; run eager warmup again"
}
Sm120CacheAction::CaptureUntracked => {
"prepared SM100 Triad automatic capture requires managed allocations; run eager warmup again"
}
Sm120CacheAction::UsePrepared | Sm120CacheAction::Validate | Sm120CacheAction::Prepare => {
"SM100 cache action is not a capture error"
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(in crate::mamba_ssm::gpu) struct Sm100AutoBranchSeal {
pub(in crate::mamba_ssm::gpu) route: Sm100ForcedRoute,
}
impl Sm100PreparedLaunchCache {
fn managed_epoch_state(entry: Option<&Sm100PreparedCacheEntry>) -> Sm120ManagedEpochState {
match entry {
None => Sm120ManagedEpochState::Missing,
Some(entry) => match entry.managed_epoch.as_ref() {
Some(epoch) if epoch.is_current() => Sm120ManagedEpochState::Current,
Some(_) => Sm120ManagedEpochState::Stale,
None => Sm120ManagedEpochState::Untracked,
},
}
}
fn ensure_sm100_prepared(
&mut self,
ctx: &GpuCtx,
key: Sm100PreparedKey,
route: Sm100ForcedRoute,
request: Sm100AutoRequest,
) -> Result<&Sm100PreparedLaunch, String> {
let capturing = ctx
.stream
.capture_status()
.map_err(|error| format!("query SM100 TCGEN capture status: {error:?}"))?
!= cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE;
let epoch = Self::managed_epoch_state(self.entries.get(&key));
let action = sm120_cache_action(capturing, self.entries.contains_key(&key), epoch);
match action {
Sm120CacheAction::UsePrepared => {
return Ok(&self.entries.get(&key).expect("cache hit above").prepared);
}
Sm120CacheAction::CaptureMissing
| Sm120CacheAction::CaptureStale
| Sm120CacheAction::CaptureUntracked => {
return Err(sm100_capture_cache_error(action).into());
}
Sm120CacheAction::Validate => {
let validation = validate_sm100_graph_replay(
&ctx.stream,
&ctx.kernels,
&self.entries.get(&key).expect("cache hit above").prepared,
);
if validation.is_ok() {
return Ok(&self
.entries
.get(&key)
.expect("validated cache hit")
.prepared);
}
self.entries.remove(&key);
}
Sm120CacheAction::Prepare => {}
}
let maps = prepare_sm100_tensor_maps(
&ctx.stream,
&ctx.kernels,
Sm100MapRequest {
op: route.op,
dtype: route.dtype,
tile: route.physical.tile,
a_ptr: request.a_ptr,
b_ptr: request.b_ptr,
shape: route.shape,
},
)?;
let prepared =
prepare_sm100_tcgen_forced(&ctx.stream, &ctx.kernels, route, &maps, request.operands)?;
let entry = Sm100PreparedCacheEntry {
managed_epoch: prepared.managed_epoch(),
prepared: Box::new(prepared),
};
make_room_in_bounded_cache(
&mut self.entries,
&key,
SM100_PREPARED_CACHE_LIMIT,
|cached| {
cached
.managed_epoch
.as_ref()
.is_none_or(ManagedAllocationEpochStamp::is_current)
},
);
if !self.entries.contains_key(&key) {
self.entries
.try_reserve(1)
.map_err(|error| format!("reserve prepared SM100 TCGEN cache: {error}"))?;
}
self.entries.insert(key, entry);
Ok(&self
.entries
.get(&key)
.expect("prepared SM100 TCGEN cache entry was inserted above")
.prepared)
}
}
fn sm100_policy_dtype(dtype: WeightDtype) -> Result<PolicyDtype, String> {
match dtype {
WeightDtype::Bf16 => Ok(PolicyDtype::Bf16),
WeightDtype::F16 => Ok(PolicyDtype::F16),
WeightDtype::F32 | WeightDtype::Tf32 => {
Err("SM100 automatic route requires BF16 or F16".into())
}
}
}
unsafe fn enqueue_sm100_tcgen_prepared_observed<O: PhysicalLaunchObserver>(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
prepared: &Sm100PreparedLaunch,
observer: &mut O,
observation: Option<PhysicalLaunchObservation>,
) -> Result<(), String> {
validate_sm100_prepared_binding(stream, kernels, prepared)?;
let spec = prepared.route.kernel_spec()?;
if spec.symbol != prepared.identity.symbol {
return Err("SM100 prepared symbol no longer matches its physical route".into());
}
let function = kernels
.sm100_function(spec.symbol)
.ok_or_else(|| format!("SM100 kernel {} is unavailable", spec.symbol))?;
let (rows, columns) = match prepared.route.op {
Sm100Op::Nn => (prepared.route.shape.m, prepared.route.shape.n),
Sm100Op::Tn => (prepared.route.shape.k, prepared.route.shape.n),
Sm100Op::Nt => (prepared.route.shape.m, prepared.route.shape.k),
};
let rows = checked_u32(rows, "SM100 output rows")?;
let columns = checked_u32(columns, "SM100 output columns")?;
let grid = checked_grid_product(
rows.div_ceil(prepared.route.physical.tile.output_rows()),
columns.div_ceil(prepared.route.physical.tile.output_columns()),
1,
)?;
let config = cudarc::driver::LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (spec.threads, 1, 1),
shared_mem_bytes: spec.dynamic_shared_bytes,
};
let mut builder = stream.launch_builder(function);
builder.arg(&prepared.operands.output_ptr);
builder.arg(&prepared.maps.a);
builder.arg(&prepared.maps.b);
builder.arg(&prepared.operands.bias_ptr);
let params = Sm100KernelParams::from_words(prepared.params);
builder.arg(¶ms);
unsafe { enqueue_with_physical_observation(observer, &mut builder, config, observation) }
.map_err(|error| error.with_driver_context(format_args!("launch {}", spec.symbol)))
}
pub(in crate::mamba_ssm::gpu) fn launch_sm100_auto_observed<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
request: Sm100AutoRequest,
) -> Result<Option<Sm100AutoBranchSeal>, String> {
let Some(target) = ctx.kernels.sm100_target_candidate() else {
return Ok(None);
};
let Some(route) = resolve_sm100_auto(target.device_cc, Some(target), request) else {
static NO_CELL: std::sync::Once = std::sync::Once::new();
crate::mamba_ssm::gpu::diagnostics::warn_once(&NO_CELL, || {
format!(
"no measured SM100 cell for {:?} {:?} {:?}; the portable tensor-core tiles serve \
it (reported once; later uncovered shapes are silent)",
request.op, request.dtype, request.shape
)
});
return Ok(None);
};
let key = Sm100PreparedKey::new(ctx.gemm_route(), route, request);
let caps = query_specialized_device_caps(
&ctx.stream,
target.nvrtc_arch,
crate::mamba_ssm::gpu::kernels::nvrtc_version(),
)?;
ctx.with_sm100_prepared_launches(|cache| {
let prepared = cache.ensure_sm100_prepared(ctx, key, route, request)?;
let resolved = prepared.identity().resolved_route(caps)?;
unsafe {
enqueue_sm100_tcgen_prepared_observed(
&ctx.stream,
&ctx.kernels,
prepared,
observer,
Some(PhysicalLaunchObservation::gemm(
sm100_policy_dtype(route.dtype)?,
None,
resolved,
)),
)
}?;
ctx.record_resolved_gemm_route(resolved)?;
Ok(Some(Sm100AutoBranchSeal { route }))
})
}
pub(in crate::mamba_ssm::gpu) fn prepare_sm100_auto_graph_sequence<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &O,
request: Sm100AutoRequest,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
let target = ctx
.kernels
.sm100_target_candidate()
.ok_or_else(|| "prepared SM100 graph route has no module target".to_string())?;
let route = resolve_sm100_auto(target.device_cc, Some(target), request)
.ok_or_else(|| "prepared SM100 graph request is no longer qualified".to_string())?;
let key = Sm100PreparedKey::new(ctx.gemm_route(), route, request);
let caps = query_specialized_device_caps(
&ctx.stream,
target.nvrtc_arch,
crate::mamba_ssm::gpu::kernels::nvrtc_version(),
)?;
ctx.with_sm100_prepared_launches(|cache| {
let prepared = &cache
.entries
.get(&key)
.ok_or_else(|| {
"prepared SM100 graph cache entry is missing; run eager warmup again".to_string()
})?
.prepared;
validate_sm100_graph_replay(&ctx.stream, &ctx.kernels, prepared)?;
let resolved = prepared.identity().resolved_route(caps)?;
let config = LaunchConfig {
grid_dim: resolved.launch.grid_dim,
block_dim: resolved.launch.block_dim,
shared_mem_bytes: resolved.launch.shared_mem_bytes,
};
let observation =
PhysicalLaunchObservation::gemm(sm100_policy_dtype(route.dtype)?, None, resolved);
let node = resolve_physical_launch_observation(observer, observation, config)?;
let function = ctx
.kernels
.sm100_function(resolved.symbol)
.ok_or_else(|| format!("qualified SM100 symbol {} is unavailable", resolved.symbol))?
.clone();
let mut arguments = PhysicalScalarKernelArguments::new();
arguments.push(prepared.operands.output_ptr)?;
arguments.push(prepared.maps.a)?;
arguments.push(prepared.maps.b)?;
arguments.push(prepared.operands.bias_ptr)?;
arguments.push(Sm100KernelParams::from_words(prepared.params))?;
validate_sm100_graph_replay(&ctx.stream, &ctx.kernels, prepared)?;
Ok(PreparedTriadPhysicalGraphSequence {
launches: vec![PreparedTriadPhysicalGraphLaunch {
function,
config,
node,
arguments: Box::new(arguments),
}]
.into_boxed_slice(),
})
})
}
const SM90A_PREPARED_CACHE_LIMIT: usize = 1024;
#[derive(Clone, Copy, Debug, PartialEq)]
struct Sm90aPreparedKey {
context: GemmRouteIdentity,
route: Sm90aForcedRoute,
a: CUptr,
b: CUptr,
output: CUptr,
bias: CUptr,
alpha_bits: u32,
beta_bits: u32,
}
impl Eq for Sm90aPreparedKey {}
impl std::hash::Hash for Sm90aPreparedKey {
fn hash<H: std::hash::Hasher>(&self, state: &mut H) {
self.context.hash(state);
self.route.op.hash(state);
match self.route.dtype {
WeightDtype::F32 | WeightDtype::Tf32 => 0_u8,
WeightDtype::F16 => 1,
WeightDtype::Bf16 => 2,
}
.hash(state);
self.route.schedule.hash(state);
self.route.shape.hash(state);
self.a.hash(state);
self.b.hash(state);
self.output.hash(state);
self.bias.hash(state);
self.alpha_bits.hash(state);
self.beta_bits.hash(state);
}
}
impl Sm90aPreparedKey {
fn new(context: GemmRouteIdentity, route: Sm90aForcedRoute, request: Sm90aAutoRequest) -> Self {
Self {
context,
route,
a: request.a_ptr,
b: request.b_ptr,
output: request.operands.output_ptr,
bias: request.operands.bias_ptr,
alpha_bits: request.operands.alpha.to_bits(),
beta_bits: request.operands.beta.to_bits(),
}
}
}
struct Sm90aPreparedCacheEntry {
maps: Box<Sm90aPreparedTensorMaps>,
identity: Sm90aRouteIdentity,
route: Sm90aForcedRoute,
operands: Sm90aLaunchOperands,
managed_epoch: Option<ManagedAllocationEpochStamp>,
}
#[derive(Default)]
pub(crate) struct Sm90aPreparedLaunchCache {
entries: HashMap<Sm90aPreparedKey, Sm90aPreparedCacheEntry>,
}
fn sm90a_capture_cache_error(action: Sm120CacheAction) -> &'static str {
match action {
Sm120CacheAction::CaptureMissing => {
"prepared SM90a Triad cache entry is missing during graph capture; run eager warmup again"
}
Sm120CacheAction::CaptureStale => {
"prepared SM90a Triad allocation epoch changed during graph capture; run eager warmup again"
}
Sm120CacheAction::CaptureUntracked => {
"prepared SM90a Triad automatic capture requires managed allocations; run eager warmup again"
}
Sm120CacheAction::UsePrepared | Sm120CacheAction::Validate | Sm120CacheAction::Prepare => {
"SM90a cache action is not a capture error"
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(in crate::mamba_ssm::gpu) struct Sm90aAutoBranchSeal {
pub(in crate::mamba_ssm::gpu) route: Sm90aForcedRoute,
}
impl Sm90aPreparedLaunchCache {
fn managed_epoch_state(entry: Option<&Sm90aPreparedCacheEntry>) -> Sm120ManagedEpochState {
match entry {
None => Sm120ManagedEpochState::Missing,
Some(entry) => match entry.managed_epoch.as_ref() {
Some(epoch) if epoch.is_current() => Sm120ManagedEpochState::Current,
Some(_) => Sm120ManagedEpochState::Stale,
None => Sm120ManagedEpochState::Untracked,
},
}
}
fn ensure_sm90a_prepared(
&mut self,
ctx: &GpuCtx,
key: Sm90aPreparedKey,
route: Sm90aForcedRoute,
request: Sm90aAutoRequest,
) -> Result<&Sm90aPreparedCacheEntry, String> {
let capturing = ctx
.stream
.capture_status()
.map_err(|error| format!("query SM90a WGMMA capture status: {error:?}"))?
!= cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE;
let epoch = Self::managed_epoch_state(self.entries.get(&key));
let action = sm120_cache_action(capturing, self.entries.contains_key(&key), epoch);
match action {
Sm120CacheAction::UsePrepared => {
return Ok(self.entries.get(&key).expect("cache hit above"));
}
Sm120CacheAction::CaptureMissing
| Sm120CacheAction::CaptureStale
| Sm120CacheAction::CaptureUntracked => {
return Err(sm90a_capture_cache_error(action).into());
}
Sm120CacheAction::Validate => {
let entry = self.entries.get(&key).expect("cache hit above");
let validation = validate_sm90a_graph_replay(
&ctx.stream,
&ctx.kernels,
entry.route,
&entry.maps,
entry.operands,
entry.identity,
);
if validation.is_ok() {
return Ok(self.entries.get(&key).expect("validated cache hit"));
}
self.entries.remove(&key);
}
Sm120CacheAction::Prepare => {}
}
let maps = prepare_sm90a_tensor_maps(
&ctx.stream,
&ctx.kernels,
Sm90aMapRequest {
op: route.op,
dtype: route.dtype,
a_ptr: request.a_ptr,
b_ptr: request.b_ptr,
shape: route.shape,
},
)?;
let identity =
sm90a_forced_identity(&ctx.stream, &ctx.kernels, route, &maps, request.operands)?;
let managed_epoch = maps.managed_epoch(route, request.operands)?;
let entry = Sm90aPreparedCacheEntry {
maps: Box::new(maps),
identity,
route,
operands: request.operands,
managed_epoch,
};
make_room_in_bounded_cache(
&mut self.entries,
&key,
SM90A_PREPARED_CACHE_LIMIT,
|cached| {
cached
.managed_epoch
.as_ref()
.is_none_or(ManagedAllocationEpochStamp::is_current)
},
);
if !self.entries.contains_key(&key) {
self.entries
.try_reserve(1)
.map_err(|error| format!("reserve prepared SM90a WGMMA cache: {error}"))?;
}
self.entries.insert(key, entry);
Ok(self
.entries
.get(&key)
.expect("prepared SM90a WGMMA cache entry was inserted above"))
}
}
fn sm90a_policy_dtype(dtype: WeightDtype) -> Result<PolicyDtype, String> {
match dtype {
WeightDtype::Bf16 => Ok(PolicyDtype::Bf16),
WeightDtype::F16 => Ok(PolicyDtype::F16),
WeightDtype::F32 | WeightDtype::Tf32 => {
Err("SM90a automatic route requires BF16 or F16".into())
}
}
}
fn sm90a_launch_config(route: Sm90aForcedRoute) -> Result<cudarc::driver::LaunchConfig, String> {
let (rows, columns) = match route.op {
Sm90aOp::Nn => (route.shape.m, route.shape.n),
Sm90aOp::Tn => (route.shape.k, route.shape.n),
Sm90aOp::Nt => (route.shape.m, route.shape.k),
};
let rows = checked_u32(rows, "SM90a output rows")?;
let columns = checked_u32(columns, "SM90a output columns")?;
let grid = checked_grid_product(
rows.div_ceil(SM90A_TILE.0),
columns.div_ceil(SM90A_TILE.1),
1,
)?;
Ok(cudarc::driver::LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (route.schedule.threads(), 1, 1),
shared_mem_bytes: SM90A_DYNAMIC_SHARED_BYTES,
})
}
unsafe fn enqueue_sm90a_wgmma_prepared_observed<O: PhysicalLaunchObserver>(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
entry: &Sm90aPreparedCacheEntry,
observer: &mut O,
observation: Option<PhysicalLaunchObservation>,
) -> Result<(), String> {
let live = sm90a_forced_identity(stream, kernels, entry.route, &entry.maps, entry.operands)?;
entry
.identity
.ensure_current(live, "SM90a automatic launch")?;
let symbol = entry.route.symbol();
let function = kernels
.sm90a_function(symbol)
.ok_or_else(|| format!("SM90a kernel {symbol} is unavailable"))?;
let config = sm90a_launch_config(entry.route)?;
let m = checked_i32(entry.route.shape.m, "M")?;
let k = checked_i32(entry.route.shape.k, "K")?;
let n = checked_i32(entry.route.shape.n, "N")?;
let ldc = checked_i32(entry.route.shape.ldc, "ldc")?;
let mut builder = stream.launch_builder(function);
builder.arg(&entry.operands.output_ptr);
builder.arg(&entry.maps.a);
builder.arg(&entry.maps.b);
builder.arg(&entry.operands.bias_ptr);
builder.arg(&entry.operands.alpha);
builder.arg(&entry.operands.beta);
builder.arg(&m);
builder.arg(&k);
builder.arg(&n);
builder.arg(&ldc);
unsafe { enqueue_with_physical_observation(observer, &mut builder, config, observation) }
.map_err(|error| error.with_driver_context(format_args!("launch {symbol}")))
}
pub(in crate::mamba_ssm::gpu) fn launch_sm90a_auto_observed<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
request: Sm90aAutoRequest,
) -> Result<Option<Sm90aAutoBranchSeal>, String> {
if !ctx.kernels.has_sm90a_wgmma() {
return Ok(None);
}
let device_cc = ctx
.stream
.context()
.compute_capability()
.map_err(|error| format!("query compute capability for the SM90a branch: {error:?}"))?;
let Some(route) = resolve_sm90a_auto(device_cc, true, request) else {
static NO_ROUTE: std::sync::Once = std::sync::Once::new();
crate::mamba_ssm::gpu::diagnostics::warn_once(&NO_ROUTE, || {
format!(
"no SM90a route for {:?} {:?} {:?}; the portable tensor-core tiles serve it \
(reported once; later unserved shapes are silent)",
request.op, request.dtype, request.shape
)
});
return Ok(None);
};
let key = Sm90aPreparedKey::new(ctx.gemm_route(), route, request);
let caps = query_specialized_device_caps(
&ctx.stream,
"sm_90a",
crate::mamba_ssm::gpu::kernels::nvrtc_version(),
)?;
ctx.with_sm90a_prepared_launches(|cache| {
let entry = cache.ensure_sm90a_prepared(ctx, key, route, request)?;
let resolved = entry.identity.resolved_route(caps)?;
unsafe {
enqueue_sm90a_wgmma_prepared_observed(
&ctx.stream,
&ctx.kernels,
entry,
observer,
Some(PhysicalLaunchObservation::gemm(
sm90a_policy_dtype(route.dtype)?,
None,
resolved,
)),
)
}?;
ctx.record_resolved_gemm_route(resolved)?;
Ok(Some(Sm90aAutoBranchSeal { route }))
})
}
pub(in crate::mamba_ssm::gpu) fn prepare_sm90a_auto_graph_sequence<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &O,
request: Sm90aAutoRequest,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
if !ctx.kernels.has_sm90a_wgmma() {
return Err("prepared SM90a graph route has no module".into());
}
let device_cc = ctx
.stream
.context()
.compute_capability()
.map_err(|error| format!("query compute capability for the SM90a graph: {error:?}"))?;
let route = resolve_sm90a_auto(device_cc, true, request)
.ok_or_else(|| "prepared SM90a graph request is no longer qualified".to_string())?;
let key = Sm90aPreparedKey::new(ctx.gemm_route(), route, request);
let caps = query_specialized_device_caps(
&ctx.stream,
"sm_90a",
crate::mamba_ssm::gpu::kernels::nvrtc_version(),
)?;
ctx.with_sm90a_prepared_launches(|cache| {
let entry = cache.entries.get(&key).ok_or_else(|| {
"prepared SM90a graph cache entry is missing; run eager warmup again".to_string()
})?;
validate_sm90a_graph_replay(
&ctx.stream,
&ctx.kernels,
entry.route,
&entry.maps,
entry.operands,
entry.identity,
)?;
let resolved = entry.identity.resolved_route(caps)?;
let config = LaunchConfig {
grid_dim: resolved.launch.grid_dim,
block_dim: resolved.launch.block_dim,
shared_mem_bytes: resolved.launch.shared_mem_bytes,
};
let observation =
PhysicalLaunchObservation::gemm(sm90a_policy_dtype(route.dtype)?, None, resolved);
let node = resolve_physical_launch_observation(observer, observation, config)?;
let function = ctx
.kernels
.sm90a_function(resolved.symbol)
.ok_or_else(|| format!("qualified SM90a symbol {} is unavailable", resolved.symbol))?
.clone();
let mut arguments = PhysicalScalarKernelArguments::new();
arguments.push(entry.operands.output_ptr)?;
arguments.push(entry.maps.a)?;
arguments.push(entry.maps.b)?;
arguments.push(entry.operands.bias_ptr)?;
arguments.push(entry.operands.alpha)?;
arguments.push(entry.operands.beta)?;
arguments.push(checked_i32(entry.route.shape.m, "M")?)?;
arguments.push(checked_i32(entry.route.shape.k, "K")?)?;
arguments.push(checked_i32(entry.route.shape.n, "N")?)?;
arguments.push(checked_i32(entry.route.shape.ldc, "ldc")?)?;
Ok(PreparedTriadPhysicalGraphSequence {
launches: vec![PreparedTriadPhysicalGraphLaunch {
function,
config,
node,
arguments: Box::new(arguments),
}]
.into_boxed_slice(),
})
})
}
pub(in crate::mamba_ssm::gpu) struct ScalarLaunchControl<'a> {
ctx: &'a GpuCtx,
routes: &'a [ResolvedGemmRoute],
plan: ScalarDispatchPlan,
operands: F32TriadOperands,
next: usize,
}
fn validate_prepared_scalar_binding(
expected_symbol: &str,
expected: ResolvedKernelLaunch,
symbol: &str,
config: cudarc::driver::LaunchConfig,
) -> Result<Sha256Digest, String> {
if expected.arguments_digest == [0; 32] {
return Err(format!(
"prepared scalar kernel {expected_symbol} has a zero argument digest"
));
}
if expected_symbol != symbol
|| expected.grid_dim != config.grid_dim
|| expected.block_dim != config.block_dim
|| expected.shared_mem_bytes != config.shared_mem_bytes
{
return Err(format!(
"prepared scalar kernel {expected_symbol} physical launch changed before {symbol}"
));
}
Ok(expected.arguments_digest)
}
impl ScalarLaunchControl<'_> {
fn bind(
&mut self,
symbol: &str,
config: cudarc::driver::LaunchConfig,
) -> Result<ResolvedGemmRoute, String> {
let route = self
.routes
.get(self.next)
.ok_or_else(|| format!("unexpected scalar kernel {symbol} after prepared route end"))?;
let _frozen_arguments_digest =
validate_prepared_scalar_binding(route.symbol, route.launch, symbol, config)?;
self.ctx.record_resolved_gemm_route(*route)?;
self.next += 1;
Ok(*route)
}
fn finish(&self) -> Result<(), String> {
if self.next == self.routes.len() {
Ok(())
} else {
Err(format!(
"prepared scalar route expected {} kernels but enqueued {}",
self.routes.len(),
self.next
))
}
}
pub(super) fn operands(&self) -> F32TriadOperands {
self.operands
}
pub(super) fn plan(&self) -> ScalarDispatchPlan {
self.plan
}
fn validate_operands(&self, actual: F32TriadOperands) -> Result<(), String> {
let expected = self.operands;
if actual.output != expected.output
|| actual.a != expected.a
|| actual.b != expected.b
|| actual.bias != expected.bias
|| actual.alpha.to_bits() != expected.alpha.to_bits()
|| actual.beta.to_bits() != expected.beta.to_bits()
{
return Err("scalar operands changed after f32 Triad preparation".into());
}
Ok(())
}
}
trait ScalarLaunchController {
fn enqueue(
&mut self,
symbol: &'static str,
config: LaunchConfig,
builder: &mut ScalarLaunchArgs<'_>,
) -> Result<(), PhysicalCudaLaunchError>;
fn plan(&self) -> ScalarDispatchPlan;
fn operands(&self) -> F32TriadOperands;
fn validate_operands(&self, actual: F32TriadOperands) -> Result<(), String>;
fn prepares_physical_graph(&self) -> bool {
false
}
}
impl ScalarLaunchController for ScalarLaunchControl<'_> {
#[inline(always)]
fn enqueue(
&mut self,
symbol: &'static str,
config: LaunchConfig,
builder: &mut ScalarLaunchArgs<'_>,
) -> Result<(), PhysicalCudaLaunchError> {
self.bind(symbol, config)
.map_err(PhysicalCudaLaunchError::from)?;
let mut observer = NoPhysicalObserver;
unsafe {
enqueue_with_physical_observation(&mut observer, builder.launch_args(), config, None)
}
}
fn plan(&self) -> ScalarDispatchPlan {
ScalarLaunchControl::plan(self)
}
fn operands(&self) -> F32TriadOperands {
ScalarLaunchControl::operands(self)
}
fn validate_operands(&self, actual: F32TriadOperands) -> Result<(), String> {
ScalarLaunchControl::validate_operands(self, actual)
}
}
struct PhysicalScalarLaunchControl<'a, O> {
base: ScalarLaunchControl<'a>,
logical_dtype: PolicyDtype,
physical_resources_digest: Sha256Digest,
observer: &'a mut O,
}
impl<O: PhysicalLaunchObserver> PhysicalScalarLaunchControl<'_, O> {
fn finish(&self) -> Result<(), String> {
self.base.finish()
}
}
impl<O: PhysicalLaunchObserver> ScalarLaunchController for PhysicalScalarLaunchControl<'_, O> {
#[inline(always)]
fn enqueue(
&mut self,
symbol: &'static str,
config: LaunchConfig,
builder: &mut ScalarLaunchArgs<'_>,
) -> Result<(), PhysicalCudaLaunchError> {
let route = self
.base
.bind(symbol, config)
.map_err(PhysicalCudaLaunchError::from)?;
unsafe {
enqueue_with_physical_observation(
self.observer,
builder.launch_args(),
config,
Some(PhysicalLaunchObservation::gemm(
self.logical_dtype,
Some(self.physical_resources_digest),
route,
)),
)
}
}
fn plan(&self) -> ScalarDispatchPlan {
self.base.plan()
}
fn operands(&self) -> F32TriadOperands {
self.base.operands()
}
fn validate_operands(&self, actual: F32TriadOperands) -> Result<(), String> {
self.base.validate_operands(actual)
}
}
const PHYSICAL_SCALAR_MAX_KERNEL_ARGUMENTS: usize = 16;
const PHYSICAL_SCALAR_MAX_ARGUMENT_BYTES: usize = 128;
#[derive(Clone, Copy)]
#[repr(C, align(16))]
struct PhysicalScalarKernelArgument {
bytes: [u8; PHYSICAL_SCALAR_MAX_ARGUMENT_BYTES],
}
unsafe impl DeviceRepr for PhysicalScalarKernelArgument {}
impl PhysicalScalarKernelArgument {
fn encode<T: Copy>(value: T) -> Result<Self, String> {
let width = std::mem::size_of::<T>();
if width > PHYSICAL_SCALAR_MAX_ARGUMENT_BYTES {
return Err(format!(
"physical scalar graph argument uses {width} bytes; maximum is {PHYSICAL_SCALAR_MAX_ARGUMENT_BYTES}"
));
}
let mut encoded = Self {
bytes: [0; PHYSICAL_SCALAR_MAX_ARGUMENT_BYTES],
};
unsafe {
std::ptr::copy_nonoverlapping(
std::ptr::from_ref(&value).cast::<u8>(),
encoded.bytes.as_mut_ptr(),
width,
);
}
Ok(encoded)
}
}
struct PhysicalScalarKernelArguments {
values: [PhysicalScalarKernelArgument; PHYSICAL_SCALAR_MAX_KERNEL_ARGUMENTS],
len: usize,
}
impl PhysicalScalarKernelArguments {
fn new() -> Self {
Self {
values: [PhysicalScalarKernelArgument {
bytes: [0; PHYSICAL_SCALAR_MAX_ARGUMENT_BYTES],
}; PHYSICAL_SCALAR_MAX_KERNEL_ARGUMENTS],
len: 0,
}
}
fn push<T: Copy>(&mut self, value: T) -> Result<(), String> {
let slot = self
.values
.get_mut(self.len)
.ok_or_else(|| "physical scalar graph argument capacity exceeded".to_string())?;
*slot = PhysicalScalarKernelArgument::encode(value)?;
self.len += 1;
Ok(())
}
fn values(&self) -> &[PhysicalScalarKernelArgument] {
&self.values[..self.len]
}
}
struct ScalarLaunchArgs<'a> {
builder: LaunchArgs<'a>,
prepared_function: Option<CudaFunction>,
prepared_arguments: Option<Box<PhysicalScalarKernelArguments>>,
preparation_error: Option<String>,
}
trait ScalarInputArgument {
fn scalar_ptr(&self) -> CUptr;
fn submit<'a>(&'a self, builder: &mut LaunchArgs<'a>);
}
trait ScalarOutputArgument {
fn scalar_ptr(&self) -> CUptr;
fn submit<'a>(&'a mut self, builder: &mut LaunchArgs<'a>);
}
struct RawScalarArgument(CUptr);
impl ScalarInputArgument for GpuBuffer {
fn scalar_ptr(&self) -> CUptr {
self.cached_ptr()
}
fn submit<'a>(&'a self, builder: &mut LaunchArgs<'a>) {
builder.arg(self.inner());
}
}
impl ScalarOutputArgument for GpuBuffer {
fn scalar_ptr(&self) -> CUptr {
self.cached_ptr()
}
fn submit<'a>(&'a mut self, builder: &mut LaunchArgs<'a>) {
builder.arg(self.inner_mut());
}
}
impl ScalarInputArgument for RawScalarArgument {
fn scalar_ptr(&self) -> CUptr {
self.0
}
fn submit<'a>(&'a self, builder: &mut LaunchArgs<'a>) {
builder.arg(&self.0);
}
}
impl ScalarOutputArgument for RawScalarArgument {
fn scalar_ptr(&self) -> CUptr {
self.0
}
fn submit<'a>(&'a mut self, builder: &mut LaunchArgs<'a>) {
builder.arg(&self.0);
}
}
impl<'a> ScalarLaunchArgs<'a> {
fn new(
stream: &'a Arc<CudaStream>,
function: &'a CudaFunction,
prepare_physical_graph: bool,
) -> Self {
Self {
builder: stream.launch_builder(function),
prepared_function: prepare_physical_graph.then(|| function.clone()),
prepared_arguments: prepare_physical_graph
.then(|| Box::new(PhysicalScalarKernelArguments::new())),
preparation_error: None,
}
}
fn capture_argument<T: Copy>(&mut self, argument: T) {
if let Some(arguments) = self.prepared_arguments.as_mut()
&& self.preparation_error.is_none()
&& let Err(error) = arguments.push(argument)
{
self.preparation_error = Some(error);
}
}
fn arg<T: Copy + DeviceRepr>(&mut self, argument: &'a T) -> &mut Self {
self.capture_argument(*argument);
self.builder.arg(argument);
self
}
fn arg_buffer<Input: ScalarInputArgument>(&mut self, buffer: &'a Input) -> &mut Self {
self.capture_argument(buffer.scalar_ptr());
buffer.submit(&mut self.builder);
self
}
fn arg_buffer_mut<Output: ScalarOutputArgument>(
&mut self,
buffer: &'a mut Output,
) -> &mut Self {
self.capture_argument(buffer.scalar_ptr());
buffer.submit(&mut self.builder);
self
}
fn launch_args(&mut self) -> &mut LaunchArgs<'a> {
&mut self.builder
}
fn take_prepared(
&mut self,
) -> Result<(CudaFunction, Box<PhysicalScalarKernelArguments>), String> {
if let Some(error) = self.preparation_error.take() {
return Err(error);
}
let function = self.prepared_function.take().ok_or_else(|| {
"scalar launch arguments were not configured for physical graph preparation".to_string()
})?;
let arguments = self
.prepared_arguments
.take()
.ok_or_else(|| "scalar launch argument storage was already consumed".to_string())?;
Ok((function, arguments))
}
}
fn scalar_launch_builder<'a, C: ScalarLaunchController>(
stream: &'a Arc<CudaStream>,
function: &'a CudaFunction,
control: &Option<&mut C>,
) -> ScalarLaunchArgs<'a> {
let prepare = control
.as_deref()
.is_some_and(ScalarLaunchController::prepares_physical_graph);
ScalarLaunchArgs::new(stream, function, prepare)
}
struct PreparedTriadPhysicalGraphLaunch {
function: CudaFunction,
config: LaunchConfig,
node: ResolvedPhysicalKernelLaunch,
arguments: Box<PhysicalScalarKernelArguments>,
}
pub(in crate::mamba_ssm::gpu) struct PreparedTriadPhysicalGraphSequence {
launches: Box<[PreparedTriadPhysicalGraphLaunch]>,
}
impl PreparedTriadPhysicalGraphSequence {
pub(in crate::mamba_ssm::gpu) fn len(&self) -> usize {
self.launches.len()
}
pub(in crate::mamba_ssm::gpu) fn bind<'a>(
&'a self,
stream: &'a Arc<CudaStream>,
) -> Result<BoundTriadPhysicalGraphSequence<'a>, String> {
let mut launches = Vec::new();
launches
.try_reserve_exact(self.launches.len())
.map_err(|error| format!("reserve bound scalar physical launches: {error}"))?;
for launch in &self.launches {
let mut builder = stream.launch_builder(&launch.function);
for argument in launch.arguments.values() {
builder.arg(argument);
}
launches.push(BoundTriadPhysicalGraphLaunch {
builder,
config: launch.config,
node: launch.node,
});
}
if launches.len() != self.launches.len() || launches.capacity() != self.launches.len() {
return Err("bound scalar physical launch backing capacity is not exact".into());
}
Ok(BoundTriadPhysicalGraphSequence { launches })
}
}
struct BoundTriadPhysicalGraphLaunch<'a> {
builder: LaunchArgs<'a>,
config: LaunchConfig,
node: ResolvedPhysicalKernelLaunch,
}
pub(in crate::mamba_ssm::gpu) struct BoundTriadPhysicalGraphSequence<'a> {
launches: Vec<BoundTriadPhysicalGraphLaunch<'a>>,
}
impl BoundTriadPhysicalGraphSequence<'_> {
#[inline(always)]
pub(in crate::mamba_ssm::gpu) unsafe fn enqueue(
&mut self,
observer: &mut RecordingPhysicalObserver,
) -> Result<(), PhysicalCudaLaunchError> {
for launch in &mut self.launches {
unsafe {
enqueue_prepared_physical_launch(
observer,
&mut launch.builder,
launch.config,
launch.node,
)?;
}
}
Ok(())
}
}
struct PreparedPhysicalScalarLaunchControl<'a, O> {
base: ScalarLaunchControl<'a>,
logical_dtype: PolicyDtype,
physical_resources_digest: Sha256Digest,
observer: &'a O,
launches: Vec<PreparedTriadPhysicalGraphLaunch>,
}
impl<O: PhysicalLaunchObserver> PreparedPhysicalScalarLaunchControl<'_, O> {
fn finish(self) -> Result<PreparedTriadPhysicalGraphSequence, String> {
self.base.finish()?;
if self.launches.len() != self.base.routes.len()
|| self.launches.capacity() != self.base.routes.len()
{
return Err("prepared scalar physical launch capacity is not exact".into());
}
Ok(PreparedTriadPhysicalGraphSequence {
launches: self.launches.into_boxed_slice(),
})
}
}
impl<O: PhysicalLaunchObserver> ScalarLaunchController
for PreparedPhysicalScalarLaunchControl<'_, O>
{
#[inline(always)]
fn enqueue(
&mut self,
symbol: &'static str,
config: LaunchConfig,
builder: &mut ScalarLaunchArgs<'_>,
) -> Result<(), PhysicalCudaLaunchError> {
let route = self
.base
.bind(symbol, config)
.map_err(PhysicalCudaLaunchError::from)?;
let observation = PhysicalLaunchObservation::gemm(
self.logical_dtype,
Some(self.physical_resources_digest),
route,
);
let node = resolve_physical_launch_observation(self.observer, observation, config)
.map_err(PhysicalCudaLaunchError::from)?;
let (function, arguments) = builder
.take_prepared()
.map_err(PhysicalCudaLaunchError::from)?;
self.launches.push(PreparedTriadPhysicalGraphLaunch {
function,
config,
node,
arguments,
});
Ok(())
}
fn plan(&self) -> ScalarDispatchPlan {
self.base.plan()
}
fn operands(&self) -> F32TriadOperands {
self.base.operands()
}
fn validate_operands(&self, actual: F32TriadOperands) -> Result<(), String> {
self.base.validate_operands(actual)
}
fn prepares_physical_graph(&self) -> bool {
true
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct ScalarNodeSpec {
symbol: &'static str,
tile: (u32, u32),
bk: u32,
stages: u8,
launch: ResolvedKernelLaunch,
}
fn scalar_node_count(plan: ScalarDispatchPlan) -> usize {
match plan {
ScalarDispatchPlan::NtSplitKTail { k_tail, .. } => 3 + k_tail,
ScalarDispatchPlan::NnSplitKThinTail { .. }
| ScalarDispatchPlan::NnSplitKThin
| ScalarDispatchPlan::NnSplitKSlim { .. }
| ScalarDispatchPlan::NnM32N64SplitK32Qualified
| ScalarDispatchPlan::TnNarrowSplitM { .. }
| ScalarDispatchPlan::TnSplitM { .. }
| ScalarDispatchPlan::TnD768InSm89DualChunkQualified
| ScalarDispatchPlan::TnD768OutSm89DirectBk16Qualified
| ScalarDispatchPlan::TnPrismSm89DirectBk16Qualified => 2,
ScalarDispatchPlan::NtD768TransposeM64N64Qualified
| ScalarDispatchPlan::NtD768OutTransposeM64N64Qualified
| ScalarDispatchPlan::NtD768OutSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtD768InSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtPrismSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtLargeDeepSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtLargeDeepTransposeM64N64Qualified
| ScalarDispatchPlan::NtPrismVectorQualified
| ScalarDispatchPlan::NtD128OutTransposeM64N64Qualified => 2,
ScalarDispatchPlan::NtSplitKMain { .. } | ScalarDispatchPlan::NtSplitKSlim { .. } => 3,
_ => 1,
}
}
fn scalar_plan_requires_zero_beta(plan: ScalarDispatchPlan) -> bool {
matches!(
plan,
ScalarDispatchPlan::NnSplitKThinTail { .. }
| ScalarDispatchPlan::NnSplitKThin
| ScalarDispatchPlan::NnSplitKSlim { .. }
| ScalarDispatchPlan::NnM32N64SplitK32Qualified
| ScalarDispatchPlan::NnSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtD768TransposeM64N64Qualified
| ScalarDispatchPlan::NtD768OutTransposeM64N64Qualified
| ScalarDispatchPlan::NtD768OutSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtD768InSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtPrismSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtLargeDeepSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtLargeDeepTransposeM64N64Qualified
| ScalarDispatchPlan::NtPrismVectorQualified
| ScalarDispatchPlan::NtD128OutTransposeM64N64Qualified
| ScalarDispatchPlan::NtM2N16SplitK32Qualified
)
}
fn scalar_plan_fields(plan: ScalarDispatchPlan) -> (u8, u64, u64) {
match plan {
ScalarDispatchPlan::NnUltraThin => (1, 0, 0),
ScalarDispatchPlan::NnNarrowSmall => (2, 0, 0),
ScalarDispatchPlan::NnNarrow => (3, 0, 0),
ScalarDispatchPlan::NnGemv => (4, 0, 0),
ScalarDispatchPlan::NnSplitKThinTail { k_main, k_tail } => {
(5, k_main as u64, k_tail as u64)
}
ScalarDispatchPlan::NnSplitKThin => (6, 0, 0),
ScalarDispatchPlan::NnSplitKSlim { chunks } => (7, chunks as u64, 0),
ScalarDispatchPlan::NnM32N64SplitK32Qualified => (36, 0, 0),
ScalarDispatchPlan::NnM64N64Qualified => (24, 0, 0),
ScalarDispatchPlan::NnSm89FixedCopyPlanQualified => (37, 0, 0),
ScalarDispatchPlan::NnFinal { slim } => (8, u64::from(slim), 0),
ScalarDispatchPlan::TnGemv => (9, 0, 0),
ScalarDispatchPlan::TnNarrow => (10, 0, 0),
ScalarDispatchPlan::TnNarrowSplitM { m_chunk, chunks } => {
(22, m_chunk as u64, chunks as u64)
}
ScalarDispatchPlan::TnSplitM { m_chunk, chunks } => (11, m_chunk as u64, chunks as u64),
ScalarDispatchPlan::TnD768InSm89DualChunkQualified => (41, 1_024, 2),
ScalarDispatchPlan::TnD768OutSm89DirectBk16Qualified => (42, 512, 4),
ScalarDispatchPlan::TnPrismSm89DirectBk16Qualified => (43, 784, 6),
ScalarDispatchPlan::TnD128InSm89DirectFoldQualified => (44, 16, 64),
ScalarDispatchPlan::TnD128OutSm89DirectFoldQualified => (45, 16, 64),
ScalarDispatchPlan::TnM16N16SplitM16Qualified => (35, 0, 0),
ScalarDispatchPlan::TnFinal { slim } => (12, u64::from(slim), 0),
ScalarDispatchPlan::NtNarrow => (13, 0, 0),
ScalarDispatchPlan::NtSmallBatchWide => (14, 0, 0),
ScalarDispatchPlan::NtGemv => (15, 0, 0),
ScalarDispatchPlan::NtSplitKTail { k_main, k_tail } => (16, k_main as u64, k_tail as u64),
ScalarDispatchPlan::NtSplitKMain { n_main, n_tail } => (17, n_main as u64, n_tail as u64),
ScalarDispatchPlan::NtSplitKSlim { chunks } => (18, chunks as u64, 0),
ScalarDispatchPlan::NtMidBatchWide => (19, 0, 0),
ScalarDispatchPlan::NtM2N16SplitK32Qualified => (31, 0, 0),
ScalarDispatchPlan::NtD768TransposeM64N64Qualified => (25, 0, 0),
ScalarDispatchPlan::NtD768OutTransposeM64N64Qualified => (26, 0, 0),
ScalarDispatchPlan::NtD768OutSm89FixedCopyPlanQualified => (38, 0, 0),
ScalarDispatchPlan::NtD768InSm89FixedCopyPlanQualified => (39, 0, 0),
ScalarDispatchPlan::NtPrismSm89FixedCopyPlanQualified => (40, 0, 0),
ScalarDispatchPlan::NtLargeDeepSm89FixedCopyPlanQualified => (46, 0, 0),
ScalarDispatchPlan::NtLargeDeepTransposeM64N64Qualified => (27, 0, 0),
ScalarDispatchPlan::NtPrismVectorQualified => (30, 0, 0),
ScalarDispatchPlan::NtD128OutTransposeM64N64Qualified => (29, 0, 0),
ScalarDispatchPlan::NtFinal { slim } => (20, u64::from(slim), 0),
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
struct ScalarArgumentLayout {
output_offset: u64,
a_offset: u64,
b_offset: u64,
null_pointer_mask: u64,
output_column: Option<u64>,
}
fn scalar_argument_layout(
request: F32TriadRequest,
operands: F32TriadOperands,
plan: ScalarDispatchPlan,
index: usize,
) -> ScalarArgumentLayout {
let bias_null = u64::from(operands.bias.is_none());
match (plan, index) {
(ScalarDispatchPlan::NnSplitKThinTail { .. }, 0)
| (ScalarDispatchPlan::NnSplitKThin, 0)
| (ScalarDispatchPlan::NnSplitKSlim { .. }, 0)
| (ScalarDispatchPlan::NnM32N64SplitK32Qualified, 0) => ScalarArgumentLayout::default(),
(ScalarDispatchPlan::NnSplitKThinTail { k_main, .. }, 1) => ScalarArgumentLayout {
a_offset: k_main as u64 * 4,
b_offset: k_main as u64 * request.shape.n as u64 * 4,
null_pointer_mask: bias_null << 2,
..ScalarArgumentLayout::default()
},
(ScalarDispatchPlan::NnSplitKThin, 1)
| (ScalarDispatchPlan::NnSplitKSlim { .. }, 1)
| (ScalarDispatchPlan::NnM32N64SplitK32Qualified, 1) => ScalarArgumentLayout {
null_pointer_mask: 0b11000 | (bias_null << 2),
..ScalarArgumentLayout::default()
},
(
ScalarDispatchPlan::NnM64N64Qualified
| ScalarDispatchPlan::NnSm89FixedCopyPlanQualified,
_,
) => ScalarArgumentLayout {
null_pointer_mask: bias_null << 3,
..ScalarArgumentLayout::default()
},
(ScalarDispatchPlan::NtM2N16SplitK32Qualified, _) => ScalarArgumentLayout::default(),
(ScalarDispatchPlan::TnM16N16SplitM16Qualified, _) => ScalarArgumentLayout::default(),
(ScalarDispatchPlan::NtD768TransposeM64N64Qualified, 1)
| (ScalarDispatchPlan::NtD768OutTransposeM64N64Qualified, 1)
| (ScalarDispatchPlan::NtD768OutSm89FixedCopyPlanQualified, 1)
| (ScalarDispatchPlan::NtD768InSm89FixedCopyPlanQualified, 1)
| (ScalarDispatchPlan::NtPrismSm89FixedCopyPlanQualified, 1)
| (ScalarDispatchPlan::NtLargeDeepSm89FixedCopyPlanQualified, 1)
| (ScalarDispatchPlan::NtLargeDeepTransposeM64N64Qualified, 1)
| (ScalarDispatchPlan::NtPrismVectorQualified, 1)
| (ScalarDispatchPlan::NtD128OutTransposeM64N64Qualified, 1) => ScalarArgumentLayout {
null_pointer_mask: bias_null << 3,
..ScalarArgumentLayout::default()
},
(ScalarDispatchPlan::NtD768TransposeM64N64Qualified, 0)
| (ScalarDispatchPlan::NtD768OutTransposeM64N64Qualified, 0)
| (ScalarDispatchPlan::NtD768OutSm89FixedCopyPlanQualified, 0)
| (ScalarDispatchPlan::NtD768InSm89FixedCopyPlanQualified, 0)
| (ScalarDispatchPlan::NtPrismSm89FixedCopyPlanQualified, 0)
| (ScalarDispatchPlan::NtLargeDeepSm89FixedCopyPlanQualified, 0)
| (ScalarDispatchPlan::NtLargeDeepTransposeM64N64Qualified, 0)
| (ScalarDispatchPlan::NtPrismVectorQualified, 0)
| (ScalarDispatchPlan::NtD128OutTransposeM64N64Qualified, 0) => {
ScalarArgumentLayout::default()
}
(ScalarDispatchPlan::TnGemv, _)
| (ScalarDispatchPlan::TnNarrow, _)
| (ScalarDispatchPlan::TnNarrowSplitM { .. }, _)
| (ScalarDispatchPlan::TnSplitM { .. }, _)
| (ScalarDispatchPlan::TnD768InSm89DualChunkQualified, _)
| (ScalarDispatchPlan::TnD768OutSm89DirectBk16Qualified, _)
| (ScalarDispatchPlan::TnPrismSm89DirectBk16Qualified, _)
| (ScalarDispatchPlan::TnD128InSm89DirectFoldQualified, _)
| (ScalarDispatchPlan::TnD128OutSm89DirectFoldQualified, _)
| (ScalarDispatchPlan::TnFinal { .. }, _)
| (ScalarDispatchPlan::NtNarrow, _)
| (ScalarDispatchPlan::NtSmallBatchWide, _)
| (ScalarDispatchPlan::NtGemv, _)
| (ScalarDispatchPlan::NtMidBatchWide, _)
| (ScalarDispatchPlan::NtFinal { .. }, _) => ScalarArgumentLayout::default(),
(ScalarDispatchPlan::NtSplitKTail { k_main, .. }, tail) if tail >= 3 => {
let column = k_main + tail - 3;
ScalarArgumentLayout {
b_offset: column as u64 * request.shape.n as u64 * 4,
output_column: Some(column as u64),
..ScalarArgumentLayout::default()
}
}
(ScalarDispatchPlan::NtSplitKMain { n_main, n_tail }, 2) if n_tail > 0 => {
ScalarArgumentLayout {
a_offset: n_main as u64 * 4,
b_offset: n_main as u64 * request.shape.k as u64 * 4,
null_pointer_mask: 0b100,
..ScalarArgumentLayout::default()
}
}
(ScalarDispatchPlan::NtSplitKMain { .. }, 2)
| (ScalarDispatchPlan::NtSplitKTail { .. }, 2)
| (ScalarDispatchPlan::NtSplitKSlim { .. }, 2) => ScalarArgumentLayout {
null_pointer_mask: 0b11100,
..ScalarArgumentLayout::default()
},
(ScalarDispatchPlan::NtSplitKMain { .. }, _)
| (ScalarDispatchPlan::NtSplitKTail { .. }, _)
| (ScalarDispatchPlan::NtSplitKSlim { .. }, _) => ScalarArgumentLayout::default(),
_ => ScalarArgumentLayout {
null_pointer_mask: bias_null << 3,
..ScalarArgumentLayout::default()
},
}
}
fn scalar_arguments_digest(
request: F32TriadRequest,
operands: F32TriadOperands,
plan: ScalarDispatchPlan,
index: usize,
symbol: &str,
) -> Sha256Digest {
let (plan_tag, plan_a, plan_b) = scalar_plan_fields(plan);
let layout = scalar_argument_layout(request, operands, plan, index);
let output_column = layout.output_column.map(u64::to_le_bytes);
FramedSha256::new(b"triad-scalar-kernel-arguments.v1")
.required(b"symbol", symbol.as_bytes())
.required(b"node-index", &(index as u64).to_le_bytes())
.required(b"plan", &[plan_tag])
.required(b"plan-a", &plan_a.to_le_bytes())
.required(b"plan-b", &plan_b.to_le_bytes())
.required(b"m", &(request.shape.m as u64).to_le_bytes())
.required(b"k", &(request.shape.k as u64).to_le_bytes())
.required(b"n", &(request.shape.n as u64).to_le_bytes())
.required(b"lda", &(request.shape.lda as u64).to_le_bytes())
.required(b"ldb", &(request.shape.ldb as u64).to_le_bytes())
.required(b"ldc", &(request.shape.ldc as u64).to_le_bytes())
.required(b"alpha", &operands.alpha.to_bits().to_le_bytes())
.required(b"beta", &operands.beta.to_bits().to_le_bytes())
.required(b"output-offset", &layout.output_offset.to_le_bytes())
.required(b"a-offset", &layout.a_offset.to_le_bytes())
.required(b"b-offset", &layout.b_offset.to_le_bytes())
.required(
b"null-pointer-mask",
&layout.null_pointer_mask.to_le_bytes(),
)
.optional(
b"output-column",
output_column.as_ref().map(<[u8; 8]>::as_slice),
)
.finish()
}
fn push_scalar_node(
nodes: &mut Vec<ScalarNodeSpec>,
context: (F32TriadRequest, F32TriadOperands, ScalarDispatchPlan),
symbol: &'static str,
tile: (u32, u32),
bk_stages: (u32, u8),
config: cudarc::driver::LaunchConfig,
) {
let (request, operands, plan) = context;
let index = nodes.len();
nodes.push(ScalarNodeSpec {
symbol,
tile,
bk: bk_stages.0,
stages: bk_stages.1,
launch: ResolvedKernelLaunch {
grid_dim: config.grid_dim,
block_dim: config.block_dim,
shared_mem_bytes: config.shared_mem_bytes,
arguments_digest: scalar_arguments_digest(request, operands, plan, index, symbol),
},
});
}
fn f32_base_is_vector_aligned(pointer: CUptr) -> bool {
pointer & 15 == 0
}
fn scalar_tn_kernel_symbol(plan: ScalarDispatchPlan, operands: F32TriadOperands) -> &'static str {
match plan {
ScalarDispatchPlan::TnNarrowSplitM { .. }
if f32_base_is_vector_aligned(operands.a) && f32_base_is_vector_aligned(operands.b) =>
{
"tn_narrow_splitm_partial_aligned"
}
ScalarDispatchPlan::TnNarrowSplitM { .. } => "tn_narrow_splitm_partial",
ScalarDispatchPlan::TnSplitM { .. }
if f32_base_is_vector_aligned(operands.a) && f32_base_is_vector_aligned(operands.b) =>
{
"tn_splitm_partial_aligned"
}
ScalarDispatchPlan::TnSplitM { .. } => "tn_splitm_partial",
ScalarDispatchPlan::TnFinal { slim: false }
if f32_base_is_vector_aligned(operands.output)
&& f32_base_is_vector_aligned(operands.a)
&& f32_base_is_vector_aligned(operands.b) =>
{
"tn_aligned"
}
ScalarDispatchPlan::TnFinal { slim: false } => "tn_big",
ScalarDispatchPlan::TnFinal { slim: true } => "tn_slim",
_ => unreachable!("TN vector-aligned symbol requested for a non-Big TN plan"),
}
}
fn scalar_physical_nodes(
request: F32TriadRequest,
operands: F32TriadOperands,
plan: ScalarDispatchPlan,
) -> Result<Box<[ScalarNodeSpec]>, String> {
let context = (request, operands, plan);
let shape = request.shape;
let m = checked_u32(shape.m, "scalar M")?;
let k = checked_u32(shape.k, "scalar K")?;
let n = checked_u32(shape.n, "scalar N")?;
let mut nodes = Vec::new();
nodes
.try_reserve_exact(scalar_node_count(plan))
.map_err(|error| format!("reserve scalar physical plan: {error}"))?;
let cfg = |grid_dim, block_dim, shared_mem_bytes| cudarc::driver::LaunchConfig {
grid_dim,
block_dim,
shared_mem_bytes,
};
match plan {
ScalarDispatchPlan::NnUltraThin => push_scalar_node(
&mut nodes,
context,
"nn_ultra_thin",
(1, 32),
(32, 1),
cfg((n.div_ceil(32), m, 1), (256, 1, 1), k * 4),
),
ScalarDispatchPlan::NnNarrowSmall => push_scalar_node(
&mut nodes,
context,
"nn_narrow_small",
(16, 16),
(16, 1),
cfg(
(
checked_grid_product(m.div_ceil(16), n.div_ceil(16), 1)?,
1,
1,
),
(64, 1, 1),
0,
),
),
ScalarDispatchPlan::NnNarrow => push_scalar_node(
&mut nodes,
context,
"nn_narrow",
(64, 32),
(16, 1),
cfg(
(
checked_grid_product(m.div_ceil(64), n.div_ceil(32), 1)?,
1,
1,
),
(128, 1, 1),
0,
),
),
ScalarDispatchPlan::NnGemv => push_scalar_node(
&mut nodes,
context,
"nn_gemv",
(4, 1),
(32, 1),
cfg((m.div_ceil(4), 1, 1), (128, 1, 1), 0),
),
ScalarDispatchPlan::NnSplitKThinTail { .. } | ScalarDispatchPlan::NnSplitKThin => {
let k_main = match plan {
ScalarDispatchPlan::NnSplitKThinTail { k_main, .. } => {
checked_u32(k_main, "NN K main")?
}
_ => k,
};
push_scalar_node(
&mut nodes,
context,
"nn_splitk32_partial",
(32, 64),
(32, 1),
cfg(
(
checked_grid_product(m.div_ceil(32), n.div_ceil(64), k_main / 32)?,
1,
1,
),
(128, 1, 1),
0,
),
);
push_scalar_node(
&mut nodes,
context,
"splitk_reduce",
(1, 1),
(1, 1),
cfg(
(
checked_u32_product(m, n, "NN reducer outputs")?.div_ceil(256),
1,
1,
),
(256, 1, 1),
0,
),
);
}
ScalarDispatchPlan::NnSplitKSlim { chunks } => {
push_scalar_node(
&mut nodes,
context,
"nn_splitk_slim_partial",
(128, 64),
(32, 1),
cfg(
(
checked_grid_product(m.div_ceil(128), n.div_ceil(64), 1)?,
1,
chunks,
),
(128, 1, 1),
0,
),
);
push_scalar_node(
&mut nodes,
context,
"splitk_reduce",
(1, 1),
(1, 1),
cfg(
(
checked_u32_product(m, n, "NN slim reducer outputs")?.div_ceil(256),
1,
1,
),
(256, 1, 1),
0,
),
);
}
ScalarDispatchPlan::NnM32N64SplitK32Qualified => {
push_scalar_node(
&mut nodes,
context,
"nn_splitk32_m32n64_exact",
(32, 64),
(32, 1),
cfg((2_048, 1, 1), (128, 1, 1), 0),
);
push_scalar_node(
&mut nodes,
context,
"splitk_reduce",
(1, 1),
(1, 1),
cfg((64, 1, 1), (256, 1, 1), 0),
);
}
ScalarDispatchPlan::NnM64N64Qualified => push_scalar_node(
&mut nodes,
context,
"nn_m64n64_bk16_s2",
(64, 64),
(16, 2),
cfg(
(
checked_grid_product(m.div_ceil(64), n.div_ceil(64), 1)?,
1,
1,
),
(128, 1, 1),
super::contract::SCALAR_NN_M64N64_DYNAMIC_SHARED_BYTES,
),
),
ScalarDispatchPlan::NnSm89FixedCopyPlanQualified => push_scalar_node(
&mut nodes,
context,
"nn_sm89_f32_n64_copyplan",
(64, 64),
(32, 2),
cfg(
(
checked_grid_product(m.div_ceil(64), n.div_ceil(64), 1)?,
1,
1,
),
(128, 1, 1),
0,
),
),
ScalarDispatchPlan::NnFinal { slim } => {
let bn = if slim { 64 } else { 128 };
push_scalar_node(
&mut nodes,
context,
if slim { "nn_slim" } else { "nn_big" },
(128, bn),
(if slim { 32 } else { 16 }, if slim { 1 } else { 2 }),
cfg(
(
checked_grid_product(m.div_ceil(128), n.div_ceil(bn), 1)?,
1,
1,
),
(if slim { 128 } else { 256 }, 1, 1),
if slim { 0 } else { 34 * 1024 },
),
);
}
ScalarDispatchPlan::TnGemv => push_scalar_node(
&mut nodes,
context,
"tn_gemv",
(4, 1),
(32, 1),
cfg((k.div_ceil(4), 1, 1), (128, 1, 1), 0),
),
ScalarDispatchPlan::TnNarrow => push_scalar_node(
&mut nodes,
context,
"tn_narrow",
(64, 32),
(16, 1),
cfg(
(
checked_grid_product(k.div_ceil(64), n.div_ceil(32), 1)?,
1,
1,
),
(128, 1, 1),
0,
),
),
ScalarDispatchPlan::TnNarrowSplitM { chunks, .. } => {
push_scalar_node(
&mut nodes,
context,
scalar_tn_kernel_symbol(plan, operands),
(64, 32),
(16, 1),
cfg(
(
k.div_ceil(64),
n.div_ceil(32),
checked_u32(chunks, "TN narrow split-M chunks")?,
),
(128, 1, 1),
0,
),
);
push_scalar_node(
&mut nodes,
context,
"splitm_reduce",
(1, 1),
(1, 1),
cfg(
(
checked_u32_product(k, n, "TN narrow reducer outputs")?.div_ceil(256),
1,
1,
),
(256, 1, 1),
0,
),
);
}
ScalarDispatchPlan::TnM16N16SplitM16Qualified => push_scalar_node(
&mut nodes,
context,
"tn_m16n16_bk16_s2_splitm16",
(16, 16),
(16, 2),
cfg(
(
checked_grid_product(k.div_ceil(16), n.div_ceil(16), 1)?,
1,
1,
),
(super::contract::SCALAR_TN_M16N16_THREADS, 1, 1),
super::contract::SCALAR_TN_M16N16_DYNAMIC_SHARED_BYTES,
),
),
ScalarDispatchPlan::TnD768InSm89DualChunkQualified => {
push_scalar_node(
&mut nodes,
context,
"transpose_f32_32x16_d768",
(32, 32),
(1, 1),
cfg((24, 64, 1), (32, 16, 1), 0),
);
push_scalar_node(
&mut nodes,
context,
super::D768_IN_FUSED_SYMBOL,
(64, 64),
(32, 2),
cfg((576, 1, 1), (128, 1, 1), 0),
);
}
ScalarDispatchPlan::TnD768OutSm89DirectBk16Qualified
| ScalarDispatchPlan::TnPrismSm89DirectBk16Qualified => {
let (symbol, raw_grid, chunks) = match plan {
ScalarDispatchPlan::TnD768OutSm89DirectBk16Qualified => {
(super::D768_OUT_RAW_SYMBOL, 288, 4)
}
ScalarDispatchPlan::TnPrismSm89DirectBk16Qualified => {
(super::PRISM_RAW_SYMBOL, 186, 6)
}
_ => unreachable!(),
};
push_scalar_node(
&mut nodes,
context,
symbol,
(64, 64),
(16, 2),
cfg((raw_grid, 1, chunks), (128, 1, 1), 0),
);
push_scalar_node(
&mut nodes,
context,
"splitm_reduce",
(1, 1),
(1, 1),
cfg(
(
checked_u32_product(k, n, "SM89 exact-F32 TN reducer outputs")?
.div_ceil(256),
1,
1,
),
(256, 1, 1),
0,
),
);
}
ScalarDispatchPlan::TnD128InSm89DirectFoldQualified
| ScalarDispatchPlan::TnD128OutSm89DirectFoldQualified => {
let symbol = match plan {
ScalarDispatchPlan::TnD128InSm89DirectFoldQualified => super::D128_IN_SYMBOL,
ScalarDispatchPlan::TnD128OutSm89DirectFoldQualified => super::D128_OUT_SYMBOL,
_ => unreachable!(),
};
let spec = super::sm89_exact_f32_d128_source::kernel_spec(symbol)
.ok_or_else(|| format!("missing exact-F32 d128 spec for {symbol}"))?;
push_scalar_node(
&mut nodes,
context,
symbol,
spec.tile,
(16, 2),
cfg(spec.grid, spec.block, spec.dynamic_shared_bytes),
);
}
ScalarDispatchPlan::TnSplitM { chunks, .. } => {
push_scalar_node(
&mut nodes,
context,
scalar_tn_kernel_symbol(plan, operands),
(128, 128),
(16, 1),
cfg(
(
checked_grid_product(k.div_ceil(128), n.div_ceil(128), 1)?,
1,
checked_u32(chunks, "TN split-M chunks")?,
),
(256, 1, 1),
0,
),
);
push_scalar_node(
&mut nodes,
context,
"splitm_reduce",
(1, 1),
(1, 1),
cfg(
(
checked_u32_product(k, n, "TN reducer outputs")?.div_ceil(256),
1,
1,
),
(256, 1, 1),
0,
),
);
}
ScalarDispatchPlan::TnFinal { slim } => {
let bn = if slim { 64 } else { 128 };
push_scalar_node(
&mut nodes,
context,
scalar_tn_kernel_symbol(plan, operands),
(128, bn),
(if slim { 32 } else { 16 }, if slim { 1 } else { 2 }),
cfg(
(
checked_grid_product(k.div_ceil(128), n.div_ceil(bn), 1)?,
1,
1,
),
(if slim { 128 } else { 256 }, 1, 1),
if slim { 0 } else { 34 * 1024 },
),
);
}
ScalarDispatchPlan::NtNarrow
| ScalarDispatchPlan::NtSmallBatchWide
| ScalarDispatchPlan::NtMidBatchWide => push_scalar_node(
&mut nodes,
context,
"nt_narrow",
(64, 32),
(16, 1),
cfg(
(
checked_grid_product(m.div_ceil(64), k.div_ceil(32), 1)?,
1,
1,
),
(128, 1, 1),
0,
),
),
ScalarDispatchPlan::NtGemv => push_scalar_node(
&mut nodes,
context,
"nt_gemv",
(1, 1),
(1, 1),
cfg(
(
checked_u32_product(m, k, "NT GEMV outputs")?.div_ceil(256),
1,
1,
),
(256, 1, 1),
0,
),
),
ScalarDispatchPlan::NtM2N16SplitK32Qualified => push_scalar_node(
&mut nodes,
context,
"nt_m2n16_bk64_splitk32",
(2, 16),
(64, 2),
cfg(
(256, 1, 1),
(super::contract::SCALAR_NT_M2N16_THREADS, 1, 1),
super::contract::SCALAR_NT_M2N16_DYNAMIC_SHARED_BYTES,
),
),
ScalarDispatchPlan::NtD768TransposeM64N64Qualified
| ScalarDispatchPlan::NtD768OutTransposeM64N64Qualified
| ScalarDispatchPlan::NtD768OutSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtD768InSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtPrismSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtLargeDeepSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtLargeDeepTransposeM64N64Qualified
| ScalarDispatchPlan::NtPrismVectorQualified
| ScalarDispatchPlan::NtD128OutTransposeM64N64Qualified => {
let (m64_symbol, bk, shared_mem_bytes) =
if plan == ScalarDispatchPlan::NtPrismVectorQualified {
(
"nn_prism_m64n64_bk16_s2",
16,
super::contract::SCALAR_NN_M64N64_DYNAMIC_SHARED_BYTES,
)
} else if matches!(
plan,
ScalarDispatchPlan::NtD768OutSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtD768InSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtPrismSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtLargeDeepSm89FixedCopyPlanQualified
) {
("nn_sm89_f32_n64_copyplan", 32, 0)
} else {
(
"nn_m64n64_bk16_s2",
16,
super::contract::SCALAR_NN_M64N64_DYNAMIC_SHARED_BYTES,
)
};
push_scalar_node(
&mut nodes,
context,
"transpose_f32_32x16_d768",
(32, 32),
(1, 1),
cfg((n.div_ceil(32), k.div_ceil(32), 1), (32, 16, 1), 0),
);
push_scalar_node(
&mut nodes,
context,
m64_symbol,
(64, 64),
(bk, 2),
cfg(
(
checked_grid_product(m.div_ceil(64), k.div_ceil(64), 1)?,
1,
1,
),
(128, 1, 1),
shared_mem_bytes,
),
);
}
ScalarDispatchPlan::NtSplitKTail { k_main, k_tail } => {
let k_main_u32 = checked_u32(k_main, "NT K main")?;
push_scalar_node(
&mut nodes,
context,
"transpose_f32_2d",
(32, 32),
(1, 1),
cfg((n.div_ceil(32), k_main_u32.div_ceil(32), 1), (32, 32, 1), 0),
);
push_scalar_node(
&mut nodes,
context,
"nn_splitk32_partial",
(32, 64),
(32, 1),
cfg(
(
checked_grid_product(m.div_ceil(32), k_main_u32.div_ceil(64), n / 32)?,
1,
1,
),
(128, 1, 1),
0,
),
);
push_scalar_node(
&mut nodes,
context,
"splitk_reduce",
(1, 1),
(1, 1),
cfg(
(
checked_u32_product(m, k_main_u32, "NT K-tail reducer outputs")?
.div_ceil(256),
1,
1,
),
(256, 1, 1),
0,
),
);
for _ in 0..k_tail {
push_scalar_node(
&mut nodes,
context,
"dx_col_gemv",
(1, 1),
(1, 1),
cfg((m.div_ceil(128), 1, 1), (128, 1, 1), 0),
);
}
}
ScalarDispatchPlan::NtSplitKMain { n_main, .. } => {
let n_main = checked_u32(n_main, "NT N main")?;
push_scalar_node(
&mut nodes,
context,
"transpose_f32_2d",
(32, 32),
(1, 1),
cfg((n.div_ceil(32), k.div_ceil(32), 1), (32, 32, 1), 0),
);
push_scalar_node(
&mut nodes,
context,
"nn_splitk32_partial",
(32, 64),
(32, 1),
cfg(
(
checked_grid_product(m.div_ceil(32), k.div_ceil(64), n_main / 32)?,
1,
1,
),
(128, 1, 1),
0,
),
);
push_scalar_node(
&mut nodes,
context,
"splitk_reduce",
(1, 1),
(1, 1),
cfg(
(
checked_u32_product(m, k, "NT reducer outputs")?.div_ceil(256),
1,
1,
),
(256, 1, 1),
0,
),
);
}
ScalarDispatchPlan::NtSplitKSlim { chunks } => {
push_scalar_node(
&mut nodes,
context,
"transpose_f32_2d",
(32, 32),
(1, 1),
cfg((n.div_ceil(32), k.div_ceil(32), 1), (32, 32, 1), 0),
);
push_scalar_node(
&mut nodes,
context,
"nn_splitk_slim_partial",
(128, 64),
(32, 1),
cfg(
(
checked_grid_product(m.div_ceil(128), k.div_ceil(64), 1)?,
1,
chunks,
),
(128, 1, 1),
0,
),
);
push_scalar_node(
&mut nodes,
context,
"splitk_reduce",
(1, 1),
(1, 1),
cfg(
(
checked_u32_product(m, k, "NT slim reducer outputs")?.div_ceil(256),
1,
1,
),
(256, 1, 1),
0,
),
);
}
ScalarDispatchPlan::NtFinal { slim } => {
let bn = if slim { 64 } else { 128 };
push_scalar_node(
&mut nodes,
context,
if slim { "nt_slim" } else { "nt_big" },
(128, bn),
(if slim { 32 } else { 16 }, if slim { 1 } else { 2 }),
cfg(
(
checked_grid_product(m.div_ceil(128), k.div_ceil(bn), 1)?,
1,
1,
),
(if slim { 128 } else { 256 }, 1, 1),
if slim {
0
} else {
SCALAR_BIG_NT_DYNAMIC_SHARED_BYTES
},
),
);
}
}
debug_assert_eq!(nodes.len(), scalar_node_count(plan));
Ok(nodes.into_boxed_slice())
}
fn scalar_route_contract(
symbol: &str,
) -> (
PhysicalGemmBackend,
ResolvedNumericContract,
ResolvedOutputOwnership,
) {
if symbol == super::D128_IN_SYMBOL || symbol == super::D128_OUT_SYMBOL {
return (
PhysicalGemmBackend::ScalarFmaTnDirectF64FoldSm89,
ResolvedNumericContract::ScalarFmaTnSplitMF64Reduce,
ResolvedOutputOwnership::OneCtaPerOutputTile,
);
}
if symbol == super::D768_IN_FUSED_SYMBOL {
return (
PhysicalGemmBackend::ScalarFmaSm89ExactF32DualChunkFused,
ResolvedNumericContract::ScalarFmaTnSplitMF64Reduce,
ResolvedOutputOwnership::OneCtaPerOutputTile,
);
}
if symbol == super::D768_OUT_RAW_SYMBOL || symbol == super::PRISM_RAW_SYMBOL {
return (
PhysicalGemmBackend::ScalarFmaSm89ExactF32DirectSplitMPartial,
ResolvedNumericContract::ScalarFmaTnSplitMPartial,
ResolvedOutputOwnership::OneCtaPerOutputTilePerSplitMPartition,
);
}
if symbol == "nn_sm89_f32_n64_copyplan" {
return (
PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan,
ResolvedNumericContract::ScalarFma,
ResolvedOutputOwnership::OneCtaPerOutputTile,
);
}
if matches!(
symbol,
"nn_splitk32_partial" | "nn_splitk_slim_partial" | "nn_splitk32_m32n64_exact"
) {
return (
PhysicalGemmBackend::ScalarFmaSplitKPartial,
ResolvedNumericContract::ScalarFmaSplitKPartial,
ResolvedOutputOwnership::OneCtaPerOutputTilePerSplitKPartition,
);
}
if symbol == "splitk_reduce" {
return (
PhysicalGemmBackend::ScalarFmaSplitKF32Reduce,
ResolvedNumericContract::ScalarFmaSplitKF32Reduce,
ResolvedOutputOwnership::OneThreadPerOutputElementFixedSplitKReduce,
);
}
if symbol.starts_with("tn_narrow_splitm_partial") {
return (
PhysicalGemmBackend::ScalarFmaTnNarrowSplitMPartial,
ResolvedNumericContract::ScalarFmaTnNarrowSplitMPartial,
ResolvedOutputOwnership::OneCtaPerOutputTile,
);
}
if symbol == "splitm_reduce" {
return (
PhysicalGemmBackend::ScalarFmaTnSplitMF64Reduce,
ResolvedNumericContract::ScalarFmaTnSplitMF64Reduce,
ResolvedOutputOwnership::OneThreadPerOutputElementFixedSplitMReduce,
);
}
if symbol == "tn_m16n16_bk16_s2_splitm16" {
return (
PhysicalGemmBackend::ScalarFmaTnSplitMF64Reduce,
ResolvedNumericContract::ScalarFmaTnSplitMF64Reduce,
ResolvedOutputOwnership::OneCtaPerOutputTile,
);
}
(
PhysicalGemmBackend::ScalarFma,
ResolvedNumericContract::ScalarFma,
ResolvedOutputOwnership::OneCtaPerOutputTile,
)
}
fn scalar_resolved_routes(
ctx: &GpuCtx,
request: F32TriadRequest,
resources_digest: Sha256Digest,
nodes: &[ScalarNodeSpec],
) -> Result<Box<[ResolvedGemmRoute]>, String> {
let context = ctx.gemm_route();
let mut routes = Vec::new();
routes
.try_reserve_exact(nodes.len())
.map_err(|error| format!("reserve scalar route plan: {error}"))?;
for node in nodes {
let (backend, numeric_contract, ownership) = scalar_route_contract(node.symbol);
let (module_kind, artifact, compiler, tuning_table_revision) = if backend
== PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan
{
(
ModuleKind::Fixed,
context.artifacts.fixed,
ctx.kernels.compiler_identity(),
SM89_FIXED_COPYPLAN_ROUTE_REVISION,
)
} else if backend == PhysicalGemmBackend::ScalarFmaTnDirectF64FoldSm89 {
(
ModuleKind::TriadSm89ExactF32D128,
context.artifacts.sm89_exact_f32_d128.ok_or_else(|| {
"SM89 exact-F32 d128 route lost its artifact identity".to_string()
})?,
ctx.kernels
.triad_sm89_exact_f32_d128_compiler_identity()
.ok_or_else(|| {
"SM89 exact-F32 d128 route lost its compiler identity".to_string()
})?,
SM89_EXACT_F32_D128_ROUTE_REVISION,
)
} else if matches!(
backend,
PhysicalGemmBackend::ScalarFmaSm89ExactF32DualChunkFused
| PhysicalGemmBackend::ScalarFmaSm89ExactF32DirectSplitMPartial
) {
(
ModuleKind::TriadSm89ExactF32,
context
.artifacts
.sm89_exact_f32
.ok_or_else(|| "SM89 exact-F32 route lost its artifact identity".to_string())?,
ctx.kernels
.triad_sm89_exact_f32_compiler_identity()
.ok_or_else(|| "SM89 exact-F32 route lost its compiler identity".to_string())?,
SM89_EXACT_F32_TN_ROUTE_REVISION,
)
} else {
(
ModuleKind::TriadScalar,
context.artifacts.triad_scalar,
ctx.kernels.triad_scalar_compiler_identity(),
TUNING_TABLE_REVISION,
)
};
routes.push(ResolvedGemmRoute {
op: request.op,
dtype: PolicyDtype::F32,
backend,
numeric_contract,
instruction_family: ResolvedInstructionFamily::ScalarFma,
instruction_shape: ResolvedInstructionShape { m: 1, n: 1, k: 1 },
operand_conversion: ResolvedOperandConversion::None,
ownership,
symbol: node.symbol,
module_kind,
target: compiler.target,
artifact,
compiler,
device: context.device,
device_caps: context.device_caps,
shape: (request.shape.m, request.shape.k, request.shape.n),
strides: (request.shape.lda, request.shape.ldb, request.shape.ldc),
tile: node.tile,
bk: node.bk,
stages: node.stages,
threads: node.launch.block_dim.0 * node.launch.block_dim.1 * node.launch.block_dim.2,
launch: node.launch,
tensor_map_revision: 0,
tensor_maps_digest: [0; 32],
resources_digest,
tuning_table_revision,
schedule_revision: SCHEDULE_REVISION,
});
}
Ok(routes.into_boxed_slice())
}
fn validate_f32_triad_operands(
request: F32TriadRequest,
operands: F32TriadOperands,
) -> Result<(), String> {
let alignment = std::mem::align_of::<f32>() as CUptr;
if operands.output == 0 || !operands.output.is_multiple_of(alignment) {
return Err("f32 Triad output pointer must be non-null and 4-byte aligned".into());
}
if operands
.bias
.is_some_and(|bias| bias == 0 || !bias.is_multiple_of(alignment))
{
return Err("f32 Triad bias pointer must be non-null and 4-byte aligned".into());
}
if request.shape.reduction(request.op) != 0 {
for (name, pointer) in [("A", operands.a), ("B", operands.b)] {
if pointer == 0 || !pointer.is_multiple_of(alignment) {
return Err(format!(
"f32 Triad {name} pointer must be non-null and 4-byte aligned"
));
}
}
}
match request.op {
ResolvedGemmOp::Nn if operands.bias.is_some() && operands.alpha != 1.0 => {
Err("f32 Triad NN bias requires alpha == 1.0".into())
}
ResolvedGemmOp::Tn if operands.bias.is_some() || operands.beta != 1.0 => {
Err("f32 Triad TN requires no bias and beta == 1.0".into())
}
ResolvedGemmOp::Nt if operands.bias.is_some() || operands.beta != 0.0 => {
Err("f32 Triad NT requires no bias and beta == 0.0".into())
}
_ => Ok(()),
}
}
pub(in crate::mamba_ssm::gpu) fn validate_f32_triad_pointer_request(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
) -> Result<(), String> {
request.shape.validate(request.op)?;
validate_f32_triad_operands(request, operands)?;
let allocation_domain = validated_allocation_domain(&ctx.stream, &ctx.kernels, "f32 Triad")?;
let resources = F32LaunchResourceSnapshot::query_output(request, operands, allocation_domain)?;
if request.shape.reduction(request.op) == 0 {
Ok(())
} else {
resources
.with_inputs(request, operands, allocation_domain)
.map(|_| ())
}
}
fn require_f32_preparation_outside_capture(ctx: &GpuCtx) -> Result<(), String> {
let status = ctx
.stream
.capture_status()
.map_err(|error| format!("query f32 Triad capture status: {error:?}"))?;
if status != cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE {
return Err("f32 Triad launch must be prepared before graph capture".into());
}
Ok(())
}
fn validated_allocation_domain(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
backend: &str,
) -> Result<AllocationDomain, String> {
let allocation_domain = AllocationDomain::from_context(stream.context())?;
let kernel_domain = kernels.allocation_domain();
if allocation_domain != kernel_domain {
return Err(format!(
"{backend} stream allocation domain {allocation_domain:?} does not match kernel module allocation domain {kernel_domain:?}"
));
}
Ok(allocation_domain)
}
fn f32_map_binding(ctx: &GpuCtx, route: Tf32PhysicalRoute) -> Result<Tf32MapBinding, String> {
let qualified = match route {
Tf32PhysicalRoute::MmaTf32Rna(_)
| Tf32PhysicalRoute::MmaTf32RnaSplitK2(_)
| Tf32PhysicalRoute::MmaTf32RnaSplitK4(_)
| Tf32PhysicalRoute::MmaTf32RnaSplitK8(_) => ctx.kernels.f32_triad_availability().portable,
Tf32PhysicalRoute::Sm89MmaTf32Compact8 => ctx.kernels.f32_triad_availability().finalist,
Tf32PhysicalRoute::Sm89TnPreRnaN96
| Tf32PhysicalRoute::Sm89TnPreRnaM64N64
| Tf32PhysicalRoute::Sm89TnPreRnaM64N96S2
| Tf32PhysicalRoute::Sm89NnDirectN96
| Tf32PhysicalRoute::Sm89NnN96
| Tf32PhysicalRoute::Sm89NtALdmatrixN96
| Tf32PhysicalRoute::Sm89NtRnaM144N96S2
| Tf32PhysicalRoute::Sm89NtRowstageM128N192S2
| Tf32PhysicalRoute::Sm89TnDirectM192N192S2
| Tf32PhysicalRoute::Sm89TnPreRnaM96N192S2
| Tf32PhysicalRoute::Sm89TnPreRnaM96N96S3 => ctx.kernels.f32_triad_availability().joint,
_ => ctx.kernels.f32_triad_availability().specialized,
}
.ok_or_else(|| format!("TF32 route {route:?} has no qualified module"))?;
Ok(Tf32MapBinding {
allocation_domain: validated_allocation_domain(&ctx.stream, &ctx.kernels, "TF32")?,
qualified,
})
}
fn prepare_specialized_tf32_maps(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
route: Tf32PhysicalRoute,
binding: Tf32MapBinding,
) -> Result<F32PreparedTensorMaps, String> {
if matches!(route, Tf32PhysicalRoute::Sm90aWgmmaTf32Tma(_)) {
let descriptor = build_sm90a_tf32_descriptor(0, 0, 64);
if decode_sm90a_tf32_descriptor(descriptor) != (0, 0, 64) {
return Err("SM90a TF32 descriptor constants changed".into());
}
}
if matches!(route, Tf32PhysicalRoute::Sm100Tcgen05Tf32Tma(_)) {
sm100_tf32_instruction_descriptor(request.op, tf32_kernel_spec(request.op, route)?.tile.1)?;
}
let capturing = ctx
.stream
.capture_status()
.map_err(|error| format!("query CUDA capture status: {error:?}"))?
!= cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE;
prepare_f32_maps_with(
request,
operands,
route,
binding,
|request, operands, route| {
tf32_tensor_map_plan(request, operands, route, binding.allocation_domain)
},
|plan| {
ctx.kernels
.triad_kernels()
.prepare_tf32_tensor_maps(request, route, plan, capturing, binding)
},
)
}
fn tf32_params(
request: F32TriadRequest,
operands: F32TriadOperands,
origins: Tf32TensorOrigins,
route: Tf32PhysicalRoute,
) -> Result<PreparedTf32Params, String> {
let shape = request.shape;
let m = checked_i32(shape.m, "M")?;
let k = checked_i32(shape.k, "K")?;
let n = checked_i32(shape.n, "N")?;
let ldc = checked_i32(shape.ldc, "ldc")?;
Ok(match route {
Tf32PhysicalRoute::MmaTf32Rna(_)
| Tf32PhysicalRoute::MmaTf32RnaSplitK2(_)
| Tf32PhysicalRoute::MmaTf32RnaSplitK4(_)
| Tf32PhysicalRoute::MmaTf32RnaSplitK8(_)
| Tf32PhysicalRoute::Sm89MmaTf32Compact8
| Tf32PhysicalRoute::Sm89NnDirectN96
| Tf32PhysicalRoute::Sm89NnN96
| Tf32PhysicalRoute::Sm89NtALdmatrixN96
| Tf32PhysicalRoute::Sm89NtRnaM144N96S2
| Tf32PhysicalRoute::Sm89NtRowstageM128N192S2
| Tf32PhysicalRoute::Sm89TnDirectM192N192S2 => {
PreparedTf32Params::Sm80(Sm80Tf32KernelParams {
alpha: operands.alpha,
beta: operands.beta,
m,
k,
n,
lda: checked_i32(shape.lda, "lda")?,
ldb: checked_i32(shape.ldb, "ldb")?,
ldc,
})
}
Tf32PhysicalRoute::Sm89TnPreRnaN96
| Tf32PhysicalRoute::Sm89TnPreRnaM64N64
| Tf32PhysicalRoute::Sm89TnPreRnaM64N96S2
| Tf32PhysicalRoute::Sm89TnPreRnaM96N192S2
| Tf32PhysicalRoute::Sm89TnPreRnaM96N96S3 => {
return Err("Ada TF32 TN parameters require the two-node pre-RNA pipeline".into());
}
Tf32PhysicalRoute::Sm90aWgmmaTf32Tma(_) => {
PreparedTf32Params::Sm90a(Sm90aTf32KernelParams {
a_x: origins.a_x,
a_y: origins.a_y,
b_x: origins.b_x,
b_y: origins.b_y,
alpha: operands.alpha,
beta: operands.beta,
m,
k,
n,
ldc,
})
}
Tf32PhysicalRoute::Sm100Tcgen05Tf32Tma(_) => PreparedTf32Params::Sm100(Sm100KernelParams {
a_x: origins.a_x,
a_y: origins.a_y,
b_x: origins.b_x,
b_y: origins.b_y,
alpha: operands.alpha,
beta: operands.beta,
m,
k,
n,
ldc,
}),
Tf32PhysicalRoute::Sm120TmaMmaTf32Rna(_)
| Tf32PhysicalRoute::Sm120TmaMmaTf32RnaStreamKV1(_) => {
PreparedTf32Params::Sm120(Sm120KernelParams {
a_x: origins.a_x,
a_y: origins.a_y,
b_x: origins.b_x,
b_y: origins.b_y,
alpha: operands.alpha,
beta: operands.beta,
m,
k,
n,
ldc,
})
}
Tf32PhysicalRoute::Sm120TmaFmaExact(exact) => {
let plan = sm120_fma_launch_plan(request, exact)?;
PreparedTf32Params::Sm120Fma(Sm120FmaKernelParams {
alpha: operands.alpha,
beta: operands.beta,
m: checked_i32(shape.output_rows(request.op), "exact-F32 rows")?,
n: checked_i32(shape.output_columns(request.op), "exact-F32 columns")?,
k: checked_i32(shape.reduction(request.op), "exact-F32 reduction")?,
ldc,
splits: i32::from(exact.splits),
tiles_per_split: checked_i32(plan.tiles_per_split, "exact-F32 tiles per split")?,
})
}
})
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(super) struct Sm120FmaLaunchPlan {
pub tiles: usize,
pub tiles_per_split: usize,
pub units: u32,
pub slab_elements: usize,
pub flag_elements: usize,
}
pub(super) fn sm120_fma_launch_plan(
request: F32TriadRequest,
route: Sm120FmaRoute,
) -> Result<Sm120FmaLaunchPlan, String> {
request.shape.validate(request.op)?;
route.validate(request.op)?;
let (bm, bn) = route.tile.dims();
let rows = request.shape.output_rows(request.op);
let columns = request.shape.output_columns(request.op);
let reduction = request.shape.reduction(request.op);
if reduction == 0 {
return Err("exact-F32 SM120 route requires a nonzero reduction".into());
}
let tiles = rows
.div_ceil(bm as usize)
.checked_mul(columns.div_ceil(bn as usize))
.ok_or_else(|| invalid_gemm_dimensions("exact-F32 tile count overflows usize"))?;
let splits = usize::from(route.splits);
let k_tiles = reduction.div_ceil(SM120_FMA_BK as usize);
let tiles_per_split = k_tiles.div_ceil(splits);
if tiles_per_split
.checked_mul(splits - 1)
.is_none_or(|covered| covered >= k_tiles)
{
return Err(invalid_gemm_dimensions(
"exact-F32 split leaves a reduction range empty",
));
}
let units = tiles
.checked_mul(splits)
.and_then(|units| u32::try_from(units).ok())
.filter(|units| *units <= i32::MAX as u32)
.ok_or_else(|| {
invalid_gemm_dimensions("exact-F32 unit count exceeds the kernel's range")
})?;
let slab_elements = tiles
.checked_mul((bm * bn) as usize)
.and_then(|slab| slab.checked_mul(splits - 1))
.ok_or_else(|| invalid_gemm_dimensions("exact-F32 slab extent overflows usize"))?;
if slab_elements > SPLITK_SCRATCH_CAP {
return Err(invalid_gemm_dimensions(
"exact-F32 slabs exceed the fixed workspace",
));
}
let flag_elements = tiles * (splits - 1);
if flag_elements > TF32_SPLITK_COUNTER_CAP {
return Err(invalid_gemm_dimensions(
"exact-F32 flags exceed the fixed counter workspace",
));
}
Ok(Sm120FmaLaunchPlan {
tiles,
tiles_per_split,
units,
slab_elements,
flag_elements,
})
}
#[derive(Clone, Copy)]
struct Tf32LaunchDigests {
maps: Sha256Digest,
resources: Sha256Digest,
arguments: Sha256Digest,
}
#[derive(Clone, Copy, Debug)]
struct Tf32SplitKLaunchPlan {
scratch_elements: usize,
counter_elements: usize,
fused: cudarc::driver::LaunchConfig,
}
#[derive(Clone, Copy)]
struct Tf32SplitKWorkspace {
partial: CUptr,
counters: CUptr,
}
fn tf32_splitk_launch_plan(
request: F32TriadRequest,
spec: &Tf32SplitKSpec,
) -> Result<Tf32SplitKLaunchPlan, String> {
request.shape.validate(request.op)?;
if request.op != spec.op {
return Err(format!(
"portable TF32 split-K specification {:?} does not match {:?}",
spec.op, request.op
));
}
let reduction = request.shape.reduction(request.op);
if reduction == 0 {
return Err("portable TF32 split-K requires a nonzero reduction".into());
}
let (_, covered) =
tf32_splitk_partition_bounds(reduction, spec.partitions, spec.partitions - 1)?;
if covered != reduction {
return Err("portable TF32 split-K partitions do not cover the reduction".into());
}
let rows = request.shape.output_rows(request.op);
let columns = request.shape.output_columns(request.op);
let output_elements = rows
.checked_mul(columns)
.ok_or_else(|| invalid_gemm_dimensions("TF32 split-K output size overflows usize"))?;
let scratch_elements = output_elements
.checked_mul(spec.partitions as usize)
.ok_or_else(|| invalid_gemm_dimensions("TF32 split-K scratch size overflows usize"))?;
if scratch_elements > SPLITK_SCRATCH_CAP {
return Err(invalid_gemm_dimensions(
"TF32 split-K scratch exceeds the fixed workspace",
));
}
let rows = checked_u32(rows, "TF32 split-K output rows")?;
let columns = checked_u32(columns, "TF32 split-K output columns")?;
let row_tiles = rows.div_ceil(spec.tile.0);
if row_tiles > 65_535 {
return Err(invalid_gemm_dimensions(
"TF32 split-K row grid exceeds the CUDA y-dimension limit",
));
}
let column_tiles = columns.div_ceil(spec.tile.1);
let counter_elements = usize::try_from(row_tiles)
.ok()
.and_then(|rows| {
usize::try_from(column_tiles)
.ok()
.and_then(|columns| rows.checked_mul(columns))
})
.ok_or_else(|| invalid_gemm_dimensions("TF32 split-K counter size overflows usize"))?;
if counter_elements > TF32_SPLITK_COUNTER_CAP {
return Err(invalid_gemm_dimensions(
"TF32 split-K counters exceed the fixed coordination workspace",
));
}
Ok(Tf32SplitKLaunchPlan {
scratch_elements,
counter_elements,
fused: cudarc::driver::LaunchConfig {
grid_dim: (column_tiles, row_tiles, spec.partitions),
block_dim: (spec.threads, 1, 1),
shared_mem_bytes: spec.dynamic_shared_bytes,
},
})
}
fn tf32_splitk_arguments_digest(
request: F32TriadRequest,
operands: F32TriadOperands,
symbol: &str,
node_index: u8,
partitions: u32,
scratch_elements: usize,
) -> Sha256Digest {
let domain: &[u8] = match partitions {
2 => b"triad-tf32-split-k2-kernel-arguments.v2",
4 => b"triad-tf32-split-k4-kernel-arguments.v2",
8 => b"triad-tf32-split-k8-kernel-arguments.v2",
_ => b"triad-tf32-invalid-split-k-kernel-arguments.v1",
};
let partition_bytes = [partitions as u8];
FramedSha256::new(domain)
.required(b"symbol", symbol.as_bytes())
.required(b"node-index", &[node_index])
.required(b"op", &[request.op as u8])
.required(b"partitions", &partition_bytes)
.required(b"m", &(request.shape.m as u64).to_le_bytes())
.required(b"k", &(request.shape.k as u64).to_le_bytes())
.required(b"n", &(request.shape.n as u64).to_le_bytes())
.required(b"lda", &(request.shape.lda as u64).to_le_bytes())
.required(b"ldb", &(request.shape.ldb as u64).to_le_bytes())
.required(b"ldc", &(request.shape.ldc as u64).to_le_bytes())
.required(b"alpha", &operands.alpha.to_bits().to_le_bytes())
.required(b"beta", &operands.beta.to_bits().to_le_bytes())
.required(b"bias-null", &[u8::from(operands.bias.is_none())])
.required(
b"scratch-elements",
&(scratch_elements as u64).to_le_bytes(),
)
.finish()
}
fn tf32_splitk_resolved_routes(
request: F32TriadRequest,
operands: F32TriadOperands,
spec: &Tf32SplitKSpec,
binding: Tf32MapBinding,
resources_digest: Sha256Digest,
plan: Tf32SplitKLaunchPlan,
) -> Box<[ResolvedGemmRoute]> {
let qualified = binding.qualified;
let shape = (request.shape.m, request.shape.k, request.shape.n);
let strides = (request.shape.lda, request.shape.ldb, request.shape.ldc);
let (backend, numeric_contract, ownership) = match spec.route {
Tf32PhysicalRoute::MmaTf32RnaSplitK2(_) => (
PhysicalGemmBackend::MmaTf32RnaSplitK2,
ResolvedNumericContract::MmaTf32RnaSplitK2,
ResolvedOutputOwnership::LastCtaPerOutputTileFixedSplitK2Reduce,
),
Tf32PhysicalRoute::MmaTf32RnaSplitK4(_) => (
PhysicalGemmBackend::MmaTf32RnaSplitK4,
ResolvedNumericContract::MmaTf32RnaSplitK4,
ResolvedOutputOwnership::LastCtaPerOutputTileFixedSplitK4Reduce,
),
Tf32PhysicalRoute::MmaTf32RnaSplitK8(_) => (
PhysicalGemmBackend::MmaTf32RnaSplitK8,
ResolvedNumericContract::MmaTf32RnaSplitK8,
ResolvedOutputOwnership::LastCtaPerOutputTileFixedSplitK8Reduce,
),
_ => unreachable!("split-K spec admitted a non-split route"),
};
let fused = ResolvedGemmRoute {
op: request.op,
dtype: PolicyDtype::F32,
backend,
numeric_contract,
instruction_family: ResolvedInstructionFamily::MmaSync,
instruction_shape: ResolvedInstructionShape { m: 16, n: 8, k: 8 },
operand_conversion: ResolvedOperandConversion::RegisterCvtRnaTf32F32,
ownership,
symbol: spec.symbol,
module_kind: ModuleKind::TriadSm80,
target: qualified.target,
artifact: qualified.artifact,
compiler: qualified.compiler,
device: qualified.device,
device_caps: qualified.device_caps,
shape,
strides,
tile: spec.tile,
bk: spec.bk,
stages: spec.stages,
threads: spec.threads,
launch: ResolvedKernelLaunch {
grid_dim: plan.fused.grid_dim,
block_dim: plan.fused.block_dim,
shared_mem_bytes: plan.fused.shared_mem_bytes,
arguments_digest: tf32_splitk_arguments_digest(
request,
operands,
spec.symbol,
0,
spec.partitions,
plan.scratch_elements,
),
},
tensor_map_revision: 0,
tensor_maps_digest: [0; 32],
resources_digest,
tuning_table_revision: F32_TF32_TUNING_REVISION,
schedule_revision: SCHEDULE_REVISION,
};
vec![fused].into_boxed_slice()
}
fn tf32_resolved_route(
request: F32TriadRequest,
spec: &Tf32KernelSpec,
binding: Tf32MapBinding,
digests: Tf32LaunchDigests,
zero_reduction: bool,
config: cudarc::driver::LaunchConfig,
) -> ResolvedGemmRoute {
let (backend, numeric_contract) = match spec.route {
Tf32PhysicalRoute::MmaTf32Rna(_) => (
PhysicalGemmBackend::MmaTf32Rna,
ResolvedNumericContract::MmaTf32Rna,
),
Tf32PhysicalRoute::Sm89MmaTf32Compact8 => (
PhysicalGemmBackend::Sm89MmaTf32Compact8,
ResolvedNumericContract::MmaTf32Rna,
),
Tf32PhysicalRoute::Sm89TnPreRnaN96
| Tf32PhysicalRoute::Sm89TnPreRnaM64N64
| Tf32PhysicalRoute::Sm89TnPreRnaM64N96S2
| Tf32PhysicalRoute::Sm89TnPreRnaM96N192S2
| Tf32PhysicalRoute::Sm89TnPreRnaM96N96S3 => (
PhysicalGemmBackend::Sm89MmaTf32PreRna,
ResolvedNumericContract::MmaTf32PreRnaAV1,
),
Tf32PhysicalRoute::Sm89NtRnaM144N96S2 => (
PhysicalGemmBackend::Sm89MmaTf32NtRna,
ResolvedNumericContract::MmaTf32Rna,
),
Tf32PhysicalRoute::Sm89TnDirectM192N192S2 => (
PhysicalGemmBackend::Sm89MmaTf32TnDirectRna,
ResolvedNumericContract::MmaTf32Rna,
),
Tf32PhysicalRoute::Sm89NnDirectN96 | Tf32PhysicalRoute::Sm89NnN96 => (
PhysicalGemmBackend::Sm89MmaTf32AddHalf,
ResolvedNumericContract::MmaTf32AddHalfUlp,
),
Tf32PhysicalRoute::Sm89NtALdmatrixN96 | Tf32PhysicalRoute::Sm89NtRowstageM128N192S2 => (
PhysicalGemmBackend::Sm89MmaTf32NtALdmatrix,
ResolvedNumericContract::MmaTf32AddHalfUlp,
),
Tf32PhysicalRoute::MmaTf32RnaSplitK2(_) => (
PhysicalGemmBackend::MmaTf32RnaSplitK2,
ResolvedNumericContract::MmaTf32RnaSplitK2,
),
Tf32PhysicalRoute::MmaTf32RnaSplitK4(_) => (
PhysicalGemmBackend::MmaTf32RnaSplitK4,
ResolvedNumericContract::MmaTf32RnaSplitK4,
),
Tf32PhysicalRoute::MmaTf32RnaSplitK8(_) => (
PhysicalGemmBackend::MmaTf32RnaSplitK8,
ResolvedNumericContract::MmaTf32RnaSplitK8,
),
Tf32PhysicalRoute::Sm90aWgmmaTf32Tma(_) => (
PhysicalGemmBackend::Sm90aWgmmaTf32Tma,
ResolvedNumericContract::Sm90aWgmmaTf32Tma,
),
Tf32PhysicalRoute::Sm100Tcgen05Tf32Tma(_) => (
PhysicalGemmBackend::Sm100Tcgen05Tf32Tma,
ResolvedNumericContract::Sm100Tcgen05Tf32Tma,
),
Tf32PhysicalRoute::Sm120TmaMmaTf32Rna(_) => (
PhysicalGemmBackend::Sm120TmaMmaTf32Rna,
ResolvedNumericContract::Sm120TmaMmaTf32Rna,
),
Tf32PhysicalRoute::Sm120TmaMmaTf32RnaStreamKV1(_) => (
PhysicalGemmBackend::Sm120TmaMmaTf32RnaStreamKV1,
ResolvedNumericContract::Sm120TmaMmaTf32RnaStreamKV1,
),
Tf32PhysicalRoute::Sm120TmaFmaExact(exact) => (
PhysicalGemmBackend::Sm120TmaFmaExact,
sm120_fma_numeric_contract(exact),
),
};
let ownership = match spec.route {
Tf32PhysicalRoute::Sm120TmaMmaTf32RnaStreamKV1(_) => {
ResolvedOutputOwnership::OwnerCtaPerOutputTileStreamKFixedOrder
}
Tf32PhysicalRoute::Sm120TmaFmaExact(exact) => sm120_fma_ownership(exact),
_ => ResolvedOutputOwnership::OneCtaPerOutputTile,
};
ResolvedGemmRoute {
op: request.op,
dtype: PolicyDtype::F32,
backend,
numeric_contract: if zero_reduction {
ResolvedNumericContract::ZeroReductionEpilogueF32
} else {
numeric_contract
},
instruction_family: if zero_reduction {
ResolvedInstructionFamily::ScalarFma
} else {
spec.instruction_family
},
instruction_shape: if zero_reduction {
ResolvedInstructionShape { m: 1, n: 1, k: 1 }
} else {
spec.instruction_shape
},
operand_conversion: if zero_reduction {
ResolvedOperandConversion::None
} else {
spec.operand_conversion
},
ownership,
symbol: spec.symbol,
module_kind: spec.module_kind,
target: binding.qualified.target,
artifact: binding.qualified.artifact,
compiler: binding.qualified.compiler,
device: binding.qualified.device,
device_caps: binding.qualified.device_caps,
shape: (request.shape.m, request.shape.k, request.shape.n),
strides: (request.shape.lda, request.shape.ldb, request.shape.ldc),
tile: spec.tile,
bk: spec.bk,
stages: spec.stages,
threads: spec.threads,
launch: ResolvedKernelLaunch {
grid_dim: config.grid_dim,
block_dim: config.block_dim,
shared_mem_bytes: config.shared_mem_bytes,
arguments_digest: digests.arguments,
},
tensor_map_revision: if zero_reduction {
ZERO_REDUCTION_MAP_REVISION
} else {
spec.tensor_map_revision
},
tensor_maps_digest: digests.maps,
resources_digest: digests.resources,
tuning_table_revision: match spec.module_kind {
ModuleKind::TriadSm89Finalist => SM89_FINALIST_TUNING_REVISION,
ModuleKind::TriadSm89Tf32Joint => SM89_TF32_JOINT_TUNING_REVISION,
_ => F32_TF32_TUNING_REVISION,
},
schedule_revision: spec.schedule_revision,
}
}
fn scalar_zero_symbol(op: ResolvedGemmOp) -> &'static str {
match op {
ResolvedGemmOp::Nn => "nn_zero_reduction",
ResolvedGemmOp::Tn => "tn_zero_reduction",
ResolvedGemmOp::Nt => "nt_zero_reduction",
}
}
fn scalar_zero_route(
ctx: &GpuCtx,
request: F32TriadRequest,
maps_digest: Sha256Digest,
resources_digest: Sha256Digest,
config: cudarc::driver::LaunchConfig,
arguments_digest: Sha256Digest,
) -> ResolvedGemmRoute {
let context = ctx.gemm_route();
let compiler = ctx.kernels.triad_scalar_compiler_identity();
ResolvedGemmRoute {
op: request.op,
dtype: PolicyDtype::F32,
backend: PhysicalGemmBackend::ScalarFma,
numeric_contract: ResolvedNumericContract::ZeroReductionEpilogueF32,
instruction_family: ResolvedInstructionFamily::ScalarFma,
instruction_shape: ResolvedInstructionShape { m: 1, n: 1, k: 1 },
operand_conversion: ResolvedOperandConversion::None,
ownership: ResolvedOutputOwnership::OneCtaPerOutputTile,
symbol: scalar_zero_symbol(request.op),
module_kind: ModuleKind::TriadScalar,
target: compiler.target,
artifact: context.artifacts.triad_scalar,
compiler,
device: context.device,
device_caps: context.device_caps,
shape: (request.shape.m, request.shape.k, request.shape.n),
strides: (request.shape.lda, request.shape.ldb, request.shape.ldc),
tile: (1, 1),
bk: 0,
stages: 1,
threads: 256,
launch: ResolvedKernelLaunch {
grid_dim: config.grid_dim,
block_dim: config.block_dim,
shared_mem_bytes: config.shared_mem_bytes,
arguments_digest,
},
tensor_map_revision: ZERO_REDUCTION_MAP_REVISION,
tensor_maps_digest: maps_digest,
resources_digest,
tuning_table_revision: TUNING_TABLE_REVISION,
schedule_revision: SCHEDULE_REVISION,
}
}
struct ScalarScratchResources {
split: Option<(CUptr, u64)>,
transpose: Option<(CUptr, u64)>,
}
fn scalar_transpose_scratch_elements(
request: F32TriadRequest,
plan: ScalarDispatchPlan,
) -> Result<Option<usize>, String> {
if !plan.needs_transpose_scratch() {
return Ok(None);
}
let rows = match plan {
ScalarDispatchPlan::NtSplitKTail { k_main, .. } => k_main,
ScalarDispatchPlan::NtSplitKMain { .. }
| ScalarDispatchPlan::NtSplitKSlim { .. }
| ScalarDispatchPlan::NtD768TransposeM64N64Qualified
| ScalarDispatchPlan::NtD768OutTransposeM64N64Qualified
| ScalarDispatchPlan::NtD768OutSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtD768InSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtPrismSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtLargeDeepSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtLargeDeepTransposeM64N64Qualified
| ScalarDispatchPlan::NtPrismVectorQualified
| ScalarDispatchPlan::NtD128OutTransposeM64N64Qualified => request.shape.k,
ScalarDispatchPlan::TnD768InSm89DualChunkQualified => request.shape.m,
_ => return Err("scalar plan declares unsupported transpose scratch".into()),
};
let columns = match plan {
ScalarDispatchPlan::TnD768InSm89DualChunkQualified => request.shape.k,
_ => request.shape.n,
};
let elements = rows
.checked_mul(columns)
.ok_or_else(|| "scalar transpose scratch extent overflows usize".to_string())?;
if elements > super::contract::SCALAR_TRANSPOSE_SCRATCH_CAP_ELEMENTS {
return Err(format!(
"scalar transpose scratch requires {elements} f32 elements, capacity is {}",
super::contract::SCALAR_TRANSPOSE_SCRATCH_CAP_ELEMENTS
));
}
Ok(Some(elements))
}
fn scalar_launch_facts(kernels: &GpuKernels) -> ScalarLaunchFacts {
ScalarLaunchFacts {
scalar_artifact: kernels.artifact_set_identity().triad_scalar,
scalar_compiler: kernels.triad_scalar_compiler_identity(),
fixed_artifact: kernels.artifact_set_identity().fixed,
fixed_compiler: kernels.compiler_identity(),
fixed_copyplan_loaded: kernels.fixed_sm89_f32_n64_copyplan.is_some(),
sm89_exact_f32_artifact: kernels.triad_sm89_exact_f32_artifact_identity(),
sm89_exact_f32_compiler: kernels.triad_sm89_exact_f32_compiler_identity(),
sm89_exact_f32_symbols_loaded: [
kernels
.triad_sm89_exact_f32_function(super::D768_IN_FUSED_SYMBOL)
.is_some(),
kernels
.triad_sm89_exact_f32_function(super::D768_OUT_RAW_SYMBOL)
.is_some(),
kernels
.triad_sm89_exact_f32_function(super::PRISM_RAW_SYMBOL)
.is_some(),
],
sm89_exact_f32_d128_artifact: kernels.triad_sm89_exact_f32_d128_artifact_identity(),
sm89_exact_f32_d128_compiler: kernels.triad_sm89_exact_f32_d128_compiler_identity(),
sm89_exact_f32_d128_symbols_loaded: [
kernels
.triad_sm89_exact_f32_d128_function(super::D128_IN_SYMBOL)
.is_some(),
kernels
.triad_sm89_exact_f32_d128_function(super::D128_OUT_SYMBOL)
.is_some(),
],
compute_capability: kernels.triad_scalar_compute_capability(),
multiprocessor_count: kernels.multiprocessor_count(),
}
}
fn scalar_scratch_resources(
ctx: &GpuCtx,
request: F32TriadRequest,
plan: ScalarDispatchPlan,
) -> Result<ScalarScratchResources, String> {
use cudarc::driver::DevicePtr;
let split = if plan.needs_split_scratch() {
let buffer = ctx.kernels.splitk_scratch_buf(&ctx.stream)?;
let (pointer, _) = buffer.device_ptr(&ctx.stream);
Some((pointer, (SPLITK_SCRATCH_CAP as u64) * 4))
} else {
None
};
let transpose = if plan.needs_transpose_scratch() {
let buffer = ctx.kernels.transpose_scratch_buf(&ctx.stream)?;
let required = scalar_transpose_scratch_elements(request, plan)?
.ok_or_else(|| "scalar transpose plan lost its scratch extent".to_string())?;
if buffer.len() != super::contract::SCALAR_TRANSPOSE_SCRATCH_CAP_ELEMENTS
|| required > buffer.len()
{
return Err(format!(
"scalar transpose scratch allocation has {} f32 elements, requires {required} with exact capacity {}",
buffer.len(),
super::contract::SCALAR_TRANSPOSE_SCRATCH_CAP_ELEMENTS
));
}
let (pointer, _) = buffer.device_ptr(&ctx.stream);
Some((
pointer,
u64::try_from(super::contract::SCALAR_TRANSPOSE_SCRATCH_CAP_ELEMENTS)
.map_err(|_| "scalar transpose scratch capacity overflows u64")?
.checked_mul(4)
.ok_or_else(|| {
"scalar transpose scratch byte capacity overflows u64".to_string()
})?,
))
} else {
None
};
Ok(ScalarScratchResources { split, transpose })
}
fn prepare_scalar_f32(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
output_resources: F32LaunchResourceSnapshot,
) -> Result<PreparedF32TriadLaunch, String> {
let plan = scalar_ledger_plan(ctx, request, operands)?;
prepare_scalar_f32_with_plan(ctx, request, operands, output_resources, plan)
}
fn prepare_scalar_f32_with_plan(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
output_resources: F32LaunchResourceSnapshot,
plan: ScalarDispatchPlan,
) -> Result<PreparedF32TriadLaunch, String> {
if operands.beta != 0.0 && scalar_plan_requires_zero_beta(plan) {
return Err("selected exact scalar plan requires beta == 0".into());
}
let allocation_domain = validated_allocation_domain(&ctx.stream, &ctx.kernels, "f32 Triad")?;
let resources = output_resources.with_inputs(request, operands, allocation_domain)?;
let scratch = scalar_scratch_resources(ctx, request, plan)?;
let resources =
resources.with_scratch(scratch.split, scratch.transpose, None, allocation_domain)?;
let resources_digest = resources.digest(request, operands, [0; 32]);
let nodes = scalar_physical_nodes(request, operands, plan)?;
let routes = scalar_resolved_routes(ctx, request, resources_digest, &nodes)?;
let resolved_launch_set = build_resolved_gemm_launch_set(&routes)?;
let managed_epoch = resources.managed_epoch();
Ok(PreparedF32TriadLaunch {
context_token: ctx.instance_token(),
stream_token: ctx.stream_token(),
request,
operands,
resources,
managed_epoch,
routes,
resolved_launch_set,
kind: PreparedF32Kind::Scalar(plan),
})
}
fn prepare_scalar_zero_f32(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
resources: F32LaunchResourceSnapshot,
maps: F32PreparedTensorMaps,
) -> Result<PreparedF32TriadLaunch, String> {
let maps_digest = maps.identity_digest();
let resources_digest = resources.digest(request, operands, maps_digest);
let rows = request.shape.output_rows(request.op);
let columns = request.shape.output_columns(request.op);
let total = rows
.checked_mul(columns)
.ok_or_else(|| invalid_gemm_dimensions("zero-reduction output size overflows usize"))?;
let config = cudarc::driver::LaunchConfig {
grid_dim: (
checked_u32(total.div_ceil(256), "zero-reduction grid.x")?,
1,
1,
),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let symbol = scalar_zero_symbol(request.op);
let null_pointer_mask = 0b110_u64 | (u64::from(operands.bias.is_none()) << 3);
let arguments_digest = FramedSha256::new(b"triad-scalar-zero-reduction-arguments.v1")
.required(b"symbol", symbol.as_bytes())
.required(b"op", &[request.op as u8])
.required(b"m", &(request.shape.m as u64).to_le_bytes())
.required(b"k", &(request.shape.k as u64).to_le_bytes())
.required(b"n", &(request.shape.n as u64).to_le_bytes())
.required(b"lda", &(request.shape.lda as u64).to_le_bytes())
.required(b"ldb", &(request.shape.ldb as u64).to_le_bytes())
.required(b"ldc", &(request.shape.ldc as u64).to_le_bytes())
.required(b"alpha", &operands.alpha.to_bits().to_le_bytes())
.required(b"beta", &operands.beta.to_bits().to_le_bytes())
.required(b"output-offset", &0_u64.to_le_bytes())
.required(b"null-pointer-mask", &null_pointer_mask.to_le_bytes())
.finish();
let route = scalar_zero_route(
ctx,
request,
maps_digest,
resources_digest,
config,
arguments_digest,
);
let routes = vec![route].into_boxed_slice();
let resolved_launch_set = build_zero_reduction_route_identity(routes[0])?;
let params = SgbZeroReductionParams {
alpha: operands.alpha,
beta: operands.beta,
m: checked_i32(request.shape.m, "M")?,
k: checked_i32(request.shape.k, "K")?,
n: checked_i32(request.shape.n, "N")?,
lda: checked_i32(request.shape.lda, "lda")?,
ldb: checked_i32(request.shape.ldb, "ldb")?,
ldc: checked_i32(request.shape.ldc, "ldc")?,
};
let managed_epoch = resources.managed_epoch();
Ok(PreparedF32TriadLaunch {
context_token: ctx.instance_token(),
stream_token: ctx.stream_token(),
request,
operands,
resources,
managed_epoch,
routes,
resolved_launch_set,
kind: PreparedF32Kind::ScalarZero {
maps,
params,
config,
},
})
}
fn prepare_tf32_f32(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
output_resources: F32LaunchResourceSnapshot,
route: Tf32PhysicalRoute,
) -> Result<PreparedF32TriadLaunch, String> {
let availability = ctx.kernels.f32_triad_availability();
let route = resolve_tf32_forced(request, availability, route).map_err(|error| {
ctx.kernels
.tf32_qualification_rejection(route)
.map(|rejection| {
format!("forced TF32 route {route:?} failed module qualification: {rejection}")
})
.unwrap_or(error)
})?;
validate_sm89_tf32_joint_operands(request, operands, route)?;
if matches!(
route,
Tf32PhysicalRoute::MmaTf32RnaSplitK2(_)
| Tf32PhysicalRoute::MmaTf32RnaSplitK4(_)
| Tf32PhysicalRoute::MmaTf32RnaSplitK8(_)
) {
return prepare_tf32_splitk_f32(ctx, request, operands, output_resources, route);
}
if matches!(route, Tf32PhysicalRoute::Sm120TmaMmaTf32RnaStreamKV1(_)) {
return prepare_tf32_streamk_f32(ctx, request, operands, output_resources, route);
}
if route.is_exact_fma() {
return prepare_sm120_fma_f32(ctx, request, operands, output_resources, route);
}
if matches!(
route,
Tf32PhysicalRoute::Sm89TnPreRnaN96
| Tf32PhysicalRoute::Sm89TnPreRnaM64N64
| Tf32PhysicalRoute::Sm89TnPreRnaM64N96S2
| Tf32PhysicalRoute::Sm89TnPreRnaM96N192S2
| Tf32PhysicalRoute::Sm89TnPreRnaM96N96S3
) {
return prepare_sm89_tf32_tn_pre_rna(ctx, request, operands, output_resources, route);
}
let spec = tf32_kernel_spec(request.op, route)?;
let binding = f32_map_binding(ctx, route)?;
let allocation_domain = binding.allocation_domain;
let zero_reduction = request.shape.reduction(request.op) == 0;
let resources = if zero_reduction {
output_resources
} else {
output_resources.with_inputs(request, operands, allocation_domain)?
};
let maps = if zero_reduction {
Some(F32PreparedTensorMaps::zero_reduction(
request,
Some(route),
Some(binding),
match route {
Tf32PhysicalRoute::Sm120TmaMmaTf32Rna(_)
| Tf32PhysicalRoute::Sm120TmaMmaTf32RnaStreamKV1(_) => Tf32TensorMapFormat::Uint32,
_ => Tf32TensorMapFormat::Tfloat32,
},
))
} else if matches!(
route,
Tf32PhysicalRoute::MmaTf32Rna(_)
| Tf32PhysicalRoute::Sm89MmaTf32Compact8
| Tf32PhysicalRoute::Sm89NnDirectN96
| Tf32PhysicalRoute::Sm89NnN96
| Tf32PhysicalRoute::Sm89NtALdmatrixN96
| Tf32PhysicalRoute::Sm89NtRnaM144N96S2
| Tf32PhysicalRoute::Sm89NtRowstageM128N192S2
| Tf32PhysicalRoute::Sm89TnDirectM192N192S2
) {
None
} else {
Some(prepare_specialized_tf32_maps(
ctx, request, operands, route, binding,
)?)
};
let origins = maps
.as_ref()
.map(F32PreparedTensorMaps::origins)
.unwrap_or_default();
let maps_digest = maps
.as_ref()
.map(F32PreparedTensorMaps::identity_digest)
.unwrap_or([0; 32]);
let resources_digest = resources.digest(request, operands, maps_digest);
let rows = checked_u32(request.shape.output_rows(request.op), "TF32 output rows")?;
let columns = checked_u32(
request.shape.output_columns(request.op),
"TF32 output columns",
)?;
let config = cudarc::driver::LaunchConfig {
grid_dim: (
checked_tile_grid(rows, spec.tile.0, columns, spec.tile.1)?,
1,
1,
),
block_dim: (spec.threads, 1, 1),
shared_mem_bytes: spec.dynamic_shared_bytes,
};
let arguments_digest =
tf32_kernel_arguments_digest(request, operands, spec.symbol, maps_digest);
let resolved = tf32_resolved_route(
request,
spec,
binding,
Tf32LaunchDigests {
maps: maps_digest,
resources: resources_digest,
arguments: arguments_digest,
},
zero_reduction,
config,
);
let routes = vec![resolved].into_boxed_slice();
let resolved_launch_set = build_resolved_gemm_launch_set(&routes)?;
let managed_epoch = resources.managed_epoch();
Ok(PreparedF32TriadLaunch {
context_token: ctx.instance_token(),
stream_token: ctx.stream_token(),
request,
operands,
resources,
managed_epoch,
routes,
resolved_launch_set,
kind: PreparedF32Kind::Tf32 {
route,
maps,
params: tf32_params(request, operands, origins, route)?,
config,
},
})
}
fn validate_sm89_tf32_joint_operands(
request: F32TriadRequest,
operands: F32TriadOperands,
route: Tf32PhysicalRoute,
) -> Result<(), String> {
if !matches!(
route,
Tf32PhysicalRoute::Sm89TnPreRnaN96
| Tf32PhysicalRoute::Sm89TnPreRnaM64N64
| Tf32PhysicalRoute::Sm89TnPreRnaM64N96S2
| Tf32PhysicalRoute::Sm89NnDirectN96
| Tf32PhysicalRoute::Sm89NnN96
| Tf32PhysicalRoute::Sm89NtALdmatrixN96
| Tf32PhysicalRoute::Sm89NtRnaM144N96S2
| Tf32PhysicalRoute::Sm89NtRowstageM128N192S2
| Tf32PhysicalRoute::Sm89TnDirectM192N192S2
| Tf32PhysicalRoute::Sm89TnPreRnaM96N192S2
| Tf32PhysicalRoute::Sm89TnPreRnaM96N96S3
) {
return Ok(());
}
let required_beta: f32 = if request.op == ResolvedGemmOp::Tn {
1.0
} else {
0.0
};
if operands.alpha.to_bits() != 1.0_f32.to_bits()
|| operands.beta.to_bits() != required_beta.to_bits()
|| operands.bias.is_some()
|| [operands.output, operands.a, operands.b]
.into_iter()
.any(|pointer| pointer == 0 || !pointer.is_multiple_of(16))
{
return Err(format!(
"Ada TF32 joint route {route:?} requires exact alpha/beta, no bias, and non-null 16-byte aligned operands"
));
}
Ok(())
}
fn prepare_sm89_tf32_tn_pre_rna(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
output_resources: F32LaunchResourceSnapshot,
route: Tf32PhysicalRoute,
) -> Result<PreparedF32TriadLaunch, String> {
use cudarc::driver::DevicePtr;
let spec = tf32_kernel_spec(request.op, route)?;
let binding = f32_map_binding(ctx, route)?;
let allocation_domain = binding.allocation_domain;
let (output_stride, scratch_elements) = sm89_tf32_tn_scratch_layout(request)?;
let scratch_buffer = ctx.kernels.transpose_scratch_buf(&ctx.stream)?;
if scratch_elements > scratch_buffer.len()
|| scratch_buffer.len() != SCALAR_TRANSPOSE_SCRATCH_CAP_ELEMENTS
{
return Err(format!(
"Ada TF32 transpose scratch has {} elements, needs {scratch_elements} with exact capacity {SCALAR_TRANSPOSE_SCRATCH_CAP_ELEMENTS}",
scratch_buffer.len()
));
}
let (scratch, _) = scratch_buffer.device_ptr(&ctx.stream);
let scratch_bytes = u64::try_from(scratch_elements)
.ok()
.and_then(|elements| elements.checked_mul(4))
.ok_or_else(|| "Ada TF32 transpose byte extent overflows u64".to_string())?;
let resources = output_resources
.with_inputs(request, operands, allocation_domain)?
.with_scratch(
None,
Some((scratch, scratch_bytes)),
None,
allocation_domain,
)?;
let transpose_params = Sm89Tf32JointTransposeParams {
rows: checked_i32(request.shape.m, "Ada TF32 transpose rows")?,
columns: checked_i32(request.shape.k, "Ada TF32 transpose columns")?,
output_stride: checked_i32(output_stride, "Ada TF32 transpose stride")?,
};
let gemm_params = Sm89Tf32JointGemmParams {
alpha: operands.alpha,
beta: operands.beta,
m: checked_i32(request.shape.k, "Ada TF32 physical M")?,
k: checked_i32(request.shape.m, "Ada TF32 physical K")?,
n: checked_i32(request.shape.n, "Ada TF32 physical N")?,
lda: checked_i32(output_stride, "Ada TF32 physical lda")?,
ldb: checked_i32(request.shape.ldb, "Ada TF32 physical ldb")?,
ldc: checked_i32(request.shape.ldc, "Ada TF32 physical ldc")?,
};
let transpose_config = cudarc::driver::LaunchConfig {
grid_dim: (
checked_u32(request.shape.k.div_ceil(32), "Ada TF32 transpose grid x")?,
checked_u32(output_stride.div_ceil(32), "Ada TF32 transpose grid y")?,
1,
),
block_dim: (32, 8, 1),
shared_mem_bytes: 0,
};
let physical_rows = checked_u32(request.shape.k, "Ada TF32 physical rows")?;
let columns = checked_u32(request.shape.n, "Ada TF32 physical columns")?;
let gemm_config = cudarc::driver::LaunchConfig {
grid_dim: (
checked_tile_grid(physical_rows, spec.tile.0, columns, spec.tile.1)?,
1,
1,
),
block_dim: (spec.threads, 1, 1),
shared_mem_bytes: spec.dynamic_shared_bytes,
};
let resources_digest = resources.digest(request, operands, [0; 32]);
let arguments_digest = FramedSha256::new(b"triad-sm89-tf32-pre-rna-gemm-arguments.v1")
.required(b"symbol", spec.symbol.as_bytes())
.required(b"output", &operands.output.to_le_bytes())
.required(b"scratch", &scratch.to_le_bytes())
.required(b"B", &operands.b.to_le_bytes())
.required(b"bias-null", &0_u64.to_le_bytes())
.required(b"alpha", &operands.alpha.to_bits().to_le_bytes())
.required(b"beta", &operands.beta.to_bits().to_le_bytes())
.required(b"physical-m", &gemm_params.m.to_le_bytes())
.required(b"physical-k", &gemm_params.k.to_le_bytes())
.required(b"physical-n", &gemm_params.n.to_le_bytes())
.required(b"physical-lda", &gemm_params.lda.to_le_bytes())
.required(b"physical-ldb", &gemm_params.ldb.to_le_bytes())
.required(b"physical-ldc", &gemm_params.ldc.to_le_bytes())
.finish();
let resolved = tf32_resolved_route(
request,
spec,
binding,
Tf32LaunchDigests {
maps: [0; 32],
resources: resources_digest,
arguments: arguments_digest,
},
false,
gemm_config,
);
let transform = ResolvedInputTransform {
numeric_contract: ResolvedNumericContract::Tf32RnaPreprocess,
operand_conversion: ResolvedOperandConversion::RegisterCvtRnaTf32F32,
output_ownership: ResolvedTransformOutputOwnership::PreparedScratchAllocation,
target: binding.qualified.target,
artifact: binding.qualified.artifact,
compiler: binding.qualified.compiler,
device: binding.qualified.device,
device_caps: binding.qualified.device_caps,
output_stride: checked_u32(output_stride, "Ada TF32 transform stride")?,
output_elements: u64::try_from(scratch_elements)
.map_err(|_| "Ada TF32 transform extent exceeds u64".to_string())?,
resources_digest: resources.physical_digest(),
tuning_table_revision: SM89_TF32_JOINT_TUNING_REVISION,
schedule_revision: SCHEDULE_REVISION,
};
let routes = vec![resolved].into_boxed_slice();
let resolved_launch_set = build_resolved_gemm_launch_set(&routes)?;
let managed_epoch = resources.managed_epoch();
Ok(PreparedF32TriadLaunch {
context_token: ctx.instance_token(),
stream_token: ctx.stream_token(),
request,
operands,
resources,
managed_epoch,
routes,
resolved_launch_set,
kind: PreparedF32Kind::Tf32TnPreRna {
route,
transpose_params,
gemm_params,
transpose_config,
gemm_config,
scratch,
scratch_elements,
transform: Box::new(transform),
},
})
}
fn sm89_tf32_tn_scratch_layout(request: F32TriadRequest) -> Result<(usize, usize), String> {
let output_stride = request
.shape
.m
.checked_add(3)
.map(|value| value & !3)
.ok_or_else(|| "Ada TF32 transpose stride overflows usize".to_string())?;
let scratch_elements = request
.shape
.k
.checked_mul(output_stride)
.ok_or_else(|| "Ada TF32 transpose extent overflows usize".to_string())?;
Ok((output_stride, scratch_elements))
}
fn prepare_tf32_streamk_f32(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
output_resources: F32LaunchResourceSnapshot,
route: Tf32PhysicalRoute,
) -> Result<PreparedF32TriadLaunch, String> {
use cudarc::driver::DevicePtr;
let spec = tf32_kernel_spec(request.op, route)?;
let zero_reduction = request.shape.reduction(request.op) == 0;
let plan = tf32_streamk_launch_plan(request, spec, ctx.kernels.multiprocessor_count())?;
let binding = f32_map_binding(ctx, route)?;
let allocation_domain = binding.allocation_domain;
let resources = if zero_reduction {
output_resources
} else {
output_resources.with_inputs(request, operands, allocation_domain)?
};
let scratch_buffer = ctx.kernels.splitk_scratch_buf(&ctx.stream)?;
let (partial, _) = scratch_buffer.device_ptr(&ctx.stream);
let flag_buffer = ctx
.kernels
.triad_kernels()
.tf32_splitk_counter_buf(&ctx.stream)?;
let (flags, _) = flag_buffer.device_ptr(&ctx.stream);
let resources = resources.with_scratch(
Some((partial, (SPLITK_SCRATCH_CAP as u64) * 4)),
None,
Some((flags, (TF32_SPLITK_COUNTER_CAP as u64) * 4)),
allocation_domain,
)?;
let maps = if zero_reduction {
F32PreparedTensorMaps::zero_reduction(
request,
Some(route),
Some(binding),
Tf32TensorMapFormat::Uint32,
)
} else {
prepare_specialized_tf32_maps(ctx, request, operands, route, binding)?
};
let origins = maps.origins();
let maps_digest = maps.identity_digest();
let resources_digest = resources.digest(request, operands, maps_digest);
let config = cudarc::driver::LaunchConfig {
grid_dim: (plan.grid, 1, 1),
block_dim: (spec.threads, 1, 1),
shared_mem_bytes: spec.dynamic_shared_bytes,
};
let arguments_digest =
tf32_kernel_arguments_digest(request, operands, spec.symbol, maps_digest);
let resolved = tf32_resolved_route(
request,
spec,
binding,
Tf32LaunchDigests {
maps: maps_digest,
resources: resources_digest,
arguments: arguments_digest,
},
zero_reduction,
config,
);
let routes = vec![resolved].into_boxed_slice();
let resolved_launch_set = build_resolved_gemm_launch_set(&routes)?;
let managed_epoch = resources.managed_epoch();
Ok(PreparedF32TriadLaunch {
context_token: ctx.instance_token(),
stream_token: ctx.stream_token(),
request,
operands,
resources,
managed_epoch,
routes,
resolved_launch_set,
kind: PreparedF32Kind::Tf32StreamK {
route,
maps,
params: tf32_params(request, operands, origins, route)?,
config,
plan,
workspace: Tf32StreamKWorkspace { partial, flags },
},
})
}
fn prepare_sm120_fma_f32(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
output_resources: F32LaunchResourceSnapshot,
route: Tf32PhysicalRoute,
) -> Result<PreparedF32TriadLaunch, String> {
use cudarc::driver::DevicePtr;
let exact = route
.exact_fma()
.ok_or_else(|| "exact-F32 preparation requires an exact route".to_string())?;
let spec = tf32_kernel_spec(request.op, route)?;
let plan = sm120_fma_launch_plan(request, exact)?;
let binding = f32_map_binding(ctx, route)?;
let allocation_domain = binding.allocation_domain;
let resources = output_resources.with_inputs(request, operands, allocation_domain)?;
let (resources, partial, flags) = if exact.splits == 1 {
debug_assert_eq!(plan.slab_elements, 0);
debug_assert_eq!(plan.flag_elements, 0);
(resources, 0, 0)
} else {
let scratch_buffer = ctx.kernels.splitk_scratch_buf(&ctx.stream)?;
let (partial, _) = scratch_buffer.device_ptr(&ctx.stream);
let flag_buffer = ctx
.kernels
.triad_kernels()
.tf32_splitk_counter_buf(&ctx.stream)?;
let (flags, _) = flag_buffer.device_ptr(&ctx.stream);
let resources = resources.with_scratch(
Some((partial, (SPLITK_SCRATCH_CAP as u64) * 4)),
None,
Some((flags, (TF32_SPLITK_COUNTER_CAP as u64) * 4)),
allocation_domain,
)?;
(resources, partial, flags)
};
let maps = prepare_specialized_tf32_maps(ctx, request, operands, route, binding)?;
let origins = maps.origins();
let maps_digest = maps.identity_digest();
let resources_digest = resources.digest(request, operands, maps_digest);
let config = cudarc::driver::LaunchConfig {
grid_dim: (plan.units, 1, 1),
block_dim: (spec.threads, 1, 1),
shared_mem_bytes: spec.dynamic_shared_bytes,
};
let arguments_digest =
tf32_kernel_arguments_digest(request, operands, spec.symbol, maps_digest);
let resolved = tf32_resolved_route(
request,
spec,
binding,
Tf32LaunchDigests {
maps: maps_digest,
resources: resources_digest,
arguments: arguments_digest,
},
false,
config,
);
let resolved = ResolvedGemmRoute {
numeric_contract: sm120_fma_numeric_contract(exact),
ownership: sm120_fma_ownership(exact),
..resolved
};
let routes = vec![resolved].into_boxed_slice();
let resolved_launch_set = build_resolved_gemm_launch_set(&routes)?;
let managed_epoch = resources.managed_epoch();
Ok(PreparedF32TriadLaunch {
context_token: ctx.instance_token(),
stream_token: ctx.stream_token(),
request,
operands,
resources,
managed_epoch,
routes,
resolved_launch_set,
kind: PreparedF32Kind::Tf32StreamK {
route,
maps,
params: tf32_params(request, operands, origins, route)?,
config,
plan: Tf32StreamKLaunchPlan {
grid: plan.units,
partial_elements: plan.slab_elements,
flag_elements: plan.flag_elements,
},
workspace: Tf32StreamKWorkspace { partial, flags },
},
})
}
fn sm120_fma_numeric_contract(route: Sm120FmaRoute) -> ResolvedNumericContract {
if route.splits == 1 {
ResolvedNumericContract::ScalarFma
} else {
ResolvedNumericContract::ScalarFmaFixedSplitFold
}
}
fn sm120_fma_ownership(route: Sm120FmaRoute) -> ResolvedOutputOwnership {
if route.splits == 1 {
ResolvedOutputOwnership::OneCtaPerOutputTile
} else {
ResolvedOutputOwnership::OwnerCtaPerOutputTileFixedSplitFold
}
}
fn prepare_tf32_splitk_f32(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
output_resources: F32LaunchResourceSnapshot,
route: Tf32PhysicalRoute,
) -> Result<PreparedF32TriadLaunch, String> {
use cudarc::driver::DevicePtr;
let spec = tf32_splitk_spec(request.op, route)?;
let plan = tf32_splitk_launch_plan(request, spec)?;
let binding = f32_map_binding(ctx, route)?;
let allocation_domain = binding.allocation_domain;
let resources = output_resources.with_inputs(request, operands, allocation_domain)?;
let scratch_buffer = ctx.kernels.splitk_scratch_buf(&ctx.stream)?;
let (scratch, _) = scratch_buffer.device_ptr(&ctx.stream);
let counter_buffer = ctx
.kernels
.triad_kernels()
.tf32_splitk_counter_buf(&ctx.stream)?;
let (counters, _) = counter_buffer.device_ptr(&ctx.stream);
let resources = resources.with_scratch(
Some((scratch, (SPLITK_SCRATCH_CAP as u64) * 4)),
None,
Some((counters, (TF32_SPLITK_COUNTER_CAP as u64) * 4)),
allocation_domain,
)?;
let resources_digest = resources.digest(request, operands, [0; 32]);
let routes =
tf32_splitk_resolved_routes(request, operands, spec, binding, resources_digest, plan);
let resolved_launch_set = build_resolved_gemm_launch_set(&routes)?;
let PreparedTf32Params::Sm80(params) =
tf32_params(request, operands, Tf32TensorOrigins::default(), route)?
else {
return Err("portable TF32 split-K resolved a non-SM80 parameter ABI".into());
};
let managed_epoch = resources.managed_epoch();
Ok(PreparedF32TriadLaunch {
context_token: ctx.instance_token(),
stream_token: ctx.stream_token(),
request,
operands,
resources,
managed_epoch,
routes,
resolved_launch_set,
kind: PreparedF32Kind::Tf32SplitK {
route,
params,
plan,
workspace: Tf32SplitKWorkspace {
partial: scratch,
counters,
},
},
})
}
fn tf32_kernel_arguments_digest(
request: F32TriadRequest,
operands: F32TriadOperands,
symbol: &str,
maps_digest: Sha256Digest,
) -> Sha256Digest {
FramedSha256::new(b"triad-tf32-kernel-arguments.v2")
.required(b"symbol", symbol.as_bytes())
.required(b"op", &[request.op as u8])
.required(b"m", &(request.shape.m as u64).to_le_bytes())
.required(b"k", &(request.shape.k as u64).to_le_bytes())
.required(b"n", &(request.shape.n as u64).to_le_bytes())
.required(b"lda", &(request.shape.lda as u64).to_le_bytes())
.required(b"ldb", &(request.shape.ldb as u64).to_le_bytes())
.required(b"ldc", &(request.shape.ldc as u64).to_le_bytes())
.required(b"alpha", &operands.alpha.to_bits().to_le_bytes())
.required(b"beta", &operands.beta.to_bits().to_le_bytes())
.required(b"bias-null", &[u8::from(operands.bias.is_none())])
.required(b"tensor-maps", &maps_digest)
.finish()
}
fn physical_prepared_f32_route(
prepared: &PreparedF32TriadLaunch,
route: ResolvedGemmRoute,
) -> ResolvedGemmRoute {
let mut physical = route;
physical.resources_digest = prepared.resources.physical_digest();
if let PreparedF32Kind::Tf32 { maps, .. } = &prepared.kind {
let maps_digest = maps
.as_ref()
.map(F32PreparedTensorMaps::physical_identity_digest)
.unwrap_or([0; 32]);
physical.tensor_maps_digest = maps_digest;
}
physical
}
fn stream_is_capturing(ctx: &GpuCtx) -> Result<bool, String> {
let status = ctx
.stream
.capture_status()
.map_err(|error| format!("query route proof capture status: {error:?}"))?;
Ok(status != cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE)
}
fn resident_ctas(
function: &cudarc::driver::CudaFunction,
symbol: &str,
threads: u32,
dynamic_shared_bytes: u32,
) -> Result<u32, String> {
function
.occupancy_max_active_blocks_per_multiprocessor(
threads,
dynamic_shared_bytes as usize,
None,
)
.map_err(|error| format!("query {symbol} occupancy: {error:?}"))
}
fn tc64_reference_footprint(
ctx: &GpuCtx,
op: ResolvedGemmOp,
dtype: WeightDtype,
) -> Result<super::proof::TileFootprint, String> {
let (symbol, function) = match op {
ResolvedGemmOp::Nn => ("nn_tc64", ctx.kernels.gemm_bi_nn_tc64_typed.get(dtype)),
ResolvedGemmOp::Tn => ("tn_tc64", ctx.kernels.gemm_bi_tn_tc64_typed.get(dtype)),
ResolvedGemmOp::Nt => ("nt_tc64", ctx.kernels.gemm_bi_nt_tc64_typed.get(dtype)),
};
Ok(super::proof::TileFootprint {
tile: TcTile::Tile64.extents(),
resident: resident_ctas(function, symbol, TcTile::Tile64.block_dim(), 0)?,
})
}
fn decline_unproven_route(
ctx: &GpuCtx,
key: super::proof::RouteProofKey,
reason: &str,
) -> Result<(), String> {
ctx.with_route_proofs(|ledger| ledger.record(key, super::proof::RouteProofVerdict::Declined))?;
static DECLINED_BY_WAVES: std::sync::Once = std::sync::Once::new();
crate::mamba_ssm::gpu::diagnostics::warn_once(&DECLINED_BY_WAVES, || {
format!(
"route {} is declined on this board before its proof: {reason}; the reference \
route serves its shapes",
key.candidate
)
});
Ok(())
}
fn half_wave_guard_admits(
ctx: &GpuCtx,
spec: super::sm89_half_source::Sm89HalfRuntimeSpec,
key: super::proof::RouteProofKey,
output: (usize, usize),
) -> Result<bool, String> {
if spec.schedule != super::sm89_half_source::Sm89HalfSchedule::Tiled {
return Ok(true);
}
let Some(function) = ctx.kernels.triad_sm89_half_runtime_function(spec.symbol) else {
return Ok(false);
};
let candidate = super::proof::TileFootprint {
tile: spec.tile,
resident: resident_ctas(
function,
spec.symbol,
spec.threads,
spec.dynamic_shared_bytes,
)?,
};
let reference = tc64_reference_footprint(ctx, key.op, key.dtype)?;
let multiprocessors = ctx.kernels.multiprocessor_count();
if super::proof::wave_guard_admits(output.0, output.1, candidate, reference, multiprocessors) {
return Ok(true);
}
decline_unproven_route(
ctx,
key,
&format!(
"its {}x{} tile takes more waves on {multiprocessors} multiprocessors than the tiled \
64x64 reference",
spec.tile.0, spec.tile.1
),
)?;
Ok(false)
}
pub(in crate::mamba_ssm::gpu) fn proven_candidate<Candidate, Reference>(
ctx: &GpuCtx,
key: super::proof::RouteProofKey,
output: CUptr,
elements: usize,
dtype: WeightDtype,
candidate: Candidate,
reference: Reference,
) -> Result<bool, String>
where
Candidate: FnOnce(CUptr) -> Result<(), String>,
Reference: FnOnce(CUptr) -> Result<(), String>,
{
use super::proof::RouteProofVerdict;
if let Some(verdict) = ctx.with_route_proofs(|ledger| ledger.verdict(key))? {
return Ok(verdict == RouteProofVerdict::Admitted);
}
if stream_is_capturing(ctx)? {
return Ok(false);
}
let proof = ctx.with_gemm_route_recording_suspended(|| {
super::proof::prove_bits(&ctx.stream, output, elements, dtype, candidate, reference)
})?;
let verdict = match proof {
Ok(verdict) => verdict,
Err(reason) => {
static FAILED: std::sync::Once = std::sync::Once::new();
crate::mamba_ssm::gpu::diagnostics::warn_once(&FAILED, || {
format!(
"the first-use proof of route {} could not run on this board ({reason}); \
the reference route serves it",
key.candidate
)
});
RouteProofVerdict::Declined
}
};
ctx.with_route_proofs(|ledger| ledger.record(key, verdict))?;
if verdict == RouteProofVerdict::Declined {
static DECLINED: std::sync::Once = std::sync::Once::new();
crate::mamba_ssm::gpu::diagnostics::warn_once(&DECLINED, || {
format!(
"route {} was declined on this board: its output words differ from the \
reference route of the same contract; the reference serves its shapes",
key.candidate
)
});
}
Ok(verdict == RouteProofVerdict::Admitted)
}
fn tf32_route_footprint(
ctx: &GpuCtx,
op: ResolvedGemmOp,
route: Tf32PhysicalRoute,
) -> Result<Option<super::proof::TileFootprint>, String> {
let spec = super::contract::tf32_kernel_spec(op, route)?;
let Some(function) = ctx.kernels.tf32_function(spec.symbol) else {
return Ok(None);
};
Ok(Some(super::proof::TileFootprint {
tile: spec.tile,
resident: resident_ctas(
function,
spec.symbol,
spec.threads,
spec.dynamic_shared_bytes,
)?,
}))
}
fn tf32_wave_guard_admits(
ctx: &GpuCtx,
request: F32TriadRequest,
key: super::proof::RouteProofKey,
candidate: Tf32PhysicalRoute,
reference: Tf32PhysicalRoute,
) -> Result<bool, String> {
let (Some(candidate_footprint), Some(reference_footprint)) = (
tf32_route_footprint(ctx, request.op, candidate)?,
tf32_route_footprint(ctx, request.op, reference)?,
) else {
return Ok(false);
};
let shape = request.shape;
let (rows, columns) = match request.op {
ResolvedGemmOp::Nn => (shape.m, shape.n),
ResolvedGemmOp::Tn => (shape.k, shape.n),
ResolvedGemmOp::Nt => (shape.m, shape.k),
};
let multiprocessors = ctx.kernels.multiprocessor_count();
if super::proof::wave_guard_admits(
rows,
columns,
candidate_footprint,
reference_footprint,
multiprocessors,
) {
return Ok(true);
}
decline_unproven_route(
ctx,
key,
&format!(
"its {}x{} tile takes more waves on {multiprocessors} multiprocessors than the \
portable {}x{} tile",
candidate_footprint.tile.0,
candidate_footprint.tile.1,
reference_footprint.tile.0,
reference_footprint.tile.1
),
)?;
Ok(false)
}
fn f32_triad_output_elements(request: F32TriadRequest) -> Result<usize, String> {
let shape = request.shape;
let (rows, columns) = match request.op {
ResolvedGemmOp::Nn => (shape.m, shape.n),
ResolvedGemmOp::Tn => (shape.k, shape.n),
ResolvedGemmOp::Nt => (shape.m, shape.k),
};
if rows == 0 || columns == 0 {
return Ok(0);
}
(rows - 1)
.checked_mul(shape.ldc)
.and_then(|span| span.checked_add(columns))
.ok_or_else(|| "f32 Triad output span overflows usize".to_string())
}
fn proven_tf32_route(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
candidate: Tf32PhysicalRoute,
reference: Tf32PhysicalRoute,
) -> Result<Tf32PhysicalRoute, String> {
let spec = super::contract::tf32_kernel_spec(request.op, candidate)?;
let key = super::proof::RouteProofKey {
candidate: spec.symbol,
op: request.op,
dtype: WeightDtype::F32,
dims: (request.shape.m, request.shape.k, request.shape.n),
};
if let Some(verdict) = ctx.with_route_proofs(|ledger| ledger.verdict(key))? {
return Ok(if verdict == super::proof::RouteProofVerdict::Admitted {
candidate
} else {
reference
});
}
if stream_is_capturing(ctx)? {
return Ok(reference);
}
if !tf32_wave_guard_admits(ctx, request, key, candidate, reference)? {
return Ok(reference);
}
let elements = f32_triad_output_elements(request)?;
let arm = |output: CUptr, route: Tf32PhysicalRoute| -> Result<(), String> {
let prepared =
prepare_f32_triad_forced(ctx, request, F32TriadOperands { output, ..operands }, route)?;
unsafe {
launch_prepared_f32_triad(ctx, &prepared, |_| {
Err("a TF32 proof arm has no scalar fallback".into())
})
}
};
let admitted = proven_candidate(
ctx,
key,
operands.output,
elements,
WeightDtype::F32,
|scratch| arm(scratch, candidate),
|scratch| arm(scratch, reference),
)?;
Ok(if admitted { candidate } else { reference })
}
struct ProvenScalarLaunch {
plan: ScalarDispatchPlan,
operands: F32TriadOperands,
}
impl ScalarLaunchController for ProvenScalarLaunch {
#[inline(always)]
fn enqueue(
&mut self,
_symbol: &'static str,
config: LaunchConfig,
builder: &mut ScalarLaunchArgs<'_>,
) -> Result<(), PhysicalCudaLaunchError> {
let mut observer = NoPhysicalObserver;
unsafe {
enqueue_with_physical_observation(&mut observer, builder.launch_args(), config, None)
}
}
fn plan(&self) -> ScalarDispatchPlan {
self.plan
}
fn operands(&self) -> F32TriadOperands {
self.operands
}
fn validate_operands(&self, actual: F32TriadOperands) -> Result<(), String> {
let expected = self.operands;
if actual.output != expected.output
|| actual.a != expected.a
|| actual.b != expected.b
|| actual.bias != expected.bias
|| actual.alpha.to_bits() != expected.alpha.to_bits()
|| actual.beta.to_bits() != expected.beta.to_bits()
{
return Err(
"resolved scalar operands differ from the physical launch arguments".into(),
);
}
Ok(())
}
}
fn launch_scalar_plan_raw(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
plan: ScalarDispatchPlan,
) -> Result<(), String> {
let dims = (request.shape.m, request.shape.k, request.shape.n);
let mut control = ProvenScalarLaunch { plan, operands };
match request.op {
ResolvedGemmOp::Nn => {
let mut output = RawScalarArgument(operands.output);
let scalar_operands = GemmBiFwdSubOperands {
x_ptr: operands.a,
lda: request.shape.lda,
w_ptr: operands.b,
bias_ptr: operands.bias.unwrap_or(0),
};
gemm_bi_forward_sub_with_control(
&ctx.stream,
&ctx.kernels,
&mut output,
&scalar_operands,
dims,
Some(&mut control),
)
}
ResolvedGemmOp::Tn => {
let x_saved = RawScalarArgument(operands.a);
let dy = RawScalarArgument(operands.b);
gemm_bi_backward_dw_with_control(
&ctx.stream,
&ctx.kernels,
operands.output,
&dy,
&x_saved,
dims,
Some(&mut control),
)
}
ResolvedGemmOp::Nt => {
let mut output = RawScalarArgument(operands.output);
let dy = RawScalarArgument(operands.a);
gemm_bi_backward_dx_with_control(
&ctx.stream,
&ctx.kernels,
&mut output,
&dy,
operands.b,
dims,
Some(&mut control),
)
}
}
}
fn scalar_arms_reproduce_the_request(request: F32TriadRequest, operands: F32TriadOperands) -> bool {
let dims = (request.shape.m, request.shape.k, request.shape.n);
let contiguous = F32TriadShape::contiguous(request.op, dims);
match request.op {
ResolvedGemmOp::Nn => {
request.shape
== F32TriadShape {
lda: request.shape.lda,
..contiguous
}
}
ResolvedGemmOp::Tn => {
request.shape == contiguous
&& operands.bias.is_none()
&& operands.beta.to_bits() == 1.0f32.to_bits()
}
ResolvedGemmOp::Nt => {
request.shape == contiguous
&& operands.bias.is_none()
&& operands.beta.to_bits() == 0.0f32.to_bits()
}
}
}
fn scalar_proof_key(
request: F32TriadRequest,
operands: F32TriadOperands,
candidate: ScalarDispatchPlan,
) -> Result<Option<super::proof::RouteProofKey>, String> {
let nodes = scalar_physical_nodes(request, operands, candidate)?;
Ok(nodes.first().map(|node| super::proof::RouteProofKey {
candidate: node.symbol,
op: request.op,
dtype: WeightDtype::F32,
dims: (request.shape.m, request.shape.k, request.shape.n),
}))
}
fn scalar_ledger_plan(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
) -> Result<ScalarDispatchPlan, String> {
let facts = scalar_launch_facts(&ctx.kernels);
let plan = scalar_launch_plan(facts, request, operands)?;
if plan != scalar_dispatch_plan(request, facts.multiprocessor_count)? {
return Ok(plan);
}
let Some(candidate) = scalar_proof_plan(facts, request, operands)? else {
return Ok(plan);
};
let Some(key) = scalar_proof_key(request, operands, candidate)? else {
return Ok(plan);
};
Ok(match ctx.with_route_proofs(|ledger| ledger.verdict(key))? {
Some(super::proof::RouteProofVerdict::Admitted) => candidate,
_ => plan,
})
}
fn prove_scalar_plan<Arm>(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
arm: Arm,
) -> Result<ScalarDispatchPlan, String>
where
Arm: Fn(CUptr, ScalarDispatchPlan) -> Result<(), String>,
{
let facts = scalar_launch_facts(&ctx.kernels);
let plan = scalar_launch_plan(facts, request, operands)?;
if plan != scalar_dispatch_plan(request, facts.multiprocessor_count)? {
return Ok(plan);
}
let Some(candidate) = scalar_proof_plan(facts, request, operands)? else {
return Ok(plan);
};
if !scalar_arms_reproduce_the_request(request, operands) {
return Ok(plan);
}
let Some(key) = scalar_proof_key(request, operands, candidate)? else {
return Ok(plan);
};
let elements = f32_triad_output_elements(request)?;
let admitted = proven_candidate(
ctx,
key,
operands.output,
elements,
WeightDtype::F32,
|scratch| arm(scratch, candidate),
|scratch| arm(scratch, plan),
)?;
Ok(if admitted { candidate } else { plan })
}
fn prove_scalar_route_for_cache(
ctx: &GpuCtx,
selection: F32PreparedSelection,
request: F32TriadRequest,
operands: F32TriadOperands,
) -> Result<(), String> {
let scalar = match selection {
F32PreparedSelection::ExactScalar => true,
F32PreparedSelection::Automatic => matches!(
resolve_f32_triad_auto_with_operands(
ctx.f32_triad_policy(),
request,
operands,
ctx.kernels.f32_triad_availability(),
)?,
F32TriadSelection::ScalarFma
),
F32PreparedSelection::Forced(_) => false,
};
if !scalar || request.shape.reduction(request.op) == 0 {
return Ok(());
}
prove_scalar_plan(ctx, request, operands, |output, plan| {
launch_scalar_plan_raw(ctx, request, F32TriadOperands { output, ..operands }, plan)
})
.map(drop)
}
pub(in crate::mamba_ssm::gpu) fn prepare_f32_triad(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
) -> Result<PreparedF32TriadLaunch, String> {
request.shape.validate(request.op)?;
validate_f32_triad_operands(request, operands)?;
require_f32_preparation_outside_capture(ctx)?;
let allocation_domain = validated_allocation_domain(&ctx.stream, &ctx.kernels, "f32 Triad")?;
let output_resources =
F32LaunchResourceSnapshot::query_output(request, operands, allocation_domain)?;
let reduction = request.shape.reduction(request.op);
let zero_reduction = reduction == 0;
if zero_reduction {
let maps = F32PreparedTensorMaps::zero_reduction(
request,
None,
None,
Tf32TensorMapFormat::Tfloat32,
);
let prepared = prepare_scalar_zero_f32(ctx, request, operands, output_resources, maps)?;
return Ok(prepared);
}
match resolve_f32_triad_auto_with_operands(
ctx.f32_triad_policy(),
request,
operands,
ctx.kernels.f32_triad_availability(),
)? {
F32TriadSelection::ScalarFma => {
prepare_scalar_f32(ctx, request, operands, output_resources)
}
F32TriadSelection::Tf32Proof {
candidate,
reference,
} => {
let route = proven_tf32_route(ctx, request, operands, candidate, reference)?;
prepare_tf32_f32(ctx, request, operands, output_resources, route)
}
F32TriadSelection::Tf32(route) => {
if let Ok(spec) = super::contract::tf32_kernel_spec(request.op, route)
&& let Some(reason) = ctx
.kernels
.triad_kernels()
.tf32_symbol_exclusion(spec.symbol)
{
static EXCLUDED: std::sync::Once = std::sync::Once::new();
crate::mamba_ssm::gpu::diagnostics::warn_once(&EXCLUDED, || {
format!(
"TF32 route {} is excluded on this toolkit ({reason}); the exact \
family serves this shape",
spec.symbol
)
});
return match super::dispatch::exact_or_scalar_selection(
request,
Some(operands),
ctx.kernels.f32_triad_availability(),
) {
F32TriadSelection::ExactSm120Fma(route) => prepare_tf32_f32(
ctx,
request,
operands,
output_resources,
Tf32PhysicalRoute::Sm120TmaFmaExact(route),
),
_ => prepare_scalar_f32(ctx, request, operands, output_resources),
};
}
prepare_tf32_f32(ctx, request, operands, output_resources, route)
}
F32TriadSelection::ExactSm120Fma(route) => prepare_tf32_f32(
ctx,
request,
operands,
output_resources,
Tf32PhysicalRoute::Sm120TmaFmaExact(route),
),
}
}
fn prepare_exact_scalar_f32_triad(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
) -> Result<PreparedF32TriadLaunch, String> {
request.shape.validate(request.op)?;
validate_f32_triad_operands(request, operands)?;
require_f32_preparation_outside_capture(ctx)?;
let allocation_domain = validated_allocation_domain(&ctx.stream, &ctx.kernels, "f32 Triad")?;
let output_resources =
F32LaunchResourceSnapshot::query_output(request, operands, allocation_domain)?;
if request.shape.reduction(request.op) == 0 {
let maps = F32PreparedTensorMaps::zero_reduction(
request,
None,
None,
Tf32TensorMapFormat::Tfloat32,
);
return prepare_scalar_zero_f32(ctx, request, operands, output_resources, maps);
}
prepare_scalar_f32(ctx, request, operands, output_resources)
}
impl F32PreparedLaunchCache {
fn prepare(
ctx: &GpuCtx,
selection: F32PreparedSelection,
request: F32TriadRequest,
operands: F32TriadOperands,
) -> Result<PreparedF32TriadLaunch, String> {
match selection {
F32PreparedSelection::Automatic => prepare_f32_triad(ctx, request, operands),
F32PreparedSelection::ExactScalar => {
prepare_exact_scalar_f32_triad(ctx, request, operands)
}
F32PreparedSelection::Forced(route) => {
prepare_f32_triad_forced(ctx, request, operands, route)
}
}
}
fn ensure_prepared(
&mut self,
ctx: &GpuCtx,
selection: F32PreparedSelection,
request: F32TriadRequest,
operands: F32TriadOperands,
) -> Result<&PreparedF32TriadLaunch, String> {
let key = PreparedF32Key::new(
ctx.instance_token(),
ctx.gemm_policy(),
selection,
request,
operands,
);
let stale = match self.entries.get_mut(&key) {
Some(prepared) => {
let mut managed_epoch = prepared.managed_epoch.take();
let validation = refresh_cached_validation(
&mut managed_epoch,
|| {
ctx.stream
.capture_status()
.map(|status| {
status
!= cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE
})
.map_err(|error| {
format!("query f32 Triad capture status: {error:?}")
})
},
|| {
let refreshed = prepared.resources.managed_epoch();
validate_prepared_f32_triad(ctx, prepared)?;
Ok(refreshed)
},
);
prepared.managed_epoch = managed_epoch;
validation.is_err()
}
None => false,
};
if !self.entries.contains_key(&key) || stale {
let capturing = ctx
.stream
.capture_status()
.map_err(|error| format!("query f32 Triad capture status: {error:?}"))?
!= cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE;
if capturing {
let reason = if stale { "stale" } else { "missing" };
return Err(format!(
"prepared f32 Triad cache entry is {reason} during graph capture; run eager warmup again"
));
}
prove_scalar_route_for_cache(ctx, selection, request, operands)?;
let prepared = Self::prepare(ctx, selection, request, operands)?;
make_room_in_bounded_cache(
&mut self.entries,
&key,
F32_PREPARED_CACHE_LIMIT,
|cached| {
cached
.managed_epoch
.as_ref()
.is_none_or(ManagedAllocationEpochStamp::is_current)
},
);
if !self.entries.contains_key(&key) {
self.entries
.try_reserve(1)
.map_err(|error| format!("reserve prepared f32 Triad cache: {error}"))?;
}
self.entries.insert(key, Box::new(prepared));
}
Ok(self
.entries
.get(&key)
.expect("prepared f32 Triad cache entry was inserted above"))
}
fn launch<Scalar>(
&mut self,
ctx: &GpuCtx,
selection: F32PreparedSelection,
request: F32TriadRequest,
operands: F32TriadOperands,
scalar: Scalar,
) -> Result<(), String>
where
Scalar: FnOnce(&mut ScalarLaunchControl<'_>) -> Result<(), String>,
{
let prepared = self.ensure_prepared(ctx, selection, request, operands)?;
unsafe { enqueue_validated_prepared_f32_triad(ctx, prepared, scalar) }
}
fn launch_observed<O, Scalar>(
&mut self,
ctx: &GpuCtx,
selection: F32PreparedSelection,
request: F32TriadRequest,
operands: F32TriadOperands,
physical: (&mut O, PolicyDtype),
scalar: Scalar,
) -> Result<(), String>
where
O: PhysicalLaunchObserver,
Scalar: FnOnce(&mut PhysicalScalarLaunchControl<'_, O>) -> Result<(), String>,
{
let prepared = self.ensure_prepared(ctx, selection, request, operands)?;
unsafe {
enqueue_validated_prepared_f32_triad_observed(
ctx, prepared, physical.0, physical.1, scalar,
)
}
}
}
fn refresh_cached_validation<Capture, Validate>(
stamp: &mut Option<ManagedAllocationEpochStamp>,
capture_status: Capture,
validate: Validate,
) -> Result<(), String>
where
Capture: FnOnce() -> Result<bool, String>,
Validate: FnOnce() -> Result<Option<ManagedAllocationEpochStamp>, String>,
{
if stamp
.as_ref()
.is_some_and(ManagedAllocationEpochStamp::is_current)
{
return Ok(());
}
if stamp.is_some() && capture_status()? {
return Err(
"prepared f32 Triad allocation epoch changed during graph capture; run eager warmup again"
.into(),
);
}
let refreshed = validate()?;
if refreshed
.as_ref()
.is_some_and(|candidate| !candidate.is_current())
{
return Err("managed CUDA allocation epoch changed during f32 Triad validation".into());
}
*stamp = refreshed;
Ok(())
}
fn launch_cached_f32_triad<Scalar>(
ctx: &GpuCtx,
selection: F32PreparedSelection,
request: F32TriadRequest,
operands: F32TriadOperands,
scalar: Scalar,
) -> Result<(), String>
where
Scalar: FnOnce(&mut ScalarLaunchControl<'_>) -> Result<(), String>,
{
ctx.with_f32_prepared_launches(|cache| cache.launch(ctx, selection, request, operands, scalar))
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum FixedSm120ExactTmaTile {
M128N64,
M64N128,
}
pub(in crate::mamba_ssm::gpu) fn launch_cached_fixed_sm120_exact_tma<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
shape: (usize, usize, usize),
operands: F32TriadOperands,
tile: FixedSm120ExactTmaTile,
observer: &mut O,
) -> Result<bool, String> {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, shape),
};
let route = Tf32PhysicalRoute::Sm120TmaFmaExact(Sm120FmaRoute {
tile: match tile {
FixedSm120ExactTmaTile::M128N64 => Sm120FmaTile::M128N64,
FixedSm120ExactTmaTile::M64N128 => Sm120FmaTile::M64N128,
},
kvec: false,
splits: 1,
});
if resolve_tf32_forced(request, ctx.kernels.f32_triad_availability(), route).is_err() {
return Ok(false);
}
let spec = tf32_kernel_spec(request.op, route)?;
if ctx
.kernels
.triad_kernels()
.tf32_function(spec.symbol)
.is_none()
{
return Ok(false);
}
launch_cached_f32_triad_observed(
ctx,
F32PreparedSelection::Forced(route),
request,
operands,
(observer, PolicyDtype::F32),
|_| Err("Fixed exact-TMA bridge unexpectedly entered the scalar fallback".into()),
)?;
Ok(true)
}
pub(in crate::mamba_ssm::gpu) fn with_cached_f32_triad_prepared<R>(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
use_prepared: impl FnOnce(&PreparedF32TriadLaunch) -> Result<R, String>,
) -> Result<R, String> {
ctx.with_f32_prepared_launches(|cache| {
let prepared =
cache.ensure_prepared(ctx, F32PreparedSelection::Automatic, request, operands)?;
use_prepared(prepared)
})
}
fn launch_cached_f32_triad_observed<O, Scalar>(
ctx: &GpuCtx,
selection: F32PreparedSelection,
request: F32TriadRequest,
operands: F32TriadOperands,
physical: (&mut O, PolicyDtype),
scalar: Scalar,
) -> Result<(), String>
where
O: PhysicalLaunchObserver,
Scalar: FnOnce(&mut PhysicalScalarLaunchControl<'_, O>) -> Result<(), String>,
{
ctx.with_f32_prepared_launches(|cache| {
cache.launch_observed(ctx, selection, request, operands, physical, scalar)
})
}
unsafe fn launch_cached_f32_forward_ptrs_selected(
ctx: &GpuCtx,
y: CUptr,
x: CUptr,
w_ptr: CUptr,
bias_ptr: CUptr,
dims: (usize, usize, usize),
selection: F32PreparedSelection,
) -> Result<(), String> {
let shape = F32TriadShape::contiguous(ResolvedGemmOp::Nn, dims);
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape,
};
let reduction_is_zero = shape.reduction(request.op) == 0;
let output = y;
let x_ptr = if reduction_is_zero { 0 } else { x };
let w_ptr = if reduction_is_zero { 0 } else { w_ptr };
let operands = F32TriadOperands {
output,
a: x_ptr,
b: w_ptr,
bias: (bias_ptr != 0).then_some(bias_ptr),
alpha: 1.0,
beta: 0.0,
};
let scalar_operands = GemmBiFwdSubOperands {
x_ptr,
lda: shape.lda,
w_ptr,
bias_ptr,
};
let mut output_arg = RawScalarArgument(output);
launch_cached_f32_triad(ctx, selection, request, operands, |control| {
gemm_bi_forward_sub_with_control(
&ctx.stream,
&ctx.kernels,
&mut output_arg,
&scalar_operands,
dims,
Some(control),
)
})
}
pub(crate) unsafe fn launch_cached_f32_forward_ptrs(
ctx: &GpuCtx,
y: CUptr,
x: CUptr,
w: CUptr,
bias: CUptr,
dims: (usize, usize, usize),
) -> Result<(), String> {
unsafe {
launch_cached_f32_forward_ptrs_selected(
ctx,
y,
x,
w,
bias,
dims,
F32PreparedSelection::Automatic,
)
}
}
#[cfg(test)]
pub(crate) fn launch_cached_f32_forward(
ctx: &GpuCtx,
y: &mut GpuBuffer,
x: &GpuBuffer,
w_ptr: CUptr,
bias_ptr: CUptr,
dims: (usize, usize, usize),
) -> Result<(), String> {
unsafe {
launch_cached_f32_forward_ptrs(ctx, y.cached_ptr(), x.cached_ptr(), w_ptr, bias_ptr, dims)
}
}
#[derive(Clone, Copy)]
pub(crate) struct ScalarFallbackPhysicalContext {
pub(crate) dims: (usize, usize, usize),
pub(crate) dtype: WeightDtype,
}
enum ScalarPhysicalGraphDispatch<'a> {
Nn {
y: &'a mut GpuBuffer,
x: &'a GpuBuffer,
w_ptr: CUptr,
bias_ptr: CUptr,
dims: (usize, usize, usize),
},
Tn {
dw_ptr: CUptr,
dy: &'a GpuBuffer,
x_saved: &'a GpuBuffer,
dims: (usize, usize, usize),
},
Nt {
dx: &'a mut GpuBuffer,
dy: &'a GpuBuffer,
w_ptr: CUptr,
dims: (usize, usize, usize),
},
}
fn prepare_exact_scalar_physical_graph_sequence<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &O,
request: F32TriadRequest,
operands: F32TriadOperands,
logical_dtype: PolicyDtype,
dispatch: ScalarPhysicalGraphDispatch<'_>,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
ctx.with_f32_prepared_launches(|cache| {
let prepared =
cache.ensure_prepared(ctx, F32PreparedSelection::ExactScalar, request, operands)?;
prepare_prepared_scalar_physical_graph_sequence(
ctx,
observer,
prepared,
logical_dtype,
dispatch,
)
})
}
fn prepare_prepared_scalar_physical_graph_sequence<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &O,
prepared: &PreparedF32TriadLaunch,
logical_dtype: PolicyDtype,
dispatch: ScalarPhysicalGraphDispatch<'_>,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
validate_prepared_f32_triad(ctx, prepared)?;
let PreparedF32Kind::Scalar(plan) = prepared.kind else {
return Err("prepared scalar graph package resolved a non-scalar F32 launch".into());
};
let mut launches = Vec::new();
launches
.try_reserve_exact(prepared.routes.len())
.map_err(|error| format!("reserve prepared scalar physical launches: {error}"))?;
let base = ScalarLaunchControl {
ctx,
routes: &prepared.routes,
plan,
operands: prepared.operands,
next: 0,
};
let mut control = PreparedPhysicalScalarLaunchControl {
base,
logical_dtype,
physical_resources_digest: prepared.resources.physical_digest(),
observer,
launches,
};
match dispatch {
ScalarPhysicalGraphDispatch::Nn {
y,
x,
w_ptr,
bias_ptr,
dims,
} => {
let scalar_operands = GemmBiFwdSubOperands {
x_ptr: prepared.operands.a,
lda: prepared.request.shape.lda,
w_ptr,
bias_ptr,
};
gemm_bi_forward_sub_with_control(
&ctx.stream,
&ctx.kernels,
y,
&scalar_operands,
dims,
Some(&mut control),
)?;
if x.raw_ptr(&ctx.stream) != prepared.operands.a {
return Err("prepared scalar NN input binding changed".into());
}
}
ScalarPhysicalGraphDispatch::Tn {
dw_ptr,
dy,
x_saved,
dims,
} => gemm_bi_backward_dw_with_control(
&ctx.stream,
&ctx.kernels,
dw_ptr,
dy,
x_saved,
dims,
Some(&mut control),
)?,
ScalarPhysicalGraphDispatch::Nt {
dx,
dy,
w_ptr,
dims,
} => gemm_bi_backward_dx_with_control(
&ctx.stream,
&ctx.kernels,
dx,
dy,
w_ptr,
dims,
Some(&mut control),
)?,
}
control.finish()
}
pub(in crate::mamba_ssm::gpu) fn prepare_prepared_f32_forward_graph_sequence<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &O,
prepared: &PreparedF32TriadLaunch,
y: &mut GpuBuffer,
x: &GpuBuffer,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
if prepared.request.op != ResolvedGemmOp::Nn {
return Err("prepared F32 forward graph sequence requires an NN launch".into());
}
prepare_prepared_scalar_physical_graph_sequence(
ctx,
observer,
prepared,
PolicyDtype::F32,
ScalarPhysicalGraphDispatch::Nn {
y,
x,
w_ptr: prepared.operands.b,
bias_ptr: prepared.operands.bias.unwrap_or(0),
dims: (
prepared.request.shape.m,
prepared.request.shape.k,
prepared.request.shape.n,
),
},
)
}
pub(in crate::mamba_ssm::gpu) fn prepare_prepared_f32_backward_dw_graph_sequence<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &O,
prepared: &PreparedF32TriadLaunch,
dy: &GpuBuffer,
x_saved: &GpuBuffer,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
if prepared.request.op != ResolvedGemmOp::Tn {
return Err("prepared F32 backward-dW graph sequence requires a TN launch".into());
}
prepare_prepared_scalar_physical_graph_sequence(
ctx,
observer,
prepared,
PolicyDtype::F32,
ScalarPhysicalGraphDispatch::Tn {
dw_ptr: prepared.operands.output,
dy,
x_saved,
dims: (
prepared.request.shape.m,
prepared.request.shape.k,
prepared.request.shape.n,
),
},
)
}
pub(in crate::mamba_ssm::gpu) fn prepare_prepared_f32_backward_dx_graph_sequence<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &O,
prepared: &PreparedF32TriadLaunch,
dx: &mut GpuBuffer,
dy: &GpuBuffer,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
if prepared.request.op != ResolvedGemmOp::Nt {
return Err("prepared F32 backward-dX graph sequence requires an NT launch".into());
}
prepare_prepared_scalar_physical_graph_sequence(
ctx,
observer,
prepared,
PolicyDtype::F32,
ScalarPhysicalGraphDispatch::Nt {
dx,
dy,
w_ptr: prepared.operands.b,
dims: (
prepared.request.shape.m,
prepared.request.shape.k,
prepared.request.shape.n,
),
},
)
}
pub(in crate::mamba_ssm::gpu) fn prepare_prepared_f32_direct_graph_sequence<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &O,
prepared: &PreparedF32TriadLaunch,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
validate_prepared_f32_triad(ctx, prepared)?;
if matches!(prepared.kind, PreparedF32Kind::Tf32TnPreRna { .. }) {
return prepare_sm89_tf32_tn_pre_rna_graph_sequence(ctx, observer, prepared);
}
if let PreparedF32Kind::Tf32SplitK {
params,
plan,
workspace,
..
} = &prepared.kind
{
return prepare_tf32_splitk_direct_graph_sequence(
ctx,
observer,
prepared,
params,
*plan,
workspace.partial,
workspace.counters,
);
}
if prepared.routes.len() != 1 {
return Err("prepared direct F32 graph sequence requires exactly one route".into());
}
let route = physical_prepared_f32_route(prepared, prepared.routes[0]);
let config = LaunchConfig {
grid_dim: route.launch.grid_dim,
block_dim: route.launch.block_dim,
shared_mem_bytes: route.launch.shared_mem_bytes,
};
let observation = PhysicalLaunchObservation::gemm(PolicyDtype::F32, None, route);
let node = resolve_physical_launch_observation(observer, observation, config)?;
let output = prepared.operands.output;
let bias = prepared.operands.bias.unwrap_or(0);
let mut arguments = PhysicalScalarKernelArguments::new();
let function = match &prepared.kind {
PreparedF32Kind::Scalar(_) => {
return Err(
"prepared scalar F32 graph sequence requires its operation-level producer".into(),
);
}
PreparedF32Kind::ScalarZero { params, .. } => {
let null_input = 0_u64;
arguments.push(output)?;
arguments.push(null_input)?;
arguments.push(null_input)?;
arguments.push(bias)?;
arguments.push(*params)?;
let kernels = ctx.kernels.triad_kernels();
match prepared.request.op {
ResolvedGemmOp::Nn => kernels.gemm_bi_nn_zero_reduction.clone(),
ResolvedGemmOp::Tn => kernels.gemm_bi_tn_zero_reduction.clone(),
ResolvedGemmOp::Nt => kernels.gemm_bi_nt_zero_reduction.clone(),
}
}
PreparedF32Kind::Tf32 {
route: physical_route,
maps,
params,
..
} => {
match (*physical_route, *params) {
(
Tf32PhysicalRoute::MmaTf32Rna(_)
| Tf32PhysicalRoute::Sm89MmaTf32Compact8
| Tf32PhysicalRoute::Sm89NnDirectN96
| Tf32PhysicalRoute::Sm89NnN96
| Tf32PhysicalRoute::Sm89NtALdmatrixN96
| Tf32PhysicalRoute::Sm89NtRnaM144N96S2
| Tf32PhysicalRoute::Sm89NtRowstageM128N192S2
| Tf32PhysicalRoute::Sm89TnDirectM192N192S2,
PreparedTf32Params::Sm80(params),
) => {
let reduction_is_zero =
prepared.request.shape.reduction(prepared.request.op) == 0;
arguments.push(output)?;
arguments.push(if reduction_is_zero {
0
} else {
prepared.operands.a
})?;
arguments.push(if reduction_is_zero {
0
} else {
prepared.operands.b
})?;
arguments.push(bias)?;
arguments.push(params)?;
}
(Tf32PhysicalRoute::Sm90aWgmmaTf32Tma(_), PreparedTf32Params::Sm90a(params)) => {
let maps = maps
.as_ref()
.ok_or_else(|| "SM90a TF32 graph sequence has no tensor maps".to_string())?
.maps();
arguments.push(output)?;
arguments.push(maps[0])?;
arguments.push(maps[1])?;
arguments.push(bias)?;
arguments.push(params)?;
}
(Tf32PhysicalRoute::Sm100Tcgen05Tf32Tma(_), PreparedTf32Params::Sm100(params)) => {
let maps = maps
.as_ref()
.ok_or_else(|| "SM100 TF32 graph sequence has no tensor maps".to_string())?
.maps();
arguments.push(output)?;
arguments.push(maps[0])?;
arguments.push(maps[1])?;
arguments.push(bias)?;
arguments.push(params)?;
}
(Tf32PhysicalRoute::Sm120TmaMmaTf32Rna(_), PreparedTf32Params::Sm120(params)) => {
let maps = maps
.as_ref()
.ok_or_else(|| "SM120 TF32 graph sequence has no tensor maps".to_string())?
.maps();
arguments.push(output)?;
arguments.push(maps[0])?;
arguments.push(maps[1])?;
arguments.push(bias)?;
arguments.push(params)?;
}
_ => return Err("prepared TF32 graph route and parameter ABI disagree".into()),
}
ctx.kernels
.triad_kernels()
.tf32_function(route.symbol)
.ok_or_else(|| format!("qualified TF32 symbol {} is unavailable", route.symbol))?
.clone()
}
PreparedF32Kind::Tf32SplitK { .. } => {
return Err("prepared TF32 split-K graph sequence was not expanded".into());
}
PreparedF32Kind::Tf32StreamK {
maps,
params,
workspace,
..
} => {
let maps = maps.maps();
arguments.push(output)?;
arguments.push(workspace.partial)?;
arguments.push(workspace.flags)?;
arguments.push(maps[0])?;
arguments.push(maps[1])?;
arguments.push(bias)?;
match params {
PreparedTf32Params::Sm120(params) => arguments.push(*params)?,
PreparedTf32Params::Sm120Fma(params) => arguments.push(*params)?,
_ => {
return Err("prepared TF32 stream-K route and parameter ABI disagree".into());
}
}
ctx.kernels
.triad_kernels()
.tf32_function(route.symbol)
.ok_or_else(|| format!("qualified TF32 symbol {} is unavailable", route.symbol))?
.clone()
}
PreparedF32Kind::Tf32TnPreRna { .. } => {
return Err("Ada TF32 pre-RNA graph sequence was not expanded".into());
}
};
Ok(PreparedTriadPhysicalGraphSequence {
launches: vec![PreparedTriadPhysicalGraphLaunch {
function,
config,
node,
arguments: Box::new(arguments),
}]
.into_boxed_slice(),
})
}
fn prepare_sm89_tf32_tn_pre_rna_graph_sequence<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &O,
prepared: &PreparedF32TriadLaunch,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
let PreparedF32Kind::Tf32TnPreRna {
transpose_params,
gemm_params,
transpose_config,
gemm_config,
scratch,
scratch_elements,
transform,
..
} = &prepared.kind
else {
return Err("prepared launch is not an Ada TF32 pre-RNA pipeline".into());
};
let [resolved] = prepared.routes.as_ref() else {
return Err("Ada TF32 pre-RNA graph requires exactly one GEMM route".into());
};
let scratch_bytes = u64::try_from(*scratch_elements)
.ok()
.and_then(|elements| elements.checked_mul(4))
.ok_or_else(|| "Ada TF32 graph scratch span overflows u64".to_string())?;
let transform_observation = PhysicalLaunchObservation::input_transform(
TN_PRE_RNA_TRANSPOSE_SYMBOL,
prepared.request.op,
PolicyDtype::F32,
(
prepared.request.shape.m,
prepared.request.shape.k,
prepared.request.shape.n,
),
(
prepared.request.shape.lda,
prepared.request.shape.ldb,
prepared.request.shape.ldc,
),
**transform,
PhysicalConversionArguments::new(
prepared.operands.a,
sm89_tf32_tn_source_bytes(prepared.request)?,
*scratch,
scratch_bytes,
),
);
let transform_node =
resolve_physical_launch_observation(observer, transform_observation, *transpose_config)?;
let physical = physical_prepared_f32_route(prepared, *resolved);
let gemm_node = resolve_physical_launch_observation(
observer,
PhysicalLaunchObservation::gemm(PolicyDtype::F32, None, physical),
*gemm_config,
)?;
let mut transform_arguments = PhysicalScalarKernelArguments::new();
transform_arguments.push(prepared.operands.a)?;
transform_arguments.push(*scratch)?;
transform_arguments.push(*transpose_params)?;
let mut gemm_arguments = PhysicalScalarKernelArguments::new();
gemm_arguments.push(prepared.operands.output)?;
gemm_arguments.push(*scratch)?;
gemm_arguments.push(prepared.operands.b)?;
gemm_arguments.push(0_u64)?;
gemm_arguments.push(*gemm_params)?;
let transform_function = ctx
.kernels
.triad_sm89_tf32_joint_function(TN_PRE_RNA_TRANSPOSE_SYMBOL)
.ok_or_else(|| "qualified Ada TF32 transpose symbol is unavailable".to_string())?
.clone();
let gemm_function = ctx
.kernels
.triad_sm89_tf32_joint_function(resolved.symbol)
.ok_or_else(|| {
format!(
"qualified Ada TF32 symbol {} is unavailable",
resolved.symbol
)
})?
.clone();
Ok(PreparedTriadPhysicalGraphSequence {
launches: vec![
PreparedTriadPhysicalGraphLaunch {
function: transform_function,
config: *transpose_config,
node: transform_node,
arguments: Box::new(transform_arguments),
},
PreparedTriadPhysicalGraphLaunch {
function: gemm_function,
config: *gemm_config,
node: gemm_node,
arguments: Box::new(gemm_arguments),
},
]
.into_boxed_slice(),
})
}
fn prepare_tf32_splitk_direct_graph_sequence<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &O,
prepared: &PreparedF32TriadLaunch,
params: &Sm80Tf32KernelParams,
plan: Tf32SplitKLaunchPlan,
scratch: CUptr,
counters: CUptr,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
let [resolved] = prepared.routes.as_ref() else {
return Err("prepared fused TF32 split-K graph sequence requires one route".into());
};
let config = plan.fused;
let route = physical_prepared_f32_route(prepared, *resolved);
if route.launch.grid_dim != config.grid_dim
|| route.launch.block_dim != config.block_dim
|| route.launch.shared_mem_bytes != config.shared_mem_bytes
{
return Err("prepared fused TF32 split-K graph launch configuration changed".into());
}
let observation = PhysicalLaunchObservation::gemm(PolicyDtype::F32, None, route);
let node = resolve_physical_launch_observation(observer, observation, config)?;
let mut arguments = PhysicalScalarKernelArguments::new();
arguments.push(prepared.operands.output)?;
arguments.push(scratch)?;
arguments.push(counters)?;
arguments.push(prepared.operands.a)?;
arguments.push(prepared.operands.b)?;
arguments.push(prepared.operands.bias.unwrap_or(0))?;
arguments.push(*params)?;
let function = ctx
.kernels
.triad_kernels()
.tf32_splitk_function(route.symbol)
.ok_or_else(|| {
format!(
"qualified fused TF32 split-K symbol {} is unavailable",
route.symbol
)
})?
.clone();
Ok(PreparedTriadPhysicalGraphSequence {
launches: vec![PreparedTriadPhysicalGraphLaunch {
function,
config,
node,
arguments: Box::new(arguments),
}]
.into_boxed_slice(),
})
}
pub(in crate::mamba_ssm::gpu) fn prepare_exact_scalar_f32_forward_graph_sequence<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &O,
y: &mut GpuBuffer,
x: &GpuBuffer,
w_ptr: CUptr,
bias_ptr: CUptr,
physical: ScalarFallbackPhysicalContext,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
let ScalarFallbackPhysicalContext { dims, dtype } = physical;
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, dims),
};
let operands = F32TriadOperands {
output: y.raw_ptr(&ctx.stream),
a: x.raw_ptr(&ctx.stream),
b: w_ptr,
bias: (bias_ptr != 0).then_some(bias_ptr),
alpha: 1.0,
beta: 0.0,
};
prepare_exact_scalar_physical_graph_sequence(
ctx,
observer,
request,
operands,
half_policy_dtype(dtype)?,
ScalarPhysicalGraphDispatch::Nn {
y,
x,
w_ptr,
bias_ptr,
dims,
},
)
}
pub(in crate::mamba_ssm::gpu) fn prepare_exact_scalar_f32_backward_dw_graph_sequence<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &O,
dw_ptr: CUptr,
dy: &GpuBuffer,
x_saved: &GpuBuffer,
physical: ScalarFallbackPhysicalContext,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
let ScalarFallbackPhysicalContext { dims, dtype } = physical;
let request = F32TriadRequest {
op: ResolvedGemmOp::Tn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Tn, dims),
};
let operands = F32TriadOperands {
output: dw_ptr,
a: x_saved.raw_ptr(&ctx.stream),
b: dy.raw_ptr(&ctx.stream),
bias: None,
alpha: 1.0,
beta: 1.0,
};
prepare_exact_scalar_physical_graph_sequence(
ctx,
observer,
request,
operands,
half_policy_dtype(dtype)?,
ScalarPhysicalGraphDispatch::Tn {
dw_ptr,
dy,
x_saved,
dims,
},
)
}
pub(in crate::mamba_ssm::gpu) fn prepare_exact_scalar_f32_backward_dx_graph_sequence<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &O,
dx: &mut GpuBuffer,
dy: &GpuBuffer,
w_ptr: CUptr,
physical: ScalarFallbackPhysicalContext,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
let ScalarFallbackPhysicalContext { dims, dtype } = physical;
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, dims),
};
let operands = F32TriadOperands {
output: dx.raw_ptr(&ctx.stream),
a: dy.raw_ptr(&ctx.stream),
b: w_ptr,
bias: None,
alpha: 1.0,
beta: 0.0,
};
prepare_exact_scalar_physical_graph_sequence(
ctx,
observer,
request,
operands,
half_policy_dtype(dtype)?,
ScalarPhysicalGraphDispatch::Nt {
dx,
dy,
w_ptr,
dims,
},
)
}
pub(in crate::mamba_ssm::gpu) fn record_physical_exact_scalar_f32_forward<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &mut O,
y: &mut GpuBuffer,
x: &GpuBuffer,
w_ptr: CUptr,
bias_ptr: CUptr,
physical: ScalarFallbackPhysicalContext,
) -> Result<(), String> {
let ScalarFallbackPhysicalContext { dims, dtype } = physical;
let shape = F32TriadShape::contiguous(ResolvedGemmOp::Nn, dims);
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape,
};
let operands = F32TriadOperands {
output: y.raw_ptr(&ctx.stream),
a: x.raw_ptr(&ctx.stream),
b: w_ptr,
bias: (bias_ptr != 0).then_some(bias_ptr),
alpha: 1.0,
beta: 0.0,
};
let scalar_operands = GemmBiFwdSubOperands {
x_ptr: operands.a,
lda: shape.lda,
w_ptr,
bias_ptr,
};
if O::ENABLED {
let logical_dtype = half_policy_dtype(dtype)?;
return launch_cached_f32_triad_observed(
ctx,
F32PreparedSelection::ExactScalar,
request,
operands,
(observer, logical_dtype),
|control| {
gemm_bi_forward_sub_with_control(
&ctx.stream,
&ctx.kernels,
y,
&scalar_operands,
dims,
Some(control),
)
},
);
}
launch_cached_f32_triad(
ctx,
F32PreparedSelection::ExactScalar,
request,
operands,
|control| {
gemm_bi_forward_sub_with_control(
&ctx.stream,
&ctx.kernels,
y,
&scalar_operands,
dims,
Some(control),
)
},
)
}
fn launch_cached_f32_backward_dw_selected(
ctx: &GpuCtx,
dw_ptr: CUptr,
dy: &GpuBuffer,
x_saved: &GpuBuffer,
dims: (usize, usize, usize),
selection: F32PreparedSelection,
) -> Result<(), String> {
let shape = F32TriadShape::contiguous(ResolvedGemmOp::Tn, dims);
let request = F32TriadRequest {
op: ResolvedGemmOp::Tn,
shape,
};
let reduction_is_zero = shape.reduction(request.op) == 0;
let operands = F32TriadOperands {
output: dw_ptr,
a: if reduction_is_zero {
0
} else {
x_saved.raw_ptr(&ctx.stream)
},
b: if reduction_is_zero {
0
} else {
dy.raw_ptr(&ctx.stream)
},
bias: None,
alpha: 1.0,
beta: 1.0,
};
launch_cached_f32_triad(ctx, selection, request, operands, |control| {
gemm_bi_backward_dw_with_control(
&ctx.stream,
&ctx.kernels,
dw_ptr,
dy,
x_saved,
dims,
Some(control),
)
})
}
pub(crate) fn launch_cached_f32_backward_dw(
ctx: &GpuCtx,
dw_ptr: CUptr,
dy: &GpuBuffer,
x_saved: &GpuBuffer,
dims: (usize, usize, usize),
) -> Result<(), String> {
launch_cached_f32_backward_dw_selected(
ctx,
dw_ptr,
dy,
x_saved,
dims,
F32PreparedSelection::Automatic,
)
}
pub(in crate::mamba_ssm::gpu) fn record_physical_exact_scalar_f32_backward_dw<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &mut O,
dw_ptr: CUptr,
dy: &GpuBuffer,
x_saved: &GpuBuffer,
physical: ScalarFallbackPhysicalContext,
) -> Result<(), String> {
let ScalarFallbackPhysicalContext { dims, dtype } = physical;
let shape = F32TriadShape::contiguous(ResolvedGemmOp::Tn, dims);
let request = F32TriadRequest {
op: ResolvedGemmOp::Tn,
shape,
};
let operands = F32TriadOperands {
output: dw_ptr,
a: x_saved.raw_ptr(&ctx.stream),
b: dy.raw_ptr(&ctx.stream),
bias: None,
alpha: 1.0,
beta: 1.0,
};
if O::ENABLED {
let logical_dtype = half_policy_dtype(dtype)?;
return launch_cached_f32_triad_observed(
ctx,
F32PreparedSelection::ExactScalar,
request,
operands,
(observer, logical_dtype),
|control| {
gemm_bi_backward_dw_with_control(
&ctx.stream,
&ctx.kernels,
dw_ptr,
dy,
x_saved,
dims,
Some(control),
)
},
);
}
launch_cached_f32_triad(
ctx,
F32PreparedSelection::ExactScalar,
request,
operands,
|control| {
gemm_bi_backward_dw_with_control(
&ctx.stream,
&ctx.kernels,
dw_ptr,
dy,
x_saved,
dims,
Some(control),
)
},
)
}
unsafe fn launch_cached_f32_backward_dx_ptrs_selected(
ctx: &GpuCtx,
dx: CUptr,
dy: CUptr,
w_ptr: CUptr,
dims: (usize, usize, usize),
selection: F32PreparedSelection,
) -> Result<(), String> {
let shape = F32TriadShape::contiguous(ResolvedGemmOp::Nt, dims);
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape,
};
let reduction_is_zero = shape.reduction(request.op) == 0;
let operands = F32TriadOperands {
output: dx,
a: if reduction_is_zero { 0 } else { dy },
b: if reduction_is_zero { 0 } else { w_ptr },
bias: None,
alpha: 1.0,
beta: 0.0,
};
let mut output_arg = RawScalarArgument(dx);
let input_arg = RawScalarArgument(operands.a);
launch_cached_f32_triad(ctx, selection, request, operands, |control| {
gemm_bi_backward_dx_with_control(
&ctx.stream,
&ctx.kernels,
&mut output_arg,
&input_arg,
w_ptr,
dims,
Some(control),
)
})
}
pub(crate) unsafe fn launch_cached_f32_backward_dx_ptrs(
ctx: &GpuCtx,
dx: CUptr,
dy: CUptr,
w: CUptr,
dims: (usize, usize, usize),
) -> Result<(), String> {
unsafe {
launch_cached_f32_backward_dx_ptrs_selected(
ctx,
dx,
dy,
w,
dims,
F32PreparedSelection::Automatic,
)
}
}
pub(crate) fn launch_cached_f32_backward_dx(
ctx: &GpuCtx,
dx: &mut GpuBuffer,
dy: &GpuBuffer,
w_ptr: CUptr,
dims: (usize, usize, usize),
) -> Result<(), String> {
unsafe {
launch_cached_f32_backward_dx_ptrs(ctx, dx.cached_ptr(), dy.cached_ptr(), w_ptr, dims)
}
}
pub(in crate::mamba_ssm::gpu) fn record_physical_exact_scalar_f32_backward_dx<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &mut O,
dx: &mut GpuBuffer,
dy: &GpuBuffer,
w_ptr: CUptr,
physical: ScalarFallbackPhysicalContext,
) -> Result<(), String> {
record_physical_exact_scalar_f32_backward_dx_with_arguments(
ctx, observer, dx, dy, w_ptr, physical,
)
}
fn record_physical_exact_scalar_f32_backward_dx_with_arguments<
O: PhysicalLaunchObserver,
Output: ScalarOutputArgument,
Input: ScalarInputArgument,
>(
ctx: &GpuCtx,
observer: &mut O,
output: &mut Output,
input: &Input,
w_ptr: CUptr,
physical: ScalarFallbackPhysicalContext,
) -> Result<(), String> {
let ScalarFallbackPhysicalContext { dims, dtype } = physical;
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, dims),
};
let operands = F32TriadOperands {
output: output.scalar_ptr(),
a: input.scalar_ptr(),
b: w_ptr,
bias: None,
alpha: 1.0,
beta: 0.0,
};
if O::ENABLED {
let logical_dtype = half_policy_dtype(dtype)?;
return launch_cached_f32_triad_observed(
ctx,
F32PreparedSelection::ExactScalar,
request,
operands,
(observer, logical_dtype),
|control| {
gemm_bi_backward_dx_with_control(
&ctx.stream,
&ctx.kernels,
output,
input,
w_ptr,
dims,
Some(control),
)
},
);
}
launch_cached_f32_triad(
ctx,
F32PreparedSelection::ExactScalar,
request,
operands,
|control| {
gemm_bi_backward_dx_with_control(
&ctx.stream,
&ctx.kernels,
output,
input,
w_ptr,
dims,
Some(control),
)
},
)
}
pub(in crate::mamba_ssm::gpu) unsafe fn record_physical_exact_scalar_f32_backward_dx_ptrs<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &mut O,
dx: CUptr,
dy: CUptr,
w_ptr: CUptr,
physical: ScalarFallbackPhysicalContext,
) -> Result<(), String> {
let dims = physical.dims;
let shape = F32TriadShape::contiguous(ResolvedGemmOp::Nt, dims);
let reduction_is_zero = shape.reduction(ResolvedGemmOp::Nt) == 0;
let dy = if reduction_is_zero { 0 } else { dy };
let w_ptr = if reduction_is_zero { 0 } else { w_ptr };
let mut output = RawScalarArgument(dx);
let input = RawScalarArgument(dy);
record_physical_exact_scalar_f32_backward_dx_with_arguments(
ctx,
observer,
&mut output,
&input,
w_ptr,
physical,
)
}
pub(in crate::mamba_ssm::gpu) fn prepare_f32_triad_forced(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
route: Tf32PhysicalRoute,
) -> Result<PreparedF32TriadLaunch, String> {
request.shape.validate(request.op)?;
validate_f32_triad_operands(request, operands)?;
require_f32_preparation_outside_capture(ctx)?;
let allocation_domain = validated_allocation_domain(&ctx.stream, &ctx.kernels, "f32 Triad")?;
let output_resources =
F32LaunchResourceSnapshot::query_output(request, operands, allocation_domain)?;
prepare_tf32_f32(ctx, request, operands, output_resources, route)
}
pub(in crate::mamba_ssm::gpu) fn prepare_sm89_exact_f32_tn_forced(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
route: super::Sm89ExactF32TnRoute,
) -> Result<PreparedF32TriadLaunch, String> {
request.shape.validate(request.op)?;
validate_f32_triad_operands(request, operands)?;
require_f32_preparation_outside_capture(ctx)?;
let plan =
forced_sm89_exact_f32_plan(scalar_launch_facts(&ctx.kernels), request, operands, route)?;
let allocation_domain = validated_allocation_domain(&ctx.stream, &ctx.kernels, "f32 Triad")?;
let output_resources =
F32LaunchResourceSnapshot::query_output(request, operands, allocation_domain)?;
prepare_scalar_f32_with_plan(ctx, request, operands, output_resources, plan)
}
pub(in crate::mamba_ssm::gpu) fn prepare_sm89_exact_f32_d128_tn_forced(
ctx: &GpuCtx,
request: F32TriadRequest,
operands: F32TriadOperands,
route: super::Sm89ExactF32D128Route,
) -> Result<PreparedF32TriadLaunch, String> {
request.shape.validate(request.op)?;
validate_f32_triad_operands(request, operands)?;
require_f32_preparation_outside_capture(ctx)?;
let plan = forced_sm89_exact_f32_d128_plan(
scalar_launch_facts(&ctx.kernels),
request,
operands,
route,
)?;
let allocation_domain = validated_allocation_domain(&ctx.stream, &ctx.kernels, "f32 Triad")?;
let output_resources =
F32LaunchResourceSnapshot::query_output(request, operands, allocation_domain)?;
prepare_scalar_f32_with_plan(ctx, request, operands, output_resources, plan)
}
pub(in crate::mamba_ssm::gpu) unsafe fn launch_sm89_exact_f32_tn_forced(
ctx: &GpuCtx,
prepared: &PreparedF32TriadLaunch,
output: CUptr,
a: &GpuBuffer,
b: &GpuBuffer,
dims: (usize, usize, usize),
) -> Result<(), String> {
unsafe {
launch_prepared_f32_triad(ctx, prepared, |control| {
gemm_bi_backward_dw_with_control(
&ctx.stream,
&ctx.kernels,
output,
b,
a,
dims,
Some(control),
)
})
}
}
pub(in crate::mamba_ssm::gpu) unsafe fn launch_sm89_exact_f32_d128_tn_forced(
ctx: &GpuCtx,
prepared: &PreparedF32TriadLaunch,
output: CUptr,
a: &GpuBuffer,
b: &GpuBuffer,
dims: (usize, usize, usize),
) -> Result<(), String> {
unsafe {
launch_prepared_f32_triad(ctx, prepared, |control| {
gemm_bi_backward_dw_with_control(
&ctx.stream,
&ctx.kernels,
output,
b,
a,
dims,
Some(control),
)
})
}
}
fn validate_prepared_f32_triad(
ctx: &GpuCtx,
prepared: &PreparedF32TriadLaunch,
) -> Result<(), String> {
if prepared.context_token != ctx.instance_token() {
return Err("prepared f32 Triad launch belongs to another GPU context".into());
}
if prepared.stream_token != ctx.stream_token() {
return Err("prepared f32 Triad launch belongs to another CUDA stream".into());
}
prepared.resources.validate_live()?;
match &prepared.kind {
PreparedF32Kind::Scalar(_) => {}
PreparedF32Kind::ScalarZero { maps, .. } => maps.validate_live_allocations()?,
PreparedF32Kind::Tf32 { route, maps, .. } => {
if let Some(maps) = maps {
maps.validate_live_allocations()?;
let binding = f32_map_binding(ctx, *route)?;
if !maps.matches_binding(binding) {
return Err("prepared TF32 tensor-map binding changed before launch".into());
}
}
}
PreparedF32Kind::Tf32SplitK { route, plan, .. } => {
f32_map_binding(ctx, *route)?;
validate_tf32_splitk_prepared_layout(prepared, *route, *plan)?;
}
PreparedF32Kind::Tf32StreamK {
route, maps, plan, ..
} => {
maps.validate_live_allocations()?;
let binding = f32_map_binding(ctx, *route)?;
if !maps.matches_binding(binding) {
return Err("prepared TF32 tensor-map binding changed before launch".into());
}
validate_tf32_streamk_prepared_layout(ctx, prepared, *route, *plan)?;
}
PreparedF32Kind::Tf32TnPreRna {
route,
scratch,
scratch_elements,
transform,
..
} => {
f32_map_binding(ctx, *route)?;
let identity = prepared
.resources
.transpose_scratch_identity()
.ok_or_else(|| "prepared Ada TF32 route lost transpose scratch".to_string())?;
let required_bytes = u64::try_from(*scratch_elements)
.ok()
.and_then(|elements| elements.checked_mul(4))
.ok_or_else(|| "prepared Ada TF32 scratch extent overflows u64".to_string())?;
if !identity.matches_requested_range(*scratch, required_bytes)
|| transform.output_elements != *scratch_elements as u64
{
return Err("prepared Ada TF32 scratch identity changed before launch".into());
}
}
}
let mut live = ResolvedGemmLaunchSetBuilder::new(prepared.routes.len())?;
for route in &prepared.routes {
ctx.validate_resolved_gemm_route(route, "prepared f32 Triad")?;
live.push(route)?;
}
prepared
.resolved_launch_set
.ensure_current(live.finish()?, "prepared f32 Triad")
}
pub(in crate::mamba_ssm::gpu) fn validate_prepared_f32_triad_for_timing(
ctx: &GpuCtx,
prepared: &PreparedF32TriadLaunch,
) -> Result<(), String> {
validate_prepared_f32_triad(ctx, prepared)
}
fn validate_tf32_streamk_prepared_layout(
ctx: &GpuCtx,
prepared: &PreparedF32TriadLaunch,
route: Tf32PhysicalRoute,
plan: Tf32StreamKLaunchPlan,
) -> Result<(), String> {
let spec = tf32_kernel_spec(prepared.request.op, route)?;
let [resolved] = prepared.routes.as_ref() else {
return Err("prepared TF32 stream-K launch requires exactly one route".into());
};
if resolved.symbol != spec.symbol {
return Err("prepared TF32 stream-K route changed".into());
}
let expected = match route.exact_fma() {
Some(exact) => {
let plan = sm120_fma_launch_plan(prepared.request, exact)?;
Tf32StreamKLaunchPlan {
grid: plan.units,
partial_elements: plan.slab_elements,
flag_elements: plan.flag_elements,
}
}
None => {
tf32_streamk_launch_plan(prepared.request, spec, ctx.kernels.multiprocessor_count())?
}
};
if expected.grid != plan.grid
|| expected.partial_elements != plan.partial_elements
|| expected.flag_elements != plan.flag_elements
{
return Err("prepared TF32 stream-K launch plan changed".into());
}
if resolved.launch.grid_dim != (plan.grid, 1, 1)
|| resolved.launch.block_dim != (spec.threads, 1, 1)
|| resolved.launch.shared_mem_bytes != spec.dynamic_shared_bytes
{
return Err("prepared TF32 stream-K launch configuration changed".into());
}
Ok(())
}
fn validate_tf32_splitk_prepared_layout(
prepared: &PreparedF32TriadLaunch,
route: Tf32PhysicalRoute,
plan: Tf32SplitKLaunchPlan,
) -> Result<(), String> {
let spec = tf32_splitk_spec(prepared.request.op, route)?;
let [fused] = prepared.routes.as_ref() else {
return Err("prepared fused TF32 split-K launch requires exactly one route".into());
};
if fused.symbol != spec.symbol {
return Err("prepared fused TF32 split-K route changed".into());
}
let expected = tf32_splitk_launch_plan(prepared.request, spec)?;
if expected.scratch_elements != plan.scratch_elements
|| expected.counter_elements != plan.counter_elements
{
return Err("prepared fused TF32 split-K workspace extent changed".into());
}
if expected.fused.grid_dim != plan.fused.grid_dim
|| expected.fused.block_dim != plan.fused.block_dim
|| expected.fused.shared_mem_bytes != plan.fused.shared_mem_bytes
|| fused.launch.grid_dim != plan.fused.grid_dim
|| fused.launch.block_dim != plan.fused.block_dim
|| fused.launch.shared_mem_bytes != plan.fused.shared_mem_bytes
{
return Err("prepared fused TF32 split-K launch configuration changed".into());
}
Ok(())
}
unsafe fn enqueue_scalar_zero_f32<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
prepared: &PreparedF32TriadLaunch,
params: &SgbZeroReductionParams,
config: cudarc::driver::LaunchConfig,
observer: &mut O,
observation: Option<PhysicalLaunchObservation>,
) -> Result<(), String> {
let kernels = ctx.kernels.triad_kernels();
let function = match prepared.request.op {
ResolvedGemmOp::Nn => &kernels.gemm_bi_nn_zero_reduction,
ResolvedGemmOp::Tn => &kernels.gemm_bi_tn_zero_reduction,
ResolvedGemmOp::Nt => &kernels.gemm_bi_nt_zero_reduction,
};
let output = prepared.operands.output;
let null_input = 0_u64;
let bias = prepared.operands.bias.unwrap_or(0);
let mut builder = ctx.stream.launch_builder(function);
builder.arg(&output);
builder.arg(&null_input);
builder.arg(&null_input);
builder.arg(&bias);
builder.arg(params);
unsafe { enqueue_with_physical_observation(observer, &mut builder, config, observation) }
.map_err(|error| error.with_driver_context(format_args!("{}", prepared.routes[0].symbol)))
}
unsafe fn enqueue_tf32_f32<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
prepared: &PreparedF32TriadLaunch,
launch: Tf32RawLaunch<'_>,
observer: &mut O,
) -> Result<(), String> {
let function = ctx
.kernels
.triad_kernels()
.tf32_function(prepared.routes[0].symbol)
.ok_or_else(|| {
format!(
"qualified TF32 symbol {} is unavailable",
prepared.routes[0].symbol
)
})?;
unsafe { enqueue_tf32_raw(&ctx.stream, function, launch, observer) }
}
fn sm89_tf32_tn_source_bytes(request: F32TriadRequest) -> Result<u64, String> {
u64::try_from(request.shape.m)
.ok()
.and_then(|rows| rows.checked_mul(request.shape.lda as u64))
.and_then(|elements| elements.checked_mul(4))
.ok_or_else(|| "Ada TF32 transpose input span overflows u64".to_string())
}
unsafe fn enqueue_sm89_tf32_tn_pre_rna<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
prepared: &PreparedF32TriadLaunch,
observer: &mut O,
logical_dtype: Option<PolicyDtype>,
) -> Result<(), String> {
let PreparedF32Kind::Tf32TnPreRna {
route,
transpose_params,
gemm_params,
transpose_config,
gemm_config,
scratch,
scratch_elements,
transform,
} = &prepared.kind
else {
return Err("prepared launch is not an Ada TF32 pre-RNA pipeline".into());
};
let [resolved] = prepared.routes.as_ref() else {
return Err("Ada TF32 pre-RNA pipeline requires exactly one GEMM route".into());
};
let transpose = ctx
.kernels
.triad_sm89_tf32_joint_function(TN_PRE_RNA_TRANSPOSE_SYMBOL)
.ok_or_else(|| "qualified Ada TF32 transpose symbol is unavailable".to_string())?;
let gemm = ctx
.kernels
.triad_sm89_tf32_joint_function(resolved.symbol)
.ok_or_else(|| {
format!(
"qualified Ada TF32 symbol {} is unavailable",
resolved.symbol
)
})?;
ctx.record_resolved_gemm_route(*resolved)?;
let scratch_bytes = u64::try_from(*scratch_elements)
.ok()
.and_then(|elements| elements.checked_mul(4))
.ok_or_else(|| "Ada TF32 scratch span overflows u64".to_string())?;
let source_bytes = sm89_tf32_tn_source_bytes(prepared.request)?;
let transform_observation = logical_dtype.map(|dtype| {
PhysicalLaunchObservation::input_transform(
TN_PRE_RNA_TRANSPOSE_SYMBOL,
prepared.request.op,
dtype,
(
prepared.request.shape.m,
prepared.request.shape.k,
prepared.request.shape.n,
),
(
prepared.request.shape.lda,
prepared.request.shape.ldb,
prepared.request.shape.ldc,
),
**transform,
PhysicalConversionArguments::new(
prepared.operands.a,
source_bytes,
*scratch,
scratch_bytes,
),
)
});
let mut transpose_builder = ctx.stream.launch_builder(transpose);
transpose_builder.arg(&prepared.operands.a);
transpose_builder.arg(scratch);
transpose_builder.arg(transpose_params);
unsafe {
enqueue_with_physical_observation(
observer,
&mut transpose_builder,
*transpose_config,
transform_observation,
)
}
.map_err(|error| error.with_driver_context(format_args!("{TN_PRE_RNA_TRANSPOSE_SYMBOL}")))?;
let physical = physical_prepared_f32_route(prepared, *resolved);
let gemm_observation =
logical_dtype.map(|dtype| PhysicalLaunchObservation::gemm(dtype, None, physical));
let output = prepared.operands.output;
let b = prepared.operands.b;
let bias = 0_u64;
let mut gemm_builder = ctx.stream.launch_builder(gemm);
gemm_builder.arg(&output);
gemm_builder.arg(scratch);
gemm_builder.arg(&b);
gemm_builder.arg(&bias);
gemm_builder.arg(gemm_params);
unsafe {
enqueue_with_physical_observation(
observer,
&mut gemm_builder,
*gemm_config,
gemm_observation,
)
}
.map_err(|error| error.with_driver_context(format_args!("{:?}", route)))
}
unsafe fn enqueue_tf32_splitk_f32<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
prepared: &PreparedF32TriadLaunch,
params: &Sm80Tf32KernelParams,
plan: Tf32SplitKLaunchPlan,
workspace: Tf32SplitKWorkspace,
observer: &mut O,
logical_dtype: Option<PolicyDtype>,
) -> Result<(), String> {
let [resolved] = prepared.routes.as_ref() else {
return Err("prepared fused TF32 split-K launch requires exactly one route".into());
};
let function = ctx
.kernels
.triad_kernels()
.tf32_splitk_function(resolved.symbol)
.ok_or_else(|| {
format!(
"qualified fused TF32 split-K symbol {} is unavailable",
resolved.symbol
)
})?;
ctx.record_resolved_gemm_route(*resolved)?;
let physical = physical_prepared_f32_route(prepared, *resolved);
let observation =
logical_dtype.map(|dtype| PhysicalLaunchObservation::gemm(dtype, None, physical));
let output = prepared.operands.output;
let bias = prepared.operands.bias.unwrap_or(0);
let mut builder = ctx.stream.launch_builder(function);
builder.arg(&output);
builder.arg(&workspace.partial);
builder.arg(&workspace.counters);
builder.arg(&prepared.operands.a);
builder.arg(&prepared.operands.b);
builder.arg(&bias);
builder.arg(params);
unsafe { enqueue_with_physical_observation(observer, &mut builder, plan.fused, observation) }
.map_err(|error| error.with_driver_context(format_args!("{}", resolved.symbol)))
}
unsafe fn enqueue_tf32_raw<O: PhysicalLaunchObserver>(
stream: &Arc<CudaStream>,
function: &CudaFunction,
launch: Tf32RawLaunch<'_>,
observer: &mut O,
) -> Result<(), String> {
let output = launch.operands.output;
let bias = launch.operands.bias.unwrap_or(0);
match (launch.route, launch.params) {
(
Tf32PhysicalRoute::MmaTf32Rna(_)
| Tf32PhysicalRoute::Sm89MmaTf32Compact8
| Tf32PhysicalRoute::Sm89NnDirectN96
| Tf32PhysicalRoute::Sm89NnN96
| Tf32PhysicalRoute::Sm89NtALdmatrixN96
| Tf32PhysicalRoute::Sm89NtRnaM144N96S2
| Tf32PhysicalRoute::Sm89NtRowstageM128N192S2
| Tf32PhysicalRoute::Sm89TnDirectM192N192S2,
PreparedTf32Params::Sm80(params),
) => {
let a = if launch.zero_reduction {
0
} else {
launch.operands.a
};
let b = if launch.zero_reduction {
0
} else {
launch.operands.b
};
let mut builder = stream.launch_builder(function);
builder.arg(&output);
builder.arg(&a);
builder.arg(&b);
builder.arg(&bias);
builder.arg(¶ms);
unsafe {
enqueue_with_physical_observation(
observer,
&mut builder,
launch.config,
launch.observation,
)
}
.map_err(|error| error.with_driver_context(format_args!("{}", launch.symbol)))
}
(Tf32PhysicalRoute::Sm90aWgmmaTf32Tma(_), PreparedTf32Params::Sm90a(params)) => {
let maps = launch
.maps
.ok_or_else(|| "SM90a TF32 launch has no tensor maps".to_string())?
.maps();
let mut builder = stream.launch_builder(function);
builder.arg(&output);
builder.arg(&maps[0]);
builder.arg(&maps[1]);
builder.arg(&bias);
builder.arg(¶ms);
unsafe {
enqueue_with_physical_observation(
observer,
&mut builder,
launch.config,
launch.observation,
)
}
.map_err(|error| error.with_driver_context(format_args!("{}", launch.symbol)))
}
(Tf32PhysicalRoute::Sm100Tcgen05Tf32Tma(_), PreparedTf32Params::Sm100(params)) => {
let maps = launch
.maps
.ok_or_else(|| "SM100 TF32 launch has no tensor maps".to_string())?
.maps();
let mut builder = stream.launch_builder(function);
builder.arg(&output);
builder.arg(&maps[0]);
builder.arg(&maps[1]);
builder.arg(&bias);
builder.arg(¶ms);
unsafe {
enqueue_with_physical_observation(
observer,
&mut builder,
launch.config,
launch.observation,
)
}
.map_err(|error| error.with_driver_context(format_args!("{}", launch.symbol)))
}
(
Tf32PhysicalRoute::Sm120TmaMmaTf32Rna(_)
| Tf32PhysicalRoute::Sm120TmaMmaTf32RnaStreamKV1(_),
PreparedTf32Params::Sm120(_),
)
| (Tf32PhysicalRoute::Sm120TmaFmaExact(_), PreparedTf32Params::Sm120Fma(_)) => {
let maps = launch
.maps
.ok_or_else(|| "SM120 TF32 launch has no tensor maps".to_string())?
.maps();
let workspace_first = match launch.route {
Tf32PhysicalRoute::Sm120TmaMmaTf32RnaStreamKV1(_) => {
if launch.streamk.is_none() && launch.config.grid_dim != (1, 1, 1) {
return Err(
"SM120 TF32 stream-K launch without a workspace must be one CTA".into(),
);
}
true
}
Tf32PhysicalRoute::Sm120TmaFmaExact(exact) => {
if launch.streamk.is_none() && exact.splits != 1 {
return Err("exact-F32 SM120 split launch requires a workspace".into());
}
true
}
_ => false,
};
let workspace = launch.streamk.unwrap_or(Tf32StreamKWorkspace {
partial: 0,
flags: 0,
});
let params = launch.params;
let mut builder = stream.launch_builder(function);
builder.arg(&output);
if workspace_first {
builder.arg(&workspace.partial);
builder.arg(&workspace.flags);
}
builder.arg(&maps[0]);
builder.arg(&maps[1]);
builder.arg(&bias);
match ¶ms {
PreparedTf32Params::Sm120(params) => {
builder.arg(params);
}
PreparedTf32Params::Sm120Fma(params) => {
builder.arg(params);
}
_ => return Err("prepared TF32 route and parameter ABI disagree".into()),
}
unsafe {
enqueue_with_physical_observation(
observer,
&mut builder,
launch.config,
launch.observation,
)
}
.map_err(|error| error.with_driver_context(format_args!("{}", launch.symbol)))
}
_ => Err("prepared TF32 route and parameter ABI disagree".into()),
}
}
pub(super) unsafe fn enqueue_tf32_qualification_probe(
stream: &Arc<CudaStream>,
function: &CudaFunction,
request: F32TriadRequest,
operands: F32TriadOperands,
route: Tf32PhysicalRoute,
maps: Option<&F32PreparedTensorMaps>,
config: LaunchConfig,
) -> Result<(), String> {
let spec = tf32_kernel_spec(request.op, route)?;
let origins = maps.map(F32PreparedTensorMaps::origins).unwrap_or_default();
let params = tf32_params(request, operands, origins, route)?;
let mut observer = NoPhysicalObserver;
unsafe {
enqueue_tf32_raw(
stream,
function,
Tf32RawLaunch {
operands,
route,
maps,
params,
config,
zero_reduction: request.shape.reduction(request.op) == 0,
symbol: spec.symbol,
observation: None,
streamk: None,
},
&mut observer,
)
}
}
pub(in crate::mamba_ssm::gpu) unsafe fn launch_prepared_f32_triad(
ctx: &GpuCtx,
prepared: &PreparedF32TriadLaunch,
scalar: impl FnOnce(&mut ScalarLaunchControl<'_>) -> Result<(), String>,
) -> Result<(), String> {
validate_prepared_f32_triad(ctx, prepared)?;
unsafe { enqueue_validated_prepared_f32_triad(ctx, prepared, scalar) }
}
pub(in crate::mamba_ssm::gpu) unsafe fn enqueue_validated_prepared_f32_triad(
ctx: &GpuCtx,
prepared: &PreparedF32TriadLaunch,
scalar: impl FnOnce(&mut ScalarLaunchControl<'_>) -> Result<(), String>,
) -> Result<(), String> {
match &prepared.kind {
PreparedF32Kind::Scalar(plan) => {
let mut control = ScalarLaunchControl {
ctx,
routes: &prepared.routes,
plan: *plan,
operands: prepared.operands,
next: 0,
};
scalar(&mut control)?;
control.finish()
}
PreparedF32Kind::ScalarZero { params, config, .. } => {
ctx.record_resolved_gemm_route(prepared.routes[0])?;
let mut observer = NoPhysicalObserver;
unsafe { enqueue_scalar_zero_f32(ctx, prepared, params, *config, &mut observer, None) }
}
PreparedF32Kind::Tf32 { .. } | PreparedF32Kind::Tf32StreamK { .. } => {
let (route, maps, params, config, streamk) = tf32_prepared_kernel_parts(prepared)?;
ctx.record_resolved_gemm_route(prepared.routes[0])?;
let mut observer = NoPhysicalObserver;
unsafe {
enqueue_tf32_f32(
ctx,
prepared,
Tf32RawLaunch {
operands: prepared.operands,
route,
maps,
params,
config: *config,
zero_reduction: streamk.is_none()
&& prepared.request.shape.reduction(prepared.request.op) == 0,
symbol: prepared.routes[0].symbol,
observation: None,
streamk,
},
&mut observer,
)
}
}
PreparedF32Kind::Tf32SplitK {
params,
plan,
workspace,
..
} => {
let mut observer = NoPhysicalObserver;
unsafe {
enqueue_tf32_splitk_f32(
ctx,
prepared,
params,
*plan,
*workspace,
&mut observer,
None,
)
}
}
PreparedF32Kind::Tf32TnPreRna { .. } => {
let mut observer = NoPhysicalObserver;
unsafe { enqueue_sm89_tf32_tn_pre_rna(ctx, prepared, &mut observer, None) }
}
}
}
type Tf32PreparedKernelParts<'a> = (
Tf32PhysicalRoute,
Option<&'a F32PreparedTensorMaps>,
PreparedTf32Params,
&'a cudarc::driver::LaunchConfig,
Option<Tf32StreamKWorkspace>,
);
fn tf32_prepared_kernel_parts(
prepared: &PreparedF32TriadLaunch,
) -> Result<Tf32PreparedKernelParts<'_>, String> {
match &prepared.kind {
PreparedF32Kind::Tf32 {
route,
maps,
params,
config,
} => Ok((*route, maps.as_ref(), *params, config, None)),
PreparedF32Kind::Tf32StreamK {
route,
maps,
params,
config,
workspace,
..
} => Ok((*route, Some(maps), *params, config, Some(*workspace))),
_ => Err("prepared launch is not a TF32 kernel launch".into()),
}
}
unsafe fn enqueue_validated_prepared_f32_triad_observed<O, Scalar>(
ctx: &GpuCtx,
prepared: &PreparedF32TriadLaunch,
observer: &mut O,
logical_dtype: PolicyDtype,
scalar: Scalar,
) -> Result<(), String>
where
O: PhysicalLaunchObserver,
Scalar: FnOnce(&mut PhysicalScalarLaunchControl<'_, O>) -> Result<(), String>,
{
let physical_resources_digest = prepared.resources.physical_digest();
match &prepared.kind {
PreparedF32Kind::Scalar(plan) => {
let base = ScalarLaunchControl {
ctx,
routes: &prepared.routes,
plan: *plan,
operands: prepared.operands,
next: 0,
};
let mut control = PhysicalScalarLaunchControl {
base,
logical_dtype,
physical_resources_digest,
observer,
};
scalar(&mut control)?;
control.finish()
}
PreparedF32Kind::ScalarZero { params, config, .. } => {
let route = prepared.routes[0];
ctx.record_resolved_gemm_route(route)?;
unsafe {
enqueue_scalar_zero_f32(
ctx,
prepared,
params,
*config,
observer,
Some(PhysicalLaunchObservation::gemm(
logical_dtype,
Some(physical_resources_digest),
route,
)),
)
}
}
PreparedF32Kind::Tf32 { .. } | PreparedF32Kind::Tf32StreamK { .. } => {
let (route, maps, params, config, streamk) = tf32_prepared_kernel_parts(prepared)?;
let resolved = prepared.routes[0];
let physical = physical_prepared_f32_route(prepared, resolved);
ctx.record_resolved_gemm_route(resolved)?;
unsafe {
enqueue_tf32_f32(
ctx,
prepared,
Tf32RawLaunch {
operands: prepared.operands,
route,
maps,
params,
config: *config,
zero_reduction: streamk.is_none()
&& prepared.request.shape.reduction(prepared.request.op) == 0,
symbol: resolved.symbol,
observation: Some(PhysicalLaunchObservation::gemm(
logical_dtype,
None,
physical,
)),
streamk,
},
observer,
)
}
}
PreparedF32Kind::Tf32SplitK {
params,
plan,
workspace,
..
} => unsafe {
enqueue_tf32_splitk_f32(
ctx,
prepared,
params,
*plan,
*workspace,
observer,
Some(logical_dtype),
)
},
PreparedF32Kind::Tf32TnPreRna { .. } => unsafe {
enqueue_sm89_tf32_tn_pre_rna(ctx, prepared, observer, Some(logical_dtype))
},
}
}
fn require_sm90a_tensor_map_access(stream: &Arc<cudarc::driver::CudaStream>) -> Result<(), String> {
let supported = stream
.context()
.attribute(
cudarc::driver::sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_TENSOR_MAP_ACCESS_SUPPORTED,
)
.map_err(|error| format!("query CUDA tensor-map support: {error:?}"))?;
if supported == 0 {
return Err("SM90a WGMMA requires CUDA tensor-map access support".into());
}
Ok(())
}
fn specialized_device_identity(
compute_capability: (u32, u32),
multiprocessor_count: u32,
target: crate::mamba_ssm::gpu::kernel_identity::CudaTarget,
driver: crate::mamba_ssm::gpu::kernel_identity::DriverIdentity,
) -> crate::mamba_ssm::gpu::kernel_identity::DeviceIdentity {
crate::mamba_ssm::gpu::kernel_identity::DeviceIdentity {
compute_capability,
multiprocessor_count,
target,
driver,
}
}
fn sm90a_map_binding(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
) -> Result<Sm90aMapBinding, String> {
require_sm90a_tensor_map_access(stream)?;
let allocation_domain = validated_allocation_domain(stream, kernels, "SM90a")?;
let compiler = kernels
.sm90a_compiler_identity()
.filter(|compiler| compiler.target.as_str() == "sm_90a")
.ok_or_else(|| "exact-sm_90a compiler identity is unavailable".to_string())?;
let artifact = kernels
.artifact_set_identity()
.specialized
.filter(|artifact| {
artifact.module_kind == crate::mamba_ssm::gpu::kernel_identity::ModuleKind::TriadSm90a
})
.ok_or_else(|| "exact-sm_90a artifact identity is unavailable".to_string())?;
Ok(Sm90aMapBinding {
allocation_domain,
artifact,
compiler,
device: sm90a_device_identity(stream, kernels)?,
})
}
fn sm90a_device_identity(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
) -> Result<crate::mamba_ssm::gpu::kernel_identity::DeviceIdentity, String> {
let (major, minor) = stream
.context()
.compute_capability()
.map_err(|error| format!("query CUDA compute capability: {error:?}"))?;
Ok(specialized_device_identity(
(
u32::try_from(major).map_err(|_| format!("negative CUDA CC major {major}"))?,
u32::try_from(minor).map_err(|_| format!("negative CUDA CC minor {minor}"))?,
),
kernels.multiprocessor_count(),
crate::mamba_ssm::gpu::kernel_identity::CudaTarget::new("sm_90a")?,
crate::mamba_ssm::gpu::kernel_identity::query_driver_identity()?,
))
}
pub fn prepare_sm90a_tensor_maps(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
request: Sm90aMapRequest,
) -> Result<Sm90aPreparedTensorMaps, String> {
stream
.context()
.bind_to_thread()
.map_err(|error| format!("bind CUDA context for tensor-map encoding: {error:?}"))?;
let binding = sm90a_map_binding(stream, kernels)?;
if stream.context().compute_capability().ok() != Some((9, 0)) || !kernels.has_sm90a_wgmma() {
return Err("SM90a tensor maps require a loaded exact-sm_90a module".into());
}
let keys = sm90a_tensor_map_keys(request)?;
let allocations = sm90a_allocation_identities(keys, binding.allocation_domain)?;
let capturing = stream
.capture_status()
.map_err(|error| format!("query CUDA capture status: {error:?}"))?
!= cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE;
kernels.prepare_sm90a_tensor_maps(request, keys, allocations, capturing, binding)
}
fn sm90a_forced_identity(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
route: Sm90aForcedRoute,
maps: &Sm90aPreparedTensorMaps,
operands: Sm90aLaunchOperands,
) -> Result<Sm90aRouteIdentity, String> {
stream
.context()
.bind_to_thread()
.map_err(|error| format!("bind CUDA context for SM90a launch: {error:?}"))?;
route.shape.validate(route.op)?;
let binding = sm90a_map_binding(stream, kernels)?;
if !maps.matches_binding(binding) {
return Err(
"SM90a prepared tensor maps belong to a different CUDA context or module".into(),
);
}
maps.validate_live_allocations()?;
if maps.request.op != route.op
|| maps.request.dtype != route.dtype
|| maps.request.shape != route.shape
{
return Err("SM90a prepared tensor maps do not match the forced route".into());
}
let resolved = resolve_sm90a_forced(
(
i32::try_from(binding.device.compute_capability.0)
.map_err(|_| "CUDA CC major exceeds i32::MAX".to_string())?,
i32::try_from(binding.device.compute_capability.1)
.map_err(|_| "CUDA CC minor exceeds i32::MAX".to_string())?,
),
kernels.has_sm90a_wgmma(),
route.op,
route.dtype,
route.schedule,
route.shape,
)?;
if resolved != Some(route) {
return Err("SM90a forced route is unavailable; use the resolved baseline".into());
}
if operands.output_ptr == 0 {
return Err("SM90a output pointer must be non-null".into());
}
let output_alignment = if route.op == Sm90aOp::Tn { 4 } else { 2 };
if !operands.output_ptr.is_multiple_of(output_alignment) {
return Err(format!(
"SM90a output pointer must be {output_alignment}-byte aligned"
));
}
if operands.bias_ptr != 0 && !operands.bias_ptr.is_multiple_of(4) {
return Err("SM90a bias pointer must be 4-byte aligned".into());
}
match route.op {
Sm90aOp::Nn => {}
Sm90aOp::Tn if operands.bias_ptr != 0 || operands.beta != 1.0 => {
return Err("SM90a TN requires no bias and beta == 1.0".into());
}
Sm90aOp::Nt if operands.bias_ptr != 0 || operands.beta != 0.0 => {
return Err("SM90a NT requires no bias and beta == 0.0".into());
}
_ => {}
}
let tensor_maps_digest = maps.identity_digest();
Ok(Sm90aRouteIdentity {
numeric_contract: Sm90aNumericContract::Wgmma,
op: route.op,
dtype: route.dtype,
schedule: route.schedule,
shape: route.shape,
tile: SM90A_TILE,
stages: SM90A_STAGES,
cluster: (1, 1, 1),
symbol: route.symbol(),
module_kind: crate::mamba_ssm::gpu::kernel_identity::ModuleKind::TriadSm90a,
exact_target: "sm_90a",
artifact: binding.artifact,
compiler: binding.compiler,
device: binding.device,
tensor_maps_digest,
resources_digest: sm90a_resources_digest(
route,
operands,
binding.allocation_domain,
tensor_maps_digest,
)?,
tuning_revision: 0,
})
}
pub fn validate_sm90a_graph_replay(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
route: Sm90aForcedRoute,
maps: &Sm90aPreparedTensorMaps,
operands: Sm90aLaunchOperands,
captured: Sm90aRouteIdentity,
) -> Result<(), String> {
let live = sm90a_forced_identity(stream, kernels, route, maps, operands)?;
captured.ensure_current(live, "SM90a graph replay")
}
pub fn launch_sm90a_wgmma_forced(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
route: Sm90aForcedRoute,
maps: &Sm90aPreparedTensorMaps,
operands: Sm90aLaunchOperands,
) -> Result<Sm90aRouteIdentity, String> {
let identity = sm90a_forced_identity(stream, kernels, route, maps, operands)?;
let function = kernels
.sm90a_function(route.symbol())
.ok_or_else(|| format!("SM90a kernel {} is unavailable", route.symbol()))?;
let config = sm90a_launch_config(route)?;
let m = checked_i32(route.shape.m, "M")?;
let k = checked_i32(route.shape.k, "K")?;
let n = checked_i32(route.shape.n, "N")?;
let ldc = checked_i32(route.shape.ldc, "ldc")?;
let mut builder = stream.launch_builder(function);
builder.arg(&operands.output_ptr);
builder.arg(&maps.a);
builder.arg(&maps.b);
builder.arg(&operands.bias_ptr);
builder.arg(&operands.alpha);
builder.arg(&operands.beta);
builder.arg(&m);
builder.arg(&k);
builder.arg(&n);
builder.arg(&ldc);
unsafe { builder.launch(config) }
.map(|_| identity)
.map_err(|error| format!("launch {}: {error:?}", route.symbol()))
}
fn require_sm100_tensor_map_access(stream: &Arc<cudarc::driver::CudaStream>) -> Result<(), String> {
let supported = stream
.context()
.attribute(
cudarc::driver::sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_TENSOR_MAP_ACCESS_SUPPORTED,
)
.map_err(|error| format!("query CUDA tensor-map support: {error:?}"))?;
if supported == 0 {
return Err("SM100 TCGEN requires CUDA tensor-map access support".into());
}
Ok(())
}
fn sm100_map_binding(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
) -> Result<Sm100MapBinding, String> {
require_sm100_tensor_map_access(stream)?;
let allocation_domain = validated_allocation_domain(stream, kernels, "SM100")?;
let target = kernels
.sm100_target_candidate()
.ok_or_else(|| "SM100 target identity is unavailable".to_string())?;
let compiler = kernels
.sm100_compiler_identity()
.filter(|compiler| compiler.target.as_str() == target.nvrtc_arch)
.ok_or_else(|| "SM100 compiler identity is unavailable".to_string())?;
let artifact = kernels
.artifact_set_identity()
.specialized
.filter(|artifact| {
artifact.module_kind == crate::mamba_ssm::gpu::kernel_identity::ModuleKind::TriadSm100
})
.ok_or_else(|| "SM100 artifact identity is unavailable".to_string())?;
let (major, minor) = stream
.context()
.compute_capability()
.map_err(|error| format!("query CUDA compute capability: {error:?}"))?;
if (major, minor) != target.device_cc {
return Err("SM100 module target does not match the CUDA device minor".into());
}
let compute_capability = (
u32::try_from(major).map_err(|_| format!("negative CUDA CC major {major}"))?,
u32::try_from(minor).map_err(|_| format!("negative CUDA CC minor {minor}"))?,
);
let device = specialized_device_identity(
compute_capability,
kernels.multiprocessor_count(),
crate::mamba_ssm::gpu::kernel_identity::CudaTarget::new(
crate::mamba_ssm::gpu::device::GpuDevice::resolve_nvrtc_target(compute_capability)?,
)?,
crate::mamba_ssm::gpu::kernel_identity::query_driver_identity()?,
);
Ok(Sm100MapBinding {
allocation_domain,
artifact,
compiler,
device,
target,
})
}
pub fn prepare_sm100_tensor_maps(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
request: Sm100MapRequest,
) -> Result<Sm100PreparedTensorMaps, String> {
stream
.context()
.bind_to_thread()
.map_err(|error| format!("bind CUDA context for SM100 tensor maps: {error:?}"))?;
if stream
.capture_status()
.map_err(|error| format!("query CUDA capture status: {error:?}"))?
!= cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE
{
return Err("SM100 tensor maps must be prepared before graph capture".into());
}
let binding = sm100_map_binding(stream, kernels)?;
if !kernels.has_sm100_tcgen() {
return Err("SM100 tensor maps require a complete specialized module".into());
}
let plan = sm100_tensor_map_plan(request, binding.allocation_domain)?;
kernels.prepare_sm100_tensor_maps(
request,
plan.keys,
plan.allocations,
plan.origins,
false,
binding,
)
}
pub fn prepare_sm100_tcgen_forced(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
route: Sm100ForcedRoute,
maps: &Sm100PreparedTensorMaps,
operands: Sm100LaunchOperands,
) -> Result<Sm100PreparedLaunch, String> {
stream
.context()
.bind_to_thread()
.map_err(|error| format!("bind CUDA context for SM100 preparation: {error:?}"))?;
if stream
.capture_status()
.map_err(|error| format!("query CUDA capture status: {error:?}"))?
!= cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE
{
return Err("SM100 launch must be prepared before graph capture".into());
}
route.shape.validate(route.op)?;
let binding = sm100_map_binding(stream, kernels)?;
if !maps.matches_binding(binding) {
return Err(
"SM100 prepared tensor maps belong to a different CUDA context or module".into(),
);
}
maps.validate_live_allocations()?;
if maps.request.op != route.op
|| maps.request.dtype != route.dtype
|| maps.request.tile != route.physical.tile
|| maps.request.shape != route.shape
{
return Err("SM100 tensor maps do not match the forced physical route".into());
}
if resolve_sm100_forced(binding.target.device_cc, Some(binding.target), route)? != Some(route) {
return Err("SM100 forced route is unavailable; use the resolved baseline".into());
}
validate_sm100_operands(route, operands)?;
let spec = route.kernel_spec()?;
let tensor_maps_digest = maps.identity_digest();
let resources = Sm100LaunchResourceSnapshot::query(route, operands, binding.allocation_domain)?;
let params = Sm100KernelParams {
a_x: maps.origins.a_x,
a_y: maps.origins.a_y,
b_x: maps.origins.b_x,
b_y: maps.origins.b_y,
alpha: operands.alpha,
beta: operands.beta,
m: checked_i32(route.shape.m, "M")?,
k: checked_i32(route.shape.k, "K")?,
n: checked_i32(route.shape.n, "N")?,
ldc: checked_i32(route.shape.ldc, "ldc")?,
};
let identity = Sm100RouteIdentity {
numeric_contract: Sm100NumericContract::Tcgen05F32,
op: route.op,
dtype: route.dtype,
physical: route.physical,
shape: route.shape,
symbol: spec.symbol,
module_kind: crate::mamba_ssm::gpu::kernel_identity::ModuleKind::TriadSm100,
target: binding.target,
artifact: binding.artifact,
compiler: binding.compiler,
device: binding.device,
tensor_map_revision: SM100_TENSOR_MAP_REVISION,
tensor_maps_digest,
resources_digest: resources.digest(route, operands, tensor_maps_digest),
tuning_revision: SM100_TUNING_REVISION,
};
Ok(Sm100PreparedLaunch {
route,
maps: *maps,
operands,
params: params.into_words(),
identity,
resources,
})
}
fn validate_sm100_operands(
route: Sm100ForcedRoute,
operands: Sm100LaunchOperands,
) -> Result<(), String> {
if operands.output_ptr == 0 {
return Err("SM100 output pointer must be non-null".into());
}
let output_alignment = if route.op == Sm100Op::Tn { 4 } else { 2 };
if !operands.output_ptr.is_multiple_of(output_alignment) {
return Err(format!(
"SM100 output pointer must be {output_alignment}-byte aligned"
));
}
if operands.bias_ptr != 0 && !operands.bias_ptr.is_multiple_of(4) {
return Err("SM100 bias pointer must be 4-byte aligned".into());
}
match route.op {
Sm100Op::Nn if operands.bias_ptr != 0 && operands.alpha != 1.0 => {
Err("SM100 NN bias seeding requires alpha == 1.0".into())
}
Sm100Op::Tn if operands.bias_ptr != 0 || operands.beta != 1.0 => {
Err("SM100 TN requires no bias and beta == 1.0".into())
}
Sm100Op::Nt if operands.bias_ptr != 0 || operands.beta != 0.0 => {
Err("SM100 NT requires no bias and beta == 0.0".into())
}
_ => Ok(()),
}
}
fn validate_sm100_prepared_binding(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
prepared: &Sm100PreparedLaunch,
) -> Result<(), String> {
if AllocationDomain::from_context(stream.context())? != prepared.maps.binding.allocation_domain
|| kernels.allocation_domain() != prepared.maps.binding.allocation_domain
|| kernels.sm100_compiler_identity() != Some(prepared.identity.compiler)
|| kernels.artifact_set_identity().specialized != Some(prepared.identity.artifact)
|| kernels.sm100_target_candidate() != Some(prepared.identity.target)
|| prepared.maps.binding.artifact != prepared.identity.artifact
|| prepared.maps.binding.compiler != prepared.identity.compiler
|| prepared.maps.binding.device != prepared.identity.device
|| prepared.maps.binding.target != prepared.identity.target
{
return Err("SM100 prepared launch no longer matches its module context".into());
}
Ok(())
}
pub fn validate_sm100_graph_replay(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
prepared: &Sm100PreparedLaunch,
) -> Result<(), String> {
stream
.context()
.bind_to_thread()
.map_err(|error| format!("bind CUDA context for SM100 replay guard: {error:?}"))?;
validate_sm100_prepared_binding(stream, kernels, prepared)?;
prepared.maps.validate_live_allocations()?;
let live = Sm100LaunchResourceSnapshot::query(
prepared.route,
prepared.operands,
prepared.maps.binding.allocation_domain,
)?;
if live != prepared.resources {
return Err("SM100 graph replay allocation identity changed since capture".into());
}
let digest = live.digest(
prepared.route,
prepared.operands,
prepared.maps.identity_digest(),
);
if digest != prepared.identity.resources_digest {
return Err("SM100 graph replay resource identity changed since capture".into());
}
Ok(())
}
pub fn launch_sm100_tcgen_prepared(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
prepared: &Sm100PreparedLaunch,
) -> Result<Sm100RouteIdentity, String> {
validate_sm100_prepared_binding(stream, kernels, prepared)?;
let spec = prepared.route.kernel_spec()?;
if spec.symbol != prepared.identity.symbol {
return Err("SM100 prepared symbol no longer matches its physical route".into());
}
let function = kernels
.sm100_function(spec.symbol)
.ok_or_else(|| format!("SM100 kernel {} is unavailable", spec.symbol))?;
let (rows, columns) = match prepared.route.op {
Sm100Op::Nn => (prepared.route.shape.m, prepared.route.shape.n),
Sm100Op::Tn => (prepared.route.shape.k, prepared.route.shape.n),
Sm100Op::Nt => (prepared.route.shape.m, prepared.route.shape.k),
};
let rows = checked_u32(rows, "SM100 output rows")?;
let columns = checked_u32(columns, "SM100 output columns")?;
let grid = checked_grid_product(
rows.div_ceil(prepared.route.physical.tile.output_rows()),
columns.div_ceil(prepared.route.physical.tile.output_columns()),
1,
)?;
let config = cudarc::driver::LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (spec.threads, 1, 1),
shared_mem_bytes: spec.dynamic_shared_bytes,
};
let mut builder = stream.launch_builder(function);
builder.arg(&prepared.operands.output_ptr);
builder.arg(&prepared.maps.a);
builder.arg(&prepared.maps.b);
builder.arg(&prepared.operands.bias_ptr);
let params = Sm100KernelParams::from_words(prepared.params);
builder.arg(¶ms);
unsafe { builder.launch(config) }
.map(|_| prepared.identity)
.map_err(|error| format!("launch {}: {error:?}", spec.symbol))
}
pub fn prepare_sm120_tensor_maps(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
request: Sm120MapRequest,
) -> Result<Sm120PreparedTensorMaps, String> {
if stream
.capture_status()
.map_err(|error| format!("query CUDA capture status: {error:?}"))?
!= cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE
{
return Err("SM120 tensor maps must be prepared before graph capture".into());
}
stream
.context()
.bind_to_thread()
.map_err(|error| format!("bind CUDA context for SM120 tensor maps: {error:?}"))?;
let binding = sm120_map_binding(stream, kernels)?;
if !kernels.has_sm120_tma_mma16() {
return Err("SM120 tensor maps require a complete specialized module".into());
}
let plan = sm120_tensor_map_plan(request, binding.allocation_domain)?;
kernels.prepare_sm120_tensor_maps(
request,
plan.keys,
plan.allocations,
plan.origins,
false,
binding,
)
}
fn sm120_map_binding(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
) -> Result<Sm120MapBinding, String> {
let allocation_domain = validated_allocation_domain(stream, kernels, "SM120")?;
let tensor_map_access = stream
.context()
.attribute(
cudarc::driver::sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_TENSOR_MAP_ACCESS_SUPPORTED,
)
.map_err(|error| format!("query SM120 tensor-map support: {error:?}"))?
!= 0;
let optin_shared = stream
.context()
.attribute(
cudarc::driver::sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_MAX_SHARED_MEMORY_PER_BLOCK_OPTIN,
)
.map_err(|error| format!("query SM120 opt-in shared memory: {error:?}"))?;
let (major, minor) = stream
.context()
.compute_capability()
.map_err(|error| format!("query SM120 compute capability: {error:?}"))?;
let device_cc = (major, minor);
let compiler = kernels
.sm120_compiler_identity()
.ok_or_else(|| "SM120 compiler identity is unavailable".to_string())?;
let target = kernels
.sm120_target_candidate()
.ok_or_else(|| "SM120 generic target identity is unavailable".to_string())?;
if target.device_cc != device_cc {
return Err("SM120 module target does not match the CUDA device minor".into());
}
if compiler.target.as_str() != target.nvrtc_arch {
return Err("SM120 compiler target does not match the accepted transaction".into());
}
let artifact = kernels
.artifact_set_identity()
.specialized
.filter(|artifact| {
artifact.module_kind == crate::mamba_ssm::gpu::kernel_identity::ModuleKind::TriadSm120
})
.ok_or_else(|| "SM120 artifact identity is unavailable".to_string())?;
let compute_capability = (
u32::try_from(major).map_err(|_| format!("negative CUDA CC major {major}"))?,
u32::try_from(minor).map_err(|_| format!("negative CUDA CC minor {minor}"))?,
);
let accepted_target =
crate::mamba_ssm::gpu::kernel_identity::CudaTarget::new(target.nvrtc_arch)?;
let device_caps = crate::mamba_ssm::gpu::kernel_identity::DeviceCaps {
compute_capability,
nvrtc_version: compiler.nvrtc_version,
accepted_target: Some(accepted_target),
optin_shared_bytes: u32::try_from(optin_shared)
.map_err(|_| format!("negative SM120 opt-in shared memory {optin_shared}"))?,
tensor_map_access,
};
if kernels.sm120_device_caps() != Some(device_caps) {
return Err("SM120 device capabilities changed since module qualification".into());
}
let device = specialized_device_identity(
compute_capability,
kernels.multiprocessor_count(),
crate::mamba_ssm::gpu::kernel_identity::CudaTarget::new(target.ptx_target)?,
crate::mamba_ssm::gpu::kernel_identity::query_driver_identity()?,
);
Ok(Sm120MapBinding {
allocation_domain,
artifact,
compiler,
device,
device_caps,
target,
})
}
pub fn prepare_sm120_tma_forced(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
route: Sm120ForcedRoute,
maps: &Sm120PreparedTensorMaps,
operands: Sm120LaunchOperands,
) -> Result<Sm120PreparedLaunch, String> {
if stream
.capture_status()
.map_err(|error| format!("query CUDA capture status: {error:?}"))?
!= cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE
{
return Err("SM120 launch must be prepared before graph capture".into());
}
stream
.context()
.bind_to_thread()
.map_err(|error| format!("bind CUDA context for SM120 preparation: {error:?}"))?;
route.shape.validate(route.op)?;
let binding = sm120_map_binding(stream, kernels)?;
if !maps.matches_binding(binding) {
return Err(
"SM120 prepared tensor maps belong to a different CUDA context or module".into(),
);
}
maps.validate_live_allocations()?;
if maps.request.op != route.op
|| maps.request.dtype != route.dtype
|| maps.request.tile != route.physical.tile
|| maps.request.bk != route.physical.bk
|| maps.request.shape != route.shape
{
return Err("SM120 tensor maps do not match the forced physical route".into());
}
if resolve_sm120_forced(binding.device_caps, Some(binding.target), route)? != Some(route) {
return Err("SM120 forced route is unavailable; use the resolved baseline".into());
}
validate_sm120_operands(route, operands)?;
let spec = route.kernel_spec()?;
let kernel_resources = kernels
.sm120_kernel_resources(spec.symbol)
.ok_or_else(|| format!("SM120 resource census for {} is unavailable", spec.symbol))?;
let tensor_maps_digest = maps.identity_digest();
let resources = Sm120LaunchResourceSnapshot::query(route, operands, binding.allocation_domain)?;
let params = Sm120KernelParams {
a_x: maps.origins.a_x,
a_y: maps.origins.a_y,
b_x: maps.origins.b_x,
b_y: maps.origins.b_y,
alpha: operands.alpha,
beta: operands.beta,
m: checked_i32(route.shape.m, "M")?,
k: checked_i32(route.shape.k, "K")?,
n: checked_i32(route.shape.n, "N")?,
ldc: checked_i32(route.shape.ldc, "ldc")?,
};
let identity = Sm120RouteIdentity {
numeric_contract: Sm120NumericContract::for_schedule(route.physical.schedule),
op: route.op,
dtype: route.dtype,
physical: route.physical,
shape: route.shape,
symbol: spec.symbol,
module_kind: crate::mamba_ssm::gpu::kernel_identity::ModuleKind::TriadSm120,
target: binding.target,
artifact: binding.artifact,
compiler: binding.compiler,
device: binding.device,
device_caps: binding.device_caps,
tensor_map_revision: SM120_TENSOR_MAP_REVISION,
tensor_maps_digest,
resources_digest: resources.digest(route, operands, tensor_maps_digest),
tuning_revision: SM120_TUNING_REVISION,
schedule_revision: SM120_SCHEDULE_REVISION,
};
let resolved_launch_set = build_resolved_gemm_launch_set(&[identity.resolved_route()?])?;
Ok(Sm120PreparedLaunch {
stream_handle: stream.cu_stream() as usize,
route,
maps: *maps,
operands,
params: params.into_words(),
identity,
resolved_launch_set,
resources,
kernel_resources,
})
}
fn validate_sm120_operands(
route: Sm120ForcedRoute,
operands: Sm120LaunchOperands,
) -> Result<(), String> {
if operands.output_ptr == 0 {
return Err("SM120 output pointer must be non-null".into());
}
let output_alignment = if route.op == Sm120Op::Tn { 4 } else { 2 };
if !operands.output_ptr.is_multiple_of(output_alignment) {
return Err(format!(
"SM120 output pointer must be {output_alignment}-byte aligned"
));
}
if operands.bias_ptr != 0 && !operands.bias_ptr.is_multiple_of(4) {
return Err("SM120 bias pointer must be 4-byte aligned".into());
}
match route.op {
Sm120Op::Nn if operands.bias_ptr != 0 && operands.alpha != 1.0 => {
Err("SM120 NN bias seeding requires alpha == 1.0".into())
}
Sm120Op::Tn if operands.bias_ptr != 0 || operands.beta != 1.0 => {
Err("SM120 TN requires no bias and beta == 1.0".into())
}
Sm120Op::Nt if operands.bias_ptr != 0 || operands.beta != 0.0 => {
Err("SM120 NT requires no bias and beta == 0.0".into())
}
_ => Ok(()),
}
}
fn validate_sm120_prepared_binding(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
prepared: &Sm120PreparedLaunch,
) -> Result<(), String> {
if stream.cu_stream() as usize != prepared.stream_handle {
return Err("SM120 prepared launch belongs to a different CUDA stream".into());
}
if AllocationDomain::from_context(stream.context())? != prepared.maps.binding.allocation_domain
|| kernels.allocation_domain() != prepared.maps.binding.allocation_domain
|| kernels.sm120_compiler_identity() != Some(prepared.identity.compiler)
|| kernels.artifact_set_identity().specialized != Some(prepared.identity.artifact)
|| prepared.maps.binding.artifact != prepared.identity.artifact
|| prepared.maps.binding.compiler != prepared.identity.compiler
|| prepared.maps.binding.device != prepared.identity.device
|| prepared.maps.binding.device_caps != prepared.identity.device_caps
|| prepared.maps.binding.target != prepared.identity.target
{
return Err("SM120 prepared launch no longer matches its module context".into());
}
Ok(())
}
pub fn validate_sm120_graph_replay(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
prepared: &Sm120PreparedLaunch,
) -> Result<(), String> {
stream
.context()
.bind_to_thread()
.map_err(|error| format!("bind CUDA context for SM120 replay guard: {error:?}"))?;
validate_sm120_prepared_binding(stream, kernels, prepared)?;
let binding = sm120_map_binding(stream, kernels)?;
if binding != prepared.maps.binding {
return Err("SM120 graph replay device or module identity changed since capture".into());
}
prepared.maps.validate_live_allocations()?;
let live = Sm120LaunchResourceSnapshot::query(
prepared.route,
prepared.operands,
prepared.maps.binding.allocation_domain,
)?;
if live != prepared.resources {
return Err("SM120 graph replay allocation identity changed since capture".into());
}
let digest = live.digest(
prepared.route,
prepared.operands,
prepared.maps.identity_digest(),
);
if digest != prepared.identity.resources_digest {
return Err("SM120 graph replay resource identity changed since capture".into());
}
let live_identity = Sm120RouteIdentity {
artifact: binding.artifact,
compiler: binding.compiler,
device: binding.device,
device_caps: binding.device_caps,
tensor_maps_digest: prepared.maps.identity_digest(),
resources_digest: digest,
..prepared.identity
};
let live_launch_set = build_resolved_gemm_launch_set(&[live_identity.resolved_route()?])?;
prepared
.resolved_launch_set()
.ensure_current(live_launch_set, "SM120 graph replay")?;
Ok(())
}
pub fn launch_sm120_tma_prepared(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
prepared: &Sm120PreparedLaunch,
) -> Result<Sm120RouteIdentity, String> {
let mut observer = NoPhysicalObserver;
unsafe { enqueue_sm120_tma_prepared_observed(stream, kernels, prepared, &mut observer, None) }
.map(|()| prepared.identity)
}
unsafe fn enqueue_sm120_tma_prepared_observed<O: PhysicalLaunchObserver>(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
prepared: &Sm120PreparedLaunch,
observer: &mut O,
observation: Option<PhysicalLaunchObservation>,
) -> Result<(), String> {
validate_sm120_prepared_binding(stream, kernels, prepared)?;
let spec = prepared.route.kernel_spec()?;
if spec.symbol != prepared.identity.symbol {
return Err("SM120 prepared symbol no longer matches its physical route".into());
}
let function = kernels
.sm120_function(spec.symbol)
.ok_or_else(|| format!("SM120 kernel {} is unavailable", spec.symbol))?;
let (rows, columns) = match prepared.route.op {
Sm120Op::Nn => (prepared.route.shape.m, prepared.route.shape.n),
Sm120Op::Tn => (prepared.route.shape.k, prepared.route.shape.n),
Sm120Op::Nt => (prepared.route.shape.m, prepared.route.shape.k),
};
let rows = checked_u32(rows, "SM120 output rows")?;
let columns = checked_u32(columns, "SM120 output columns")?;
let grid = sm120_launch_grid(kernels, prepared.route, spec, rows, columns)?;
let config = cudarc::driver::LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (spec.threads, 1, 1),
shared_mem_bytes: spec.dynamic_shared_bytes,
};
let workspace = sm120_streamk_workspace(stream, kernels, prepared.route, spec, grid)?;
let mut builder = stream.launch_builder(function);
builder.arg(&prepared.operands.output_ptr);
if let Some(workspace) = &workspace {
builder.arg(&workspace.partial);
builder.arg(&workspace.flags);
}
builder.arg(&prepared.maps.a);
builder.arg(&prepared.maps.b);
builder.arg(&prepared.operands.bias_ptr);
let params = Sm120KernelParams::from_words(prepared.params);
builder.arg(¶ms);
unsafe { enqueue_with_physical_observation(observer, &mut builder, config, observation) }
.map_err(|error| error.with_driver_context(format_args!("launch {}", spec.symbol)))
}
fn sm120_launch_grid(
kernels: &GpuKernels,
route: Sm120ForcedRoute,
spec: &Sm120KernelSpec,
rows: u32,
columns: u32,
) -> Result<u32, String> {
match route.physical.schedule {
Sm120Schedule::Tiled => checked_grid_product(
rows.div_ceil(route.physical.tile.output_rows()),
columns.div_ceil(route.physical.tile.output_columns()),
1,
),
Sm120Schedule::StreamK => {
let grid = kernels.multiprocessor_count();
if grid == 0 {
return Err("SM120 stream-K launch requires at least one multiprocessor".into());
}
let reduction = match route.op {
Sm120Op::Nn => route.shape.k,
Sm120Op::Tn => route.shape.m,
Sm120Op::Nt => route.shape.n,
};
let k_tiles = reduction
.div_ceil(route.physical.bk.elements() as usize)
.max(1);
let tiles = (rows.div_ceil(route.physical.tile.output_rows()) as usize)
.checked_mul(columns.div_ceil(route.physical.tile.output_columns()) as usize)
.ok_or_else(|| "SM120 stream-K tile count overflows usize".to_string())?;
let units = tiles
.checked_mul(k_tiles)
.ok_or_else(|| "SM120 stream-K unit count overflows usize".to_string())?;
if units > i32::MAX as usize {
return Err("SM120 stream-K unit count exceeds the kernel's 32-bit range".into());
}
let _ = spec;
Ok(grid)
}
}
}
struct Sm120StreamKWorkspace {
partial: CUptr,
flags: CUptr,
}
fn sm120_streamk_workspace(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
route: Sm120ForcedRoute,
spec: &Sm120KernelSpec,
grid: u32,
) -> Result<Option<Sm120StreamKWorkspace>, String> {
use cudarc::driver::DevicePtr;
if route.physical.schedule != Sm120Schedule::StreamK {
return Ok(None);
}
let m_atoms = if route.physical.wide_m_warp() { 4 } else { 2 };
let slab_floats = (spec.threads as usize) * m_atoms * 16;
let partial_floats = (grid as usize)
.checked_mul(SM120_TF32_STREAMK_SLOTS_PER_CTA)
.and_then(|slots| slots.checked_mul(slab_floats))
.ok_or_else(|| "SM120 stream-K slab extent overflows usize".to_string())?;
if partial_floats > SPLITK_SCRATCH_CAP {
return Err("SM120 stream-K slabs exceed the fixed workspace".into());
}
if (grid as usize) * SM120_TF32_STREAMK_SLOTS_PER_CTA > TF32_SPLITK_COUNTER_CAP {
return Err("SM120 stream-K flags exceed the fixed counter workspace".into());
}
let (partial, _) = kernels.splitk_scratch_buf(stream)?.device_ptr(stream);
let (flags, _) = kernels
.triad_kernels()
.tf32_splitk_counter_buf(stream)?
.device_ptr(stream);
Ok(Some(Sm120StreamKWorkspace { partial, flags }))
}
fn sm120_policy_dtype(dtype: WeightDtype) -> Result<PolicyDtype, String> {
match dtype {
WeightDtype::Bf16 => Ok(PolicyDtype::Bf16),
WeightDtype::F16 => Ok(PolicyDtype::F16),
WeightDtype::F32 | WeightDtype::Tf32 => {
Err("SM120 automatic route requires BF16 or F16".into())
}
}
}
pub(in crate::mamba_ssm::gpu) fn launch_sm120_auto_observed<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
request: Sm120AutoRequest,
) -> Result<Option<Sm120AutoBranchSeal>, String> {
let family = ctx.kernels.triad_kernels().serves_sm120_family();
let Some(caps) = ctx.kernels.sm120_device_caps() else {
if family {
static NO_CAPS: std::sync::Once = std::sync::Once::new();
crate::mamba_ssm::gpu::diagnostics::warn_once(&NO_CAPS, || {
"this SM120 board reports no SM120 device capabilities; the portable \
tensor-core tiles serve every half GEMM"
.to_string()
});
}
return Ok(None);
};
let Some(target) = ctx.kernels.sm120_target_candidate() else {
if family {
static NO_TARGET: std::sync::Once = std::sync::Once::new();
crate::mamba_ssm::gpu::diagnostics::warn_once(&NO_TARGET, || {
"this SM120 board has no bound SM120 module target; the portable \
tensor-core tiles serve every half GEMM"
.to_string()
});
}
return Ok(None);
};
let Some(route) = resolve_sm120_auto(caps, Some(target), request) else {
if family {
static NO_CELL: std::sync::Once = std::sync::Once::new();
crate::mamba_ssm::gpu::diagnostics::warn_once(&NO_CELL, || {
format!(
"no measured SM120 half cell for {:?} {:?} {:?}; the portable tensor-core \
tiles serve it (reported once; later uncovered shapes are silent)",
request.op, request.dtype, request.shape
)
});
}
return Ok(None);
};
let key = Sm120PreparedKey::new(ctx.gemm_route(), route, request);
ctx.with_sm120_prepared_launches(|cache| {
let prepared = cache.ensure_sm120_prepared(ctx, key, route, request)?;
let identity = prepared.identity();
let resolved = identity.resolved_route()?;
unsafe {
enqueue_sm120_tma_prepared_observed(
&ctx.stream,
&ctx.kernels,
prepared,
observer,
Some(PhysicalLaunchObservation::gemm(
sm120_policy_dtype(route.dtype)?,
None,
resolved,
)),
)
}?;
ctx.record_resolved_gemm_route(resolved)?;
Ok(Some(Sm120AutoBranchSeal { route }))
})
}
pub(in crate::mamba_ssm::gpu) fn prepare_sm120_auto_graph_sequence<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &O,
request: Sm120AutoRequest,
) -> Result<PreparedTriadPhysicalGraphSequence, String> {
let caps = ctx
.kernels
.sm120_device_caps()
.ok_or_else(|| "prepared SM120 graph route has no device capabilities".to_string())?;
let target = ctx
.kernels
.sm120_target_candidate()
.ok_or_else(|| "prepared SM120 graph route has no module target".to_string())?;
let route = resolve_sm120_auto(caps, Some(target), request)
.ok_or_else(|| "prepared SM120 graph request is no longer qualified".to_string())?;
let key = Sm120PreparedKey::new(ctx.gemm_route(), route, request);
ctx.with_sm120_prepared_launches(|cache| {
let prepared = &cache
.entries
.get(&key)
.ok_or_else(|| {
"prepared SM120 graph cache entry is missing; run eager warmup again".to_string()
})?
.prepared;
validate_sm120_graph_replay(&ctx.stream, &ctx.kernels, prepared)?;
let resolved = prepared.identity().resolved_route()?;
let config = LaunchConfig {
grid_dim: resolved.launch.grid_dim,
block_dim: resolved.launch.block_dim,
shared_mem_bytes: resolved.launch.shared_mem_bytes,
};
let observation =
PhysicalLaunchObservation::gemm(sm120_policy_dtype(route.dtype)?, None, resolved);
let node = resolve_physical_launch_observation(observer, observation, config)?;
let function = ctx
.kernels
.sm120_function(resolved.symbol)
.ok_or_else(|| format!("qualified SM120 symbol {} is unavailable", resolved.symbol))?
.clone();
let mut arguments = PhysicalScalarKernelArguments::new();
arguments.push(prepared.operands.output_ptr)?;
let spec = prepared.route.kernel_spec()?;
if let Some(workspace) = sm120_streamk_workspace(
&ctx.stream,
&ctx.kernels,
prepared.route,
spec,
resolved.launch.grid_dim.0,
)? {
arguments.push(workspace.partial)?;
arguments.push(workspace.flags)?;
}
arguments.push(prepared.maps.a)?;
arguments.push(prepared.maps.b)?;
arguments.push(prepared.operands.bias_ptr)?;
arguments.push(Sm120KernelParams::from_words(prepared.params))?;
validate_sm120_graph_replay(&ctx.stream, &ctx.kernels, prepared)?;
Ok(PreparedTriadPhysicalGraphSequence {
launches: vec![PreparedTriadPhysicalGraphLaunch {
function,
config,
node,
arguments: Box::new(arguments),
}]
.into_boxed_slice(),
})
})
}
pub fn gemm_bi_forward(
ctx: &GpuCtx,
y: &mut GpuBuffer,
x: &GpuBuffer,
w_ptr: CUptr,
bias_ptr: CUptr, dims: (usize, usize, usize),
) -> Result<(), String> {
GemmDims::nn(dims, dims.1)?;
let x_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _r) = x.inner().device_ptr(&ctx.stream);
ptr
};
let operands = GemmBiFwdSubOperands {
x_ptr,
lda: dims.1,
w_ptr,
bias_ptr,
};
gemm_bi_forward_sub(ctx, y, &operands, dims)
}
pub fn gemm_bi_forward_sub(
ctx: &GpuCtx,
y: &mut GpuBuffer,
operands: &GemmBiFwdSubOperands,
dims: (usize, usize, usize),
) -> Result<(), String> {
let checked_dims = GemmDims::nn(dims, operands.lda)?;
let (batch, n_in, n_out) = checked_dims.tuple();
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape {
m: batch,
k: n_in,
n: n_out,
lda: operands.lda,
ldb: n_out,
ldc: n_out,
},
};
let triad = F32TriadOperands {
output: y.cached_ptr(),
a: operands.x_ptr,
b: operands.w_ptr,
bias: (operands.bias_ptr != 0).then_some(operands.bias_ptr),
alpha: 1.0,
beta: 0.0,
};
let plan = prove_scalar_plan(ctx, request, triad, |output, plan| {
launch_scalar_plan_raw(ctx, request, F32TriadOperands { output, ..triad }, plan)
})?;
let mut control = ProvenScalarLaunch {
plan,
operands: triad,
};
gemm_bi_forward_sub_with_control(
&ctx.stream,
&ctx.kernels,
y,
operands,
dims,
Some(&mut control),
)
}
fn enqueue_scalar_forward<C: ScalarLaunchController>(
control: &mut Option<&mut C>,
_request: F32TriadRequest,
_operands: F32TriadOperands,
symbol: &'static str,
config: cudarc::driver::LaunchConfig,
builder: &mut ScalarLaunchArgs<'_>,
) -> Result<(), PhysicalCudaLaunchError> {
if let Some(control) = control.as_deref_mut() {
return control.enqueue(symbol, config, builder);
}
let mut observer = NoPhysicalObserver;
unsafe { enqueue_with_physical_observation(&mut observer, builder.launch_args(), config, None) }
}
fn gemm_bi_forward_sub_with_control<C: ScalarLaunchController, Output: ScalarOutputArgument>(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
y: &mut Output,
operands: &GemmBiFwdSubOperands,
dims: (usize, usize, usize),
mut control: Option<&mut C>,
) -> Result<(), String> {
let GemmBiFwdSubOperands {
x_ptr,
lda,
w_ptr,
bias_ptr,
} = *operands;
let checked_dims = GemmDims::nn(dims, lda)?;
let (batch, n_in, n_out) = checked_dims.tuple();
let request = F32TriadRequest {
op: crate::mamba_ssm::gpu::kernel_identity::ResolvedGemmOp::Nn,
shape: F32TriadShape {
m: batch,
k: n_in,
n: n_out,
lda,
ldb: n_out,
ldc: n_out,
},
};
let lda_i = checked_dims.lda;
let alpha = control
.as_ref()
.map(|control| control.operands().alpha)
.unwrap_or(1.0);
let beta = control
.as_ref()
.map(|control| control.operands().beta)
.unwrap_or(0.0);
let actual_operands = F32TriadOperands {
output: y.scalar_ptr(),
a: x_ptr,
b: w_ptr,
bias: (bias_ptr != 0).then_some(bias_ptr),
alpha,
beta,
};
if let Some(control) = control.as_deref() {
control.validate_operands(actual_operands)?;
}
let scalar_plan = match control.as_ref() {
Some(prepared) => prepared.plan(),
None => scalar_launch_plan(scalar_launch_facts(kernels), request, actual_operands)?,
};
validate_bias_preseed(alpha, bias_ptr, "gemm_bi_forward_sub")?;
if matches!(scalar_plan, ScalarDispatchPlan::NnUltraThin) {
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_i = checked_dims.k_i32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_dims.n_u32.div_ceil(32), checked_dims.m_u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: checked_u32_product(
checked_dims.k_u32,
checked_u32(std::mem::size_of::<f32>(), "f32 byte width")?,
"ultra-thin shared memory",
)?,
};
let mut builder = scalar_launch_builder(stream, &kernels.gemm_bi_nn_ultra_thin, &control);
builder.arg_buffer_mut(y);
builder.arg(&x_ptr);
builder.arg(&w_ptr);
builder.arg(&bias_ptr);
builder.arg(&alpha);
builder.arg(&beta);
builder.arg(&m_i);
builder.arg(&n_i);
builder.arg(&k_i);
builder.arg(&lda_i); builder.arg(&n_i); builder.arg(&n_i); enqueue_scalar_forward(
&mut control,
request,
actual_operands,
"nn_ultra_thin",
cfg,
&mut builder,
)
.map_err(|error| error.with_driver_context(format_args!("nn_ultra_thin forward")))?;
return Ok(());
}
if matches!(scalar_plan, ScalarDispatchPlan::NnNarrowSmall) {
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_i = checked_dims.k_i32;
let post_op: i32 = 0;
let num_pid_m = checked_dims.m_u32.div_ceil(16);
let num_pid_n = checked_dims.n_u32.div_ceil(16);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_grid_product(num_pid_m, num_pid_n, 1)?, 1, 1),
block_dim: (64, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = scalar_launch_builder(stream, &kernels.gemm_bi_nn_narrow_small, &control);
builder.arg_buffer_mut(y);
builder.arg(&x_ptr);
builder.arg(&w_ptr);
builder.arg(&bias_ptr);
builder.arg(&alpha);
builder.arg(&beta);
builder.arg(&m_i);
builder.arg(&n_i);
builder.arg(&k_i);
builder.arg(&lda_i);
builder.arg(&n_i);
builder.arg(&n_i);
builder.arg(&post_op);
enqueue_scalar_forward(
&mut control,
request,
actual_operands,
"nn_narrow_small",
cfg,
&mut builder,
)
.map_err(|error| error.with_driver_context(format_args!("nn_narrow_small forward")))?;
return Ok(());
}
if matches!(scalar_plan, ScalarDispatchPlan::NnNarrow) {
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_i = checked_dims.k_i32;
let post_op: i32 = 0;
let num_pid_m = checked_dims.m_u32.div_ceil(64);
let num_pid_n = checked_dims.n_u32.div_ceil(32);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_grid_product(num_pid_m, num_pid_n, 1)?, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = scalar_launch_builder(stream, &kernels.gemm_bi_nn_narrow, &control);
builder.arg_buffer_mut(y);
builder.arg(&x_ptr);
builder.arg(&w_ptr);
builder.arg(&bias_ptr);
builder.arg(&alpha);
builder.arg(&beta);
builder.arg(&m_i);
builder.arg(&n_i);
builder.arg(&k_i);
builder.arg(&lda_i);
builder.arg(&n_i);
builder.arg(&n_i);
builder.arg(&post_op);
enqueue_scalar_forward(
&mut control,
request,
actual_operands,
"nn_narrow",
cfg,
&mut builder,
)
.map_err(|error| error.with_driver_context(format_args!("nn_narrow forward")))?;
return Ok(());
}
if matches!(scalar_plan, ScalarDispatchPlan::NnGemv) {
let m_i = checked_dims.m_i32;
let k_i = checked_dims.k_i32;
let ldy_i: i32 = 1;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_dims.m_u32.div_ceil(4), 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = scalar_launch_builder(stream, &kernels.gemm_bi_nn_gemv, &control);
builder.arg_buffer_mut(y);
builder.arg(&x_ptr);
builder.arg(&w_ptr);
builder.arg(&bias_ptr);
builder.arg(&alpha);
builder.arg(&beta);
builder.arg(&m_i);
builder.arg(&k_i);
builder.arg(&lda_i);
builder.arg(&ldy_i);
enqueue_scalar_forward(
&mut control,
request,
actual_operands,
"nn_gemv",
cfg,
&mut builder,
)
.map_err(|error| error.with_driver_context(format_args!("nn_gemv forward")))?;
return Ok(());
}
if let ScalarDispatchPlan::NnSplitKThinTail { k_main, k_tail } = scalar_plan {
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_chunks = checked_i32(k_main / 32, "NN K-tail chunks")?;
let num_pid_m = checked_dims.m_u32.div_ceil(32);
let num_pid_n = checked_dims.n_u32.div_ceil(64);
let partial_cfg = cudarc::driver::LaunchConfig {
grid_dim: (
checked_grid_product(
num_pid_m,
num_pid_n,
checked_u32(k_main / 32, "NN K-tail chunks")?,
)?,
1,
1,
),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let partial_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _r) = kernels.splitk_scratch_buf(stream)?.device_ptr(stream);
ptr
};
let mut pb = scalar_launch_builder(stream, &kernels.gemm_bi_nn_splitk32_partial, &control);
pb.arg(&partial_ptr);
pb.arg(&x_ptr);
pb.arg(&w_ptr);
pb.arg(&m_i);
pb.arg(&n_i);
pb.arg(&k_chunks);
pb.arg(&lda_i);
enqueue_scalar_forward(
&mut control,
request,
actual_operands,
"nn_splitk32_partial",
partial_cfg,
&mut pb,
)
.map_err(|error| {
error.with_driver_context(format_args!("nn_splitk32_partial (K-tail main)"))
})?;
let total = checked_dims.mn_u32;
let reduce_cfg = cudarc::driver::LaunchConfig {
grid_dim: (total.div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let zero_i32: i32 = 0;
let tail_cnt_i = checked_i32(k_tail, "NN K-tail count")?;
let x_tail_ptr = checked_ptr_add(
x_ptr,
checked_byte_offset(k_main, std::mem::size_of::<f32>(), "X tail")?,
"X tail",
)?;
let w_tail_ptr = checked_ptr_add(
w_ptr,
checked_byte_offset(
k_main.checked_mul(n_out).ok_or_else(|| {
invalid_gemm_dimensions("W tail element offset overflows usize")
})?,
std::mem::size_of::<f32>(),
"W tail",
)?,
"W tail",
)?;
let x_tail_stride_i = lda_i; let mut rb = scalar_launch_builder(stream, &kernels.gemm_bi_splitk_reduce, &control);
rb.arg_buffer_mut(y);
rb.arg(&partial_ptr);
rb.arg(&bias_ptr);
rb.arg(&x_tail_ptr);
rb.arg(&w_tail_ptr);
rb.arg(&alpha);
rb.arg(&m_i);
rb.arg(&n_i);
rb.arg(&k_chunks);
rb.arg(&x_tail_stride_i);
rb.arg(&zero_i32); rb.arg(&tail_cnt_i);
enqueue_scalar_forward(
&mut control,
request,
actual_operands,
"splitk_reduce",
reduce_cfg,
&mut rb,
)
.map_err(|error| error.with_driver_context(format_args!("splitk_reduce (K-tail)")))?;
return Ok(());
}
if matches!(
scalar_plan,
ScalarDispatchPlan::NnSplitKThin | ScalarDispatchPlan::NnM32N64SplitK32Qualified
) {
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_chunks = checked_i32(n_in / 32, "NN split-K chunks")?;
let num_pid_m = checked_dims.m_u32.div_ceil(32);
let num_pid_n = checked_dims.n_u32.div_ceil(64);
let partial_cfg = cudarc::driver::LaunchConfig {
grid_dim: (
checked_grid_product(
num_pid_m,
num_pid_n,
checked_u32(n_in / 32, "NN split-K chunks")?,
)?,
1,
1,
),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let partial_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _r) = kernels.splitk_scratch_buf(stream)?.device_ptr(stream);
ptr
};
let (partial_function, partial_symbol) =
if scalar_plan == ScalarDispatchPlan::NnM32N64SplitK32Qualified {
(
&kernels.gemm_bi_nn_splitk32_m32n64_exact,
"nn_splitk32_m32n64_exact",
)
} else {
(&kernels.gemm_bi_nn_splitk32_partial, "nn_splitk32_partial")
};
let mut pb = scalar_launch_builder(stream, partial_function, &control);
pb.arg(&partial_ptr);
pb.arg(&x_ptr);
pb.arg(&w_ptr);
pb.arg(&m_i);
pb.arg(&n_i);
pb.arg(&k_chunks);
pb.arg(&lda_i);
enqueue_scalar_forward(
&mut control,
request,
actual_operands,
partial_symbol,
partial_cfg,
&mut pb,
)
.map_err(|error| error.with_driver_context(format_args!("{partial_symbol}")))?;
let total = checked_dims.mn_u32;
let reduce_cfg = cudarc::driver::LaunchConfig {
grid_dim: (total.div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let null_tail: u64 = 0;
let zero_i32: i32 = 0;
let mut rb = scalar_launch_builder(stream, &kernels.gemm_bi_splitk_reduce, &control);
rb.arg_buffer_mut(y);
rb.arg(&partial_ptr);
rb.arg(&bias_ptr);
rb.arg(&null_tail); rb.arg(&null_tail); rb.arg(&alpha);
rb.arg(&m_i);
rb.arg(&n_i);
rb.arg(&k_chunks);
rb.arg(&zero_i32); rb.arg(&zero_i32); rb.arg(&zero_i32); enqueue_scalar_forward(
&mut control,
request,
actual_operands,
"splitk_reduce",
reduce_cfg,
&mut rb,
)
.map_err(|error| error.with_driver_context(format_args!("splitk_reduce")))?;
return Ok(());
}
const SPLITK_SLIM_K_CHUNK: u32 = 64; if let ScalarDispatchPlan::NnSplitKSlim { chunks: f_final } = scalar_plan {
let base_blocks = checked_tile_grid(checked_dims.m_u32, 128, checked_dims.n_u32, 64)?;
let k_chunk = SPLITK_SLIM_K_CHUNK;
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_i = checked_dims.k_i32;
let ldb_i = checked_dims.n_i32; let k_chunk_i = checked_i32(
checked_usize(k_chunk, "NN slim split-K chunk")?,
"NN slim split-K chunk",
)?;
let partial_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _r) = kernels.splitk_scratch_buf(stream)?.device_ptr(stream);
ptr
};
let partial_cfg = cudarc::driver::LaunchConfig {
grid_dim: (base_blocks, 1, f_final),
block_dim: (128, 1, 1), shared_mem_bytes: 0, };
let mut pb =
scalar_launch_builder(stream, &kernels.gemm_bi_nn_splitk_slim_partial, &control);
pb.arg(&partial_ptr);
pb.arg(&x_ptr);
pb.arg(&w_ptr);
pb.arg(&m_i);
pb.arg(&n_i);
pb.arg(&k_i);
pb.arg(&lda_i);
pb.arg(&ldb_i);
pb.arg(&k_chunk_i);
enqueue_scalar_forward(
&mut control,
request,
actual_operands,
"nn_splitk_slim_partial",
partial_cfg,
&mut pb,
)
.map_err(|error| error.with_driver_context(format_args!("nn_splitk_slim_partial")))?;
let total = checked_dims.mn_u32;
let reduce_cfg = cudarc::driver::LaunchConfig {
grid_dim: (total.div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let null_tail: u64 = 0;
let zero_i32_local: i32 = 0;
let f_i = checked_i32(
checked_usize(f_final, "NN slim split-K chunks")?,
"NN slim split-K chunks",
)?;
let mut rb = scalar_launch_builder(stream, &kernels.gemm_bi_splitk_reduce, &control);
rb.arg_buffer_mut(y);
rb.arg(&partial_ptr);
rb.arg(&bias_ptr);
rb.arg(&null_tail); rb.arg(&null_tail); rb.arg(&alpha);
rb.arg(&m_i);
rb.arg(&n_i);
rb.arg(&f_i); rb.arg(&zero_i32_local); rb.arg(&zero_i32_local); rb.arg(&zero_i32_local); enqueue_scalar_forward(
&mut control,
request,
actual_operands,
"splitk_reduce",
reduce_cfg,
&mut rb,
)
.map_err(|error| error.with_driver_context(format_args!("splitk_reduce (slim)")))?;
return Ok(());
}
if matches!(
scalar_plan,
ScalarDispatchPlan::NnSm89FixedCopyPlanQualified
) {
let params = SgbNnM64N64Params {
alpha,
beta,
m: checked_dims.m_i32,
n: checked_dims.n_i32,
k: checked_dims.k_i32,
lda: lda_i,
ldb: checked_dims.n_i32,
ldc: checked_dims.n_i32,
};
let total_tiles = checked_tile_grid(checked_dims.m_u32, 64, checked_dims.n_u32, 64)?;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (total_tiles, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let function = kernels
.fixed_sm89_f32_n64_copyplan
.as_ref()
.ok_or_else(|| "qualified Fixed CopyPlan kernel is unavailable".to_string())?;
let mut builder = scalar_launch_builder(stream, function, &control);
builder.arg_buffer_mut(y);
builder.arg(&x_ptr);
builder.arg(&w_ptr);
builder.arg(&bias_ptr);
builder.arg(¶ms);
enqueue_scalar_forward(
&mut control,
request,
actual_operands,
"nn_sm89_f32_n64_copyplan",
cfg,
&mut builder,
)
.map_err(|error| {
error.with_driver_context(format_args!("nn_sm89_f32_n64_copyplan forward"))
})?;
return Ok(());
}
if matches!(scalar_plan, ScalarDispatchPlan::NnM64N64Qualified) {
let params = SgbNnM64N64Params {
alpha,
beta,
m: checked_dims.m_i32,
n: checked_dims.n_i32,
k: checked_dims.k_i32,
lda: lda_i,
ldb: checked_dims.n_i32,
ldc: checked_dims.n_i32,
};
let total_tiles = checked_tile_grid(checked_dims.m_u32, 64, checked_dims.n_u32, 64)?;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (total_tiles, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: super::contract::SCALAR_NN_M64N64_DYNAMIC_SHARED_BYTES,
};
let mut builder =
scalar_launch_builder(stream, &kernels.gemm_bi_nn_m64n64_bk16_s2, &control);
builder.arg_buffer_mut(y);
builder.arg(&x_ptr);
builder.arg(&w_ptr);
builder.arg(&bias_ptr);
builder.arg(¶ms);
enqueue_scalar_forward(
&mut control,
request,
actual_operands,
"nn_m64n64_bk16_s2",
cfg,
&mut builder,
)
.map_err(|error| error.with_driver_context(format_args!("nn_m64n64_bk16_s2 forward")))?;
return Ok(());
}
if let ScalarDispatchPlan::NnFinal { slim } = scalar_plan {
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_i = checked_dims.k_i32;
let (func, bn) = if slim {
(&kernels.gemm_bi_nn_slim, 64)
} else {
(&kernels.gemm_bi_nn, 128)
};
let threads = if slim { 128u32 } else { 256u32 };
let smem_bytes: u32 = if slim { 0 } else { 34 * 1024 };
let total_tiles = checked_tile_grid(checked_dims.m_u32, 128, checked_dims.n_u32, bn)?;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (total_tiles, 1, 1),
block_dim: (threads, 1, 1),
shared_mem_bytes: smem_bytes,
};
let mut builder = scalar_launch_builder(stream, func, &control);
builder.arg_buffer_mut(y);
builder.arg(&x_ptr);
builder.arg(&w_ptr);
builder.arg(&bias_ptr);
builder.arg(&alpha);
builder.arg(&beta);
builder.arg(&m_i);
builder.arg(&n_i);
builder.arg(&k_i);
builder.arg(&lda_i); builder.arg(&n_i); builder.arg(&n_i); enqueue_scalar_forward(
&mut control,
request,
actual_operands,
if slim { "nn_slim" } else { "nn_big" },
cfg,
&mut builder,
)
.map_err(|error| {
error.with_driver_context(format_args!(
"nn_big{} forward",
if slim { "_slim" } else { "" }
))
})?;
return Ok(());
}
panic!(
"gpu_gemm_bi_forward: cuBLAS fallback hit (shape M={batch} K={n_in} N={n_out}). \
The zero-cuBLAS contract requires every shape to route through a custom \
kernel — add a dispatcher branch in this function for this shape."
);
}
#[cfg(test)]
fn scalar_backward_launch_plan(
op: ResolvedGemmOp,
dims: (usize, usize, usize),
multiprocessor_count: u32,
) -> Result<(GemmDims, F32TriadRequest, ScalarDispatchPlan), String> {
let (checked_dims, request) = scalar_backward_request(op, dims)?;
let plan = scalar_dispatch_plan(request, multiprocessor_count)?;
Ok((checked_dims, request, plan))
}
fn scalar_backward_request(
op: ResolvedGemmOp,
dims: (usize, usize, usize),
) -> Result<(GemmDims, F32TriadRequest), String> {
let checked_dims = match op {
ResolvedGemmOp::Tn => GemmDims::tn(dims)?,
ResolvedGemmOp::Nt => GemmDims::nt(dims)?,
ResolvedGemmOp::Nn => return Err("backward scalar launch plan requires TN or NT".into()),
};
let request = F32TriadRequest {
op,
shape: F32TriadShape::contiguous(op, dims),
};
Ok((checked_dims, request))
}
fn enqueue_scalar_backward<C: ScalarLaunchController>(
control: &mut Option<&mut C>,
_request: F32TriadRequest,
_operands: F32TriadOperands,
symbol: &'static str,
config: cudarc::driver::LaunchConfig,
builder: &mut ScalarLaunchArgs<'_>,
) -> Result<(), PhysicalCudaLaunchError> {
if let Some(control) = control.as_deref_mut() {
return control.enqueue(symbol, config, builder);
}
let mut observer = NoPhysicalObserver;
unsafe { enqueue_with_physical_observation(&mut observer, builder.launch_args(), config, None) }
}
pub fn gemm_bi_backward_dw(
ctx: &GpuCtx,
dw_ptr: CUptr, dy: &GpuBuffer,
x_saved: &GpuBuffer,
dims: (usize, usize, usize),
) -> Result<(), String> {
let (_, request) = scalar_backward_request(ResolvedGemmOp::Tn, dims)?;
let operands = F32TriadOperands {
output: dw_ptr,
a: x_saved.cached_ptr(),
b: dy.cached_ptr(),
bias: None,
alpha: 1.0,
beta: 1.0,
};
let plan = prove_scalar_plan(ctx, request, operands, |output, plan| {
launch_scalar_plan_raw(ctx, request, F32TriadOperands { output, ..operands }, plan)
})?;
let mut control = ProvenScalarLaunch { plan, operands };
gemm_bi_backward_dw_with_control(
&ctx.stream,
&ctx.kernels,
dw_ptr,
dy,
x_saved,
dims,
Some(&mut control),
)
}
fn gemm_bi_backward_dw_with_control<C: ScalarLaunchController, Input: ScalarInputArgument>(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
dw_ptr: CUptr,
dy: &Input,
x_saved: &Input,
dims: (usize, usize, usize),
mut control: Option<&mut C>,
) -> Result<(), String> {
let (checked_dims, request) = scalar_backward_request(ResolvedGemmOp::Tn, dims)?;
let (batch, n_in, n_out) = checked_dims.tuple();
let alpha = control
.as_ref()
.map(|control| control.operands().alpha)
.unwrap_or(1.0);
let operands = F32TriadOperands {
output: dw_ptr,
a: x_saved.scalar_ptr(),
b: dy.scalar_ptr(),
bias: None,
alpha,
beta: 1.0,
};
let scalar_plan = match control.as_ref() {
Some(prepared) => prepared.plan(),
None => scalar_launch_plan(scalar_launch_facts(kernels), request, operands)?,
};
if let Some(control) = control.as_ref() {
let prepared = control.operands();
if prepared.output != operands.output
|| prepared.a != operands.a
|| prepared.b != operands.b
|| prepared.bias != operands.bias
|| prepared.alpha.to_bits() != operands.alpha.to_bits()
|| prepared.beta.to_bits() != operands.beta.to_bits()
{
return Err("prepared TN operands differ from the physical launch arguments".into());
}
}
if matches!(scalar_plan, ScalarDispatchPlan::TnGemv) {
let m_i = checked_dims.m_i32;
let k_i = checked_dims.k_i32;
let lda_i = checked_dims.k_i32;
let ldy_i: i32 = 1;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_dims.k_u32.div_ceil(4), 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = scalar_launch_builder(stream, &kernels.gemm_bi_tn_gemv, &control);
builder.arg(&dw_ptr);
builder.arg_buffer(x_saved);
builder.arg_buffer(dy);
builder.arg(&alpha);
builder.arg(&m_i);
builder.arg(&k_i);
builder.arg(&lda_i);
builder.arg(&ldy_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"tn_gemv",
cfg,
&mut builder,
)
.map_err(|error| error.with_driver_context(format_args!("tn_gemv backward_dw")))?;
return Ok(());
}
if matches!(scalar_plan, ScalarDispatchPlan::TnNarrow) {
let m_i = checked_dims.m_i32;
let k_i = checked_dims.k_i32;
let n_i = checked_dims.n_i32;
let num_pid_m = checked_dims.k_u32.div_ceil(64);
let num_pid_n = checked_dims.n_u32.div_ceil(32);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_grid_product(num_pid_m, num_pid_n, 1)?, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = scalar_launch_builder(stream, &kernels.gemm_bi_tn_narrow, &control);
builder.arg(&dw_ptr);
builder.arg_buffer(x_saved);
builder.arg_buffer(dy);
builder.arg(&alpha);
builder.arg(&m_i);
builder.arg(&k_i);
builder.arg(&n_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"tn_narrow",
cfg,
&mut builder,
)
.map_err(|error| error.with_driver_context(format_args!("tn_narrow backward_dw")))?;
return Ok(());
}
if let ScalarDispatchPlan::TnNarrowSplitM {
m_chunk,
chunks: f_final,
} = scalar_plan
{
let grid_m = checked_dims.k_u32.div_ceil(64);
let grid_n = checked_dims.n_u32.div_ceil(32);
let m_i = checked_dims.m_i32;
let k_i = checked_dims.k_i32;
let n_i = checked_dims.n_i32;
let m_chunk_i = checked_i32(m_chunk, "TN narrow split-M chunk")?;
let f_i = checked_i32(f_final, "TN narrow split-M partitions")?;
let f_u32 = checked_u32(f_final, "TN narrow split-M partitions")?;
let partial_ptr = {
use cudarc::driver::DevicePtr;
kernels.splitk_scratch_buf(stream)?.device_ptr(stream).0
};
let partial_cfg = cudarc::driver::LaunchConfig {
grid_dim: (grid_m, grid_n, f_u32),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let symbol = scalar_tn_kernel_symbol(scalar_plan, operands);
let function = if symbol == "tn_narrow_splitm_partial_aligned" {
&kernels.gemm_bi_tn_narrow_splitm_partial_aligned
} else {
&kernels.gemm_bi_tn_narrow_splitm_partial
};
let mut partial_builder = scalar_launch_builder(stream, function, &control);
partial_builder.arg(&partial_ptr);
partial_builder.arg_buffer(x_saved);
partial_builder.arg_buffer(dy);
partial_builder.arg(&m_i);
partial_builder.arg(&k_i);
partial_builder.arg(&n_i);
partial_builder.arg(&m_chunk_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
symbol,
partial_cfg,
&mut partial_builder,
)
.map_err(|error| error.with_driver_context(format_args!("{symbol}")))?;
let reduce_cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_dims.kn_u32.div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let mut reducer = scalar_launch_builder(stream, &kernels.gemm_bi_splitm_reduce, &control);
reducer.arg(&dw_ptr);
reducer.arg(&partial_ptr);
reducer.arg(&alpha);
reducer.arg(&k_i);
reducer.arg(&n_i);
reducer.arg(&f_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"splitm_reduce",
reduce_cfg,
&mut reducer,
)
.map_err(|error| error.with_driver_context(format_args!("splitm_reduce")))?;
return Ok(());
}
if scalar_plan == ScalarDispatchPlan::TnM16N16SplitM16Qualified {
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (
checked_tile_grid(checked_dims.k_u32, 16, checked_dims.n_u32, 16)?,
1,
1,
),
block_dim: (super::contract::SCALAR_TN_M16N16_THREADS, 1, 1),
shared_mem_bytes: super::contract::SCALAR_TN_M16N16_DYNAMIC_SHARED_BYTES,
};
let mut builder = scalar_launch_builder(
stream,
&kernels.gemm_bi_tn_m16n16_bk16_s2_splitm16,
&control,
);
builder.arg(&dw_ptr);
builder.arg_buffer(x_saved);
builder.arg_buffer(dy);
builder.arg(&alpha);
builder.arg(&checked_dims.m_i32);
builder.arg(&checked_dims.k_i32);
builder.arg(&checked_dims.n_i32);
enqueue_scalar_backward(
&mut control,
request,
operands,
"tn_m16n16_bk16_s2_splitm16",
cfg,
&mut builder,
)
.map_err(|error| {
error.with_driver_context(format_args!("tn_m16n16_bk16_s2_splitm16 backward_dw"))
})?;
return Ok(());
}
if matches!(
scalar_plan,
ScalarDispatchPlan::TnD128InSm89DirectFoldQualified
| ScalarDispatchPlan::TnD128OutSm89DirectFoldQualified
) {
let symbol = match scalar_plan {
ScalarDispatchPlan::TnD128InSm89DirectFoldQualified => super::D128_IN_SYMBOL,
ScalarDispatchPlan::TnD128OutSm89DirectFoldQualified => super::D128_OUT_SYMBOL,
_ => unreachable!(),
};
let spec = super::sm89_exact_f32_d128_source::kernel_spec(symbol)
.ok_or_else(|| format!("missing exact-F32 d128 spec for {symbol}"))?;
let function = kernels
.triad_sm89_exact_f32_d128_function(symbol)
.ok_or_else(|| {
format!("qualified SM89 exact-F32 d128 kernel {symbol} is unavailable")
})?;
let mut builder = scalar_launch_builder(stream, function, &control);
builder.arg(&dw_ptr);
builder.arg_buffer(x_saved);
builder.arg_buffer(dy);
builder.arg(&alpha);
builder.arg(&checked_dims.m_i32);
builder.arg(&checked_dims.k_i32);
builder.arg(&checked_dims.n_i32);
enqueue_scalar_backward(
&mut control,
request,
operands,
symbol,
cudarc::driver::LaunchConfig {
grid_dim: spec.grid,
block_dim: spec.block,
shared_mem_bytes: spec.dynamic_shared_bytes,
},
&mut builder,
)
.map_err(|error| error.with_driver_context(format_args!("{symbol}")))?;
return Ok(());
}
if scalar_plan == ScalarDispatchPlan::TnD768InSm89DualChunkQualified {
let transposed_ptr = {
use cudarc::driver::DevicePtr;
kernels.transpose_scratch_buf(stream)?.device_ptr(stream).0
};
let transpose_cfg = cudarc::driver::LaunchConfig {
grid_dim: (24, 64, 1),
block_dim: (32, 16, 1),
shared_mem_bytes: 0,
};
let rows = checked_dims.m_i32;
let columns = checked_dims.k_i32;
let x_ptr = x_saved.scalar_ptr();
let mut transpose =
scalar_launch_builder(stream, &kernels.gemm_bi_transpose_f32_32x16_d768, &control);
transpose.arg(&transposed_ptr);
transpose.arg(&x_ptr);
transpose.arg(&rows);
transpose.arg(&columns);
enqueue_scalar_backward(
&mut control,
request,
operands,
"transpose_f32_32x16_d768",
transpose_cfg,
&mut transpose,
)
.map_err(|error| {
error.with_driver_context(format_args!("transpose X for SM89 exact-F32 TN"))
})?;
let function = kernels
.triad_sm89_exact_f32_function(super::D768_IN_FUSED_SYMBOL)
.ok_or_else(|| "qualified SM89 exact-F32 d768-in kernel is unavailable".to_string())?;
let params = super::Sm89ExactF32DualChunkParams {
alpha,
m: checked_dims.k_i32,
n: checked_dims.n_i32,
k0: 1_024,
k1: 1_024,
lda: checked_dims.m_i32,
ldb: checked_dims.n_i32,
ldc: checked_dims.n_i32,
};
let dy_ptr = dy.scalar_ptr();
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (576, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let mut fused = scalar_launch_builder(stream, function, &control);
fused.arg(&dw_ptr);
fused.arg(&transposed_ptr);
fused.arg(&dy_ptr);
fused.arg(¶ms);
enqueue_scalar_backward(
&mut control,
request,
operands,
super::D768_IN_FUSED_SYMBOL,
cfg,
&mut fused,
)
.map_err(|error| {
error.with_driver_context(format_args!("{}", super::D768_IN_FUSED_SYMBOL))
})?;
return Ok(());
}
if matches!(
scalar_plan,
ScalarDispatchPlan::TnD768OutSm89DirectBk16Qualified
| ScalarDispatchPlan::TnPrismSm89DirectBk16Qualified
) {
let (symbol, raw_grid, chunks, m_chunk) = match scalar_plan {
ScalarDispatchPlan::TnD768OutSm89DirectBk16Qualified => {
(super::D768_OUT_RAW_SYMBOL, 288, 4, 512)
}
ScalarDispatchPlan::TnPrismSm89DirectBk16Qualified => {
(super::PRISM_RAW_SYMBOL, 186, 6, 784)
}
_ => unreachable!(),
};
let partial_ptr = {
use cudarc::driver::DevicePtr;
kernels.splitk_scratch_buf(stream)?.device_ptr(stream).0
};
let function = kernels
.triad_sm89_exact_f32_function(symbol)
.ok_or_else(|| format!("qualified SM89 exact-F32 kernel {symbol} is unavailable"))?;
let m_chunk_i = checked_i32(m_chunk, "SM89 exact-F32 TN chunk")?;
let chunks_i = checked_i32(chunks, "SM89 exact-F32 TN partitions")?;
let raw_cfg = cudarc::driver::LaunchConfig {
grid_dim: (raw_grid, 1, chunks as u32),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let mut raw = scalar_launch_builder(stream, function, &control);
raw.arg(&partial_ptr);
raw.arg_buffer(x_saved);
raw.arg_buffer(dy);
raw.arg(&checked_dims.m_i32);
raw.arg(&checked_dims.k_i32);
raw.arg(&checked_dims.n_i32);
raw.arg(&m_chunk_i);
enqueue_scalar_backward(&mut control, request, operands, symbol, raw_cfg, &mut raw)
.map_err(|error| error.with_driver_context(format_args!("{symbol}")))?;
let reduce_cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_dims.kn_u32.div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let mut reducer = scalar_launch_builder(stream, &kernels.gemm_bi_splitm_reduce, &control);
reducer.arg(&dw_ptr);
reducer.arg(&partial_ptr);
reducer.arg(&alpha);
reducer.arg(&checked_dims.k_i32);
reducer.arg(&checked_dims.n_i32);
reducer.arg(&chunks_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"splitm_reduce",
reduce_cfg,
&mut reducer,
)
.map_err(|error| error.with_driver_context(format_args!("splitm_reduce")))?;
return Ok(());
}
if let ScalarDispatchPlan::TnSplitM {
m_chunk,
chunks: f_final,
} = scalar_plan
{
let base_blocks = checked_tile_grid(checked_dims.k_u32, 128, checked_dims.n_u32, 128)?;
let m_i = checked_dims.m_i32;
let k_i = checked_dims.k_i32;
let n_i = checked_dims.n_i32;
let m_chunk_i = checked_i32(m_chunk, "TN split-M chunk")?;
let f_i = checked_i32(f_final, "TN split-M partitions")?;
let f_final_u32 = checked_u32(f_final, "TN split-M partitions")?;
let partial_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _r) = kernels.splitk_scratch_buf(stream)?.device_ptr(stream);
ptr
};
let partial_cfg = cudarc::driver::LaunchConfig {
grid_dim: (base_blocks, 1, f_final_u32),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let partial_symbol = scalar_tn_kernel_symbol(scalar_plan, operands);
let partial_function = if partial_symbol == "tn_splitm_partial_aligned" {
&kernels.gemm_bi_tn_splitm_partial_aligned
} else {
&kernels.gemm_bi_tn_splitm_partial
};
let mut pb = scalar_launch_builder(stream, partial_function, &control);
pb.arg(&partial_ptr);
pb.arg_buffer(x_saved);
pb.arg_buffer(dy);
pb.arg(&m_i);
pb.arg(&k_i);
pb.arg(&n_i);
pb.arg(&m_chunk_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
partial_symbol,
partial_cfg,
&mut pb,
)
.map_err(|error| error.with_driver_context(format_args!("{partial_symbol}")))?;
let total = checked_dims.kn_u32;
let reduce_cfg = cudarc::driver::LaunchConfig {
grid_dim: (total.div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let mut rb = scalar_launch_builder(stream, &kernels.gemm_bi_splitm_reduce, &control);
rb.arg(&dw_ptr);
rb.arg(&partial_ptr);
rb.arg(&alpha);
rb.arg(&k_i);
rb.arg(&n_i);
rb.arg(&f_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"splitm_reduce",
reduce_cfg,
&mut rb,
)
.map_err(|error| error.with_driver_context(format_args!("splitm_reduce")))?;
return Ok(());
}
if let ScalarDispatchPlan::TnFinal { slim } = scalar_plan {
let m_i = checked_dims.m_i32;
let k_i = checked_dims.k_i32;
let n_i = checked_dims.n_i32;
let symbol = scalar_tn_kernel_symbol(scalar_plan, operands);
let (func, bn) = match symbol {
"tn_slim" => (&kernels.gemm_bi_tn_slim, 64),
"tn_aligned" => (&kernels.gemm_bi_tn_aligned, 128),
"tn_big" => (&kernels.gemm_bi_tn, 128),
_ => unreachable!("unexpected Big TN symbol"),
};
let threads = if slim { 128u32 } else { 256u32 };
let smem_bytes: u32 = if slim { 0 } else { 34 * 1024 };
let total_tiles = checked_tile_grid(checked_dims.k_u32, 128, checked_dims.n_u32, bn)?;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (total_tiles, 1, 1),
block_dim: (threads, 1, 1),
shared_mem_bytes: smem_bytes,
};
let mut builder = scalar_launch_builder(stream, func, &control);
builder.arg(&dw_ptr);
builder.arg_buffer(x_saved);
builder.arg_buffer(dy);
builder.arg(&alpha);
builder.arg(&m_i);
builder.arg(&k_i);
builder.arg(&n_i);
enqueue_scalar_backward(&mut control, request, operands, symbol, cfg, &mut builder)
.map_err(|error| {
error.with_driver_context(format_args!(
"tn_big{} backward_dw",
if slim { "_slim" } else { "" }
))
})?;
return Ok(());
}
panic!(
"gpu_gemm_bi_backward_dw: cuBLAS fallback hit (shape M={batch} K={n_in} N={n_out}). \
The zero-cuBLAS contract requires every shape to route through a custom \
kernel — add a dispatcher branch in this function for this shape."
);
}
pub fn gemm_bi_backward_dx(
ctx: &GpuCtx,
dx: &mut GpuBuffer,
dy: &GpuBuffer,
w_ptr: CUptr,
dims: (usize, usize, usize),
) -> Result<(), String> {
let (_, request) = scalar_backward_request(ResolvedGemmOp::Nt, dims)?;
let operands = F32TriadOperands {
output: dx.cached_ptr(),
a: dy.cached_ptr(),
b: w_ptr,
bias: None,
alpha: 1.0,
beta: 0.0,
};
let plan = prove_scalar_plan(ctx, request, operands, |output, plan| {
launch_scalar_plan_raw(ctx, request, F32TriadOperands { output, ..operands }, plan)
})?;
let mut control = ProvenScalarLaunch { plan, operands };
gemm_bi_backward_dx_with_control(
&ctx.stream,
&ctx.kernels,
dx,
dy,
w_ptr,
dims,
Some(&mut control),
)
}
fn gemm_bi_backward_dx_with_control<
C: ScalarLaunchController,
Output: ScalarOutputArgument,
Input: ScalarInputArgument,
>(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
dx: &mut Output,
dy: &Input,
w_ptr: CUptr,
dims: (usize, usize, usize),
mut control: Option<&mut C>,
) -> Result<(), String> {
let (checked_dims, request) = scalar_backward_request(ResolvedGemmOp::Nt, dims)?;
let (batch, n_in, n_out) = checked_dims.tuple();
let alpha = control
.as_ref()
.map(|control| control.operands().alpha)
.unwrap_or(1.0);
let operands = F32TriadOperands {
output: dx.scalar_ptr(),
a: dy.scalar_ptr(),
b: w_ptr,
bias: None,
alpha,
beta: 0.0,
};
let scalar_plan = match control.as_ref() {
Some(prepared) => prepared.plan(),
None => scalar_launch_plan(scalar_launch_facts(kernels), request, operands)?,
};
if let Some(control) = control.as_ref() {
let prepared = control.operands();
if prepared.output != operands.output
|| prepared.a != operands.a
|| prepared.b != operands.b
|| prepared.bias != operands.bias
|| prepared.alpha.to_bits() != operands.alpha.to_bits()
|| prepared.beta.to_bits() != operands.beta.to_bits()
{
return Err("prepared NT operands differ from the physical launch arguments".into());
}
}
if scalar_plan == ScalarDispatchPlan::NtM2N16SplitK32Qualified {
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (256, 1, 1),
block_dim: (super::contract::SCALAR_NT_M2N16_THREADS, 1, 1),
shared_mem_bytes: super::contract::SCALAR_NT_M2N16_DYNAMIC_SHARED_BYTES,
};
let mut builder =
scalar_launch_builder(stream, &kernels.gemm_bi_nt_m2n16_bk64_splitk32, &control);
builder.arg_buffer_mut(dx);
builder.arg_buffer(dy);
builder.arg(&w_ptr);
builder.arg(&alpha);
builder.arg(&checked_dims.m_i32);
builder.arg(&checked_dims.n_i32);
builder.arg(&checked_dims.k_i32);
enqueue_scalar_backward(
&mut control,
request,
operands,
"nt_m2n16_bk64_splitk32",
cfg,
&mut builder,
)
.map_err(|error| {
error.with_driver_context(format_args!("nt_m2n16_bk64_splitk32 backward_dx"))
})?;
return Ok(());
}
if matches!(
scalar_plan,
ScalarDispatchPlan::NtD768TransposeM64N64Qualified
| ScalarDispatchPlan::NtD768OutTransposeM64N64Qualified
| ScalarDispatchPlan::NtD768OutSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtD768InSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtPrismSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtLargeDeepSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtLargeDeepTransposeM64N64Qualified
| ScalarDispatchPlan::NtPrismVectorQualified
| ScalarDispatchPlan::NtD128OutTransposeM64N64Qualified
) {
use cudarc::driver::DevicePtr;
let transpose_scratch = kernels.transpose_scratch_buf(stream)?;
let required =
scalar_transpose_scratch_elements(request, scalar_plan)?.ok_or_else(|| {
"qualified scalar NT transpose plan lost its scratch extent".to_string()
})?;
if transpose_scratch.len() != super::contract::SCALAR_TRANSPOSE_SCRATCH_CAP_ELEMENTS
|| required > transpose_scratch.len()
{
return Err(format!(
"qualified scalar NT transpose scratch has {} f32 elements, requires {required} with exact capacity {}",
transpose_scratch.len(),
super::contract::SCALAR_TRANSPOSE_SCRATCH_CAP_ELEMENTS
));
}
let (w_t_ptr, _) = transpose_scratch.device_ptr(stream);
let rows = checked_dims.k_i32;
let columns = checked_dims.n_i32;
let transpose_cfg = cudarc::driver::LaunchConfig {
grid_dim: (
checked_dims.n_u32.div_ceil(32),
checked_dims.k_u32.div_ceil(32),
1,
),
block_dim: (32, 16, 1),
shared_mem_bytes: 0,
};
let mut transpose =
scalar_launch_builder(stream, &kernels.gemm_bi_transpose_f32_32x16_d768, &control);
transpose.arg(&w_t_ptr);
transpose.arg(&w_ptr);
transpose.arg(&rows);
transpose.arg(&columns);
enqueue_scalar_backward(
&mut control,
request,
operands,
"transpose_f32_32x16_d768",
transpose_cfg,
&mut transpose,
)
.map_err(|error| {
error.with_driver_context(format_args!("transpose_f32_32x16_d768 backward_dx"))
})?;
let params = SgbNnM64N64Params {
alpha,
beta: 0.0,
m: checked_dims.m_i32,
n: checked_dims.k_i32,
k: checked_dims.n_i32,
lda: checked_dims.n_i32,
ldb: checked_dims.k_i32,
ldc: checked_dims.k_i32,
};
let uses_fixed_copyplan = matches!(
scalar_plan,
ScalarDispatchPlan::NtD768OutSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtD768InSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtPrismSm89FixedCopyPlanQualified
| ScalarDispatchPlan::NtLargeDeepSm89FixedCopyPlanQualified
);
let m64_cfg = cudarc::driver::LaunchConfig {
grid_dim: (
checked_grid_product(
checked_dims.m_u32.div_ceil(64),
checked_dims.k_u32.div_ceil(64),
1,
)?,
1,
1,
),
block_dim: (128, 1, 1),
shared_mem_bytes: if uses_fixed_copyplan {
0
} else {
super::contract::SCALAR_NN_M64N64_DYNAMIC_SHARED_BYTES
},
};
let bias = 0_u64;
let (m64_function, m64_symbol) = if uses_fixed_copyplan {
(
kernels
.fixed_sm89_f32_n64_copyplan
.as_ref()
.ok_or_else(|| "qualified Fixed CopyPlan kernel is unavailable".to_string())?,
"nn_sm89_f32_n64_copyplan",
)
} else if scalar_plan == ScalarDispatchPlan::NtPrismVectorQualified {
(
&kernels.gemm_bi_nn_prism_m64n64_bk16_s2,
"nn_prism_m64n64_bk16_s2",
)
} else {
(&kernels.gemm_bi_nn_m64n64_bk16_s2, "nn_m64n64_bk16_s2")
};
let mut m64 = scalar_launch_builder(stream, m64_function, &control);
m64.arg_buffer_mut(dx);
m64.arg_buffer(dy);
m64.arg(&w_t_ptr);
m64.arg(&bias);
m64.arg(¶ms);
enqueue_scalar_backward(
&mut control,
request,
operands,
m64_symbol,
m64_cfg,
&mut m64,
)
.map_err(|error| {
error.with_driver_context(format_args!(
"{m64_symbol} qualified NT transpose backward_dx"
))
})?;
return Ok(());
}
if matches!(scalar_plan, ScalarDispatchPlan::NtNarrow) {
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_i = checked_dims.k_i32;
let num_pid_m = checked_dims.m_u32.div_ceil(64);
let num_pid_n = checked_dims.k_u32.div_ceil(32);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_grid_product(num_pid_m, num_pid_n, 1)?, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = scalar_launch_builder(stream, &kernels.gemm_bi_nt_narrow, &control);
builder.arg_buffer_mut(dx);
builder.arg_buffer(dy);
builder.arg(&w_ptr);
builder.arg(&alpha);
builder.arg(&m_i);
builder.arg(&n_i);
builder.arg(&k_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"nt_narrow",
cfg,
&mut builder,
)
.map_err(|error| error.with_driver_context(format_args!("nt_narrow backward_dx")))?;
return Ok(());
}
if matches!(scalar_plan, ScalarDispatchPlan::NtSmallBatchWide) {
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_i = checked_dims.k_i32;
let num_pid_m = checked_dims.m_u32.div_ceil(64);
let num_pid_n = checked_dims.k_u32.div_ceil(32);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_grid_product(num_pid_m, num_pid_n, 1)?, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = scalar_launch_builder(stream, &kernels.gemm_bi_nt_narrow, &control);
builder.arg_buffer_mut(dx);
builder.arg_buffer(dy);
builder.arg(&w_ptr);
builder.arg(&alpha);
builder.arg(&m_i);
builder.arg(&n_i);
builder.arg(&k_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"nt_narrow",
cfg,
&mut builder,
)
.map_err(|error| {
error.with_driver_context(format_args!("nt_narrow (small-batch wide-N) backward_dx"))
})?;
return Ok(());
}
if matches!(scalar_plan, ScalarDispatchPlan::NtGemv) {
let m_i = checked_dims.m_i32;
let k_i = checked_dims.k_i32;
let ldx_i = checked_dims.k_i32;
let ldy_i: i32 = 1;
let total = checked_dims.mk_u32;
let block = 256u32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (total.div_ceil(block), 1, 1),
block_dim: (block, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = scalar_launch_builder(stream, &kernels.gemm_bi_nt_gemv, &control);
builder.arg_buffer_mut(dx);
builder.arg_buffer(dy);
builder.arg(&w_ptr);
builder.arg(&alpha);
builder.arg(&m_i);
builder.arg(&k_i);
builder.arg(&ldx_i);
builder.arg(&ldy_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"nt_gemv",
cfg,
&mut builder,
)
.map_err(|error| error.with_driver_context(format_args!("nt_gemv backward_dx")))?;
return Ok(());
}
if let ScalarDispatchPlan::NtSplitKTail {
k_main,
k_tail: k_tail_cnt,
} = scalar_plan
{
let rows_i = checked_i32(k_main, "NT K-tail rows")?;
let cols_i = checked_dims.n_i32;
let t_grid_x = checked_dims.n_u32.div_ceil(32);
let t_grid_y = checked_u32(k_main, "NT K-tail rows")?.div_ceil(32);
let t_cfg = cudarc::driver::LaunchConfig {
grid_dim: (t_grid_x, t_grid_y, 1),
block_dim: (32, 32, 1),
shared_mem_bytes: 0,
};
let w_t_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _r) = kernels.transpose_scratch_buf(stream)?.device_ptr(stream);
ptr
};
let mut tb = scalar_launch_builder(stream, &kernels.gemm_bi_transpose_f32_2d, &control);
tb.arg(&w_t_ptr);
tb.arg(&w_ptr);
tb.arg(&rows_i);
tb.arg(&cols_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"transpose_f32_2d",
t_cfg,
&mut tb,
)
.map_err(|error| error.with_driver_context(format_args!("transpose_f32_2d (K-tail)")))?;
let m_i = checked_dims.m_i32;
let k_main_i = checked_i32(k_main, "NT K-tail columns")?;
let k_chunks = checked_i32(n_out / 32, "NT K-tail chunks")?;
let lda_dy_i = checked_dims.n_i32;
let num_pid_m = checked_dims.m_u32.div_ceil(32);
let num_pid_n = checked_u32(k_main, "NT K-tail columns")?.div_ceil(64);
let partial_cfg = cudarc::driver::LaunchConfig {
grid_dim: (
checked_grid_product(
num_pid_m,
num_pid_n,
checked_u32(n_out / 32, "NT K-tail chunks")?,
)?,
1,
1,
),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let partial_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _r) = kernels.splitk_scratch_buf(stream)?.device_ptr(stream);
ptr
};
let mut pb = scalar_launch_builder(stream, &kernels.gemm_bi_nn_splitk32_partial, &control);
pb.arg(&partial_ptr);
pb.arg_buffer(dy);
pb.arg(&w_t_ptr);
pb.arg(&m_i);
pb.arg(&k_main_i);
pb.arg(&k_chunks);
pb.arg(&lda_dy_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"nn_splitk32_partial",
partial_cfg,
&mut pb,
)
.map_err(|error| {
error.with_driver_context(format_args!("nn_splitk32_partial (NT K-tail main)"))
})?;
let null_tail: u64 = 0;
let null_bias: u64 = 0;
let zero_i32: i32 = 0;
let out_stride_i = checked_dims.k_i32;
let total_main = checked_u32(
batch
.checked_mul(k_main)
.ok_or_else(|| invalid_gemm_dimensions("NT K-tail output total overflows usize"))?,
"NT K-tail output total",
)?;
let reduce_cfg = cudarc::driver::LaunchConfig {
grid_dim: (total_main.div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let mut rb = scalar_launch_builder(stream, &kernels.gemm_bi_splitk_reduce, &control);
rb.arg_buffer_mut(dx);
rb.arg(&partial_ptr);
rb.arg(&null_bias);
rb.arg(&null_tail);
rb.arg(&null_tail);
rb.arg(&alpha);
rb.arg(&m_i);
rb.arg(&k_main_i);
rb.arg(&k_chunks);
rb.arg(&zero_i32);
rb.arg(&out_stride_i); rb.arg(&zero_i32); enqueue_scalar_backward(
&mut control,
request,
operands,
"splitk_reduce",
reduce_cfg,
&mut rb,
)
.map_err(|error| {
error.with_driver_context(format_args!("splitk_reduce (NT K-tail main)"))
})?;
let w_base_ptr = w_ptr;
let n_i = checked_dims.n_i32;
let block = 128u32;
let tail_cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_dims.m_u32.div_ceil(block), 1, 1),
block_dim: (block, 1, 1),
shared_mem_bytes: 0,
};
for k in 0..k_tail_cnt {
let k_tail_col = k_main + k;
let row_elements = k_tail_col
.checked_mul(n_out)
.ok_or_else(|| invalid_gemm_dimensions("W tail row offset overflows usize"))?;
let w_tail_row_ptr = checked_ptr_add(
w_base_ptr,
checked_byte_offset(row_elements, std::mem::size_of::<f32>(), "W tail row")?,
"W tail row",
)?;
let col_idx_i = checked_i32(k_tail_col, "NT K-tail column")?;
let mut gb = scalar_launch_builder(stream, &kernels.gemm_bi_dx_col_gemv, &control);
gb.arg_buffer_mut(dx);
gb.arg_buffer(dy);
gb.arg(&w_tail_row_ptr);
gb.arg(&m_i);
gb.arg(&n_i);
gb.arg(&col_idx_i);
gb.arg(&out_stride_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"dx_col_gemv",
tail_cfg,
&mut gb,
)
.map_err(|error| {
error.with_driver_context(format_args!("dx_col_gemv (NT K-tail col={})", k))
})?;
}
return Ok(());
}
if let ScalarDispatchPlan::NtSplitKMain {
n_main: n_main_nt,
n_tail: n_tail_nt,
} = scalar_plan
{
let rows_i = checked_dims.k_i32;
let cols_i = checked_dims.n_i32;
let t_grid_x = checked_dims.n_u32.div_ceil(32);
let t_grid_y = checked_dims.k_u32.div_ceil(32);
let t_cfg = cudarc::driver::LaunchConfig {
grid_dim: (t_grid_x, t_grid_y, 1),
block_dim: (32, 32, 1),
shared_mem_bytes: 0,
};
let w_t_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _r) = kernels.transpose_scratch_buf(stream)?.device_ptr(stream);
ptr
};
let mut tb = scalar_launch_builder(stream, &kernels.gemm_bi_transpose_f32_2d, &control);
tb.arg(&w_t_ptr);
tb.arg(&w_ptr);
tb.arg(&rows_i);
tb.arg(&cols_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"transpose_f32_2d",
t_cfg,
&mut tb,
)
.map_err(|error| error.with_driver_context(format_args!("transpose_f32_2d")))?;
let m_i = checked_dims.m_i32;
let k_out_i = checked_dims.k_i32;
let k_chunks = checked_i32(n_main_nt / 32, "NT split-K chunks")?;
let num_pid_m = checked_dims.m_u32.div_ceil(32);
let num_pid_n = checked_dims.k_u32.div_ceil(64);
let partial_cfg = cudarc::driver::LaunchConfig {
grid_dim: (
checked_grid_product(
num_pid_m,
num_pid_n,
checked_u32(n_main_nt / 32, "NT split-K chunks")?,
)?,
1,
1,
),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let partial_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _r) = kernels.splitk_scratch_buf(stream)?.device_ptr(stream);
ptr
};
let lda_i = checked_dims.n_i32; let mut pb = scalar_launch_builder(stream, &kernels.gemm_bi_nn_splitk32_partial, &control);
pb.arg(&partial_ptr);
pb.arg_buffer(dy);
pb.arg(&w_t_ptr);
pb.arg(&m_i);
pb.arg(&k_out_i);
pb.arg(&k_chunks);
pb.arg(&lda_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"nn_splitk32_partial",
partial_cfg,
&mut pb,
)
.map_err(|error| {
error.with_driver_context(format_args!("nn_splitk32_partial (NT-via-T N-tail)"))
})?;
let null_bias: u64 = 0;
let total = checked_dims.mk_u32;
let reduce_cfg = cudarc::driver::LaunchConfig {
grid_dim: (total.div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let zero_i32: i32 = 0;
let tail_cnt_i = checked_i32(n_tail_nt, "NT reduction tail")?;
let dy_tail_stride_i = checked_dims.n_i32; let (dy_tail_ptr, wt_tail_ptr): (u64, u64) = if n_tail_nt > 0 {
let dyp = checked_ptr_add(
dy.scalar_ptr(),
checked_byte_offset(n_main_nt, std::mem::size_of::<f32>(), "dY tail")?,
"dY tail",
)?;
let wt_elements = n_main_nt.checked_mul(n_in).ok_or_else(|| {
invalid_gemm_dimensions("transposed W tail offset overflows usize")
})?;
let wtp = checked_ptr_add(
w_t_ptr,
checked_byte_offset(wt_elements, std::mem::size_of::<f32>(), "transposed W tail")?,
"transposed W tail",
)?;
(dyp, wtp)
} else {
(0, 0)
};
let mut rb = scalar_launch_builder(stream, &kernels.gemm_bi_splitk_reduce, &control);
rb.arg_buffer_mut(dx);
rb.arg(&partial_ptr);
rb.arg(&null_bias);
rb.arg(&dy_tail_ptr);
rb.arg(&wt_tail_ptr);
rb.arg(&alpha);
rb.arg(&m_i);
rb.arg(&k_out_i);
rb.arg(&k_chunks);
rb.arg(&dy_tail_stride_i);
rb.arg(&zero_i32); rb.arg(&tail_cnt_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"splitk_reduce",
reduce_cfg,
&mut rb,
)
.map_err(|error| {
error.with_driver_context(format_args!("splitk_reduce (NT-via-T N-tail)"))
})?;
return Ok(());
}
const SLIM_NT_K_CHUNK: u32 = 64;
if let ScalarDispatchPlan::NtSplitKSlim { chunks: f_final } = scalar_plan {
let base_blocks = checked_grid_product(
checked_dims.m_u32.div_ceil(128),
checked_dims.k_u32.div_ceil(64),
1,
)?;
let k_chunk = SLIM_NT_K_CHUNK;
let rows_i = checked_dims.k_i32;
let cols_i = checked_dims.n_i32;
let t_grid_x = checked_dims.n_u32.div_ceil(32);
let t_grid_y = checked_dims.k_u32.div_ceil(32);
let t_cfg = cudarc::driver::LaunchConfig {
grid_dim: (t_grid_x, t_grid_y, 1),
block_dim: (32, 32, 1),
shared_mem_bytes: 0,
};
let w_t_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _r) = kernels.transpose_scratch_buf(stream)?.device_ptr(stream);
ptr
};
let mut tb = scalar_launch_builder(stream, &kernels.gemm_bi_transpose_f32_2d, &control);
tb.arg(&w_t_ptr);
tb.arg(&w_ptr);
tb.arg(&rows_i);
tb.arg(&cols_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"transpose_f32_2d",
t_cfg,
&mut tb,
)
.map_err(|error| error.with_driver_context(format_args!("transpose_f32_2d (slim NT)")))?;
let m_i = checked_dims.m_i32;
let k_out_i = checked_dims.k_i32; let k_full_i = checked_dims.n_i32; let lda_i = checked_dims.n_i32; let ldb_i = checked_dims.k_i32; let k_chunk_i = checked_i32(
checked_usize(k_chunk, "NT slim split-K chunk")?,
"NT slim split-K chunk",
)?;
let partial_ptr = {
use cudarc::driver::DevicePtr;
let (ptr, _r) = kernels.splitk_scratch_buf(stream)?.device_ptr(stream);
ptr
};
let partial_cfg = cudarc::driver::LaunchConfig {
grid_dim: (base_blocks, 1, f_final),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let mut pb =
scalar_launch_builder(stream, &kernels.gemm_bi_nn_splitk_slim_partial, &control);
pb.arg(&partial_ptr);
pb.arg_buffer(dy);
pb.arg(&w_t_ptr);
pb.arg(&m_i);
pb.arg(&k_out_i);
pb.arg(&k_full_i);
pb.arg(&lda_i);
pb.arg(&ldb_i);
pb.arg(&k_chunk_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"nn_splitk_slim_partial",
partial_cfg,
&mut pb,
)
.map_err(|error| {
error.with_driver_context(format_args!("nn_splitk_slim_partial (slim NT)"))
})?;
let null_bias: u64 = 0;
let null_tail: u64 = 0;
let zero_i32_nt: i32 = 0;
let f_i = checked_i32(
checked_usize(f_final, "NT slim split-K chunks")?,
"NT slim split-K chunks",
)?;
let total = checked_dims.mk_u32;
let reduce_cfg = cudarc::driver::LaunchConfig {
grid_dim: (total.div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let mut rb = scalar_launch_builder(stream, &kernels.gemm_bi_splitk_reduce, &control);
rb.arg_buffer_mut(dx);
rb.arg(&partial_ptr);
rb.arg(&null_bias);
rb.arg(&null_tail);
rb.arg(&null_tail);
rb.arg(&alpha);
rb.arg(&m_i);
rb.arg(&k_out_i);
rb.arg(&f_i);
rb.arg(&zero_i32_nt);
rb.arg(&zero_i32_nt);
rb.arg(&zero_i32_nt);
enqueue_scalar_backward(
&mut control,
request,
operands,
"splitk_reduce",
reduce_cfg,
&mut rb,
)
.map_err(|error| error.with_driver_context(format_args!("splitk_reduce (slim NT)")))?;
return Ok(());
}
if matches!(scalar_plan, ScalarDispatchPlan::NtMidBatchWide) {
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_i = checked_dims.k_i32;
let num_pid_m = checked_dims.m_u32.div_ceil(64);
let num_pid_n = checked_dims.k_u32.div_ceil(32);
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_grid_product(num_pid_m, num_pid_n, 1)?, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let mut builder = scalar_launch_builder(stream, &kernels.gemm_bi_nt_narrow, &control);
builder.arg_buffer_mut(dx);
builder.arg_buffer(dy);
builder.arg(&w_ptr);
builder.arg(&alpha);
builder.arg(&m_i);
builder.arg(&n_i);
builder.arg(&k_i);
enqueue_scalar_backward(
&mut control,
request,
operands,
"nt_narrow",
cfg,
&mut builder,
)
.map_err(|error| {
error.with_driver_context(format_args!("nt_narrow (gap-fill mid-batch wide-N)"))
})?;
return Ok(());
}
if let ScalarDispatchPlan::NtFinal { slim } = scalar_plan {
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_i = checked_dims.k_i32;
let (func, bn) = if slim {
(&kernels.gemm_bi_nt_slim, 64)
} else {
(&kernels.gemm_bi_nt, 128)
};
let threads = if slim { 128u32 } else { 256u32 };
let smem_bytes = if slim {
0
} else {
SCALAR_BIG_NT_DYNAMIC_SHARED_BYTES
};
let total_tiles = checked_tile_grid(checked_dims.m_u32, 128, checked_dims.k_u32, bn)?;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (total_tiles, 1, 1),
block_dim: (threads, 1, 1),
shared_mem_bytes: smem_bytes,
};
let mut builder = scalar_launch_builder(stream, func, &control);
builder.arg_buffer_mut(dx);
builder.arg_buffer(dy);
builder.arg(&w_ptr);
builder.arg(&alpha);
builder.arg(&m_i);
builder.arg(&n_i);
builder.arg(&k_i);
let symbol = if slim { "nt_slim" } else { "nt_big" };
enqueue_scalar_backward(&mut control, request, operands, symbol, cfg, &mut builder)
.map_err(|error| {
error.with_driver_context(format_args!(
"nt_big{} backward_dx",
if slim { "_slim" } else { "" }
))
})?;
return Ok(());
}
panic!(
"gpu_gemm_bi_backward_dx: cuBLAS fallback hit (shape M={batch} K={n_in} N={n_out}). \
The zero-cuBLAS contract requires every shape to route through a custom \
kernel — add a dispatcher branch in this function for this shape."
);
}
use super::super::blas::{HalfPhysicalTraceRequest, TypedPtr};
use super::super::dtype::WeightDtype;
fn require_half(dt: WeightDtype, what: &str) -> Result<(), String> {
if dt == WeightDtype::F32 {
return Err(format!(
"gemm_bi typed dispatch: {what} is f32 — use the f32 entry points"
));
}
Ok(())
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum HalfSchedule {
Tiled,
StreamKFixedOrder,
RelayChain,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct HalfKernelIdentity {
symbol: &'static str,
module_kind: ModuleKind,
schedule: HalfSchedule,
}
impl HalfKernelIdentity {
fn resolve(base: &str, dtype: WeightDtype) -> Result<Self, String> {
let module_kind = if matches!(
base,
"nn_sm89_m128n128_bk64_s3"
| "tn_sm89_m64n64_bk64_s2_compact_bxor"
| "tn_sm89_m64n64_bk64_s2_regpipe_vec2"
| "tn_sm89_m16n16_bk64_s2_ldb72"
| "tn_sm89_half_d128_in_m32n16_bk64_s4_cg"
| "tn_sm89_half_d128_out_m32n16_bk64_s4_cg"
| "nt_sm89_m16n64_bk64_s4"
| "nn_sm89_m16n64_bk64_s4"
| "nt_sm89_m128n128_bk64_s3_bxor"
| "nt_sm89_m96n128_bk64_s3"
| "tn_sm89_relay_m64n64_bk64_s3"
) {
ModuleKind::TriadSm89Half
} else if matches!(
base,
"nn_tc"
| "nn_tc64"
| "nn_tc16"
| "tn_tc"
| "tn_tc64"
| "tn_tc64_streamk"
| "tn_tc128x64"
| "nt_tc"
| "nt_tc64"
) {
ModuleKind::TriadSm80
} else {
ModuleKind::TriadScalar
};
let schedule = match base {
"tn_tc64_streamk" => HalfSchedule::StreamKFixedOrder,
SM89_HALF_RELAY_BASE => HalfSchedule::RelayChain,
_ => HalfSchedule::Tiled,
};
let symbol = match (base, dtype) {
("nn_gemv", WeightDtype::Bf16) => "nn_gemv_bf16",
("nn_gemv", WeightDtype::F16) => "nn_gemv_f16",
("nn_ultra_thin", WeightDtype::Bf16) => "nn_ultra_thin_bf16",
("nn_ultra_thin", WeightDtype::F16) => "nn_ultra_thin_f16",
("nn_narrow", WeightDtype::Bf16) => "nn_narrow_bf16",
("nn_narrow", WeightDtype::F16) => "nn_narrow_f16",
("nn_narrow_small", WeightDtype::Bf16) => "nn_narrow_small_bf16",
("nn_narrow_small", WeightDtype::F16) => "nn_narrow_small_f16",
("nn_big", WeightDtype::Bf16) => "nn_big_bf16",
("nn_big", WeightDtype::F16) => "nn_big_f16",
("tn_gemv", WeightDtype::Bf16) => "tn_gemv_bf16",
("tn_gemv", WeightDtype::F16) => "tn_gemv_f16",
("tn_narrow", WeightDtype::Bf16) => "tn_narrow_bf16",
("tn_narrow", WeightDtype::F16) => "tn_narrow_f16",
("tn_big", WeightDtype::Bf16) => "tn_big_bf16",
("tn_big", WeightDtype::F16) => "tn_big_f16",
("nt_gemv", WeightDtype::Bf16) => "nt_gemv_bf16",
("nt_gemv", WeightDtype::F16) => "nt_gemv_f16",
("nt_narrow", WeightDtype::Bf16) => "nt_narrow_bf16",
("nt_narrow", WeightDtype::F16) => "nt_narrow_f16",
("nt_big", WeightDtype::Bf16) => "nt_big_bf16",
("nt_big", WeightDtype::F16) => "nt_big_f16",
("nn_tc", WeightDtype::Bf16) => "nn_tc_bf16",
("nn_tc", WeightDtype::F16) => "nn_tc_f16",
("nn_tc64", WeightDtype::Bf16) => "nn_tc64_bf16",
("nn_tc64", WeightDtype::F16) => "nn_tc64_f16",
("nn_tc16", WeightDtype::Bf16) => "nn_tc16_bf16",
("nn_tc16", WeightDtype::F16) => "nn_tc16_f16",
("tn_tc", WeightDtype::Bf16) => "tn_tc_bf16",
("tn_tc", WeightDtype::F16) => "tn_tc_f16",
("tn_tc64", WeightDtype::Bf16) => "tn_tc64_bf16",
("tn_tc64", WeightDtype::F16) => "tn_tc64_f16",
("tn_tc64_streamk", WeightDtype::Bf16) => "tn_tc64_streamk_bf16",
("tn_tc64_streamk", WeightDtype::F16) => "tn_tc64_streamk_f16",
("tn_tc128x64", WeightDtype::Bf16) => "tn_tc128x64_bf16",
("tn_tc128x64", WeightDtype::F16) => "tn_tc128x64_f16",
("nt_tc", WeightDtype::Bf16) => "nt_tc_bf16",
("nt_tc", WeightDtype::F16) => "nt_tc_f16",
("nt_tc64", WeightDtype::Bf16) => "nt_tc64_bf16",
("nt_tc64", WeightDtype::F16) => "nt_tc64_f16",
("nn_sm89_m128n128_bk64_s3", WeightDtype::Bf16) => "nn_sm89_m128n128_bk64_s3_bf16",
("nn_sm89_m128n128_bk64_s3", WeightDtype::F16) => "nn_sm89_m128n128_bk64_s3_f16",
("tn_sm89_m64n64_bk64_s2_compact_bxor", WeightDtype::Bf16) => {
"tn_sm89_m64n64_bk64_s2_compact_bxor_bf16"
}
("tn_sm89_m64n64_bk64_s2_compact_bxor", WeightDtype::F16) => {
"tn_sm89_m64n64_bk64_s2_compact_bxor_f16"
}
("tn_sm89_m64n64_bk64_s2_regpipe_vec2", WeightDtype::Bf16) => {
"tn_sm89_m64n64_bk64_s2_regpipe_vec2_bf16"
}
("tn_sm89_m64n64_bk64_s2_regpipe_vec2", WeightDtype::F16) => {
"tn_sm89_m64n64_bk64_s2_regpipe_vec2_f16"
}
("tn_sm89_m16n16_bk64_s2_ldb72", WeightDtype::Bf16) => {
"tn_sm89_m16n16_bk64_s2_ldb72_bf16"
}
("tn_sm89_m16n16_bk64_s2_ldb72", WeightDtype::F16) => {
"tn_sm89_m16n16_bk64_s2_ldb72_f16"
}
("tn_sm89_half_d128_in_m32n16_bk64_s4_cg", WeightDtype::Bf16) => {
super::sm89_half_d128_source::D128_IN_BF16_SYMBOL
}
("tn_sm89_half_d128_in_m32n16_bk64_s4_cg", WeightDtype::F16) => {
super::sm89_half_d128_source::D128_IN_F16_SYMBOL
}
("tn_sm89_half_d128_out_m32n16_bk64_s4_cg", WeightDtype::Bf16) => {
super::sm89_half_d128_source::D128_OUT_BF16_SYMBOL
}
("tn_sm89_half_d128_out_m32n16_bk64_s4_cg", WeightDtype::F16) => {
super::sm89_half_d128_source::D128_OUT_F16_SYMBOL
}
("tn_sm89_relay_m64n64_bk64_s3", WeightDtype::Bf16) => {
super::sm89_half_relay_source::RELAY_BF16_SYMBOL
}
("tn_sm89_relay_m64n64_bk64_s3", WeightDtype::F16) => {
super::sm89_half_relay_source::RELAY_F16_SYMBOL
}
("nt_sm89_m16n64_bk64_s4", WeightDtype::Bf16) => {
super::sm89_half_small_source::NT_BF16_SYMBOL
}
("nt_sm89_m16n64_bk64_s4", WeightDtype::F16) => {
super::sm89_half_small_source::NT_F16_SYMBOL
}
("nn_sm89_m16n64_bk64_s4", WeightDtype::Bf16) => {
super::sm89_half_small_source::NN_BF16_SYMBOL
}
("nn_sm89_m16n64_bk64_s4", WeightDtype::F16) => {
super::sm89_half_small_source::NN_F16_SYMBOL
}
("nt_sm89_m128n128_bk64_s3_bxor", WeightDtype::Bf16) => {
"nt_sm89_m128n128_bk64_s3_bxor_bf16"
}
("nt_sm89_m128n128_bk64_s3_bxor", WeightDtype::F16) => {
"nt_sm89_m128n128_bk64_s3_bxor_f16"
}
("nt_sm89_m96n128_bk64_s3", WeightDtype::Bf16) => "nt_sm89_m96n128_bk64_s3_bf16",
("nt_sm89_m96n128_bk64_s3", WeightDtype::F16) => "nt_sm89_m96n128_bk64_s3_f16",
(_, WeightDtype::F32) => {
return Err("half kernel identity does not accept f32".into());
}
_ => return Err(format!("unknown half kernel base {base}")),
};
Ok(Self {
symbol,
module_kind,
schedule,
})
}
#[cfg(test)]
fn validate(self, module_kind: ModuleKind, symbol: &str) -> Result<(), String> {
if module_kind != self.module_kind || symbol != self.symbol {
return Err(
"half kernel identity does not match its exact symbol and module owner".into(),
);
}
Ok(())
}
}
struct HalfLaunchEnvironment<'a, O> {
stream: &'a Arc<CudaStream>,
kernels: &'a GpuKernels,
context: Option<&'a GpuCtx>,
observer: O,
}
impl<'a> HalfLaunchEnvironment<'a, NoPhysicalObserver> {
fn production(stream: &'a Arc<CudaStream>, kernels: &'a GpuKernels) -> Self {
Self {
stream,
kernels,
context: None,
observer: NoPhysicalObserver,
}
}
}
impl<'a, O: PhysicalLaunchObserver> HalfLaunchEnvironment<'a, O> {
fn observed(ctx: &'a GpuCtx, observer: O) -> Self {
Self {
stream: &ctx.stream,
kernels: &ctx.kernels,
context: Some(ctx),
observer,
}
}
}
const _: [(); 3 * std::mem::size_of::<usize>()] =
[(); std::mem::size_of::<HalfLaunchEnvironment<'static, NoPhysicalObserver>>()];
#[derive(Clone, Copy)]
struct HalfGemmArguments {
output: CUptr,
a: CUptr,
b: CUptr,
bias: CUptr,
}
#[derive(Clone, Copy)]
struct HalfGemmObservation {
base: &'static str,
op: ResolvedGemmOp,
dtype: WeightDtype,
dims: (usize, usize, usize),
strides: (usize, usize, usize),
tile: (u32, u32),
bk_stages: (u32, u8),
arguments: HalfGemmArguments,
}
impl HalfGemmObservation {
fn shape(self) -> F32TriadShape {
F32TriadShape {
m: self.dims.0,
k: self.dims.1,
n: self.dims.2,
lda: self.strides.0,
ldb: self.strides.1,
ldc: self.strides.2,
}
}
}
#[derive(Clone, Copy)]
struct HalfKernelChoice<'a> {
base: &'static str,
function: &'a CudaFunction,
}
impl<'a> HalfKernelChoice<'a> {
fn new(base: &'static str, function: &'a CudaFunction) -> Self {
Self { base, function }
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(in crate::mamba_ssm::gpu) struct HalfNativeBranchSeal {
pub(in crate::mamba_ssm::gpu) base: &'static str,
pub(in crate::mamba_ssm::gpu) op: ResolvedGemmOp,
pub(in crate::mamba_ssm::gpu) dtype: WeightDtype,
pub(in crate::mamba_ssm::gpu) dims: (usize, usize, usize),
pub(in crate::mamba_ssm::gpu) strides: (usize, usize, usize),
pub(in crate::mamba_ssm::gpu) tile: (u32, u32),
pub(in crate::mamba_ssm::gpu) bk_stages: (u32, u8),
pub(in crate::mamba_ssm::gpu) grid_dim: (u32, u32, u32),
pub(in crate::mamba_ssm::gpu) block_dim: (u32, u32, u32),
pub(in crate::mamba_ssm::gpu) shared_mem_bytes: u32,
}
fn half_native_branch_seal(
observation: HalfGemmObservation,
config: LaunchConfig,
) -> HalfNativeBranchSeal {
HalfNativeBranchSeal {
base: observation.base,
op: observation.op,
dtype: observation.dtype,
dims: observation.dims,
strides: observation.strides,
tile: observation.tile,
bk_stages: observation.bk_stages,
grid_dim: config.grid_dim,
block_dim: config.block_dim,
shared_mem_bytes: config.shared_mem_bytes,
}
}
fn half_policy_dtype(dtype: WeightDtype) -> Result<PolicyDtype, String> {
match dtype {
WeightDtype::Bf16 => Ok(PolicyDtype::Bf16),
WeightDtype::F16 => Ok(PolicyDtype::F16),
WeightDtype::F32 | WeightDtype::Tf32 => {
Err("half physical launch identity does not accept f32".into())
}
}
}
fn half_gemm_arguments_digest(
argument_identity: impl Fn(CUptr, u64) -> Result<Sha256Digest, String>,
observation: HalfGemmObservation,
identity: HalfKernelIdentity,
dtype: PolicyDtype,
) -> Result<Sha256Digest, String> {
let shape = observation.shape();
let element_bytes = u64::try_from(observation.dtype.size_bytes())
.map_err(|_| "half element width exceeds u64::MAX".to_string())?;
let matrix_bytes = |elements: usize, name: &str| {
u64::try_from(elements)
.ok()
.and_then(|elements| elements.checked_mul(element_bytes))
.ok_or_else(|| format!("physical half {name} span overflows u64"))
};
let output_element_bytes = if observation.op == ResolvedGemmOp::Tn {
4
} else {
element_bytes
};
let output_elements = shape
.output_rows(observation.op)
.checked_mul(shape.output_columns(observation.op))
.ok_or_else(|| "physical half output span overflows usize".to_string())?;
let output_bytes = u64::try_from(output_elements)
.ok()
.and_then(|elements| elements.checked_mul(output_element_bytes))
.ok_or_else(|| "physical half output span overflows u64".to_string())?;
let elements = |rows: usize, columns: usize, name: &str| {
rows.checked_mul(columns)
.ok_or_else(|| format!("physical half {name} span overflows usize"))
};
let (a_elements, b_elements) = match observation.op {
ResolvedGemmOp::Nn => (
elements(shape.m, shape.k, "A")?,
elements(shape.k, shape.n, "B")?,
),
ResolvedGemmOp::Tn => (
elements(shape.m, shape.k, "A")?,
elements(shape.m, shape.n, "B")?,
),
ResolvedGemmOp::Nt => (
elements(shape.m, shape.n, "A")?,
elements(shape.k, shape.n, "B")?,
),
};
let output_identity = argument_identity(observation.arguments.output, output_bytes)?;
let a_identity = argument_identity(observation.arguments.a, matrix_bytes(a_elements, "A")?)?;
let b_identity = argument_identity(observation.arguments.b, matrix_bytes(b_elements, "B")?)?;
let bias_identity = if observation.arguments.bias == 0 {
None
} else {
let bytes = u64::try_from(shape.output_columns(observation.op))
.ok()
.and_then(|columns| columns.checked_mul(4))
.ok_or_else(|| "physical half bias span overflows u64".to_string())?;
Some(argument_identity(observation.arguments.bias, bytes)?)
};
let alpha = 1.0_f32;
let beta = if observation.op == ResolvedGemmOp::Tn {
1.0_f32
} else {
0.0_f32
};
Ok(FramedSha256::new(b"triad-half-kernel-arguments.v2")
.required(b"symbol", identity.symbol.as_bytes())
.required(b"op", &[observation.op as u8])
.required(b"dtype", &[dtype as u8])
.required(b"m", &(shape.m as u64).to_le_bytes())
.required(b"k", &(shape.k as u64).to_le_bytes())
.required(b"n", &(shape.n as u64).to_le_bytes())
.required(b"lda", &(shape.lda as u64).to_le_bytes())
.required(b"ldb", &(shape.ldb as u64).to_le_bytes())
.required(b"ldc", &(shape.ldc as u64).to_le_bytes())
.required(b"alpha", &alpha.to_bits().to_le_bytes())
.required(b"beta", &beta.to_bits().to_le_bytes())
.required(
b"null-pointer-mask",
&(u64::from(observation.arguments.bias == 0) << 3).to_le_bytes(),
)
.required(b"output-allocation", &output_identity)
.required(b"a-allocation", &a_identity)
.required(b"b-allocation", &b_identity)
.optional(
b"bias-allocation",
bias_identity.as_ref().map(Sha256Digest::as_slice),
)
.required(b"tile-m", &observation.tile.0.to_le_bytes())
.required(b"tile-n", &observation.tile.1.to_le_bytes())
.required(b"bk", &observation.bk_stages.0.to_le_bytes())
.required(b"stages", &[observation.bk_stages.1])
.finish())
}
fn half_gemm_resources_digest(identity: HalfKernelIdentity, config: LaunchConfig) -> Sha256Digest {
let static_shared_bytes = if identity.module_kind == ModuleKind::TriadSm89Half {
super::sm89_half_source::runtime_kernel_spec(identity.symbol)
.map(|spec| spec.static_shared_bytes)
.unwrap_or(0)
} else {
0
};
FramedSha256::new(b"triad-half-kernel-resources.v2")
.required(b"symbol", identity.symbol.as_bytes())
.required(b"module-kind", &[identity.module_kind as u8])
.required(
b"threads",
&(config.block_dim.0 * config.block_dim.1 * config.block_dim.2).to_le_bytes(),
)
.required(
b"dynamic-shared-memory-bytes",
&config.shared_mem_bytes.to_le_bytes(),
)
.required(
b"static-shared-memory-bytes",
&static_shared_bytes.to_le_bytes(),
)
.finish()
}
fn half_kernel_compiler_identity(
kernels: &GpuKernels,
identity: HalfKernelIdentity,
) -> Result<crate::mamba_ssm::gpu::kernel_identity::CompilerIdentity, String> {
match identity.module_kind {
ModuleKind::TriadScalar => Ok(kernels.triad_scalar_compiler_identity()),
ModuleKind::TriadSm80 => Ok(kernels.triad_sm80_compiler_identity()),
ModuleKind::TriadSm89Half => kernels
.triad_sm89_half_compiler_identity()
.ok_or_else(|| "SM89 half route has no compiler identity".to_string()),
module_kind => Err(format!(
"half physical GEMM has unsupported module owner {module_kind:?}"
)),
}
}
fn resolved_half_gemm_route(
context: GemmRouteIdentity,
kernels: &GpuKernels,
argument_identity: impl Fn(CUptr, u64) -> Result<Sha256Digest, String>,
observation: HalfGemmObservation,
identity: HalfKernelIdentity,
config: LaunchConfig,
) -> Result<ResolvedGemmRoute, String> {
resolved_half_gemm_route_with_compiler(
context,
half_kernel_compiler_identity(kernels, identity)?,
argument_identity,
observation,
identity,
config,
)
}
fn resolved_half_gemm_route_with_compiler(
context: GemmRouteIdentity,
compiler: crate::mamba_ssm::gpu::kernel_identity::CompilerIdentity,
argument_identity: impl Fn(CUptr, u64) -> Result<Sha256Digest, String>,
observation: HalfGemmObservation,
identity: HalfKernelIdentity,
config: LaunchConfig,
) -> Result<ResolvedGemmRoute, String> {
let dtype = half_policy_dtype(observation.dtype)?;
let shape = observation.shape();
let (artifact, backend, numeric_contract, instruction_family, instruction_shape) =
match identity.module_kind {
ModuleKind::TriadScalar => (
context.artifacts.triad_scalar,
PhysicalGemmBackend::ScalarFma,
ResolvedNumericContract::ScalarFma,
ResolvedInstructionFamily::ScalarFma,
ResolvedInstructionShape { m: 1, n: 1, k: 1 },
),
ModuleKind::TriadSm80 => (
context.artifacts.triad_sm80,
PhysicalGemmBackend::Sm80Mma16,
match identity.schedule {
HalfSchedule::Tiled => ResolvedNumericContract::MmaSyncF32,
HalfSchedule::StreamKFixedOrder => {
ResolvedNumericContract::MmaSyncF32StreamKFixedOrder
}
HalfSchedule::RelayChain => {
return Err("TriadSm80 owns no relay schedule".into());
}
},
ResolvedInstructionFamily::MmaSync,
ResolvedInstructionShape { m: 16, n: 8, k: 16 },
),
ModuleKind::TriadSm89Half => (
context
.artifacts
.sm89_half
.ok_or_else(|| "SM89 half route has no artifact identity".to_string())?,
match super::sm89_half_source::runtime_kernel_spec(identity.symbol)
.map(|spec| spec.stages)
{
Some(2) => PhysicalGemmBackend::Sm89Mma16HalfS2,
Some(4) => PhysicalGemmBackend::Sm89Mma16HalfS4,
_ => PhysicalGemmBackend::Sm89Mma16HalfS3,
},
ResolvedNumericContract::MmaSyncF32,
ResolvedInstructionFamily::MmaSync,
ResolvedInstructionShape { m: 16, n: 8, k: 16 },
),
module_kind => {
return Err(format!(
"half physical GEMM has unsupported module owner {module_kind:?}"
));
}
};
let launch = ResolvedKernelLaunch {
grid_dim: config.grid_dim,
block_dim: config.block_dim,
shared_mem_bytes: config.shared_mem_bytes,
arguments_digest: half_gemm_arguments_digest(
argument_identity,
observation,
identity,
dtype,
)?,
};
Ok(ResolvedGemmRoute {
op: observation.op,
dtype,
backend,
numeric_contract,
instruction_family,
instruction_shape,
operand_conversion: ResolvedOperandConversion::None,
ownership: match identity.schedule {
HalfSchedule::Tiled => ResolvedOutputOwnership::OneCtaPerOutputTile,
HalfSchedule::StreamKFixedOrder => {
ResolvedOutputOwnership::OwnerCtaPerOutputTileStreamKFixedOrder
}
HalfSchedule::RelayChain => ResolvedOutputOwnership::RelayCtaChainPerOutputTile,
},
symbol: identity.symbol,
module_kind: identity.module_kind,
target: compiler.target,
artifact,
compiler,
device: context.device,
device_caps: context.device_caps,
shape: (shape.m, shape.k, shape.n),
strides: (shape.lda, shape.ldb, shape.ldc),
tile: observation.tile,
bk: observation.bk_stages.0,
stages: observation.bk_stages.1,
threads: launch.block_dim.0 * launch.block_dim.1 * launch.block_dim.2,
launch,
tensor_map_revision: 0,
tensor_maps_digest: [0; 32],
resources_digest: half_gemm_resources_digest(identity, config),
tuning_table_revision: if identity.module_kind == ModuleKind::TriadSm89Half {
SM89_HALF_ROUTE_REVISION
} else {
TUNING_TABLE_REVISION
},
schedule_revision: SCHEDULE_REVISION,
})
}
fn resolve_half_gemm_observation_with_context<O: PhysicalLaunchObserver>(
observer: &O,
context: GemmRouteIdentity,
compiler: crate::mamba_ssm::gpu::kernel_identity::CompilerIdentity,
observation: HalfGemmObservation,
identity: HalfKernelIdentity,
config: LaunchConfig,
) -> Result<PhysicalLaunchObservation, String> {
if !context.policy.batch_invariant
|| context.policy.bi_gemm_family != crate::mamba_ssm::gpu::context::BiGemmFamily::Triad
{
return Err(
"recording a half launch requires the live batch-invariant Triad policy".into(),
);
}
if identity.module_kind == ModuleKind::TriadSm80 && !context.policy.bi_tensor_cores {
return Err(
"recording a forced half Tensor Core launch requires the live Triad Tensor Core policy"
.into(),
);
}
let route = resolved_half_gemm_route_with_compiler(
context,
compiler,
|pointer, bytes| observer.argument_identity_digest(pointer, bytes),
observation,
identity,
config,
)?;
Ok(PhysicalLaunchObservation::gemm(
half_policy_dtype(observation.dtype)?,
None,
route,
))
}
fn resolve_half_gemm_observation<O: PhysicalLaunchObserver>(
observer: &O,
kernels: &GpuKernels,
observation: HalfGemmObservation,
config: LaunchConfig,
) -> Result<PhysicalLaunchObservation, String> {
let context = observer
.route_context()
.ok_or_else(|| "recording half launch requires a GEMM route context".to_string())?;
let identity = HalfKernelIdentity::resolve(observation.base, observation.dtype)?;
resolve_half_gemm_observation_with_context(
observer,
context,
half_kernel_compiler_identity(kernels, identity)?,
observation,
identity,
config,
)
}
pub(in crate::mamba_ssm::gpu) struct PreparedHalfGraphIdentity {
function: CudaFunction,
config: LaunchConfig,
node: ResolvedPhysicalKernelLaunch,
base: &'static str,
}
impl PreparedHalfGraphIdentity {
pub(in crate::mamba_ssm::gpu) fn into_parts(
self,
) -> (
CudaFunction,
LaunchConfig,
ResolvedPhysicalKernelLaunch,
&'static str,
) {
(self.function, self.config, self.node, self.base)
}
}
fn prepared_half_graph_base(
expected: ResolvedPhysicalKernelLaunch,
dtype: WeightDtype,
) -> Result<&'static str, String> {
let suffix = match dtype {
WeightDtype::Bf16 => "_bf16",
WeightDtype::F16 => "_f16",
WeightDtype::F32 | WeightDtype::Tf32 => {
return Err("native half graph identity does not accept f32".into());
}
};
expected
.symbol()
.strip_suffix(suffix)
.ok_or_else(|| "native half graph symbol has the wrong dtype suffix".to_string())
}
fn prepared_half_graph_observation(
expected: ResolvedPhysicalKernelLaunch,
request: HalfPhysicalTraceRequest,
base: &'static str,
) -> Result<(LaunchConfig, HalfGemmObservation), String> {
let launch = expected.launch();
let config = LaunchConfig {
grid_dim: launch.grid_dim,
block_dim: launch.block_dim,
shared_mem_bytes: launch.shared_mem_bytes,
};
let route = expected
.gemm_route()
.ok_or_else(|| "native half graph node has no GEMM route".to_string())?;
let observation = HalfGemmObservation {
base,
op: request.op,
dtype: request.dtype,
dims: request.dims,
strides: request.nn_strides.unwrap_or_else(|| {
let shape = F32TriadShape::contiguous(request.op, request.dims);
(shape.lda, shape.ldb, shape.ldc)
}),
tile: expected
.tile()
.ok_or_else(|| "native half graph node has no tile".to_string())?,
bk_stages: (route.bk, route.stages),
arguments: HalfGemmArguments {
output: request.output,
a: request.a,
b: request.b,
bias: if request.op == ResolvedGemmOp::Nn {
request.bias
} else {
0
},
},
};
Ok((config, observation))
}
fn resolve_prepared_half_graph_node_with_context<O: PhysicalLaunchObserver>(
observer: &O,
context: GemmRouteIdentity,
compiler: crate::mamba_ssm::gpu::kernel_identity::CompilerIdentity,
expected: ResolvedPhysicalKernelLaunch,
request: HalfPhysicalTraceRequest,
base: &'static str,
) -> Result<(LaunchConfig, ResolvedPhysicalKernelLaunch), String> {
let (config, observation) = prepared_half_graph_observation(expected, request, base)?;
let identity = HalfKernelIdentity::resolve(base, request.dtype)?;
if expected.module_kind() != identity.module_kind || expected.symbol() != identity.symbol {
return Err("half kernel identity does not match its exact symbol and module owner".into());
}
let resolved = resolve_half_gemm_observation_with_context(
observer,
context,
compiler,
observation,
identity,
config,
)?;
let node = resolve_physical_launch_observation(observer, resolved, config)?;
Ok((config, node))
}
fn resolve_prepared_half_graph_node<O: PhysicalLaunchObserver>(
observer: &O,
kernels: &GpuKernels,
expected: ResolvedPhysicalKernelLaunch,
request: HalfPhysicalTraceRequest,
base: &'static str,
) -> Result<(LaunchConfig, ResolvedPhysicalKernelLaunch), String> {
let context = observer
.route_context()
.ok_or_else(|| "recording half launch requires a GEMM route context".to_string())?;
let identity = HalfKernelIdentity::resolve(base, request.dtype)?;
resolve_prepared_half_graph_node_with_context(
observer,
context,
half_kernel_compiler_identity(kernels, identity)?,
expected,
request,
base,
)
}
pub(in crate::mamba_ssm::gpu) fn prepare_native_half_graph_identity<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &O,
expected: ResolvedPhysicalKernelLaunch,
request: HalfPhysicalTraceRequest,
) -> Result<PreparedHalfGraphIdentity, String> {
let base = prepared_half_graph_base(expected, request.dtype)?;
let choice = match base {
"nn_gemv" => {
HalfKernelChoice::new(base, ctx.kernels.gemm_bi_nn_gemv_typed.get(request.dtype))
}
"nn_ultra_thin" => HalfKernelChoice::new(
base,
ctx.kernels.gemm_bi_nn_ultra_thin_typed.get(request.dtype),
),
"nn_narrow_small" => HalfKernelChoice::new(
base,
ctx.kernels.gemm_bi_nn_narrow_small_typed.get(request.dtype),
),
"nn_narrow" => {
HalfKernelChoice::new(base, ctx.kernels.gemm_bi_nn_narrow_typed.get(request.dtype))
}
"nn_big" => {
HalfKernelChoice::new(base, ctx.kernels.gemm_bi_nn_big_typed.get(request.dtype))
}
"tn_gemv" => {
HalfKernelChoice::new(base, ctx.kernels.gemm_bi_tn_gemv_typed.get(request.dtype))
}
"tn_narrow" => {
HalfKernelChoice::new(base, ctx.kernels.gemm_bi_tn_narrow_typed.get(request.dtype))
}
"tn_big" => {
HalfKernelChoice::new(base, ctx.kernels.gemm_bi_tn_big_typed.get(request.dtype))
}
"nt_gemv" => {
HalfKernelChoice::new(base, ctx.kernels.gemm_bi_nt_gemv_typed.get(request.dtype))
}
"nt_narrow" => {
HalfKernelChoice::new(base, ctx.kernels.gemm_bi_nt_narrow_typed.get(request.dtype))
}
"nt_big" => {
HalfKernelChoice::new(base, ctx.kernels.gemm_bi_nt_big_typed.get(request.dtype))
}
"nn_tc" => HalfKernelChoice::new(base, ctx.kernels.gemm_bi_nn_tc_typed.get(request.dtype)),
"nn_tc64" => {
HalfKernelChoice::new(base, ctx.kernels.gemm_bi_nn_tc64_typed.get(request.dtype))
}
"nn_tc16" => {
HalfKernelChoice::new(base, ctx.kernels.gemm_bi_nn_tc16_typed.get(request.dtype))
}
"tn_tc" => HalfKernelChoice::new(base, ctx.kernels.gemm_bi_tn_tc_typed.get(request.dtype)),
"tn_tc64" => {
HalfKernelChoice::new(base, ctx.kernels.gemm_bi_tn_tc64_typed.get(request.dtype))
}
"tn_tc128x64" => HalfKernelChoice::new(
base,
ctx.kernels.gemm_bi_tn_tc128x64_typed.get(request.dtype),
),
"tn_tc64_streamk" => HalfKernelChoice::new(
base,
ctx.kernels
.gemm_bi_tn_tc64_streamk_typed
.as_ref()
.ok_or_else(|| {
"the portable stream-K dW kernel is not composed for this target".to_string()
})?
.get(request.dtype),
),
"nt_tc" => HalfKernelChoice::new(base, ctx.kernels.gemm_bi_nt_tc_typed.get(request.dtype)),
"nt_tc64" => {
HalfKernelChoice::new(base, ctx.kernels.gemm_bi_nt_tc64_typed.get(request.dtype))
}
"nn_sm89_m128n128_bk64_s3" => HalfKernelChoice::new(
base,
ctx.kernels
.triad_sm89_half_function(
super::sm89_half_source::Sm89HalfRoute::NnM128N128Bk64S3,
request.dtype,
)
.ok_or_else(|| "prepared SM89 half NN symbol is unavailable".to_string())?,
),
"tn_sm89_m64n64_bk64_s2_compact_bxor" => HalfKernelChoice::new(
base,
ctx.kernels
.triad_sm89_half_function(
super::sm89_half_source::Sm89HalfRoute::TnM64N64Bk64S2CompactBxor,
request.dtype,
)
.ok_or_else(|| "prepared SM89 half TN compact symbol is unavailable".to_string())?,
),
"tn_sm89_m64n64_bk64_s2_regpipe_vec2" => HalfKernelChoice::new(
base,
ctx.kernels
.triad_sm89_half_function(
super::sm89_half_source::Sm89HalfRoute::TnM64N64Bk64S2RegpipeVec2,
request.dtype,
)
.ok_or_else(|| {
"prepared SM89 half TN regpipe+vec2 symbol is unavailable".to_string()
})?,
),
"tn_sm89_m16n16_bk64_s2_ldb72"
| "tn_sm89_half_d128_in_m32n16_bk64_s4_cg"
| "tn_sm89_half_d128_out_m32n16_bk64_s4_cg"
| "nt_sm89_m16n64_bk64_s4"
| "nn_sm89_m16n64_bk64_s4"
| "tn_sm89_relay_m64n64_bk64_s3" => HalfKernelChoice::new(
base,
ctx.kernels
.triad_sm89_half_runtime_function(expected.symbol())
.ok_or_else(|| "prepared SM89 half runtime symbol is unavailable".to_string())?,
),
"nt_sm89_m128n128_bk64_s3_bxor" => HalfKernelChoice::new(
base,
ctx.kernels
.triad_sm89_half_function(
super::sm89_half_source::Sm89HalfRoute::NtM128N128Bk64S3Bxor,
request.dtype,
)
.ok_or_else(|| "prepared SM89 half NT Bxor symbol is unavailable".to_string())?,
),
"nt_sm89_m96n128_bk64_s3" => HalfKernelChoice::new(
base,
ctx.kernels
.triad_sm89_half_function(
super::sm89_half_source::Sm89HalfRoute::NtM96N128Bk64S3,
request.dtype,
)
.ok_or_else(|| "prepared SM89 half NT M96 symbol is unavailable".to_string())?,
),
_ => return Err("native half graph symbol is not a prepared typed route".into()),
};
let (config, node) =
resolve_prepared_half_graph_node(observer, &ctx.kernels, expected, request, choice.base)?;
Ok(PreparedHalfGraphIdentity {
function: choice.function.clone(),
config,
node,
base: choice.base,
})
}
#[inline(always)]
unsafe fn enqueue_half_gemm<O: PhysicalLaunchObserver>(
observer: &mut O,
bindings: (&GpuKernels, Option<&GpuCtx>),
builder: &mut LaunchArgs<'_>,
config: LaunchConfig,
observation: HalfGemmObservation,
driver_context: std::fmt::Arguments<'_>,
) -> Result<HalfNativeBranchSeal, String> {
let (kernels, context) = bindings;
if let Some(ctx) = context {
ctx.ensure_gemm_usable()?;
if ctx.gemm_route_recording_active()? {
let identity = HalfKernelIdentity::resolve(observation.base, observation.dtype)?;
let route = resolved_half_gemm_route(
ctx.gemm_route(),
kernels,
|pointer, bytes| {
if pointer == 0 && bytes != 0 {
return Err("half GEMM context route has a null nonempty span".into());
}
pointer
.checked_add(bytes)
.ok_or("half GEMM context span overflows u64")?;
Ok(FramedSha256::new(b"triad-half-context-pointer-span.v1")
.required(b"pointer", &pointer.to_le_bytes())
.required(b"bytes", &bytes.to_le_bytes())
.finish())
},
observation,
identity,
config,
)?;
ctx.validate_resolved_gemm_route(&route, "native half terminal")?;
ctx.record_resolved_gemm_route(route)?;
}
}
let physical_observation = if O::ENABLED {
Some(resolve_half_gemm_observation(
observer,
kernels,
observation,
config,
)?)
} else {
None
};
unsafe { enqueue_with_physical_observation(observer, builder, config, physical_observation) }
.map_err(|error| error.with_driver_context(driver_context))?;
Ok(half_native_branch_seal(observation, config))
}
pub const SM89_HALF_RELAY_BASE: &str = "tn_sm89_relay_m64n64_bk64_s3";
pub fn half_tn_graph_parameter_count(base: &str) -> usize {
match base {
"tn_gemv" => 8,
"tn_tc64_streamk" | SM89_HALF_RELAY_BASE => 9,
_ => 7,
}
}
fn sm89_half_base(route: super::sm89_half_source::Sm89HalfRuntimeRoute) -> &'static str {
match route {
super::sm89_half_source::Sm89HalfRuntimeRoute::Legacy(
super::sm89_half_source::Sm89HalfRoute::NnM128N128Bk64S3,
) => "nn_sm89_m128n128_bk64_s3",
super::sm89_half_source::Sm89HalfRuntimeRoute::Legacy(
super::sm89_half_source::Sm89HalfRoute::TnM64N64Bk64S2CompactBxor,
) => "tn_sm89_m64n64_bk64_s2_compact_bxor",
super::sm89_half_source::Sm89HalfRuntimeRoute::Legacy(
super::sm89_half_source::Sm89HalfRoute::TnM64N64Bk64S2RegpipeVec2,
) => "tn_sm89_m64n64_bk64_s2_regpipe_vec2",
super::sm89_half_source::Sm89HalfRuntimeRoute::TnSmall16Bk64S2Ldb72 => {
"tn_sm89_m16n16_bk64_s2_ldb72"
}
super::sm89_half_source::Sm89HalfRuntimeRoute::TnD128InM32N16Bk64S4 => {
"tn_sm89_half_d128_in_m32n16_bk64_s4_cg"
}
super::sm89_half_source::Sm89HalfRuntimeRoute::TnD128OutM32N16Bk64S4 => {
"tn_sm89_half_d128_out_m32n16_bk64_s4_cg"
}
super::sm89_half_source::Sm89HalfRuntimeRoute::NtSmallM16N64Bk64S4 => {
"nt_sm89_m16n64_bk64_s4"
}
super::sm89_half_source::Sm89HalfRuntimeRoute::NnSmallM16N64Bk64S4 => {
"nn_sm89_m16n64_bk64_s4"
}
super::sm89_half_source::Sm89HalfRuntimeRoute::TnRelayM64N64Bk64S3 => SM89_HALF_RELAY_BASE,
super::sm89_half_source::Sm89HalfRuntimeRoute::Legacy(
super::sm89_half_source::Sm89HalfRoute::NtM128N128Bk64S3Bxor,
) => "nt_sm89_m128n128_bk64_s3_bxor",
super::sm89_half_source::Sm89HalfRuntimeRoute::Legacy(
super::sm89_half_source::Sm89HalfRoute::NtM96N128Bk64S3,
) => "nt_sm89_m96n128_bk64_s3",
}
}
fn sm89_half_launch_config(
spec: super::sm89_half_source::Sm89HalfRuntimeSpec,
dims: (usize, usize, usize),
) -> Result<LaunchConfig, String> {
let (output_rows, output_columns) = match spec.op {
ResolvedGemmOp::Nn => (dims.0, dims.2),
ResolvedGemmOp::Tn => (dims.1, dims.2),
ResolvedGemmOp::Nt => (dims.0, dims.1),
};
Ok(LaunchConfig {
grid_dim: (
checked_tile_grid(
checked_u32(output_rows, "SM89 half output rows")?,
spec.tile.0,
checked_u32(output_columns, "SM89 half output columns")?,
spec.tile.1,
)?,
1,
1,
),
block_dim: (spec.threads, 1, 1),
shared_mem_bytes: spec.dynamic_shared_bytes,
})
}
#[derive(Clone, Copy, Debug, PartialEq)]
enum Sm89HalfTnArgument {
Pointer(CUptr),
ScalarF32(f32),
ScalarI32(i32),
}
#[derive(Clone, Copy, Debug, PartialEq)]
struct Sm89HalfRelayWorkspace {
partial: CUptr,
flags: CUptr,
}
#[derive(Clone, Copy, Debug, PartialEq)]
struct Sm89HalfRelayLaunch {
grid: u32,
workspace: Sm89HalfRelayWorkspace,
}
pub fn sm89_half_relay_grid(
kernels: &GpuKernels,
dims: (usize, usize, usize),
) -> Result<u32, String> {
let checked = GemmDims::tn(dims)?;
let (batch, n_in, n_out) = checked.tuple();
let tiles = checked_tile_grid(
checked_u32(n_in, "half relay tile rows")?,
64,
checked_u32(n_out, "half relay tile columns")?,
64,
)?;
let slabs = checked_u32(batch, "half relay reduction rows")?.div_ceil(64);
let units = u64::from(tiles) * u64::from(slabs);
let resident = u64::from(kernels.multiprocessor_count().max(1))
* u64::from(kernels.sm89_half_relay_resident_ctas().max(1));
u32::try_from(units.min(resident).max(1))
.map_err(|_| "the half relay grid exceeds u32".to_string())
}
pub fn sm89_half_relay_workspace(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
grid: u32,
) -> Result<(CUptr, CUptr), String> {
use cudarc::driver::DevicePtr;
let slots = grid as usize;
let partial_floats = slots
.checked_mul(super::sm89_half_relay_source::SLAB_FLOATS)
.ok_or_else(|| "the half relay slab extent overflows usize".to_string())?;
if partial_floats > SPLITK_SCRATCH_CAP {
return Err("the half relay slabs exceed the fixed workspace".into());
}
if slots > TF32_SPLITK_COUNTER_CAP {
return Err("the half relay flags exceed the fixed counter workspace".into());
}
let (partial, _) = kernels.splitk_scratch_buf(stream)?.device_ptr(stream);
let (flags, _) = kernels
.triad_kernels()
.tf32_splitk_counter_buf(stream)?
.device_ptr(stream);
Ok((partial, flags))
}
#[derive(Clone, Copy, Debug, PartialEq)]
struct Sm89HalfTnArguments {
core: [Sm89HalfTnArgument; 7],
relay: Option<Sm89HalfRelayWorkspace>,
}
impl Sm89HalfTnArguments {
fn bind<'a>(&'a self, builder: &mut LaunchArgs<'a>) {
for argument in &self.core {
match argument {
Sm89HalfTnArgument::Pointer(value) => builder.arg(value),
Sm89HalfTnArgument::ScalarF32(value) => builder.arg(value),
Sm89HalfTnArgument::ScalarI32(value) => builder.arg(value),
};
}
if let Some(workspace) = &self.relay {
builder.arg(&workspace.partial);
builder.arg(&workspace.flags);
}
}
}
#[derive(Clone, Copy)]
struct Sm89HalfTnLaunchPlan {
arguments: Sm89HalfTnArguments,
config: LaunchConfig,
observation: HalfGemmObservation,
}
fn sm89_half_tn_launch_plan(
spec: super::sm89_half_source::Sm89HalfRuntimeSpec,
dw_ptr: CUptr,
dy: TypedPtr,
x_saved: TypedPtr,
dims: (usize, usize, usize),
relay: Option<Sm89HalfRelayLaunch>,
) -> Result<Sm89HalfTnLaunchPlan, String> {
if spec.op != ResolvedGemmOp::Tn || spec.dtype != dy.dtype || x_saved.dtype != dy.dtype {
return Err("SM89 half TN launch plan does not match the selected dtype and op".into());
}
if relay.is_some() != (spec.schedule == super::sm89_half_source::Sm89HalfSchedule::Relay) {
return Err(
"SM89 half TN launch plan pairs the relay schedule with its persistent grid".into(),
);
}
let shape = F32TriadShape::contiguous(ResolvedGemmOp::Tn, dims);
let checked = GemmDims::tn(dims)?;
let base = sm89_half_base(spec.route);
Ok(Sm89HalfTnLaunchPlan {
arguments: Sm89HalfTnArguments {
core: [
Sm89HalfTnArgument::Pointer(dw_ptr),
Sm89HalfTnArgument::Pointer(x_saved.ptr),
Sm89HalfTnArgument::Pointer(dy.ptr),
Sm89HalfTnArgument::ScalarF32(1.0),
Sm89HalfTnArgument::ScalarI32(checked.m_i32),
Sm89HalfTnArgument::ScalarI32(checked.k_i32),
Sm89HalfTnArgument::ScalarI32(checked.n_i32),
],
relay: relay.map(|launch| launch.workspace),
},
config: match relay {
Some(launch) => LaunchConfig {
grid_dim: (launch.grid, 1, 1),
block_dim: (spec.threads, 1, 1),
shared_mem_bytes: spec.dynamic_shared_bytes,
},
None => sm89_half_launch_config(spec, dims)?,
},
observation: HalfGemmObservation {
base,
op: ResolvedGemmOp::Tn,
dtype: dy.dtype,
dims,
strides: (shape.lda, shape.ldb, shape.ldc),
tile: spec.tile,
bk_stages: (spec.bk, spec.stages),
arguments: HalfGemmArguments {
output: dw_ptr,
a: x_saved.ptr,
b: dy.ptr,
bias: 0,
},
},
})
}
fn sm89_half_auto_context(ctx: &GpuCtx) -> super::sm89_half_source::Sm89HalfAutoContext {
super::sm89_half_source::Sm89HalfAutoContext {
compiler: ctx.kernels.triad_sm89_half_compiler_identity(),
artifact: ctx.kernels.triad_sm89_half_artifact_identity(),
compute_capability: ctx.compute_capability(),
multiprocessor_count: ctx.kernels.multiprocessor_count(),
}
}
pub(in crate::mamba_ssm::gpu) fn launch_sm89_half_nn_auto_observed<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
ops: &TcFwdOperands,
dims: (usize, usize, usize),
) -> Result<Option<HalfNativeBranchSeal>, String> {
require_half(ops.y.dtype, "output")?;
if ops.x.dtype != ops.y.dtype || ops.w.dtype != ops.y.dtype {
return Err("SM89 half NN AUTO: mixed dtypes not supported".into());
}
let shape = F32TriadShape::contiguous(ResolvedGemmOp::Nn, dims);
let operands = F32TriadOperands {
output: ops.y.ptr,
a: ops.x.ptr,
b: ops.w.ptr,
bias: (ops.bias_ptr != 0).then_some(ops.bias_ptr),
alpha: 1.0,
beta: 0.0,
};
let Some((spec, admission)) =
super::sm89_half_source::select_sm89_half_auto_cell_with_admission(
sm89_half_auto_context(ctx),
super::sm89_half_source::Sm89HalfAutoRequest {
request: F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape,
},
operands,
dtype: ops.y.dtype,
half_policy: ctx.half_triad_policy(),
},
)
else {
return Ok(None);
};
if ctx
.kernels
.triad_sm89_half_runtime_function(spec.symbol)
.is_none()
{
return Ok(None);
}
if admission == super::sm89_half_source::Sm89HalfAdmission::Proof {
let key = super::proof::RouteProofKey {
candidate: spec.symbol,
op: ResolvedGemmOp::Nn,
dtype: ops.y.dtype,
dims,
};
let elements = dims
.0
.checked_mul(dims.2)
.ok_or_else(|| "SM89 half NN output span overflows usize".to_string())?;
if !half_wave_guard_admits(ctx, spec, key, (dims.0, dims.2))? {
return Ok(None);
}
let admitted = proven_candidate(
ctx,
key,
ops.y.ptr,
elements,
ops.y.dtype,
|scratch| {
let scratch_ops = TcFwdOperands {
y: TypedPtr {
ptr: scratch,
dtype: ops.y.dtype,
},
..*ops
};
enqueue_sm89_half_nn(ctx, &mut NoPhysicalObserver, spec, &scratch_ops, dims)
.map(drop)
},
|scratch| {
let scratch_ops = TcFwdOperands {
y: TypedPtr {
ptr: scratch,
dtype: ops.y.dtype,
},
..*ops
};
gemm_bi_forward_tc_with_tile(
&ctx.stream,
&ctx.kernels,
&scratch_ops,
dims,
TcTile::Tile64,
)
},
)?;
if !admitted {
return Ok(None);
}
}
enqueue_sm89_half_nn(ctx, observer, spec, ops, dims).map(Some)
}
fn enqueue_sm89_half_nn<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
spec: super::sm89_half_source::Sm89HalfRuntimeSpec,
ops: &TcFwdOperands,
dims: (usize, usize, usize),
) -> Result<HalfNativeBranchSeal, String> {
let shape = F32TriadShape::contiguous(ResolvedGemmOp::Nn, dims);
let Some(function) = ctx.kernels.triad_sm89_half_runtime_function(spec.symbol) else {
return Err("prepared SM89 half NN symbol is unavailable".into());
};
let cfg = sm89_half_launch_config(spec, dims)?;
let checked = GemmDims::nn(dims, shape.lda)?;
let params = Sm89HalfNnParams {
alpha: 1.0,
beta: 0.0,
m: checked.m_i32,
n: checked.n_i32,
k: checked.k_i32,
lda: checked_i32(shape.lda, "SM89 half NN lda")?,
ldb: checked_i32(shape.ldb, "SM89 half NN ldb")?,
ldc: checked_i32(shape.ldc, "SM89 half NN ldc")?,
};
let base = sm89_half_base(spec.route);
let mut builder = ctx.stream.launch_builder(function);
builder.arg(&ops.y.ptr);
builder.arg(&ops.x.ptr);
builder.arg(&ops.w.ptr);
builder.arg(&ops.bias_ptr);
builder.arg(¶ms);
let seal = unsafe {
enqueue_half_gemm(
observer,
(&ctx.kernels, Some(ctx)),
&mut builder,
cfg,
HalfGemmObservation {
base,
op: ResolvedGemmOp::Nn,
dtype: ops.y.dtype,
dims,
strides: (shape.lda, shape.ldb, shape.ldc),
tile: spec.tile,
bk_stages: (spec.bk, spec.stages),
arguments: HalfGemmArguments {
output: ops.y.ptr,
a: ops.x.ptr,
b: ops.w.ptr,
bias: ops.bias_ptr,
},
},
format_args!("{base}"),
)
}?;
Ok(seal)
}
pub(in crate::mamba_ssm::gpu) fn launch_sm89_half_tn_auto_observed<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
dw_ptr: CUptr,
dy: TypedPtr,
x_saved: TypedPtr,
dims: (usize, usize, usize),
) -> Result<Option<HalfNativeBranchSeal>, String> {
require_half(dy.dtype, "dY")?;
if x_saved.dtype != dy.dtype {
return Err("SM89 half TN AUTO: mixed dtypes not supported".into());
}
let shape = F32TriadShape::contiguous(ResolvedGemmOp::Tn, dims);
let operands = F32TriadOperands {
output: dw_ptr,
a: x_saved.ptr,
b: dy.ptr,
bias: None,
alpha: 1.0,
beta: 1.0,
};
let Some((spec, admission)) =
super::sm89_half_source::select_sm89_half_auto_cell_with_admission(
sm89_half_auto_context(ctx),
super::sm89_half_source::Sm89HalfAutoRequest {
request: F32TriadRequest {
op: ResolvedGemmOp::Tn,
shape,
},
operands,
dtype: dy.dtype,
half_policy: ctx.half_triad_policy(),
},
)
else {
return Ok(None);
};
if ctx
.kernels
.triad_sm89_half_runtime_function(spec.symbol)
.is_none()
{
return Ok(None);
}
if admission == super::sm89_half_source::Sm89HalfAdmission::Proof {
let key = super::proof::RouteProofKey {
candidate: spec.symbol,
op: ResolvedGemmOp::Tn,
dtype: dy.dtype,
dims,
};
let elements = dims
.1
.checked_mul(dims.2)
.ok_or_else(|| "SM89 half TN output span overflows usize".to_string())?;
if !half_wave_guard_admits(ctx, spec, key, (dims.1, dims.2))? {
return Ok(None);
}
let admitted = proven_candidate(
ctx,
key,
dw_ptr,
elements,
WeightDtype::F32,
|scratch| {
enqueue_sm89_half_tn(
ctx,
&mut NoPhysicalObserver,
spec,
scratch,
dy,
x_saved,
dims,
)
.map(drop)
},
|scratch| {
gemm_bi_backward_dw_tc_with_tile(
&ctx.stream,
&ctx.kernels,
scratch,
dy,
x_saved,
dims,
TcTile::Tile64,
)
},
)?;
if !admitted {
return Ok(None);
}
}
enqueue_sm89_half_tn(ctx, observer, spec, dw_ptr, dy, x_saved, dims).map(Some)
}
fn enqueue_sm89_half_tn<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
spec: super::sm89_half_source::Sm89HalfRuntimeSpec,
dw_ptr: CUptr,
dy: TypedPtr,
x_saved: TypedPtr,
dims: (usize, usize, usize),
) -> Result<HalfNativeBranchSeal, String> {
let Some(function) = ctx.kernels.triad_sm89_half_runtime_function(spec.symbol) else {
return Err("prepared SM89 half TN symbol is unavailable".into());
};
let relay = if spec.schedule == super::sm89_half_source::Sm89HalfSchedule::Relay {
let grid = sm89_half_relay_grid(&ctx.kernels, dims)?;
let (partial, flags) = sm89_half_relay_workspace(&ctx.stream, &ctx.kernels, grid)?;
Some(Sm89HalfRelayLaunch {
grid,
workspace: Sm89HalfRelayWorkspace { partial, flags },
})
} else {
None
};
let plan = sm89_half_tn_launch_plan(spec, dw_ptr, dy, x_saved, dims, relay)?;
let mut builder = ctx.stream.launch_builder(function);
plan.arguments.bind(&mut builder);
let seal = unsafe {
enqueue_half_gemm(
observer,
(&ctx.kernels, Some(ctx)),
&mut builder,
plan.config,
plan.observation,
format_args!("{}", plan.observation.base),
)
}?;
Ok(seal)
}
pub(in crate::mamba_ssm::gpu) fn launch_sm89_half_nt_auto_observed<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
dx: TypedPtr,
dy: TypedPtr,
w: TypedPtr,
dims: (usize, usize, usize),
) -> Result<Option<HalfNativeBranchSeal>, String> {
require_half(dx.dtype, "dX")?;
if dx.dtype != dy.dtype || dy.dtype != w.dtype {
return Err("SM89 half NT AUTO: mixed dtypes not supported".into());
}
let shape = F32TriadShape::contiguous(ResolvedGemmOp::Nt, dims);
let operands = F32TriadOperands {
output: dx.ptr,
a: dy.ptr,
b: w.ptr,
bias: None,
alpha: 1.0,
beta: 0.0,
};
let Some((spec, admission)) =
super::sm89_half_source::select_sm89_half_auto_cell_with_admission(
sm89_half_auto_context(ctx),
super::sm89_half_source::Sm89HalfAutoRequest {
request: F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape,
},
operands,
dtype: dx.dtype,
half_policy: ctx.half_triad_policy(),
},
)
else {
return Ok(None);
};
if ctx
.kernels
.triad_sm89_half_runtime_function(spec.symbol)
.is_none()
{
return Ok(None);
}
if admission == super::sm89_half_source::Sm89HalfAdmission::Proof {
let key = super::proof::RouteProofKey {
candidate: spec.symbol,
op: ResolvedGemmOp::Nt,
dtype: dx.dtype,
dims,
};
let elements = dims
.0
.checked_mul(dims.1)
.ok_or_else(|| "SM89 half NT output span overflows usize".to_string())?;
if !half_wave_guard_admits(ctx, spec, key, (dims.0, dims.1))? {
return Ok(None);
}
let admitted = proven_candidate(
ctx,
key,
dx.ptr,
elements,
dx.dtype,
|scratch| {
let scratch_dx = TypedPtr {
ptr: scratch,
dtype: dx.dtype,
};
enqueue_sm89_half_nt(ctx, &mut NoPhysicalObserver, spec, scratch_dx, dy, w, dims)
.map(drop)
},
|scratch| {
let scratch_dx = TypedPtr {
ptr: scratch,
dtype: dx.dtype,
};
gemm_bi_backward_dx_tc_with_tile(
&ctx.stream,
&ctx.kernels,
scratch_dx,
dy,
w,
dims,
TcTile::Tile64,
)
},
)?;
if !admitted {
return Ok(None);
}
}
enqueue_sm89_half_nt(ctx, observer, spec, dx, dy, w, dims).map(Some)
}
fn enqueue_sm89_half_nt<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
spec: super::sm89_half_source::Sm89HalfRuntimeSpec,
dx: TypedPtr,
dy: TypedPtr,
w: TypedPtr,
dims: (usize, usize, usize),
) -> Result<HalfNativeBranchSeal, String> {
let shape = F32TriadShape::contiguous(ResolvedGemmOp::Nt, dims);
let Some(function) = ctx.kernels.triad_sm89_half_runtime_function(spec.symbol) else {
return Err("prepared SM89 half NT symbol is unavailable".into());
};
let cfg = sm89_half_launch_config(spec, dims)?;
let checked = GemmDims::nt(dims)?;
let alpha = 1.0_f32;
let base = sm89_half_base(spec.route);
let mut builder = ctx.stream.launch_builder(function);
builder.arg(&dx.ptr);
builder.arg(&dy.ptr);
builder.arg(&w.ptr);
builder.arg(&alpha);
builder.arg(&checked.m_i32);
builder.arg(&checked.n_i32);
builder.arg(&checked.k_i32);
let seal = unsafe {
enqueue_half_gemm(
observer,
(&ctx.kernels, Some(ctx)),
&mut builder,
cfg,
HalfGemmObservation {
base,
op: ResolvedGemmOp::Nt,
dtype: dx.dtype,
dims,
strides: (shape.lda, shape.ldb, shape.ldc),
tile: spec.tile,
bk_stages: (spec.bk, spec.stages),
arguments: HalfGemmArguments {
output: dx.ptr,
a: dy.ptr,
b: w.ptr,
bias: 0,
},
},
format_args!("{base}"),
)
}?;
Ok(seal)
}
impl TcTile {
fn extents(self) -> (u32, u32) {
match self {
TcTile::Tile128 => (128, 128),
TcTile::Tile64 => (64, 64),
TcTile::Thin16 => (16, 32),
TcTile::Rect128x64 => (128, 64),
TcTile::Tile64StreamK => (64, 64),
}
}
fn block_dim(self) -> u32 {
match self {
TcTile::Tile128 => 256,
TcTile::Tile64 | TcTile::Thin16 | TcTile::Tile64StreamK => 128,
TcTile::Rect128x64 => 256,
}
}
fn bk_stages(self) -> (u32, u8) {
match self {
TcTile::Tile128 | TcTile::Tile64 | TcTile::Tile64StreamK => (64, 2),
TcTile::Thin16 => (64, 4),
TcTile::Rect128x64 => (32, 3),
}
}
fn launch_cfg(
self,
rows: usize,
cols: usize,
dyn_bytes128: u32,
) -> Result<cudarc::driver::LaunchConfig, String> {
if self == TcTile::Tile64StreamK {
return Err(
"Tile64StreamK launches a persistent grid; use the stream-K dW path".into(),
);
}
let (bm, bn) = self.extents();
let total_tiles = checked_tile_grid(
checked_u32(rows, "tile rows")?,
bm,
checked_u32(cols, "tile columns")?,
bn,
)?;
Ok(cudarc::driver::LaunchConfig {
grid_dim: (total_tiles, 1, 1),
block_dim: (self.block_dim(), 1, 1),
shared_mem_bytes: match self {
TcTile::Tile128 => dyn_bytes128,
TcTile::Tile64 | TcTile::Thin16 | TcTile::Tile64StreamK => 0,
TcTile::Rect128x64 => 0,
},
})
}
}
const SM80_STREAMK_SLAB_FLOATS: usize = 128 * 32;
const SM80_STREAMK_SLOTS_PER_CTA: usize = 2;
pub fn sm80_streamk_grid(kernels: &GpuKernels, dims: (usize, usize, usize)) -> Result<u32, String> {
let checked = GemmDims::tn(dims)?;
let (batch, n_in, n_out) = checked.tuple();
let tiles = checked_tile_grid(
checked_u32(n_in, "tile rows")?,
64,
checked_u32(n_out, "tile columns")?,
64,
)?;
let slabs = checked_u32(batch, "reduction rows")?.div_ceil(64);
let units = u64::from(tiles) * u64::from(slabs);
let resident = u64::from(kernels.multiprocessor_count().max(1))
* u64::from(kernels.tc64_streamk_resident_ctas().max(1));
u32::try_from(units.min(resident).max(1)).map_err(|_| "stream-K grid exceeds u32".to_string())
}
pub fn sm80_streamk_workspace(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
grid: u32,
) -> Result<(CUptr, CUptr), String> {
use cudarc::driver::DevicePtr;
let slots = (grid as usize)
.checked_mul(SM80_STREAMK_SLOTS_PER_CTA)
.ok_or_else(|| "sm80 stream-K slot count overflows usize".to_string())?;
let partial_floats = slots
.checked_mul(SM80_STREAMK_SLAB_FLOATS)
.ok_or_else(|| "sm80 stream-K slab extent overflows usize".to_string())?;
if partial_floats > SPLITK_SCRATCH_CAP {
return Err("sm80 stream-K slabs exceed the fixed workspace".into());
}
if slots > TF32_SPLITK_COUNTER_CAP {
return Err("sm80 stream-K flags exceed the fixed counter workspace".into());
}
let (partial, _) = kernels.splitk_scratch_buf(stream)?.device_ptr(stream);
let (flags, _) = kernels
.triad_kernels()
.tf32_splitk_counter_buf(stream)?
.device_ptr(stream);
Ok((partial, flags))
}
pub fn gemm_bi_forward_tc(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
y: TypedPtr,
x: TypedPtr,
w: TypedPtr,
bias_ptr: CUptr,
dims: (usize, usize, usize),
) -> Result<TcTile, String> {
let mut environment = HalfLaunchEnvironment::production(stream, kernels);
gemm_bi_forward_tc_in(&mut environment, y, x, w, bias_ptr, dims).map(|(tile, _)| tile)
}
fn gemm_bi_forward_tc_in<O: PhysicalLaunchObserver>(
environment: &mut HalfLaunchEnvironment<'_, O>,
y: TypedPtr,
x: TypedPtr,
w: TypedPtr,
bias_ptr: CUptr,
dims: (usize, usize, usize),
) -> Result<(TcTile, HalfNativeBranchSeal), String> {
GemmDims::nn(dims, dims.1)?;
let tile = tc_pick_tile_forward(dims, environment.kernels.multiprocessor_count()).ok_or_else(|| {
let (batch, n_in, n_out) = dims;
format!(
"UNCOVERED gemm_bi_forward_tc: shape M={batch} K={n_in} N={n_out} below the TC tile gate"
)
})?;
let ops = TcFwdOperands { y, x, w, bias_ptr };
let seal = gemm_bi_forward_tc_with_tile_in(
environment,
&ops,
F32TriadShape::contiguous(ResolvedGemmOp::Nn, dims),
tile,
)?;
Ok((tile, seal))
}
pub fn gemm_bi_forward_tc_with_tile(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
ops: &TcFwdOperands,
dims: (usize, usize, usize),
tile: TcTile,
) -> Result<(), String> {
let mut environment = HalfLaunchEnvironment::production(stream, kernels);
gemm_bi_forward_tc_with_tile_in(
&mut environment,
ops,
F32TriadShape::contiguous(ResolvedGemmOp::Nn, dims),
tile,
)
.map(drop)
}
pub(in crate::mamba_ssm::gpu) fn gemm_bi_forward_tc_with_tile_shape(
ctx: &GpuCtx,
ops: &TcFwdOperands,
shape: F32TriadShape,
tile: TcTile,
) -> Result<(), String> {
let mut environment = HalfLaunchEnvironment::observed(ctx, NoPhysicalObserver);
gemm_bi_forward_tc_with_tile_in(&mut environment, ops, shape, tile).map(drop)
}
fn gemm_bi_forward_tc_with_tile_in<O: PhysicalLaunchObserver>(
environment: &mut HalfLaunchEnvironment<'_, O>,
ops: &TcFwdOperands,
shape: F32TriadShape,
tile: TcTile,
) -> Result<HalfNativeBranchSeal, String> {
shape.validate(ResolvedGemmOp::Nn)?;
let dims = (shape.m, shape.k, shape.n);
let checked_dims = GemmDims::nn(dims, shape.lda)?;
let (batch, _n_in, n_out) = checked_dims.tuple();
require_half(ops.y.dtype, "output")?;
if ops.x.dtype != ops.y.dtype || ops.w.dtype != ops.y.dtype {
return Err("gemm_bi_forward_tc: mixed dtypes not supported".into());
}
let dt = ops.y.dtype;
let alpha: f32 = 1.0;
let beta: f32 = 0.0;
validate_bias_preseed(alpha, ops.bias_ptr, "gemm_bi_forward_tc")?;
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_i = checked_dims.k_i32;
let cfg = tile.launch_cfg(batch, n_out, 71_680)?;
let choice = match tile {
TcTile::Tile128 => {
HalfKernelChoice::new("nn_tc", environment.kernels.gemm_bi_nn_tc_typed.get(dt))
}
TcTile::Tile64 => {
HalfKernelChoice::new("nn_tc64", environment.kernels.gemm_bi_nn_tc64_typed.get(dt))
}
TcTile::Thin16 => {
HalfKernelChoice::new("nn_tc16", environment.kernels.gemm_bi_nn_tc16_typed.get(dt))
}
TcTile::Rect128x64 => {
return Err("Rect128x64 is a forced TN dW tile; NN has no rectangular route".into());
}
TcTile::Tile64StreamK => {
return Err("Tile64StreamK is a TN dW schedule; NN has no stream-K route".into());
}
};
let mut b = environment.stream.launch_builder(choice.function);
b.arg(&ops.y.ptr);
b.arg(&ops.x.ptr);
b.arg(&ops.w.ptr);
b.arg(&ops.bias_ptr);
b.arg(&alpha);
b.arg(&beta);
b.arg(&m_i);
b.arg(&n_i);
b.arg(&k_i);
let lda_i = checked_i32(shape.lda, "lda")?;
let ldb_i = checked_i32(shape.ldb, "ldb")?;
let ldc_i = checked_i32(shape.ldc, "ldc")?;
b.arg(&lda_i);
b.arg(&ldb_i);
b.arg(&ldc_i);
unsafe {
enqueue_half_gemm(
&mut environment.observer,
(environment.kernels, environment.context),
&mut b,
cfg,
HalfGemmObservation {
base: choice.base,
op: ResolvedGemmOp::Nn,
dtype: dt,
dims,
strides: (shape.lda, shape.ldb, shape.ldc),
tile: tile.extents(),
bk_stages: tile.bk_stages(),
arguments: HalfGemmArguments {
output: ops.y.ptr,
a: ops.x.ptr,
b: ops.w.ptr,
bias: ops.bias_ptr,
},
},
format_args!("nn_tc ({tile:?})"),
)
}
}
pub fn gemm_bi_backward_dw_tc(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
dw_ptr: CUptr,
dy: TypedPtr,
x_saved: TypedPtr,
dims: (usize, usize, usize),
) -> Result<TcTile, String> {
let mut environment = HalfLaunchEnvironment::production(stream, kernels);
gemm_bi_backward_dw_tc_in(
&mut environment,
dw_ptr,
dy,
x_saved,
dims,
HalfTriadPolicy::TiledParity,
)
.map(|(tile, _)| tile)
}
fn gemm_bi_backward_dw_tc_in<O: PhysicalLaunchObserver>(
environment: &mut HalfLaunchEnvironment<'_, O>,
dw_ptr: CUptr,
dy: TypedPtr,
x_saved: TypedPtr,
dims: (usize, usize, usize),
half_policy: HalfTriadPolicy,
) -> Result<(TcTile, HalfNativeBranchSeal), String> {
let checked_dims = GemmDims::tn(dims)?;
let (batch, n_in, n_out) = checked_dims.tuple();
let (major, minor) = environment.kernels.triad_scalar_compute_capability();
let compute_capability = (
i32::try_from(major).map_err(|error| format!("compute capability major: {error}"))?,
i32::try_from(minor).map_err(|error| format!("compute capability minor: {error}"))?,
);
let tile = tc_pick_tile_backward_for_device(
super::super::kernel_identity::PolicyOp::Dw,
dims,
environment.kernels.multiprocessor_count(),
compute_capability,
half_policy,
).ok_or_else(|| {
format!(
"UNCOVERED gemm_bi_backward_dw_tc: shape M={batch} K={n_in} N={n_out} outside the automatic TC route"
)
})?;
let seal = gemm_bi_backward_dw_tc_with_tile_in(environment, dw_ptr, dy, x_saved, dims, tile)?;
Ok((tile, seal))
}
pub fn gemm_bi_backward_dw_tc_with_tile(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
dw_ptr: CUptr,
dy: TypedPtr,
x_saved: TypedPtr,
dims: (usize, usize, usize),
tile: TcTile,
) -> Result<(), String> {
let mut environment = HalfLaunchEnvironment::production(stream, kernels);
gemm_bi_backward_dw_tc_with_tile_in(&mut environment, dw_ptr, dy, x_saved, dims, tile).map(drop)
}
fn gemm_bi_backward_dw_tc_with_tile_in<O: PhysicalLaunchObserver>(
environment: &mut HalfLaunchEnvironment<'_, O>,
dw_ptr: CUptr,
dy: TypedPtr,
x_saved: TypedPtr,
dims: (usize, usize, usize),
tile: TcTile,
) -> Result<HalfNativeBranchSeal, String> {
let checked_dims = GemmDims::tn(dims)?;
let (_batch, n_in, n_out) = checked_dims.tuple();
require_half(dy.dtype, "dY")?;
if dy.dtype != x_saved.dtype {
return Err("gemm_bi_backward_dw_tc: mixed dtypes not supported".into());
}
let dt = dy.dtype;
let alpha: f32 = 1.0;
let m_red_i = checked_dims.m_i32;
let k_out_i = checked_dims.k_i32;
let n_i = checked_dims.n_i32;
let (cfg, choice, workspace) = match tile {
TcTile::Tile128 => (
tile.launch_cfg(n_in, n_out, 69_632)?,
HalfKernelChoice::new("tn_tc", environment.kernels.gemm_bi_tn_tc_typed.get(dt)),
None,
),
TcTile::Tile64 => (
tile.launch_cfg(n_in, n_out, 69_632)?,
HalfKernelChoice::new("tn_tc64", environment.kernels.gemm_bi_tn_tc64_typed.get(dt)),
None,
),
TcTile::Rect128x64 => (
tile.launch_cfg(n_in, n_out, 69_632)?,
HalfKernelChoice::new(
"tn_tc128x64",
environment.kernels.gemm_bi_tn_tc128x64_typed.get(dt),
),
None,
),
TcTile::Tile64StreamK => {
let grid = sm80_streamk_grid(environment.kernels, dims)?;
let workspace = sm80_streamk_workspace(environment.stream, environment.kernels, grid)?;
(
cudarc::driver::LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (tile.block_dim(), 1, 1),
shared_mem_bytes: 0,
},
HalfKernelChoice::new(
"tn_tc64_streamk",
environment
.kernels
.gemm_bi_tn_tc64_streamk_typed
.as_ref()
.ok_or_else(|| {
"the portable stream-K dW kernel is not composed for this target"
.to_string()
})?
.get(dt),
),
Some(workspace),
)
}
TcTile::Thin16 => {
return Err("Thin16 is an NN-forward rung; the TN dW path has no thin tile".into());
}
};
let mut b = environment.stream.launch_builder(choice.function);
b.arg(&dw_ptr);
b.arg(&x_saved.ptr);
b.arg(&dy.ptr);
b.arg(&alpha);
b.arg(&m_red_i);
b.arg(&k_out_i);
b.arg(&n_i);
if let Some((partial, flags)) = &workspace {
b.arg(partial);
b.arg(flags);
}
unsafe {
enqueue_half_gemm(
&mut environment.observer,
(environment.kernels, environment.context),
&mut b,
cfg,
HalfGemmObservation {
base: choice.base,
op: ResolvedGemmOp::Tn,
dtype: dt,
dims,
strides: (dims.1, dims.2, dims.2),
tile: tile.extents(),
bk_stages: tile.bk_stages(),
arguments: HalfGemmArguments {
output: dw_ptr,
a: x_saved.ptr,
b: dy.ptr,
bias: 0,
},
},
format_args!("tn_tc ({tile:?})"),
)
}
}
pub fn gemm_bi_backward_dx_tc(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
dx: TypedPtr,
dy: TypedPtr,
w: TypedPtr,
dims: (usize, usize, usize),
) -> Result<TcTile, String> {
let mut environment = HalfLaunchEnvironment::production(stream, kernels);
gemm_bi_backward_dx_tc_in(&mut environment, dx, dy, w, dims).map(|(tile, _)| tile)
}
fn gemm_bi_backward_dx_tc_in<O: PhysicalLaunchObserver>(
environment: &mut HalfLaunchEnvironment<'_, O>,
dx: TypedPtr,
dy: TypedPtr,
w: TypedPtr,
dims: (usize, usize, usize),
) -> Result<(TcTile, HalfNativeBranchSeal), String> {
let checked_dims = GemmDims::nt(dims)?;
let (batch, n_in, n_out) = checked_dims.tuple();
let (major, minor) = environment.kernels.triad_scalar_compute_capability();
let compute_capability = (
i32::try_from(major).map_err(|error| format!("compute capability major: {error}"))?,
i32::try_from(minor).map_err(|error| format!("compute capability minor: {error}"))?,
);
let tile = tc_pick_tile_backward_for_device(
super::super::kernel_identity::PolicyOp::Dx,
dims,
environment.kernels.multiprocessor_count(),
compute_capability,
HalfTriadPolicy::TiledParity,
).ok_or_else(|| {
format!(
"UNCOVERED gemm_bi_backward_dx_tc: shape M={batch} K={n_in} N={n_out} outside the automatic TC route"
)
})?;
let seal = gemm_bi_backward_dx_tc_with_tile_in(environment, dx, dy, w, dims, tile)?;
Ok((tile, seal))
}
pub fn gemm_bi_backward_dx_tc_with_tile(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
dx: TypedPtr,
dy: TypedPtr,
w: TypedPtr,
dims: (usize, usize, usize),
tile: TcTile,
) -> Result<(), String> {
let mut environment = HalfLaunchEnvironment::production(stream, kernels);
gemm_bi_backward_dx_tc_with_tile_in(&mut environment, dx, dy, w, dims, tile).map(drop)
}
fn gemm_bi_backward_dx_tc_with_tile_in<O: PhysicalLaunchObserver>(
environment: &mut HalfLaunchEnvironment<'_, O>,
dx: TypedPtr,
dy: TypedPtr,
w: TypedPtr,
dims: (usize, usize, usize),
tile: TcTile,
) -> Result<HalfNativeBranchSeal, String> {
let checked_dims = GemmDims::nt(dims)?;
let (batch, n_in, _n_out) = checked_dims.tuple();
require_half(dx.dtype, "dX")?;
if dx.dtype != dy.dtype || dy.dtype != w.dtype {
return Err("gemm_bi_backward_dx_tc: mixed dtypes not supported".into());
}
let dt = dx.dtype;
let alpha: f32 = 1.0;
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_out_i = checked_dims.k_i32;
let cfg = tile.launch_cfg(batch, n_in, 73_728)?;
let choice = match tile {
TcTile::Tile128 => {
HalfKernelChoice::new("nt_tc", environment.kernels.gemm_bi_nt_tc_typed.get(dt))
}
TcTile::Tile64 => {
HalfKernelChoice::new("nt_tc64", environment.kernels.gemm_bi_nt_tc64_typed.get(dt))
}
TcTile::Thin16 => {
return Err("Thin16 is an NN-forward rung; the NT dX path has no thin tile".into());
}
TcTile::Rect128x64 => {
return Err("Rect128x64 is a forced TN dW tile; NT has no rectangular route".into());
}
TcTile::Tile64StreamK => {
return Err("Tile64StreamK is a TN dW schedule; NT has no stream-K route".into());
}
};
let mut b = environment.stream.launch_builder(choice.function);
b.arg(&dx.ptr);
b.arg(&dy.ptr);
b.arg(&w.ptr);
b.arg(&alpha);
b.arg(&m_i);
b.arg(&n_i);
b.arg(&k_out_i);
unsafe {
enqueue_half_gemm(
&mut environment.observer,
(environment.kernels, environment.context),
&mut b,
cfg,
HalfGemmObservation {
base: choice.base,
op: ResolvedGemmOp::Nt,
dtype: dt,
dims,
strides: (dims.2, dims.2, dims.1),
tile: tile.extents(),
bk_stages: tile.bk_stages(),
arguments: HalfGemmArguments {
output: dx.ptr,
a: dy.ptr,
b: w.ptr,
bias: 0,
},
},
format_args!("nt_tc ({tile:?})"),
)
}
}
pub fn gemm_bi_forward_typed_native(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
y: TypedPtr,
x: TypedPtr,
w: TypedPtr,
bias_ptr: CUptr, dims: (usize, usize, usize),
) -> Result<(), String> {
let mut environment = HalfLaunchEnvironment::production(stream, kernels);
gemm_bi_forward_typed_in(&mut environment, y, x, w, bias_ptr, dims).map(drop)
}
fn gemm_bi_forward_typed_in<O: PhysicalLaunchObserver>(
environment: &mut HalfLaunchEnvironment<'_, O>,
y: TypedPtr,
x: TypedPtr,
w: TypedPtr,
bias_ptr: CUptr,
dims: (usize, usize, usize),
) -> Result<HalfNativeBranchSeal, String> {
let checked_dims = GemmDims::nn(dims, dims.1)?;
let (batch, n_in, n_out) = checked_dims.tuple();
require_half(y.dtype, "output")?;
if x.dtype != y.dtype || w.dtype != y.dtype {
return Err("gemm_bi_forward_typed_native: mixed dtypes not supported".into());
}
let dt = y.dtype;
let alpha: f32 = 1.0;
let beta: f32 = 0.0;
validate_bias_preseed(alpha, bias_ptr, "gemm_bi_forward_typed_native")?;
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_i = checked_dims.k_i32;
if n_out == 1 && batch >= 1 && n_in >= 32 {
let lda_i = checked_dims.k_i32;
let ldy_i: i32 = 1;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_dims.m_u32.div_ceil(4), 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let choice =
HalfKernelChoice::new("nn_gemv", environment.kernels.gemm_bi_nn_gemv_typed.get(dt));
let mut b = environment.stream.launch_builder(choice.function);
b.arg(&y.ptr);
b.arg(&x.ptr);
b.arg(&w.ptr);
b.arg(&bias_ptr);
b.arg(&alpha);
b.arg(&beta);
b.arg(&m_i);
b.arg(&k_i);
b.arg(&lda_i);
b.arg(&ldy_i);
let seal = unsafe {
enqueue_half_gemm(
&mut environment.observer,
(environment.kernels, environment.context),
&mut b,
cfg,
HalfGemmObservation {
base: choice.base,
op: ResolvedGemmOp::Nn,
dtype: dt,
dims,
strides: (dims.1, dims.2, dims.2),
tile: (4, 1),
bk_stages: (32, 1),
arguments: HalfGemmArguments {
output: y.ptr,
a: x.ptr,
b: w.ptr,
bias: bias_ptr,
},
},
format_args!("nn_gemv typed"),
)
}?;
return Ok(seal);
}
if (1..32).contains(&batch) && (32..=2048).contains(&n_in) && n_out >= 32 {
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_dims.n_u32.div_ceil(32), checked_dims.m_u32, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: checked_u32_product(
checked_dims.k_u32,
checked_u32(std::mem::size_of::<f32>(), "f32 byte width")?,
"typed ultra-thin shared memory",
)?,
};
let choice = HalfKernelChoice::new(
"nn_ultra_thin",
environment.kernels.gemm_bi_nn_ultra_thin_typed.get(dt),
);
let mut b = environment.stream.launch_builder(choice.function);
b.arg(&y.ptr);
b.arg(&x.ptr);
b.arg(&w.ptr);
b.arg(&bias_ptr);
b.arg(&alpha);
b.arg(&beta);
b.arg(&m_i);
b.arg(&n_i);
b.arg(&k_i);
b.arg(&k_i);
b.arg(&n_i);
b.arg(&n_i);
let seal = unsafe {
enqueue_half_gemm(
&mut environment.observer,
(environment.kernels, environment.context),
&mut b,
cfg,
HalfGemmObservation {
base: choice.base,
op: ResolvedGemmOp::Nn,
dtype: dt,
dims,
strides: (dims.1, dims.2, dims.2),
tile: (1, 32),
bk_stages: (32, 1),
arguments: HalfGemmArguments {
output: y.ptr,
a: x.ptr,
b: w.ptr,
bias: bias_ptr,
},
},
format_args!("nn_ultra_thin typed"),
)
}?;
return Ok(seal);
}
if (2..=127).contains(&n_out) && batch >= 1 && n_in >= 1 {
let post_op: i32 = 0;
let small = batch <= 64;
let (grid, block, choice) = if small {
(
checked_tile_grid(checked_dims.m_u32, 16, checked_dims.n_u32, 16)?,
64u32,
HalfKernelChoice::new(
"nn_narrow_small",
environment.kernels.gemm_bi_nn_narrow_small_typed.get(dt),
),
)
} else {
(
checked_tile_grid(checked_dims.m_u32, 64, checked_dims.n_u32, 32)?,
128u32,
HalfKernelChoice::new(
"nn_narrow",
environment.kernels.gemm_bi_nn_narrow_typed.get(dt),
),
)
};
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (block, 1, 1),
shared_mem_bytes: 0,
};
let mut b = environment.stream.launch_builder(choice.function);
b.arg(&y.ptr);
b.arg(&x.ptr);
b.arg(&w.ptr);
b.arg(&bias_ptr);
b.arg(&alpha);
b.arg(&beta);
b.arg(&m_i);
b.arg(&n_i);
b.arg(&k_i);
b.arg(&k_i);
b.arg(&n_i);
b.arg(&n_i);
b.arg(&post_op);
let seal = unsafe {
enqueue_half_gemm(
&mut environment.observer,
(environment.kernels, environment.context),
&mut b,
cfg,
HalfGemmObservation {
base: choice.base,
op: ResolvedGemmOp::Nn,
dtype: dt,
dims,
strides: (dims.1, dims.2, dims.2),
tile: if small { (16, 16) } else { (64, 32) },
bk_stages: (16, 1),
arguments: HalfGemmArguments {
output: y.ptr,
a: x.ptr,
b: w.ptr,
bias: bias_ptr,
},
},
format_args!("nn_narrow typed"),
)
}?;
return Ok(seal);
}
if nn_routes_to_big(
batch,
n_in,
n_out,
environment.kernels.multiprocessor_count(),
) {
let total_tiles = checked_tile_grid(checked_dims.m_u32, 128, checked_dims.n_u32, 128)?;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (total_tiles, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 34 * 1024,
};
let choice =
HalfKernelChoice::new("nn_big", environment.kernels.gemm_bi_nn_big_typed.get(dt));
let mut b = environment.stream.launch_builder(choice.function);
b.arg(&y.ptr);
b.arg(&x.ptr);
b.arg(&w.ptr);
b.arg(&bias_ptr);
b.arg(&alpha);
b.arg(&beta);
b.arg(&m_i);
b.arg(&n_i);
b.arg(&k_i);
b.arg(&k_i);
b.arg(&n_i);
b.arg(&n_i);
let seal = unsafe {
enqueue_half_gemm(
&mut environment.observer,
(environment.kernels, environment.context),
&mut b,
cfg,
HalfGemmObservation {
base: choice.base,
op: ResolvedGemmOp::Nn,
dtype: dt,
dims,
strides: (dims.1, dims.2, dims.2),
tile: (128, 128),
bk_stages: (16, 2),
arguments: HalfGemmArguments {
output: y.ptr,
a: x.ptr,
b: w.ptr,
bias: bias_ptr,
},
},
format_args!("nn_big typed"),
)
}?;
return Ok(seal);
}
Err(format!(
"UNCOVERED gemm_bi_forward_typed_native: Big/Slim buckets not yet implemented — \
shape M={batch} K={n_in} N={n_out}. Disable the batch-invariant flag for \
this configuration."
))
}
pub fn gemm_bi_backward_dw_typed_native(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
dw_ptr: CUptr, dy: TypedPtr,
x_saved: TypedPtr,
dims: (usize, usize, usize),
) -> Result<(), String> {
let mut environment = HalfLaunchEnvironment::production(stream, kernels);
gemm_bi_backward_dw_typed_in(&mut environment, dw_ptr, dy, x_saved, dims).map(drop)
}
fn gemm_bi_backward_dw_typed_in<O: PhysicalLaunchObserver>(
environment: &mut HalfLaunchEnvironment<'_, O>,
dw_ptr: CUptr,
dy: TypedPtr,
x_saved: TypedPtr,
dims: (usize, usize, usize),
) -> Result<HalfNativeBranchSeal, String> {
let checked_dims = GemmDims::tn(dims)?;
let (batch, n_in, n_out) = checked_dims.tuple();
require_half(dy.dtype, "dY")?;
if x_saved.dtype != dy.dtype {
return Err("gemm_bi_backward_dw_typed_native: mixed dtypes not supported".into());
}
let dt = dy.dtype;
let alpha: f32 = 1.0;
if n_out == 1 && n_in >= 4 && batch >= 32 {
let m_i = checked_dims.m_i32;
let k_i = checked_dims.k_i32;
let lda_i = checked_dims.k_i32;
let ldy_i: i32 = 1;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (checked_dims.k_u32.div_ceil(4), 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let choice =
HalfKernelChoice::new("tn_gemv", environment.kernels.gemm_bi_tn_gemv_typed.get(dt));
let mut b = environment.stream.launch_builder(choice.function);
b.arg(&dw_ptr);
b.arg(&x_saved.ptr);
b.arg(&dy.ptr);
b.arg(&alpha);
b.arg(&m_i);
b.arg(&k_i);
b.arg(&lda_i);
b.arg(&ldy_i);
let seal = unsafe {
enqueue_half_gemm(
&mut environment.observer,
(environment.kernels, environment.context),
&mut b,
cfg,
HalfGemmObservation {
base: choice.base,
op: ResolvedGemmOp::Tn,
dtype: dt,
dims,
strides: (dims.1, dims.2, dims.2),
tile: (4, 1),
bk_stages: (32, 1),
arguments: HalfGemmArguments {
output: dw_ptr,
a: x_saved.ptr,
b: dy.ptr,
bias: 0,
},
},
format_args!("tn_gemv typed"),
)
}?;
return Ok(seal);
}
if (2..=127).contains(&n_out) && batch >= 1 && n_in >= 1 {
let m_red_i = checked_dims.m_i32;
let k_out_i = checked_dims.k_i32;
let n_i = checked_dims.n_i32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (
checked_tile_grid(checked_dims.k_u32, 64, checked_dims.n_u32, 32)?,
1,
1,
),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let choice = HalfKernelChoice::new(
"tn_narrow",
environment.kernels.gemm_bi_tn_narrow_typed.get(dt),
);
let mut b = environment.stream.launch_builder(choice.function);
b.arg(&dw_ptr);
b.arg(&x_saved.ptr);
b.arg(&dy.ptr);
b.arg(&alpha);
b.arg(&m_red_i);
b.arg(&k_out_i);
b.arg(&n_i);
let seal = unsafe {
enqueue_half_gemm(
&mut environment.observer,
(environment.kernels, environment.context),
&mut b,
cfg,
HalfGemmObservation {
base: choice.base,
op: ResolvedGemmOp::Tn,
dtype: dt,
dims,
strides: (dims.1, dims.2, dims.2),
tile: (64, 32),
bk_stages: (16, 1),
arguments: HalfGemmArguments {
output: dw_ptr,
a: x_saved.ptr,
b: dy.ptr,
bias: 0,
},
},
format_args!("tn_narrow typed"),
)
}?;
return Ok(seal);
}
if tn_routes_to_big(
batch,
n_in,
n_out,
environment.kernels.multiprocessor_count(),
) {
let alpha: f32 = 1.0;
let m_red_i = checked_dims.m_i32;
let k_out_i = checked_dims.k_i32;
let n_i = checked_dims.n_i32;
let total_tiles = checked_tile_grid(checked_dims.k_u32, 128, checked_dims.n_u32, 128)?;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (total_tiles, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 34 * 1024,
};
let choice =
HalfKernelChoice::new("tn_big", environment.kernels.gemm_bi_tn_big_typed.get(dt));
let mut b = environment.stream.launch_builder(choice.function);
b.arg(&dw_ptr);
b.arg(&x_saved.ptr);
b.arg(&dy.ptr);
b.arg(&alpha);
b.arg(&m_red_i);
b.arg(&k_out_i);
b.arg(&n_i);
let seal = unsafe {
enqueue_half_gemm(
&mut environment.observer,
(environment.kernels, environment.context),
&mut b,
cfg,
HalfGemmObservation {
base: choice.base,
op: ResolvedGemmOp::Tn,
dtype: dt,
dims,
strides: (dims.1, dims.2, dims.2),
tile: (128, 128),
bk_stages: (16, 2),
arguments: HalfGemmArguments {
output: dw_ptr,
a: x_saved.ptr,
b: dy.ptr,
bias: 0,
},
},
format_args!("tn_big typed"),
)
}?;
return Ok(seal);
}
Err(format!(
"UNCOVERED gemm_bi_backward_dw_typed_native: split-M/Slim buckets are upcast-fallback territory — \
shape M={batch} K={n_in} N={n_out}."
))
}
pub fn gemm_bi_backward_dx_typed_native(
stream: &Arc<cudarc::driver::CudaStream>,
kernels: &GpuKernels,
dx: TypedPtr,
dy: TypedPtr,
w: TypedPtr,
dims: (usize, usize, usize),
) -> Result<(), String> {
let mut environment = HalfLaunchEnvironment::production(stream, kernels);
gemm_bi_backward_dx_typed_in(&mut environment, dx, dy, w, dims).map(drop)
}
fn gemm_bi_backward_dx_typed_in<O: PhysicalLaunchObserver>(
environment: &mut HalfLaunchEnvironment<'_, O>,
dx: TypedPtr,
dy: TypedPtr,
w: TypedPtr,
dims: (usize, usize, usize),
) -> Result<HalfNativeBranchSeal, String> {
let checked_dims = GemmDims::nt(dims)?;
let (batch, n_in, n_out) = checked_dims.tuple();
require_half(dx.dtype, "dX")?;
if dy.dtype != dx.dtype || w.dtype != dx.dtype {
return Err("gemm_bi_backward_dx_typed_native: mixed dtypes not supported".into());
}
let dt = dx.dtype;
let alpha: f32 = 1.0;
if n_out == 1 && batch >= 1 && n_in >= 1 {
let m_i = checked_dims.m_i32;
let k_i = checked_dims.k_i32;
let ldx_i = checked_dims.k_i32;
let ldy_i: i32 = 1;
let total = checked_dims.mk_u32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (total.div_ceil(256), 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 0,
};
let choice =
HalfKernelChoice::new("nt_gemv", environment.kernels.gemm_bi_nt_gemv_typed.get(dt));
let mut b = environment.stream.launch_builder(choice.function);
b.arg(&dx.ptr);
b.arg(&dy.ptr);
b.arg(&w.ptr);
b.arg(&alpha);
b.arg(&m_i);
b.arg(&k_i);
b.arg(&ldx_i);
b.arg(&ldy_i);
let seal = unsafe {
enqueue_half_gemm(
&mut environment.observer,
(environment.kernels, environment.context),
&mut b,
cfg,
HalfGemmObservation {
base: choice.base,
op: ResolvedGemmOp::Nt,
dtype: dt,
dims,
strides: (dims.2, dims.2, dims.1),
tile: (1, 1),
bk_stages: (1, 1),
arguments: HalfGemmArguments {
output: dx.ptr,
a: dy.ptr,
b: w.ptr,
bias: 0,
},
},
format_args!("nt_gemv typed"),
)
}?;
return Ok(seal);
}
if (2..=127).contains(&n_out) && batch >= 1 && n_in >= 1 {
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_out_i = checked_dims.k_i32;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (
checked_tile_grid(checked_dims.m_u32, 64, checked_dims.k_u32, 32)?,
1,
1,
),
block_dim: (128, 1, 1),
shared_mem_bytes: 0,
};
let choice = HalfKernelChoice::new(
"nt_narrow",
environment.kernels.gemm_bi_nt_narrow_typed.get(dt),
);
let mut b = environment.stream.launch_builder(choice.function);
b.arg(&dx.ptr);
b.arg(&dy.ptr);
b.arg(&w.ptr);
b.arg(&alpha);
b.arg(&m_i);
b.arg(&n_i);
b.arg(&k_out_i);
let seal = unsafe {
enqueue_half_gemm(
&mut environment.observer,
(environment.kernels, environment.context),
&mut b,
cfg,
HalfGemmObservation {
base: choice.base,
op: ResolvedGemmOp::Nt,
dtype: dt,
dims,
strides: (dims.2, dims.2, dims.1),
tile: (64, 32),
bk_stages: (16, 1),
arguments: HalfGemmArguments {
output: dx.ptr,
a: dy.ptr,
b: w.ptr,
bias: 0,
},
},
format_args!("nt_narrow typed"),
)
}?;
return Ok(seal);
}
if nt_routes_to_big(
batch,
n_in,
n_out,
environment.kernels.multiprocessor_count(),
) {
let alpha: f32 = 1.0;
let m_i = checked_dims.m_i32;
let n_i = checked_dims.n_i32;
let k_out_i = checked_dims.k_i32;
let total_tiles = checked_tile_grid(checked_dims.m_u32, 128, checked_dims.k_u32, 128)?;
let cfg = cudarc::driver::LaunchConfig {
grid_dim: (total_tiles, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 34 * 1024,
};
let choice =
HalfKernelChoice::new("nt_big", environment.kernels.gemm_bi_nt_big_typed.get(dt));
let mut b = environment.stream.launch_builder(choice.function);
b.arg(&dx.ptr);
b.arg(&dy.ptr);
b.arg(&w.ptr);
b.arg(&alpha);
b.arg(&m_i);
b.arg(&n_i);
b.arg(&k_out_i);
let seal = unsafe {
enqueue_half_gemm(
&mut environment.observer,
(environment.kernels, environment.context),
&mut b,
cfg,
HalfGemmObservation {
base: choice.base,
op: ResolvedGemmOp::Nt,
dtype: dt,
dims,
strides: (dims.2, dims.2, dims.1),
tile: (128, 128),
bk_stages: (16, 2),
arguments: HalfGemmArguments {
output: dx.ptr,
a: dy.ptr,
b: w.ptr,
bias: 0,
},
},
format_args!("nt_big typed"),
)
}?;
return Ok(seal);
}
Err(format!(
"UNCOVERED gemm_bi_backward_dx_typed_native: split-N/Slim buckets are upcast-fallback territory — \
shape M={batch} K={n_in} N={n_out}."
))
}
pub(in crate::mamba_ssm::gpu) fn gemm_bi_forward_tc_observed<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
ops: &TcFwdOperands,
dims: (usize, usize, usize),
) -> Result<(TcTile, HalfNativeBranchSeal), String> {
let mut environment = HalfLaunchEnvironment::observed(ctx, observer);
gemm_bi_forward_tc_in(&mut environment, ops.y, ops.x, ops.w, ops.bias_ptr, dims)
}
pub(in crate::mamba_ssm::gpu) fn gemm_bi_forward_tc_with_tile_observed<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &mut O,
ops: &TcFwdOperands,
shape: F32TriadShape,
tile: TcTile,
) -> Result<HalfNativeBranchSeal, String> {
let mut environment = HalfLaunchEnvironment::observed(ctx, observer);
gemm_bi_forward_tc_with_tile_in(&mut environment, ops, shape, tile)
}
pub(in crate::mamba_ssm::gpu) fn gemm_bi_backward_dw_tc_observed<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
dw_ptr: CUptr,
dy: TypedPtr,
x_saved: TypedPtr,
dims: (usize, usize, usize),
) -> Result<(TcTile, HalfNativeBranchSeal), String> {
let mut environment = HalfLaunchEnvironment::observed(ctx, observer);
gemm_bi_backward_dw_tc_in(
&mut environment,
dw_ptr,
dy,
x_saved,
dims,
ctx.half_triad_policy(),
)
}
pub(in crate::mamba_ssm::gpu) fn gemm_bi_backward_dw_tc_with_tile_observed<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &mut O,
dw_ptr: CUptr,
dy: TypedPtr,
x_saved: TypedPtr,
dims: (usize, usize, usize),
tile: TcTile,
) -> Result<HalfNativeBranchSeal, String> {
let mut environment = HalfLaunchEnvironment::observed(ctx, observer);
gemm_bi_backward_dw_tc_with_tile_in(&mut environment, dw_ptr, dy, x_saved, dims, tile)
}
pub(in crate::mamba_ssm::gpu) fn gemm_bi_backward_dx_tc_observed<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
dx: TypedPtr,
dy: TypedPtr,
w: TypedPtr,
dims: (usize, usize, usize),
) -> Result<(TcTile, HalfNativeBranchSeal), String> {
let mut environment = HalfLaunchEnvironment::observed(ctx, observer);
gemm_bi_backward_dx_tc_in(&mut environment, dx, dy, w, dims)
}
pub(in crate::mamba_ssm::gpu) fn gemm_bi_backward_dx_tc_with_tile_observed<
O: PhysicalLaunchObserver,
>(
ctx: &GpuCtx,
observer: &mut O,
dx: TypedPtr,
dy: TypedPtr,
w: TypedPtr,
dims: (usize, usize, usize),
tile: TcTile,
) -> Result<HalfNativeBranchSeal, String> {
let mut environment = HalfLaunchEnvironment::observed(ctx, observer);
gemm_bi_backward_dx_tc_with_tile_in(&mut environment, dx, dy, w, dims, tile)
}
pub(in crate::mamba_ssm::gpu) fn gemm_bi_forward_typed_observed<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
ops: &TcFwdOperands,
dims: (usize, usize, usize),
) -> Result<HalfNativeBranchSeal, String> {
let mut environment = HalfLaunchEnvironment::observed(ctx, observer);
gemm_bi_forward_typed_in(&mut environment, ops.y, ops.x, ops.w, ops.bias_ptr, dims)
}
pub(in crate::mamba_ssm::gpu) fn gemm_bi_backward_dw_typed_observed<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
dw_ptr: CUptr,
dy: TypedPtr,
x_saved: TypedPtr,
dims: (usize, usize, usize),
) -> Result<HalfNativeBranchSeal, String> {
let mut environment = HalfLaunchEnvironment::observed(ctx, observer);
gemm_bi_backward_dw_typed_in(&mut environment, dw_ptr, dy, x_saved, dims)
}
pub(in crate::mamba_ssm::gpu) fn gemm_bi_backward_dx_typed_observed<O: PhysicalLaunchObserver>(
ctx: &GpuCtx,
observer: &mut O,
dx: TypedPtr,
dy: TypedPtr,
w: TypedPtr,
dims: (usize, usize, usize),
) -> Result<HalfNativeBranchSeal, String> {
let mut environment = HalfLaunchEnvironment::observed(ctx, observer);
gemm_bi_backward_dx_typed_in(&mut environment, dx, dy, w, dims)
}
#[cfg(test)]
mod prepared_f32_launch_tests {
use super::*;
use crate::mamba_ssm::gpu::buffers::{
managed_allocation_epoch_for_ranges, register_managed_allocation_range,
};
use crate::mamba_ssm::gpu::context::{BiGemmFamily, F32TriadPolicy, HalfTriadPolicy};
use crate::mamba_ssm::gpu::device::GpuDevice;
use crate::mamba_ssm::gpu::gemm_bi_triad::contract::{
TF32_NT_SPLITK4_S3_SPEC, TF32_NT_SPLITK4_S4_SPEC, TF32_NT_SPLITK8_S3_SPEC,
TF32_NT_SPLITK8_S4_SPEC, TF32_SPLITK2_SPEC, TF32_SPLITK4_SPEC,
allocation_identity_query_count, reset_allocation_identity_query_count,
};
use crate::mamba_ssm::gpu::kernel_identity::{
ArtifactIdentity, ArtifactKind, BackendSet, COMPILER_REVISION, COMPOSER_REVISION,
CompilerIdentity, CudaTarget, DeviceCaps, DeviceIdentity, DriverIdentity, GemmPolicy,
GemmRouteIdentity, ModuleKind, NUMERIC_ABI_REVISION, NumericContractSet, POLICY_REVISION,
ResolvedGemmOp, SCHEDULE_REVISION, TUNING_TABLE_REVISION, build_artifact_set,
};
use std::cell::Cell;
fn scalar_fixture(
op: ResolvedGemmOp,
dims: (usize, usize, usize),
) -> (F32TriadRequest, F32TriadOperands, ScalarDispatchPlan) {
let request = F32TriadRequest {
op,
shape: F32TriadShape::contiguous(op, dims),
};
let operands = F32TriadOperands {
output: 0x1000,
a: 0x2000,
b: 0x3000,
bias: (op == ResolvedGemmOp::Nn).then_some(0x4000),
alpha: 1.0,
beta: if op == ResolvedGemmOp::Tn { 1.0 } else { 0.0 },
};
let plan = scalar_dispatch_plan(request, 142).unwrap();
(request, operands, plan)
}
#[test]
fn tf32_splitk_launch_plan_freezes_one_fused_node_and_workspaces() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, (64, 1_536, 384)),
};
let plan = tf32_splitk_launch_plan(request, &TF32_SPLITK4_SPEC).unwrap();
assert_eq!(plan.scratch_elements, 98_304);
assert_eq!(plan.counter_elements, 48);
assert_eq!(plan.fused.grid_dim, (12, 4, 4));
assert_eq!(plan.fused.block_dim, (128, 1, 1));
assert_eq!(plan.fused.shared_mem_bytes, 29_696);
let oversized = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, (4_096, 128, 1_024)),
};
assert!(tf32_splitk_launch_plan(oversized, &TF32_SPLITK4_SPEC).is_err());
let oversized_row_grid = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, (1_048_561, 1, 1)),
};
assert!(tf32_splitk_launch_plan(oversized_row_grid, &TF32_SPLITK4_SPEC).is_err());
let wrong_op = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (64, 1_536, 384)),
};
assert!(tf32_splitk_launch_plan(wrong_op, &TF32_SPLITK4_SPEC).is_err());
}
#[test]
fn tf32_splitk2_launch_plan_and_route_freeze_two_partitions() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, (64, 833, 384)),
};
let operands = scalar_fixture(ResolvedGemmOp::Nn, (64, 833, 384)).1;
let plan = tf32_splitk_launch_plan(request, &TF32_SPLITK2_SPEC).unwrap();
assert_eq!(plan.scratch_elements, 49_152);
assert_eq!(plan.counter_elements, 48);
assert_eq!(plan.fused.grid_dim, (12, 4, 2));
assert_eq!(plan.fused.block_dim, (128, 1, 1));
assert_eq!(plan.fused.shared_mem_bytes, 29_696);
let splitk2 = tf32_splitk_resolved_routes(
request,
operands,
&TF32_SPLITK2_SPEC,
portable_binding(),
[9; 32],
plan,
);
let splitk4_plan = tf32_splitk_launch_plan(request, &TF32_SPLITK4_SPEC).unwrap();
let splitk4 = tf32_splitk_resolved_routes(
request,
operands,
&TF32_SPLITK4_SPEC,
portable_binding(),
[9; 32],
splitk4_plan,
);
assert_eq!(splitk2.len(), 1);
assert_eq!(splitk2[0].backend, PhysicalGemmBackend::MmaTf32RnaSplitK2);
assert_eq!(
splitk2[0].numeric_contract,
ResolvedNumericContract::MmaTf32RnaSplitK2
);
assert_eq!(
splitk2[0].ownership,
ResolvedOutputOwnership::LastCtaPerOutputTileFixedSplitK2Reduce
);
assert_ne!(splitk2[0].backend, splitk4[0].backend);
assert_ne!(splitk2[0].numeric_contract, splitk4[0].numeric_contract);
assert_ne!(splitk2[0].ownership, splitk4[0].ownership);
assert_ne!(
splitk2[0].launch.arguments_digest,
splitk4[0].launch.arguments_digest
);
}
#[test]
fn tf32_nt_splitk_live_and_reverse_plans_freeze_p4_and_p8_geometry() {
let cases = [
((64, 384, 1_536), (12, 4, 4), (12, 2, 8)),
((384, 64, 1_536), (2, 24, 4), (2, 12, 8)),
];
for (dims, p4_grid, p8_grid) in cases {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, dims),
};
for spec in [&TF32_NT_SPLITK4_S3_SPEC, &TF32_NT_SPLITK4_S4_SPEC] {
let plan = tf32_splitk_launch_plan(request, spec).unwrap();
assert_eq!(plan.scratch_elements, 98_304);
assert_eq!(plan.counter_elements, 48);
assert_eq!(plan.fused.grid_dim, p4_grid);
assert_eq!(plan.fused.block_dim, (128, 1, 1));
assert_eq!(plan.fused.shared_mem_bytes, spec.dynamic_shared_bytes);
}
for spec in [&TF32_NT_SPLITK8_S3_SPEC, &TF32_NT_SPLITK8_S4_SPEC] {
let plan = tf32_splitk_launch_plan(request, spec).unwrap();
assert_eq!(plan.scratch_elements, 196_608);
assert_eq!(plan.counter_elements, 24);
assert_eq!(plan.fused.grid_dim, p8_grid);
assert_eq!(plan.fused.block_dim, (128, 1, 1));
assert_eq!(plan.fused.shared_mem_bytes, spec.dynamic_shared_bytes);
}
}
}
#[test]
fn tf32_nt_splitk8_has_distinct_numeric_and_ownership_identity() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (64, 384, 1_536)),
};
let operands = scalar_fixture(ResolvedGemmOp::Nt, (64, 384, 1_536)).1;
let p4_plan = tf32_splitk_launch_plan(request, &TF32_NT_SPLITK4_S3_SPEC).unwrap();
let p8_plan = tf32_splitk_launch_plan(request, &TF32_NT_SPLITK8_S3_SPEC).unwrap();
let p4 = tf32_splitk_resolved_routes(
request,
operands,
&TF32_NT_SPLITK4_S3_SPEC,
portable_binding(),
[9; 32],
p4_plan,
);
let p8 = tf32_splitk_resolved_routes(
request,
operands,
&TF32_NT_SPLITK8_S3_SPEC,
portable_binding(),
[9; 32],
p8_plan,
);
assert_eq!(p8[0].backend, PhysicalGemmBackend::MmaTf32RnaSplitK8);
assert_eq!(
p8[0].numeric_contract,
ResolvedNumericContract::MmaTf32RnaSplitK8
);
assert_eq!(
p8[0].ownership,
ResolvedOutputOwnership::LastCtaPerOutputTileFixedSplitK8Reduce
);
assert_eq!(p8[0].tile, (32, 32));
assert_eq!(p8[0].launch.grid_dim, (12, 2, 8));
assert_ne!(p4[0].backend, p8[0].backend);
assert_ne!(p4[0].numeric_contract, p8[0].numeric_contract);
assert_ne!(p4[0].ownership, p8[0].ownership);
assert_ne!(p4[0].launch.arguments_digest, p8[0].launch.arguments_digest);
}
fn sm120_context_fixture() -> GemmRouteIdentity {
let compiler = CompilerIdentity {
source_digest: [1; 32],
invocation_digest: [2; 32],
header_manifest_digest: [3; 32],
target: CudaTarget::new("compute_120").unwrap(),
nvrtc_version: (12, 8),
nvrtc_library_domain: [4; 32],
nvrtc_library_known: true,
output_kind: ArtifactKind::Ptx,
composer_revision: COMPOSER_REVISION,
compiler_revision: COMPILER_REVISION,
numeric_abi_revision: NUMERIC_ABI_REVISION,
schedule_revision: SCHEDULE_REVISION,
};
let artifact = |module_kind, seed| ArtifactIdentity {
module_kind,
artifact_kind: ArtifactKind::Ptx,
compile_key: [seed; 32],
artifact_digest: [seed + 1; 32],
};
GemmRouteIdentity {
policy: GemmPolicy {
batch_invariant: true,
bi_tensor_cores: true,
fast_gemm: false,
cublas_tf32: false,
f32_triad_policy: F32TriadPolicy::ExactScalarFma,
half_triad_policy: HalfTriadPolicy::TiledParity,
bi_gemm_family: BiGemmFamily::Triad,
},
backend_set: BackendSet::TRIAD,
numeric_contracts: NumericContractSet::TRIAD_MMA_SYNC,
compiler,
artifacts: build_artifact_set(&[
artifact(ModuleKind::Fixed, 5),
artifact(ModuleKind::TriadScalar, 7),
artifact(ModuleKind::TriadSm80, 9),
artifact(ModuleKind::TriadSm120, 11),
])
.unwrap(),
policy_revision: POLICY_REVISION,
policy_hash: [13; 32],
device: DeviceIdentity {
compute_capability: (12, 0),
multiprocessor_count: 24,
target: CudaTarget::new("sm_120").unwrap(),
driver: DriverIdentity {
api_version: 12_800,
build_sources: 1,
build_digest: [14; 32],
},
},
device_caps: DeviceCaps {
compute_capability: (12, 0),
nvrtc_version: (12, 8),
accepted_target: Some(CudaTarget::new("compute_120").unwrap()),
optin_shared_bytes: 101_376,
tensor_map_access: true,
},
tuning_table_revision: TUNING_TABLE_REVISION,
schedule_set_revision: SCHEDULE_REVISION,
state_capacity: 64,
}
}
fn sm120_key_fixture() -> (Sm120ForcedRoute, Sm120AutoRequest, Sm120PreparedKey) {
let route = SM120_AUTO_CELLS_CC120[0];
let request = Sm120AutoRequest {
op: route.op,
dtype: route.dtype,
shape: route.shape,
a_ptr: 0x1_0000,
b_ptr: 0x2_0000,
multiprocessors: 170,
half_policy: HalfTriadPolicy::TiledParity,
operands: Sm120LaunchOperands {
output_ptr: 0x3_0000,
bias_ptr: 0x4_0000,
alpha: 1.0,
beta: 0.0,
},
};
let key = Sm120PreparedKey::new(sm120_context_fixture(), route, request);
(route, request, key)
}
#[test]
fn f32_prepared_key_seals_context_request_pointers_and_scalar_bits() {
let context = sm120_context_fixture();
let context_token = 7;
let policy = context.policy;
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, (49, 65, 129)),
};
let operands = F32TriadOperands {
output: 0x1_0000,
a: 0x2_0000,
b: 0x3_0000,
bias: Some(0x4_0000),
alpha: 1.0,
beta: 0.0,
};
let expected = PreparedF32Key::new(
context_token,
policy,
F32PreparedSelection::Automatic,
request,
operands,
);
let assert_distinct = |context_token, policy, selection, request, operands| {
assert_ne!(
PreparedF32Key::new(context_token, policy, selection, request, operands),
expected
);
};
assert_distinct(
context_token + 1,
policy,
F32PreparedSelection::Automatic,
request,
operands,
);
for changed_policy in [
GemmPolicy {
batch_invariant: !policy.batch_invariant,
..policy
},
GemmPolicy {
bi_tensor_cores: !policy.bi_tensor_cores,
..policy
},
GemmPolicy {
fast_gemm: !policy.fast_gemm,
..policy
},
GemmPolicy {
cublas_tf32: !policy.cublas_tf32,
..policy
},
GemmPolicy {
f32_triad_policy: F32TriadPolicy::AllowDeterministicTf32,
..policy
},
GemmPolicy {
half_triad_policy: HalfTriadPolicy::AllowStreamKFixedOrder,
..policy
},
GemmPolicy {
bi_gemm_family: BiGemmFamily::Inference,
..policy
},
] {
assert_distinct(
context_token,
changed_policy,
F32PreparedSelection::Automatic,
request,
operands,
);
}
assert_distinct(
context_token,
policy,
F32PreparedSelection::ExactScalar,
request,
operands,
);
let forced = Tf32PhysicalRoute::Sm120TmaFmaExact(Sm120FmaRoute {
tile: Sm120FmaTile::M128N64,
kvec: false,
splits: 1,
});
assert_distinct(
context_token,
policy,
F32PreparedSelection::Forced(forced),
request,
operands,
);
assert_ne!(
PreparedF32Key::new(
context_token,
policy,
F32PreparedSelection::Forced(forced),
request,
operands,
),
PreparedF32Key::new(
context_token,
policy,
F32PreparedSelection::Forced(Tf32PhysicalRoute::Sm120TmaFmaExact(Sm120FmaRoute {
tile: Sm120FmaTile::M64N128,
kvec: false,
splits: 1,
},)),
request,
operands,
),
);
let changed_request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (49, 65, 129)),
};
assert_distinct(
context_token,
policy,
F32PreparedSelection::Automatic,
changed_request,
operands,
);
for pointer in 0..4 {
let mut changed = operands;
match pointer {
0 => changed.output += 16,
1 => changed.a += 16,
2 => changed.b += 16,
_ => changed.bias = changed.bias.map(|bias| bias + 16),
}
assert_distinct(
context_token,
policy,
F32PreparedSelection::Automatic,
request,
changed,
);
}
let changed_bias_nullness = F32TriadOperands {
bias: None,
..operands
};
assert_distinct(
context_token,
policy,
F32PreparedSelection::Automatic,
request,
changed_bias_nullness,
);
let changed_alpha = F32TriadOperands {
alpha: f32::from_bits(operands.alpha.to_bits() ^ 1),
..operands
};
assert_distinct(
context_token,
policy,
F32PreparedSelection::Automatic,
request,
changed_alpha,
);
let changed_beta = F32TriadOperands {
beta: -0.0,
..operands
};
assert_distinct(
context_token,
policy,
F32PreparedSelection::Automatic,
request,
changed_beta,
);
}
#[test]
fn sm120_prepared_key_seals_route_context_pointers_and_scalar_bits() {
let (route, request, expected) = sm120_key_fixture();
let context = sm120_context_fixture();
let assert_distinct = |route, request, context| {
assert_ne!(Sm120PreparedKey::new(context, route, request), expected);
};
let mut changed_route = route;
changed_route.physical.stages = Sm120Stages::S3;
assert_distinct(changed_route, request, context);
let mut changed_context = context;
changed_context.policy_hash[0] ^= 1;
assert_distinct(route, request, changed_context);
for pointer in 0..4 {
let mut changed = request;
match pointer {
0 => changed.a_ptr += 128,
1 => changed.b_ptr += 128,
2 => changed.operands.output_ptr += 128,
_ => changed.operands.bias_ptr += 128,
}
assert_distinct(route, changed, context);
}
let mut changed_alpha = request;
changed_alpha.operands.alpha = f32::from_bits(request.operands.alpha.to_bits() ^ 1);
assert_distinct(route, changed_alpha, context);
let mut changed_beta = request;
changed_beta.operands.beta = -0.0;
assert_distinct(route, changed_beta, context);
}
#[test]
fn sm120_cache_action_is_prepared_only_and_fail_closed_during_capture() {
assert_eq!(
sm120_cache_action(true, false, Sm120ManagedEpochState::Missing),
Sm120CacheAction::CaptureMissing
);
assert_eq!(
sm120_cache_action(true, true, Sm120ManagedEpochState::Stale),
Sm120CacheAction::CaptureStale
);
assert_eq!(
sm120_cache_action(true, true, Sm120ManagedEpochState::Untracked),
Sm120CacheAction::CaptureUntracked
);
assert_eq!(
sm120_cache_action(true, true, Sm120ManagedEpochState::Current),
Sm120CacheAction::UsePrepared
);
assert_eq!(
sm120_cache_action(false, false, Sm120ManagedEpochState::Missing),
Sm120CacheAction::Prepare
);
assert_eq!(
sm120_cache_action(false, true, Sm120ManagedEpochState::Stale),
Sm120CacheAction::Validate
);
assert_eq!(
sm120_capture_cache_error(Sm120CacheAction::CaptureMissing),
"prepared SM120 Triad cache entry is missing during graph capture; run eager warmup again"
);
assert_eq!(
sm120_capture_cache_error(Sm120CacheAction::CaptureStale),
"prepared SM120 Triad allocation epoch changed during graph capture; run eager warmup again"
);
assert_eq!(
sm120_capture_cache_error(Sm120CacheAction::CaptureUntracked),
"prepared SM120 Triad automatic capture requires managed allocations; run eager warmup again"
);
}
#[test]
fn physical_tf32_route_preserves_semantic_arguments_while_projecting_resources() {
let source = include_str!("launch.rs");
let start = source
.find("fn physical_prepared_f32_route(")
.expect("physical F32 route helper");
let end = source[start..]
.find("pub(in crate::mamba_ssm::gpu) fn prepare_f32_triad(")
.map(|offset| start + offset)
.expect("prepare entry after physical route helper");
let helper = &source[start..end];
let required = [
"physical.resources_digest = prepared.resources.physical_digest();",
".map(F32PreparedTensorMaps::physical_identity_digest)",
"physical.tensor_maps_digest = maps_digest;",
];
for fragment in required {
assert!(helper.contains(fragment), "missing {fragment}");
let mutated = helper.replacen(fragment, "physical_digest_step_removed();", 1);
assert!(
required.iter().any(|required| !mutated.contains(required)),
"physical TF32 route accepted removal of {fragment}"
);
}
assert!(
!helper.contains("physical.launch.arguments_digest ="),
"physical TF32 projection must preserve the semantic kernel-argument digest"
);
let production = source
.split_once("\n#[cfg(test)]\nmod prepared_f32_launch_tests {")
.map(|(production, _)| production)
.expect("prepared F32 launch test-module boundary");
for caller in [
"fn prepare_prepared_f32_direct_graph_sequence",
"fn prepare_sm89_tf32_tn_pre_rna_graph_sequence",
"fn prepare_tf32_splitk_direct_graph_sequence",
"fn enqueue_sm89_tf32_tn_pre_rna",
"fn enqueue_tf32_splitk_f32",
"fn enqueue_validated_prepared_f32_triad_observed",
] {
assert!(
production.contains(caller),
"missing physical F32 route producer {caller}"
);
}
assert_eq!(
production
.matches("physical_prepared_f32_route(prepared,")
.count(),
6,
"eager observation and direct graph preparation must share the physical route"
);
for (call, expected) in [
(
"physical_prepared_f32_route(prepared, prepared.routes[0])",
1,
),
("physical_prepared_f32_route(prepared, *resolved)", 4),
("physical_prepared_f32_route(prepared, resolved)", 1),
] {
assert_eq!(
production.matches(call).count(),
expected,
"physical F32 route call form changed: {call}"
);
}
}
#[test]
fn tf32_argument_digests_bind_bias_nullness() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, (64, 1_536, 384)),
};
let without_bias = F32TriadOperands {
output: 0x1000,
a: 0x2000,
b: 0x3000,
bias: None,
alpha: 1.0,
beta: 0.0,
};
let with_bias = F32TriadOperands {
bias: Some(0x4000),
..without_bias
};
assert_ne!(
tf32_kernel_arguments_digest(request, without_bias, "tf32", [0; 32]),
tf32_kernel_arguments_digest(request, with_bias, "tf32", [0; 32])
);
assert_ne!(
tf32_splitk_arguments_digest(request, without_bias, "splitk", 0, 4, 98_304),
tf32_splitk_arguments_digest(request, with_bias, "splitk", 0, 4, 98_304)
);
}
#[test]
fn specialized_device_identity_preserves_the_loaded_kernel_sm_count() {
let target = CudaTarget::new("sm_100").unwrap();
let driver = DriverIdentity {
api_version: 13_020,
build_sources: 1,
build_digest: [9; 32],
};
let identity = specialized_device_identity((10, 0), 73, target, driver);
assert_eq!(
identity,
DeviceIdentity {
compute_capability: (10, 0),
multiprocessor_count: 73,
target,
driver,
}
);
}
#[test]
fn current_managed_epoch_skips_cached_full_validation() {
let registration = register_managed_allocation_range(0x3301, 0x90_0000, 4096).unwrap();
let mut stamp = Some(
managed_allocation_epoch_for_ranges(0x3301, &[(0x90_0000, 4096)])
.expect("managed allocation stamp"),
);
let validations = Cell::new(0);
refresh_cached_validation(
&mut stamp,
|| panic!("stable epoch must not query capture status"),
|| {
validations.set(validations.get() + 1);
Ok(None)
},
)
.unwrap();
assert_eq!(validations.get(), 0);
drop(registration);
}
#[test]
fn changed_managed_epoch_refreshes_once_outside_capture() {
let first = register_managed_allocation_range(0x3302, 0xa0_0000, 4096).unwrap();
let mut stamp = Some(
managed_allocation_epoch_for_ranges(0x3302, &[(0xa0_0000, 4096)])
.expect("first managed allocation stamp"),
);
drop(first);
let second = register_managed_allocation_range(0x3302, 0xa0_0000, 4096).unwrap();
let refreshed = managed_allocation_epoch_for_ranges(0x3302, &[(0xa0_0000, 4096)])
.expect("replacement managed allocation stamp");
let validations = Cell::new(0);
refresh_cached_validation(
&mut stamp,
|| Ok(false),
|| {
validations.set(validations.get() + 1);
Ok(Some(refreshed))
},
)
.unwrap();
assert_eq!(validations.get(), 1);
assert!(stamp.as_ref().is_some_and(|stamp| stamp.is_current()));
drop(second);
}
#[test]
fn changed_managed_epoch_requires_warmup_during_capture() {
let registration = register_managed_allocation_range(0x3303, 0xb0_0000, 4096).unwrap();
let mut stamp = Some(
managed_allocation_epoch_for_ranges(0x3303, &[(0xb0_0000, 4096)])
.expect("managed allocation stamp"),
);
drop(registration);
let error = refresh_cached_validation(
&mut stamp,
|| Ok(true),
|| panic!("capture must not run full allocation validation"),
)
.expect_err("stale managed epoch must reject capture");
assert!(error.contains("eager warmup"), "{error}");
}
#[test]
fn bounded_cache_sweep_removes_stale_entries_before_insertion() {
let mut entries = HashMap::from([(1_u32, true), (2, false)]);
make_room_in_bounded_cache(&mut entries, &3, F32_PREPARED_CACHE_LIMIT, |current| {
*current
});
assert_eq!(entries, HashMap::from([(1, true)]));
}
#[test]
fn bounded_cache_clears_before_a_distinct_entry_exceeds_the_limit() {
let mut entries = (0..F32_PREPARED_CACHE_LIMIT)
.map(|key| (key, true))
.collect::<HashMap<_, _>>();
make_room_in_bounded_cache(
&mut entries,
&F32_PREPARED_CACHE_LIMIT,
F32_PREPARED_CACHE_LIMIT,
|current| *current,
);
assert!(entries.is_empty());
}
#[test]
fn untracked_cached_resources_keep_full_validation() {
let mut stamp = None;
let validations = Cell::new(0);
for _ in 0..2 {
refresh_cached_validation(
&mut stamp,
|| panic!("untracked validation keeps the existing capture behavior"),
|| {
validations.set(validations.get() + 1);
Ok(None)
},
)
.unwrap();
}
assert_eq!(validations.get(), 2);
}
#[test]
#[ignore = "requires a CUDA device"]
fn managed_f32_cache_hit_skips_allocation_identity_queries() {
let device = GpuDevice::new(0).expect("open CUDA device");
let ctx = GpuCtx::new(&device).expect("create GPU context");
ctx.set_gemm_mode(crate::mamba_ssm::gpu::GemmMode::Deterministic)
.unwrap();
ctx.set_bi_gemm_family(BiGemmFamily::Triad);
ctx.set_f32_triad_policy(F32TriadPolicy::ExactScalarFma);
let dims = (8, 8, 8);
let x = GpuBuffer::from_cpu(&ctx.stream, &vec![0.25; dims.0 * dims.1]).expect("allocate X");
let w = GpuBuffer::from_cpu(&ctx.stream, &vec![0.5; dims.1 * dims.2]).expect("allocate W");
let mut y = GpuBuffer::zeros(&ctx.stream, dims.0 * dims.2).expect("allocate Y");
ctx.stream.synchronize().expect("finish allocations");
reset_allocation_identity_query_count();
launch_cached_f32_forward(&ctx, &mut y, &x, w.cached_ptr(), 0, dims)
.expect("warm managed f32 cache");
assert!(allocation_identity_query_count() >= 3);
reset_allocation_identity_query_count();
launch_cached_f32_forward(&ctx, &mut y, &x, w.cached_ptr(), 0, dims)
.expect("hit managed f32 cache");
assert_eq!(allocation_identity_query_count(), 0);
ctx.stream.synchronize().expect("finish cached launches");
}
fn alignment_fixture_values(len: usize, salt: usize) -> Vec<f32> {
(0..len)
.map(|index| ((index * 17 + salt * 13) % 31) as f32 / 32.0 - 0.5)
.collect()
}
fn assert_managed_nn_b_subview(policy: F32TriadPolicy) {
let device = GpuDevice::new(0).expect("open CUDA device");
assert_eq!(device.compute_capability, (8, 9), "Ada alignment gate");
let ctx = GpuCtx::new(&device).expect("create GPU context");
ctx.set_gemm_mode(crate::mamba_ssm::gpu::GemmMode::Deterministic)
.unwrap();
ctx.set_bi_gemm_family(BiGemmFamily::Triad);
ctx.set_f32_triad_policy(policy);
let dims = (1024, 16, 128);
let x = GpuBuffer::from_cpu(&ctx.stream, &alignment_fixture_values(dims.0 * dims.1, 3))
.expect("allocate X");
let weights = alignment_fixture_values(dims.1 * dims.2, 11);
let aligned_w = GpuBuffer::from_cpu(&ctx.stream, &weights).expect("allocate aligned W");
let mut shifted_weights = vec![123.0; weights.len() + 1];
shifted_weights[1..].copy_from_slice(&weights);
let shifted_w =
GpuBuffer::from_cpu(&ctx.stream, &shifted_weights).expect("allocate shifted W");
let mut expected =
GpuBuffer::zeros(&ctx.stream, dims.0 * dims.2).expect("allocate expected output");
let mut actual =
GpuBuffer::zeros(&ctx.stream, dims.0 * dims.2).expect("allocate actual output");
ctx.stream
.synchronize()
.expect("finish fixture allocations");
launch_cached_f32_forward(&ctx, &mut expected, &x, aligned_w.cached_ptr(), 0, dims)
.expect("launch aligned reference");
ctx.stream.synchronize().expect("finish aligned reference");
let expected = expected
.to_cpu(&ctx.stream)
.expect("download aligned output");
let trace = ctx
.record_eager_gemm_trace(|| {
launch_cached_f32_forward(
&ctx,
&mut actual,
&x,
shifted_w.raw_ptr_at(&ctx.stream, 1),
0,
dims,
)
})
.expect("record shifted-B launch");
let symbols = trace
.routes()
.iter()
.map(|route| route.symbol)
.collect::<Vec<_>>();
assert_eq!(symbols, ["nn_slim"]);
ctx.stream.synchronize().expect("finish shifted-B launch");
let actual = actual
.to_cpu(&ctx.stream)
.expect("download shifted-B output");
assert_eq!(
actual
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>(),
expected
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>()
);
}
#[test]
#[ignore = "requires an SM89 CUDA device"]
fn exact_scalar_prepared_nn_accepts_managed_b_plus_one() {
assert_managed_nn_b_subview(F32TriadPolicy::ExactScalarFma);
}
#[test]
#[ignore = "requires an SM89 CUDA device"]
fn automatic_scalar_prepared_nn_accepts_managed_b_plus_one() {
assert_managed_nn_b_subview(F32TriadPolicy::AllowDeterministicTf32);
}
#[test]
#[ignore = "requires an SM89 CUDA device"]
fn operand_aware_auto_prepares_and_repeats_the_measured_tf32_route() {
let device = GpuDevice::new(0).expect("open CUDA device");
assert_eq!(device.compute_capability, (8, 9), "Ada route gate");
let ctx = GpuCtx::new(&device).expect("create GPU context");
ctx.set_gemm_mode(crate::mamba_ssm::gpu::GemmMode::Deterministic)
.unwrap();
ctx.set_bi_gemm_family(BiGemmFamily::Triad);
ctx.set_f32_triad_policy(F32TriadPolicy::AllowDeterministicTf32);
let dims = (49, 65, 129);
let a = GpuBuffer::from_cpu(&ctx.stream, &alignment_fixture_values(dims.0 * dims.1, 5))
.expect("allocate A");
let b = GpuBuffer::from_cpu(&ctx.stream, &alignment_fixture_values(dims.1 * dims.2, 17))
.expect("allocate B");
let mut output = GpuBuffer::zeros(&ctx.stream, dims.0 * dims.2).expect("allocate output");
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, dims),
};
let operands = F32TriadOperands {
output: output.cached_ptr(),
a: a.cached_ptr(),
b: b.cached_ptr(),
bias: None,
alpha: 1.0,
beta: 0.0,
};
let prepared = prepare_f32_triad(&ctx, request, operands).expect("prepare TF32 AUTO");
assert!(matches!(
prepared.kind,
PreparedF32Kind::Tf32 {
route: Tf32PhysicalRoute::MmaTf32Rna(Tf32PortableRoute {
tile: Tf32PortableTile::M16N32,
stages: Tf32PortableStages::S4,
}),
..
}
));
assert_eq!(
prepared
.routes
.iter()
.map(|route| route.symbol)
.collect::<Vec<_>>(),
["nn_sm80_mma_tf32_m16n32_bk32_s4"]
);
let mut repeat_bits = None;
for _ in 0..2 {
output
.upload(&ctx.stream, &vec![0.0; dims.0 * dims.2])
.expect("reset output");
unsafe {
launch_prepared_f32_triad(&ctx, &prepared, |_| {
Err("TF32 AUTO unexpectedly requested scalar execution".into())
})
}
.expect("launch prepared TF32 AUTO");
ctx.stream.synchronize().expect("finish TF32 AUTO");
let bits = output
.to_cpu(&ctx.stream)
.expect("download TF32 AUTO output")
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>();
if let Some(expected) = repeat_bits.replace(bits.clone()) {
assert_eq!(bits, expected, "TF32 AUTO repeat changed output bits");
}
}
}
fn launch_scalar_nn_slim_raw(
ctx: &GpuCtx,
output: CUptr,
a: CUptr,
b: CUptr,
dims: (usize, usize, usize),
config: cudarc::driver::LaunchConfig,
) -> Result<(), String> {
let (m, k, n) = dims;
let bias = 0_u64;
let alpha = 1.0_f32;
let beta = 0.0_f32;
let m = i32::try_from(m).expect("M fits i32");
let k = i32::try_from(k).expect("K fits i32");
let n = i32::try_from(n).expect("N fits i32");
let mut builder = ctx.stream.launch_builder(&ctx.kernels.gemm_bi_nn_slim);
builder.arg(&output);
builder.arg(&a);
builder.arg(&b);
builder.arg(&bias);
builder.arg(&alpha);
builder.arg(&beta);
builder.arg(&m);
builder.arg(&n);
builder.arg(&k);
builder.arg(&k);
builder.arg(&n);
builder.arg(&n);
unsafe { builder.launch(config) }
.map(|_| ())
.map_err(|error| format!("launch nn_slim alignment fixture: {error:?}"))
}
#[test]
#[ignore = "requires an SM89 CUDA device"]
fn exact_scalar_prepared_nn_accepts_managed_output_plus_one() {
let device = GpuDevice::new(0).expect("open CUDA device");
assert_eq!(device.compute_capability, (8, 9), "Ada alignment gate");
let ctx = GpuCtx::new(&device).expect("create GPU context");
ctx.set_gemm_mode(crate::mamba_ssm::gpu::GemmMode::Deterministic)
.unwrap();
ctx.set_bi_gemm_family(BiGemmFamily::Triad);
ctx.set_f32_triad_policy(F32TriadPolicy::ExactScalarFma);
let dims = (1024, 16, 128);
let a = GpuBuffer::from_cpu(&ctx.stream, &alignment_fixture_values(dims.0 * dims.1, 5))
.expect("allocate A");
let b = GpuBuffer::from_cpu(&ctx.stream, &alignment_fixture_values(dims.1 * dims.2, 17))
.expect("allocate B");
let aligned =
GpuBuffer::zeros(&ctx.stream, dims.0 * dims.2).expect("allocate aligned output");
let shifted = GpuBuffer::zeros(&ctx.stream, 1 + dims.0 * dims.2 + 8)
.expect("allocate shifted output");
ctx.stream
.synchronize()
.expect("finish fixture allocations");
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, dims),
};
let output = shifted.raw_ptr_at(&ctx.stream, 1);
let operands = F32TriadOperands {
output,
a: a.cached_ptr(),
b: b.cached_ptr(),
bias: None,
alpha: 1.0,
beta: 0.0,
};
let prepared = prepare_f32_triad(&ctx, request, operands)
.expect("prepare exact scalar output subview");
let symbols = prepared
.routes
.iter()
.map(|route| route.symbol)
.collect::<Vec<_>>();
assert_eq!(symbols, ["nn_slim"]);
assert!(matches!(
prepared.kind,
PreparedF32Kind::Scalar(ScalarDispatchPlan::NnFinal { slim: true })
));
let route = prepared.routes[0];
let config = cudarc::driver::LaunchConfig {
grid_dim: route.launch.grid_dim,
block_dim: route.launch.block_dim,
shared_mem_bytes: route.launch.shared_mem_bytes,
};
launch_scalar_nn_slim_raw(
&ctx,
aligned.cached_ptr(),
a.cached_ptr(),
b.cached_ptr(),
dims,
config,
)
.expect("launch aligned output reference");
ctx.stream
.synchronize()
.expect("finish aligned output reference");
let expected = aligned
.to_cpu(&ctx.stream)
.expect("download aligned output");
launch_scalar_nn_slim_raw(&ctx, output, a.cached_ptr(), b.cached_ptr(), dims, config)
.expect("launch shifted output");
ctx.stream
.synchronize()
.expect("finish shifted output launch");
let shifted = shifted
.to_cpu(&ctx.stream)
.expect("download shifted output");
assert_eq!(shifted[0].to_bits(), 0.0_f32.to_bits());
assert!(
shifted[1 + dims.0 * dims.2..]
.iter()
.all(|value| value.to_bits() == 0.0_f32.to_bits())
);
assert_eq!(
shifted[1..1 + dims.0 * dims.2]
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>(),
expected
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>()
);
}
#[derive(Clone, Copy, Debug)]
enum ScalarAlignmentOperand {
Output,
A,
B,
Bias,
}
#[derive(Clone, Copy, Debug)]
struct ScalarAlignmentCase {
op: ResolvedGemmOp,
dims: (usize, usize, usize),
slim: bool,
shifted: ScalarAlignmentOperand,
}
#[derive(Clone, Copy)]
struct ScalarAlignmentPointers {
output: CUptr,
a: CUptr,
b: CUptr,
bias: Option<CUptr>,
}
fn scalar_alignment_extents(
op: ResolvedGemmOp,
dims: (usize, usize, usize),
) -> (usize, usize, usize, usize) {
let (m, k, n) = dims;
match op {
ResolvedGemmOp::Nn => (m * k, k * n, m * n, n),
ResolvedGemmOp::Tn => (m * k, m * n, k * n, 0),
ResolvedGemmOp::Nt => (m * n, k * n, m * k, 0),
}
}
fn launch_scalar_alignment_case(
ctx: &GpuCtx,
case: ScalarAlignmentCase,
pointers: ScalarAlignmentPointers,
) -> Result<(), String> {
let (m, k, n) = case.dims;
let m_i = i32::try_from(m).expect("M fits i32");
let k_i = i32::try_from(k).expect("K fits i32");
let n_i = i32::try_from(n).expect("N fits i32");
let alpha = 1.0_f32;
let beta = 0.0_f32;
let bias = pointers.bias.unwrap_or(0);
let bn = if case.slim { 64_u32 } else { 128_u32 };
let threads = if case.slim { 128_u32 } else { 256_u32 };
let (rows, columns) = match case.op {
ResolvedGemmOp::Nn => (m, n),
ResolvedGemmOp::Tn => (k, n),
ResolvedGemmOp::Nt => (m, k),
};
let grid = u32::try_from(rows.div_ceil(128) * columns.div_ceil(bn as usize))
.expect("alignment fixture grid fits u32");
let shared_mem_bytes = match (case.op, case.slim) {
(_, true) => 0,
(ResolvedGemmOp::Nt, false) => SCALAR_BIG_NT_DYNAMIC_SHARED_BYTES,
(_, false) => 34 * 1024,
};
let config = cudarc::driver::LaunchConfig {
grid_dim: (grid, 1, 1),
block_dim: (threads, 1, 1),
shared_mem_bytes,
};
let function = match (case.op, case.slim) {
(ResolvedGemmOp::Nn, false) => &ctx.kernels.gemm_bi_nn,
(ResolvedGemmOp::Nn, true) => &ctx.kernels.gemm_bi_nn_slim,
(ResolvedGemmOp::Tn, false) => &ctx.kernels.gemm_bi_tn,
(ResolvedGemmOp::Tn, true) => &ctx.kernels.gemm_bi_tn_slim,
(ResolvedGemmOp::Nt, false) => &ctx.kernels.gemm_bi_nt,
(ResolvedGemmOp::Nt, true) => &ctx.kernels.gemm_bi_nt_slim,
};
let mut builder = ctx.stream.launch_builder(function);
builder.arg(&pointers.output);
builder.arg(&pointers.a);
builder.arg(&pointers.b);
if case.op == ResolvedGemmOp::Nn {
builder.arg(&bias);
builder.arg(&alpha);
builder.arg(&beta);
builder.arg(&m_i);
builder.arg(&n_i);
builder.arg(&k_i);
builder.arg(&k_i);
builder.arg(&n_i);
builder.arg(&n_i);
} else if case.op == ResolvedGemmOp::Tn {
builder.arg(&alpha);
builder.arg(&m_i);
builder.arg(&k_i);
builder.arg(&n_i);
} else {
builder.arg(&alpha);
builder.arg(&m_i);
builder.arg(&n_i);
builder.arg(&k_i);
}
unsafe { builder.launch(config) }
.map(|_| ())
.map_err(|error| format!("launch scalar alignment case {case:?}: {error:?}"))
}
fn shifted_alignment_storage(values: &[f32], guard: f32) -> Vec<f32> {
let mut storage = Vec::with_capacity(values.len() + 2);
storage.push(guard);
storage.extend_from_slice(values);
storage.push(guard);
storage
}
fn assert_scalar_alignment_case(ctx: &GpuCtx, case: ScalarAlignmentCase) {
let (a_len, b_len, output_len, bias_len) = scalar_alignment_extents(case.op, case.dims);
let a_values = alignment_fixture_values(a_len, 31);
let b_values = alignment_fixture_values(b_len, 37);
let output_values = if case.op == ResolvedGemmOp::Tn {
alignment_fixture_values(output_len, 41)
} else {
vec![0.0; output_len]
};
let bias_values = alignment_fixture_values(bias_len, 43);
let guard = 19.25_f32;
let reference_a =
GpuBuffer::from_cpu(&ctx.stream, &a_values).expect("allocate reference A");
let reference_b =
GpuBuffer::from_cpu(&ctx.stream, &b_values).expect("allocate reference B");
let reference_bias = (!bias_values.is_empty()).then(|| {
GpuBuffer::from_cpu(&ctx.stream, &bias_values).expect("allocate reference bias")
});
let reference_output =
GpuBuffer::from_cpu(&ctx.stream, &output_values).expect("allocate reference output");
let mut actual_a = GpuBuffer::from_cpu(
&ctx.stream,
&if matches!(case.shifted, ScalarAlignmentOperand::A) {
shifted_alignment_storage(&a_values, guard)
} else {
a_values.clone()
},
)
.expect("allocate actual A");
let mut actual_b = GpuBuffer::from_cpu(
&ctx.stream,
&if matches!(case.shifted, ScalarAlignmentOperand::B) {
shifted_alignment_storage(&b_values, guard)
} else {
b_values.clone()
},
)
.expect("allocate actual B");
let mut actual_output = GpuBuffer::from_cpu(
&ctx.stream,
&if matches!(case.shifted, ScalarAlignmentOperand::Output) {
shifted_alignment_storage(&output_values, guard)
} else {
output_values.clone()
},
)
.expect("allocate actual output");
let mut actual_bias = (!bias_values.is_empty()).then(|| {
GpuBuffer::from_cpu(
&ctx.stream,
&if matches!(case.shifted, ScalarAlignmentOperand::Bias) {
shifted_alignment_storage(&bias_values, guard)
} else {
bias_values.clone()
},
)
.expect("allocate actual bias")
});
launch_scalar_alignment_case(
ctx,
case,
ScalarAlignmentPointers {
output: reference_output.cached_ptr(),
a: reference_a.cached_ptr(),
b: reference_b.cached_ptr(),
bias: reference_bias.as_ref().map(GpuBuffer::cached_ptr),
},
)
.expect("launch aligned scalar reference");
ctx.stream.synchronize().expect("finish scalar reference");
let expected = reference_output
.to_cpu(&ctx.stream)
.expect("download scalar reference");
let actual_a_pointer = actual_a.cached_ptr()
+ u64::from(matches!(case.shifted, ScalarAlignmentOperand::A)) * 4;
let actual_b_pointer = actual_b.cached_ptr()
+ u64::from(matches!(case.shifted, ScalarAlignmentOperand::B)) * 4;
let actual_output_pointer = actual_output.cached_ptr()
+ u64::from(matches!(case.shifted, ScalarAlignmentOperand::Output)) * 4;
let actual_bias_pointer = actual_bias.as_ref().map(|bias| {
bias.cached_ptr() + u64::from(matches!(case.shifted, ScalarAlignmentOperand::Bias)) * 4
});
let mut first = None;
for _ in 0..2 {
actual_a
.upload(
&ctx.stream,
&if matches!(case.shifted, ScalarAlignmentOperand::A) {
shifted_alignment_storage(&a_values, guard)
} else {
a_values.clone()
},
)
.expect("restore actual A");
actual_b
.upload(
&ctx.stream,
&if matches!(case.shifted, ScalarAlignmentOperand::B) {
shifted_alignment_storage(&b_values, guard)
} else {
b_values.clone()
},
)
.expect("restore actual B");
actual_output
.upload(
&ctx.stream,
&if matches!(case.shifted, ScalarAlignmentOperand::Output) {
shifted_alignment_storage(&output_values, guard)
} else {
output_values.clone()
},
)
.expect("restore actual output");
if let Some(bias) = actual_bias.as_mut() {
bias.upload(
&ctx.stream,
&if matches!(case.shifted, ScalarAlignmentOperand::Bias) {
shifted_alignment_storage(&bias_values, guard)
} else {
bias_values.clone()
},
)
.expect("restore actual bias");
}
launch_scalar_alignment_case(
ctx,
case,
ScalarAlignmentPointers {
output: actual_output_pointer,
a: actual_a_pointer,
b: actual_b_pointer,
bias: actual_bias_pointer,
},
)
.expect("launch shifted scalar case");
ctx.stream
.synchronize()
.expect("finish shifted scalar case");
let storage = actual_output
.to_cpu(&ctx.stream)
.expect("download shifted scalar output");
let active = if matches!(case.shifted, ScalarAlignmentOperand::Output) {
assert_eq!(storage[0].to_bits(), guard.to_bits(), "{case:?} prefix");
assert_eq!(
storage[output_len + 1].to_bits(),
guard.to_bits(),
"{case:?} suffix"
);
storage[1..output_len + 1].to_vec()
} else {
storage
};
assert_eq!(
active
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>(),
expected
.iter()
.map(|value| value.to_bits())
.collect::<Vec<_>>(),
"{case:?}"
);
if let Some(previous) = first.replace(active) {
assert_eq!(
first.as_ref().expect("second run"),
&previous,
"{case:?} is not deterministic"
);
}
}
}
#[test]
#[ignore = "requires an SM89 CUDA device"]
fn scalar_f32_big_and_slim_accept_every_four_byte_aligned_external_base() {
let device = GpuDevice::new(0).expect("open CUDA device");
assert_eq!(device.compute_capability, (8, 9), "Ada alignment gate");
let ctx = GpuCtx::new(&device).expect("create GPU context");
let dims = (128, 64, 128);
for slim in [false, true] {
for op in [ResolvedGemmOp::Nn, ResolvedGemmOp::Tn, ResolvedGemmOp::Nt] {
let operands: &[ScalarAlignmentOperand] = match op {
ResolvedGemmOp::Nn => &[
ScalarAlignmentOperand::Output,
ScalarAlignmentOperand::A,
ScalarAlignmentOperand::B,
ScalarAlignmentOperand::Bias,
],
ResolvedGemmOp::Tn | ResolvedGemmOp::Nt => &[
ScalarAlignmentOperand::Output,
ScalarAlignmentOperand::A,
ScalarAlignmentOperand::B,
],
};
for &shifted in operands {
assert_scalar_alignment_case(
&ctx,
ScalarAlignmentCase {
op,
dims,
slim,
shifted,
},
);
}
}
}
}
#[test]
fn scalar_physical_plan_freezes_maximum_node_counts_and_neighbors() {
for (op, dims, expected) in [
(ResolvedGemmOp::Nn, (32, 32, 128), 2),
(ResolvedGemmOp::Tn, (256, 128, 128), 2),
(ResolvedGemmOp::Nt, (32, 95, 128), 34),
(ResolvedGemmOp::Nt, (32, 96, 128), 3),
(ResolvedGemmOp::Nt, (31, 95, 128), 1),
(ResolvedGemmOp::Nt, (32, 95, 127), 1),
(ResolvedGemmOp::Nt, (32, 95, 129), 1),
] {
let (request, operands, plan) = scalar_fixture(op, dims);
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(nodes.len(), expected, "{op:?} {dims:?} {plan:?}");
}
}
#[test]
fn scalar_nt_thin_output_cells_have_the_qualified_splitk_node_identity() {
for (dims, expected_grids) in [
((512, 16, 2048), [(64, 1, 1), (1024, 1, 1), (32, 1, 1)]),
((16, 512, 2048), [(64, 16, 1), (512, 1, 1), (32, 1, 1)]),
] {
let (request, operands, plan) = scalar_fixture(ResolvedGemmOp::Nt, dims);
assert_eq!(
plan,
ScalarDispatchPlan::NtSplitKMain {
n_main: 2048,
n_tail: 0,
}
);
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(nodes.len(), 3);
assert_eq!(
nodes.iter().map(|node| node.symbol).collect::<Vec<_>>(),
["transpose_f32_2d", "nn_splitk32_partial", "splitk_reduce",]
);
assert_eq!(
nodes
.iter()
.map(|node| node.launch.grid_dim)
.collect::<Vec<_>>(),
expected_grids
);
for index in 0..nodes.len() {
assert!(
nodes[..index].iter().all(|node| {
node.launch.arguments_digest != nodes[index].launch.arguments_digest
}),
"{dims:?} physical argument digest collision at node {index}"
);
}
assert_ne!(
scalar_route_contract(nodes[1].symbol),
scalar_route_contract(nodes[2].symbol),
"{dims:?} partial and reducer identities"
);
}
}
#[test]
fn scalar_big_nt_uses_its_exact_shared_memory_contract() {
for (op, dims, expected_plan, symbol, shared_mem_bytes) in [
(
ResolvedGemmOp::Nn,
(2048, 128, 1024),
ScalarDispatchPlan::NnFinal { slim: false },
"nn_big",
34 * 1024,
),
(
ResolvedGemmOp::Tn,
(128, 1024, 1024),
ScalarDispatchPlan::TnFinal { slim: false },
"tn_aligned",
34 * 1024,
),
(
ResolvedGemmOp::Nt,
(2048, 1024, 129),
ScalarDispatchPlan::NtFinal { slim: false },
"nt_big",
33_376,
),
(
ResolvedGemmOp::Nt,
(2048, 512, 129),
ScalarDispatchPlan::NtFinal { slim: true },
"nt_slim",
0,
),
] {
let (request, operands, plan) = scalar_fixture(op, dims);
assert_eq!(plan, expected_plan);
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(nodes.len(), 1);
assert_eq!(nodes[0].symbol, symbol);
assert_eq!(nodes[0].launch.shared_mem_bytes, shared_mem_bytes);
}
}
#[test]
fn scalar_nn_m64n64_plan_has_the_qualified_seven_cell_identity() {
let operands = F32TriadOperands {
output: 0x1000,
a: 0x2000,
b: 0x3000,
bias: None,
alpha: 1.0,
beta: 0.0,
};
let mut digests = Vec::new();
for (dims, expected_grid) in [
((2_048, 3_072, 768), 384),
((4_096, 3_072, 1_536), 1_536),
((512, 3_072, 768), 96),
((4_096, 512, 768), 768),
((2_048, 768, 3_072), 1_536),
((2_048, 1_536, 768), 384),
((4_621, 384, 1_928), 2_263),
] {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, dims),
};
let nodes =
scalar_physical_nodes(request, operands, ScalarDispatchPlan::NnM64N64Qualified)
.unwrap();
assert_eq!(nodes.len(), 1);
assert_eq!(nodes[0].symbol, "nn_m64n64_bk16_s2");
assert_eq!(nodes[0].tile, (64, 64));
assert_eq!((nodes[0].bk, nodes[0].stages), (16, 2));
assert_eq!(nodes[0].launch.grid_dim, (expected_grid, 1, 1));
assert_eq!(nodes[0].launch.block_dim, (128, 1, 1));
assert_eq!(nodes[0].launch.shared_mem_bytes, 17_408);
assert_ne!(nodes[0].launch.arguments_digest, [0; 32]);
assert!(!digests.contains(&nodes[0].launch.arguments_digest));
digests.push(nodes[0].launch.arguments_digest);
}
assert_eq!(
scalar_plan_fields(ScalarDispatchPlan::NnM64N64Qualified),
(24, 0, 0)
);
assert_eq!(
scalar_argument_layout(
F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, (2_048, 768, 3_072),),
},
operands,
ScalarDispatchPlan::NnM64N64Qualified,
0,
)
.null_pointer_mask,
0b1000,
);
}
#[test]
fn scalar_nn_fixed_copyplan_plan_has_a_distinct_fixed_physical_identity() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, (2_048, 1_536, 768)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nn, 1.0);
let plan = ScalarDispatchPlan::NnSm89FixedCopyPlanQualified;
assert_eq!(scalar_plan_fields(plan), (37, 0, 0));
assert!(scalar_plan_requires_zero_beta(plan));
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(nodes.len(), 1);
assert_eq!(nodes[0].symbol, "nn_sm89_f32_n64_copyplan");
assert_eq!(nodes[0].tile, (64, 64));
assert_eq!((nodes[0].bk, nodes[0].stages), (32, 2));
assert_eq!(nodes[0].launch.grid_dim, (384, 1, 1));
assert_eq!(nodes[0].launch.block_dim, (128, 1, 1));
assert_eq!(nodes[0].launch.shared_mem_bytes, 0);
assert_eq!(
scalar_route_contract(nodes[0].symbol),
(
PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan,
ResolvedNumericContract::ScalarFma,
ResolvedOutputOwnership::OneCtaPerOutputTile,
)
);
let layout = scalar_argument_layout(request, operands, plan, 0);
assert_eq!(layout.null_pointer_mask, 1 << 3);
}
#[test]
fn scalar_nn_m32n64_splitk32_plan_has_two_exact_ordered_nodes() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, (128, 8_192, 128)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nn, 1.0);
let plan = ScalarDispatchPlan::NnM32N64SplitK32Qualified;
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(nodes.len(), 2);
assert_eq!(
nodes.iter().map(|node| node.symbol).collect::<Vec<_>>(),
["nn_splitk32_m32n64_exact", "splitk_reduce",]
);
assert_eq!(nodes[0].tile, (32, 64));
assert_eq!((nodes[0].bk, nodes[0].stages), (32, 1));
assert_eq!(nodes[0].launch.grid_dim, (2_048, 1, 1));
assert_eq!(nodes[0].launch.block_dim, (128, 1, 1));
assert_eq!(nodes[0].launch.shared_mem_bytes, 0);
assert_eq!(nodes[1].launch.grid_dim, (64, 1, 1));
assert_eq!(nodes[1].launch.block_dim, (256, 1, 1));
assert_eq!(nodes[1].launch.shared_mem_bytes, 0);
assert_ne!(nodes[0].launch.arguments_digest, [0; 32]);
assert_ne!(nodes[1].launch.arguments_digest, [0; 32]);
assert_ne!(
nodes[0].launch.arguments_digest,
nodes[1].launch.arguments_digest
);
assert_eq!(scalar_plan_fields(plan), (36, 0, 0));
assert!(scalar_plan_requires_zero_beta(plan));
assert!(plan.needs_split_scratch());
assert!(!plan.needs_transpose_scratch());
let partial_elements = request.shape.k / 32 * request.shape.m * request.shape.n;
assert_eq!(partial_elements, 4_194_304);
assert!(partial_elements <= SPLITK_SCRATCH_CAP);
assert_eq!(
scalar_route_contract(nodes[0].symbol),
(
PhysicalGemmBackend::ScalarFmaSplitKPartial,
ResolvedNumericContract::ScalarFmaSplitKPartial,
ResolvedOutputOwnership::OneCtaPerOutputTilePerSplitKPartition,
)
);
}
#[test]
fn scalar_nt_d768_transpose_plan_has_two_exact_ordered_nodes() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (2_048, 768, 3_072)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 0.0);
let plan = ScalarDispatchPlan::NtD768TransposeM64N64Qualified;
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(nodes.len(), 2);
assert_eq!(
nodes.iter().map(|node| node.symbol).collect::<Vec<_>>(),
["transpose_f32_32x16_d768", "nn_m64n64_bk16_s2",]
);
assert_eq!(nodes[0].tile, (32, 32));
assert_eq!((nodes[0].bk, nodes[0].stages), (1, 1));
assert_eq!(nodes[0].launch.grid_dim, (96, 24, 1));
assert_eq!(nodes[0].launch.block_dim, (32, 16, 1));
assert_eq!(nodes[0].launch.shared_mem_bytes, 0);
assert_eq!(nodes[1].tile, (64, 64));
assert_eq!((nodes[1].bk, nodes[1].stages), (16, 2));
assert_eq!(nodes[1].launch.grid_dim, (384, 1, 1));
assert_eq!(nodes[1].launch.block_dim, (128, 1, 1));
assert_eq!(
nodes[1].launch.shared_mem_bytes,
crate::mamba_ssm::gpu::gemm_bi_triad::contract::SCALAR_NN_M64N64_DYNAMIC_SHARED_BYTES
);
assert_ne!(nodes[0].launch.arguments_digest, [0; 32]);
assert_ne!(nodes[1].launch.arguments_digest, [0; 32]);
assert_ne!(
nodes[0].launch.arguments_digest,
nodes[1].launch.arguments_digest
);
assert_eq!(scalar_plan_fields(plan), (25, 0, 0));
assert!(scalar_plan_requires_zero_beta(plan));
assert_eq!(
scalar_argument_layout(request, operands, plan, 0),
ScalarArgumentLayout::default()
);
assert_eq!(
scalar_argument_layout(request, operands, plan, 1).null_pointer_mask,
0b1000
);
}
#[test]
fn scalar_nt_d768_out_transpose_plan_has_distinct_tag_and_exact_nodes() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (2_048, 1_536, 768)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 0.0);
let plan = ScalarDispatchPlan::NtD768OutTransposeM64N64Qualified;
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(nodes.len(), 2);
assert_eq!(
nodes.iter().map(|node| node.symbol).collect::<Vec<_>>(),
["transpose_f32_32x16_d768", "nn_m64n64_bk16_s2",]
);
assert_eq!(nodes[0].launch.grid_dim, (24, 48, 1));
assert_eq!(nodes[0].launch.block_dim, (32, 16, 1));
assert_eq!(nodes[0].launch.shared_mem_bytes, 0);
assert_eq!(nodes[1].launch.grid_dim, (768, 1, 1));
assert_eq!(nodes[1].launch.block_dim, (128, 1, 1));
assert_eq!(
nodes[1].launch.shared_mem_bytes,
crate::mamba_ssm::gpu::gemm_bi_triad::contract::SCALAR_NN_M64N64_DYNAMIC_SHARED_BYTES
);
assert_ne!(nodes[0].launch.arguments_digest, [0; 32]);
assert_ne!(nodes[1].launch.arguments_digest, [0; 32]);
assert_ne!(
nodes[0].launch.arguments_digest,
nodes[1].launch.arguments_digest
);
assert_eq!(scalar_plan_fields(plan), (26, 0, 0));
assert!(scalar_plan_requires_zero_beta(plan));
assert_eq!(
scalar_argument_layout(request, operands, plan, 0),
ScalarArgumentLayout::default()
);
assert_eq!(
scalar_argument_layout(request, operands, plan, 1).null_pointer_mask,
0b1000
);
}
#[test]
fn scalar_nt_large_deep_transpose_plan_has_the_qualified_revision_and_two_exact_nodes() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (4_096, 3_072, 1_536)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
let plan = ScalarDispatchPlan::NtLargeDeepTransposeM64N64Qualified;
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(nodes.len(), 2);
assert_eq!(
nodes.iter().map(|node| node.symbol).collect::<Vec<_>>(),
["transpose_f32_32x16_d768", "nn_m64n64_bk16_s2",]
);
assert_eq!(nodes[0].tile, (32, 32));
assert_eq!((nodes[0].bk, nodes[0].stages), (1, 1));
assert_eq!(nodes[0].launch.grid_dim, (48, 96, 1));
assert_eq!(nodes[0].launch.block_dim, (32, 16, 1));
assert_eq!(nodes[0].launch.shared_mem_bytes, 0);
assert_eq!(nodes[1].tile, (64, 64));
assert_eq!((nodes[1].bk, nodes[1].stages), (16, 2));
assert_eq!(nodes[1].launch.grid_dim, (3_072, 1, 1));
assert_eq!(nodes[1].launch.block_dim, (128, 1, 1));
assert_eq!(
nodes[1].launch.shared_mem_bytes,
crate::mamba_ssm::gpu::gemm_bi_triad::contract::SCALAR_NN_M64N64_DYNAMIC_SHARED_BYTES
);
assert_ne!(nodes[0].launch.arguments_digest, [0; 32]);
assert_ne!(nodes[1].launch.arguments_digest, [0; 32]);
assert_ne!(
nodes[0].launch.arguments_digest,
nodes[1].launch.arguments_digest
);
assert_eq!(scalar_plan_fields(plan), (27, 0, 0));
assert!(scalar_plan_requires_zero_beta(plan));
assert_eq!(scalar_node_count(plan), 2);
assert_eq!(
scalar_argument_layout(request, operands, plan, 0),
ScalarArgumentLayout::default()
);
assert_eq!(
scalar_argument_layout(request, operands, plan, 1).null_pointer_mask,
0b1000
);
}
#[test]
fn scalar_nt_prism_vector_plan_has_the_qualified_revision_and_two_exact_nodes() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (4_621, 384, 1_928)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
let plan = ScalarDispatchPlan::NtPrismVectorQualified;
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(nodes.len(), 2);
assert_eq!(
nodes.iter().map(|node| node.symbol).collect::<Vec<_>>(),
["transpose_f32_32x16_d768", "nn_prism_m64n64_bk16_s2",]
);
assert_eq!(nodes[0].tile, (32, 32));
assert_eq!((nodes[0].bk, nodes[0].stages), (1, 1));
assert_eq!(nodes[0].launch.grid_dim, (61, 12, 1));
assert_eq!(nodes[0].launch.block_dim, (32, 16, 1));
assert_eq!(nodes[0].launch.shared_mem_bytes, 0);
assert_eq!(nodes[1].tile, (64, 64));
assert_eq!((nodes[1].bk, nodes[1].stages), (16, 2));
assert_eq!(nodes[1].launch.grid_dim, (438, 1, 1));
assert_eq!(nodes[1].launch.block_dim, (128, 1, 1));
assert_eq!(
nodes[1].launch.shared_mem_bytes,
crate::mamba_ssm::gpu::gemm_bi_triad::contract::SCALAR_NN_M64N64_DYNAMIC_SHARED_BYTES
);
assert_ne!(nodes[0].launch.arguments_digest, [0; 32]);
assert_ne!(nodes[1].launch.arguments_digest, [0; 32]);
assert_ne!(
nodes[0].launch.arguments_digest,
nodes[1].launch.arguments_digest
);
assert_eq!(scalar_plan_fields(plan), (30, 0, 0));
assert!(scalar_plan_requires_zero_beta(plan));
assert_eq!(scalar_node_count(plan), 2);
assert_eq!(
scalar_argument_layout(request, operands, plan, 0),
ScalarArgumentLayout::default()
);
assert_eq!(
scalar_argument_layout(request, operands, plan, 1).null_pointer_mask,
0b1000
);
assert_eq!(
scalar_transpose_scratch_elements(request, plan).unwrap(),
Some(740_352)
);
assert_eq!(
crate::mamba_ssm::gpu::gemm_bi_triad::contract::SCALAR_TRANSPOSE_SCRATCH_CAP_ELEMENTS,
12_582_912
);
}
#[test]
fn scalar_nt_d128_out_transpose_plan_has_the_qualified_revision_and_two_exact_nodes() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (1_024, 256, 128)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
let plan = ScalarDispatchPlan::NtD128OutTransposeM64N64Qualified;
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(nodes.len(), 2);
assert_eq!(
nodes.iter().map(|node| node.symbol).collect::<Vec<_>>(),
["transpose_f32_32x16_d768", "nn_m64n64_bk16_s2",]
);
assert_eq!(
(nodes[0].tile, nodes[0].bk, nodes[0].stages),
((32, 32), 1, 1)
);
assert_eq!(nodes[0].launch.grid_dim, (4, 8, 1));
assert_eq!(nodes[0].launch.block_dim, (32, 16, 1));
assert_eq!(nodes[0].launch.shared_mem_bytes, 0);
assert_eq!(
(nodes[1].tile, nodes[1].bk, nodes[1].stages),
((64, 64), 16, 2)
);
assert_eq!(nodes[1].launch.grid_dim, (64, 1, 1));
assert_eq!(nodes[1].launch.block_dim, (128, 1, 1));
assert_eq!(
nodes[1].launch.shared_mem_bytes,
crate::mamba_ssm::gpu::gemm_bi_triad::contract::SCALAR_NN_M64N64_DYNAMIC_SHARED_BYTES
);
assert_ne!(nodes[0].launch.arguments_digest, [0; 32]);
assert_ne!(nodes[1].launch.arguments_digest, [0; 32]);
assert_ne!(
nodes[0].launch.arguments_digest,
nodes[1].launch.arguments_digest
);
assert_eq!(scalar_plan_fields(plan), (29, 0, 0));
assert!(scalar_plan_requires_zero_beta(plan));
assert_eq!(scalar_node_count(plan), 2);
assert_eq!(
scalar_argument_layout(request, operands, plan, 0),
ScalarArgumentLayout::default()
);
assert_eq!(
scalar_argument_layout(request, operands, plan, 1).null_pointer_mask,
0b1000
);
assert_eq!(
scalar_transpose_scratch_elements(request, plan).unwrap(),
Some(32_768)
);
}
#[test]
fn prepared_scalar_binding_uses_frozen_digest_and_rejects_symbol_config_or_order_changes() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (1_024, 256, 128)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
let nodes = scalar_physical_nodes(
request,
operands,
ScalarDispatchPlan::NtD128OutTransposeM64N64Qualified,
)
.unwrap();
assert_eq!(nodes.len(), 2);
for node in &nodes {
let config = cudarc::driver::LaunchConfig {
grid_dim: node.launch.grid_dim,
block_dim: node.launch.block_dim,
shared_mem_bytes: node.launch.shared_mem_bytes,
};
assert_eq!(
validate_prepared_scalar_binding(node.symbol, node.launch, node.symbol, config)
.unwrap(),
node.launch.arguments_digest
);
assert_ne!(node.launch.arguments_digest, [0; 32]);
for wrong_config in [
cudarc::driver::LaunchConfig {
grid_dim: (config.grid_dim.0 + 1, config.grid_dim.1, config.grid_dim.2),
..config
},
cudarc::driver::LaunchConfig {
block_dim: (
config.block_dim.0 + 1,
config.block_dim.1,
config.block_dim.2,
),
..config
},
cudarc::driver::LaunchConfig {
shared_mem_bytes: config.shared_mem_bytes + 4,
..config
},
] {
assert!(
validate_prepared_scalar_binding(
node.symbol,
node.launch,
node.symbol,
wrong_config,
)
.is_err()
);
}
}
let first_config = cudarc::driver::LaunchConfig {
grid_dim: nodes[0].launch.grid_dim,
block_dim: nodes[0].launch.block_dim,
shared_mem_bytes: nodes[0].launch.shared_mem_bytes,
};
assert!(
validate_prepared_scalar_binding(
nodes[0].symbol,
nodes[0].launch,
nodes[1].symbol,
first_config,
)
.is_err()
);
let mut zero_digest = nodes[0].launch;
zero_digest.arguments_digest = [0; 32];
assert!(
validate_prepared_scalar_binding(
nodes[0].symbol,
zero_digest,
nodes[0].symbol,
first_config,
)
.is_err()
);
}
#[test]
fn prepared_scalar_enqueue_does_not_rehash_frozen_arguments() {
let source = include_str!("launch.rs");
for (start, end) in [
(
"fn enqueue_scalar_forward",
"fn gemm_bi_forward_sub_with_control",
),
("fn enqueue_scalar_backward", "/// Weight gradient"),
] {
let body = source
.split_once(start)
.and_then(|(_, tail)| tail.split_once(end).map(|(body, _)| body))
.expect("controlled scalar enqueue body");
assert!(!body.contains("scalar_arguments_digest"));
assert!(body.contains("control.enqueue(symbol, config, builder)"));
}
}
#[test]
fn prepared_cache_lookup_does_not_build_full_route_identity() {
let source = include_str!("launch.rs");
let body = source
.split_once("fn ensure_prepared(")
.and_then(|(_, tail)| tail.split_once("fn launch<Scalar>(").map(|(body, _)| body))
.expect("prepared cache lookup body");
assert!(!body.contains("ctx.gemm_route()"));
assert!(body.contains("ctx.instance_token()"));
assert!(body.contains("ctx.gemm_policy()"));
}
#[test]
fn prepared_backward_enqueue_skips_fallback_dispatch() {
let source = include_str!("launch.rs");
for (start, end) in [
("fn gemm_bi_backward_dw_with_control", "/// Input gradient"),
(
"fn gemm_bi_backward_dx_with_control",
"/// Typed input gradient",
),
] {
let body = source
.split_once(start)
.and_then(|(_, tail)| tail.split_once(end).map(|(body, _)| body))
.expect("prepared backward enqueue body");
assert!(!body.contains("scalar_backward_launch_plan"));
assert!(body.contains("scalar_backward_request"));
assert!(body.contains("Some(prepared) => prepared.plan()"));
}
}
#[test]
fn scalar_nt_d128_out_raw_and_prepared_identity_match() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (1_024, 256, 128)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
let raw_plan = scalar_launch_plan(scalar_tn_admission_facts(), request, operands).unwrap();
let prepared_plan =
scalar_launch_plan(scalar_tn_admission_facts(), request, operands).unwrap();
assert_eq!(
raw_plan,
ScalarDispatchPlan::NtD128OutTransposeM64N64Qualified
);
assert_eq!(prepared_plan, raw_plan);
assert_eq!(
scalar_physical_nodes(request, operands, prepared_plan).unwrap(),
scalar_physical_nodes(request, operands, raw_plan).unwrap()
);
}
#[test]
fn scalar_tn_m16n16_raw_and_prepared_one_node_identity_match() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Tn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Tn, (256, 512, 384)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Tn, 1.0);
let mut ada = super::super::dispatch::scalar_sm89_composed_test_facts(2);
ada.fixed_copyplan_loaded = false;
for (environment, facts) in [("SM120", scalar_tn_admission_facts()), ("Ada", ada)] {
let raw_plan = scalar_launch_plan(facts, request, operands).unwrap();
let prepared_plan = scalar_launch_plan(facts, request, operands).unwrap();
assert_eq!(
raw_plan,
ScalarDispatchPlan::TnM16N16SplitM16Qualified,
"{environment}"
);
assert_eq!(prepared_plan, raw_plan, "{environment}");
let raw_nodes = scalar_physical_nodes(request, operands, raw_plan).unwrap();
let prepared_nodes = scalar_physical_nodes(request, operands, prepared_plan).unwrap();
assert_eq!(prepared_nodes, raw_nodes, "{environment}");
assert_eq!(raw_nodes.len(), scalar_node_count(raw_plan));
assert_eq!(raw_nodes.len(), 1);
let node = raw_nodes[0];
assert_eq!(node.symbol, "tn_m16n16_bk16_s2_splitm16");
assert_eq!(node.tile, (16, 16));
assert_eq!((node.bk, node.stages), (16, 2));
assert_eq!(node.launch.grid_dim, (768, 1, 1));
assert_eq!(node.launch.block_dim, (64, 1, 1));
assert_eq!(node.launch.shared_mem_bytes, 4_096);
assert_ne!(node.launch.arguments_digest, [0; 32]);
assert_eq!(scalar_plan_fields(raw_plan), (35, 0, 0));
assert_eq!(
scalar_argument_layout(request, operands, raw_plan, 0),
ScalarArgumentLayout::default()
);
assert_eq!(
scalar_route_contract(node.symbol),
(
PhysicalGemmBackend::ScalarFmaTnSplitMF64Reduce,
ResolvedNumericContract::ScalarFmaTnSplitMF64Reduce,
ResolvedOutputOwnership::OneCtaPerOutputTile,
)
);
assert!(!raw_plan.needs_transpose_scratch());
assert!(!raw_plan.needs_split_scratch());
}
}
#[test]
fn scalar_nn_m32n64_splitk32_raw_and_prepared_identity_match() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, (128, 8_192, 128)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nn, 1.0);
let raw_plan = scalar_launch_plan(scalar_tn_admission_facts(), request, operands).unwrap();
let prepared_plan =
scalar_launch_plan(scalar_tn_admission_facts(), request, operands).unwrap();
assert_eq!(raw_plan, ScalarDispatchPlan::NnM32N64SplitK32Qualified);
assert_eq!(prepared_plan, raw_plan);
let raw_nodes = scalar_physical_nodes(request, operands, raw_plan).unwrap();
let prepared_nodes = scalar_physical_nodes(request, operands, prepared_plan).unwrap();
assert_eq!(prepared_nodes, raw_nodes);
assert_eq!(raw_nodes.len(), scalar_node_count(raw_plan));
assert_eq!(raw_nodes.len(), 2);
assert_eq!(
raw_nodes.iter().map(|node| node.symbol).collect::<Vec<_>>(),
["nn_splitk32_m32n64_exact", "splitk_reduce",]
);
assert!(
raw_nodes
.iter()
.all(|node| node.launch.arguments_digest != [0; 32])
);
assert_ne!(
raw_nodes[0].launch.arguments_digest,
raw_nodes[1].launch.arguments_digest
);
}
#[test]
fn scalar_nt_m2n16_raw_and_prepared_one_node_identity_match() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (512, 16, 2_048)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
let raw_plan = scalar_launch_plan(scalar_tn_admission_facts(), request, operands).unwrap();
let prepared_plan =
scalar_launch_plan(scalar_tn_admission_facts(), request, operands).unwrap();
assert_eq!(raw_plan, ScalarDispatchPlan::NtM2N16SplitK32Qualified);
assert_eq!(prepared_plan, raw_plan);
let raw_nodes = scalar_physical_nodes(request, operands, raw_plan).unwrap();
let prepared_nodes = scalar_physical_nodes(request, operands, prepared_plan).unwrap();
assert_eq!(prepared_nodes, raw_nodes);
assert_eq!(raw_nodes.len(), scalar_node_count(raw_plan));
assert_eq!(raw_nodes.len(), 1);
let node = raw_nodes[0];
assert_eq!(node.symbol, "nt_m2n16_bk64_splitk32");
assert_eq!(node.tile, (2, 16));
assert_eq!((node.bk, node.stages), (64, 2));
assert_eq!(node.launch.grid_dim, (256, 1, 1));
assert_eq!(node.launch.block_dim, (64, 1, 1));
assert_eq!(node.launch.shared_mem_bytes, 17_984);
assert_ne!(node.launch.arguments_digest, [0; 32]);
assert_eq!(scalar_plan_fields(raw_plan), (31, 0, 0));
assert!(!raw_plan.needs_transpose_scratch());
assert!(!raw_plan.needs_split_scratch());
}
#[test]
fn scalar_nt_prism_raw_and_prepared_identity_match() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (4_621, 384, 1_928)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
let raw_plan = scalar_launch_plan(scalar_tn_admission_facts(), request, operands).unwrap();
let prepared_plan =
scalar_launch_plan(scalar_tn_admission_facts(), request, operands).unwrap();
assert_eq!(raw_plan, ScalarDispatchPlan::NtPrismVectorQualified);
assert_eq!(prepared_plan, raw_plan);
assert_eq!(
scalar_physical_nodes(request, operands, prepared_plan).unwrap(),
scalar_physical_nodes(request, operands, raw_plan).unwrap()
);
}
#[test]
fn scalar_nt_large_deep_raw_and_prepared_identity_match() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (4_096, 3_072, 1_536)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
let raw_plan = scalar_launch_plan(scalar_tn_admission_facts(), request, operands).unwrap();
let prepared_plan =
scalar_launch_plan(scalar_tn_admission_facts(), request, operands).unwrap();
assert_eq!(
raw_plan,
ScalarDispatchPlan::NtLargeDeepTransposeM64N64Qualified
);
assert_eq!(prepared_plan, raw_plan);
assert_eq!(
scalar_physical_nodes(request, operands, prepared_plan).unwrap(),
scalar_physical_nodes(request, operands, raw_plan).unwrap()
);
let nodes = scalar_physical_nodes(request, operands, raw_plan).unwrap();
assert_eq!(nodes.len(), 2);
assert_eq!(nodes[1].symbol, "nn_m64n64_bk16_s2");
assert_eq!((nodes[1].bk, nodes[1].stages), (16, 2));
}
#[test]
fn triad_retained_scalar_ada_large_deep_freezes_two_fixed_owned_nodes() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (4_096, 3_072, 1_536)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
let facts = super::super::dispatch::scalar_sm89_composed_test_facts(2);
let raw_plan = scalar_launch_plan(facts, request, operands).unwrap();
let prepared_plan = scalar_launch_plan(facts, request, operands).unwrap();
assert_eq!(
raw_plan,
ScalarDispatchPlan::NtLargeDeepSm89FixedCopyPlanQualified
);
assert_eq!(prepared_plan, raw_plan);
assert_eq!(scalar_plan_fields(raw_plan), (46, 0, 0));
assert!(scalar_plan_requires_zero_beta(raw_plan));
assert!(raw_plan.needs_transpose_scratch());
assert!(!raw_plan.needs_split_scratch());
assert_eq!(
scalar_transpose_scratch_elements(request, raw_plan).unwrap(),
Some(4_718_592)
);
let raw_nodes = scalar_physical_nodes(request, operands, raw_plan).unwrap();
let prepared_nodes = scalar_physical_nodes(request, operands, prepared_plan).unwrap();
assert_eq!(prepared_nodes, raw_nodes);
assert_eq!(raw_nodes.len(), scalar_node_count(raw_plan));
assert_eq!(raw_nodes.len(), 2);
assert_eq!(raw_nodes[0].symbol, "transpose_f32_32x16_d768");
assert_eq!(raw_nodes[0].launch.grid_dim, (48, 96, 1));
assert_eq!(raw_nodes[0].launch.block_dim, (32, 16, 1));
assert_eq!(raw_nodes[0].launch.shared_mem_bytes, 0);
assert_eq!(raw_nodes[1].symbol, "nn_sm89_f32_n64_copyplan");
assert_eq!(raw_nodes[1].launch.grid_dim, (3_072, 1, 1));
assert_eq!(raw_nodes[1].launch.block_dim, (128, 1, 1));
assert_eq!(raw_nodes[1].launch.shared_mem_bytes, 0);
assert_eq!(
(raw_nodes[1].tile, raw_nodes[1].bk, raw_nodes[1].stages),
((64, 64), 32, 2)
);
assert_eq!(
scalar_route_contract(raw_nodes[0].symbol),
(
PhysicalGemmBackend::ScalarFma,
ResolvedNumericContract::ScalarFma,
ResolvedOutputOwnership::OneCtaPerOutputTile,
)
);
assert_eq!(
scalar_route_contract(raw_nodes[1].symbol),
(
PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan,
ResolvedNumericContract::ScalarFma,
ResolvedOutputOwnership::OneCtaPerOutputTile,
)
);
assert!(
raw_nodes
.iter()
.all(|node| node.launch.arguments_digest != [0; 32])
);
assert_ne!(
raw_nodes[0].launch.arguments_digest,
raw_nodes[1].launch.arguments_digest
);
}
#[test]
fn scalar_transpose_workspace_extent_is_checked_per_plan() {
let large_deep = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (4_096, 3_072, 1_536)),
};
assert_eq!(
scalar_transpose_scratch_elements(
large_deep,
ScalarDispatchPlan::NtLargeDeepTransposeM64N64Qualified,
)
.unwrap(),
Some(4_718_592)
);
let oversized = F32TriadRequest {
shape: F32TriadShape {
k: large_deep.shape.k * 3,
..large_deep.shape
},
..large_deep
};
assert!(
scalar_transpose_scratch_elements(
oversized,
ScalarDispatchPlan::NtLargeDeepTransposeM64N64Qualified,
)
.is_err()
);
assert_eq!(
scalar_transpose_scratch_elements(
large_deep,
ScalarDispatchPlan::NtFinal { slim: false },
)
.unwrap(),
None
);
}
#[test]
fn scalar_nt_d768_out_raw_and_prepared_paths_freeze_the_same_identity() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (2_048, 1_536, 768)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
let raw_plan = scalar_launch_plan(scalar_tn_admission_facts(), request, operands).unwrap();
let prepared_plan =
scalar_launch_plan(scalar_tn_admission_facts(), request, operands).unwrap();
assert_eq!(
raw_plan,
ScalarDispatchPlan::NtD768OutTransposeM64N64Qualified
);
assert_eq!(prepared_plan, raw_plan);
let raw_nodes = scalar_physical_nodes(request, operands, raw_plan).unwrap();
let prepared_nodes = scalar_physical_nodes(request, operands, prepared_plan).unwrap();
assert_eq!(prepared_nodes, raw_nodes);
assert_eq!(raw_nodes.len(), scalar_node_count(raw_plan));
assert_eq!(raw_nodes.len(), 2);
assert_eq!(scalar_plan_fields(raw_plan), (26, 0, 0));
assert!(raw_plan.needs_transpose_scratch());
assert!(!raw_plan.needs_split_scratch());
}
#[test]
fn scalar_nt_d768_out_fixed_copyplan_preserves_outer_and_inner_geometry() {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (2_048, 1_536, 768)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
let plan = ScalarDispatchPlan::NtD768OutSm89FixedCopyPlanQualified;
assert_eq!(
(request.shape.m, request.shape.k, request.shape.n),
(2_048, 1_536, 768),
"frozen outer NT evidence dimensions"
);
assert_eq!(
(request.shape.m, request.shape.n, request.shape.k),
(2_048, 768, 1_536),
"inner NN M/K/N after the physical transpose"
);
assert_eq!((request.shape.m, request.shape.k), (2_048, 1_536));
assert_eq!(scalar_plan_fields(plan), (38, 0, 0));
assert!(scalar_plan_requires_zero_beta(plan));
assert_eq!(
scalar_transpose_scratch_elements(request, plan).unwrap(),
Some(1_179_648)
);
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(nodes.len(), 2);
assert_eq!(nodes[0].symbol, "transpose_f32_32x16_d768");
assert_eq!(nodes[0].launch.grid_dim, (24, 48, 1));
assert_eq!(nodes[0].launch.block_dim, (32, 16, 1));
assert_eq!(nodes[1].symbol, "nn_sm89_f32_n64_copyplan");
assert_eq!(nodes[1].launch.grid_dim, (768, 1, 1));
assert_eq!(nodes[1].launch.block_dim, (128, 1, 1));
assert_eq!(nodes[1].launch.shared_mem_bytes, 0);
assert_eq!(
(nodes[1].tile, nodes[1].bk, nodes[1].stages),
((64, 64), 32, 2)
);
assert_eq!(
scalar_route_contract(nodes[0].symbol).0,
PhysicalGemmBackend::ScalarFma
);
assert_eq!(
scalar_route_contract(nodes[1].symbol).0,
PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan
);
assert_ne!(
nodes[0].launch.arguments_digest,
nodes[1].launch.arguments_digest
);
}
#[test]
fn scalar_nt_fixed_copyplan_siblings_preserve_two_node_identity_and_geometry() {
for (dims, plan, tag, scratch, transpose_grid, fixed_grid) in [
(
(2_048, 768, 3_072),
ScalarDispatchPlan::NtD768InSm89FixedCopyPlanQualified,
39,
2_359_296,
(96, 24, 1),
(384, 1, 1),
),
(
(4_621, 384, 1_928),
ScalarDispatchPlan::NtPrismSm89FixedCopyPlanQualified,
40,
740_352,
(61, 12, 1),
(438, 1, 1),
),
] {
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, dims),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
assert_eq!(scalar_plan_fields(plan), (tag, 0, 0));
assert!(scalar_plan_requires_zero_beta(plan));
assert!(plan.needs_transpose_scratch());
assert!(!plan.needs_split_scratch());
assert_eq!(
scalar_transpose_scratch_elements(request, plan).unwrap(),
Some(scratch)
);
let raw_nodes = scalar_physical_nodes(request, operands, plan).unwrap();
let prepared_nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(prepared_nodes, raw_nodes, "raw/prepared plan {plan:?}");
assert_eq!(raw_nodes.len(), 2);
assert_eq!(raw_nodes[0].symbol, "transpose_f32_32x16_d768");
assert_eq!(raw_nodes[0].launch.grid_dim, transpose_grid);
assert_eq!(raw_nodes[0].launch.block_dim, (32, 16, 1));
assert_eq!(raw_nodes[0].launch.shared_mem_bytes, 0);
assert_eq!(raw_nodes[1].symbol, "nn_sm89_f32_n64_copyplan");
assert_eq!(raw_nodes[1].launch.grid_dim, fixed_grid);
assert_eq!(raw_nodes[1].launch.block_dim, (128, 1, 1));
assert_eq!(raw_nodes[1].launch.shared_mem_bytes, 0);
assert_eq!(
(raw_nodes[1].tile, raw_nodes[1].bk, raw_nodes[1].stages),
((64, 64), 32, 2)
);
assert_eq!(
scalar_route_contract(raw_nodes[0].symbol).0,
PhysicalGemmBackend::ScalarFma
);
assert_eq!(
scalar_route_contract(raw_nodes[1].symbol).0,
PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan
);
assert_ne!(
raw_nodes[0].launch.arguments_digest,
raw_nodes[1].launch.arguments_digest
);
}
}
#[test]
fn scalar_node_identity_distinguishes_equal_symbol_sequences() {
let left = scalar_fixture(ResolvedGemmOp::Nn, (32, 64, 128));
let right = scalar_fixture(ResolvedGemmOp::Nn, (32, 96, 128));
let left_nodes = scalar_physical_nodes(left.0, left.1, left.2).unwrap();
let right_nodes = scalar_physical_nodes(right.0, right.1, right.2).unwrap();
let left_symbols: Vec<_> = left_nodes.iter().map(|node| node.symbol).collect();
let right_symbols: Vec<_> = right_nodes.iter().map(|node| node.symbol).collect();
assert_eq!(left_symbols, right_symbols);
assert_ne!(left_nodes, right_nodes);
assert_ne!(
left_nodes[0].launch.arguments_digest,
right_nodes[0].launch.arguments_digest
);
}
#[test]
fn scalar_splitk_symbols_have_stage_specific_route_contracts() {
let partial = (
PhysicalGemmBackend::ScalarFmaSplitKPartial,
ResolvedNumericContract::ScalarFmaSplitKPartial,
ResolvedOutputOwnership::OneCtaPerOutputTilePerSplitKPartition,
);
assert_eq!(partial.0 as u8, 17);
assert_eq!(partial.1 as u8, 16);
assert_eq!(partial.2 as u8, 9);
for symbol in ["nn_splitk32_partial", "nn_splitk_slim_partial"] {
assert_eq!(scalar_route_contract(symbol), partial, "{symbol}");
}
let reducer = (
PhysicalGemmBackend::ScalarFmaSplitKF32Reduce,
ResolvedNumericContract::ScalarFmaSplitKF32Reduce,
ResolvedOutputOwnership::OneThreadPerOutputElementFixedSplitKReduce,
);
assert_eq!(reducer.0 as u8, 18);
assert_eq!(reducer.1 as u8, 17);
assert_eq!(reducer.2 as u8, 10);
assert_eq!(scalar_route_contract("splitk_reduce"), reducer);
assert_ne!(partial, reducer);
assert_ne!(partial, scalar_route_contract("nn_splitk32_partial_typo"));
}
#[test]
fn scalar_argument_layout_tracks_the_physical_pointer_abi() {
let (nn_request, mut nn_operands, nn_plan) =
scalar_fixture(ResolvedGemmOp::Nn, (32, 32, 128));
nn_operands.bias = None;
assert!(scalar_plan_requires_zero_beta(nn_plan));
assert_eq!(
scalar_argument_layout(nn_request, nn_operands, nn_plan, 0),
ScalarArgumentLayout::default()
);
assert_eq!(
scalar_argument_layout(nn_request, nn_operands, nn_plan, 1).null_pointer_mask,
0b11100
);
let (tn_request, tn_operands, tn_plan) =
scalar_fixture(ResolvedGemmOp::Tn, (256, 128, 128));
assert!(!scalar_plan_requires_zero_beta(tn_plan));
for index in 0..scalar_node_count(tn_plan) {
assert_eq!(
scalar_argument_layout(tn_request, tn_operands, tn_plan, index),
ScalarArgumentLayout::default()
);
}
let (nt_request, nt_operands, nt_plan) = scalar_fixture(ResolvedGemmOp::Nt, (32, 95, 128));
assert_eq!(
scalar_argument_layout(nt_request, nt_operands, nt_plan, 2).null_pointer_mask,
0b11100
);
let tail = scalar_argument_layout(nt_request, nt_operands, nt_plan, 3);
assert_eq!(tail.output_offset, 0);
assert_eq!(tail.output_column, Some(64));
assert_eq!(tail.b_offset, 64 * 128 * 4);
assert_eq!(tail.null_pointer_mask, 0);
}
fn portable_binding() -> Tf32MapBinding {
let target = CudaTarget::new("sm_89").unwrap();
let nvrtc_version = (13, 2);
Tf32MapBinding {
allocation_domain: AllocationDomain {
context_handle: 7,
device_ordinal: 0,
},
qualified: Tf32QualifiedModule {
module_kind: ModuleKind::TriadSm80,
target,
artifact: ArtifactIdentity {
module_kind: ModuleKind::TriadSm80,
artifact_kind: ArtifactKind::Ptx,
compile_key: [4; 32],
artifact_digest: [2; 32],
},
compiler: CompilerIdentity {
source_digest: [3; 32],
invocation_digest: [4; 32],
header_manifest_digest: [5; 32],
target,
nvrtc_version,
nvrtc_library_domain: [6; 32],
nvrtc_library_known: true,
output_kind: ArtifactKind::Ptx,
composer_revision: COMPOSER_REVISION,
compiler_revision: COMPILER_REVISION,
numeric_abi_revision: NUMERIC_ABI_REVISION,
schedule_revision: SCHEDULE_REVISION,
},
device: DeviceIdentity {
compute_capability: (8, 9),
multiprocessor_count: 142,
target,
driver: DriverIdentity {
api_version: 13_020,
build_sources: 1,
build_digest: [7; 32],
},
},
device_caps: DeviceCaps {
compute_capability: (8, 9),
nvrtc_version,
accepted_target: Some(target),
optin_shared_bytes: 99_000,
tensor_map_access: false,
},
sm120_fma_exclusions: Default::default(),
},
}
}
#[test]
fn triad_retained_tf32_prism_nt_freezes_one_rna_mma_node() {
let route = Tf32PhysicalRoute::Sm89MmaTf32Compact8;
let spec = tf32_kernel_spec(ResolvedGemmOp::Nt, route).unwrap();
let request = F32TriadRequest {
op: ResolvedGemmOp::Nt,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nt, (4_621, 384, 1_928)),
};
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
let mut binding = portable_binding();
binding.qualified.module_kind = ModuleKind::TriadSm89Finalist;
binding.qualified.artifact.module_kind = ModuleKind::TriadSm89Finalist;
let config = cudarc::driver::LaunchConfig {
grid_dim: (222, 1, 1),
block_dim: (256, 1, 1),
shared_mem_bytes: 49_152,
};
let arguments = tf32_kernel_arguments_digest(request, operands, spec.symbol, [0; 32]);
let resolved = tf32_resolved_route(
request,
spec,
binding,
Tf32LaunchDigests {
maps: [0; 32],
resources: [2; 32],
arguments,
},
false,
config,
);
assert_eq!(resolved.symbol, "nt_sm89_mma_tf32_compact8_m128n64_bk32_s2");
assert_eq!(resolved.module_kind, ModuleKind::TriadSm89Finalist);
assert_eq!(resolved.backend, PhysicalGemmBackend::Sm89MmaTf32Compact8);
assert_eq!(
resolved.numeric_contract,
ResolvedNumericContract::MmaTf32Rna
);
assert_eq!(
resolved.instruction_family,
ResolvedInstructionFamily::MmaSync
);
assert_eq!(
resolved.instruction_shape,
ResolvedInstructionShape { m: 16, n: 8, k: 8 }
);
assert_eq!(
resolved.operand_conversion,
ResolvedOperandConversion::RegisterCvtRnaTf32F32
);
assert_eq!(
resolved.ownership,
ResolvedOutputOwnership::OneCtaPerOutputTile
);
assert_eq!(resolved.shape, (4_621, 384, 1_928));
assert_eq!(resolved.strides, (1_928, 1_928, 384));
assert_eq!(
(resolved.tile, resolved.bk, resolved.stages),
((128, 64), 32, 2)
);
assert_eq!(resolved.threads, 256);
assert_eq!(resolved.launch.grid_dim, (222, 1, 1));
assert_eq!(resolved.launch.block_dim, (256, 1, 1));
assert_eq!(resolved.launch.shared_mem_bytes, 49_152);
assert_eq!(resolved.launch.arguments_digest, arguments);
assert_ne!(resolved.launch.arguments_digest, [0; 32]);
assert_eq!(resolved.tuning_table_revision, 4);
let raw = build_resolved_gemm_launch_set(&[resolved]).unwrap();
let prepared = build_resolved_gemm_launch_set(&[resolved]).unwrap();
assert_eq!(prepared, raw);
assert_eq!(raw.launch_count, 1);
assert_ne!(raw.ordered_digest, [0; 32]);
}
fn rect_wide_sm120_binding() -> Tf32MapBinding {
let mut binding = portable_binding();
let module_target = CudaTarget::new("compute_120").unwrap();
let device_target = CudaTarget::new("sm_120").unwrap();
binding.qualified.module_kind = ModuleKind::TriadSm120;
binding.qualified.target = module_target;
binding.qualified.artifact.module_kind = ModuleKind::TriadSm120;
binding.qualified.compiler.target = module_target;
binding.qualified.device.compute_capability = (12, 0);
binding.qualified.device.multiprocessor_count = 170;
binding.qualified.device.target = device_target;
binding.qualified.device_caps.compute_capability = (12, 0);
binding.qualified.device_caps.accepted_target = Some(module_target);
binding.qualified.device_caps.tensor_map_access = true;
binding
}
fn tn_underfill_portable_binding() -> Tf32MapBinding {
let mut binding = portable_binding();
let module_target = CudaTarget::new("compute_120").unwrap();
let device_target = CudaTarget::new("sm_120").unwrap();
binding.qualified.target = module_target;
binding.qualified.compiler.target = module_target;
binding.qualified.device.compute_capability = (12, 0);
binding.qualified.device.multiprocessor_count = 170;
binding.qualified.device.target = device_target;
binding.qualified.device_caps.compute_capability = (12, 0);
binding.qualified.device_caps.accepted_target = Some(module_target);
binding
}
#[test]
fn tn_underfill_route_freezes_physical_and_graph_identity() {
let route = Tf32PhysicalRoute::MmaTf32Rna(Tf32PortableRoute {
tile: Tf32PortableTile::M16N32,
stages: Tf32PortableStages::S4,
});
let spec = tf32_kernel_spec(ResolvedGemmOp::Tn, route).unwrap();
let request = F32TriadRequest {
op: ResolvedGemmOp::Tn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Tn, (256, 512, 384)),
};
let config = cudarc::driver::LaunchConfig {
grid_dim: (384, 1, 1),
block_dim: (128, 1, 1),
shared_mem_bytes: 32_768,
};
let resolved = tf32_resolved_route(
request,
spec,
tn_underfill_portable_binding(),
Tf32LaunchDigests {
maps: [0; 32],
resources: [2; 32],
arguments: [3; 32],
},
false,
config,
);
assert_eq!(resolved.symbol, "tn_sm80_mma_tf32_m16n32_bk32_s4");
assert_eq!(resolved.module_kind, ModuleKind::TriadSm80);
assert_eq!(resolved.backend, PhysicalGemmBackend::MmaTf32Rna);
assert_eq!(
resolved.numeric_contract,
ResolvedNumericContract::MmaTf32Rna
);
assert_eq!(resolved.shape, (256, 512, 384));
assert_eq!(resolved.strides, (512, 384, 384));
assert_eq!(resolved.tile, (16, 32));
assert_eq!(resolved.bk, 32);
assert_eq!(resolved.stages, 4);
assert_eq!(resolved.threads, 128);
assert_eq!(resolved.launch.grid_dim, config.grid_dim);
assert_eq!(resolved.launch.block_dim, config.block_dim);
assert_eq!(resolved.launch.shared_mem_bytes, config.shared_mem_bytes);
assert_eq!(resolved.tensor_maps_digest, [0; 32]);
assert_eq!(resolved.resources_digest, [2; 32]);
assert_eq!(resolved.launch.arguments_digest, [3; 32]);
assert_eq!(resolved.tuning_table_revision, 46);
let eager = build_resolved_gemm_launch_set(&[resolved]).unwrap();
let graph = build_resolved_gemm_launch_set(&[resolved]).unwrap();
assert_eq!(eager, graph);
assert_ne!(eager.ordered_digest, [0; 32]);
let mut mutated = resolved;
mutated.launch.arguments_digest[0] ^= 1;
assert_ne!(
eager.ordered_digest,
build_resolved_gemm_launch_set(&[mutated])
.unwrap()
.ordered_digest
);
}
#[test]
fn rect_wide_route_freezes_physical_and_graph_identity() {
let route = Tf32PhysicalRoute::Sm120TmaMmaTf32Rna(Tf32Sm120Route {
tile: Tf32Sm120Tile::M80N32Bk64,
stages: Tf32Sm120Stages::S2,
});
let spec = tf32_kernel_spec(ResolvedGemmOp::Nn, route).unwrap();
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, (512, 3_072, 768)),
};
let config = cudarc::driver::LaunchConfig {
grid_dim: (168, 1, 1),
block_dim: (160, 1, 1),
shared_mem_bytes: 57_472,
};
let resolved = tf32_resolved_route(
request,
spec,
rect_wide_sm120_binding(),
Tf32LaunchDigests {
maps: [1; 32],
resources: [2; 32],
arguments: [3; 32],
},
false,
config,
);
assert_eq!(resolved.symbol, "nn_sm120_tma_mma_tf32_m80n32_bk64_s2");
assert_eq!(resolved.module_kind, ModuleKind::TriadSm120);
assert_eq!(resolved.backend, PhysicalGemmBackend::Sm120TmaMmaTf32Rna);
assert_eq!(
resolved.numeric_contract,
ResolvedNumericContract::Sm120TmaMmaTf32Rna
);
assert_eq!(resolved.shape, (512, 3_072, 768));
assert_eq!(resolved.strides, (3_072, 768, 768));
assert_eq!(resolved.tile, (80, 32));
assert_eq!(resolved.bk, 64);
assert_eq!(resolved.stages, 2);
assert_eq!(resolved.threads, 160);
assert_eq!(resolved.launch.grid_dim, config.grid_dim);
assert_eq!(resolved.launch.block_dim, config.block_dim);
assert_eq!(resolved.launch.shared_mem_bytes, config.shared_mem_bytes);
assert_eq!(resolved.tensor_maps_digest, [1; 32]);
assert_eq!(resolved.resources_digest, [2; 32]);
assert_eq!(resolved.launch.arguments_digest, [3; 32]);
assert_eq!(resolved.tuning_table_revision, 46);
let launch_set = build_resolved_gemm_launch_set(&[resolved]).unwrap();
assert_eq!(launch_set.launch_count, 1);
assert_ne!(launch_set.ordered_digest, [0; 32]);
let mut mutated = resolved;
mutated.launch.arguments_digest[0] ^= 1;
let mutated_set = build_resolved_gemm_launch_set(&[mutated]).unwrap();
assert_ne!(launch_set.ordered_digest, mutated_set.ordered_digest);
}
fn zero_request(op: ResolvedGemmOp) -> F32TriadRequest {
let dims = match op {
ResolvedGemmOp::Nn => (3, 0, 5),
ResolvedGemmOp::Tn => (0, 3, 5),
ResolvedGemmOp::Nt => (3, 5, 0),
};
F32TriadRequest {
op,
shape: F32TriadShape::contiguous(op, dims),
}
}
fn portable_route() -> Tf32PhysicalRoute {
Tf32PhysicalRoute::MmaTf32Rna(Tf32PortableRoute {
tile: Tf32PortableTile::M16N32,
stages: Tf32PortableStages::S4,
})
}
#[test]
fn tf32_kernel_param_abi_matches_cuda() {
macro_rules! check_abi {
($ty:ty, $size:literal, $($field:ident => $offset:literal),+ $(,)?) => {{
assert_eq!(std::mem::size_of::<$ty>(), $size);
assert_eq!(std::mem::align_of::<$ty>(), 4);
$(assert_eq!(std::mem::offset_of!($ty, $field), $offset);)+
}};
}
check_abi!(Sm80Tf32KernelParams, 32,
alpha => 0, beta => 4, m => 8, k => 12, n => 16, lda => 20, ldb => 24, ldc => 28);
check_abi!(Sm90aTf32KernelParams, 40,
a_x => 0, a_y => 4, b_x => 8, b_y => 12, alpha => 16, beta => 20,
m => 24, k => 28, n => 32, ldc => 36);
check_abi!(Sm100KernelParams, 40,
a_x => 0, a_y => 4, b_x => 8, b_y => 12, alpha => 16, beta => 20,
m => 24, k => 28, n => 32, ldc => 36);
check_abi!(Sm120KernelParams, 40,
a_x => 0, a_y => 4, b_x => 8, b_y => 12, alpha => 16, beta => 20,
m => 24, k => 28, n => 32, ldc => 36);
}
#[test]
fn f32_triad_operand_validation_rejects_null_and_misaligned_pointers() {
for op in [ResolvedGemmOp::Nn, ResolvedGemmOp::Tn, ResolvedGemmOp::Nt] {
let (request, operands, _) = scalar_fixture(op, (3, 4, 5));
validate_f32_triad_operands(request, operands)
.unwrap_or_else(|error| panic!("aligned {op:?} operands were rejected: {error}"));
validate_f32_triad_operands(
request,
F32TriadOperands {
b: operands.b + std::mem::size_of::<f32>() as u64,
..operands
},
)
.unwrap_or_else(|error| {
panic!("4-byte-aligned B+1 f32 {op:?} operand was rejected: {error}")
});
for (name, misaligned) in [
(
"A",
F32TriadOperands {
a: operands.a + 2,
..operands
},
),
(
"B",
F32TriadOperands {
b: operands.b + 2,
..operands
},
),
] {
let error = validate_f32_triad_operands(request, misaligned)
.expect_err("misaligned f32 inputs must be rejected");
assert!(error.contains(&format!("{name} pointer")), "{error}");
}
let zero_reduction_operands = F32TriadOperands {
a: 0,
b: 0,
..operands
};
validate_f32_triad_operands(zero_request(op), zero_reduction_operands)
.unwrap_or_else(|error| panic!("K=0 {op:?} operands were rejected: {error}"));
}
let (request, operands, _) = scalar_fixture(ResolvedGemmOp::Nn, (3, 4, 5));
for output in [0, 0x1002] {
let error =
validate_f32_triad_operands(request, F32TriadOperands { output, ..operands })
.expect_err("null and misaligned outputs must be rejected");
assert!(error.contains("output pointer"), "{error}");
}
for bias in [0, 0x2002] {
let error = validate_f32_triad_operands(
request,
F32TriadOperands {
bias: Some(bias),
..operands
},
)
.expect_err("null and misaligned biases must be rejected");
assert!(error.contains("bias pointer"), "{error}");
}
}
#[test]
fn zero_reduction_preparation_never_queries_input_allocations() {
let a_input_queries = Cell::new(0);
let b_input_queries = Cell::new(0);
let tensor_map_plan_queries = Cell::new(0);
let tensor_map_encodes = Cell::new(0);
for op in [ResolvedGemmOp::Nn, ResolvedGemmOp::Tn, ResolvedGemmOp::Nt] {
let request = zero_request(op);
let maps = prepare_f32_maps_with(
request,
F32TriadOperands {
output: 0x3000,
a: 0,
b: 0,
bias: None,
alpha: 1.0,
beta: if op == ResolvedGemmOp::Tn { 1.0 } else { 0.0 },
},
portable_route(),
portable_binding(),
|_, _, _| {
a_input_queries.set(a_input_queries.get() + 1);
b_input_queries.set(b_input_queries.get() + 1);
tensor_map_plan_queries.set(tensor_map_plan_queries.get() + 1);
Err("K=0 must not build a tensor-map plan".into())
},
|_| {
tensor_map_encodes.set(tensor_map_encodes.get() + 1);
Err("K=0 must not encode tensor maps".into())
},
)
.unwrap();
assert!(matches!(maps, F32PreparedTensorMaps::ZeroReduction { .. }));
}
assert_eq!(a_input_queries.get(), 0);
assert_eq!(b_input_queries.get(), 0);
assert_eq!(tensor_map_plan_queries.get(), 0);
assert_eq!(tensor_map_encodes.get(), 0);
let request = F32TriadRequest {
op: ResolvedGemmOp::Nn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Nn, (3, 4, 5)),
};
let error = prepare_f32_maps_with(
request,
F32TriadOperands {
output: 0x3000,
a: 0x1000,
b: 0x2000,
bias: None,
alpha: 1.0,
beta: 0.0,
},
portable_route(),
portable_binding(),
|_, _, _| {
a_input_queries.set(a_input_queries.get() + 1);
b_input_queries.set(b_input_queries.get() + 1);
tensor_map_plan_queries.set(tensor_map_plan_queries.get() + 1);
Err("nonzero control reached the plan boundary".into())
},
|_| {
tensor_map_encodes.set(tensor_map_encodes.get() + 1);
Err("nonzero control must stop at the plan boundary".into())
},
)
.unwrap_err();
assert!(error.contains("plan boundary"));
assert_eq!(a_input_queries.get(), 1);
assert_eq!(b_input_queries.get(), 1);
assert_eq!(tensor_map_plan_queries.get(), 1);
assert_eq!(tensor_map_encodes.get(), 0);
}
fn scalar_test_operands(op: ResolvedGemmOp, alpha: f32) -> F32TriadOperands {
F32TriadOperands {
output: 0x3000,
a: 0x1000,
b: 0x2000,
bias: None,
alpha,
beta: if op == ResolvedGemmOp::Tn { 1.0 } else { 0.0 },
}
}
fn scalar_tn_admission_facts() -> ScalarLaunchFacts {
let compiler = CompilerIdentity {
source_digest: [1; 32],
invocation_digest: [2; 32],
header_manifest_digest: [3; 32],
target: CudaTarget::new("compute_120").unwrap(),
nvrtc_version: (13, 2),
nvrtc_library_domain: [4; 32],
nvrtc_library_known: true,
output_kind: ArtifactKind::Ptx,
composer_revision: COMPOSER_REVISION,
compiler_revision: COMPILER_REVISION,
numeric_abi_revision: NUMERIC_ABI_REVISION,
schedule_revision: SCHEDULE_REVISION,
};
ScalarLaunchFacts {
scalar_artifact: ArtifactIdentity {
module_kind: ModuleKind::TriadScalar,
artifact_kind: ArtifactKind::Ptx,
compile_key: compiler.invocation_digest,
artifact_digest: [5; 32],
},
scalar_compiler: compiler,
fixed_artifact: ArtifactIdentity {
module_kind: ModuleKind::Fixed,
artifact_kind: ArtifactKind::Ptx,
compile_key: [0; 32],
artifact_digest: [0; 32],
},
fixed_compiler: compiler,
fixed_copyplan_loaded: false,
sm89_exact_f32_artifact: None,
sm89_exact_f32_compiler: None,
sm89_exact_f32_symbols_loaded: [false; 3],
sm89_exact_f32_d128_artifact: None,
sm89_exact_f32_d128_compiler: None,
sm89_exact_f32_d128_symbols_loaded: [false; 2],
compute_capability: (12, 0),
multiprocessor_count: 170,
}
}
#[test]
fn backward_launch_plan_uses_the_canonical_op_and_contiguous_abi() {
let (tn_dims, tn_request, tn_plan) =
scalar_backward_launch_plan(ResolvedGemmOp::Tn, (32, 4, 1), 142).unwrap();
assert_eq!(
scalar_backward_request(ResolvedGemmOp::Tn, (32, 4, 1)).unwrap(),
(tn_dims, tn_request)
);
assert_eq!(tn_dims.tuple(), (32, 4, 1));
assert_eq!(
tn_request.shape,
F32TriadShape {
m: 32,
k: 4,
n: 1,
lda: 4,
ldb: 1,
ldc: 1,
}
);
assert_eq!(tn_plan, ScalarDispatchPlan::TnGemv);
let (nt_dims, nt_request, nt_plan) =
scalar_backward_launch_plan(ResolvedGemmOp::Nt, (1, 7, 1), 142).unwrap();
assert_eq!(
scalar_backward_request(ResolvedGemmOp::Nt, (1, 7, 1)).unwrap(),
(nt_dims, nt_request)
);
assert_eq!(nt_dims.tuple(), (1, 7, 1));
assert_eq!(
nt_request.shape,
F32TriadShape {
m: 1,
k: 7,
n: 1,
lda: 1,
ldb: 1,
ldc: 7,
}
);
assert_eq!(nt_plan, ScalarDispatchPlan::NtGemv);
}
#[test]
fn tn_splitm_plan_matches_enqueue_geometry_at_the_batch_boundary() {
let (_, request, plan) =
scalar_backward_launch_plan(ResolvedGemmOp::Tn, (256, 128, 128), 142).unwrap();
assert_eq!(
plan,
ScalarDispatchPlan::TnSplitM {
m_chunk: 16,
chunks: 16,
}
);
let nodes =
scalar_physical_nodes(request, scalar_test_operands(ResolvedGemmOp::Tn, 1.0), plan)
.unwrap();
assert_eq!(nodes.len(), 2);
assert_eq!(nodes[0].symbol, "tn_splitm_partial_aligned");
assert_eq!(nodes[0].launch.grid_dim, (1, 1, 16));
assert_eq!(nodes[0].launch.block_dim, (256, 1, 1));
assert_eq!(nodes[1].symbol, "splitm_reduce");
assert_eq!(nodes[1].launch.grid_dim, (64, 1, 1));
assert_eq!(nodes[1].launch.block_dim, (256, 1, 1));
let (_, below_request, below_plan) =
scalar_backward_launch_plan(ResolvedGemmOp::Tn, (255, 128, 128), 142).unwrap();
assert_eq!(below_plan, ScalarDispatchPlan::TnFinal { slim: true });
let below = scalar_physical_nodes(
below_request,
scalar_test_operands(ResolvedGemmOp::Tn, 1.0),
below_plan,
)
.unwrap();
assert_eq!(below.len(), 1);
assert_eq!(below[0].symbol, "tn_slim");
}
#[test]
fn tn_narrow_splitm_qualified_manifests_have_exact_two_node_geometry() {
for (dims, partition, partial_grid, reducer_grid) in [
((1024, 47, 17), (32, 32), (1, 1, 32), (4, 1, 1)),
((1024, 128, 25), (32, 32), (2, 1, 32), (13, 1, 1)),
((4096, 64, 64), (48, 86), (1, 2, 86), (16, 1, 1)),
] {
let (_, request, automatic) =
scalar_backward_launch_plan(ResolvedGemmOp::Tn, dims, 170).unwrap();
assert_eq!(automatic, ScalarDispatchPlan::TnNarrow, "{dims:?}");
let aligned = scalar_test_operands(ResolvedGemmOp::Tn, 1.0);
let plan = scalar_launch_plan(scalar_tn_admission_facts(), request, aligned).unwrap();
assert_eq!(
plan,
ScalarDispatchPlan::TnNarrowSplitM {
m_chunk: partition.0,
chunks: partition.1,
},
"{dims:?}"
);
let aligned_nodes = scalar_physical_nodes(request, aligned, plan).unwrap();
assert_eq!(aligned_nodes.len(), 2, "{dims:?}");
assert_eq!(
aligned_nodes[0].symbol, "tn_narrow_splitm_partial_aligned",
"{dims:?}"
);
assert_eq!(aligned_nodes[0].launch.grid_dim, partial_grid, "{dims:?}");
assert_eq!(aligned_nodes[0].launch.block_dim, (128, 1, 1), "{dims:?}");
assert_eq!(aligned_nodes[1].symbol, "splitm_reduce", "{dims:?}");
assert_eq!(aligned_nodes[1].launch.grid_dim, reducer_grid, "{dims:?}");
assert_eq!(aligned_nodes[1].launch.block_dim, (256, 1, 1), "{dims:?}");
for unaligned in [
F32TriadOperands {
a: aligned.a + std::mem::size_of::<f32>() as u64,
..aligned
},
F32TriadOperands {
b: aligned.b + std::mem::size_of::<f32>() as u64,
..aligned
},
F32TriadOperands {
a: aligned.a + std::mem::size_of::<f32>() as u64,
b: aligned.b + std::mem::size_of::<f32>() as u64,
..aligned
},
] {
let unaligned_nodes = scalar_physical_nodes(request, unaligned, plan).unwrap();
assert_eq!(unaligned_nodes.len(), 2, "{dims:?}");
assert_eq!(
unaligned_nodes[0].symbol, "tn_narrow_splitm_partial",
"{dims:?}"
);
assert_eq!(unaligned_nodes[0].launch.grid_dim, partial_grid, "{dims:?}");
assert_eq!(unaligned_nodes[1].symbol, "splitm_reduce", "{dims:?}");
assert_eq!(unaligned_nodes[1].launch.grid_dim, reducer_grid, "{dims:?}");
}
let output_offset = F32TriadOperands {
output: aligned.output + std::mem::size_of::<f32>() as u64,
..aligned
};
let output_offset_nodes = scalar_physical_nodes(request, output_offset, plan).unwrap();
assert_eq!(
output_offset_nodes[0].symbol, "tn_narrow_splitm_partial_aligned",
"{dims:?}"
);
}
}
#[test]
fn tn_big_and_splitm_physical_symbols_separate_aligned_hot_paths() {
for (dims, aligned_symbol, fallback_symbol) in [
(
(256, 128, 128),
"tn_splitm_partial_aligned",
"tn_splitm_partial",
),
((128, 1024, 1024), "tn_aligned", "tn_big"),
] {
let (_, request, plan) =
scalar_backward_launch_plan(ResolvedGemmOp::Tn, dims, 142).unwrap();
let aligned = scalar_test_operands(ResolvedGemmOp::Tn, 1.0);
let mut unaligned = aligned;
unaligned.a += std::mem::size_of::<f32>() as u64;
let aligned_nodes = scalar_physical_nodes(request, aligned, plan).unwrap();
let unaligned_nodes = scalar_physical_nodes(request, unaligned, plan).unwrap();
assert_eq!(aligned_nodes[0].symbol, aligned_symbol, "{dims:?}");
assert_eq!(unaligned_nodes[0].symbol, fallback_symbol, "{dims:?}");
assert_ne!(
aligned_nodes[0].launch.arguments_digest,
unaligned_nodes[0].launch.arguments_digest,
"{dims:?}"
);
}
}
#[test]
fn nt_k_tail_34_node_plan_preserves_repeated_symbol_arguments_and_neighbors() {
let (_, request, plan) =
scalar_backward_launch_plan(ResolvedGemmOp::Nt, (32, 95, 128), 142).unwrap();
assert_eq!(
plan,
ScalarDispatchPlan::NtSplitKTail {
k_main: 64,
k_tail: 31,
}
);
let operands = scalar_test_operands(ResolvedGemmOp::Nt, 1.0);
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(nodes.len(), 34);
assert_eq!(nodes[0].symbol, "transpose_f32_2d");
assert_eq!(nodes[0].launch.grid_dim, (4, 2, 1));
assert_eq!(nodes[0].launch.block_dim, (32, 32, 1));
assert_eq!(nodes[1].symbol, "nn_splitk32_partial");
assert_eq!(nodes[1].launch.grid_dim, (4, 1, 1));
assert_eq!(nodes[2].symbol, "splitk_reduce");
assert_eq!(nodes[2].launch.grid_dim, (8, 1, 1));
for node in &nodes[3..] {
assert_eq!(node.symbol, "dx_col_gemv");
assert_eq!(node.launch.grid_dim, (1, 1, 1));
assert_eq!(node.launch.block_dim, (128, 1, 1));
}
for (index, node) in nodes[3..].iter().enumerate() {
assert!(
nodes[3..index + 3]
.iter()
.all(|prior| prior.launch.arguments_digest != node.launch.arguments_digest),
"tail node {index} collided with an earlier physical argument digest"
);
}
let (_, tail_30_request, tail_30_plan) =
scalar_backward_launch_plan(ResolvedGemmOp::Nt, (32, 94, 128), 142).unwrap();
assert_eq!(
scalar_physical_nodes(tail_30_request, operands, tail_30_plan)
.unwrap()
.len(),
33
);
let (_, aligned_request, aligned_plan) =
scalar_backward_launch_plan(ResolvedGemmOp::Nt, (32, 96, 128), 142).unwrap();
assert!(matches!(
aligned_plan,
ScalarDispatchPlan::NtSplitKMain {
n_main: 128,
n_tail: 0
}
));
assert_eq!(
scalar_physical_nodes(aligned_request, operands, aligned_plan)
.unwrap()
.len(),
3
);
for (dims, expected_plan) in [
((31, 95, 128), ScalarDispatchPlan::NtSmallBatchWide),
((32, 95, 127), ScalarDispatchPlan::NtNarrow),
((32, 95, 129), ScalarDispatchPlan::NtMidBatchWide),
] {
let (_, neighbor_request, neighbor_plan) =
scalar_backward_launch_plan(ResolvedGemmOp::Nt, dims, 142).unwrap();
assert_eq!(neighbor_plan, expected_plan, "neighbor {dims:?}");
let neighbor =
scalar_physical_nodes(neighbor_request, operands, neighbor_plan).unwrap();
assert_eq!(neighbor.len(), 1, "neighbor {dims:?}");
assert_eq!(neighbor[0].symbol, "nt_narrow", "neighbor {dims:?}");
}
let alpha_mutation =
scalar_physical_nodes(request, scalar_test_operands(ResolvedGemmOp::Nt, 2.0), plan)
.unwrap();
assert_ne!(
nodes[0].launch.arguments_digest,
alpha_mutation[0].launch.arguments_digest
);
}
}
#[cfg(test)]
mod half_physical_trace_tests {
use super::*;
use crate::mamba_ssm::gpu::context::{BiGemmFamily, F32TriadPolicy};
use crate::mamba_ssm::gpu::kernel_identity::{
ArtifactIdentity, ArtifactKind, COMPILER_REVISION, COMPOSER_REVISION, CompilerIdentity,
CudaTarget, DeviceCaps, DeviceIdentity, DriverIdentity, NUMERIC_ABI_REVISION,
PhysicalLaunchKind, build_artifact_set, route_backend_contract_sets,
};
fn small16_compiler_fixture() -> CompilerIdentity {
let target = CudaTarget::new("sm_89").unwrap();
CompilerIdentity {
source_digest: [31; 32],
invocation_digest: [32; 32],
header_manifest_digest: [33; 32],
target,
nvrtc_version: (13, 2),
nvrtc_library_domain: [34; 32],
nvrtc_library_known: true,
output_kind: ArtifactKind::Ptx,
composer_revision: COMPOSER_REVISION,
compiler_revision: COMPILER_REVISION,
numeric_abi_revision: NUMERIC_ABI_REVISION,
schedule_revision: SCHEDULE_REVISION,
}
}
fn small16_context_fixture(compiler: CompilerIdentity) -> GemmRouteIdentity {
let artifact = |module_kind, seed| ArtifactIdentity {
module_kind,
artifact_kind: ArtifactKind::Ptx,
compile_key: [seed; 32],
artifact_digest: [seed + 1; 32],
};
let policy = GemmPolicy {
batch_invariant: true,
bi_tensor_cores: true,
fast_gemm: false,
cublas_tf32: false,
f32_triad_policy: F32TriadPolicy::ExactScalarFma,
half_triad_policy: HalfTriadPolicy::TiledParity,
bi_gemm_family: BiGemmFamily::Triad,
};
let (backend_set, numeric_contracts) = route_backend_contract_sets(policy);
GemmRouteIdentity {
policy,
backend_set,
numeric_contracts,
compiler,
artifacts: build_artifact_set(&[
artifact(ModuleKind::Fixed, 1),
artifact(ModuleKind::TriadScalar, 3),
artifact(ModuleKind::TriadSm80, 5),
artifact(ModuleKind::TriadSm89Half, 7),
])
.unwrap(),
policy_revision: 1,
policy_hash: [35; 32],
device: DeviceIdentity {
compute_capability: (8, 9),
multiprocessor_count: 142,
target: CudaTarget::new("sm_89").unwrap(),
driver: DriverIdentity {
api_version: 13_020,
build_sources: 1,
build_digest: [36; 32],
},
},
device_caps: DeviceCaps {
compute_capability: (8, 9),
nvrtc_version: (13, 2),
accepted_target: Some(CudaTarget::new("sm_89").unwrap()),
optin_shared_bytes: 99_000,
tensor_map_access: false,
},
tuning_table_revision: TUNING_TABLE_REVISION,
schedule_set_revision: SCHEDULE_REVISION,
state_capacity: 64,
}
}
fn small16_observer() -> RecordingPhysicalObserver {
crate::mamba_ssm::gpu::kernel_identity::inference_test_support::observer(
|pointer, bytes| {
Ok(FramedSha256::new(b"small16-test-allocation.v1")
.required(b"pointer", &pointer.to_le_bytes())
.required(b"bytes", &bytes.to_le_bytes())
.finish())
},
)
}
#[test]
fn no_physical_observer_is_zero_sized_and_never_records() {
assert_eq!(std::mem::size_of::<NoPhysicalObserver>(), 0);
assert_eq!(
std::mem::size_of::<HalfLaunchEnvironment<'static, NoPhysicalObserver>>(),
3 * std::mem::size_of::<usize>()
);
}
#[test]
fn sm89_half_auto_launch_geometry_and_nn_parameter_abi_are_frozen() {
assert_eq!(std::mem::size_of::<Sm89HalfNnParams>(), 32);
assert_eq!(std::mem::align_of::<Sm89HalfNnParams>(), 4);
for (route, dtype, dims, expected_grid, expected_block, expected_shared) in [
(
super::super::sm89_half_source::Sm89HalfRoute::NnM128N128Bk64S3,
WeightDtype::Bf16,
(2048, 768, 3072),
384,
256,
98_304,
),
(
super::super::sm89_half_source::Sm89HalfRoute::NnM128N128Bk64S3,
WeightDtype::F16,
(4621, 384, 1928),
592,
256,
98_304,
),
(
super::super::sm89_half_source::Sm89HalfRoute::NtM128N128Bk64S3Bxor,
WeightDtype::F16,
(2048, 768, 3072),
96,
256,
98_304,
),
(
super::super::sm89_half_source::Sm89HalfRoute::NtM128N128Bk64S3Bxor,
WeightDtype::Bf16,
(4621, 384, 1928),
111,
256,
98_304,
),
(
super::super::sm89_half_source::Sm89HalfRoute::NtM96N128Bk64S3,
WeightDtype::F16,
(2048, 1536, 768),
264,
384,
86_016,
),
(
super::super::sm89_half_source::Sm89HalfRoute::TnM64N64Bk64S2RegpipeVec2,
WeightDtype::F16,
(2048, 768, 3072),
576,
128,
0,
),
(
super::super::sm89_half_source::Sm89HalfRoute::TnM64N64Bk64S2CompactBxor,
WeightDtype::Bf16,
(4621, 384, 1928),
186,
128,
0,
),
] {
let spec = (*super::super::sm89_half_source::kernel_spec(route, dtype).unwrap()).into();
let config = sm89_half_launch_config(spec, dims).unwrap();
assert_eq!(config.grid_dim, (expected_grid, 1, 1), "{route:?}");
assert_eq!(config.block_dim, (expected_block, 1, 1), "{route:?}");
assert_eq!(config.shared_mem_bytes, expected_shared, "{route:?}");
assert_eq!(
sm89_half_base(super::super::sm89_half_source::Sm89HalfRuntimeRoute::Legacy(route)),
spec.symbol
.trim_end_matches("_bf16")
.trim_end_matches("_f16")
);
}
}
#[test]
fn half_kernel_identity_uses_exact_dtype_suffix_and_module_owner() {
for (base, module_kind) in [
("nn_gemv", ModuleKind::TriadScalar),
("nn_ultra_thin", ModuleKind::TriadScalar),
("nn_narrow", ModuleKind::TriadScalar),
("nn_narrow_small", ModuleKind::TriadScalar),
("nn_big", ModuleKind::TriadScalar),
("tn_gemv", ModuleKind::TriadScalar),
("tn_narrow", ModuleKind::TriadScalar),
("tn_big", ModuleKind::TriadScalar),
("nt_gemv", ModuleKind::TriadScalar),
("nt_narrow", ModuleKind::TriadScalar),
("nt_big", ModuleKind::TriadScalar),
("nn_tc", ModuleKind::TriadSm80),
("nn_tc64", ModuleKind::TriadSm80),
("nn_tc16", ModuleKind::TriadSm80),
("tn_tc", ModuleKind::TriadSm80),
("tn_tc64", ModuleKind::TriadSm80),
("tn_tc128x64", ModuleKind::TriadSm80),
("nt_tc", ModuleKind::TriadSm80),
("nt_tc64", ModuleKind::TriadSm80),
("nn_sm89_m128n128_bk64_s3", ModuleKind::TriadSm89Half),
(
"tn_sm89_m64n64_bk64_s2_compact_bxor",
ModuleKind::TriadSm89Half,
),
(
"tn_sm89_m64n64_bk64_s2_regpipe_vec2",
ModuleKind::TriadSm89Half,
),
("nt_sm89_m128n128_bk64_s3_bxor", ModuleKind::TriadSm89Half),
("nt_sm89_m96n128_bk64_s3", ModuleKind::TriadSm89Half),
] {
for (dtype, suffix) in [(WeightDtype::Bf16, "_bf16"), (WeightDtype::F16, "_f16")] {
let identity = HalfKernelIdentity::resolve(base, dtype).unwrap();
assert_eq!(identity.symbol, format!("{base}{suffix}"));
assert_eq!(identity.module_kind, module_kind);
}
}
let tensor_core = HalfKernelIdentity::resolve("nn_tc64", WeightDtype::F16).unwrap();
assert!(
tensor_core
.validate(ModuleKind::TriadScalar, tensor_core.symbol)
.is_err()
);
assert!(
tensor_core
.validate(tensor_core.module_kind, "nn_tc64")
.is_err()
);
assert!(HalfKernelIdentity::resolve("nn_tc64", WeightDtype::F32).is_err());
assert!(HalfKernelIdentity::resolve("gemm_bi_unknown", WeightDtype::Bf16).is_err());
}
#[test]
fn triad_retained_half_small16_identity_uses_exact_suffix_and_owner() {
let base = "tn_sm89_m16n16_bk64_s2_ldb72";
for (dtype, symbol, resources_digest) in [
(
WeightDtype::Bf16,
"tn_sm89_m16n16_bk64_s2_ldb72_bf16",
[
49, 212, 166, 201, 100, 151, 56, 95, 31, 152, 208, 83, 18, 42, 33, 208, 193,
176, 220, 120, 125, 136, 61, 14, 36, 190, 210, 121, 95, 190, 204, 183,
],
),
(
WeightDtype::F16,
"tn_sm89_m16n16_bk64_s2_ldb72_f16",
[
100, 197, 96, 141, 108, 35, 8, 242, 247, 195, 64, 103, 194, 218, 40, 248, 79,
167, 23, 191, 35, 167, 248, 211, 107, 30, 98, 187, 145, 242, 0, 215,
],
),
] {
let identity =
HalfKernelIdentity::resolve(base, dtype).expect("retained small16 half identity");
assert_eq!(identity.symbol, symbol);
assert_eq!(identity.module_kind, ModuleKind::TriadSm89Half);
assert_eq!(identity.schedule, HalfSchedule::Tiled);
let spec = super::super::sm89_half_source::runtime_kernel_spec(symbol).unwrap();
let config = sm89_half_launch_config(spec, (1024, 256, 128)).unwrap();
assert_eq!(sm89_half_base(spec.route), base);
assert_eq!(config.grid_dim, (128, 1, 1));
assert_eq!(config.block_dim, (32, 1, 1));
assert_eq!(config.shared_mem_bytes, 0);
assert_eq!(
half_gemm_resources_digest(identity, config),
resources_digest
);
}
}
#[test]
fn triad_retained_half_small16_tn_argument_spans_match_the_physical_abi() {
let dims = (1024, 256, 128);
for dtype in [WeightDtype::Bf16, WeightDtype::F16] {
let base = "tn_sm89_m16n16_bk64_s2_ldb72";
let identity = HalfKernelIdentity::resolve(base, dtype).unwrap();
let spec =
super::super::sm89_half_source::runtime_kernel_spec(identity.symbol).unwrap();
let plan = sm89_half_tn_launch_plan(
spec,
0x3000,
TypedPtr { ptr: 0x2000, dtype },
TypedPtr { ptr: 0x1000, dtype },
dims,
None,
)
.unwrap();
assert_eq!(
plan.arguments,
Sm89HalfTnArguments {
core: [
Sm89HalfTnArgument::Pointer(0x3000),
Sm89HalfTnArgument::Pointer(0x1000),
Sm89HalfTnArgument::Pointer(0x2000),
Sm89HalfTnArgument::ScalarF32(1.0),
Sm89HalfTnArgument::ScalarI32(1024),
Sm89HalfTnArgument::ScalarI32(256),
Sm89HalfTnArgument::ScalarI32(128),
],
relay: None,
}
);
let ranges = std::cell::RefCell::new(Vec::new());
half_gemm_arguments_digest(
|pointer, bytes| {
ranges.borrow_mut().push((pointer, bytes));
Ok(FramedSha256::new(b"small16-test-allocation.v1")
.required(b"pointer", &pointer.to_le_bytes())
.required(b"bytes", &bytes.to_le_bytes())
.finish())
},
plan.observation,
identity,
half_policy_dtype(dtype).unwrap(),
)
.unwrap();
assert_eq!(
*ranges.borrow(),
[
(0x3000, 256 * 128 * 4),
(0x1000, 1024 * 256 * 2),
(0x2000, 1024 * 128 * 2),
]
);
assert_eq!(plan.config.grid_dim, (128, 1, 1));
assert_eq!(plan.config.block_dim, (32, 1, 1));
assert_eq!(plan.config.shared_mem_bytes, 0);
}
}
#[test]
fn triad_retained_half_relay_plan_carries_the_persistent_grid_and_its_workspaces() {
let dims = (2048, 1536, 768);
for (dtype, symbol) in [
(
WeightDtype::Bf16,
super::super::sm89_half_relay_source::RELAY_BF16_SYMBOL,
),
(
WeightDtype::F16,
super::super::sm89_half_relay_source::RELAY_F16_SYMBOL,
),
] {
let spec = super::super::sm89_half_source::runtime_kernel_spec(symbol).unwrap();
let relay = Sm89HalfRelayLaunch {
grid: 284,
workspace: Sm89HalfRelayWorkspace {
partial: 0x5000,
flags: 0x6000,
},
};
let plan = sm89_half_tn_launch_plan(
spec,
0x3000,
TypedPtr { ptr: 0x2000, dtype },
TypedPtr { ptr: 0x1000, dtype },
dims,
Some(relay),
)
.unwrap();
assert_eq!(plan.config.grid_dim, (284, 1, 1));
assert_eq!(plan.config.block_dim, (128, 1, 1));
assert_eq!(plan.config.shared_mem_bytes, 49_152);
assert_eq!(plan.arguments.relay, Some(relay.workspace));
assert_eq!(
plan.arguments.core,
[
Sm89HalfTnArgument::Pointer(0x3000),
Sm89HalfTnArgument::Pointer(0x1000),
Sm89HalfTnArgument::Pointer(0x2000),
Sm89HalfTnArgument::ScalarF32(1.0),
Sm89HalfTnArgument::ScalarI32(2048),
Sm89HalfTnArgument::ScalarI32(1536),
Sm89HalfTnArgument::ScalarI32(768),
]
);
let identity = HalfKernelIdentity::resolve(plan.observation.base, dtype).unwrap();
assert_eq!(identity.symbol, symbol);
assert_eq!(identity.module_kind, ModuleKind::TriadSm89Half);
assert_eq!(identity.schedule, HalfSchedule::RelayChain);
assert!(
sm89_half_tn_launch_plan(
spec,
0x3000,
TypedPtr { ptr: 0x2000, dtype },
TypedPtr { ptr: 0x1000, dtype },
dims,
None,
)
.is_err()
);
let tiled = super::super::sm89_half_source::runtime_kernel_spec(
super::super::sm89_half_tn_source::SMALL16_BF16_SYMBOL,
)
.unwrap();
assert!(
sm89_half_tn_launch_plan(
tiled,
0x3000,
TypedPtr {
ptr: 0x2000,
dtype: WeightDtype::Bf16
},
TypedPtr {
ptr: 0x1000,
dtype: WeightDtype::Bf16
},
(1024, 256, 128),
Some(relay),
)
.is_err()
);
}
}
#[test]
fn triad_retained_half_small16_eager_and_prepared_physical_nodes_match() {
let compiler = small16_compiler_fixture();
let context = small16_context_fixture(compiler);
let mut argument_digests = Vec::new();
for (dtype, policy_dtype, symbol, resources_digest) in [
(
WeightDtype::Bf16,
PolicyDtype::Bf16,
super::super::sm89_half_tn_source::SMALL16_BF16_SYMBOL,
[
49, 212, 166, 201, 100, 151, 56, 95, 31, 152, 208, 83, 18, 42, 33, 208, 193,
176, 220, 120, 125, 136, 61, 14, 36, 190, 210, 121, 95, 190, 204, 183,
],
),
(
WeightDtype::F16,
PolicyDtype::F16,
super::super::sm89_half_tn_source::SMALL16_F16_SYMBOL,
[
100, 197, 96, 141, 108, 35, 8, 242, 247, 195, 64, 103, 194, 218, 40, 248, 79,
167, 23, 191, 35, 167, 248, 211, 107, 30, 98, 187, 145, 242, 0, 215,
],
),
] {
let spec = super::super::sm89_half_source::runtime_kernel_spec(symbol).unwrap();
let plan = sm89_half_tn_launch_plan(
spec,
0x3000,
TypedPtr { ptr: 0x2000, dtype },
TypedPtr { ptr: 0x1000, dtype },
(1024, 256, 128),
None,
)
.unwrap();
let identity = HalfKernelIdentity::resolve(plan.observation.base, dtype).unwrap();
let observer = small16_observer();
let physical = resolve_half_gemm_observation_with_context(
&observer,
context,
compiler,
plan.observation,
identity,
plan.config,
)
.unwrap();
let node =
resolve_physical_launch_observation(&observer, physical, plan.config).unwrap();
let route = node.gemm_route().unwrap();
assert_eq!(node.kind(), PhysicalLaunchKind::Gemm);
assert_eq!(node.symbol(), symbol);
assert_eq!(node.module_kind(), ModuleKind::TriadSm89Half);
assert_eq!(node.logical_op(), ResolvedGemmOp::Tn);
assert_eq!(node.logical_dtype(), policy_dtype);
assert_eq!(node.execution_dtype(), policy_dtype);
assert_eq!(node.shape(), (1024, 256, 128));
assert_eq!(node.strides(), (256, 128, 128));
assert_eq!(node.tile(), Some((16, 16)));
assert_eq!(route.op, ResolvedGemmOp::Tn);
assert_eq!(route.dtype, policy_dtype);
assert_eq!(route.backend, PhysicalGemmBackend::Sm89Mma16HalfS2);
assert_eq!(route.numeric_contract, ResolvedNumericContract::MmaSyncF32);
assert_eq!(route.instruction_family, ResolvedInstructionFamily::MmaSync);
assert_eq!(
route.instruction_shape,
ResolvedInstructionShape { m: 16, n: 8, k: 16 }
);
assert_eq!(route.operand_conversion, ResolvedOperandConversion::None);
assert_eq!(
route.ownership,
ResolvedOutputOwnership::OneCtaPerOutputTile
);
assert_eq!(route.symbol, symbol);
assert_eq!(route.module_kind, ModuleKind::TriadSm89Half);
assert_eq!(route.target, compiler.target);
assert_eq!(route.artifact, context.artifacts.sm89_half.unwrap());
assert_eq!(route.compiler, compiler);
assert_eq!(route.device, context.device);
assert_eq!(route.device_caps, context.device_caps);
assert_eq!(route.shape, (1024, 256, 128));
assert_eq!(route.strides, (256, 128, 128));
assert_eq!(route.tile, (16, 16));
assert_eq!((route.bk, route.stages, route.threads), (64, 2, 32));
assert_eq!(route.launch.grid_dim, (128, 1, 1));
assert_eq!(route.launch.block_dim, (32, 1, 1));
assert_eq!(route.launch.shared_mem_bytes, 0);
assert_ne!(route.launch.arguments_digest, [0; 32]);
assert_eq!(route.tensor_map_revision, 0);
assert_eq!(route.tensor_maps_digest, [0; 32]);
assert_eq!(route.resources_digest, resources_digest);
assert_eq!(route.tuning_table_revision, SM89_HALF_ROUTE_REVISION);
assert_eq!(route.schedule_revision, SCHEDULE_REVISION);
argument_digests.push(route.launch.arguments_digest);
let seal = half_native_branch_seal(plan.observation, plan.config);
assert_eq!(seal.base, plan.observation.base);
assert_eq!(seal.op, node.logical_op());
assert_eq!(half_policy_dtype(seal.dtype).unwrap(), node.logical_dtype());
assert_eq!(seal.dims, node.shape());
assert_eq!(seal.strides, node.strides());
assert_eq!(Some(seal.tile), node.tile());
assert_eq!(seal.bk_stages, (route.bk, route.stages));
assert_eq!(seal.grid_dim, node.launch().grid_dim);
assert_eq!(seal.block_dim, node.launch().block_dim);
assert_eq!(seal.shared_mem_bytes, node.launch().shared_mem_bytes);
let request = HalfPhysicalTraceRequest {
op: ResolvedGemmOp::Tn,
output: 0x3000,
a: 0x1000,
b: 0x2000,
bias: 0,
dtype,
dims: (1024, 256, 128),
nn_strides: None,
forced_tile: None,
capacity: 1,
};
let base = prepared_half_graph_base(node, dtype).unwrap();
let (prepared_config, prepared_node) = resolve_prepared_half_graph_node_with_context(
&observer, context, compiler, node, request, base,
)
.unwrap();
assert_eq!(prepared_config.grid_dim, plan.config.grid_dim);
assert_eq!(prepared_config.block_dim, plan.config.block_dim);
assert_eq!(
prepared_config.shared_mem_bytes,
plan.config.shared_mem_bytes
);
assert_eq!(prepared_node, node);
let wrong_dtype = match dtype {
WeightDtype::Bf16 => WeightDtype::F16,
WeightDtype::F16 => WeightDtype::Bf16,
WeightDtype::F32 | WeightDtype::Tf32 => unreachable!(),
};
assert!(prepared_half_graph_base(node, wrong_dtype).is_err());
let mut wrong_owner = node;
wrong_owner.module_kind = ModuleKind::TriadSm80;
assert!(
identity
.validate(wrong_owner.module_kind(), wrong_owner.symbol())
.is_err()
);
let error = resolve_prepared_half_graph_node_with_context(
&observer,
context,
compiler,
wrong_owner,
request,
base,
)
.expect_err("prepared small16 identity must reject the wrong module owner");
assert!(error.contains("exact symbol and module owner"), "{error}");
}
assert_ne!(argument_digests[0], argument_digests[1]);
}
#[test]
fn rect128x64_tn_forced_identity_is_exact() {
let tile = TcTile::Rect128x64;
assert_eq!(tile.extents(), (128, 64));
assert_eq!(tile.block_dim(), 256);
assert_eq!(tile.bk_stages(), (32, 3));
let config = tile.launch_cfg(1536, 768, 69_632).unwrap();
assert_eq!(config.grid_dim, (144, 1, 1));
assert_eq!(config.block_dim, (256, 1, 1));
assert_eq!(config.shared_mem_bytes, 0);
for (dtype, symbol) in [
(WeightDtype::Bf16, "tn_tc128x64_bf16"),
(WeightDtype::F16, "tn_tc128x64_f16"),
] {
let identity = HalfKernelIdentity::resolve("tn_tc128x64", dtype).unwrap();
assert_eq!(identity.symbol, symbol);
assert_eq!(identity.module_kind, ModuleKind::TriadSm80);
}
}
}
#[cfg(test)]
#[path = "scalar_nt_tests.rs"]
mod scalar_nt_tests;
#[cfg(test)]
#[path = "scalar_nn_tn_tests.rs"]
mod scalar_nn_tn_tests;
#[cfg(test)]
mod sm89_exact_f32_tn_route_tests {
use super::*;
#[test]
fn large_tn_routes_freeze_nodes_scratch_and_numeric_identity() {
let cases = [
(
(2_048, 768, 3_072),
ScalarDispatchPlan::TnD768InSm89DualChunkQualified,
[
"transpose_f32_32x16_d768",
super::super::D768_IN_FUSED_SYMBOL,
],
[(24, 64, 1), (576, 1, 1)],
Some(2_048 * 768),
),
(
(2_048, 1_536, 768),
ScalarDispatchPlan::TnD768OutSm89DirectBk16Qualified,
[super::super::D768_OUT_RAW_SYMBOL, "splitm_reduce"],
[(288, 1, 4), (4_608, 1, 1)],
None,
),
(
(4_621, 384, 1_928),
ScalarDispatchPlan::TnPrismSm89DirectBk16Qualified,
[super::super::PRISM_RAW_SYMBOL, "splitm_reduce"],
[(186, 1, 6), (2_892, 1, 1)],
None,
),
];
for (dims, plan, symbols, grids, transpose_elements) in cases {
let request = F32TriadRequest {
op: ResolvedGemmOp::Tn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Tn, dims),
};
let operands = F32TriadOperands {
output: 0x3000,
a: 0x1000,
b: 0x2000,
bias: None,
alpha: 1.0,
beta: 1.0,
};
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(nodes.len(), 2);
assert_eq!([nodes[0].symbol, nodes[1].symbol], symbols);
assert_eq!([nodes[0].launch.grid_dim, nodes[1].launch.grid_dim], grids);
assert_ne!(
nodes[0].launch.arguments_digest,
nodes[1].launch.arguments_digest
);
assert_eq!(
scalar_transpose_scratch_elements(request, plan).unwrap(),
transpose_elements
);
assert_eq!(plan.needs_split_scratch(), transpose_elements.is_none());
}
assert_eq!(
scalar_route_contract(super::super::D768_IN_FUSED_SYMBOL),
(
PhysicalGemmBackend::ScalarFmaSm89ExactF32DualChunkFused,
ResolvedNumericContract::ScalarFmaTnSplitMF64Reduce,
ResolvedOutputOwnership::OneCtaPerOutputTile,
)
);
for symbol in [
super::super::D768_OUT_RAW_SYMBOL,
super::super::PRISM_RAW_SYMBOL,
] {
assert_eq!(
scalar_route_contract(symbol),
(
PhysicalGemmBackend::ScalarFmaSm89ExactF32DirectSplitMPartial,
ResolvedNumericContract::ScalarFmaTnSplitMPartial,
ResolvedOutputOwnership::OneCtaPerOutputTilePerSplitMPartition,
)
);
}
}
#[test]
fn d128_direct_folds_are_one_node_exact_routes_without_scratch() {
for (dims, plan, symbol, grid) in [
(
(1_024, 128, 512),
ScalarDispatchPlan::TnD128InSm89DirectFoldQualified,
super::super::D128_IN_SYMBOL,
256,
),
(
(1_024, 256, 128),
ScalarDispatchPlan::TnD128OutSm89DirectFoldQualified,
super::super::D128_OUT_SYMBOL,
128,
),
] {
let request = F32TriadRequest {
op: ResolvedGemmOp::Tn,
shape: F32TriadShape::contiguous(ResolvedGemmOp::Tn, dims),
};
let operands = F32TriadOperands {
output: 0x3000,
a: 0x1000,
b: 0x2000,
bias: None,
alpha: 1.0,
beta: 1.0,
};
let nodes = scalar_physical_nodes(request, operands, plan).unwrap();
assert_eq!(scalar_node_count(plan), 1);
assert_eq!(nodes.len(), 1);
assert_eq!(nodes[0].symbol, symbol);
assert_eq!(nodes[0].tile, (16, 16));
assert_eq!(nodes[0].bk, 16);
assert_eq!(nodes[0].stages, 2);
assert_eq!(nodes[0].launch.grid_dim, (grid, 1, 1));
assert_eq!(nodes[0].launch.block_dim, (256, 1, 1));
assert_eq!(nodes[0].launch.shared_mem_bytes, 49_152);
assert!(!plan.needs_split_scratch());
assert!(!plan.needs_transpose_scratch());
assert_eq!(
scalar_plan_fields(plan).0,
if grid == 256 { 44 } else { 45 }
);
assert_eq!(
scalar_route_contract(symbol),
(
PhysicalGemmBackend::ScalarFmaTnDirectF64FoldSm89,
ResolvedNumericContract::ScalarFmaTnSplitMF64Reduce,
ResolvedOutputOwnership::OneCtaPerOutputTile,
)
);
}
}
}
#[cfg(test)]
mod sm89_tf32_joint_route_tests {
use super::*;
fn request(op: ResolvedGemmOp, dims: (usize, usize, usize)) -> F32TriadRequest {
F32TriadRequest {
op,
shape: F32TriadShape::contiguous(op, dims),
}
}
fn operands(op: ResolvedGemmOp) -> F32TriadOperands {
F32TriadOperands {
output: 0x3000,
a: 0x1000,
b: 0x2000,
bias: None,
alpha: 1.0,
beta: if op == ResolvedGemmOp::Tn { 1.0 } else { 0.0 },
}
}
#[test]
fn exact_cells_freeze_transpose_layout_and_physical_gemm_geometry() {
for (dims, route, stride, elements, physical, grid) in [
(
(2_048, 768, 3_072),
Tf32PhysicalRoute::Sm89TnPreRnaN96,
2_048,
1_572_864,
(768, 2_048, 3_072),
192,
),
(
(2_048, 1_536, 768),
Tf32PhysicalRoute::Sm89TnPreRnaN96,
2_048,
3_145_728,
(1_536, 2_048, 768),
96,
),
(
(4_621, 384, 1_928),
Tf32PhysicalRoute::Sm89TnPreRnaM64N96S2,
4_624,
1_775_616,
(384, 4_621, 1_928),
126,
),
] {
let request = request(ResolvedGemmOp::Tn, dims);
assert_eq!(
sm89_tf32_tn_scratch_layout(request).unwrap(),
(stride, elements)
);
let spec = tf32_kernel_spec(request.op, route).unwrap();
assert_eq!(
(request.shape.k, request.shape.m, request.shape.n,),
physical
);
assert_eq!(
checked_tile_grid(
u32::try_from(physical.0).unwrap(),
spec.tile.0,
u32::try_from(physical.2).unwrap(),
spec.tile.1,
)
.unwrap(),
grid
);
}
}
#[test]
fn exact_cells_reject_alpha_beta_bias_null_and_alignment_drift() {
for (op, dims, route) in [
(
ResolvedGemmOp::Tn,
(2_048, 768, 3_072),
Tf32PhysicalRoute::Sm89TnPreRnaN96,
),
(
ResolvedGemmOp::Tn,
(4_621, 384, 1_928),
Tf32PhysicalRoute::Sm89TnPreRnaM64N96S2,
),
(
ResolvedGemmOp::Nn,
(4_621, 384, 1_928),
Tf32PhysicalRoute::Sm89NnDirectN96,
),
(
ResolvedGemmOp::Nn,
(2_048, 1_536, 768),
Tf32PhysicalRoute::Sm89NnN96,
),
(
ResolvedGemmOp::Nt,
(2_048, 768, 3_072),
Tf32PhysicalRoute::Sm89NtALdmatrixN96,
),
(
ResolvedGemmOp::Tn,
(2_048, 768, 3_072),
Tf32PhysicalRoute::Sm89TnPreRnaM96N192S2,
),
(
ResolvedGemmOp::Tn,
(2_048, 1_536, 768),
Tf32PhysicalRoute::Sm89TnPreRnaM96N96S3,
),
(
ResolvedGemmOp::Tn,
(4_096, 3_072, 1_536),
Tf32PhysicalRoute::Sm89TnDirectM192N192S2,
),
(
ResolvedGemmOp::Nt,
(4_096, 3_072, 1_536),
Tf32PhysicalRoute::Sm89NtRowstageM128N192S2,
),
(
ResolvedGemmOp::Nt,
(4_621, 384, 1_928),
Tf32PhysicalRoute::Sm89NtRnaM144N96S2,
),
] {
let request = request(op, dims);
let valid = operands(op);
validate_sm89_tf32_joint_operands(request, valid, route).unwrap();
let mut mutations = Vec::new();
mutations.push(F32TriadOperands {
alpha: 0.5,
..valid
});
mutations.push(F32TriadOperands {
beta: if op == ResolvedGemmOp::Tn { 0.0 } else { 1.0 },
..valid
});
mutations.push(F32TriadOperands {
bias: Some(0x4000),
..valid
});
mutations.push(F32TriadOperands { output: 0, ..valid });
mutations.push(F32TriadOperands { a: 0x1004, ..valid });
mutations.push(F32TriadOperands { b: 0x2004, ..valid });
for mutation in mutations {
assert!(
validate_sm89_tf32_joint_operands(request, mutation, route).is_err(),
"joint route {route:?} accepted operand drift {mutation:?}"
);
}
}
}
}
#[cfg(test)]
mod sm120_api_tests {
use super::super::super::kernels::MambaKernels;
use super::super::contract::{
Sm120ForcedRoute, Sm120LaunchOperands, Sm120MapRequest, Sm120PreparedLaunch,
Sm120PreparedTensorMaps, Sm120RouteIdentity,
};
use super::{
launch_sm120_tma_prepared, prepare_sm120_tensor_maps, prepare_sm120_tma_forced,
validate_sm120_graph_replay,
};
use std::sync::Arc;
type Stream = Arc<cudarc::driver::CudaStream>;
type PrepareMaps =
fn(&Stream, &MambaKernels, Sm120MapRequest) -> Result<Sm120PreparedTensorMaps, String>;
type PrepareLaunch = fn(
&Stream,
&MambaKernels,
Sm120ForcedRoute,
&Sm120PreparedTensorMaps,
Sm120LaunchOperands,
) -> Result<Sm120PreparedLaunch, String>;
type LaunchPrepared =
fn(&Stream, &MambaKernels, &Sm120PreparedLaunch) -> Result<Sm120RouteIdentity, String>;
type ValidateReplay = fn(&Stream, &MambaKernels, &Sm120PreparedLaunch) -> Result<(), String>;
const _: PrepareMaps = prepare_sm120_tensor_maps;
const _: PrepareLaunch = prepare_sm120_tma_forced;
const _: LaunchPrepared = launch_sm120_tma_prepared;
const _: ValidateReplay = validate_sm120_graph_replay;
}
pub fn query_specialized_device_caps(
stream: &Arc<cudarc::driver::CudaStream>,
nvrtc_arch: &'static str,
nvrtc_version: (i32, i32),
) -> Result<crate::mamba_ssm::gpu::kernel_identity::DeviceCaps, String> {
let ctx = stream.context();
let (major, minor) = ctx
.compute_capability()
.map_err(|error| format!("query compute capability for {nvrtc_arch}: {error:?}"))?;
let optin_shared = ctx
.attribute(
cudarc::driver::sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_MAX_SHARED_MEMORY_PER_BLOCK_OPTIN,
)
.map_err(|error| format!("query opt-in shared memory for {nvrtc_arch}: {error:?}"))?;
let tensor_map_access = ctx
.attribute(
cudarc::driver::sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_TENSOR_MAP_ACCESS_SUPPORTED,
)
.map_err(|error| format!("query tensor-map support for {nvrtc_arch}: {error:?}"))?
!= 0;
Ok(crate::mamba_ssm::gpu::kernel_identity::DeviceCaps {
compute_capability: (
u32::try_from(major).map_err(|_| format!("negative CUDA CC major {major}"))?,
u32::try_from(minor).map_err(|_| format!("negative CUDA CC minor {minor}"))?,
),
nvrtc_version,
accepted_target: Some(crate::mamba_ssm::gpu::kernel_identity::CudaTarget::new(
nvrtc_arch,
)?),
optin_shared_bytes: u32::try_from(optin_shared)
.map_err(|_| format!("negative opt-in shared memory {optin_shared}"))?,
tensor_map_access,
})
}