burn-cubecl 0.22.0

Generic backend that can be compiled just-in-time to any shader language target
use crate::ops::numeric::empty_device_dtype;
use crate::tensor::CubeTensor;
use alloc::{vec, vec::Vec};
use burn_backend::cubecl::dtype_to_storage_type;
use burn_backend::{DType, TensorMetadata};

/// Convert the tensor back to a higher precision data type.
///
/// # Panics
///
/// A storage-tiled tensor: the kernel reads its values as rows.
pub fn dequantize(tensor: CubeTensor, dtype: DType) -> CubeTensor {
    let scheme = match tensor.dtype {
        DType::QFloat(scheme) => scheme,
        _ => return tensor,
    };
    assert!(
        !tensor.meta.is_tiled(),
        "dequantize: a storage-tiled quantized tensor is read only by the kernel it was tiled for"
    );
    let (tensor, inverse_axes) = match scheme.store {
        cubecl::quant::scheme::QuantStore::PackedU32(dim)
        | cubecl::quant::scheme::QuantStore::PackedNative(dim)
            if dim != 0 =>
        {
            let rank = tensor.rank();
            let packed_axis = rank - dim - 1;
            let mut axes = (0..rank)
                .filter(|axis| *axis != packed_axis)
                .collect::<Vec<_>>();
            axes.push(packed_axis);

            let mut inverse_axes = vec![0; rank];
            for (axis, source_axis) in axes.iter().enumerate() {
                inverse_axes[*source_axis] = axis;
            }

            let tensor = (packed_axis..rank - 1).fold(tensor, |tensor, axis| {
                crate::ops::swap_dims(tensor, axis, axis + 1)
            });

            (tensor, Some(inverse_axes))
        }
        _ => (tensor, None),
    };
    let scheme = tensor.scheme();

    let output = empty_device_dtype(
        tensor.client.clone(),
        tensor.device.clone(),
        tensor.shape(),
        dtype,
    );
    let global = tensor.global();
    let (values, params) = tensor.quantized_handles().unwrap();

    // Innermost first: the block scales, then the per-tensor scale they are normalized against.
    let mut scales = vec![params.binding()];
    scales.extend(global.map(|g| g.binding()));

    cubek::quantization::dequantize::launch_ref(
        &output.client,
        values.binding(),
        output.clone().binding(),
        &scales,
        &scheme,
        dtype_to_storage_type(dtype),
    )
    .expect("Kernel to never fail");

    match inverse_axes {
        Some(axes) => crate::ops::permute(output, &axes),
        None => output,
    }
}