Skip to main content

ruprim_host/quantization/
transfer.rs

1use super::*;
2
3pub fn q_from_data(data: TensorData) -> HostQTensor {
4    let scheme = match data.dtype {
5        DType::QFloat(scheme) => scheme,
6        _ => panic!("Expected quantized dtype, got {:?}", data.dtype),
7    };
8
9    let shape = data.shape.clone();
10    let num_elements = data.num_elements();
11
12    let q_bytes = QuantizedBytes {
13        bytes: data.into_bytes(),
14        scheme,
15        num_elements,
16    };
17
18    let (values, qparams) = q_bytes.into_vec_i8_with_shape(&shape);
19    let tensor_data = TensorData::new(values, shape);
20    let tensor = HostTensor::from_data(tensor_data);
21
22    // Use native storage since we've unpacked to i8
23    let scheme = scheme.with_store(QuantStore::Native);
24
25    HostQTensor::new(tensor, scheme, qparams.scales)
26}
27
28pub async fn q_into_data(tensor: HostQTensor) -> Result<TensorData, ExecutionError> {
29    let shape = tensor.tensor.shape();
30    let scheme = tensor.scheme;
31    let qt = tensor.tensor.to_contiguous();
32    let values: Vec<i8> = qt.storage::<i8>().to_vec();
33
34    Ok(TensorData::quantized(
35        values,
36        shape.to_vec(),
37        scheme,
38        &tensor.scales,
39    ))
40}
41