#[cfg(test)]
mod tests {
use oxgpu::{Buffer, ComputeKernel, Context};
#[tokio::test]
async fn test_matrix_scaling() {
let ctx = Context::new().await.unwrap();
let rows = 4usize;
let cols = 4usize;
let input_data: Vec<f32> = (0..rows * cols).map(|x| x as f32).collect();
let buf_input = Buffer::from_slice(&ctx, &input_data).await;
let buf_output = Buffer::<f32>::zeros(&ctx, rows * cols).await;
let scale_factor_data = [2.0f32];
let buf_scale = Buffer::from_slice(&ctx, &[0.0f32]).await; buf_scale.write(&ctx, &scale_factor_data);
let shader = r#"
@group(0) @binding(0) var<storage, read> input: array<f32>;
@group(0) @binding(1) var<storage, read_write> output: array<f32>;
@group(0) @binding(2) var<storage, read> scale: array<f32>;
@compute @workgroup_size(1, 1)
fn main(@builtin(global_invocation_id) global_id: vec3<u32>) {
let x = global_id.x;
let y = global_id.y;
let width = 4u;
if (x < width && y < 4u) {
let index = y * width + x;
output[index] = input[index] * scale[0];
}
}
"#;
let kernel = ComputeKernel::builder()
.source(shader)
.entry_point("main")
.label("MatrixScaleKernel")
.add_storage_read(0) .add_storage_read_write(1) .add_storage_read(2) .build(&ctx)
.await
.unwrap();
kernel.run(&ctx, (4, 4), &[&buf_input, &buf_output, &buf_scale]);
let result = buf_output.read(&ctx).await.unwrap();
let expected: Vec<f32> = input_data.iter().map(|&x| x * 2.0).collect();
assert_eq!(result, expected);
println!(
"Matrix scaling test passed: Input {:?} * 2.0 = {:?}",
input_data, result
);
}
}