Skip to main content

ruda_tensor_device/dispatch/
quantized.rs

1use ruda_tensor::{
2    DType, ExecutionError, QTensorPrimitive, Shape, Slice, TensorData, TensorMetadata,
3    TensorPrimitive,
4    ops::{FloatTensorOps, QTensorOps},
5    quantization::{
6        QuantLevel, QuantPropagation, QuantScheme, QuantValue,
7        QuantizationParametersPrimitive,
8    },
9    tensor::{Device, FloatTensor, IntTensor, QuantizedTensor},
10};
11use ruda_core::tensor::FloatDType;
12use ruda_core::{ir::features::Plane as PlaneFeature, quant::scheme::QuantStore};
13
14use crate::{DeviceBackend, DeviceRuntime, FloatElement, IntElement, element::BoolElement, RudaTensor};
15use rublas::tensor_matmul::MatmulStrategy;
16
17use super::{permute, swap_dims};
18
19fn maybe_dequantize_native_fp4<R: DeviceRuntime>(
20    tensor: RudaTensor<R>,
21    dtype: DType,
22    scalar_matmul: bool,
23) -> RudaTensor<R> {
24    let DType::QFloat(scheme) = tensor.dtype else {
25        return tensor;
26    };
27    let QuantStore::PackedNative(packed_dim) = scheme.store else {
28        return tensor;
29    };
30    if scheme.value != QuantValue::E2M1 {
31        return tensor;
32    }
33
34    let packed_axis = tensor.rank() - packed_dim - 1;
35    if scalar_matmul || !tensor.shape()[packed_axis].is_multiple_of(scheme.num_quants()) {
36        ruda_kernel::tensor::dequantize::dequantize(tensor, dtype)
37    } else {
38        tensor
39    }
40}
41
42pub use ruda_kernel::tensor::allocation::{empty_qtensor, empty_qtensor_optimized};
43
44impl<R, F, I, BT> QTensorOps<Self> for DeviceBackend<R, F, I, BT>
45where
46    R: DeviceRuntime,
47    F: FloatElement,
48    I: IntElement,
49    BT: BoolElement,
50{
51    fn q_from_data(data: TensorData, device: &Device<Self>) -> QuantizedTensor<Self> {
52        ruda_kernel::tensor::transfer::q_from_data(data, device)
53    }
54
55    // TODO: quantize_dynamic (we can compute min-max on the fly and scale, especially when not per-tensor)
56
57    fn quantize(
58        tensor: FloatTensor<Self>,
59        scheme: &QuantScheme,
60        qparams: QuantizationParametersPrimitive<Self>,
61    ) -> QuantizedTensor<Self> {
62        ruda_kernel::tensor::quantize::quantize(tensor, scheme, qparams.scales)
63    }
64
65    fn dequantize(tensor: QuantizedTensor<Self>, dtype: FloatDType) -> FloatTensor<Self> {
66        ruda_kernel::tensor::dequantize::dequantize(tensor, dtype.into())
67    }
68
69    fn q_device(tensor: &QuantizedTensor<Self>) -> Device<Self> {
70        tensor.device.clone()
71    }
72
73    fn q_to_device(tensor: QuantizedTensor<Self>, device: &Device<Self>) -> QuantizedTensor<Self> {
74        super::to_device(tensor, device)
75    }
76
77    fn q_reshape(tensor: QuantizedTensor<Self>, shape: Shape) -> QuantizedTensor<Self> {
78        let scheme = *tensor.scheme();
79        match ruda_kernel::tensor::reshape::try_q_reshape(tensor, shape) {
80            Ok(tensor) => tensor,
81            Err((tensor, shape)) => {
82                let tensor = Self::dequantize(tensor, FloatDType::F32);
83                let output = Self::float_reshape(tensor, shape);
84                Self::quantize_dynamic(output, &scheme)
85            }
86        }
87    }
88
89    async fn q_into_data(tensor: QuantizedTensor<Self>) -> Result<TensorData, ExecutionError> {
90        ruda_kernel::tensor::transfer::q_into_data(tensor).await
91    }
92
93    fn q_swap_dims(
94        tensor: QuantizedTensor<Self>,
95        dim1: usize,
96        dim2: usize,
97    ) -> QuantizedTensor<Self> {
98        swap_dims(tensor, dim1, dim2)
99    }
100
101    fn q_permute(tensor: QuantizedTensor<Self>, axes: &[usize]) -> QuantizedTensor<Self> {
102        permute(tensor, axes)
103    }
104
105    fn q_flip(tensor: QuantizedTensor<Self>, axes: &[usize]) -> QuantizedTensor<Self> {
106        let scheme = *tensor.scheme();
107        match scheme.level {
108            QuantLevel::Tensor => ruprim::indexing::quantized_flip(tensor, axes),
109            QuantLevel::Block(_) => {
110                let tensor = Self::dequantize(tensor, FloatDType::F32);
111                let output = Self::float_flip(tensor, axes);
112                Self::quantize_dynamic(output, &scheme)
113            }
114        }
115    }
116
117    fn q_gather(
118        dim: usize,
119        tensor: QuantizedTensor<Self>,
120        indices: IntTensor<Self>,
121    ) -> QuantizedTensor<Self> {
122        let scheme = *tensor.scheme();
123        match scheme.level {
124            QuantLevel::Tensor => ruprim::indexing::quantized_gather(dim, tensor, indices),
125            QuantLevel::Block(_) => {
126                let dtype = ruda_tensor::get_device_settings::<Self>(&tensor.device).float_dtype;
127                let tensor = Self::dequantize(tensor, dtype);
128                let output = Self::float_gather(dim, tensor, indices);
129                Self::quantize_dynamic(output, &scheme)
130            }
131        }
132    }
133
134    fn q_select(
135        tensor: QuantizedTensor<Self>,
136        dim: usize,
137        indices: IntTensor<Self>,
138    ) -> QuantizedTensor<Self> {
139        let scheme = *tensor.scheme();
140        match scheme.level {
141            QuantLevel::Tensor => ruprim::indexing::quantized_select(tensor, dim, indices),
142            QuantLevel::Block(_) => {
143                let tensor = Self::dequantize(tensor, FloatDType::F32);
144                let output = Self::float_select(tensor, dim, indices);
145                Self::quantize_dynamic(output, &scheme)
146            }
147        }
148    }
149
150    fn q_slice(tensor: QuantizedTensor<Self>, slices: &[Slice]) -> QuantizedTensor<Self> {
151        let scheme = *tensor.scheme();
152        match scheme.level {
153            QuantLevel::Tensor => ruprim::indexing::quantized_slice(tensor, slices),
154            QuantLevel::Block(_) => {
155                let tensor = Self::dequantize(tensor, FloatDType::F32);
156                let output = Self::float_slice(tensor, slices);
157                Self::quantize_dynamic(output, &scheme)
158            }
159        }
160    }
161
162    fn q_expand(tensor: QuantizedTensor<Self>, shape: Shape) -> QuantizedTensor<Self> {
163        super::expand(tensor, shape)
164    }
165
166    fn q_matmul(lhs: TensorPrimitive<Self>, rhs: TensorPrimitive<Self>) -> TensorPrimitive<Self> {
167        let (propagation, scheme) = match (&lhs, &rhs) {
168            (TensorPrimitive::QFloat(lhs), _) => (lhs.propagation(), *lhs.scheme()),
169            (_, TensorPrimitive::QFloat(rhs)) => (rhs.propagation(), *rhs.scheme()),
170            _ => unreachable!(),
171        };
172
173        // Inherit precision for mixed inputs, default to `FloatElem` for fully quantized.
174        let out_dtype = match (&lhs, &rhs) {
175            (TensorPrimitive::Float(lhs), _) => lhs.dtype,
176            (_, TensorPrimitive::Float(rhs)) => rhs.dtype,
177            _ => F::dtype(),
178        };
179
180        let (_lhs_dtype, lhs) = match lhs {
181            TensorPrimitive::Float(lhs) => (lhs.dtype, lhs),
182            TensorPrimitive::QFloat(lhs) => (out_dtype, lhs),
183        };
184        let (_rhs_dtype, rhs) = match rhs {
185            TensorPrimitive::Float(rhs) => (rhs.dtype, rhs),
186            TensorPrimitive::QFloat(rhs) => (out_dtype, rhs),
187        };
188        let has_plane_ops = lhs
189            .client
190            .properties()
191            .features
192            .plane
193            .contains(PlaneFeature::Ops);
194        let lhs = maybe_dequantize_native_fp4(lhs, out_dtype, !has_plane_ops);
195        let rhs = maybe_dequantize_native_fp4(rhs, out_dtype, !has_plane_ops);
196
197        let strategy = if has_plane_ops {
198            MatmulStrategy::default()
199        } else {
200            MatmulStrategy::Naive
201        };
202        let out = rublas::tensor_matmul::matmul(lhs, rhs, None, strategy, out_dtype).unwrap();
203
204        match propagation {
205            QuantPropagation::Propagate => {
206                TensorPrimitive::QFloat(Self::quantize_dynamic(out, &scheme))
207            }
208            QuantPropagation::Inhibit => TensorPrimitive::Float(out),
209        }
210    }
211}