#![allow(clippy::upper_case_acronyms)]
#![allow(clippy::needless_pass_by_ref_mut)]
use super::{DTypeCapability, Dev, DeviceInfo, DeviceProgramId, LaunchArg, Pool, ProgramId};
use crate::{
DType, Set,
error::{BackendError, ErrorStatus},
graph::{Graph, Node},
kernel::{Kernel, OpId},
shape::Dim,
slab::{Slab, SlabId},
};
use libloading::Library;
use nanoserde::DeJson;
use std::{
collections::BTreeSet,
sync::{Arc, Mutex, OnceLock},
};
static CBLAS_DEVICE: OnceLock<Mutex<CblasDevice>> = OnceLock::new();
type SgemmFn = unsafe extern "C" fn(
order: i32,
transa: i32,
transb: i32,
m: i32,
n: i32,
k: i32,
alpha: f32,
a: *const f32,
lda: i32,
b: *const f32,
ldb: i32,
beta: f32,
c: *mut f32,
ldc: i32,
);
const CBLAS_ROW_MAJOR: i32 = 101;
const CBLAS_NO_TRANS: i32 = 111;
const OPENBLAS_PATH: &str = "/usr/lib/x86_64-linux-gnu/libopenblas.so";
#[derive(Debug, DeJson)]
#[nserde(default)]
pub struct CblasConfig {
pub enabled: bool,
}
impl Default for CblasConfig {
fn default() -> Self {
Self { enabled: true }
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct CblasKernelId(u32);
impl From<usize> for CblasKernelId {
fn from(value: usize) -> Self {
CblasKernelId(u32::try_from(value).unwrap())
}
}
impl From<CblasKernelId> for usize {
fn from(value: CblasKernelId) -> Self {
value.0 as usize
}
}
impl SlabId for CblasKernelId {
const ZERO: Self = Self(0);
const NULL: Self = Self(u32::MAX);
fn inc(&mut self) {
self.0 += 1;
}
}
#[derive(Debug)]
pub struct CblasKernel {
sgemm: SgemmFn,
}
#[derive(Debug)]
pub struct CblasProgram {
kernel: CblasKernelId,
m: Dim,
n: Dim,
k: Dim,
}
#[derive(Debug)]
pub struct CblasDevice {
device_info: Arc<DeviceInfo>,
#[allow(dead_code)]
lib: Library,
kernels: Slab<CblasKernelId, CblasKernel>,
programs: Slab<DeviceProgramId, CblasProgram>,
}
fn device_with(config: &CblasConfig, debug_dev: bool) -> Result<&'static Mutex<CblasDevice>, BackendError> {
if let Some(dev) = CBLAS_DEVICE.get() {
return Ok(dev);
}
if !config.enabled {
if debug_dev {
println!("[cblas] configured out");
}
return Err(BackendError { status: ErrorStatus::Initialization, context: "CBLAS backend configured out".into() });
}
let lib = unsafe { Library::new(OPENBLAS_PATH) }?;
let sgemm: SgemmFn = *unsafe { lib.get(b"cblas_sgemm") }?;
let mut kernels = Slab::new();
kernels.push(CblasKernel { sgemm });
let dev = Mutex::new(CblasDevice {
device_info: Arc::new(DeviceInfo {
compute: 1,
max_global_work_dims: vec![Dim::from(0i64); 3],
max_local_threads: 1,
max_local_work_dims: vec![1, 1, 1],
preferred_vector_size: 8,
local_mem_size: 0,
max_register_bytes: 0,
tensor_cores: false,
warp_size: 1,
cc: [0, 0],
dtype_capability: [DTypeCapability::none(); DType::N_DTYPES],
has_native_exp2: false,
supported_vec_lens: vec![],
tenstorrent: false,
tile: [1, 1],
tile_sizes: vec![],
wmma_layouts: vec![],
num_circular_buffers: 0,
has_openmp: false,
}),
lib,
kernels,
programs: Slab::new(),
});
if debug_dev {
println!("[cblas] initialized from {OPENBLAS_PATH}");
}
let _ = CBLAS_DEVICE.set(dev);
Ok(CBLAS_DEVICE.get().unwrap())
}
pub(super) fn device() -> Result<&'static Mutex<CblasDevice>, BackendError> {
device_with(&super::config().cblas, super::debug_backends())
}
impl CblasDevice {
pub fn info(&self) -> Arc<DeviceInfo> {
self.device_info.clone()
}
pub fn free_compute(&self) -> u128 {
self.device_info.compute
}
pub fn release(&mut self, program_id: DeviceProgramId) {
self.programs.remove(program_id);
}
pub fn compile(&mut self, _kernel: &Kernel, _debug_asm: bool) -> Result<DeviceProgramId, BackendError> {
Err(BackendError {
status: ErrorStatus::KernelCompilation,
context: "cblas device only runs AOT matmul kernels, it does not compile generic kernels.".into(),
})
}
pub fn match_graph(&mut self, graph: &mut Graph, outputs: &BTreeSet<OpId>) {
let order = graph.topo_sort_classes::<true>(&Set::default(), outputs, None);
for &cid in &order {
let Some(mm) = graph.match_matmul(cid) else {
continue;
};
if mm.in_dtype != DType::F32 || mm.acc_dtype != DType::F32 {
continue;
}
println!("[cblas] matched matmul m={}, n={}, k={}", mm.m, mm.n, mm.k);
let program_id = self.programs.push(CblasProgram { kernel: CblasKernelId::ZERO, m: mm.m, n: mm.n, k: mm.k });
graph.mint_node(
Node::Kernel {
inputs: Box::new([mm.a, mm.b]),
outputs: Box::new([mm.out]),
program_id: ProgramId { dev: Dev::Cblas, program_id },
time: 1,
},
mm.out,
);
}
}
#[allow(clippy::needless_pass_by_value)]
pub fn launch(&mut self, program_id: DeviceProgramId, pool_handle: Pool, args: &[LaunchArg]) -> Result<(), BackendError> {
debug_assert_eq!(pool_handle, Pool::Host);
let host = super::host::pool();
let mut memory_pool = super::lock(pool_handle, host);
let program = &self.programs[program_id];
let kernel = &self.kernels[program.kernel];
let m: i32 = i32::try_from(program.m)
.map_err(|_| BackendError { status: ErrorStatus::IncorrectKernelArg, context: "m exceeds i32 range".into() })?;
let n: i32 = i32::try_from(program.n)
.map_err(|_| BackendError { status: ErrorStatus::IncorrectKernelArg, context: "n exceeds i32 range".into() })?;
let k: i32 = i32::try_from(program.k)
.map_err(|_| BackendError { status: ErrorStatus::IncorrectKernelArg, context: "k exceeds i32 range".into() })?;
let LaunchArg::Buffer(b0) = args[0] else {
unreachable!("cblas sgemm args are plain buffers")
};
let LaunchArg::Buffer(b1) = args[1] else {
unreachable!("cblas sgemm args are plain buffers")
};
let LaunchArg::Buffer(b2) = args[2] else {
unreachable!("cblas sgemm args are plain buffers")
};
let a = memory_pool.buffer_ptr_mut(b0) as *mut f32;
let b = memory_pool.buffer_ptr_mut(b1) as *mut f32;
let c = memory_pool.buffer_ptr_mut(b2) as *mut f32;
unsafe {
(kernel.sgemm)(CBLAS_ROW_MAJOR, CBLAS_NO_TRANS, CBLAS_NO_TRANS, m, n, k, 1.0, a, k, b, n, 0.0, c, n);
}
Ok(())
}
}