use std::any::{Any, TypeId};
use std::sync::Arc;
use burn::backend::wgpu::{CubeBackend, CubeTensor, WgpuDevice, WgpuRuntime};
use burn::tensor::backend::Backend;
use burn::tensor::{DType, Device, FloatDType, Shape, Tensor, TensorPrimitive};
use burn_cubecl::fusion::FusionCubeRuntime;
use burn_cubecl::kernel::into_contiguous;
use burn_cubecl_fusion::CubeFusionHandle;
use burn_fusion::Fusion;
use burn_fusion::stream::{Operation, OperationStreams};
use burn_ir::{CustomOpIr, HandleContainer, OperationIr, TensorIr, TensorStatus};
use combs_formats::ModelSource;
use crate::llama::linear as dense_linear;
use crate::qmatmul::QuantWeight;
use crate::{ModelError, Result};
type FusedF32 = Fusion<CubeBackend<WgpuRuntime, f32, i32, u32>>;
type UnfusedF32 = CubeBackend<WgpuRuntime, f32, i32, u32>;
type UnfusedF16 = CubeBackend<WgpuRuntime, burn::tensor::f16, i32, u32>;
type InnerF32 = CubeBackend<WgpuRuntime, f32, i32, u32>;
pub trait QuantLinearOp<B: Backend>: Send + Sync {
fn forward(&self, x: Tensor<B, 3>) -> Tensor<B, 3>;
fn dims(&self) -> [usize; 2];
fn vram_bytes(&self) -> usize;
}
pub enum Linear<B: Backend> {
Dense(Tensor<B, 2>),
Quant(Box<dyn QuantLinearOp<B>>),
}
impl<B: Backend> Linear<B> {
pub fn dims(&self) -> [usize; 2] {
match self {
Linear::Dense(w) => w.dims(),
Linear::Quant(op) => op.dims(),
}
}
pub fn forward(&self, x: Tensor<B, 3>, bias: Option<&Tensor<B, 1>>) -> Tensor<B, 3> {
match self {
Linear::Dense(w) => dense_linear(x, w, bias),
Linear::Quant(op) => {
let out = op.forward(x);
match bias {
Some(b) => {
let [batch, seq, dim] = out.dims();
out + b.clone().reshape([1, 1, dim]).expand([batch, seq, dim])
}
None => out,
}
}
}
}
}
struct CubeQuantLinear {
w: Arc<QuantWeight>,
}
impl CubeQuantLinear {
fn dims(&self) -> [usize; 2] {
[self.w.n_out(), self.w.k()]
}
fn forward_cube(&self, x: CubeTensor<WgpuRuntime>, batch: usize, seq: usize) -> CubeTensor<WgpuRuntime> {
let x = into_contiguous(x);
let out_h = self.w.matmul_device(&x.client, x.handle.clone(), batch * seq);
CubeTensor::new_contiguous(
x.client.clone(),
x.device.clone(),
Shape::from([batch, seq, self.w.n_out()]),
out_h,
DType::F32,
)
}
}
fn to_f32<B: Backend>(x: Tensor<B, 3>) -> Tensor<B, 3> {
match x.dtype() {
DType::F32 => x,
_ => x.cast(FloatDType::F32),
}
}
fn to_dtype<B: Backend>(out: Tensor<B, 3>, dtype: DType) -> Tensor<B, 3> {
match dtype {
DType::F16 => out.cast(FloatDType::F16),
DType::BF16 => out.cast(FloatDType::BF16),
_ => out,
}
}
impl QuantLinearOp<UnfusedF32> for CubeQuantLinear {
fn forward(&self, x: Tensor<UnfusedF32, 3>) -> Tensor<UnfusedF32, 3> {
let in_dtype = x.dtype();
let [batch, seq, _] = x.dims();
let prim = to_f32(x).into_primitive().tensor();
let out = self.forward_cube(prim, batch, seq);
to_dtype(
Tensor::from_primitive(TensorPrimitive::Float(out)),
in_dtype,
)
}
fn dims(&self) -> [usize; 2] {
CubeQuantLinear::dims(self)
}
fn vram_bytes(&self) -> usize {
self.w.vram_bytes()
}
}
impl QuantLinearOp<UnfusedF16> for CubeQuantLinear {
fn forward(&self, x: Tensor<UnfusedF16, 3>) -> Tensor<UnfusedF16, 3> {
let in_dtype = x.dtype();
let [batch, seq, _] = x.dims();
let prim = to_f32(x).into_primitive().tensor();
let out = self.forward_cube(prim, batch, seq);
to_dtype(
Tensor::<UnfusedF16, 3>::from_primitive(TensorPrimitive::Float(out)),
in_dtype,
)
}
fn dims(&self) -> [usize; 2] {
CubeQuantLinear::dims(self)
}
fn vram_bytes(&self) -> usize {
self.w.vram_bytes()
}
}
struct QuantMatmulOp {
desc: CustomOpIr,
w: Arc<QuantWeight>,
batch: usize,
seq: usize,
}
impl core::fmt::Debug for QuantMatmulOp {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
write!(
f,
"QuantMatmulOp {{ w: [{}, {}], m: {} }}",
self.w.n_out(),
self.w.k(),
self.batch * self.seq
)
}
}
impl Operation<FusionCubeRuntime<WgpuRuntime>> for QuantMatmulOp {
fn execute(&self, handles: &mut HandleContainer<CubeFusionHandle<WgpuRuntime>>) {
let ([input], [output]) = self.desc.as_fixed::<1, 1>();
let x: CubeTensor<WgpuRuntime> = handles.get_float_tensor::<InnerF32>(input);
let x = into_contiguous(x);
let out_h = self.w.matmul_device(&x.client, x.handle.clone(), self.batch * self.seq);
let out = CubeTensor::new_contiguous(
x.client.clone(),
x.device.clone(),
Shape::from([self.batch, self.seq, self.w.n_out()]),
out_h,
DType::F32,
);
handles.register_float_tensor::<InnerF32>(&output.id, out);
}
}
impl QuantLinearOp<FusedF32> for CubeQuantLinear {
fn forward(&self, x: Tensor<FusedF32, 3>) -> Tensor<FusedF32, 3> {
let in_dtype = x.dtype();
let [batch, seq, _] = x.dims();
let prim = to_f32(x).into_primitive().tensor();
let client = prim.client.clone();
let mut streams = OperationStreams::default();
streams.tensor(&prim);
let input_ir = prim.into_ir();
let out_ir = TensorIr {
id: client.create_empty_handle(),
shape: Shape::from([batch, seq, self.w.n_out()]),
status: TensorStatus::NotInit,
dtype: DType::F32,
};
let desc = CustomOpIr::new("combs_quant_matmul", &[input_ir], &[out_ir]);
let op = QuantMatmulOp {
desc: desc.clone(),
w: self.w.clone(),
batch,
seq,
};
let mut outputs = client.register(streams, OperationIr::Custom(desc), op);
let out = outputs.pop().expect("custom op declares one output");
to_dtype(
Tensor::from_primitive(TensorPrimitive::Float(out)),
in_dtype,
)
}
fn dims(&self) -> [usize; 2] {
CubeQuantLinear::dims(self)
}
fn vram_bytes(&self) -> usize {
self.w.vram_bytes()
}
}
fn cast_op<B: Backend, T: Backend>(op: Box<dyn QuantLinearOp<T>>) -> Option<Box<dyn QuantLinearOp<B>>> {
let any: Box<dyn Any> = Box::new(op);
any.downcast::<Box<dyn QuantLinearOp<B>>>().ok().map(|b| *b)
}
fn debug_quant(name: &str, outcome: &str) {
if std::env::var_os("COMBS_DEBUG_QUANT").is_some() {
eprintln!("quant-linear {name}: {outcome}");
}
}
pub fn try_quant_linear<B: Backend>(
source: &dyn ModelSource,
name: &str,
device: &Device<B>,
) -> Result<Option<Box<dyn QuantLinearOp<B>>>> {
if std::env::var_os("COMBS_NO_QUANT_KERNELS").is_some_and(|v| v != "0") {
return Ok(None);
}
let supported = [
TypeId::of::<FusedF32>(),
TypeId::of::<UnfusedF32>(),
TypeId::of::<UnfusedF16>(),
];
if !supported.contains(&TypeId::of::<B>()) {
debug_quant(name, "backend not wgpu f32/f16 — dense fallback");
return Ok(None);
}
let device_any: &dyn Any = device;
let Some(wgpu_device) = device_any.downcast_ref::<WgpuDevice>() else {
debug_quant(name, "device not WgpuDevice — dense fallback");
return Ok(None);
};
let Some(qt) = source.open_tensor_quant(name).map_err(ModelError::Format)? else {
debug_quant(name, "no packed quant tensor — dense fallback");
return Ok(None);
};
let &[n_out, k] = qt.shape.as_slice() else {
debug_quant(name, "not rank-2 — dense fallback");
return Ok(None);
};
let client = <WgpuRuntime as cubecl::prelude::Runtime>::client(wgpu_device);
let Ok(w) = QuantWeight::from_quant_tensor(&client, qt.format, &qt.data, n_out, k) else {
debug_quant(name, "kernel-incompatible shape — dense fallback");
return Ok(None);
};
debug_quant(name, "packed on device");
let lin = CubeQuantLinear { w: Arc::new(w) };
if TypeId::of::<B>() == TypeId::of::<FusedF32>() {
return Ok(cast_op::<B, FusedF32>(Box::new(lin)));
}
if TypeId::of::<B>() == TypeId::of::<UnfusedF32>() {
return Ok(cast_op::<B, UnfusedF32>(Box::new(lin)));
}
if TypeId::of::<B>() == TypeId::of::<UnfusedF16>() {
return Ok(cast_op::<B, UnfusedF16>(Box::new(lin)));
}
Ok(None)
}
#[cfg(test)]
mod tests {
use super::*;
use burn::tensor::TensorData;
use combs_formats::QuantFormat;
use cubecl::prelude::Runtime;
fn synth_q4_0(n_blocks: usize) -> Vec<u8> {
let mut out = Vec::with_capacity(n_blocks * 18);
let mut s = 0x12345678u32;
for b in 0..n_blocks {
let scale = burn::tensor::f16::from_f32(0.003 * ((b % 11) as f32 + 1.0));
out.extend_from_slice(&scale.to_le_bytes());
for _ in 0..16 {
s = s.wrapping_mul(1664525).wrapping_add(1013904223);
out.push((s >> 24) as u8);
}
}
out
}
fn pin_device_dtypes() {
use std::sync::Once;
static PIN: Once = Once::new();
PIN.call_once(|| {
let device = WgpuDevice::default();
let _ = burn::tensor::set_default_dtypes::<UnfusedF32>(
&device,
FloatDType::F32,
burn::tensor::IntDType::I32,
);
});
}
fn quant_and_dense<B: Backend>(
device: &Device<B>,
n_out: usize,
k: usize,
dtype: FloatDType,
) -> (Linear<B>, Linear<B>)
where
CubeQuantLinear: QuantLinearOp<B>,
{
let data = synth_q4_0(n_out * k / 32);
let client = <WgpuRuntime as Runtime>::client(&Default::default());
let w = Arc::new(
QuantWeight::from_quant_tensor(&client, QuantFormat::Q4_0, &data, n_out, k).unwrap(),
);
let quant = Linear::Quant(Box::new(CubeQuantLinear { w }) as Box<dyn QuantLinearOp<B>>);
let wf = combs_formats::quants::dequantize_q4_0(&data, n_out * k).unwrap();
let dense = Linear::Dense(
Tensor::<B, 2>::from_data(TensorData::new(wf, [n_out, k]), device).cast(dtype),
);
(quant, dense)
}
fn assert_close(got: &[f32], expect: &[f32], rel: f32) {
assert_eq!(got.len(), expect.len());
for (i, (g, e)) in got.iter().zip(expect.iter()).enumerate() {
let tol = rel * e.abs().max(1.0);
assert!((g - e).abs() <= tol, "[{i}]: got {g}, expect {e}");
}
}
#[test]
fn fused_backend_matches_dense() {
if crate::skip_no_gpu() {
return;
}
pin_device_dtypes();
let device: Device<FusedF32> = Default::default();
let (n_out, k) = (48, 64);
let (quant, dense) = quant_and_dense::<FusedF32>(&device, n_out, k, FloatDType::F32);
assert_eq!(quant.dims(), [n_out, k]);
let x: Vec<f32> = (0..3 * k).map(|i| ((i % 32) as f32) / 16.0 - 1.0).collect();
let x = Tensor::<FusedF32, 3>::from_data(TensorData::new(x, [1, 3, k]), &device)
.cast(FloatDType::F32);
let b: Vec<f32> = (0..n_out).map(|i| (i as f32) / 100.0).collect();
let bias = Tensor::<FusedF32, 1>::from_data(TensorData::new(b, [n_out]), &device)
.cast(FloatDType::F32);
let got: Vec<f32> = quant
.forward(x.clone(), Some(&bias))
.into_data()
.to_vec()
.unwrap();
let expect: Vec<f32> = dense
.forward(x, Some(&bias))
.into_data()
.to_vec()
.unwrap();
assert_close(&got, &expect, 1e-4);
}
#[test]
fn f16_backend_matches_dense() {
if crate::skip_no_gpu() {
return;
}
pin_device_dtypes();
let device: Device<UnfusedF16> = Default::default();
let (n_out, k) = (48, 64);
let (quant, dense) = quant_and_dense::<UnfusedF16>(&device, n_out, k, FloatDType::F16);
let x: Vec<f32> = (0..3 * k).map(|i| ((i % 32) as f32) / 16.0 - 1.0).collect();
let x = Tensor::<UnfusedF16, 3>::from_data(TensorData::new(x, [1, 3, k]), &device)
.cast(FloatDType::F16);
let got: Vec<f32> = quant
.forward(x.clone(), None)
.into_data()
.convert::<f32>()
.to_vec()
.unwrap();
let expect: Vec<f32> = dense
.forward(x, None)
.into_data()
.convert::<f32>()
.to_vec()
.unwrap();
assert_close(&got, &expect, 1e-2);
}
}