mod context;
mod gemm;
mod pool;
use std::sync::OnceLock;
use crate::GemmTask;
use crate::backend::BackendUnavailable;
use crate::backend::operand::{Operand, classify};
use self::context::{Context, SetupError};
const FLOP_THRESHOLD: usize = 1 << 24;
static CONTEXT: OnceLock<Result<Context, SetupError>> = OnceLock::new();
static POISON: OnceLock<String> = OnceLock::new();
fn initialized() -> &'static Result<Context, SetupError> {
CONTEXT.get_or_init(Context::new)
}
fn context() -> Result<&'static Context, BackendUnavailable> {
if let Some(reason) = POISON.get() {
return Err(BackendUnavailable::Poisoned(reason.clone()));
}
match initialized() {
Ok(context) => Ok(context),
Err(error) => Err(BackendUnavailable::Initialization(error.to_string())),
}
}
pub(super) fn status() -> Result<(), BackendUnavailable> {
context().map(|_| ())
}
fn eligible<Element>(task: &GemmTask<'_, Element>) -> Option<(Operand, Operand, i32, i32, i32)> {
if task.m() == 1 || task.n() == 1 {
return None;
}
let flops = 2usize
.saturating_mul(task.m())
.saturating_mul(task.n())
.saturating_mul(task.k());
if flops < FLOP_THRESHOLD {
return None;
}
task.m().checked_mul(task.n())?;
let a = classify(task.a_strides(), task.m(), task.k())?;
let b = classify(task.b_strides(), task.k(), task.n())?;
let m = i32::try_from(task.m()).ok()?;
let n = i32::try_from(task.n()).ok()?;
let k = i32::try_from(task.k()).ok()?;
Some((a, b, m, n, k))
}
pub(super) fn gemm_f32(task: &GemmTask<'_, f32>) -> Option<Vec<f32>> {
let (a, b, m, n, k) = eligible(task)?;
let context = context().ok()?;
match gemm::executed_f32(context, task, &a, &b, m, n, k) {
Ok(product) => Some(product),
Err(reason) => {
let _ = POISON.set(reason);
None
}
}
}
pub(super) fn gemm_f64(task: &GemmTask<'_, f64>) -> Option<Vec<f64>> {
let (a, b, m, n, k) = eligible(task)?;
let context = context().ok()?;
match gemm::executed_f64(context, task, &a, &b, m, n, k) {
Ok(product) => Some(product),
Err(reason) => {
let _ = POISON.set(reason);
None
}
}
}
#[cfg(test)]
#[path = "tests/cuda_tests.rs"]
mod tests;