ruda_tensor_device/dispatch/
quantized.rs1use 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 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 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}