use super::device::GpuDevice;
use super::dtype::WeightDtype;
use super::gemm_bi_triad::{
F32PreparedLaunchCache, Sm90aPreparedLaunchCache, Sm100PreparedLaunchCache,
Sm120PreparedLaunchCache,
};
pub use super::gemm_mode::GemmMode;
use super::gemm_mode::{MathModeBackend, MathTransitionError, change_math_mode};
use super::kernel_identity::{
ArtifactIdentity, BackendSet, CapturedGemmGraphPlan, CompilerIdentity, ModuleKind,
PhysicalGemmBackend, PolicyDtype, PreparedGemmCaptureManifest, RecordedGemmTrace,
ResolvedGemmLaunchSet, ResolvedGemmRoute, ResolvedNumericContract,
build_resolved_gemm_launch_set,
};
use super::kernels::MambaKernels;
use crate::config::MambaConfig;
use std::cell::{Cell, RefCell};
use std::rc::Rc;
use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
static NEXT_GPU_CTX_TOKEN: AtomicU64 = AtomicU64::new(1);
fn m1_mixed_graph_max_dim(dims: &super::forward::GpuMambaDims) -> usize {
dims.d_model
.max(2 * dims.d_inner)
.max(dims.xdbl_dim)
.max(dims.mamba_input_dim)
}
fn next_gpu_ctx_token() -> Result<u64, String> {
NEXT_GPU_CTX_TOKEN
.fetch_update(Ordering::Relaxed, Ordering::Relaxed, |token| {
token.checked_add(1)
})
.map_err(|_| "GpuCtx instance token space exhausted".to_string())
}
struct GemmEnvValues {
mode: Result<String, std::env::VarError>,
}
impl GemmEnvValues {
fn read() -> Self {
Self {
mode: std::env::var("MAMBA_RS_GEMM_MODE"),
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) struct GemmRole {
pub(crate) family: BiGemmFamily,
pub(crate) f32_numeric: F32TriadPolicy,
}
impl GemmRole {
pub(crate) fn triad(dtype: WeightDtype) -> Self {
Self {
family: BiGemmFamily::Triad,
f32_numeric: dtype.f32_numeric(),
}
}
pub(crate) fn inference(dtype: WeightDtype) -> Self {
Self {
family: BiGemmFamily::Inference,
f32_numeric: dtype.f32_numeric(),
}
}
}
impl WeightDtype {
pub(crate) fn f32_numeric(self) -> F32TriadPolicy {
match self {
Self::Tf32 => F32TriadPolicy::AllowDeterministicTf32,
Self::F32 | Self::Bf16 | Self::F16 => F32TriadPolicy::ExactScalarFma,
}
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
struct ResolvedGemmEnv {
mode: GemmMode,
tensor_cores: bool,
family: BiGemmFamily,
f32_policy: F32TriadPolicy,
half_policy: HalfTriadPolicy,
}
fn explicit_gemm_config(mode: GemmMode, role: GemmRole) -> ResolvedGemmEnv {
ResolvedGemmEnv {
mode,
family: role.family,
tensor_cores: true,
f32_policy: role.f32_numeric,
half_policy: HalfTriadPolicy::AllowStreamKFixedOrder,
}
}
fn gemm_mode_from_env(mode: Result<String, std::env::VarError>) -> Result<GemmMode, String> {
match mode {
Ok(value) => GemmMode::parse_env_value(&value),
Err(std::env::VarError::NotPresent) => Ok(GemmMode::Deterministic),
Err(std::env::VarError::NotUnicode(value)) => Err(format!(
"MAMBA_RS_GEMM_MODE={value:?} is not valid Unicode \
(use deterministic, cublas-fast, or cublas-pedantic)"
)),
}
}
fn resolve_gemm_env(values: GemmEnvValues, role: GemmRole) -> Result<ResolvedGemmEnv, String> {
Ok(explicit_gemm_config(gemm_mode_from_env(values.mode)?, role))
}
#[doc(hidden)]
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, Default)]
pub enum BiGemmFamily {
#[default]
Triad,
Inference,
}
#[doc(hidden)]
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
#[repr(u8)]
pub enum F32TriadPolicy {
#[default]
ExactScalarFma = 0,
AllowDeterministicTf32 = 1,
}
#[doc(hidden)]
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Hash)]
#[repr(u8)]
pub enum HalfTriadPolicy {
#[default]
TiledParity = 0,
AllowStreamKFixedOrder = 1,
}
pub use super::kernel_identity::GemmRouteIdentity;
#[doc(hidden)]
pub use super::kernel_identity::{BackendSet as GemmBackendSet, GemmPolicy, NumericContractSet};
pub type GemmRoute = GemmRouteIdentity;
struct GemmRouteRecorder {
context: GemmRouteIdentity,
mode: GemmRouteRecorderMode,
routes: Vec<ResolvedGemmRoute>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
enum GemmRouteRecorderMode {
GrowableEager,
FixedCapture { route_capacity: u32 },
}
pub(crate) struct GemmRouteRecordingGuard<'a> {
ctx: &'a GpuCtx,
}
impl GemmRouteRecordingGuard<'_> {
#[cfg(test)]
pub(crate) fn finish(self) -> Result<CapturedGemmGraphPlan, String> {
let recorder = self.ctx.take_gemm_route_recording()?;
if recorder.routes.is_empty() {
return Err("captured GEMM graph plan must not be empty".into());
}
recorder
.context
.ensure_current(self.ctx.gemm_route(), "GEMM graph capture")?;
let launches = build_resolved_gemm_launch_set(&recorder.routes)?;
let routes = recorder.routes.into_boxed_slice();
CapturedGemmGraphPlan::new(recorder.context, launches, routes)
}
pub(crate) fn finish_trace(self) -> Result<RecordedGemmTrace, String> {
let recorder = self.ctx.take_gemm_route_recording()?;
if recorder.mode != GemmRouteRecorderMode::GrowableEager {
return Err("fixed GEMM capture recording cannot produce an eager trace".into());
}
recorder
.context
.ensure_current(self.ctx.gemm_route(), "eager GEMM route recording")?;
RecordedGemmTrace::from_routes(recorder.context, recorder.routes)
}
pub(crate) fn finish_against_manifest(
self,
manifest: &PreparedGemmCaptureManifest,
) -> Result<Option<CapturedGemmGraphPlan>, String> {
let recorder = self.ctx.take_gemm_route_recording()?;
let GemmRouteRecorderMode::FixedCapture { route_capacity } = recorder.mode else {
return Err("growable eager GEMM recording cannot finish a CUDA capture".into());
};
let route_capacity =
usize::try_from(route_capacity).expect("u32 GEMM route capacity always fits in usize");
manifest.validate_capture_request(self.ctx.gemm_route(), route_capacity)?;
recorder
.context
.ensure_current(self.ctx.gemm_route(), "GEMM graph capture")?;
let launches = resolved_gemm_launches(&recorder.routes)?;
manifest.validate_capture_result(recorder.context, launches)?;
let Some(launches) = launches else {
return Ok(None);
};
let routes = recorder.routes.into_boxed_slice();
CapturedGemmGraphPlan::new(recorder.context, launches, routes).map(Some)
}
}
fn resolved_gemm_launches(
routes: &[ResolvedGemmRoute],
) -> Result<Option<ResolvedGemmLaunchSet>, String> {
if routes.is_empty() {
Ok(None)
} else {
build_resolved_gemm_launch_set(routes).map(Some)
}
}
impl Drop for GemmRouteRecordingGuard<'_> {
fn drop(&mut self) {
self.ctx.clear_gemm_route_recording();
}
}
#[doc(hidden)]
pub struct GpuCtxResources {
pub stream: Arc<cudarc::driver::CudaStream>,
pub kernels: Arc<MambaKernels>,
pub blas: cudarc::cublas::CudaBlas,
pub _blas_workspace: Arc<cudarc::driver::CudaSlice<u8>>,
half_staging: RefCell<Option<cudarc::driver::CudaSlice<u8>>>,
half_staging_ptr: RefCell<cudarc::driver::sys::CUdeviceptr>,
half_staging_bytes: RefCell<usize>,
bi_upcast_scratch: [RefCell<Option<super::buffers::GpuBuffer>>; 3],
}
pub struct GpuCtx {
resources: Rc<GpuCtxResources>,
gemm_route_recorder: RefCell<Option<GemmRouteRecorder>>,
f32_prepared_launches: RefCell<F32PreparedLaunchCache>,
route_proofs: RefCell<super::gemm_bi_triad::proof::RouteProofLedger>,
sm120_prepared_launches: RefCell<Sm120PreparedLaunchCache>,
sm100_prepared_launches: RefCell<Sm100PreparedLaunchCache>,
sm90a_prepared_launches: RefCell<Sm90aPreparedLaunchCache>,
pub(crate) fixed_tf32_maps: RefCell<super::gemm_bi_inference::FixedTf32MapCache>,
pub(crate) fixed_postbias_maps: RefCell<super::gemm_bi_inference::FixedPostBiasMapCache>,
pub(crate) fixed_half_maps: RefCell<super::gemm_bi_inference::FixedHalfMapCache>,
gemm_mode: Cell<GemmMode>,
gemm_unusable: RefCell<Option<String>>,
bi_tensor_cores: std::cell::Cell<bool>,
bi_gemm_family: std::cell::Cell<BiGemmFamily>,
f32_triad_policy: std::cell::Cell<F32TriadPolicy>,
half_triad_policy: std::cell::Cell<HalfTriadPolicy>,
state_cap: usize,
instance_token: u64,
device_identity: super::kernel_identity::DeviceIdentity,
device_caps: super::kernel_identity::DeviceCaps,
policy_hash: super::kernel_identity::Sha256Digest,
graphs_captured: std::cell::Cell<u64>,
graph_scratch_frozen: std::cell::Cell<bool>,
}
impl std::ops::Deref for GpuCtx {
type Target = GpuCtxResources;
fn deref(&self) -> &Self::Target {
&self.resources
}
}
fn validate_multiprocessor_identity(
device_multiprocessor_count: u32,
kernel_multiprocessor_count: u32,
) -> Result<(), String> {
if device_multiprocessor_count == 0 || kernel_multiprocessor_count == 0 {
return Err("CUDA topology requires a nonzero multiprocessor count".into());
}
if device_multiprocessor_count != kernel_multiprocessor_count {
return Err(format!(
"CUDA topology changed while loading kernels: device identity has {device_multiprocessor_count} multiprocessors but loaded kernels observed {kernel_multiprocessor_count}"
));
}
Ok(())
}
struct CublasMathBackend<'a> {
blas: &'a cudarc::cublas::CudaBlas,
}
impl MathModeBackend for CublasMathBackend<'_> {
type Mode = cudarc::cublas::sys::cublasMath_t;
fn query(&mut self) -> Result<Self::Mode, String> {
let mut mode = Self::Mode::CUBLAS_DEFAULT_MATH;
let status =
unsafe { cudarc::cublas::sys::cublasGetMathMode(*self.blas.handle(), &mut mode) };
if status == cudarc::cublas::sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS {
Ok(mode)
} else {
Err(format!("cublasGetMathMode failed: {status:?}"))
}
}
fn update(&mut self, mode: Self::Mode) -> Result<(), String> {
let status = unsafe { cudarc::cublas::sys::cublasSetMathMode(*self.blas.handle(), mode) };
if status == cudarc::cublas::sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS {
Ok(())
} else {
Err(format!("cublasSetMathMode({mode:?}) failed: {status:?}"))
}
}
}
impl GpuCtx {
pub fn new(device: &GpuDevice) -> Result<Self, String> {
Self::new_with_mode(device, GemmMode::default())
}
pub fn new_with_mode(device: &GpuDevice, mode: GemmMode) -> Result<Self, String> {
Self::new_with_state_cap_and_mode(device, 64, mode)
}
pub fn new_from_env(device: &GpuDevice) -> Result<Self, String> {
Self::new_from_env_with_state_cap_and_role(device, 64, GemmRole::triad(WeightDtype::F32))
}
pub fn new_from_env_with_state_cap(
device: &GpuDevice,
state_cap: usize,
) -> Result<Self, String> {
Self::new_from_env_with_state_cap_and_role(
device,
state_cap,
GemmRole::triad(WeightDtype::F32),
)
}
pub fn new_with_state_cap(device: &GpuDevice, state_cap: usize) -> Result<Self, String> {
Self::new_with_state_cap_and_mode(device, state_cap, GemmMode::default())
}
pub fn new_with_state_cap_and_mode(
device: &GpuDevice,
state_cap: usize,
mode: GemmMode,
) -> Result<Self, String> {
Self::new_with_state_cap_mode_and_role(
device,
state_cap,
mode,
GemmRole::triad(WeightDtype::F32),
)
}
pub(crate) fn new_from_env_with_state_cap_and_role(
device: &GpuDevice,
state_cap: usize,
role: GemmRole,
) -> Result<Self, String> {
let config = resolve_gemm_env(GemmEnvValues::read(), role)?;
Self::new_with_state_cap_and_config(device, state_cap, config)
}
pub(crate) fn new_with_state_cap_mode_and_role(
device: &GpuDevice,
state_cap: usize,
mode: GemmMode,
role: GemmRole,
) -> Result<Self, String> {
Self::new_with_state_cap_and_config(device, state_cap, explicit_gemm_config(mode, role))
}
fn new_with_state_cap_and_config(
device: &GpuDevice,
state_cap: usize,
config: ResolvedGemmEnv,
) -> Result<Self, String> {
unsafe {
device.context().disable_event_tracking();
}
let stream = device.fork_stream()?;
let arch = device.nvrtc_target();
let kernels = MambaKernels::compile_with_state_cap(device.context(), arch, state_cap)?;
let device_identity = device.identity();
validate_multiprocessor_identity(
device_identity.multiprocessor_count,
kernels.multiprocessor_count(),
)?;
device
.default_stream()
.synchronize()
.map_err(|e| format!("default-stream drain after kernel compile: {e:?}"))?;
let (blas, ws) = device.create_cublas(&stream, config.mode)?;
let instance_token = next_gpu_ctx_token()?;
let compiler = kernels.compiler_identity();
let optin_shared_bytes = device
.context()
.attribute(
cudarc::driver::sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_MAX_SHARED_MEMORY_PER_BLOCK_OPTIN,
)
.map_err(|error| format!("query opt-in shared memory: {error:?}"))?;
let tensor_map_access = device
.context()
.attribute(
cudarc::driver::sys::CUdevice_attribute::CU_DEVICE_ATTRIBUTE_TENSOR_MAP_ACCESS_SUPPORTED,
)
.map(|value| value != 0)
.unwrap_or(false);
let device_caps = super::kernel_identity::DeviceCaps {
compute_capability: device.compute_capability,
nvrtc_version: compiler.nvrtc_version,
accepted_target: kernels
.specialized_compiler_identity()
.map(|identity| identity.target),
optin_shared_bytes: u32::try_from(optin_shared_bytes)
.map_err(|_| format!("negative opt-in shared memory {optin_shared_bytes}"))?,
tensor_map_access,
};
let kernels = Arc::new(kernels);
Ok(Self {
resources: Rc::new(GpuCtxResources {
stream,
kernels,
blas,
_blas_workspace: Arc::new(ws),
half_staging: RefCell::new(None),
half_staging_ptr: RefCell::new(0),
half_staging_bytes: RefCell::new(0),
bi_upcast_scratch: [RefCell::new(None), RefCell::new(None), RefCell::new(None)],
}),
gemm_route_recorder: RefCell::new(None),
f32_prepared_launches: RefCell::new(F32PreparedLaunchCache::default()),
route_proofs: RefCell::new(super::gemm_bi_triad::proof::RouteProofLedger::default()),
sm120_prepared_launches: RefCell::new(Sm120PreparedLaunchCache::default()),
sm100_prepared_launches: RefCell::new(Sm100PreparedLaunchCache::default()),
sm90a_prepared_launches: RefCell::new(Sm90aPreparedLaunchCache::default()),
fixed_tf32_maps: RefCell::new(super::gemm_bi_inference::FixedTf32MapCache::default()),
fixed_postbias_maps: RefCell::new(
super::gemm_bi_inference::FixedPostBiasMapCache::default(),
),
fixed_half_maps: RefCell::new(super::gemm_bi_inference::FixedHalfMapCache::default()),
gemm_mode: Cell::new(config.mode),
gemm_unusable: RefCell::new(None),
bi_tensor_cores: Cell::new(config.tensor_cores),
bi_gemm_family: Cell::new(config.family),
f32_triad_policy: Cell::new(config.f32_policy),
half_triad_policy: Cell::new(config.half_policy),
state_cap,
instance_token,
device_identity,
device_caps,
policy_hash: super::kernel_identity::gemm_dispatch_policy_digest(
device_identity.multiprocessor_count,
),
graphs_captured: Cell::new(0),
graph_scratch_frozen: Cell::new(false),
})
}
pub(crate) fn resource_anchor(&self) -> Rc<GpuCtxResources> {
self.resources.clone()
}
pub(crate) fn instance_token(&self) -> u64 {
self.instance_token
}
pub(crate) fn stream_token(&self) -> usize {
Arc::as_ptr(&self.stream) as usize
}
pub(crate) fn freeze_graph_scratch(&self) {
self.graph_scratch_frozen.set(true);
}
pub(crate) fn with_bi_upcast_scratch<R>(
&self,
elems: (usize, usize, usize),
f: impl FnOnce(
&mut super::buffers::GpuBuffer,
&mut super::buffers::GpuBuffer,
&mut super::buffers::GpuBuffer,
) -> Result<R, String>,
) -> Result<R, String> {
let sizes = [elems.0, elems.1, elems.2];
let needs_growth = self
.bi_upcast_scratch
.iter()
.zip(&sizes)
.any(|(cell, &need)| cell.borrow().as_ref().map_or(0, |b| b.len()) < need.max(1));
if needs_growth && self.graph_scratch_frozen.get() {
return Err(
"batch-invariant upcast scratch cannot grow after CUDA graph capture; \
destroy the context or pre-size the largest shape before capture"
.into(),
);
}
for (cell, &need) in self.bi_upcast_scratch.iter().zip(&sizes) {
let mut slot = cell.borrow_mut();
let have = slot.as_ref().map_or(0, |b| b.len());
let need = need.max(1);
if have < need {
*slot = Some(super::buffers::GpuBuffer::zeros(&self.stream, need)?);
}
}
let mut a = self.bi_upcast_scratch[0].borrow_mut();
let mut b = self.bi_upcast_scratch[1].borrow_mut();
let mut c = self.bi_upcast_scratch[2].borrow_mut();
f(
a.as_mut().expect("bi_upcast_scratch[0] sized above"),
b.as_mut().expect("bi_upcast_scratch[1] sized above"),
c.as_mut().expect("bi_upcast_scratch[2] sized above"),
)
}
pub(crate) fn bi_upcast_scratch_ptrs(&self) -> [cudarc::driver::sys::CUdeviceptr; 3] {
let p = |i: usize| {
self.bi_upcast_scratch[i]
.borrow()
.as_ref()
.map_or(0, |b| b.cached_ptr())
};
[p(0), p(1), p(2)]
}
pub(crate) fn ensure_graph_scratch_ptrs(
&self,
half_staging: cudarc::driver::sys::CUdeviceptr,
bi_upcast: [cudarc::driver::sys::CUdeviceptr; 3],
label: &str,
) -> Result<(), String> {
if self.half_staging_ptr() != half_staging || self.bi_upcast_scratch_ptrs() != bi_upcast {
return Err(format!(
"{label}: graph-visible staging scratch changed since capture"
));
}
Ok(())
}
pub fn presize_bi_upcast_scratch_for_train(
&self,
cfg: &MambaConfig,
batch: usize,
seq_len: usize,
dtype: WeightDtype,
) -> Result<(), String> {
if matches!(dtype, WeightDtype::F32) || !self.batch_invariant() {
return Ok(());
}
let m = batch * seq_len;
let dm = cfg.d_model;
let di = cfg.d_inner();
let xproj_out = cfg.dt_rank() + 2 * cfg.d_state;
let max_dim = dm.max(2 * di).max(xproj_out);
let max_kn = (dm * 2 * di)
.max(di * xproj_out)
.max(cfg.dt_rank() * di)
.max(di * dm);
let elems = (m * max_dim).max(max_kn);
self.with_bi_upcast_scratch((elems, elems, elems), |_, _, _| Ok(()))
}
pub fn presize_bi_upcast_scratch_for_train_with_input(
&self,
cfg: &MambaConfig,
batch: usize,
seq_len: usize,
input_dim: usize,
dtype: WeightDtype,
) -> Result<(), String> {
if matches!(dtype, WeightDtype::F32) || !self.batch_invariant() {
return Ok(());
}
let m = batch * seq_len;
let dm = cfg.d_model;
let di = cfg.d_inner();
let xproj_out = cfg.dt_rank() + 2 * cfg.d_state;
let max_dim = dm.max(2 * di).max(xproj_out).max(input_dim);
let max_kn = (dm * 2 * di)
.max(di * xproj_out)
.max(cfg.dt_rank() * di)
.max(di * dm)
.max(input_dim * dm);
let elems = (m * max_dim).max(max_kn);
self.with_bi_upcast_scratch((elems, elems, elems), |_, _, _| Ok(()))
}
pub fn presize_bi_upcast_scratch_for_train_m3(
&self,
cfg: &crate::mamba3_siso::config::Mamba3Config,
batch: usize,
seq_len: usize,
input_dim: usize,
dtype: WeightDtype,
) -> Result<(), String> {
if matches!(dtype, WeightDtype::F32) || !self.batch_invariant() {
return Ok(());
}
let m = batch * seq_len;
let dm = cfg.d_model;
let di = cfg.d_inner();
let ip = cfg.in_proj_out_dim();
let max_dim = dm.max(ip).max(di).max(input_dim);
let max_kn = (dm * ip).max(di * dm).max(input_dim * dm);
let elems = (m * max_dim).max(max_kn);
self.with_bi_upcast_scratch((elems, elems, elems), |_, _, _| Ok(()))
}
pub(crate) fn presize_mixed_graph_scratch_m1(
&self,
dims: &super::forward::GpuMambaDims,
dtype: WeightDtype,
) -> Result<(), String> {
if matches!(dtype, WeightDtype::F32) {
return Ok(());
}
let max_dim = m1_mixed_graph_max_dim(dims);
self.ensure_half_staging(dims.bt() * max_dim * dtype.size_bytes())?;
if self.batch_invariant() {
let max_kn = (dims.d_model * 2 * dims.d_inner)
.max(dims.d_inner * dims.xdbl_dim)
.max(dims.dt_rank * dims.d_inner)
.max(dims.d_inner * dims.d_model)
.max(dims.mamba_input_dim * dims.d_model);
let elems = (dims.bt() * max_dim).max(max_kn);
self.with_bi_upcast_scratch((elems, elems, elems), |_, _, _| Ok(()))?;
}
Ok(())
}
pub(crate) fn presize_mixed_graph_scratch_m3(
&self,
dims: &crate::mamba3_siso::gpu::state::GpuMamba3Dims,
dtype: WeightDtype,
) -> Result<(), String> {
if matches!(dtype, WeightDtype::F32) {
return Ok(());
}
let max_dim = dims
.d_model
.max(dims.d_inner)
.max(dims.in_proj_dim)
.max(dims.mamba_input_dim);
self.ensure_half_staging(dims.bt() * max_dim * dtype.size_bytes())?;
if self.batch_invariant() {
let max_kn = (dims.d_model * dims.in_proj_dim)
.max(dims.d_inner * dims.d_model)
.max(dims.mamba_input_dim * dims.d_model);
let elems = (dims.bt() * max_dim).max(max_kn);
self.with_bi_upcast_scratch((elems, elems, elems), |_, _, _| Ok(()))?;
}
Ok(())
}
pub fn set_gemm_mode(&self, mode: GemmMode) -> Result<(), String> {
self.ensure_gemm_usable()?;
if mode == self.gemm_mode.get() {
return Ok(());
}
if self
.gemm_route_recorder
.try_borrow()
.map_err(|_| "GEMM route recorder is already borrowed".to_string())?
.is_some()
{
return Err("cannot change GEMM mode while GEMM route recording is active".into());
}
let capture_status = self
.stream
.capture_status()
.map_err(|error| format!("query CUDA stream capture state: {error:?}"))?;
if capture_status
!= cudarc::driver::sys::CUstreamCaptureStatus::CU_STREAM_CAPTURE_STATUS_NONE
{
return Err(format!(
"cannot change GEMM mode while CUDA stream capture state is {capture_status:?}"
));
}
let mut backend = CublasMathBackend { blas: &self.blas };
match change_math_mode(&mut backend, mode.cublas_math()) {
Ok(()) => {
if self.graphs_captured.get() > 0 {
eprintln!(
"mamba-rs WARNING: GEMM route changed after graph capture; replay will reject it"
);
}
self.gemm_mode.set(mode);
Ok(())
}
Err(MathTransitionError::Recoverable(error)) => Err(error),
Err(MathTransitionError::Unusable(error)) => {
let diagnostic = format!(
"GPU context is unusable after an unverified cuBLAS math rollback: {error}"
);
*self.gemm_unusable.borrow_mut() = Some(diagnostic.clone());
Err(diagnostic)
}
}
}
pub fn gemm_mode(&self) -> GemmMode {
self.gemm_mode.get()
}
pub(crate) fn ensure_gemm_usable(&self) -> Result<(), String> {
match self.gemm_unusable.borrow().as_ref() {
Some(error) => Err(error.clone()),
None => Ok(()),
}
}
#[cfg(test)]
pub(crate) fn poison_gemm_for_test(&self) {
*self.gemm_unusable.borrow_mut() = Some("test: GPU GEMM context is unusable".into());
}
pub(crate) fn ensure_vendor_gemm(&self, label: &str) -> Result<GemmMode, String> {
self.ensure_gemm_usable()?;
let mode = self.gemm_mode();
if mode == GemmMode::Deterministic {
Err(format!(
"{label}: deterministic GEMM mode reached a cuBLAS dispatch boundary"
))
} else {
Ok(mode)
}
}
pub(crate) fn batch_invariant(&self) -> bool {
self.gemm_mode().batch_invariant()
}
pub(crate) fn set_bi_gemm_family(&self, family: BiGemmFamily) {
if self.graphs_captured.get() > 0 {
eprintln!(
"mamba-rs WARNING: GEMM route changed after graph capture; replay will reject it"
);
}
self.bi_gemm_family.set(family);
}
pub(crate) fn bi_gemm_family(&self) -> BiGemmFamily {
self.bi_gemm_family.get()
}
pub(crate) fn set_bi_tensor_cores(&self, on: bool) {
if self.graphs_captured.get() > 0 {
eprintln!(
"mamba-rs WARNING: GEMM route changed after graph capture; replay will reject it"
);
}
self.bi_tensor_cores.set(on);
}
pub(crate) fn fast_gemm(&self) -> bool {
self.gemm_mode().fast_gemm()
}
pub(crate) fn note_graph_capture(&self) {
self.graphs_captured.set(self.graphs_captured.get() + 1);
}
pub(crate) fn begin_gemm_route_recording(
&self,
capacity: usize,
) -> Result<GemmRouteRecordingGuard<'_>, String> {
self.ensure_gemm_usable()?;
let mut active = self
.gemm_route_recorder
.try_borrow_mut()
.map_err(|_| "GEMM route recorder is already borrowed".to_string())?;
if active.is_some() {
return Err("nested GEMM route recording is not supported".into());
}
let route_capacity = u32::try_from(capacity)
.map_err(|_| "GEMM route recording capacity exceeds u32::MAX".to_string())?;
let mut routes = Vec::new();
routes
.try_reserve_exact(capacity)
.map_err(|error| format!("reserve GEMM route recording capacity: {error}"))?;
let recorder = GemmRouteRecorder {
context: self.gemm_route(),
mode: GemmRouteRecorderMode::FixedCapture { route_capacity },
routes,
};
*active = Some(recorder);
Ok(GemmRouteRecordingGuard { ctx: self })
}
fn begin_eager_gemm_route_recording(&self) -> Result<GemmRouteRecordingGuard<'_>, String> {
self.ensure_gemm_usable()?;
let mut active = self
.gemm_route_recorder
.try_borrow_mut()
.map_err(|_| "GEMM route recorder is already borrowed".to_string())?;
if active.is_some() {
return Err("nested GEMM route recording is not supported".into());
}
*active = Some(GemmRouteRecorder {
context: self.gemm_route(),
mode: GemmRouteRecorderMode::GrowableEager,
routes: Vec::new(),
});
Ok(GemmRouteRecordingGuard { ctx: self })
}
pub fn record_eager_gemm_trace<F>(&self, body: F) -> Result<RecordedGemmTrace, String>
where
F: FnOnce() -> Result<(), String>,
{
let recording = self.begin_eager_gemm_route_recording()?;
body()?;
recording.finish_trace()
}
pub fn record_eager_gemm_manifest<F>(
&self,
body: F,
) -> Result<PreparedGemmCaptureManifest, String>
where
F: FnOnce() -> Result<(), String>,
{
self.record_eager_gemm_trace(body)
.map(|trace| trace.manifest())
}
pub(crate) fn with_f32_prepared_launches<T>(
&self,
access: impl FnOnce(&mut F32PreparedLaunchCache) -> Result<T, String>,
) -> Result<T, String> {
let mut launches = self
.f32_prepared_launches
.try_borrow_mut()
.map_err(|_| "prepared f32 Triad cache is already borrowed".to_string())?;
access(&mut launches)
}
pub(crate) fn with_route_proofs<T>(
&self,
access: impl FnOnce(&mut super::gemm_bi_triad::proof::RouteProofLedger) -> T,
) -> Result<T, String> {
let mut proofs = self
.route_proofs
.try_borrow_mut()
.map_err(|_| "route proof ledger is already borrowed".to_string())?;
Ok(access(&mut proofs))
}
pub(crate) fn with_sm120_prepared_launches<T>(
&self,
access: impl FnOnce(&mut Sm120PreparedLaunchCache) -> Result<T, String>,
) -> Result<T, String> {
let mut launches = self
.sm120_prepared_launches
.try_borrow_mut()
.map_err(|_| "prepared SM120 TMA cache is already borrowed".to_string())?;
access(&mut launches)
}
pub(crate) fn with_sm90a_prepared_launches<T>(
&self,
access: impl FnOnce(&mut Sm90aPreparedLaunchCache) -> Result<T, String>,
) -> Result<T, String> {
let mut launches = self
.sm90a_prepared_launches
.try_borrow_mut()
.map_err(|_| "prepared SM90a WGMMA cache is already borrowed".to_string())?;
access(&mut launches)
}
pub(crate) fn with_sm100_prepared_launches<T>(
&self,
access: impl FnOnce(&mut Sm100PreparedLaunchCache) -> Result<T, String>,
) -> Result<T, String> {
let mut launches = self
.sm100_prepared_launches
.try_borrow_mut()
.map_err(|_| "prepared SM100 TCGEN cache is already borrowed".to_string())?;
access(&mut launches)
}
pub(crate) fn gemm_route_recording_active(&self) -> Result<bool, String> {
self.gemm_route_recorder
.try_borrow()
.map(|recorder| recorder.is_some())
.map_err(|_| "GEMM route recorder is already borrowed".to_string())
}
pub(crate) fn record_resolved_gemm_route(
&self,
route: ResolvedGemmRoute,
) -> Result<(), String> {
let mut active = self
.gemm_route_recorder
.try_borrow_mut()
.map_err(|_| "GEMM route recorder is already borrowed".to_string())?;
let Some(recorder) = active.as_mut() else {
return Ok(());
};
if let GemmRouteRecorderMode::FixedCapture { route_capacity } = recorder.mode {
let route_count = u32::try_from(recorder.routes.len())
.expect("fixed GEMM route count is bounded by u32 capacity");
if route_count >= route_capacity {
return Err(format!(
"GEMM route recording exceeded its capacity {route_capacity}"
));
}
}
recorder.routes.push(route);
Ok(())
}
pub(crate) fn with_gemm_route_recording_suspended<T>(
&self,
body: impl FnOnce() -> T,
) -> Result<T, String> {
struct Restore<'a> {
ctx: &'a GpuCtx,
suspended: Option<GemmRouteRecorder>,
}
impl Drop for Restore<'_> {
fn drop(&mut self) {
if let Ok(mut active) = self.ctx.gemm_route_recorder.try_borrow_mut()
&& active.is_none()
{
*active = self.suspended.take();
}
}
}
let suspended = self
.gemm_route_recorder
.try_borrow_mut()
.map_err(|_| "GEMM route recorder is already borrowed".to_string())?
.take();
let restore = Restore {
ctx: self,
suspended,
};
let result = body();
drop(restore);
Ok(result)
}
fn take_gemm_route_recording(&self) -> Result<GemmRouteRecorder, String> {
self.gemm_route_recorder
.try_borrow_mut()
.map_err(|_| "GEMM route recorder is already borrowed".to_string())?
.take()
.ok_or_else(|| "GEMM route recorder is not active".to_string())
}
fn clear_gemm_route_recording(&self) {
if let Ok(mut active) = self.gemm_route_recorder.try_borrow_mut() {
*active = None;
}
}
fn live_gemm_module_binding(
&self,
module_kind: ModuleKind,
) -> Option<(ArtifactIdentity, CompilerIdentity)> {
let artifacts = self.kernels.artifact_set_identity();
match module_kind {
ModuleKind::Fixed => Some((artifacts.fixed, self.kernels.compiler_identity())),
ModuleKind::TriadScalar => Some((
artifacts.triad_scalar,
self.kernels.triad_scalar_compiler_identity(),
)),
ModuleKind::TriadSm80 => Some((
self.kernels.artifact_set_identity().triad_sm80,
self.kernels.triad_sm80_compiler_identity(),
)),
ModuleKind::TriadSm89Finalist => artifacts
.sm89_finalist
.zip(self.kernels.triad_sm89_finalist_compiler_identity()),
ModuleKind::TriadSm89Half => artifacts
.sm89_half
.zip(self.kernels.triad_sm89_half_compiler_identity()),
ModuleKind::TriadSm89ExactF32 => artifacts
.sm89_exact_f32
.zip(self.kernels.triad_sm89_exact_f32_compiler_identity()),
ModuleKind::TriadSm89ExactF32D128 => artifacts
.sm89_exact_f32_d128
.zip(self.kernels.triad_sm89_exact_f32_d128_compiler_identity()),
ModuleKind::TriadSm89Tf32Joint => artifacts
.sm89_tf32_joint
.zip(self.kernels.triad_sm89_tf32_joint_compiler_identity()),
ModuleKind::InferenceSm89Cells => artifacts
.sm89_cells
.zip(self.kernels.sm89_cells_compiler_identity()),
ModuleKind::TriadSm90a | ModuleKind::TriadSm100 | ModuleKind::TriadSm120 => self
.kernels
.artifact_set_identity()
.specialized
.filter(|artifact| artifact.module_kind == module_kind)
.zip(self.kernels.specialized_compiler_identity()),
ModuleKind::Mamba3Combined => None,
}
}
fn live_qualified_tf32_binding(
&self,
module_kind: ModuleKind,
) -> Option<super::gemm_bi_triad::Tf32QualifiedModule> {
let availability = self.kernels.f32_triad_availability();
match module_kind {
ModuleKind::TriadSm80 => availability.portable,
ModuleKind::TriadSm89Finalist => availability.finalist,
ModuleKind::TriadSm89Half => None,
ModuleKind::TriadSm89ExactF32 => None,
ModuleKind::TriadSm89ExactF32D128 => None,
ModuleKind::TriadSm89Tf32Joint => availability.joint,
ModuleKind::TriadSm90a | ModuleKind::TriadSm100 | ModuleKind::TriadSm120 => {
availability.specialized
}
_ => None,
}
.filter(|binding| binding.module_kind == module_kind)
}
pub(crate) fn validate_resolved_gemm_route(
&self,
route: &ResolvedGemmRoute,
label: &str,
) -> Result<(), String> {
self.validate_resolved_gemm_route_in(&self.gemm_route(), route, label)
}
pub(crate) fn validate_resolved_gemm_route_in(
&self,
context: &GemmRouteIdentity,
route: &ResolvedGemmRoute,
label: &str,
) -> Result<(), String> {
let context = *context;
let inference_terminal = matches!(
route.backend,
PhysicalGemmBackend::InferenceScalarFma
| PhysicalGemmBackend::InferenceWmma
| PhysicalGemmBackend::InferenceMma16
| PhysicalGemmBackend::InferenceSm90aWgmma
| PhysicalGemmBackend::InferenceSm100Tcgen05
| PhysicalGemmBackend::InferenceMmaTf32Rna
| PhysicalGemmBackend::InferenceSm120TmaFma
| PhysicalGemmBackend::InferenceSm120TmaMma16
| PhysicalGemmBackend::InferenceSm120TmaMmaTf32Rna
| PhysicalGemmBackend::FixedMatvecEightWarp
) || (route.backend
== PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan
&& route.numeric_contract == ResolvedNumericContract::ScalarFmaPostDotBias)
|| (context.policy.bi_gemm_family == BiGemmFamily::Inference
&& route.backend == PhysicalGemmBackend::MmaTf32Rna
&& route.symbol == "nn_sm80_mma_tf32_m128n128_bk32_s3");
if inference_terminal {
return self.validate_inference_terminal_route(&context, route, label);
}
if context.policy.bi_gemm_family == BiGemmFamily::Inference
&& route.backend == PhysicalGemmBackend::Sm120TmaFmaExact
{
super::gemm_bi_inference::identity::validate_cached_bridge(route)?;
if self.kernels.tf32_function(route.symbol).is_none() {
return Err(format!(
"{label}: prepared exact-TMA bridge holder is not loaded"
));
}
}
if route.backend == PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan
&& (route.numeric_contract != ResolvedNumericContract::ScalarFma
|| context.policy.bi_gemm_family != BiGemmFamily::Triad
|| route.symbol != "nn_sm89_f32_n64_copyplan"
|| self
.kernels
.inference_terminal_function(route.symbol)
.is_none())
{
return Err(format!(
"{label}: Fixed copy-plan backend/family/numeric binding changed"
));
}
if !context.policy.batch_invariant || !context.backend_set.contains(BackendSet::TRIAD) {
return Err(format!(
"{label}: captured Triad backend is unavailable under the live GEMM policy"
));
}
let required_contract = match route.numeric_contract {
ResolvedNumericContract::ScalarFmaPostDotBias
| ResolvedNumericContract::WmmaF32PostDotBias
| ResolvedNumericContract::ScalarFmaEightWarpTreePostDotBias => {
return Err(format!(
"{label}: Inference numeric contract has an incompatible backend"
));
}
ResolvedNumericContract::ScalarFma
| ResolvedNumericContract::ScalarFmaSplitKPartial
| ResolvedNumericContract::ScalarFmaSplitKF32Reduce
| ResolvedNumericContract::ScalarFmaTnNarrowSplitMPartial
| ResolvedNumericContract::ScalarFmaTnNarrowSplitMF64Reduce
| ResolvedNumericContract::ScalarFmaTnSplitMF64Reduce
| ResolvedNumericContract::ScalarFmaTnSplitMPartial
| ResolvedNumericContract::ScalarFmaFixedSplitFold
| ResolvedNumericContract::ZeroReductionEpilogueF32 => {
NumericContractSet::TRIAD_SCALAR_FMA
}
ResolvedNumericContract::MmaSyncF32
| ResolvedNumericContract::WgmmaF32
| ResolvedNumericContract::Tcgen05F32 => NumericContractSet::TRIAD_MMA_SYNC,
ResolvedNumericContract::MmaSyncF32StreamKFixedOrder => {
NumericContractSet::TRIAD_MMA_SYNC_STREAM_K
}
ResolvedNumericContract::MmaTf32Rna
| ResolvedNumericContract::MmaTf32PreRnaAV1
| ResolvedNumericContract::MmaTf32AddHalfUlp
| ResolvedNumericContract::Tf32RnaPreprocess
| ResolvedNumericContract::Sm90aWgmmaTf32Tma
| ResolvedNumericContract::Sm100Tcgen05Tf32Tma
| ResolvedNumericContract::Sm120TmaMmaTf32Rna => {
NumericContractSet::TRIAD_DETERMINISTIC_TF32
}
ResolvedNumericContract::MmaTf32RnaSplitK2
| ResolvedNumericContract::MmaTf32RnaSplitK4
| ResolvedNumericContract::MmaTf32RnaSplitK8
| ResolvedNumericContract::Sm120TmaMmaTf32RnaStreamKV1 => {
NumericContractSet::TRIAD_DETERMINISTIC_TF32_SPLIT_K
}
};
if !context.numeric_contracts.contains(required_contract) {
return Err(format!(
"{label}: captured Triad numeric contract is unavailable under the live GEMM policy"
));
}
let expected_module = expected_route_module(route.backend);
if route.module_kind != expected_module || route.artifact.module_kind != expected_module {
return Err(format!(
"{label}: captured physical backend no longer matches its module binding"
));
}
let uses_qualified_tf32_module = matches!(
route.numeric_contract,
super::kernel_identity::ResolvedNumericContract::MmaTf32Rna
| super::kernel_identity::ResolvedNumericContract::MmaTf32PreRnaAV1
| super::kernel_identity::ResolvedNumericContract::MmaTf32AddHalfUlp
| super::kernel_identity::ResolvedNumericContract::MmaTf32RnaSplitK2
| super::kernel_identity::ResolvedNumericContract::MmaTf32RnaSplitK4
| super::kernel_identity::ResolvedNumericContract::MmaTf32RnaSplitK8
| super::kernel_identity::ResolvedNumericContract::Sm90aWgmmaTf32Tma
| super::kernel_identity::ResolvedNumericContract::Sm100Tcgen05Tf32Tma
| super::kernel_identity::ResolvedNumericContract::Sm120TmaMmaTf32Rna
| super::kernel_identity::ResolvedNumericContract::Sm120TmaMmaTf32RnaStreamKV1
| super::kernel_identity::ResolvedNumericContract::ZeroReductionEpilogueF32
) && route.module_kind != ModuleKind::TriadScalar
|| route.backend == super::kernel_identity::PhysicalGemmBackend::Sm120TmaFmaExact;
if uses_qualified_tf32_module {
let binding = self
.live_qualified_tf32_binding(route.module_kind)
.ok_or_else(|| format!("{label}: captured qualified TF32 module is not loaded"))?;
if route.artifact != binding.artifact
|| route.compiler != binding.compiler
|| route.target != binding.target
|| route.device != binding.device
|| route.device_caps != binding.device_caps
{
return Err(format!(
"{label}: captured GEMM route no longer matches its qualified TF32 binding"
));
}
} else {
let (artifact, compiler) = self
.live_gemm_module_binding(route.module_kind)
.ok_or_else(|| format!("{label}: captured GEMM module is not loaded"))?;
if route.artifact != artifact
|| route.compiler != compiler
|| route.target != compiler.target
|| route.device != context.device
|| route.device_caps != context.device_caps
{
return Err(format!(
"{label}: captured GEMM route no longer matches its live module binding"
));
}
}
if route.tuning_table_revision
!= expected_route_tuning_revision(route.backend, context.tuning_table_revision)
|| route.schedule_revision
!= expected_route_schedule_revision(route.backend, context.schedule_set_revision)
{
return Err(format!("{label}: captured GEMM route revision is stale"));
}
let tf32_numeric = matches!(
route.numeric_contract,
super::kernel_identity::ResolvedNumericContract::MmaTf32Rna
| super::kernel_identity::ResolvedNumericContract::MmaTf32PreRnaAV1
| super::kernel_identity::ResolvedNumericContract::MmaTf32AddHalfUlp
| super::kernel_identity::ResolvedNumericContract::MmaTf32RnaSplitK2
| super::kernel_identity::ResolvedNumericContract::MmaTf32RnaSplitK4
| super::kernel_identity::ResolvedNumericContract::MmaTf32RnaSplitK8
| super::kernel_identity::ResolvedNumericContract::Sm90aWgmmaTf32Tma
| super::kernel_identity::ResolvedNumericContract::Sm100Tcgen05Tf32Tma
| super::kernel_identity::ResolvedNumericContract::Sm120TmaMmaTf32Rna
| super::kernel_identity::ResolvedNumericContract::Sm120TmaMmaTf32RnaStreamKV1
);
if tf32_numeric
&& (route.dtype != PolicyDtype::F32
|| self.f32_triad_policy() != F32TriadPolicy::AllowDeterministicTf32)
{
return Err(format!(
"{label}: captured deterministic TF32 route is disabled by the live policy"
));
}
if route.dtype == PolicyDtype::F32
&& !tf32_numeric
&& route.numeric_contract
!= super::kernel_identity::ResolvedNumericContract::ZeroReductionEpilogueF32
&& !scalar_backend_supports_logical_f32(route.backend)
{
return Err(format!(
"{label}: captured logical-f32 route has an incompatible physical backend"
));
}
if route.numeric_contract
== super::kernel_identity::ResolvedNumericContract::ZeroReductionEpilogueF32
&& (route.dtype != PolicyDtype::F32
|| route.instruction_family
!= super::kernel_identity::ResolvedInstructionFamily::ScalarFma
|| route.instruction_shape
!= super::kernel_identity::ResolvedInstructionShape { m: 1, n: 1, k: 1 }
|| route.operand_conversion
!= super::kernel_identity::ResolvedOperandConversion::None
|| match route.op {
super::kernel_identity::ResolvedGemmOp::Nn => route.shape.1 != 0,
super::kernel_identity::ResolvedGemmOp::Tn => route.shape.0 != 0,
super::kernel_identity::ResolvedGemmOp::Nt => route.shape.2 != 0,
})
{
return Err(format!(
"{label}: captured zero-reduction epilogue identity is inconsistent"
));
}
if route.dtype != PolicyDtype::F32
&& route.backend != PhysicalGemmBackend::ScalarFma
&& !self.bi_tensor_cores()
{
return Err(format!(
"{label}: captured typed Tensor Core route is disabled by the live policy"
));
}
Ok(())
}
fn validate_inference_terminal_route(
&self,
context: &GemmRouteIdentity,
route: &ResolvedGemmRoute,
label: &str,
) -> Result<(), String> {
let spec = super::gemm_bi_inference::identity::terminal(route.symbol)
.ok_or_else(|| format!("{label}: unknown Inference terminal"))?;
spec.validate_route(route)?;
let required = spec.required_contract(context.policy.bi_gemm_family)?;
let backend = if route.module_kind == ModuleKind::Fixed {
BackendSet::FIXED
} else {
BackendSet::TRIAD
};
if !context.policy.batch_invariant
|| !context.backend_set.contains(backend)
|| !context.numeric_contracts.contains(required)
|| (required == NumericContractSet::FIXED_DETERMINISTIC_TF32
&& self.f32_triad_policy() != F32TriadPolicy::AllowDeterministicTf32)
{
return Err(format!(
"{label}: Inference terminal is disabled by the live policy"
));
}
let (artifact, compiler) = self
.live_gemm_module_binding(route.module_kind)
.ok_or_else(|| format!("{label}: Inference terminal module is not loaded"))?;
if route.artifact != artifact
|| route.compiler != compiler
|| route.target != compiler.target
|| route.device != context.device
|| route.device_caps != context.device_caps
|| self
.kernels
.inference_terminal_function(route.symbol)
.is_none()
{
return Err(format!(
"{label}: Inference terminal live function/module binding changed"
));
}
Ok(())
}
pub(crate) fn validate_resolved_input_transform(
&self,
symbol: &'static str,
transform: &super::kernel_identity::ResolvedInputTransform,
label: &str,
) -> Result<(), String> {
use super::kernel_identity::{ResolvedOperandConversion, ResolvedTransformOutputOwnership};
if self.f32_triad_policy() != F32TriadPolicy::AllowDeterministicTf32
|| transform.numeric_contract != ResolvedNumericContract::Tf32RnaPreprocess
|| transform.operand_conversion != ResolvedOperandConversion::RegisterCvtRnaTf32F32
|| transform.output_ownership
!= ResolvedTransformOutputOwnership::PreparedScratchAllocation
|| symbol != super::gemm_bi_triad::TN_PRE_RNA_TRANSPOSE_SYMBOL
{
return Err(format!(
"{label}: captured input transform contract is unavailable"
));
}
let binding = self
.live_qualified_tf32_binding(ModuleKind::TriadSm89Tf32Joint)
.ok_or_else(|| format!("{label}: captured input transform module is not loaded"))?;
if transform.artifact != binding.artifact
|| transform.compiler != binding.compiler
|| transform.target != binding.target
|| transform.device != binding.device
|| transform.device_caps != binding.device_caps
|| transform.tuning_table_revision
!= super::gemm_bi_triad::SM89_TF32_JOINT_TUNING_REVISION
|| transform.schedule_revision != super::kernel_identity::SCHEDULE_REVISION
|| transform.resources_digest == [0; 32]
|| self
.kernels
.triad_sm89_tf32_joint_function(symbol)
.is_none()
{
return Err(format!(
"{label}: captured input transform no longer matches its live module or resources"
));
}
Ok(())
}
pub(crate) fn gemm_policy(&self) -> GemmPolicy {
GemmPolicy {
batch_invariant: self.batch_invariant(),
bi_tensor_cores: self.bi_tensor_cores.get(),
fast_gemm: self.fast_gemm(),
cublas_tf32: self.tf32(),
f32_triad_policy: self.f32_triad_policy.get(),
half_triad_policy: self.half_triad_policy.get(),
bi_gemm_family: self.bi_gemm_family.get(),
}
}
pub fn gemm_route(&self) -> GemmRoute {
let policy = self.gemm_policy();
let (backend_set, numeric_contracts) =
super::kernel_identity::route_backend_contract_sets(policy);
GemmRouteIdentity {
policy,
backend_set,
numeric_contracts,
compiler: self.kernels.compiler_identity(),
artifacts: self.kernels.artifact_set_identity(),
policy_revision: super::kernel_identity::POLICY_REVISION,
policy_hash: self.policy_hash,
device: self.device_identity,
device_caps: self.device_caps,
tuning_table_revision: super::kernel_identity::TUNING_TABLE_REVISION,
schedule_set_revision: super::kernel_identity::SCHEDULE_REVISION,
state_capacity: u32::try_from(self.state_cap)
.expect("validated state capacity fits in u32"),
}
}
pub(crate) fn bi_tensor_cores(&self) -> bool {
self.bi_tensor_cores.get()
}
pub(crate) fn set_f32_triad_policy(&self, policy: F32TriadPolicy) {
if self.graphs_captured.get() > 0 {
eprintln!(
"mamba-rs WARNING: GEMM route changed after graph capture; replay will reject it"
);
}
self.f32_triad_policy.set(policy);
}
pub(crate) fn f32_triad_policy(&self) -> F32TriadPolicy {
self.f32_triad_policy.get()
}
pub(crate) fn f32_storage_dtype(&self) -> WeightDtype {
match self.f32_triad_policy.get() {
F32TriadPolicy::AllowDeterministicTf32 => WeightDtype::Tf32,
F32TriadPolicy::ExactScalarFma => WeightDtype::F32,
}
}
pub(crate) fn set_half_triad_policy(&self, policy: HalfTriadPolicy) {
if self.graphs_captured.get() > 0 {
eprintln!(
"mamba-rs WARNING: GEMM route changed after graph capture; replay will reject it"
);
}
self.half_triad_policy.set(policy);
}
pub(crate) fn half_triad_policy(&self) -> HalfTriadPolicy {
self.half_triad_policy.get()
}
pub(crate) fn tf32(&self) -> bool {
self.gemm_mode().tf32()
}
#[doc(hidden)]
pub fn route_controls(&self) -> RouteControls<'_> {
RouteControls { ctx: self }
}
pub fn state_cap(&self) -> usize {
self.state_cap
}
pub(in crate::mamba_ssm::gpu) fn compute_capability(&self) -> (u32, u32) {
self.device_identity.compute_capability
}
pub fn presize_bi_scratch(&self) -> Result<(), String> {
if self.batch_invariant() {
self.kernels.splitk_scratch_buf(&self.stream)?;
self.kernels.transpose_scratch_buf(&self.stream)?;
}
Ok(())
}
pub fn presize_half_staging_for_step(
&self,
cfg: &MambaConfig,
batch: usize,
dtype: WeightDtype,
) -> Result<(), String> {
if matches!(dtype, WeightDtype::F32) {
return Ok(());
}
let dm = cfg.d_model;
let di = cfg.d_inner();
let dt_rank = cfg.dt_rank();
let max_in_elems = batch * dm.max(di).max(dt_rank);
let bytes = max_in_elems * dtype.size_bytes();
self.ensure_half_staging(bytes)
}
pub fn presize_half_staging_for_train_m3(
&self,
cfg: &crate::mamba3_siso::config::Mamba3Config,
batch: usize,
seq_len: usize,
dtype: WeightDtype,
) -> Result<(), String> {
if matches!(dtype, WeightDtype::F32) {
return Ok(());
}
let dm = cfg.d_model;
let di = cfg.d_inner();
let ip = cfg.in_proj_out_dim();
let max_dim = dm.max(di).max(ip);
let bytes = batch * seq_len * max_dim * dtype.size_bytes();
self.ensure_half_staging(bytes)
}
pub fn presize_half_staging_for_train(
&self,
cfg: &MambaConfig,
batch: usize,
seq_len: usize,
dtype: WeightDtype,
) -> Result<(), String> {
if matches!(dtype, WeightDtype::F32) {
return Ok(());
}
let dm = cfg.d_model;
let di = cfg.d_inner();
let dt_rank = cfg.dt_rank();
let max_dim = dm.max(di).max(dt_rank);
let bytes = batch * seq_len * max_dim * dtype.size_bytes();
self.ensure_half_staging(bytes)
}
pub fn ensure_half_staging(&self, bytes: usize) -> Result<(), String> {
let mut cur = self.half_staging_bytes.borrow_mut();
if *cur >= bytes {
return Ok(());
}
if self.graph_scratch_frozen.get() {
return Err(
"half-precision staging cannot grow after CUDA graph capture; destroy the \
context or pre-size the largest shape before capture"
.into(),
);
}
let page = 4096;
let new_size = bytes.div_ceil(page) * page;
let buf = self
.stream
.alloc_zeros::<u8>(new_size)
.map_err(|e| format!("half_staging alloc {new_size}B failed: {e:?}"))?;
let ptr = {
use cudarc::driver::DevicePtr;
let (p, _g) = buf.device_ptr(&self.stream);
p
};
*self.half_staging.borrow_mut() = Some(buf);
*self.half_staging_ptr.borrow_mut() = ptr;
*cur = new_size;
Ok(())
}
pub fn half_staging_ptr(&self) -> cudarc::driver::sys::CUdeviceptr {
*self.half_staging_ptr.borrow()
}
}
fn expected_route_schedule_revision(backend: PhysicalGemmBackend, generic: u16) -> u16 {
if backend == PhysicalGemmBackend::Sm120TmaMma16 {
super::gemm_bi_triad::SM120_SCHEDULE_REVISION
} else {
generic
}
}
const fn expected_route_module(backend: PhysicalGemmBackend) -> ModuleKind {
match backend {
PhysicalGemmBackend::ScalarFma
| PhysicalGemmBackend::ScalarFmaSplitKPartial
| PhysicalGemmBackend::ScalarFmaSplitKF32Reduce
| PhysicalGemmBackend::ScalarFmaTnNarrowSplitMPartial
| PhysicalGemmBackend::ScalarFmaTnSplitMF64Reduce => ModuleKind::TriadScalar,
PhysicalGemmBackend::ScalarFmaSm89ExactF32DualChunkFused
| PhysicalGemmBackend::ScalarFmaSm89ExactF32DirectSplitMPartial => {
ModuleKind::TriadSm89ExactF32
}
PhysicalGemmBackend::ScalarFmaTnDirectF64FoldSm89 => ModuleKind::TriadSm89ExactF32D128,
PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan
| PhysicalGemmBackend::InferenceScalarFma
| PhysicalGemmBackend::InferenceWmma
| PhysicalGemmBackend::InferenceMma16
| PhysicalGemmBackend::InferenceSm90aWgmma
| PhysicalGemmBackend::InferenceSm100Tcgen05
| PhysicalGemmBackend::InferenceMmaTf32Rna
| PhysicalGemmBackend::InferenceSm120TmaFma
| PhysicalGemmBackend::InferenceSm120TmaMma16
| PhysicalGemmBackend::InferenceSm120TmaMmaTf32Rna
| PhysicalGemmBackend::FixedMatvecEightWarp => ModuleKind::Fixed,
PhysicalGemmBackend::Sm80Mma16
| PhysicalGemmBackend::MmaTf32Rna
| PhysicalGemmBackend::MmaTf32RnaSplitK2
| PhysicalGemmBackend::MmaTf32RnaSplitK4
| PhysicalGemmBackend::MmaTf32RnaSplitK8 => ModuleKind::TriadSm80,
PhysicalGemmBackend::Sm89MmaTf32Compact8 => ModuleKind::TriadSm89Finalist,
PhysicalGemmBackend::Sm89MmaTf32PreRna
| PhysicalGemmBackend::Sm89MmaTf32AddHalf
| PhysicalGemmBackend::Sm89MmaTf32NtALdmatrix
| PhysicalGemmBackend::Sm89MmaTf32NtRna
| PhysicalGemmBackend::Sm89MmaTf32TnDirectRna => ModuleKind::TriadSm89Tf32Joint,
PhysicalGemmBackend::Sm89Mma16HalfS3
| PhysicalGemmBackend::Sm89Mma16HalfS2
| PhysicalGemmBackend::Sm89Mma16HalfS4 => ModuleKind::TriadSm89Half,
PhysicalGemmBackend::Sm90aWgmma | PhysicalGemmBackend::Sm90aWgmmaTf32Tma => {
ModuleKind::TriadSm90a
}
PhysicalGemmBackend::Sm100Tcgen05 | PhysicalGemmBackend::Sm100Tcgen05Tf32Tma => {
ModuleKind::TriadSm100
}
PhysicalGemmBackend::Sm120TmaMma16
| PhysicalGemmBackend::Sm120TmaMmaTf32Rna
| PhysicalGemmBackend::Sm120TmaMmaTf32RnaStreamKV1
| PhysicalGemmBackend::Sm120TmaFmaExact => ModuleKind::TriadSm120,
}
}
fn expected_route_tuning_revision(backend: PhysicalGemmBackend, generic: u16) -> u16 {
match backend {
PhysicalGemmBackend::Sm89MmaTf32Compact8 => {
super::gemm_bi_triad::SM89_FINALIST_TUNING_REVISION
}
PhysicalGemmBackend::Sm89MmaTf32PreRna
| PhysicalGemmBackend::Sm89MmaTf32AddHalf
| PhysicalGemmBackend::Sm89MmaTf32NtALdmatrix
| PhysicalGemmBackend::Sm89MmaTf32NtRna
| PhysicalGemmBackend::Sm89MmaTf32TnDirectRna => {
super::gemm_bi_triad::SM89_TF32_JOINT_TUNING_REVISION
}
PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan => {
super::kernel_identity::SM89_FIXED_COPYPLAN_ROUTE_REVISION
}
PhysicalGemmBackend::Sm89Mma16HalfS3
| PhysicalGemmBackend::Sm89Mma16HalfS2
| PhysicalGemmBackend::Sm89Mma16HalfS4 => super::kernel_identity::SM89_HALF_ROUTE_REVISION,
PhysicalGemmBackend::ScalarFmaSm89ExactF32DualChunkFused
| PhysicalGemmBackend::ScalarFmaSm89ExactF32DirectSplitMPartial => {
super::kernel_identity::SM89_EXACT_F32_TN_ROUTE_REVISION
}
PhysicalGemmBackend::ScalarFmaTnDirectF64FoldSm89 => {
super::kernel_identity::SM89_EXACT_F32_D128_ROUTE_REVISION
}
_ => generic,
}
}
const fn scalar_backend_supports_logical_f32(backend: PhysicalGemmBackend) -> bool {
matches!(
backend,
PhysicalGemmBackend::ScalarFma
| PhysicalGemmBackend::ScalarFmaSplitKPartial
| PhysicalGemmBackend::ScalarFmaSplitKF32Reduce
| PhysicalGemmBackend::ScalarFmaTnNarrowSplitMPartial
| PhysicalGemmBackend::ScalarFmaTnSplitMF64Reduce
| PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan
| PhysicalGemmBackend::ScalarFmaSm89ExactF32DualChunkFused
| PhysicalGemmBackend::ScalarFmaSm89ExactF32DirectSplitMPartial
| PhysicalGemmBackend::ScalarFmaTnDirectF64FoldSm89
| PhysicalGemmBackend::Sm120TmaFmaExact
)
}
#[doc(hidden)]
pub struct RouteControls<'a> {
ctx: &'a GpuCtx,
}
#[doc(hidden)]
impl RouteControls<'_> {
pub fn family(&self) -> BiGemmFamily {
self.ctx.bi_gemm_family()
}
pub fn set_family(&self, family: BiGemmFamily) {
self.ctx.set_bi_gemm_family(family);
}
pub fn tensor_cores(&self) -> bool {
self.ctx.bi_tensor_cores()
}
pub fn set_tensor_cores(&self, on: bool) {
self.ctx.set_bi_tensor_cores(on);
}
pub fn f32_policy(&self) -> F32TriadPolicy {
self.ctx.f32_triad_policy()
}
pub fn set_f32_policy(&self, policy: F32TriadPolicy) {
self.ctx.set_f32_triad_policy(policy);
}
pub fn half_policy(&self) -> HalfTriadPolicy {
self.ctx.half_triad_policy()
}
pub fn set_half_policy(&self, policy: HalfTriadPolicy) {
self.ctx.set_half_triad_policy(policy);
}
pub fn policy(&self) -> GemmPolicy {
self.ctx.gemm_policy()
}
pub fn tf32(&self) -> bool {
self.ctx.tf32()
}
pub fn batch_invariant(&self) -> bool {
self.ctx.batch_invariant()
}
pub fn fast_gemm(&self) -> bool {
self.ctx.fast_gemm()
}
}
#[cfg(test)]
mod tests {
#[test]
#[ignore = "needs a CUDA device"]
fn inference_recorder_query_rejects_conflicting_borrow() {
let device = crate::mamba_ssm::gpu::device::GpuDevice::new(0).unwrap();
let ctx = super::GpuCtx::new(&device).unwrap();
assert!(!ctx.gemm_route_recording_active().unwrap());
let borrow = ctx.gemm_route_recorder.borrow_mut();
assert!(ctx.gemm_route_recording_active().is_err());
drop(borrow);
ctx.record_eager_gemm_trace(|| {
assert!(ctx.gemm_route_recording_active()?);
Ok(())
})
.unwrap();
assert!(!ctx.gemm_route_recording_active().unwrap());
}
#[test]
#[ignore = "needs a CUDA device"]
fn proof_arms_run_with_the_recorder_set_aside() {
let device = crate::mamba_ssm::gpu::device::GpuDevice::new(0).unwrap();
let ctx = super::GpuCtx::new(&device).unwrap();
ctx.record_eager_gemm_trace(|| {
assert!(ctx.gemm_route_recording_active()?);
let inside =
ctx.with_gemm_route_recording_suspended(|| ctx.gemm_route_recording_active())?;
assert!(!inside?);
assert!(ctx.gemm_route_recording_active()?);
Ok(())
})
.unwrap();
assert!(!ctx.gemm_route_recording_active().unwrap());
}
use super::{
BiGemmFamily, F32TriadPolicy, GemmEnvValues, GemmMode, GemmRole, HalfTriadPolicy,
ResolvedGemmEnv, WeightDtype, expected_route_module, expected_route_schedule_revision,
expected_route_tuning_revision, explicit_gemm_config, m1_mixed_graph_max_dim,
resolve_gemm_env, scalar_backend_supports_logical_f32, validate_multiprocessor_identity,
};
use crate::config::ScanMode;
use crate::mamba_ssm::gpu::forward::GpuMambaDims;
use crate::mamba_ssm::gpu::kernel_identity::{
ModuleKind, PhysicalGemmBackend, SCHEDULE_REVISION,
};
use std::ffi::OsString;
#[test]
fn mixed_graph_scratch_covers_the_two_inner_projection_output() {
let dims = GpuMambaDims {
batch: 1,
d_model: 64,
d_inner: 129,
d_state: 8,
d_conv: 4,
dt_rank: 4,
xdbl_dim: 20,
seq_len: 3,
mamba_input_dim: 17,
n_layers: 1,
scan_mode: ScanMode::Sequential,
rms_norm_eps: 1e-5,
};
assert_eq!(m1_mixed_graph_max_dim(&dims), 258);
}
#[test]
fn device_and_loaded_kernel_multiprocessor_counts_must_match() {
assert!(validate_multiprocessor_identity(142, 142).is_ok());
for counts in [(0, 142), (142, 0), (108, 142)] {
let error = validate_multiprocessor_identity(counts.0, counts.1)
.expect_err("incoherent CUDA topology must be rejected");
assert!(error.contains("multiprocessor"), "{error}");
}
}
#[test]
fn sm120_mma16_graph_routes_use_their_sealed_schedule_revision() {
assert_eq!(
expected_route_schedule_revision(PhysicalGemmBackend::Sm120TmaMma16, SCHEDULE_REVISION,),
super::super::gemm_bi_triad::SM120_SCHEDULE_REVISION
);
assert_eq!(
expected_route_schedule_revision(
PhysicalGemmBackend::Sm120TmaMmaTf32Rna,
SCHEDULE_REVISION,
),
SCHEDULE_REVISION
);
}
#[test]
fn sm89_finalist_routes_use_their_private_tuning_revision() {
let generic = super::super::gemm_bi_triad::F32_TF32_TUNING_REVISION;
let finalist =
expected_route_tuning_revision(PhysicalGemmBackend::Sm89MmaTf32Compact8, generic);
assert_eq!(
finalist,
super::super::gemm_bi_triad::SM89_FINALIST_TUNING_REVISION
);
assert_ne!(finalist, 0);
assert_ne!(finalist, 1);
let portable = expected_route_tuning_revision(PhysicalGemmBackend::MmaTf32Rna, generic);
assert_eq!(portable, 46);
assert_ne!(portable, finalist);
}
#[test]
fn sm89_tf32_joint_backends_use_the_joint_module_and_private_revision() {
let generic = super::super::gemm_bi_triad::F32_TF32_TUNING_REVISION;
for backend in [
PhysicalGemmBackend::Sm89MmaTf32PreRna,
PhysicalGemmBackend::Sm89MmaTf32AddHalf,
PhysicalGemmBackend::Sm89MmaTf32NtALdmatrix,
PhysicalGemmBackend::Sm89MmaTf32NtRna,
PhysicalGemmBackend::Sm89MmaTf32TnDirectRna,
] {
assert_eq!(
expected_route_module(backend),
ModuleKind::TriadSm89Tf32Joint
);
assert_eq!(
expected_route_tuning_revision(backend, generic),
super::super::gemm_bi_triad::SM89_TF32_JOINT_TUNING_REVISION
);
}
assert_eq!(generic, 46);
}
#[test]
fn sm89_fixed_copyplan_routes_use_a_private_revision_without_moving_global_46() {
let generic = super::super::gemm_bi_triad::F32_TF32_TUNING_REVISION;
let copyplan = expected_route_tuning_revision(
PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan,
generic,
);
assert_eq!(
copyplan,
crate::mamba_ssm::gpu::kernel_identity::SM89_FIXED_COPYPLAN_ROUTE_REVISION
);
assert_ne!(copyplan, 0);
assert_ne!(copyplan, 2);
assert_eq!(
expected_route_tuning_revision(PhysicalGemmBackend::ScalarFma, generic),
46
);
assert_eq!(
expected_route_module(PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan),
ModuleKind::Fixed
);
assert_eq!(
expected_route_module(PhysicalGemmBackend::ScalarFma),
ModuleKind::TriadScalar
);
}
#[test]
fn sm89_half_routes_use_the_isolated_module_and_private_revision() {
let generic = super::super::gemm_bi_triad::F32_TF32_TUNING_REVISION;
for backend in [
PhysicalGemmBackend::Sm89Mma16HalfS3,
PhysicalGemmBackend::Sm89Mma16HalfS2,
PhysicalGemmBackend::Sm89Mma16HalfS4,
] {
assert_eq!(
expected_route_tuning_revision(backend, generic),
crate::mamba_ssm::gpu::kernel_identity::SM89_HALF_ROUTE_REVISION
);
assert_eq!(expected_route_module(backend), ModuleKind::TriadSm89Half);
}
assert_eq!(generic, 46);
}
#[test]
fn sm89_exact_f32_tn_routes_use_the_isolated_module_and_private_revision() {
for backend in [
PhysicalGemmBackend::ScalarFmaSm89ExactF32DualChunkFused,
PhysicalGemmBackend::ScalarFmaSm89ExactF32DirectSplitMPartial,
] {
assert_eq!(
expected_route_module(backend),
ModuleKind::TriadSm89ExactF32
);
assert_eq!(
expected_route_tuning_revision(backend, 46),
crate::mamba_ssm::gpu::kernel_identity::SM89_EXACT_F32_TN_ROUTE_REVISION
);
assert!(scalar_backend_supports_logical_f32(backend));
}
}
#[test]
fn sm89_exact_f32_d128_routes_use_their_isolated_module_and_private_revision() {
let backend = PhysicalGemmBackend::ScalarFmaTnDirectF64FoldSm89;
assert_eq!(
expected_route_module(backend),
ModuleKind::TriadSm89ExactF32D128
);
assert_eq!(
expected_route_tuning_revision(backend, 46),
crate::mamba_ssm::gpu::kernel_identity::SM89_EXACT_F32_D128_ROUTE_REVISION
);
assert!(scalar_backend_supports_logical_f32(backend));
}
#[test]
fn logical_f32_accepts_only_scalar_triad_backends() {
for backend in [
PhysicalGemmBackend::ScalarFma,
PhysicalGemmBackend::ScalarFmaSplitKPartial,
PhysicalGemmBackend::ScalarFmaSplitKF32Reduce,
PhysicalGemmBackend::ScalarFmaTnNarrowSplitMPartial,
PhysicalGemmBackend::ScalarFmaTnSplitMF64Reduce,
PhysicalGemmBackend::ScalarFmaSm89FixedCopyPlan,
PhysicalGemmBackend::ScalarFmaSm89ExactF32DualChunkFused,
PhysicalGemmBackend::ScalarFmaSm89ExactF32DirectSplitMPartial,
PhysicalGemmBackend::ScalarFmaTnDirectF64FoldSm89,
PhysicalGemmBackend::Sm120TmaFmaExact,
] {
assert!(scalar_backend_supports_logical_f32(backend), "{backend:?}");
}
for backend in [
PhysicalGemmBackend::Sm80Mma16,
PhysicalGemmBackend::MmaTf32Rna,
PhysicalGemmBackend::Sm120TmaMma16,
] {
assert!(!scalar_backend_supports_logical_f32(backend), "{backend:?}");
}
}
#[cfg(unix)]
#[test]
fn explicit_gemm_config_carries_mode_role_and_numeric_defaults() {
for mode in [
GemmMode::Deterministic,
GemmMode::CublasFast,
GemmMode::CublasPedantic,
] {
for (dtype, f32_policy) in [
(WeightDtype::F32, F32TriadPolicy::ExactScalarFma),
(WeightDtype::Tf32, F32TriadPolicy::AllowDeterministicTf32),
(WeightDtype::Bf16, F32TriadPolicy::ExactScalarFma),
(WeightDtype::F16, F32TriadPolicy::ExactScalarFma),
] {
for (role, family) in [
(GemmRole::inference(dtype), BiGemmFamily::Inference),
(GemmRole::triad(dtype), BiGemmFamily::Triad),
] {
assert_eq!(
explicit_gemm_config(mode, role),
ResolvedGemmEnv {
mode,
family,
tensor_cores: true,
f32_policy,
half_policy: HalfTriadPolicy::AllowStreamKFixedOrder,
}
);
}
}
}
}
fn absent_env() -> Result<String, std::env::VarError> {
Err(std::env::VarError::NotPresent)
}
fn empty_gemm_env() -> GemmEnvValues {
GemmEnvValues { mode: absent_env() }
}
#[test]
fn gemm_mode_environment_reads_one_variable() {
for (value, expected) in [
("deterministic", GemmMode::Deterministic),
("cublas-fast", GemmMode::CublasFast),
("cublas-pedantic", GemmMode::CublasPedantic),
] {
let values = GemmEnvValues {
mode: Ok(value.into()),
};
let resolved = resolve_gemm_env(values, GemmRole::triad(WeightDtype::Tf32)).unwrap();
assert_eq!(resolved.mode, expected);
assert_eq!(resolved.family, BiGemmFamily::Triad);
assert_eq!(resolved.f32_policy, F32TriadPolicy::AllowDeterministicTf32);
assert!(resolved.tensor_cores);
assert_eq!(
resolved.half_policy,
HalfTriadPolicy::AllowStreamKFixedOrder
);
}
let resolved =
resolve_gemm_env(empty_gemm_env(), GemmRole::inference(WeightDtype::F32)).unwrap();
assert_eq!(resolved.mode, GemmMode::Deterministic);
assert_eq!(resolved.family, BiGemmFamily::Inference);
assert_eq!(resolved.f32_policy, F32TriadPolicy::ExactScalarFma);
let error = resolve_gemm_env(
GemmEnvValues {
mode: Ok("fast".into()),
},
GemmRole::triad(WeightDtype::F32),
)
.unwrap_err();
assert!(error.contains("MAMBA_RS_GEMM_MODE"), "{error}");
let error = resolve_gemm_env(
GemmEnvValues {
mode: Err(std::env::VarError::NotUnicode(OsString::from("x"))),
},
GemmRole::triad(WeightDtype::F32),
)
.unwrap_err();
assert!(error.contains("MAMBA_RS_GEMM_MODE"), "{error}");
}
}