Consistent formatting pass: line wrapping, import sorting, trailing whitespace removal, let-chain indentation, merged derive attributes, and unsafe block reformatting. Co-Authored-By: Claude Opus 4.6 (1M context) <[email protected]>
274 lines
8.4 KiB
Rust
274 lines
8.4 KiB
Rust
//! CPU fallback operations for `CubeCL` backend.
|
|
//!
|
|
//! These provide CPU implementations of tensor operations when GPU is unavailable.
|
|
//! When the WGPU feature is enabled, GPU kernels will be used instead.
|
|
|
|
use crate::client;
|
|
use crate::device::CubeclDevice;
|
|
use crate::tensor::CubeclTensorPrimitive;
|
|
|
|
/// Execute an elementwise binary operation.
|
|
pub fn elementwise_binary<const D: usize, F>(
|
|
lhs: CubeclTensorPrimitive<D>,
|
|
rhs: CubeclTensorPrimitive<D>,
|
|
op: F,
|
|
) -> CubeclTensorPrimitive<D>
|
|
where
|
|
F: Fn(f32, f32) -> f32,
|
|
{
|
|
// Get data from both tensors
|
|
let lhs_data = read_tensor_data(&lhs);
|
|
let rhs_data = read_tensor_data(&rhs);
|
|
let device = lhs.device.clone();
|
|
|
|
// Apply operation element-wise
|
|
let result: Vec<f32> = lhs_data
|
|
.iter()
|
|
.zip(rhs_data.iter())
|
|
.map(|(a, b)| op(*a, *b))
|
|
.collect();
|
|
|
|
// Write result to new tensor
|
|
write_tensor_data(lhs.shape, result, &device)
|
|
}
|
|
|
|
/// Execute an elementwise unary operation.
|
|
pub fn elementwise_unary<const D: usize, F>(
|
|
tensor: CubeclTensorPrimitive<D>,
|
|
op: F,
|
|
) -> CubeclTensorPrimitive<D>
|
|
where
|
|
F: Fn(f32) -> f32,
|
|
{
|
|
let data = read_tensor_data(&tensor);
|
|
let device = tensor.device.clone();
|
|
let shape = tensor.shape;
|
|
|
|
let result: Vec<f32> = data.iter().map(|x| op(*x)).collect();
|
|
|
|
write_tensor_data(shape, result, &device)
|
|
}
|
|
|
|
/// Read tensor data as f32 values.
|
|
pub fn read_tensor_data<const D: usize>(tensor: &CubeclTensorPrimitive<D>) -> Vec<f32> {
|
|
if let Some(buffer) = tensor.buffer() {
|
|
// Try to read from GPU
|
|
if let Some(gpu_handle) = buffer.as_gpu()
|
|
&& let Ok(client) = client::get_or_create(&tensor.device)
|
|
&& let Ok(data) = client.read_f32(gpu_handle)
|
|
{
|
|
return data;
|
|
}
|
|
// Read from CPU buffer
|
|
if let Some(bytes) = buffer.as_cpu() {
|
|
return bytes
|
|
.chunks_exact(4)
|
|
.map(|chunk| f32::from_le_bytes(chunk.try_into().unwrap()))
|
|
.collect();
|
|
}
|
|
}
|
|
// No buffer - return zeros
|
|
vec![0.0; tensor.numel()]
|
|
}
|
|
|
|
/// Write f32 data to a new tensor.
|
|
pub fn write_tensor_data<const D: usize>(
|
|
shape: [usize; D],
|
|
data: Vec<f32>,
|
|
device: &CubeclDevice,
|
|
) -> CubeclTensorPrimitive<D> {
|
|
// Try to write to GPU
|
|
if let Ok(client) = client::get_or_create(device)
|
|
&& let Ok(handle) = client.create_f32(&data)
|
|
{
|
|
return CubeclTensorPrimitive::from_handle(shape, handle, device.clone());
|
|
}
|
|
|
|
// CPU fallback
|
|
let bytes: Vec<u8> = data.iter().flat_map(|f| f.to_le_bytes()).collect();
|
|
CubeclTensorPrimitive::from_cpu_data(shape, bytes, device.clone())
|
|
}
|
|
|
|
/// Perform matrix multiplication (2D).
|
|
pub fn matmul(
|
|
lhs: CubeclTensorPrimitive<2>,
|
|
rhs: CubeclTensorPrimitive<2>,
|
|
) -> CubeclTensorPrimitive<2> {
|
|
let m = lhs.shape[0];
|
|
let k = lhs.shape[1];
|
|
let n = rhs.shape[1];
|
|
let device = lhs.device.clone();
|
|
|
|
let lhs_data = read_tensor_data(&lhs);
|
|
let rhs_data = read_tensor_data(&rhs);
|
|
|
|
// Simple naive matmul for correctness
|
|
let mut result = vec![0.0f32; m * n];
|
|
for i in 0..m {
|
|
for j in 0..n {
|
|
let mut sum = 0.0;
|
|
for p in 0..k {
|
|
sum += lhs_data[i * k + p] * rhs_data[p * n + j];
|
|
}
|
|
result[i * n + j] = sum;
|
|
}
|
|
}
|
|
|
|
write_tensor_data([m, n], result, &device)
|
|
}
|
|
|
|
/// Perform batched matrix multiplication (3D).
|
|
pub fn bmm(
|
|
lhs: CubeclTensorPrimitive<3>,
|
|
rhs: CubeclTensorPrimitive<3>,
|
|
) -> CubeclTensorPrimitive<3> {
|
|
let batch = lhs.shape[0];
|
|
let m = lhs.shape[1];
|
|
let k = lhs.shape[2];
|
|
let n = rhs.shape[2];
|
|
let device = lhs.device.clone();
|
|
|
|
let lhs_data = read_tensor_data(&lhs);
|
|
let rhs_data = read_tensor_data(&rhs);
|
|
|
|
let mut result = vec![0.0f32; batch * m * n];
|
|
for b in 0..batch {
|
|
for i in 0..m {
|
|
for j in 0..n {
|
|
let mut sum = 0.0;
|
|
for p in 0..k {
|
|
sum += lhs_data[b * m * k + i * k + p] * rhs_data[b * k * n + p * n + j];
|
|
}
|
|
result[b * m * n + i * n + j] = sum;
|
|
}
|
|
}
|
|
}
|
|
|
|
write_tensor_data([batch, m, n], result, &device)
|
|
}
|
|
|
|
/// Sum all elements.
|
|
pub fn sum<const D: usize>(tensor: CubeclTensorPrimitive<D>) -> CubeclTensorPrimitive<1> {
|
|
let data = read_tensor_data(&tensor);
|
|
let device = tensor.device.clone();
|
|
let result: f32 = data.iter().sum();
|
|
write_tensor_data([1], vec![result], &device)
|
|
}
|
|
|
|
/// Mean of all elements.
|
|
pub fn mean<const D: usize>(tensor: CubeclTensorPrimitive<D>) -> CubeclTensorPrimitive<1> {
|
|
let data = read_tensor_data(&tensor);
|
|
let device = tensor.device.clone();
|
|
let sum: f32 = data.iter().sum();
|
|
let mean = sum / data.len() as f32;
|
|
write_tensor_data([1], vec![mean], &device)
|
|
}
|
|
|
|
/// Max of all elements.
|
|
pub fn max<const D: usize>(tensor: CubeclTensorPrimitive<D>) -> CubeclTensorPrimitive<1> {
|
|
let data = read_tensor_data(&tensor);
|
|
let device = tensor.device.clone();
|
|
let max = data.iter().copied().fold(f32::NEG_INFINITY, f32::max);
|
|
write_tensor_data([1], vec![max], &device)
|
|
}
|
|
|
|
/// Min of all elements.
|
|
pub fn min<const D: usize>(tensor: CubeclTensorPrimitive<D>) -> CubeclTensorPrimitive<1> {
|
|
let data = read_tensor_data(&tensor);
|
|
let device = tensor.device.clone();
|
|
let min = data.iter().copied().fold(f32::INFINITY, f32::min);
|
|
write_tensor_data([1], vec![min], &device)
|
|
}
|
|
|
|
/// Variance of all elements.
|
|
pub fn var<const D: usize>(tensor: CubeclTensorPrimitive<D>) -> CubeclTensorPrimitive<1> {
|
|
let data = read_tensor_data(&tensor);
|
|
let device = tensor.device.clone();
|
|
let n = data.len() as f32;
|
|
let mean: f32 = data.iter().sum::<f32>() / n;
|
|
let variance: f32 = data.iter().map(|x| (x - mean).powi(2)).sum::<f32>() / n;
|
|
write_tensor_data([1], vec![variance], &device)
|
|
}
|
|
|
|
/// Softmax along the last dimension.
|
|
pub fn softmax<const D: usize>(
|
|
tensor: CubeclTensorPrimitive<D>,
|
|
dim: usize,
|
|
) -> CubeclTensorPrimitive<D> {
|
|
let data = read_tensor_data(&tensor);
|
|
let device = tensor.device.clone();
|
|
let shape = tensor.shape;
|
|
|
|
// Simple softmax for last dimension only
|
|
// For full implementation, would need proper dimension handling
|
|
let numel = tensor.numel();
|
|
if D == 0 || dim != D - 1 {
|
|
// Return input as-is for unsupported cases
|
|
return write_tensor_data(shape, data, &device);
|
|
}
|
|
|
|
let last_dim = shape[D - 1];
|
|
let outer = numel / last_dim;
|
|
|
|
let mut result = vec![0.0f32; numel];
|
|
for i in 0..outer {
|
|
let start = i * last_dim;
|
|
let slice = &data[start..start + last_dim];
|
|
|
|
// Find max for numerical stability
|
|
let max_val = slice.iter().copied().fold(f32::NEG_INFINITY, f32::max);
|
|
|
|
// Compute exp and sum
|
|
let exp_vals: Vec<f32> = slice.iter().map(|x| (x - max_val).exp()).collect();
|
|
let sum: f32 = exp_vals.iter().sum();
|
|
|
|
// Normalize
|
|
for (j, exp_val) in exp_vals.iter().enumerate() {
|
|
result[start + j] = exp_val / sum;
|
|
}
|
|
}
|
|
|
|
write_tensor_data(shape, result, &device)
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
|
|
#[test]
|
|
fn test_elementwise_add() {
|
|
let device = CubeclDevice::cpu();
|
|
let a = write_tensor_data([2, 2], vec![1.0, 2.0, 3.0, 4.0], &device);
|
|
let b = write_tensor_data([2, 2], vec![5.0, 6.0, 7.0, 8.0], &device);
|
|
let c = elementwise_binary(a, b, |x, y| x + y);
|
|
let data = read_tensor_data(&c);
|
|
assert_eq!(data, vec![6.0, 8.0, 10.0, 12.0]);
|
|
}
|
|
|
|
#[test]
|
|
fn test_matmul() {
|
|
let device = CubeclDevice::cpu();
|
|
// 2x3 matrix
|
|
let a = write_tensor_data([2, 3], vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &device);
|
|
// 3x2 matrix
|
|
let b = write_tensor_data([3, 2], vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0], &device);
|
|
let c = matmul(a, b);
|
|
assert_eq!(c.shape, [2, 2]);
|
|
let data = read_tensor_data(&c);
|
|
// [1,2,3] * [1,3,5]' = 1+6+15 = 22
|
|
// [1,2,3] * [2,4,6]' = 2+8+18 = 28
|
|
// [4,5,6] * [1,3,5]' = 4+15+30 = 49
|
|
// [4,5,6] * [2,4,6]' = 8+20+36 = 64
|
|
assert_eq!(data, vec![22.0, 28.0, 49.0, 64.0]);
|
|
}
|
|
|
|
#[test]
|
|
fn test_sum() {
|
|
let device = CubeclDevice::cpu();
|
|
let a = write_tensor_data([2, 2], vec![1.0, 2.0, 3.0, 4.0], &device);
|
|
let s = sum(a);
|
|
let data = read_tensor_data(&s);
|
|
assert_eq!(data, vec![10.0]);
|
|
}
|
|
}
|