use crate::{
backend::BackendStorage, CpuStorage, DType, Device, Result, Shape, Storage, Tensor, D,
};
use iq_quants::*;
use k_quants::*;
use std::borrow::Cow;
use std::sync::Arc;
#[cfg(target_feature = "avx2")]
pub mod avx;
pub mod dsv4_qat;
mod dummy_cuda;
mod dummy_metal;
pub mod expert_stream;
pub mod ggml_file;
pub mod gguf_file;
pub mod imatrix_file;
mod iq_grids;
pub mod iq_quants;
pub mod k_quants;
#[cfg(feature = "metal")]
pub mod metal;
#[cfg(not(target_arch = "wasm32"))]
pub mod tokenizer;
#[cfg(not(feature = "metal"))]
mod metal {
pub use super::dummy_metal::*;
}
#[cfg(feature = "cuda")]
pub mod cuda;
#[cfg(feature = "cuda")]
pub mod fast_mmq;
#[cfg(feature = "cuda")]
pub mod fast_mmvq;
#[cfg(not(feature = "cuda"))]
mod cuda {
pub use super::dummy_cuda::*;
}
#[cfg(target_feature = "neon")]
pub mod neon;
#[cfg(target_feature = "simd128")]
pub mod simd128;
pub mod utils;
pub mod quant_format;
use half::{bf16, f16};
pub use k_quants::GgmlType;
fn as_t_slice<T>(data: &[u8]) -> &[T] {
let size = std::mem::size_of::<T>();
assert_eq!(
data.len() % size,
0,
"Data length must be a multiple of T's size"
);
let ptr = data.as_ptr();
assert_eq!(
(ptr as usize) % std::mem::align_of::<T>(),
0,
"Data pointer must be aligned to T's alignment"
);
unsafe { std::slice::from_raw_parts(ptr as *const T, data.len() / size) }
}
#[derive(Default)]
struct ResidentBanks {
#[cfg(feature = "rocm")]
rocm: std::sync::OnceLock<std::sync::Arc<crate::RocmStorage>>,
#[cfg(feature = "vulkan")]
vulkan: std::sync::OnceLock<std::sync::Arc<crate::VulkanStorage>>,
#[cfg(feature = "vulkan")]
vulkan_split: std::sync::OnceLock<std::sync::Arc<crate::vulkan::MoeBankSplit>>,
#[cfg(feature = "wgpu")]
wgpu: std::sync::OnceLock<std::sync::Arc<crate::WgpuStorage>>,
}
pub struct QTensor {
storage: QStorage,
shape: Shape,
banks: ResidentBanks,
}
impl Device {
fn qzeros(&self, elem_count: usize, dtype: GgmlDType) -> Result<QStorage> {
match self {
Device::Cpu => {
let storage = dtype.cpu_zeros(elem_count);
Ok(QStorage::Cpu(storage))
}
Device::Metal(metal) => {
let storage = metal::QMetalStorage::zeros(metal, elem_count, dtype)?;
Ok(QStorage::Metal(storage))
}
Device::Cuda(cuda) => {
let storage = cuda::QCudaStorage::zeros(cuda, elem_count, dtype)?;
Ok(QStorage::Cuda(storage))
}
#[cfg(feature = "rocm")]
Device::Rocm(d) => {
let storage = dtype.cpu_zeros(elem_count);
Ok(QStorage::Rocm(storage, d.clone()))
}
#[cfg(feature = "vulkan")]
Device::Vulkan(d) => {
let storage = dtype.cpu_zeros(elem_count);
Ok(QStorage::Vulkan(storage, d.clone()))
}
#[cfg(feature = "wgpu")]
Device::Wgpu(d) => {
let storage = dtype.cpu_zeros(elem_count);
Ok(QStorage::Wgpu(storage, d.clone()))
}
}
}
}
pub enum QStorage {
Cpu(Box<dyn QuantizedType>),
Metal(metal::QMetalStorage),
Cuda(cuda::QCudaStorage),
#[cfg(feature = "rocm")]
Rocm(Box<dyn QuantizedType>, crate::RocmDevice),
#[cfg(feature = "vulkan")]
Vulkan(Box<dyn QuantizedType>, crate::VulkanDevice),
#[cfg(feature = "wgpu")]
Wgpu(Box<dyn QuantizedType>, crate::WgpuDevice),
Stream(Arc<expert_stream::ExpertStreamBank>),
}
impl QStorage {
pub fn from_data(data: Cow<'_, [u8]>, device: &Device, dtype: GgmlDType) -> Result<Self> {
match device {
Device::Cpu => Ok(Self::Cpu(dtype.from_data(data))),
Device::Metal(d) => match dtype {
GgmlDType::F32 => metal::load_quantized(d, as_t_slice::<f32>(&data)),
GgmlDType::F16 => metal::load_quantized(d, as_t_slice::<f16>(&data)),
GgmlDType::Q4_0 => metal::load_quantized(d, as_t_slice::<BlockQ4_0>(&data)),
GgmlDType::Q4_1 => metal::load_quantized(d, as_t_slice::<BlockQ4_1>(&data)),
GgmlDType::Q5_0 => metal::load_quantized(d, as_t_slice::<BlockQ5_0>(&data)),
GgmlDType::Q5_1 => metal::load_quantized(d, as_t_slice::<BlockQ5_1>(&data)),
GgmlDType::Q8_0 => metal::load_quantized(d, as_t_slice::<BlockQ8_0>(&data)),
GgmlDType::Q8_1 => metal::load_quantized(d, as_t_slice::<BlockQ8_1>(&data)),
GgmlDType::Q2K => metal::load_quantized(d, as_t_slice::<BlockQ2K>(&data)),
GgmlDType::Q3K => metal::load_quantized(d, as_t_slice::<BlockQ3K>(&data)),
GgmlDType::Q4K => metal::load_quantized(d, as_t_slice::<BlockQ4K>(&data)),
GgmlDType::Q5K => metal::load_quantized(d, as_t_slice::<BlockQ5K>(&data)),
GgmlDType::Q6K => metal::load_quantized(d, as_t_slice::<BlockQ6K>(&data)),
GgmlDType::Q8K => metal::load_quantized(d, as_t_slice::<BlockQ8K>(&data)),
GgmlDType::IQ4_NL => metal::load_quantized(d, as_t_slice::<BlockIQ4nl>(&data)),
GgmlDType::IQ4_XS => metal::load_quantized(d, as_t_slice::<BlockIQ4xs>(&data)),
GgmlDType::MXFP4 => metal::load_quantized(d, as_t_slice::<BlockMXFP4>(&data)),
GgmlDType::BF16 => metal::load_quantized(d, as_t_slice::<bf16>(&data)),
GgmlDType::I32 => metal::load_quantized(d, as_t_slice::<i32>(&data)),
GgmlDType::IQ2_XXS => metal::load_quantized(d, as_t_slice::<BlockIQ2xxs>(&data)),
GgmlDType::IQ2_XS => metal::load_quantized(d, as_t_slice::<BlockIQ2xs>(&data)),
GgmlDType::IQ2_S => metal::load_quantized(d, as_t_slice::<BlockIQ2s>(&data)),
GgmlDType::IQ3_XXS => metal::load_quantized(d, as_t_slice::<BlockIQ3xxs>(&data)),
GgmlDType::IQ3_S => metal::load_quantized(d, as_t_slice::<BlockIQ3s>(&data)),
GgmlDType::IQ1_S => metal::load_quantized(d, as_t_slice::<BlockIQ1s>(&data)),
GgmlDType::IQ1_M => metal::load_quantized(d, as_t_slice::<BlockIQ1m>(&data)),
other => crate::bail!("{other:?} is not supported on the Metal backend"),
},
Device::Cuda(d) => match dtype {
GgmlDType::F32 => cuda::load_quantized(d, as_t_slice::<f32>(&data)),
GgmlDType::F16 => cuda::load_quantized(d, as_t_slice::<f16>(&data)),
GgmlDType::Q4_0 => cuda::load_quantized(d, as_t_slice::<BlockQ4_0>(&data)),
GgmlDType::Q4_1 => cuda::load_quantized(d, as_t_slice::<BlockQ4_1>(&data)),
GgmlDType::Q5_0 => cuda::load_quantized(d, as_t_slice::<BlockQ5_0>(&data)),
GgmlDType::Q5_1 => cuda::load_quantized(d, as_t_slice::<BlockQ5_1>(&data)),
GgmlDType::Q8_0 => cuda::load_quantized(d, as_t_slice::<BlockQ8_0>(&data)),
GgmlDType::Q8_1 => cuda::load_quantized(d, as_t_slice::<BlockQ8_1>(&data)),
GgmlDType::Q2K => cuda::load_quantized(d, as_t_slice::<BlockQ2K>(&data)),
GgmlDType::Q3K => cuda::load_quantized(d, as_t_slice::<BlockQ3K>(&data)),
GgmlDType::Q4K => cuda::load_quantized(d, as_t_slice::<BlockQ4K>(&data)),
GgmlDType::Q5K => cuda::load_quantized(d, as_t_slice::<BlockQ5K>(&data)),
GgmlDType::Q6K => cuda::load_quantized(d, as_t_slice::<BlockQ6K>(&data)),
GgmlDType::Q8K => cuda::load_quantized(d, as_t_slice::<BlockQ8K>(&data)),
GgmlDType::IQ4_NL => cuda::load_quantized(d, as_t_slice::<BlockIQ4nl>(&data)),
GgmlDType::IQ4_XS => cuda::load_quantized(d, as_t_slice::<BlockIQ4xs>(&data)),
GgmlDType::MXFP4 => cuda::load_quantized(d, as_t_slice::<BlockMXFP4>(&data)),
GgmlDType::BF16 => cuda::load_quantized(d, as_t_slice::<bf16>(&data)),
GgmlDType::I32 => cuda::load_quantized(d, as_t_slice::<i32>(&data)),
GgmlDType::IQ2_XXS => cuda::load_quantized(d, as_t_slice::<BlockIQ2xxs>(&data)),
GgmlDType::IQ2_XS => cuda::load_quantized(d, as_t_slice::<BlockIQ2xs>(&data)),
GgmlDType::IQ2_S => cuda::load_quantized(d, as_t_slice::<BlockIQ2s>(&data)),
GgmlDType::IQ3_XXS => cuda::load_quantized(d, as_t_slice::<BlockIQ3xxs>(&data)),
GgmlDType::IQ3_S => cuda::load_quantized(d, as_t_slice::<BlockIQ3s>(&data)),
GgmlDType::IQ1_S => cuda::load_quantized(d, as_t_slice::<BlockIQ1s>(&data)),
GgmlDType::IQ1_M => cuda::load_quantized(d, as_t_slice::<BlockIQ1m>(&data)),
GgmlDType::TQ1_0 => cuda::load_quantized(d, as_t_slice::<BlockTQ1_0>(&data)),
GgmlDType::TQ2_0 => cuda::load_quantized(d, as_t_slice::<BlockTQ2_0>(&data)),
GgmlDType::NVFP4 => cuda::load_quantized(d, as_t_slice::<BlockNVFP4>(&data)),
GgmlDType::Q1_0 => cuda::load_quantized(d, as_t_slice::<BlockQ1_0>(&data)),
},
#[cfg(feature = "rocm")]
Device::Rocm(d) => Ok(Self::Rocm(dtype.from_data(data), d.clone())),
#[cfg(feature = "vulkan")]
Device::Vulkan(d) => Ok(Self::Vulkan(dtype.from_data(data), d.clone())),
#[cfg(feature = "wgpu")]
Device::Wgpu(d) => Ok(Self::Wgpu(dtype.from_data(data), d.clone())),
}
}
fn block_size(&self) -> usize {
match self {
QStorage::Cpu(storage) => storage.block_size(),
QStorage::Metal(storage) => storage.dtype().block_size(),
QStorage::Cuda(storage) => storage.dtype().block_size(),
#[cfg(feature = "rocm")]
QStorage::Rocm(storage, _) => storage.block_size(),
#[cfg(feature = "vulkan")]
QStorage::Vulkan(storage, _) => storage.block_size(),
#[cfg(feature = "wgpu")]
QStorage::Wgpu(storage, _) => storage.block_size(),
QStorage::Stream(bank) => bank.dtype().block_size(),
}
}
fn dtype(&self) -> GgmlDType {
match self {
QStorage::Cpu(storage) => storage.dtype(),
QStorage::Metal(storage) => storage.dtype(),
QStorage::Cuda(storage) => storage.dtype(),
#[cfg(feature = "rocm")]
QStorage::Rocm(storage, _) => storage.dtype(),
#[cfg(feature = "vulkan")]
QStorage::Vulkan(storage, _) => storage.dtype(),
#[cfg(feature = "wgpu")]
QStorage::Wgpu(storage, _) => storage.dtype(),
QStorage::Stream(bank) => bank.dtype(),
}
}
fn device(&self) -> Device {
match self {
QStorage::Cpu(_storage) => Device::Cpu,
QStorage::Metal(storage) => Device::Metal(storage.device().clone()),
QStorage::Cuda(storage) => Device::Cuda(storage.device().clone()),
#[cfg(feature = "rocm")]
QStorage::Rocm(_storage, device) => Device::Rocm(device.clone()),
#[cfg(feature = "vulkan")]
QStorage::Vulkan(_storage, device) => Device::Vulkan(device.clone()),
#[cfg(feature = "wgpu")]
QStorage::Wgpu(_storage, device) => Device::Wgpu(device.clone()),
QStorage::Stream(_) => Device::Cpu,
}
}
fn size_in_bytes(&self) -> usize {
match self {
QStorage::Cpu(storage) => storage.storage_size_in_bytes(),
QStorage::Metal(storage) => storage.storage_size_in_bytes(),
QStorage::Cuda(storage) => storage.storage_size_in_bytes(),
#[cfg(feature = "rocm")]
QStorage::Rocm(storage, _) => storage.storage_size_in_bytes(),
#[cfg(feature = "vulkan")]
QStorage::Vulkan(storage, _) => storage.storage_size_in_bytes(),
#[cfg(feature = "wgpu")]
QStorage::Wgpu(storage, _) => storage.storage_size_in_bytes(),
QStorage::Stream(bank) => bank.logical_bytes(),
}
}
fn quantize(&mut self, src: &Storage) -> Result<()> {
match (self, src) {
(QStorage::Cpu(storage), Storage::Cpu(src)) => {
storage.from_float(src.as_slice::<f32>()?);
}
(QStorage::Metal(storage), Storage::Metal(src)) => storage.quantize(src)?,
(QStorage::Cuda(storage), Storage::Cuda(src)) => storage.quantize(src)?,
_ => crate::bail!("Invalid quantize storage locations do not match"),
}
Ok(())
}
fn quantize_imatrix(
&mut self,
src: &Storage,
imatrix_weights: &[f32],
n_per_row: usize,
) -> Result<()> {
match (self, src) {
(QStorage::Cpu(storage), Storage::Cpu(src)) => {
storage.from_float_imatrix(src.as_slice::<f32>()?, imatrix_weights, n_per_row);
}
(QStorage::Metal(storage), Storage::Metal(src)) => {
storage.quantize_imatrix(src, imatrix_weights, n_per_row)?
}
(QStorage::Cuda(storage), Storage::Cuda(src)) => {
storage.quantize_imatrix(src, imatrix_weights, n_per_row)?
}
_ => crate::bail!("Invalid quantize storage locations do not match"),
}
Ok(())
}
fn quantize_onto(&mut self, src: &Storage) -> Result<()> {
match (self, src) {
(QStorage::Cpu(storage), Storage::Cpu(src)) => {
storage.from_float(src.as_slice::<f32>()?);
}
(QStorage::Metal(storage), Storage::Cpu(src)) => storage.quantize_onto(src)?,
(QStorage::Cuda(storage), Storage::Cpu(src)) => storage.quantize_onto(src)?,
_ => crate::bail!("Invalid quantize source storage locations: not on cpu"),
}
Ok(())
}
fn quantize_imatrix_onto(
&mut self,
src: &Storage,
imatrix_weights: &[f32],
n_per_row: usize,
) -> Result<()> {
match (self, src) {
(QStorage::Cpu(storage), Storage::Cpu(src)) => {
storage.from_float_imatrix(src.as_slice::<f32>()?, imatrix_weights, n_per_row);
}
(QStorage::Metal(storage), Storage::Cpu(src)) => {
storage.quantize_imatrix_onto(src, imatrix_weights, n_per_row)?
}
(QStorage::Cuda(storage), Storage::Cpu(src)) => {
storage.quantize_imatrix_onto(src, imatrix_weights, n_per_row)?
}
_ => crate::bail!("Invalid quantize storage locations do not match"),
}
Ok(())
}
fn dequantize(&self, elem_count: usize) -> Result<Storage> {
match self {
QStorage::Cpu(storage) => Ok(Storage::Cpu(storage.dequantize(elem_count)?)),
QStorage::Metal(storage) => Ok(Storage::Metal(storage.dequantize(elem_count)?)),
QStorage::Cuda(storage) => Ok(Storage::Cuda(storage.dequantize(elem_count)?)),
#[cfg(feature = "rocm")]
QStorage::Rocm(storage, device) => {
use crate::backend::BackendDevice;
let cpu = storage.dequantize(elem_count)?;
Ok(Storage::Rocm(device.storage_from_cpu_storage(&cpu)?))
}
#[cfg(feature = "vulkan")]
QStorage::Vulkan(storage, device) => {
let cpu = storage.dequantize(elem_count)?;
Ok(Storage::Vulkan(device.upload_f32(cpu.as_slice::<f32>()?)?))
}
#[cfg(feature = "wgpu")]
QStorage::Wgpu(storage, device) => {
let cpu = storage.dequantize(elem_count)?;
Ok(Storage::Wgpu(device.upload_f32(cpu.as_slice::<f32>()?)?))
}
QStorage::Stream(_) => {
crate::bail!("streaming expert bank has no whole-tensor dequantize; consume it via indexed_moe_forward")
}
}
}
fn data(&self) -> Result<Cow<'_, [u8]>> {
match self {
QStorage::Cpu(storage) => {
let data_ptr = storage.as_ptr();
let size_in_bytes = storage.storage_size_in_bytes();
let data = unsafe { std::slice::from_raw_parts(data_ptr, size_in_bytes) };
Ok(Cow::from(data))
}
QStorage::Cuda(storage) => Ok(Cow::from(storage.data()?)),
QStorage::Metal(storage) => Ok(Cow::from(storage.data()?)),
#[cfg(feature = "rocm")]
QStorage::Rocm(storage, _) => {
let data_ptr = storage.as_ptr();
let size_in_bytes = storage.storage_size_in_bytes();
let data = unsafe { std::slice::from_raw_parts(data_ptr, size_in_bytes) };
Ok(Cow::from(data))
}
#[cfg(feature = "vulkan")]
QStorage::Vulkan(storage, _) => {
let data_ptr = storage.as_ptr();
let size_in_bytes = storage.storage_size_in_bytes();
let data = unsafe { std::slice::from_raw_parts(data_ptr, size_in_bytes) };
Ok(Cow::from(data))
}
#[cfg(feature = "wgpu")]
QStorage::Wgpu(storage, _) => {
let data_ptr = storage.as_ptr();
let size_in_bytes = storage.storage_size_in_bytes();
let data = unsafe { std::slice::from_raw_parts(data_ptr, size_in_bytes) };
Ok(Cow::from(data))
}
QStorage::Stream(_) => {
crate::bail!("streaming expert bank is not resident; consume it via indexed_moe_forward")
}
}
}
pub fn device_ptr(&self) -> Result<*const u8> {
match self {
QStorage::Cuda(storage) => storage.device_ptr(),
#[cfg(feature = "rocm")]
QStorage::Rocm(..) => crate::bail!("not implemented"),
#[cfg(feature = "vulkan")]
QStorage::Vulkan(..) => crate::bail!("not implemented"),
#[cfg(feature = "wgpu")]
QStorage::Wgpu(..) => crate::bail!("not implemented"),
QStorage::Metal(_) | QStorage::Cpu(_) | QStorage::Stream(_) => {
crate::bail!("not implemented");
}
}
}
#[cfg(feature = "cuda")]
pub fn device_ptr_with_guard<'a>(
&'a self,
stream: &'a crate::cuda_backend::cudarc::driver::CudaStream,
) -> Result<(
*const u8,
crate::cuda_backend::cudarc::driver::SyncOnDrop<'a>,
)> {
match self {
QStorage::Cuda(storage) => storage.device_ptr_with_guard(stream),
QStorage::Metal(_) | QStorage::Cpu(_) | QStorage::Stream(_) => {
crate::bail!("not implemented");
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum GgmlDType {
F32,
F16,
BF16,
I32,
Q4_0,
Q4_1,
Q5_0,
Q5_1,
Q8_0,
Q8_1,
Q2K,
Q3K,
Q4K,
Q5K,
Q6K,
Q8K,
#[allow(non_camel_case_types)]
IQ4_NL,
#[allow(non_camel_case_types)]
IQ4_XS,
MXFP4,
#[allow(non_camel_case_types)]
IQ2_XXS,
#[allow(non_camel_case_types)]
IQ2_XS,
#[allow(non_camel_case_types)]
IQ3_XXS,
#[allow(non_camel_case_types)]
IQ1_S,
#[allow(non_camel_case_types)]
IQ3_S,
#[allow(non_camel_case_types)]
IQ2_S,
#[allow(non_camel_case_types)]
IQ1_M,
TQ1_0,
TQ2_0,
NVFP4,
Q1_0,
}
use crate::for_each_quant;
macro_rules! gen_from_u32 {
($($v:ident => $b:ident @ $id:literal),+ $(,)?) => {
pub(crate) fn from_u32(u: u32) -> Result<Self> {
let dtype = match u {
0 => Self::F32,
1 => Self::F16,
30 => Self::BF16,
26 => Self::I32,
$( $id => Self::$v, )+
_ => crate::bail!("unknown dtype for tensor {u}"),
};
Ok(dtype)
}
};
}
macro_rules! gen_to_u32 {
($($v:ident => $b:ident @ $id:literal),+ $(,)?) => {
pub fn to_u32(self) -> u32 {
match self {
Self::F32 => 0,
Self::F16 => 1,
Self::BF16 => 30,
Self::I32 => 26,
$( Self::$v => $id, )+
}
}
};
}
macro_rules! gen_cpu_zeros {
($($v:ident => $b:ident @ $id:literal),+ $(,)?) => {
pub fn cpu_zeros(&self, elem_count: usize) -> Box<dyn QuantizedType> {
match self {
Self::F32 => Box::new(vec![f32::zeros(); elem_count]),
Self::F16 => Box::new(vec![f16::zeros(); elem_count]),
Self::BF16 => Box::new(vec![bf16::zeros(); elem_count]),
Self::I32 => Box::new(vec![0i32; elem_count]),
$( Self::$v => Box::new(vec![<$b>::zeros(); elem_count / <$b>::BLCK_SIZE]), )+
}
}
};
}
macro_rules! gen_from_data {
($($v:ident => $b:ident @ $id:literal),+ $(,)?) => {
pub fn from_data(&self, data: Cow<'_, [u8]>) -> Box<dyn QuantizedType> {
match self {
Self::F32 => Box::new(as_t_slice::<f32>(&data).to_vec()),
Self::F16 => Box::new(as_t_slice::<f16>(&data).to_vec()),
Self::BF16 => Box::new(as_t_slice::<bf16>(&data).to_vec()),
Self::I32 => Box::new(as_t_slice::<i32>(&data).to_vec()),
$( Self::$v => Box::new(as_t_slice::<$b>(&data).to_vec()), )+
}
}
};
}
macro_rules! gen_type_size {
($($v:ident => $b:ident @ $id:literal),+ $(,)?) => {
pub fn type_size(&self) -> usize {
use k_quants::*;
match self {
Self::F32 => 4,
Self::F16 | Self::BF16 => 2,
Self::I32 => 4,
$( Self::$v => std::mem::size_of::<$b>(), )+
}
}
};
}
macro_rules! gen_type_align {
($($v:ident => $b:ident @ $id:literal),+ $(,)?) => {
pub fn type_align(&self) -> usize {
use k_quants::*;
match self {
Self::F32 => std::mem::align_of::<f32>(),
Self::F16 | Self::BF16 => std::mem::align_of::<f16>(),
Self::I32 => std::mem::align_of::<i32>(),
$( Self::$v => std::mem::align_of::<$b>(), )+
}
}
};
}
macro_rules! gen_from_mmap {
($($v:ident => $b:ident @ $id:literal),+ $(,)?) => {
#[allow(clippy::wrong_self_convention)] pub(crate) fn from_mmap(
&self,
mmap: Arc<memmap2::Mmap>,
offset: usize,
n_blocks: usize,
) -> Box<dyn QuantizedType> {
match self {
Self::F32 => Box::new(QMmap::<f32>::new(mmap, offset, n_blocks)),
Self::F16 => Box::new(QMmap::<f16>::new(mmap, offset, n_blocks)),
Self::BF16 => Box::new(QMmap::<bf16>::new(mmap, offset, n_blocks)),
Self::I32 => Box::new(QMmap::<i32>::new(mmap, offset, n_blocks)),
$( Self::$v => Box::new(QMmap::<$b>::new(mmap, offset, n_blocks)), )+
}
}
};
}
impl GgmlDType {
for_each_quant!(gen_from_u32);
for_each_quant!(gen_to_u32);
for_each_quant!(gen_cpu_zeros);
for_each_quant!(gen_from_data);
for_each_quant!(gen_from_mmap);
for_each_quant!(gen_type_size);
for_each_quant!(gen_type_align);
pub fn block_size(&self) -> usize {
match self {
Self::F32 => 1,
Self::F16 | Self::BF16 => 1,
Self::I32 => 1,
Self::Q4_0 => k_quants::QK4_0,
Self::Q4_1 => k_quants::QK4_1,
Self::Q5_0 => k_quants::QK5_0,
Self::Q5_1 => k_quants::QK5_1,
Self::Q8_0 => k_quants::QK8_0,
Self::Q8_1 => k_quants::QK8_1,
Self::IQ4_NL => k_quants::QK4_NL,
Self::MXFP4 => k_quants::QK_MXFP4,
Self::Q1_0 => iq_quants::QK1_0,
Self::NVFP4 => iq_quants::QK_NVFP4,
Self::Q2K
| Self::Q3K
| Self::Q4K
| Self::Q5K
| Self::Q6K
| Self::Q8K
| Self::IQ4_XS
| Self::IQ2_XXS
| Self::IQ2_XS
| Self::IQ3_XXS
| Self::IQ1_S
| Self::IQ3_S
| Self::IQ2_S
| Self::IQ1_M
| Self::TQ1_0
| Self::TQ2_0 => k_quants::QK_K,
}
}
}
pub trait QuantizedType: Send + Sync {
fn dtype(&self) -> GgmlDType;
fn matmul_t(&self, mkn: (usize, usize, usize), lhs: &[f32], dst: &mut [f32]) -> Result<()>;
fn matmul_t_f16(&self, mkn: (usize, usize, usize), lhs: &[f16], dst: &mut [f16]) -> Result<()>;
fn dequantize(&self, elem_count: usize) -> Result<CpuStorage>;
fn storage_size_in_bytes(&self) -> usize;
fn as_ptr(&self) -> *const u8;
fn block_size(&self) -> usize;
#[allow(clippy::wrong_self_convention)]
fn from_float(&mut self, xs: &[f32]);
#[allow(clippy::wrong_self_convention)]
fn from_float_imatrix(&mut self, xs: &[f32], imatrix_weights: &[f32], n_per_row: usize);
fn size(&self) -> usize;
}
impl<T: k_quants::GgmlType + Send + Sync> QuantizedType for Vec<T> {
fn matmul_t(&self, mkn: (usize, usize, usize), lhs: &[f32], dst: &mut [f32]) -> Result<()> {
k_quants::matmul(mkn, lhs, self.as_slice(), dst)
}
fn matmul_t_f16(&self, mkn: (usize, usize, usize), lhs: &[f16], dst: &mut [f16]) -> Result<()> {
k_quants::matmul_f16(mkn, lhs, self.as_slice(), dst)
}
fn size(&self) -> usize {
self.len() * core::mem::size_of::<T>()
}
fn from_float(&mut self, xs: &[f32]) {
T::from_float(xs, self)
}
fn from_float_imatrix(&mut self, xs: &[f32], imatrix_weights: &[f32], n_per_row: usize) {
T::from_float_imatrix(xs, self, imatrix_weights, n_per_row)
}
fn dtype(&self) -> GgmlDType {
T::DTYPE
}
fn block_size(&self) -> usize {
T::BLCK_SIZE
}
fn dequantize(&self, elem_count: usize) -> Result<CpuStorage> {
let mut ys = vec![0.0f32; elem_count];
T::to_float(self.as_slice(), &mut ys);
Ok(CpuStorage::F32(ys))
}
fn storage_size_in_bytes(&self) -> usize {
self.len() * std::mem::size_of::<T>()
}
fn as_ptr(&self) -> *const u8 {
self.as_ptr() as *const u8
}
}
pub struct QMmap<T> {
mmap: Arc<memmap2::Mmap>,
offset: usize,
n_blocks: usize,
_t: std::marker::PhantomData<T>,
}
impl<T> QMmap<T> {
fn new(mmap: Arc<memmap2::Mmap>, offset: usize, n_blocks: usize) -> Self {
Self {
mmap,
offset,
n_blocks,
_t: std::marker::PhantomData,
}
}
#[inline]
fn as_slice(&self) -> &[T] {
let len = self.n_blocks * std::mem::size_of::<T>();
as_t_slice::<T>(&self.mmap[self.offset..self.offset + len])
}
}
impl<T: k_quants::GgmlType + Send + Sync> QuantizedType for QMmap<T> {
fn matmul_t(&self, mkn: (usize, usize, usize), lhs: &[f32], dst: &mut [f32]) -> Result<()> {
k_quants::matmul(mkn, lhs, self.as_slice(), dst)
}
fn matmul_t_f16(&self, mkn: (usize, usize, usize), lhs: &[f16], dst: &mut [f16]) -> Result<()> {
k_quants::matmul_f16(mkn, lhs, self.as_slice(), dst)
}
fn size(&self) -> usize {
self.n_blocks * std::mem::size_of::<T>()
}
fn from_float(&mut self, _xs: &[f32]) {
panic!("QMmap is read-only: cannot quantize into a memory-mapped weight region")
}
fn from_float_imatrix(&mut self, _xs: &[f32], _imatrix_weights: &[f32], _n_per_row: usize) {
panic!("QMmap is read-only: cannot quantize into a memory-mapped weight region")
}
fn dtype(&self) -> GgmlDType {
T::DTYPE
}
fn block_size(&self) -> usize {
T::BLCK_SIZE
}
fn dequantize(&self, elem_count: usize) -> Result<CpuStorage> {
let mut ys = vec![0.0f32; elem_count];
T::to_float(self.as_slice(), &mut ys);
Ok(CpuStorage::F32(ys))
}
fn storage_size_in_bytes(&self) -> usize {
self.n_blocks * std::mem::size_of::<T>()
}
fn as_ptr(&self) -> *const u8 {
self.as_slice().as_ptr() as *const u8
}
}
impl std::fmt::Debug for QTensor {
fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
write!(f, "QTensor[{:?}; {:?}]", self.shape, self.dtype())
}
}
fn check_shape(shape: &Shape, block_size: usize) -> Result<()> {
let dims = shape.dims();
if dims.is_empty() {
crate::bail!("scalar tensor cannot be quantized {shape:?}")
}
if !dims[dims.len() - 1].is_multiple_of(block_size) {
crate::bail!(
"quantized tensor must have their last dim divisible by block size {shape:?} {}",
block_size
)
}
Ok(())
}
impl QTensor {
fn make(storage: QStorage, shape: Shape) -> Self {
Self {
storage,
shape,
banks: ResidentBanks::default(),
}
}
pub fn new<S: Into<Shape>>(storage: QStorage, shape: S) -> Result<Self> {
let shape = shape.into();
check_shape(&shape, storage.block_size())?;
Ok(Self::make(storage, shape))
}
pub fn quantize(src: &Tensor, dtype: GgmlDType) -> Result<Self> {
let shape = src.shape();
let block_size = dtype.block_size();
check_shape(shape, block_size)?;
let src = src.to_dtype(crate::DType::F32)?.flatten_all()?;
let elem_count = shape.elem_count();
if !elem_count.is_multiple_of(block_size) {
crate::bail!(
"tensor size ({shape:?}) is not divisible by block size {}",
block_size
)
}
let mut storage = src.device().qzeros(elem_count, dtype)?;
storage.quantize(&src.storage())?;
Ok(Self::make(storage, shape.clone()))
}
pub fn quantize_imatrix(
src: &Tensor,
imatrix_weights: &[f32],
dtype: GgmlDType,
) -> Result<Self> {
let n_per_row = src.dim(D::Minus1)?;
if imatrix_weights.len() != n_per_row {
crate::bail!(
"imatrix weights must have the same length {} as the last dim of src {}",
imatrix_weights.len(),
src.dim(D::Minus1)?
);
}
let shape = src.shape();
let block_size = dtype.block_size();
check_shape(shape, block_size)?;
let src = src.to_dtype(crate::DType::F32)?.flatten_all()?;
let elem_count = shape.elem_count();
if !elem_count.is_multiple_of(block_size) {
crate::bail!(
"tensor size ({shape:?}) is not divisible by block size {}",
block_size
);
}
let mut storage = src.device().qzeros(elem_count, dtype)?;
storage.quantize_imatrix(&src.storage(), imatrix_weights, n_per_row)?;
Ok(Self::make(storage, shape.clone()))
}
pub fn quantize_imatrix_onto(
src: &Tensor,
imatrix_weights: &[f32],
dtype: GgmlDType,
dev: &Device,
) -> Result<Self> {
if !src.device().is_cpu() {
crate::bail!(
"`quantize_onto` expects a `src` to be on the cpu, got {:?}.",
src.device()
)
}
let n_per_row = src.dim(D::Minus1)?;
if imatrix_weights.len() != n_per_row {
crate::bail!(
"imatrix weights must have the same length {} as the last dim of src {}",
imatrix_weights.len(),
src.dim(D::Minus1)?
);
}
let shape = src.shape();
let block_size = dtype.block_size();
check_shape(shape, block_size)?;
let src = src.to_dtype(crate::DType::F32)?.flatten_all()?;
let elem_count = shape.elem_count();
if !elem_count.is_multiple_of(block_size) {
crate::bail!(
"tensor size ({shape:?}) is not divisible by block size {}",
block_size
)
}
let mut storage = dev.qzeros(elem_count, dtype)?;
storage.quantize_imatrix_onto(&src.storage(), imatrix_weights, n_per_row)?;
Ok(Self::make(storage, shape.clone()))
}
pub fn quantize_onto(src: &Tensor, dtype: GgmlDType, dev: &Device) -> Result<Self> {
if !src.device().is_cpu() {
crate::bail!(
"`quantize_onto` expects a `src` to be on the cpu, got {:?}.",
src.device()
)
}
let shape = src.shape();
let block_size = dtype.block_size();
check_shape(shape, block_size)?;
let src = src.to_dtype(crate::DType::F32)?.flatten_all()?;
let elem_count = shape.elem_count();
if !elem_count.is_multiple_of(block_size) {
crate::bail!(
"tensor size ({shape:?}) is not divisible by block size {}",
block_size
)
}
let mut storage = dev.qzeros(elem_count, dtype)?;
storage.quantize_onto(&src.storage())?;
Ok(Self::make(storage, shape.clone()))
}
pub fn dtype(&self) -> GgmlDType {
self.storage.dtype()
}
pub fn device(&self) -> Device {
self.storage.device()
}
pub fn rank(&self) -> usize {
self.shape.rank()
}
pub fn shape(&self) -> &Shape {
&self.shape
}
pub fn dequantize(&self, device: &Device) -> Result<Tensor> {
let storage = self.storage.dequantize(self.shape.elem_count())?;
let none = crate::op::BackpropOp::none();
crate::tensor::from_storage(storage, self.shape.clone(), none, false).to_device(device)
}
pub fn dequantize_f16(&self, device: &Device) -> Result<Tensor> {
match &self.storage {
QStorage::Cuda(s) => {
let s = s.dequantize_f16(self.shape.elem_count())?;
let none = crate::op::BackpropOp::none();
crate::tensor::from_storage(Storage::Cuda(s), self.shape.clone(), none, false)
.to_device(device)
}
_ => {
let s = self.dequantize(device)?.to_dtype(crate::DType::F16)?;
Ok(s)
}
}
}
pub fn storage_size_in_bytes(&self) -> usize {
self.storage.size_in_bytes()
}
pub fn data(&self) -> Result<Cow<'_, [u8]>> {
self.storage.data()
}
#[cfg(feature = "rocm")]
fn rocm_moe_bank(&self, dev: &crate::RocmDevice) -> Result<std::sync::Arc<crate::RocmStorage>> {
use crate::backend::BackendDevice;
let bank = self.data()?;
cache_or_upload(&self.banks.rocm, bank.as_ref(), |b| {
dev.storage_from_slice(b)
})
}
#[cfg(feature = "vulkan")]
fn vulkan_moe_bank(
&self,
dev: &crate::VulkanDevice,
e_cnt: usize,
n: usize,
k: usize,
) -> Result<std::sync::Arc<crate::VulkanStorage>> {
let bank = self.data()?;
let dt = self.storage.dtype();
cache_or_upload(&self.banks.vulkan, bank.as_ref(), |b| match dt {
GgmlDType::Q8_0 => dev.quantize_q8_blocks(b, e_cnt * n, k),
GgmlDType::Q6K => dev.quantize_q6k(b, e_cnt * n, k),
_ => dev.upload_qweight(b),
})
}
#[cfg(feature = "vulkan")]
fn vulkan_moe_bank_split(
&self,
dev: &crate::VulkanDevice,
e_cnt: usize,
n: usize,
k: usize,
) -> Result<std::sync::Arc<crate::vulkan::MoeBankSplit>> {
let bank = self.data()?;
let dt = self.storage.dtype();
cache_or_upload(&self.banks.vulkan_split, bank.as_ref(), |b| match dt {
GgmlDType::Q4K => dev.quantize_q4k_split(b, e_cnt * n, k),
GgmlDType::Q6K => dev.quantize_q6k_split(b, e_cnt * n, k),
_ => crate::bail!("vulkan_moe_bank_split: unsupported dtype {dt:?}"),
})
}
#[cfg(feature = "wgpu")]
fn wgpu_moe_bank(
&self,
dev: &crate::WgpuDevice,
) -> Result<std::sync::Arc<crate::WgpuStorage>> {
let bank = self.data()?;
cache_or_upload(&self.banks.wgpu, bank.as_ref(), |b| dev.upload_qweight(b))
}
pub fn indexed_moe_forward(&self, x: &Tensor, ids: &Tensor) -> Result<Tensor> {
let ids = &ids.contiguous()?;
match &self.storage {
QStorage::Cuda(s) if cuda::QCudaStorage::supports_indexed_moe(s.dtype()) => {
let out_dtype = x.dtype();
let x = x.to_dtype(crate::DType::F32)?.contiguous()?;
let (x_guard, x_l) = x.storage_and_layout();
let (ids_guard, ids_l) = ids.storage_and_layout();
match (&*x_guard, &*ids_guard) {
(Storage::Cuda(x_storage), Storage::Cuda(ids_storage)) => {
let (storage, out_shape) = s.indexed_moe_forward(
self.shape(),
x_storage,
x_l,
ids_storage,
ids_l,
)?;
crate::tensor::from_storage(
Storage::Cuda(storage),
out_shape,
crate::op::BackpropOp::none(),
false,
)
.to_dtype(out_dtype)
}
_ => {
panic!("Non-cuda indexed_moe_forward is not implemented!");
}
}
}
#[cfg(feature = "cuda")]
QStorage::Cuda(s) if cuda::QCudaStorage::supports_iquant_moe(s.dtype()) => {
let out_dtype = x.dtype();
let (_e_cnt, n, k) = self.shape().dims3()?;
let (t, topk) = ids.dims2()?;
let nrows = t * topk;
if t > 1 {
let x_f32 = x.to_dtype(crate::DType::F32)?.contiguous()?;
let ids_u32 = ids.to_dtype(crate::DType::U32)?.contiguous()?;
let (xs, _) = x_f32.storage_and_layout();
let xc = match &*xs {
Storage::Cuda(c) => c,
_ => crate::bail!("cuda i-quant MoE: x not on cuda after contiguous()"),
};
let (ids_s, _) = ids_u32.storage_and_layout();
let idc = match &*ids_s {
Storage::Cuda(c) => c,
_ => crate::bail!("cuda i-quant MoE: ids not on cuda"),
};
if let Some((st, sh)) = s.moe_iquant_qmmq(
self.shape(),
xc.as_cuda_slice::<f32>()?,
x.shape(),
&idc.as_cuda_slice::<u32>()?.slice(0..),
ids.shape(),
)? {
return crate::tensor::from_storage(
Storage::Cuda(st),
sh,
crate::op::BackpropOp::none(),
false,
)
.to_dtype(out_dtype);
}
}
let sdim = x.dim(1)?; let x_exp = if sdim == topk {
x.clone()
} else {
x.broadcast_as((t, topk, k))?
};
let x_flat = x_exp
.reshape((nrows, k))?
.to_dtype(crate::DType::F32)?
.contiguous()?;
let ids_flat = ids
.reshape((nrows,))?
.to_dtype(crate::DType::U32)?
.contiguous()?;
let (xstore, _) = x_flat.storage_and_layout();
let xc = match &*xstore {
Storage::Cuda(c) => c,
_ => crate::bail!("cuda i-quant MoE: x not on cuda after contiguous()"),
};
let (idstore, _) = ids_flat.storage_and_layout();
let idc = match &*idstore {
Storage::Cuda(c) => c,
_ => crate::bail!("cuda i-quant MoE: ids not on cuda"),
};
let y = s.moe_iquant_dp4a(
&xc.as_cuda_slice::<f32>()?.slice(0..),
&idc.as_cuda_slice::<u32>()?.slice(0..),
nrows,
n,
k,
)?;
let out = crate::tensor::from_storage(
Storage::Cuda(y),
(nrows, n),
crate::op::BackpropOp::none(),
false,
);
out.reshape((t, topk, n))?.to_dtype(out_dtype)
}
#[cfg(feature = "vulkan")]
QStorage::Vulkan(_, vk_dev) if vk_moe_kernel(self.storage.dtype()).is_some() => {
let out_dtype = x.dtype();
let (e_cnt, n, k) = self.shape().dims3()?;
let (t, topk) = ids.dims2()?;
let s = x.dim(1)?; let x_exp = if s == topk {
x.clone()
} else {
x.broadcast_as((t, topk, k))?
};
let nrows = t * topk;
let x_flat = x_exp
.reshape((nrows, k))?
.to_dtype(crate::DType::F32)?
.contiguous()?;
let dt = self.storage.dtype();
let ids_u32 = ids.reshape((nrows,))?.to_dtype(crate::DType::U32)?.contiguous()?;
let y = {
let (store, _) = x_flat.storage_and_layout();
let xv = match &*store {
Storage::Vulkan(v) => v,
_ => crate::bail!("vulkan MoE: x not on vulkan after contiguous()"),
};
let (ids_store, _) = ids_u32.storage_and_layout();
let ids_v = match &*ids_store {
Storage::Vulkan(v) => v,
_ => crate::bail!("vulkan MoE: ids not on vulkan after contiguous()"),
};
match vk_moe_blk_dp4a_kernel(dt, n, k).filter(|_| vk_dev.has_int_dot8()) {
Some((blk, with_xsum)) => {
let bank = self.vulkan_moe_bank_split(vk_dev, e_cnt, n, k)?;
vk_dev.moe_matvec_blk_dp4a_gpu(blk, with_xsum, bank.as_ref(), xv, ids_v, nrows, n, k)?
}
None => match vk_moe_blk_kernel(dt, n, k) {
Some(blk) => {
let bank = self.vulkan_moe_bank_split(vk_dev, e_cnt, n, k)?;
vk_dev.moe_matvec_blk_gpu(blk, bank.as_ref(), xv, ids_v, nrows, n, k)?
}
None => {
let kernel = vk_moe_kernel(dt).unwrap();
let wbank = self.vulkan_moe_bank(vk_dev, e_cnt, n, k)?;
vk_dev.moe_matvec_gpu(kernel, wbank.as_ref(), xv, ids_v, nrows, n, k)?
}
},
}
};
let out = crate::tensor::from_storage(
Storage::Vulkan(y),
(nrows, n),
crate::op::BackpropOp::none(),
false,
);
out.reshape((t, topk, n))?.to_dtype(out_dtype)
}
#[cfg(feature = "wgpu")]
QStorage::Wgpu(_, wgpu_dev) if wgpu_moe_kernel(self.storage.dtype()).is_some() => {
let out_dtype = x.dtype();
let (e_cnt, n, k) = self.shape().dims3()?;
let (t, topk) = ids.dims2()?;
let s = x.dim(1)?; let x_exp = if s == topk {
x.clone()
} else {
x.broadcast_as((t, topk, k))?
};
let nrows = t * topk;
let x_flat = x_exp
.reshape((nrows, k))?
.to_dtype(crate::DType::F32)?
.contiguous()?;
let ids_vec = ids
.reshape((nrows,))?
.to_dtype(crate::DType::U32)?
.to_vec1::<u32>()?;
if let Some(&bad) = ids_vec.iter().find(|&&e| e as usize >= e_cnt) {
crate::bail!("indexed_moe_forward: expert id {bad} >= num_experts {e_cnt}");
}
let kernel = wgpu_moe_kernel(self.storage.dtype()).unwrap();
let wbank = self.wgpu_moe_bank(wgpu_dev)?;
let ids_buf = wgpu_dev.upload_ids(&ids_vec)?;
let y = {
let (store, _) = x_flat.storage_and_layout();
let xv = match &*store {
Storage::Wgpu(v) => v,
_ => crate::bail!("wgpu MoE: x not on wgpu after contiguous()"),
};
wgpu_dev.moe_matvec_gpu(kernel, wbank.as_ref(), xv, &ids_buf, nrows, n, k)?
};
let out = crate::tensor::from_storage(
Storage::Wgpu(y),
(nrows, n),
crate::op::BackpropOp::none(),
false,
);
out.reshape((t, topk, n))?.to_dtype(out_dtype)
}
#[cfg(feature = "rocm")]
QStorage::Rocm(_, rocm_dev)
if crate::RocmQuantType::from_ggml(self.storage.dtype()).is_some() =>
{
let qt = crate::RocmQuantType::from_ggml(self.storage.dtype()).unwrap();
let out_dtype = x.dtype();
let (_e_cnt, n, k) = self.shape().dims3()?;
let (t, topk) = ids.dims2()?;
let s = x.dim(1)?; let x_exp = if s == topk {
x.clone()
} else {
x.broadcast_as((t, topk, k))?
};
let nrows = t * topk;
let use_qmmq = t > 1 && qt.qmmq_capable();
let x_flat = match x_exp.dtype() {
DType::F16 | DType::F32 if use_qmmq => {
x_exp.reshape((nrows, k))?.contiguous()?
}
_ if use_qmmq => x_exp
.reshape((nrows, k))?
.to_dtype(DType::F16)?
.contiguous()?,
DType::BF16 | DType::F16 => x_exp.reshape((nrows, k))?.contiguous()?,
DType::F32 if qt.dp4a_active() => x_exp.reshape((nrows, k))?.contiguous()?,
_ => x_exp
.reshape((nrows, k))?
.to_dtype(DType::F16)?
.contiguous()?,
};
let wbank = self.rocm_moe_bank(rocm_dev)?;
let ids_u32 = ids
.reshape((nrows,))?
.to_dtype(crate::DType::U32)?
.contiguous()?;
let (store, _) = x_flat.storage_and_layout();
let xr = match &*store {
crate::Storage::Rocm(r) => r,
_ => crate::bail!("rocm MoE: x not on rocm after contiguous()"),
};
let (idstore, _) = ids_u32.storage_and_layout();
let idr = match &*idstore {
crate::Storage::Rocm(r) => r,
_ => crate::bail!("rocm MoE: ids not on rocm"),
};
let y = if use_qmmq {
rocm_dev.moe_qmmq_quant(qt, wbank.as_ref(), xr, idr, nrows, n, k)?
} else {
rocm_dev.moe_matvec_quant(qt, wbank.as_ref(), xr, idr, nrows, n, k)?
};
let out = crate::tensor::from_storage(
crate::Storage::Rocm(y),
(nrows, n),
crate::op::BackpropOp::none(),
false,
);
out.reshape((t, topk, n))?.to_dtype(out_dtype)
}
#[cfg(feature = "metal")]
QStorage::Metal(s)
if matches!(&*x.storage(), Storage::Metal(_))
&& matches!(&*ids.storage(), Storage::Metal(_)) =>
{
let out_dtype = x.dtype();
let x = x.contiguous()?;
let (xs_guard, x_l) = x.storage_and_layout();
let (ids_guard, ids_l) = ids.storage_and_layout();
let (Storage::Metal(x_storage), Storage::Metal(ids_storage)) =
(&*xs_guard, &*ids_guard)
else {
unreachable!("metal MoE arm is guarded on Metal x/ids storage");
};
let (storage, out_shape) =
s.indexed_moe_forward(self.shape(), x_storage, x_l, ids_storage, ids_l)?;
let out = crate::tensor::from_storage(
Storage::Metal(storage),
out_shape,
crate::op::BackpropOp::none(),
false,
);
out.to_dtype(out_dtype)
}
QStorage::Stream(bank) => {
let (e_cnt, n, k) = self.shape().dims3()?;
let dtype = bank.dtype();
moe_grouped_per_expert(x, ids, n, k, |eid, device| {
if eid as usize >= e_cnt {
crate::bail!("indexed_moe_forward: expert id {eid} >= num_experts {e_cnt}");
}
let bytes = bank.fetch(eid)?;
QStorage::from_data(std::borrow::Cow::Borrowed(&bytes), device, dtype)
})
}
_ => {
let (e_cnt, n, k) = self.shape().dims3()?;
let dtype = self.storage.dtype();
let all_bytes = self.data()?;
let expert_bytes = all_bytes.len() / e_cnt;
moe_grouped_per_expert(x, ids, n, k, |eid, device| {
let off = eid as usize * expert_bytes;
QStorage::from_data(
std::borrow::Cow::Borrowed(&all_bytes[off..off + expert_bytes]),
device,
dtype,
)
})
}
}
}
pub fn device_ptr(&self) -> Result<*const u8> {
match &self.storage {
QStorage::Cuda(storage) => storage.device_ptr(),
#[cfg(feature = "rocm")]
QStorage::Rocm(..) => crate::bail!("not implemented"),
#[cfg(feature = "vulkan")]
QStorage::Vulkan(..) => crate::bail!("not implemented"),
#[cfg(feature = "wgpu")]
QStorage::Wgpu(..) => crate::bail!("not implemented"),
QStorage::Metal(_) | QStorage::Cpu(_) | QStorage::Stream(_) => {
crate::bail!("not implemented");
}
}
}
#[cfg(feature = "cuda")]
pub fn device_ptr_with_guard<'a>(
&'a self,
stream: &'a crate::cuda_backend::cudarc::driver::CudaStream,
) -> Result<(
*const u8,
crate::cuda_backend::cudarc::driver::SyncOnDrop<'a>,
)> {
self.storage.device_ptr_with_guard(stream)
}
}
#[derive(Clone, Debug)]
pub enum QMatMul {
QTensor(std::sync::Arc<QTensor>),
Tensor(Tensor),
TensorF16(Tensor),
#[cfg(feature = "vulkan")]
VulkanQuant {
qtensor: std::sync::Arc<QTensor>,
wq: std::sync::Arc<crate::VulkanStorage>,
dtype: GgmlDType,
n: usize,
k: usize,
},
#[cfg(feature = "wgpu")]
WgpuQuant {
qtensor: std::sync::Arc<QTensor>,
wq: std::sync::Arc<crate::WgpuStorage>,
dtype: GgmlDType,
n: usize,
k: usize,
},
#[cfg(feature = "rocm")]
RocmQuant {
qtensor: std::sync::Arc<QTensor>,
wq: std::sync::Arc<crate::RocmStorage>,
dtype: GgmlDType,
n: usize,
k: usize,
},
}
#[cfg(any(feature = "rocm", feature = "vulkan", feature = "wgpu"))]
fn cache_or_upload<S>(
slot: &std::sync::OnceLock<std::sync::Arc<S>>,
bank: &[u8],
upload: impl FnOnce(&[u8]) -> Result<S>,
) -> Result<std::sync::Arc<S>> {
if let Some(w) = slot.get() {
return Ok(w.clone());
}
let w = std::sync::Arc::new(upload(bank)?);
Ok(slot.get_or_init(|| w).clone())
}
#[cfg(feature = "vulkan")]
fn vk_moe_kernel(dt: GgmlDType) -> Option<&'static str> {
match dt {
GgmlDType::Q4_0 => Some("moe_matvec_q4_0"),
GgmlDType::Q8_0 => Some("moe_matvec_q8_0"),
GgmlDType::Q4K => Some("moe_matvec_q4k"),
GgmlDType::Q6K => Some("moe_matvec_q6k"),
_ => None,
}
}
#[cfg(feature = "vulkan")]
fn vk_moe_blk_kernel(dt: GgmlDType, n: usize, k: usize) -> Option<&'static str> {
if std::env::var_os("VK_MOE_PACKED").is_some() {
return None;
}
match (dt, n, k) {
(GgmlDType::Q4K, 768, 2048) => Some("moe_matvec_q4k_blk_gu"),
(GgmlDType::Q4K, 2048, 768) => Some("moe_matvec_q4k_blk_dn"),
(GgmlDType::Q6K, 2048, 768) => Some("moe_matvec_q6k_blk_dn"),
_ => None,
}
}
fn vk_moe_blk_dp4a_kernel(dt: GgmlDType, n: usize, k: usize) -> Option<(&'static str, bool)> {
if std::env::var_os("VK_MOE_PACKED").is_some() || std::env::var_os("VK_MOE_DP4A_OFF").is_some() {
return None;
}
match (dt, n, k) {
(GgmlDType::Q4K, 768, 2048) => Some(("moe_matvec_q4k_dp4a_blk_gu", true)),
(GgmlDType::Q4K, 2048, 768) => Some(("moe_matvec_q4k_dp4a_blk_dn", true)),
(GgmlDType::Q6K, 2048, 768) => Some(("moe_matvec_q6k_dp4a_blk_dn", false)),
_ => None,
}
}
#[cfg(feature = "wgpu")]
fn wgpu_moe_kernel(dt: GgmlDType) -> Option<&'static str> {
match dt {
GgmlDType::Q4_0 => Some("moe_matvec_q4_0"),
GgmlDType::Q8_0 => Some("moe_matvec_q8_0"),
GgmlDType::Q4K => Some("moe_matvec_q4k"),
_ => None,
}
}
fn moe_grouped_per_expert(
x: &Tensor,
ids: &Tensor,
n: usize,
k: usize,
mut make_storage: impl FnMut(u32, &Device) -> Result<QStorage>,
) -> Result<Tensor> {
use crate::Module; use std::collections::HashMap;
use std::sync::Arc;
let device = x.device();
let out_dtype = x.dtype();
let (t, topk) = ids.dims2()?;
let s = x.dim(1)?; let x_exp = if s == topk {
x.clone()
} else {
x.broadcast_as((t, topk, k))?
};
let x_flat = x_exp
.reshape((t * topk, k))?
.to_dtype(DType::F32)?
.contiguous()?;
let ids_flat = ids.reshape((t * topk,))?.to_dtype(DType::U32)?;
let ids_vec = ids_flat.to_vec1::<u32>()?;
let mut groups: HashMap<u32, Vec<u32>> = HashMap::new();
for (slot, eid) in ids_vec.iter().enumerate() {
groups.entry(*eid).or_default().push(slot as u32);
}
let mut out_flat = Tensor::zeros((t * topk, n), DType::F32, device)?;
for (eid, slots) in groups.into_iter() {
let qs = make_storage(eid, device)?;
let shape: crate::Shape = (n, k).into();
let w_e = QTensor::make(qs, shape);
let qm = QMatMul::from_arc(Arc::new(w_e))?;
let m = slots.len();
let idx = Tensor::from_vec(slots, (m,), device)?;
let x_e = x_flat.index_select(&idx, 0)?; let y_e = qm.forward(&x_e)?.to_dtype(DType::F32)?; out_flat = out_flat.index_add(&idx, &y_e, 0)?;
}
out_flat.reshape((t, topk, n))?.to_dtype(out_dtype)
}
#[cfg_attr(not(feature = "rocm"), allow(unused_variables))]
pub fn moe_combine(ys: &Tensor, scores: &Tensor) -> Result<Tensor> {
let (t, topk, n) = ys.dims3()?;
#[cfg(feature = "rocm")]
if let Device::Rocm(dev) = ys.device() {
let ys_c = ys.contiguous()?;
let scores_c = scores.to_dtype(DType::F32)?.contiguous()?;
let (ys_store, _) = ys_c.storage_and_layout();
let yr = match &*ys_store {
Storage::Rocm(r) => r,
_ => crate::bail!("moe_combine: ys not on rocm after contiguous()"),
};
let (sc_store, _) = scores_c.storage_and_layout();
let sr = match &*sc_store {
Storage::Rocm(r) => r,
_ => crate::bail!("moe_combine: scores not on rocm after contiguous()"),
};
let out = dev.moe_combine(yr, sr, t, topk, n)?;
return Ok(crate::tensor::from_storage(
Storage::Rocm(out),
(t, n),
crate::op::BackpropOp::none(),
false,
));
}
let out_dtype = ys.dtype();
ys.to_dtype(DType::F32)?
.broadcast_mul(&scores.to_dtype(DType::F32)?.unsqueeze(D::Minus1)?)?
.sum(D::Minus2)?
.to_dtype(out_dtype)
}
#[cfg_attr(not(feature = "rocm"), allow(unused_variables))]
pub fn moe_route(logits: &Tensor, topk: usize, norm: bool) -> Result<(Tensor, Tensor)> {
let (ntok, n_experts) = logits.dims2()?;
#[cfg(feature = "rocm")]
if let Device::Rocm(dev) = logits.device() {
let logits_c = logits.to_dtype(DType::F32)?.contiguous()?;
let (lg_store, _) = logits_c.storage_and_layout();
let lr = match &*lg_store {
Storage::Rocm(r) => r,
_ => crate::bail!("moe_route: logits not on rocm after contiguous()"),
};
let (ids, w) = dev.moe_route(lr, ntok, n_experts, topk, norm)?;
let ids_t = crate::tensor::from_storage(
Storage::Rocm(ids),
(ntok, topk),
crate::op::BackpropOp::none(),
false,
);
let w_t = crate::tensor::from_storage(
Storage::Rocm(w),
(ntok, topk),
crate::op::BackpropOp::none(),
false,
);
return Ok((ids_t, w_t));
}
#[cfg(feature = "cuda")]
if let Device::Cuda(cdev) = logits.device() {
if n_experts <= 256 && topk <= 32 {
let logits_c = logits.to_dtype(DType::F32)?.contiguous()?;
let (lg_store, _) = logits_c.storage_and_layout();
let lr = match &*lg_store {
Storage::Cuda(c) => c,
_ => crate::bail!("moe_route: logits not on cuda after contiguous()"),
};
let lview = lr.as_cuda_slice::<f32>()?.slice(0..);
let (ids, w) = cuda::moe_route(&lview, ntok, n_experts, topk, norm, cdev)?;
let ids_t = crate::tensor::from_storage(
Storage::Cuda(ids),
(ntok, topk),
crate::op::BackpropOp::none(),
false,
);
let w_t = crate::tensor::from_storage(
Storage::Cuda(w),
(ntok, topk),
crate::op::BackpropOp::none(),
false,
);
return Ok((ids_t, w_t));
}
}
#[cfg(feature = "vulkan")]
if let Device::Vulkan(vdev) = logits.device() {
if norm && n_experts == 128 && topk == 8 {
let logits_c = logits.to_dtype(DType::F32)?.contiguous()?;
let (lg_store, _) = logits_c.storage_and_layout();
let lv = match &*lg_store {
Storage::Vulkan(v) => v,
_ => crate::bail!("moe_route: logits not on vulkan after contiguous()"),
};
let (ids, w) = vdev.moe_route_vk(lv, ntok, n_experts, topk)?;
let ids_t = crate::tensor::from_storage(
Storage::Vulkan(ids),
(ntok, topk),
crate::op::BackpropOp::none(),
false,
);
let w_t = crate::tensor::from_storage(
Storage::Vulkan(w),
(ntok, topk),
crate::op::BackpropOp::none(),
false,
);
return Ok((ids_t, w_t));
}
}
let lf = logits.to_dtype(DType::F32)?;
let mx = lf.max_keepdim(D::Minus1)?;
let e = lf.broadcast_sub(&mx)?.exp()?;
let z = e.sum_keepdim(D::Minus1)?;
let p = e.broadcast_div(&z)?;
let (sv, si) = p.sort_last_dim(false)?;
let ids = si.narrow(D::Minus1, 0, topk)?.contiguous()?;
let mut w = sv.narrow(D::Minus1, 0, topk)?.contiguous()?;
if norm {
w = w.broadcast_div(&w.sum_keepdim(D::Minus1)?)?;
}
Ok((ids, w))
}
pub fn moe_gate_up(
x: &Tensor,
ids: &Tensor,
gate: &QMatMul,
up: &QMatMul,
) -> Result<(Tensor, Tensor)> {
#[cfg(feature = "rocm")]
{
if let (QMatMul::QTensor(gq), QMatMul::QTensor(uq)) = (gate, up) {
if let (QStorage::Rocm(_, dev), QStorage::Rocm(..)) = (&gq.storage, &uq.storage) {
let dt = gq.storage.dtype();
if dt == uq.storage.dtype() {
if let Some(qt) = crate::RocmQuantType::from_ggml(dt) {
let (_e, n, k) = gq.shape().dims3()?;
let (t, topk) = ids.dims2()?;
let use_qmmq = t > 1 && qt.qmmq_capable();
if x.dim(1)? == 1 && !use_qmmq {
let nrows = t * topk;
let x_exp = x.broadcast_as((t, topk, k))?;
let x_flat = match x_exp.dtype() {
DType::BF16 | DType::F16 => {
x_exp.reshape((nrows, k))?.contiguous()?
}
DType::F32 if qt.dp4a_active() => {
x_exp.reshape((nrows, k))?.contiguous()?
}
_ => x_exp
.reshape((nrows, k))?
.to_dtype(DType::F16)?
.contiguous()?,
};
let out_dtype = x.dtype();
let ids_u32 =
ids.reshape((nrows,))?.to_dtype(DType::U32)?.contiguous()?;
let gwb = gq.rocm_moe_bank(dev)?;
let uwb = uq.rocm_moe_bank(dev)?;
let (xstore, _) = x_flat.storage_and_layout();
let xr = match &*xstore {
Storage::Rocm(r) => r,
_ => crate::bail!("moe_gate_up: x not on rocm after contiguous()"),
};
let (idstore, _) = ids_u32.storage_and_layout();
let idr = match &*idstore {
Storage::Rocm(r) => r,
_ => crate::bail!("moe_gate_up: ids not on rocm"),
};
let (gy, uy) = dev.moe_matvec_pair(
qt,
gwb.as_ref(),
uwb.as_ref(),
xr,
idr,
nrows,
n,
k,
)?;
let g = crate::tensor::from_storage(
Storage::Rocm(gy),
(nrows, n),
crate::op::BackpropOp::none(),
false,
)
.reshape((t, topk, n))?
.to_dtype(out_dtype)?;
let u = crate::tensor::from_storage(
Storage::Rocm(uy),
(nrows, n),
crate::op::BackpropOp::none(),
false,
)
.reshape((t, topk, n))?
.to_dtype(out_dtype)?;
return Ok((g, u));
}
}
}
}
}
}
#[cfg(feature = "vulkan")]
if std::env::var_os("VK_MOE_GU_FUSE_OFF").is_none() {
if let (QMatMul::QTensor(gq), QMatMul::QTensor(uq)) = (gate, up) {
if let (QStorage::Vulkan(_, dev), QStorage::Vulkan(..)) = (&gq.storage, &uq.storage) {
let dt = gq.storage.dtype();
if dt == uq.storage.dtype() && dev.has_int_dot8() {
let (e_cnt, n, k) = gq.shape().dims3()?;
if uq.shape().dims3()? == (e_cnt, n, k) && x.dim(1)? == 1 {
if let Some((blk, with_xsum)) = vk_moe_blk_dp4a_kernel(dt, n, k) {
let (t, topk) = ids.dims2()?;
let nrows = t * topk;
let x_flat = x
.broadcast_as((t, topk, k))?
.reshape((nrows, k))?
.to_dtype(DType::F32)?
.contiguous()?;
let ids_u32 = ids.reshape((nrows,))?.to_dtype(DType::U32)?.contiguous()?;
let (xstore, _) = x_flat.storage_and_layout();
let xv = match &*xstore {
Storage::Vulkan(v) => v,
_ => crate::bail!("moe_gate_up: x not on vulkan after contiguous()"),
};
let (idstore, _) = ids_u32.storage_and_layout();
let idv = match &*idstore {
Storage::Vulkan(v) => v,
_ => crate::bail!("moe_gate_up: ids not on vulkan"),
};
let (xq, xs, xsum) = dev.quantize_act_q8(xv, nrows, k)?;
let gbank = gq.vulkan_moe_bank_split(dev, e_cnt, n, k)?;
let ubank = uq.vulkan_moe_bank_split(dev, e_cnt, n, k)?;
let out_dtype = x.dtype();
let gy = dev.moe_matvec_blk_dp4a_pre_gpu(
blk, with_xsum, gbank.as_ref(), &xq, &xs, &xsum, idv, nrows, n,
)?;
let uy = dev.moe_matvec_blk_dp4a_pre_gpu(
blk, with_xsum, ubank.as_ref(), &xq, &xs, &xsum, idv, nrows, n,
)?;
let shape = |o| -> Result<Tensor> {
crate::tensor::from_storage(
Storage::Vulkan(o),
(nrows, n),
crate::op::BackpropOp::none(),
false,
)
.reshape((t, topk, n))?
.to_dtype(out_dtype)
};
return Ok((shape(gy)?, shape(uy)?));
}
}
}
}
}
}
Ok((
gate.indexed_moe_forward(x, ids)?,
up.indexed_moe_forward(x, ids)?,
))
}
impl QMatMul {
pub fn from_arc(qtensor: std::sync::Arc<QTensor>) -> Result<Self> {
#[cfg(feature = "vulkan")]
{
let dt = qtensor.dtype();
let native_vk = matches!(
dt,
GgmlDType::Q4_0
| GgmlDType::Q8_0
| GgmlDType::Q4K
| GgmlDType::Q5K
| GgmlDType::Q6K
| GgmlDType::Q2K
| GgmlDType::Q3K
| GgmlDType::IQ4_XS
| GgmlDType::IQ4_NL
| GgmlDType::TQ2_0
| GgmlDType::IQ2_XXS
| GgmlDType::IQ2_S
| GgmlDType::IQ3_XXS
| GgmlDType::IQ3_S
| GgmlDType::IQ1_S
| GgmlDType::IQ1_M
| GgmlDType::IQ2_XS
);
if native_vk {
if let Device::Vulkan(d) = qtensor.device() {
if let Ok((n, k)) = qtensor.shape().dims2() {
let blk = dt.block_size();
if k % blk == 0 {
let bytes = qtensor.data()?;
let wq = match dt {
GgmlDType::Q6K => d.quantize_q6k(&bytes, n, k)?,
GgmlDType::Q3K => d.quantize_q3k(&bytes, n, k)?,
GgmlDType::Q8_0 => d.quantize_q8_blocks(&bytes, n, k)?,
GgmlDType::IQ2_XXS => d.quantize_iq2xxs(&bytes, n, k)?,
GgmlDType::IQ2_XS => d.quantize_iq2xs(&bytes, n, k)?,
GgmlDType::IQ1_M => d.quantize_iq1m(&bytes, n, k)?,
GgmlDType::IQ1_S => d.quantize_iq1s(&bytes, n, k)?,
GgmlDType::IQ3_S => d.quantize_iq3s(&bytes, n, k)?,
GgmlDType::IQ3_XXS => d.quantize_iq3xxs(&bytes, n, k)?,
GgmlDType::IQ2_S => d.quantize_iq2s(&bytes, n, k)?,
_ => d.upload_qweight(&bytes)?,
};
return Ok(Self::VulkanQuant {
qtensor,
wq: std::sync::Arc::new(wq),
dtype: dt,
n,
k,
});
}
}
}
}
}
#[cfg(feature = "wgpu")]
{
let dt = qtensor.dtype();
let native_wgpu = matches!(dt, GgmlDType::Q4_0 | GgmlDType::Q8_0 | GgmlDType::Q4K);
if native_wgpu {
if let Device::Wgpu(d) = qtensor.device() {
if let Ok((n, k)) = qtensor.shape().dims2() {
let blk = dt.block_size();
if k % blk == 0 {
let bytes = qtensor.data()?;
let wq = d.upload_qweight(&bytes)?;
return Ok(Self::WgpuQuant {
qtensor,
wq: std::sync::Arc::new(wq),
dtype: dt,
n,
k,
});
}
}
}
}
}
#[cfg(feature = "rocm")]
{
let dt = qtensor.dtype();
if let Some(qt) = crate::RocmQuantType::from_ggml(dt) {
if let Device::Rocm(d) = qtensor.device() {
if let Ok((n, k)) = qtensor.shape().dims2() {
let blk_ok = k % qt.block_elems() == 0;
if blk_ok {
use crate::backend::BackendDevice;
let bytes = qtensor.data()?;
let wq = d.storage_from_slice(bytes.as_ref())?;
return Ok(Self::RocmQuant {
qtensor,
wq: std::sync::Arc::new(wq),
dtype: dt,
n,
k,
});
}
}
}
}
}
#[cfg(feature = "rocm")]
{
if qtensor.device().is_rocm()
&& qtensor.shape().dims().len() == 3
&& crate::RocmQuantType::from_ggml(qtensor.dtype()).is_some()
{
return Ok(Self::QTensor(qtensor));
}
}
#[cfg(feature = "vulkan")]
{
if qtensor.device().is_vulkan()
&& qtensor.shape().dims().len() == 3
&& vk_moe_kernel(qtensor.dtype()).is_some()
{
return Ok(Self::QTensor(qtensor));
}
}
#[cfg(feature = "wgpu")]
{
if qtensor.device().is_wgpu()
&& qtensor.shape().dims().len() == 3
&& wgpu_moe_kernel(qtensor.dtype()).is_some()
{
return Ok(Self::QTensor(qtensor));
}
}
let dequantize = match qtensor.dtype() {
GgmlDType::F32 | GgmlDType::F16 | GgmlDType::BF16 | GgmlDType::I32 => true,
_ => {
qtensor.device().is_vulkan()
|| qtensor.device().is_wgpu()
|| qtensor.device().is_rocm()
}
};
let t = if dequantize {
if qtensor.device().is_rocm() {
Self::TensorF16(qtensor.dequantize_f16(&qtensor.device())?)
} else {
Self::Tensor(qtensor.dequantize(&qtensor.device())?)
}
} else {
Self::QTensor(qtensor)
};
Ok(t)
}
pub fn from_qtensor(qtensor: QTensor) -> Result<Self> {
Self::from_arc(std::sync::Arc::new(qtensor))
}
pub fn dequantize_f16(&self) -> Result<Tensor> {
match self {
Self::QTensor(t) => t.dequantize_f16(&t.device()),
Self::Tensor(t) => t.to_dtype(DType::F16),
Self::TensorF16(t) => Ok(t.clone()),
#[cfg(feature = "rocm")]
Self::RocmQuant { qtensor, .. } => qtensor.dequantize_f16(&qtensor.device()),
#[cfg(feature = "vulkan")]
Self::VulkanQuant { qtensor, .. } => qtensor.dequantize_f16(&qtensor.device()),
#[cfg(feature = "wgpu")]
Self::WgpuQuant { qtensor, .. } => qtensor.dequantize_f16(&qtensor.device()),
}
}
pub fn forward_via_f16(&self, xs: &Tensor) -> Result<Tensor> {
let w = self.dequantize_f16()?;
let in_dtype = xs.dtype();
let w = match *xs.dims() {
[b1, b2, _, _] => w.broadcast_left((b1, b2))?.t()?,
[bsize, _, _] => w.broadcast_left(bsize)?.t()?,
_ => w.t()?,
};
xs.to_dtype(DType::F16)?.matmul(&w)?.to_dtype(in_dtype)
}
pub fn indexed_moe_forward(&self, x: &Tensor, ids: &Tensor) -> Result<Tensor> {
match self {
Self::QTensor(t) => t.indexed_moe_forward(x, ids),
#[cfg(feature = "rocm")]
Self::RocmQuant {
qtensor, wq, dtype, ..
} if crate::RocmQuantType::from_ggml(*dtype).is_some() => {
let qt = crate::RocmQuantType::from_ggml(*dtype).unwrap();
let wbank = wq.as_ref();
let (_e_cnt, n, k) = qtensor.shape().dims3()?;
let (t, topk) = ids.dims2()?;
let s = x.dim(1)?; let x_exp = if s == topk {
x.clone()
} else {
x.broadcast_as((t, topk, k))?
};
let nrows = t * topk;
let use_qmmq = t > 1 && qt.qmmq_capable();
let x_flat = match x_exp.dtype() {
DType::F16 | DType::F32 if use_qmmq => {
x_exp.reshape((nrows, k))?.contiguous()?
}
_ if use_qmmq => x_exp
.reshape((nrows, k))?
.to_dtype(DType::F16)?
.contiguous()?,
DType::BF16 | DType::F16 => x_exp.reshape((nrows, k))?.contiguous()?,
DType::F32 if qt.dp4a_active() => x_exp.reshape((nrows, k))?.contiguous()?,
_ => x_exp
.reshape((nrows, k))?
.to_dtype(DType::F16)?
.contiguous()?,
};
let out_dtype = x.dtype();
let ids_u32 = ids
.reshape((nrows,))?
.to_dtype(crate::DType::U32)?
.contiguous()?;
let (xstore, _) = x_flat.storage_and_layout();
let xr = match &*xstore {
crate::Storage::Rocm(r) => r,
_ => crate::bail!("rocm MoE: x not on rocm after contiguous()"),
};
let (idstore, _) = ids_u32.storage_and_layout();
let idr = match &*idstore {
crate::Storage::Rocm(r) => r,
_ => crate::bail!("rocm MoE: ids not on rocm"),
};
let y = if use_qmmq {
wbank
.device
.moe_qmmq_quant(qt, wbank, xr, idr, nrows, n, k)?
} else {
wbank
.device
.moe_matvec_quant(qt, wbank, xr, idr, nrows, n, k)?
};
let out = crate::tensor::from_storage(
crate::Storage::Rocm(y),
(nrows, n),
crate::op::BackpropOp::none(),
false,
);
out.reshape((t, topk, n))?.to_dtype(out_dtype)
}
#[cfg(feature = "rocm")]
Self::RocmQuant { qtensor, .. } => qtensor.indexed_moe_forward(x, ids),
#[cfg(feature = "vulkan")]
Self::VulkanQuant { qtensor, .. } => qtensor.indexed_moe_forward(x, ids),
#[cfg(feature = "wgpu")]
Self::WgpuQuant { qtensor, .. } => qtensor.indexed_moe_forward(x, ids),
_ => {
panic!("Not implemented!")
}
}
}
}
impl crate::CustomOp1 for QTensor {
fn name(&self) -> &'static str {
"qmatmul"
}
fn cpu_fwd(
&self,
storage: &crate::CpuStorage,
layout: &crate::Layout,
) -> Result<(crate::CpuStorage, Shape)> {
if !layout.is_contiguous() {
crate::bail!("input tensor is not contiguous {layout:?}")
}
let src_shape = layout.shape();
let (n, k) = self.shape.dims2()?;
if src_shape.rank() < 2 {
crate::bail!("input tensor has only one dimension {layout:?}")
}
let mut dst_shape = src_shape.dims().to_vec();
let last_k = dst_shape.pop().unwrap();
if last_k != k {
crate::bail!("input tensor {layout:?} incompatible with {:?}", self.shape)
}
dst_shape.push(n);
let dst_shape = Shape::from(dst_shape);
#[allow(clippy::infallible_destructuring_match)]
let self_storage = match &self.storage {
QStorage::Cpu(storage) => storage,
#[cfg(feature = "rocm")]
QStorage::Rocm(..) => crate::bail!("Invalid storage"),
#[cfg(feature = "vulkan")]
QStorage::Vulkan(..) => crate::bail!("Invalid storage"),
#[cfg(feature = "wgpu")]
QStorage::Wgpu(..) => crate::bail!("Invalid storage"),
QStorage::Metal(_) | QStorage::Cuda(_) | QStorage::Stream(_) => {
crate::bail!("Invalid storage")
}
};
match storage.dtype() {
DType::F32 => {
let slice = storage.as_slice::<f32>()?;
let slice =
&slice[layout.start_offset()..layout.start_offset() + src_shape.elem_count()];
let mut dst_storage = vec![0f32; dst_shape.elem_count()];
self_storage.matmul_t(
(dst_shape.elem_count() / n, k, n),
slice,
&mut dst_storage,
)?;
Ok((crate::CpuStorage::F32(dst_storage), dst_shape))
}
DType::F16 => {
let slice = storage.as_slice::<f16>()?;
let slice =
&slice[layout.start_offset()..layout.start_offset() + src_shape.elem_count()];
let mut dst_storage = vec![f16::ZERO; dst_shape.elem_count()];
self_storage.matmul_t_f16(
(dst_shape.elem_count() / n, k, n),
slice,
&mut dst_storage,
)?;
Ok((crate::CpuStorage::F16(dst_storage), dst_shape))
}
_ => crate::bail!("Expected f32/f16"),
}
}
fn metal_fwd(
&self,
storage: &crate::MetalStorage,
layout: &crate::Layout,
) -> Result<(crate::MetalStorage, Shape)> {
let self_storage = match &self.storage {
QStorage::Metal(metal) => metal,
_ => unreachable!("Cannot call metal matmul on non metal QTensor"),
};
self_storage.fwd(&self.shape, storage, layout)
}
fn cuda_fwd(
&self,
storage: &crate::CudaStorage,
layout: &crate::Layout,
) -> Result<(crate::CudaStorage, Shape)> {
let self_storage = match &self.storage {
QStorage::Cuda(cuda) => cuda,
_ => unreachable!("Cannot call cuda matmul on non cuda QTensor"),
};
self_storage.fwd(&self.shape, storage, layout)
}
}
fn dense_matmul(xs: &Tensor, w: &Tensor) -> Result<Tensor> {
let k = *w.dims().last().unwrap();
let rows = xs.elem_count() / k;
if rows == 1 && xs.device().is_rocm() {
let n = w.dim(0)?;
#[cfg(feature = "rocm")]
{
let d = match xs.device() {
Device::Rocm(d) => d.clone(),
_ => unreachable!(),
};
let xs1 = xs.reshape((k,))?.to_dtype(w.dtype())?.contiguous()?;
let w = w.contiguous()?;
let (wstore, _) = w.storage_and_layout();
let wr = match &*wstore {
crate::Storage::Rocm(r) => r,
_ => crate::bail!("dense_matmul: weight not on rocm"),
};
let (xstore, _) = xs1.storage_and_layout();
let xr = match &*xstore {
crate::Storage::Rocm(r) => r,
_ => crate::bail!("dense_matmul: x not on rocm"),
};
let y = d.dense_gemv(wr, xr, n, k)?;
let mut dims = xs.dims().to_vec();
*dims.last_mut().unwrap() = n;
return crate::tensor::from_storage(
crate::Storage::Rocm(y),
dims,
crate::op::BackpropOp::none(),
false,
)
.to_dtype(xs.dtype());
}
#[cfg(not(feature = "rocm"))]
{
let out = xs.reshape((1, k))?.broadcast_mul(w)?.sum(D::Minus1)?;
let mut dims = xs.dims().to_vec();
*dims.last_mut().unwrap() = n;
return out.reshape(dims);
}
}
let w = match *xs.dims() {
[b1, b2, _, _] => w.broadcast_left((b1, b2))?.t()?,
[bsize, _, _] => w.broadcast_left(bsize)?.t()?,
_ => w.t()?,
};
xs.matmul(&w)
}
#[cfg(feature = "vulkan")]
fn vulkan_prefill_gemm_max_rows(dtype: GgmlDType) -> usize {
match dtype {
GgmlDType::Q4_0 | GgmlDType::Q8_0 | GgmlDType::Q4K | GgmlDType::Q5K | GgmlDType::Q6K => {
usize::MAX
}
_ => 0,
}
}
#[cfg(feature = "vulkan")]
fn vulkan_act_offset0(xs: &Tensor) -> Result<Tensor> {
let xs = xs.contiguous()?;
if xs.layout().start_offset() == 0 {
Ok(xs)
} else {
xs.force_contiguous()
}
}
impl crate::Module for QMatMul {
fn forward(&self, xs: &Tensor) -> Result<Tensor> {
match self {
#[cfg(feature = "rocm")]
Self::RocmQuant {
qtensor,
wq,
dtype,
n,
k,
} => {
let xs_recovered = if xs.device().is_rocm() {
None
} else {
Some(xs.to_device(&qtensor.device())?)
};
let xs = xs_recovered.as_ref().unwrap_or(xs);
let rows: usize = xs.elem_count() / *k;
#[cfg(feature = "rocm")]
let unified_qt = crate::RocmQuantType::from_ggml(*dtype);
#[cfg(not(feature = "rocm"))]
let unified_qt: Option<()> = None;
#[cfg(feature = "rocm")]
let qmmq_ok = unified_qt.map(|qt| qt.qmmq_capable()).unwrap_or(false);
#[cfg(not(feature = "rocm"))]
let qmmq_ok = false;
if rows == 1 && unified_qt.is_some() {
#[cfg(feature = "rocm")]
let keep_f32 = unified_qt.map(|qt| qt.dp4a_active()).unwrap_or(false);
#[cfg(not(feature = "rocm"))]
let keep_f32 = false;
let xs = match xs.dtype() {
DType::BF16 | DType::F16 => xs.contiguous()?,
DType::F32 if keep_f32 => xs.contiguous()?,
_ => xs.to_dtype(DType::F16)?.contiguous()?,
};
let d = match xs.device() {
Device::Rocm(d) => d,
_ => crate::bail!("RocmQuant input not on rocm"),
};
let y = {
let (store, _) = xs.storage_and_layout();
let xr = match &*store {
crate::Storage::Rocm(r) => r,
_ => crate::bail!("RocmQuant expected rocm storage"),
};
#[cfg(feature = "rocm")]
{
d.matvec_quant(unified_qt.unwrap(), wq, xr, *n, *k)?
}
#[cfg(not(feature = "rocm"))]
{
crate::bail!("rocm feature disabled")
}
};
let mut dims = xs.dims().to_vec();
let last = dims.len() - 1;
dims[last] = *n;
Ok(crate::tensor::from_storage(
crate::Storage::Rocm(y),
dims,
crate::op::BackpropOp::none(),
false,
))
} else if let Some(qt) = unified_qt.filter(|_| qmmq_ok) {
let xs = xs.to_dtype(DType::F16)?.contiguous()?;
let d = match xs.device() {
Device::Rocm(d) => d,
_ => crate::bail!("RocmQuant input not on rocm"),
};
let m = xs.elem_count() / *k;
let y = {
let (store, _) = xs.storage_and_layout();
let xr = match &*store {
crate::Storage::Rocm(r) => r,
_ => crate::bail!("RocmQuant expected rocm storage"),
};
#[cfg(feature = "rocm")]
{
d.qmmq_quant(qt, xr, wq, m, *n, *k)?
}
#[cfg(not(feature = "rocm"))]
{
let _ = qt;
crate::bail!("rocm feature disabled")
}
};
let mut dims = xs.dims().to_vec();
let last = dims.len() - 1;
dims[last] = *n;
Ok(crate::tensor::from_storage(
crate::Storage::Rocm(y),
dims,
crate::op::BackpropOp::none(),
false,
))
} else {
let w = qtensor.dequantize_f16(&xs.device())?;
let w = match *xs.dims() {
[b1, b2, _, _] => w.broadcast_left((b1, b2))?.t()?,
[bsize, _, _] => w.broadcast_left(bsize)?.t()?,
_ => w.t()?,
};
xs.to_dtype(DType::F16)?.matmul(&w)
}
}
#[cfg(feature = "vulkan")]
Self::VulkanQuant {
qtensor,
wq,
dtype,
n,
k,
} => {
let vdev = qtensor.device();
let xs_on_vdev = if xs.device().same_device(&vdev) {
None
} else {
Some(xs.to_device(&vdev)?)
};
let xs = xs_on_vdev.as_ref().unwrap_or(xs);
let rows: usize = xs.elem_count() / *k;
if rows == 1 {
let xs = vulkan_act_offset0(xs)?;
let d = match xs.device() {
Device::Vulkan(d) => d,
_ => crate::bail!("VulkanQuant input not on vulkan"),
};
let y = {
let (store, _) = xs.storage_and_layout();
let xv = match &*store {
crate::Storage::Vulkan(v) => v,
_ => crate::bail!("VulkanQuant expected vulkan storage"),
};
match dtype {
GgmlDType::Q4_0 => d.matvec_q4_0_gpu(wq, xv, *n, *k)?,
GgmlDType::Q8_0 => d.matvec_q8_gpu(wq, xv, *n, *k)?,
GgmlDType::Q4K => d.matvec_q4k_gpu(wq, xv, *n, *k)?,
GgmlDType::Q5K => d.matvec_q5k_gpu(wq, xv, *n, *k)?,
GgmlDType::Q6K => d.matvec_q6k_gpu(wq, xv, *n, *k)?,
GgmlDType::Q2K => d.matvec_q2k_gpu(wq, xv, *n, *k)?,
GgmlDType::Q3K => d.matvec_q3k_gpu(wq, xv, *n, *k)?,
GgmlDType::IQ4_XS => d.matvec_iq4xs_gpu(wq, xv, *n, *k)?,
GgmlDType::IQ4_NL => d.matvec_iq4nl_gpu(wq, xv, *n, *k)?,
GgmlDType::IQ2_XXS => d.matvec_iq2xxs_gpu(wq, xv, *n, *k)?,
GgmlDType::IQ2_XS => d.matvec_iq2xs_gpu(wq, xv, *n, *k)?,
GgmlDType::IQ1_M => d.matvec_iq1m_gpu(wq, xv, *n, *k)?,
GgmlDType::IQ1_S => d.matvec_iq1s_gpu(wq, xv, *n, *k)?,
GgmlDType::IQ3_S => d.matvec_iq3s_gpu(wq, xv, *n, *k)?,
GgmlDType::IQ3_XXS => d.matvec_iq3xxs_gpu(wq, xv, *n, *k)?,
GgmlDType::IQ2_S => d.matvec_iq2s_gpu(wq, xv, *n, *k)?,
GgmlDType::TQ2_0 => d.matvec_tq2_0_gpu(wq, xv, *n, *k)?,
other => crate::bail!("VulkanQuant: no native matvec for {other:?}"),
}
};
let mut dims = xs.dims().to_vec();
let last = dims.len() - 1;
dims[last] = *n;
Ok(crate::tensor::from_storage(
crate::Storage::Vulkan(y),
dims,
crate::op::BackpropOp::none(),
false,
))
} else if rows <= vulkan_prefill_gemm_max_rows(*dtype) {
let m = rows;
let xs = vulkan_act_offset0(xs)?;
let d = match xs.device() {
Device::Vulkan(d) => d,
_ => crate::bail!("VulkanQuant input not on vulkan"),
};
let y = {
let (store, _) = xs.storage_and_layout();
let xv = match &*store {
crate::Storage::Vulkan(v) => v,
_ => crate::bail!("VulkanQuant expected vulkan storage"),
};
match dtype {
GgmlDType::Q4_0 => d.matmul_q4_0_gpu(wq, xv, m, *n, *k)?,
GgmlDType::Q8_0 => d.matmul_q8_gpu(wq, xv, m, *n, *k)?,
GgmlDType::Q4K
if *n == 2048
&& *k == 2048
&& std::env::var_os("VK_MMQ_Q4K").is_some() =>
{
let bank = qtensor.vulkan_moe_bank_split(d, 1, *n, *k)?;
let (xq, xsq, xsum) = d.quantize_act_q8(xv, m, *k)?;
d.mmq_q4k_gpu(&xq, &xsq, &xsum, bank.as_ref(), m, *n)?
}
GgmlDType::Q4K => d.matmul_q4k_gpu(wq, xv, m, *n, *k)?,
GgmlDType::Q5K => d.matmul_q5k_gpu(wq, xv, m, *n, *k)?,
GgmlDType::Q6K => d.matmul_q6k_gpu(wq, xv, m, *n, *k)?,
other => crate::bail!("VulkanQuant: no native matmul for {other:?}"),
}
};
let mut dims = xs.dims().to_vec();
let last = dims.len() - 1;
dims[last] = *n;
Ok(crate::tensor::from_storage(
crate::Storage::Vulkan(y),
dims,
crate::op::BackpropOp::none(),
false,
))
} else {
let w = qtensor.dequantize(&xs.device())?;
let w = match *xs.dims() {
[b1, b2, _, _] => w.broadcast_left((b1, b2))?.t()?,
[bsize, _, _] => w.broadcast_left(bsize)?.t()?,
_ => w.t()?,
};
xs.matmul(&w)
}
}
#[cfg(feature = "wgpu")]
Self::WgpuQuant {
qtensor,
wq,
dtype,
n,
k,
} => {
let vdev = qtensor.device();
let xs_on_vdev = if xs.device().same_device(&vdev) {
None
} else {
Some(xs.to_device(&vdev)?)
};
let xs = xs_on_vdev.as_ref().unwrap_or(xs);
let rows: usize = xs.elem_count() / *k;
if rows == 1 {
let xs = xs.contiguous()?;
let d = match xs.device() {
Device::Wgpu(d) => d,
_ => crate::bail!("WgpuQuant input not on wgpu"),
};
let y = {
let (store, _) = xs.storage_and_layout();
let xv = match &*store {
crate::Storage::Wgpu(v) => v,
_ => crate::bail!("WgpuQuant expected wgpu storage"),
};
match dtype {
GgmlDType::Q4_0 => d.matvec_q4_0_gpu(wq, xv, *n, *k)?,
GgmlDType::Q8_0 => d.matvec_q8_0_gpu(wq, xv, *n, *k)?,
GgmlDType::Q4K => d.matvec_q4k_gpu(wq, xv, *n, *k)?,
other => crate::bail!("WgpuQuant: no native matvec for {other:?}"),
}
};
let mut dims = xs.dims().to_vec();
let last = dims.len() - 1;
dims[last] = *n;
Ok(crate::tensor::from_storage(
crate::Storage::Wgpu(y),
dims,
crate::op::BackpropOp::none(),
false,
))
} else {
let w = qtensor.dequantize(&xs.device())?;
let w = match *xs.dims() {
[b1, b2, _, _] => w.broadcast_left((b1, b2))?.t()?,
[bsize, _, _] => w.broadcast_left(bsize)?.t()?,
_ => w.t()?,
};
xs.matmul(&w)
}
}
Self::QTensor(t) => xs.apply_op1_no_bwd(t.as_ref()),
Self::Tensor(w) => dense_matmul(xs, w),
Self::TensorF16(w) => {
let in_dtype = xs.dtype();
dense_matmul(&xs.to_dtype(DType::F16)?, w)?.to_dtype(in_dtype)
}
}
}
}