use crate::{
CubeBackend, CubeDevice, kernel, ops::numeric::empty_device_dtype, tensor::CubeTensor,
};
use burn_backend::cubecl::dtype_to_storage_type;
use burn_backend::{
DType, ExecutionError, Shape, TensorData,
quantization::{QuantStore, params_shape},
};
use burn_backend::{TensorMetadata, ops::QTensorOps, ops::unfold::calculate_unfold_shape};
use burn_std::{
Metadata, QuantValue, ReshapeAnalysis, reshape_analysis, strides,
tensor::{ReshapeAction, contiguous_strides, reshape_action},
};
use cubecl::tensor_vector_size_parallel;
use cubecl::{ir::VectorSize, server::CopyDescriptor};
pub(crate) fn from_data(data: TensorData, device: &CubeDevice) -> CubeTensor {
if data.bytes.property() == burn_std::AllocationProperty::Device {
let _ = data.bytes.read(burn_std::Reader::new());
}
let client = device.client();
let alloc = client.create_tensor(data.bytes, data.shape.clone(), data.dtype.size());
let shape: Shape = (&data.shape).into();
CubeTensor::new(
client,
alloc.memory,
Metadata::new(shape, alloc.strides),
device.clone(),
data.dtype,
)
}
pub(crate) async fn into_data(tensor: CubeTensor) -> Result<TensorData, ExecutionError> {
let tensor = kernel::into_contiguous_aligned(tensor);
let elem_size = tensor.elem_size();
let shape = tensor.meta.shape().clone();
let strides = tensor.meta.strides().clone();
let binding = CopyDescriptor::new(tensor.handle.binding(), shape, strides, elem_size);
let read = match burn_std::runtime_kind() {
burn_std::RuntimeKind::Async => tensor.client.read_one_tensor_async(binding).await,
_ => tensor.client.read_lazy_async(binding).await,
};
let bytes = read.map_err(|err| ExecutionError::WithContext {
reason: format!("{err}"),
})?;
Ok(TensorData::from_bytes(
bytes,
tensor.meta.shape.clone(),
tensor.dtype,
))
}
#[allow(unused, reason = "useful for debugging kernels")]
pub fn into_data_sync(tensor: CubeTensor) -> TensorData {
burn_std::future::block_on(into_data(tensor)).unwrap()
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(level = "trace", skip(tensor, device))
)]
pub(crate) fn to_device(tensor: CubeTensor, device: &CubeDevice) -> CubeTensor {
if &tensor.device == device {
return tensor;
}
if tensor.device.runtime() != device.runtime() {
return to_device_across_runtimes(tensor, device);
}
let mut tensor = match tensor.qparams {
Some(_) => tensor,
None => kernel::into_contiguous_aligned(tensor),
};
let client = device.client();
tensor.to_client(client, device.clone())
}
fn to_device_across_runtimes(tensor: CubeTensor, device: &CubeDevice) -> CubeTensor {
if tensor.qparams.is_some() {
let from = tensor.device.clone();
let data = burn_std::future::block_on(CubeBackend::q_into_data(tensor))
.unwrap_or_else(|err| transfer_failed(&from, device, err));
return CubeBackend::q_from_data(data, device);
}
let bytes = tensor
.client
.read_one(tensor.handle.clone())
.unwrap_or_else(|err| transfer_failed(&tensor.device, device, err));
let client = device.client();
let handle = client.create(bytes);
CubeTensor {
client,
handle,
meta: tensor.meta,
device: device.clone(),
dtype: tensor.dtype,
qparams: tensor.qparams,
}
}
fn transfer_failed(from: &CubeDevice, to: &CubeDevice, err: impl core::fmt::Display) -> ! {
panic!("Failed to read a tensor off {from:?} on the way to {to:?}: {err}")
}
pub(crate) fn empty(shape: Shape, device: &CubeDevice, dtype: DType) -> CubeTensor {
let client = device.client();
let alloc = client.empty_tensor(shape.clone(), dtype.size());
CubeTensor::new(
client,
alloc.memory,
Metadata::new(shape, alloc.strides),
device.clone(),
dtype,
)
}
pub(crate) fn swap_dims(mut tensor: CubeTensor, dim1: usize, dim2: usize) -> CubeTensor {
tensor.meta.swap(dim1, dim2);
if let DType::QFloat(scheme) = &mut tensor.dtype
&& scheme.block_size().is_some()
{
let rank = tensor.meta.num_dims();
scheme.swap_block_dims(rank, dim1, dim2);
let qparams = tensor.qparams.as_mut().unwrap();
qparams.scales.metadata.swap(dim1, dim2);
}
if let DType::QFloat(scheme) = &mut tensor.dtype
&& let QuantStore::PackedU32(packed_dim) | QuantStore::PackedNative(packed_dim) =
&mut scheme.store
{
let rank = tensor.meta.num_dims();
if *packed_dim == rank - dim1 - 1 {
*packed_dim = rank - dim2 - 1;
} else if *packed_dim == rank - dim2 - 1 {
*packed_dim = rank - dim1 - 1;
}
}
tensor
}
pub fn permute(mut tensor: CubeTensor, axes: &[usize]) -> CubeTensor {
tensor.meta.permute(axes).unwrap();
if let DType::QFloat(scheme) = &mut tensor.dtype
&& scheme.block_size().is_some()
{
let rank = tensor.meta.num_dims();
scheme.permute_block_dims(rank, axes);
let qparams = tensor.qparams.as_mut().unwrap();
qparams.scales.metadata.permute(axes).unwrap();
}
if let DType::QFloat(scheme) = &mut tensor.dtype
&& let QuantStore::PackedU32(packed_dim) | QuantStore::PackedNative(packed_dim) =
&mut scheme.store
{
let rank = tensor.meta.num_dims();
let new_pos = axes
.iter()
.position(|axis| *axis == rank - *packed_dim - 1)
.unwrap_or(0);
*packed_dim = rank - new_pos - 1;
}
tensor
}
pub fn permute_nchw_to_nhwc(tensor: CubeTensor) -> CubeTensor {
let rank = tensor.meta.num_dims();
let c_dim = 1;
let mut dims = vec![0];
dims.extend(2..rank);
dims.push(c_dim);
permute(tensor, &dims)
}
pub fn permute_nchw_to_nhwc_shape(shape: Shape) -> Shape {
let rank = shape.num_dims();
let c_dim = 1;
let mut dims = vec![0];
dims.extend(2..rank);
dims.push(c_dim);
shape.permuted(&dims).expect("Shape permute should succeed")
}
pub fn permute_nhwc_to_nchw(tensor: CubeTensor) -> CubeTensor {
let rank = tensor.meta.num_dims();
let c_dim = rank - 1;
let mut dims = vec![0];
dims.push(c_dim);
dims.extend(1..c_dim);
permute(tensor, &dims)
}
pub fn permute_nhwc_to_nchw_shape(shape: Shape) -> Shape {
let rank = shape.num_dims();
let c_dim = rank - 1;
let mut dims = vec![0];
dims.push(c_dim);
dims.extend(1..c_dim);
shape.permuted(&dims).expect("Shape permute should succeed")
}
pub(crate) fn expand(tensor: CubeTensor, target_shape: Shape) -> CubeTensor {
let tensor = crate::kernel::untile(tensor);
let ndims_in = tensor.meta.shape().num_dims();
let ndims_out = target_shape.num_dims();
let mut new_strides = strides![0usize; ndims_out];
let dim_diff = ndims_out.saturating_sub(ndims_in);
let mut tensor_dim_iter = tensor.meta.shape().iter().rev();
for i in (0..ndims_out).rev() {
if i >= dim_diff {
if let Some(&tensor_dim) = tensor_dim_iter.next() {
if tensor_dim == target_shape[i] || tensor_dim == 1 {
new_strides[i] = if tensor_dim == target_shape[i] {
tensor.meta.strides()[i - dim_diff]
} else {
0
};
} else {
panic!(
"Dimension mismatch: cannot broadcast dimension {tensor_dim} of tensor to target shape"
);
}
} else {
new_strides[i] = 0;
}
} else {
new_strides[i] = 0;
}
}
if tensor.qparams.is_some() && tensor.scheme().block_size().is_some() {
todo!()
}
CubeTensor {
client: tensor.client.clone(),
device: tensor.device.clone(),
meta: Box::new(Metadata::new(target_shape, new_strides)),
handle: tensor.handle.clone(),
dtype: tensor.dtype,
qparams: tensor.qparams.clone(),
}
}
pub fn reshape(tensor: CubeTensor, shape: Shape) -> CubeTensor {
let mut tensor = crate::kernel::untile(tensor);
let analysis = reshape_action(tensor.meta.shape(), tensor.meta.strides(), &shape);
match analysis {
ReshapeAction::UpdateStrides { strides } => {
*tensor.meta = Metadata::new(shape, strides);
return tensor;
}
ReshapeAction::NoChange => return tensor,
ReshapeAction::Recompute => (),
}
let out = empty_device_dtype(
tensor.client.clone(),
tensor.device.clone(),
shape,
tensor.dtype,
);
cubecl::std::tensor::copy_into(
&out.client,
tensor.binding(),
out.clone().binding(),
dtype_to_storage_type(out.dtype),
);
out
}
pub fn q_reshape(tensor: CubeTensor, shape: Shape) -> CubeTensor {
let mut tensor = crate::kernel::untile(tensor);
let scheme = tensor.scheme();
let curr_shape = tensor.meta.shape();
let shape_values = match scheme.store {
QuantStore::Native => shape.clone(),
QuantStore::PackedNative(packed_dim) | QuantStore::PackedU32(packed_dim) => {
let rank = shape.num_dims();
let mut shape = shape.clone();
let packed_d = rank - packed_dim - 1;
let num_quants = scheme.num_quants();
if !shape[packed_d].is_multiple_of(num_quants) {
unimplemented!(
"Cannot reshape packed tensor: inner dimension {} is not aligned with packing factor {num_quants}",
shape[packed_d]
);
}
shape[packed_d] = shape[packed_d].div_ceil(num_quants);
shape
}
};
let (values, scales) = tensor.quantized_handles().unwrap();
let analysis_values = reshape_analysis(
values.meta.shape(),
Some(values.meta.strides()),
&shape_values,
);
let action_values =
analysis_values.action(values.meta.shape(), values.meta.strides(), &shape_values);
let n_new_dims = shape.num_dims().saturating_sub(curr_shape.num_dims());
let is_unsqueeze = n_new_dims > 0 && shape[n_new_dims..] == **curr_shape;
if let ReshapeAction::UpdateStrides { .. } = &action_values {
match analysis_values {
ReshapeAnalysis::IsContiguous => {
if let Some(block_size) = scheme.block_size()
&& block_size.len() > 1
&& !is_unsqueeze
{
unimplemented!("Reshape of ND block-quantized tensor is not yet supported.");
}
}
ReshapeAnalysis::Broadcasted => {} ReshapeAnalysis::Split => {
if let Some(block_size) = scheme.block_size()
&& block_size.len() > 1
{
unimplemented!(
"Split reshape of ND block-quantized tensor is not yet supported."
);
}
}
other => unreachable!("Reshape analysis {other:?} should not update strides."),
}
}
let shape_last = *shape.last().unwrap();
let shape_scales = match scheme.block_size() {
None => scales.meta.shape().clone(), Some(block_size) if block_size.len() == 1 && shape_last < (block_size[0] as usize) => {
if scales.meta.shape().num_elements() > 1 {
unimplemented!("Reshape would split a block across multiple rows.");
}
scales.meta.shape().clone()
}
Some(_) => {
params_shape(&shape, &scheme)
}
};
let action_scales = reshape_action(scales.meta.shape(), scales.meta.strides(), &shape_scales);
match (action_values, action_scales) {
(
ReshapeAction::UpdateStrides { strides },
ReshapeAction::UpdateStrides {
strides: scales_strides,
},
) => {
let qparams = tensor.qparams.as_mut().unwrap();
*tensor.meta = Metadata::new(shape, strides);
qparams.scales.metadata = Metadata::new(shape_scales, scales_strides);
}
(ReshapeAction::UpdateStrides { strides }, ReshapeAction::NoChange) => {
*tensor.meta = Metadata::new(shape, strides);
}
(
ReshapeAction::NoChange,
ReshapeAction::UpdateStrides {
strides: scales_strides,
},
) => {
let qparams = tensor.qparams.as_mut().unwrap();
qparams.scales.metadata = Metadata::new(shape_scales, scales_strides);
}
(ReshapeAction::Recompute, _) | (_, ReshapeAction::Recompute) => {
if !is_unsqueeze
&& matches!(
scheme.value,
QuantValue::Q4S | QuantValue::Q4F | QuantValue::Q2S | QuantValue::Q2F
)
{
todo!(
"Reshape with sub-byte values is not supported when the buffer must be recomputed"
)
}
if scheme.block_size().is_some() && shape_scales.num_elements() > 1 {
unimplemented!(
"Cannot reshape a block-quantized tensor when the reshape requires recomputing the buffer."
);
}
tensor = kernel::into_contiguous(tensor);
*tensor.meta = Metadata::new(shape, contiguous_strides(&shape_values));
let qparams = tensor.qparams.as_mut().unwrap();
let strides = contiguous_strides(&shape_scales);
qparams.scales.metadata = Metadata::new(shape_scales, strides);
}
(ReshapeAction::NoChange, ReshapeAction::NoChange) => {}
}
tensor
}
pub(crate) fn max_vector_size(tensor: &CubeTensor) -> VectorSize {
tensor_vector_size_parallel(
tensor.client.io_optimized_vector_sizes(tensor.dtype.size()),
tensor.meta.shape(),
tensor.meta.strides(),
tensor.meta.num_dims() - 1,
)
}
pub(crate) fn max_vector_size_many(tensors: &[&CubeTensor], axis: usize) -> VectorSize {
let vec = tensors
.iter()
.map(|tensor| {
tensor_vector_size_parallel(
tensor.client.io_optimized_vector_sizes(tensor.dtype.size()),
tensor.meta.shape(),
tensor.meta.strides(),
axis,
)
})
.min();
vec.unwrap_or(0)
}
pub fn unfold(tensor: CubeTensor, dim: usize, size: usize, step: usize) -> CubeTensor {
let tensor = crate::kernel::untile(tensor);
let shape = calculate_unfold_shape(tensor.shape(), dim, size, step);
let d_stride = tensor.meta.strides()[dim];
let mut strides = tensor.meta.strides.clone();
strides[dim] = step * d_stride;
strides.push(d_stride);
CubeTensor {
meta: Box::new(Metadata::new(shape, strides)),
client: tensor.client.clone(),
handle: tensor.handle.clone(),
device: tensor.device.clone(),
dtype: tensor.dtype,
qparams: tensor.qparams.clone(),
}
}
#[cfg(all(test, any(feature = "wgpu", feature = "cuda")))]
mod same_runtime_tests {
use super::*;
use burn_backend::{
Tolerance,
quantization::{QuantScheme, QuantValue, ScaleDtype},
};
use burn_std::{FloatDType, TensorData};
#[cfg(feature = "wgpu")]
#[test]
#[ignore = "needs two discrete wgpu adapters"]
fn moves_between_two_wgpu_adapters() {
use cubecl::wgpu::{WgpuDevice, WgpuDeviceKind};
let first = CubeDevice::Wgpu(WgpuDevice::new(WgpuDeviceKind::DiscreteGpu(0)));
let second = CubeDevice::Wgpu(WgpuDevice::new(WgpuDeviceKind::DiscreteGpu(1)));
moves_both_ways(&first, &second);
moves_quantized_both_ways(&first, &second);
}
#[cfg(feature = "cuda")]
#[test]
#[ignore = "needs two CUDA devices"]
fn moves_between_two_cuda_devices() {
use cubecl::cuda::CudaDevice;
let first = CubeDevice::Cuda(CudaDevice { index: 0 });
let second = CubeDevice::Cuda(CudaDevice { index: 1 });
moves_both_ways(&first, &second);
moves_quantized_both_ways(&first, &second);
}
fn moves_both_ways(first: &CubeDevice, second: &CubeDevice) {
for (from, to) in [(first, second), (second, first)] {
let data = TensorData::from([[1.0f32, 2.0, 3.0], [4.0, 5.0, 6.0]]);
let moved = to_device(from_data(data.clone(), from), to);
assert_eq!(&moved.device, to);
into_data_sync(moved).assert_eq(&data, true);
}
}
fn moves_quantized_both_ways(first: &CubeDevice, second: &CubeDevice) {
let per_tensor = QuantScheme::default();
let two_level = QuantScheme::default()
.with_value(QuantValue::Q8S)
.per_block([4], ScaleDtype::F16)
.per_tensor(ScaleDtype::F32);
for (scheme, shape) in [(per_tensor, [3, 12]), (two_level, [2, 12])] {
for (from, to) in [(first, second), (second, first)] {
moves_quantized(quantize(&scheme, shape, from), to);
}
}
for (from, to) in [(first, second), (second, first)] {
let permuted = permute(quantize(&per_tensor, [3, 12], from), &[1, 0]);
moves_quantized(permuted, to);
}
}
fn quantize(scheme: &QuantScheme, shape: [usize; 2], device: &CubeDevice) -> CubeTensor {
let len = shape.iter().product::<usize>();
let values = (0..len)
.map(|i| i as f32 / len as f32 - 0.5)
.collect::<Vec<_>>();
CubeBackend::quantize_dynamic(from_data(TensorData::new(values, shape), device), scheme)
}
fn moves_quantized(quantized: CubeTensor, to: &CubeDevice) {
let regions = |tensor: &CubeTensor| {
(
tensor.handle.size_in_used(),
tensor.scales().map(|scales| scales.handle.size_in_used()),
tensor.global().map(|global| global.handle.size_in_used()),
)
};
let expected = into_data_sync(CubeBackend::dequantize(quantized.clone(), FloatDType::F32));
let source_regions = regions(&quantized);
let moved = to_device(quantized, to);
assert_eq!(&moved.device, to);
assert_eq!(regions(&moved), source_regions);
into_data_sync(CubeBackend::dequantize(moved, FloatDType::F32))
.assert_approx_eq::<f32>(&expected, Tolerance::default());
}
}
#[cfg(all(test, feature = "cpu", feature = "wgpu"))]
mod cross_runtime_tests {
use super::*;
use burn_std::TensorData;
#[test]
fn crosses_between_runtimes_in_both_directions() {
let cpu = CubeDevice::Cpu(Default::default());
let wgpu = CubeDevice::Wgpu(Default::default());
for (from, to) in [(&cpu, &wgpu), (&wgpu, &cpu)] {
let data = TensorData::from([[1.0f32, 2.0, 3.0], [4.0, 5.0, 6.0]]);
let tensor = from_data(data.clone(), from);
let moved = to_device(tensor, to);
assert_eq!(&moved.device, to);
into_data_sync(moved).assert_eq(&data, true);
}
}
#[test]
fn a_non_contiguous_tensor_keeps_its_strides() {
let cpu = CubeDevice::Cpu(Default::default());
let wgpu = CubeDevice::Wgpu(Default::default());
let data = TensorData::from([[1.0f32, 2.0, 3.0], [4.0, 5.0, 6.0]]);
let swapped = swap_dims(from_data(data, &cpu), 0, 1);
let expected = into_data_sync(swapped.clone());
let moved = to_device(swapped, &wgpu);
into_data_sync(moved).assert_eq(&expected, true);
}
}