use ruda_kernel::library::tensor::AsView;
use ruda_kernel::library::tensor::AsViewExpand;
use ruda_kernel::library::tensor::AsViewMut;
use ruda_kernel::library::tensor::AsViewMutExpand;
use ruda_kernel::dsl::zspace::shape;
use ruda_kernel::dsl as kernel_dsl;
use ruda_kernel::library::tensor::layout::Coords1d;
use ruda_kernel::library::tensor::layout::Layout;
use ruda_kernel::library::tensor::layout::LayoutExpand;
use ruda_test_runtime::TestRuntime;
use ruda_kernel::dsl::prelude::*;
use ruda_test_utils::{
DataKind, HostData, HostDataType, StrideSpec, TestInput, ValidationResult,
assert_equals_approx, assert_equals_approx_in_slice, print_tensor,
};
#[test]
fn eye_handle_row_major() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let shape = [2, 3];
let handle = TestInput::builder(client.clone(), shape).eye().generate();
let expected = TestInput::builder(client.clone(), [2, 3])
.custom(vec![1., 0., 0., 0., 1., 0.])
.f32_host_data();
let actual = HostData::from_tensor_handle(&client, handle, HostDataType::F32);
assert_equals_approx(&actual, &expected, 0.001)
.as_test_outcome()
.enforce();
}
#[test]
fn eye_handle_col_major() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let shape = [2, 3];
let handle = TestInput::builder(client.clone(), shape)
.stride(StrideSpec::ColMajor)
.eye()
.generate();
let expected = TestInput::builder(client.clone(), [2, 3])
.custom(vec![1., 0., 0., 0., 1., 0.])
.f32_host_data();
let actual = HostData::from_tensor_handle(&client, handle, HostDataType::F32);
assert_equals_approx(&actual, &expected, 0.001)
.as_test_outcome()
.enforce();
}
#[test]
fn arange_handle_row_major() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let shape = shape![2, 3];
let handle = TestInput::builder(client.clone(), shape)
.arange()
.generate();
let expected = TestInput::builder(client.clone(), shape![2, 3])
.custom(vec![0., 1., 2., 3., 4., 5.])
.f32_host_data();
let actual = HostData::from_tensor_handle(&client, handle, HostDataType::F32);
assert_equals_approx(&actual, &expected, 0.001)
.as_test_outcome()
.enforce();
}
#[test]
fn arange_handle_col_major() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let shape = shape![2, 3];
let handle = TestInput::builder(client.clone(), shape)
.stride(StrideSpec::ColMajor)
.arange()
.generate();
let expected = TestInput::builder(client.clone(), shape![2, 3])
.custom(vec![0., 1., 2., 3., 4., 5.])
.f32_host_data();
let actual = HostData::from_tensor_handle(&client, handle, HostDataType::F32);
assert_equals_approx(&actual, &expected, 0.001)
.as_test_outcome()
.enforce();
}
#[test]
fn custom_handle_row_major_col_major() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let contiguous_data = [9., 8., 7., 6., 5., 4.].to_vec();
let (_, row_major) = TestInput::builder(client.clone(), shape![2, 3])
.custom(contiguous_data.clone())
.generate_with_f32_host_data();
let (_, col_major) = TestInput::builder(client.clone(), shape![2, 3])
.stride(StrideSpec::ColMajor)
.custom(contiguous_data)
.generate_with_f32_host_data();
assert_equals_approx(&col_major, &row_major, 0.001)
.as_test_outcome()
.enforce();
}
#[test]
fn arange_handle_row_major_slice() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let shape = shape![2, 3];
let actual_data = vec![0., 1., 2., 9., 9., 9.]; let actual = TestInput::builder(client.clone(), shape.clone())
.custom(actual_data)
.f32_host_data();
let expected_data = vec![0., 1., 2., 3., 4., 5.];
let expected = TestInput::builder(client.clone(), shape)
.custom(expected_data)
.f32_host_data();
assert_equals_approx_in_slice(&actual, &expected, 0.001, vec![0..1, 0..3])
.as_test_outcome()
.enforce();
}
#[test]
fn fail_message_contains_aggregate_stats_and_examples() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let actual = TestInput::builder(client.clone(), shape![2, 3])
.custom(vec![0., 1., 2., 4., 5., 6.])
.f32_host_data();
let expected = TestInput::builder(client.clone(), shape![2, 3])
.custom(vec![0., 1., 2., 3., 4., 5.])
.f32_host_data();
let result = assert_equals_approx(&actual, &expected, 0.001);
let reason = match result {
ValidationResult::Fail(r) => r,
other => panic!("expected Fail, got {other:?}"),
};
assert!(
reason.contains("3/6 elements mismatched"),
"missing mismatch count, got: {reason}"
);
assert!(
reason.contains("max |Δ|="),
"missing max delta, got: {reason}"
);
assert!(
reason.contains("worst at [1, 0]"),
"missing worst index, got: {reason}"
);
let is_print_mode = std::env::var("RUDA_TEST_MODE")
.map(|v| v.to_lowercase().starts_with("print"))
.unwrap_or(false);
if !is_print_mode {
assert!(
reason.contains("First mismatches:"),
"missing examples header, got: {reason}"
);
}
}
#[test]
fn assert_equals_approx_in_slice_accepts_tensor_filter() {
use ruda_test_utils::DimFilter;
let client = <TestRuntime as Runtime>::client(&Default::default());
let actual = TestInput::builder(client.clone(), shape![2, 3])
.custom(vec![0., 1., 2., 9., 9., 9.])
.f32_host_data();
let expected = TestInput::builder(client.clone(), shape![2, 3])
.custom(vec![0., 1., 2., 3., 4., 5.])
.f32_host_data();
let filter = vec![DimFilter::Exact(0), DimFilter::Range { start: 0, end: 2 }];
assert_equals_approx_in_slice(&actual, &expected, 0.001, filter)
.as_test_outcome()
.enforce();
}
#[test]
fn builder_matches_constructor() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let from_builder = TestInput::builder(client.clone(), shape![2, 3])
.arange()
.f32_host_data();
let from_new = TestInput::new(
client.clone(),
shape![2, 3],
f32::as_type_native_unchecked().storage_type(),
StrideSpec::RowMajor,
DataKind::Arange { scale: None },
)
.f32_host_data();
assert_equals_approx(&from_builder, &from_new, 0.0)
.as_test_outcome()
.enforce();
}
#[test]
fn read_rowmajor_tensor_as_tiled() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let matrix_len = 4;
let shape = shape![matrix_len, matrix_len];
let input_handle = TestInput::builder(client.clone(), shape.clone())
.stride(StrideSpec::RowMajor)
.arange()
.generate();
let dtype = f32::as_type_native_unchecked().storage_type();
let output_handle = TestInput::builder(client.clone(), shape.clone())
.stride(StrideSpec::RowMajor)
.zeros()
.generate_without_host_data();
let ruda_count = RudaCount::new_single();
let ruda_dim = RudaDim::new_single();
let vector_size = 1;
launch_read_rowmajor_tensor_as_tiled::launch::<TestRuntime>(
&client,
ruda_count,
ruda_dim,
input_handle.binding().into_tensor_arg(),
output_handle.clone().binding().into_tensor_arg(),
dtype,
vector_size,
matrix_len,
);
let output = HostData::from_tensor_handle(&client, output_handle, HostDataType::F32);
#[rustfmt::skip]
let expected_values = [
0.000, 1.000, 4.000, 5.000,
2.000, 3.000, 6.000, 7.000,
8.000, 9.000, 12.000, 13.000,
10.00, 11.000, 14.000, 15.000,
].to_vec();
let (_, expected_values) = TestInput::builder(client, shape)
.custom(expected_values)
.generate_with_f32_host_data();
assert_equals_approx(&output, &expected_values, 1e-6)
.as_test_outcome()
.enforce()
}
#[derive(RudaType, Clone, Copy)]
pub struct TiledLayout {
width: usize,
height: usize,
tile_w: usize,
tile_h: usize,
}
#[ruda]
impl TiledLayout {
pub fn new(width: usize, height: usize, tile_w: usize, tile_h: usize) -> TiledLayout {
TiledLayout {
width,
height,
tile_w,
tile_h,
}
}
}
#[ruda]
impl Layout for TiledLayout {
type Coordinates = (usize, usize);
type SourceCoordinates = Coords1d;
fn to_source_pos(&self, pos: Self::Coordinates) -> Self::SourceCoordinates {
let row = pos.0;
let col = pos.1;
let tile_row = row / self.tile_h;
let tile_col = col / self.tile_w;
let local_row = row % self.tile_h;
let local_col = col % self.tile_w;
let tiles_per_row = self.width / self.tile_w;
let tile_index = tile_row * tiles_per_row + tile_col;
let tile_area = self.tile_h * self.tile_w;
let block_offset = tile_index * tile_area;
let local_offset = local_row * self.tile_w + local_col;
block_offset + local_offset
}
fn to_source_pos_checked(&self, pos: Self::Coordinates) -> (Self::SourceCoordinates, bool) {
let is_valid = pos.0 < self.height && pos.1 < self.width;
(self.to_source_pos(pos), is_valid)
}
fn shape(&self) -> Self::Coordinates {
(self.width, self.height)
}
fn is_in_bounds(&self, _pos: Self::Coordinates) -> bool {
true.runtime()
}
}
#[derive(RudaType, Clone, Copy)]
pub struct RowMajorLayout {
width: usize,
height: usize,
vector_size: usize,
}
#[ruda]
impl RowMajorLayout {
pub fn new(width: usize, height: usize, vector_size: usize) -> Self {
RowMajorLayout {
width,
height,
vector_size,
}
}
}
#[ruda]
impl Layout for RowMajorLayout {
type Coordinates = (usize, usize);
type SourceCoordinates = Coords1d;
fn to_source_pos(&self, pos: Self::Coordinates) -> Self::SourceCoordinates {
(self.width * pos.0 + pos.1) / self.vector_size
}
fn to_source_pos_checked(&self, pos: Self::Coordinates) -> (Self::SourceCoordinates, bool) {
let is_valid = pos.0 < self.height && pos.1 < self.width;
(self.to_source_pos(pos), is_valid)
}
fn shape(&self) -> Self::Coordinates {
(self.width, self.height)
}
fn is_in_bounds(&self, _pos: Self::Coordinates) -> bool {
true.runtime()
}
}
#[ruda(launch)]
fn launch_read_rowmajor_tensor_as_tiled<N: Numeric, S: Size>(
input: &Tensor<Vector<N, S>>,
output: &mut Tensor<Vector<N, S>>,
#[define(N)] _dtype: StorageType,
#[define(S)] vector_size: usize,
#[comptime] matrix_len: usize,
) {
let input_view = input.view(TiledLayout::new(
matrix_len,
matrix_len,
matrix_len / 2,
matrix_len / 2,
));
let output_view = output.view_mut(RowMajorLayout::new(matrix_len, matrix_len, vector_size));
for i in 0..matrix_len {
for j in 0..matrix_len {
let value = input_view.read((i, j));
output_view.write((i, j), value);
}
}
}
#[test]
fn builder_overrides_stride_and_dtype() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let from_builder = TestInput::builder(client.clone(), shape![2, 3])
.stride(StrideSpec::ColMajor)
.arange()
.f32_host_data();
let from_new = TestInput::new(
client.clone(),
shape![2, 3],
f32::as_type_native_unchecked().storage_type(),
StrideSpec::ColMajor,
DataKind::Arange { scale: None },
)
.f32_host_data();
assert_equals_approx(&from_builder, &from_new, 0.0)
.as_test_outcome()
.enforce();
}
#[test]
fn builder_linspace_produces_evenly_spaced_values() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let actual = TestInput::builder(client.clone(), shape![1, 5])
.linspace(0.0, 1.0)
.f32_host_data();
let expected = TestInput::builder(client.clone(), shape![1, 5])
.custom(vec![0.0, 0.25, 0.5, 0.75, 1.0])
.f32_host_data();
assert_equals_approx(&actual, &expected, 1e-6)
.as_test_outcome()
.enforce();
}
#[test]
fn builder_normal_distribution_within_statistical_bounds() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let n: usize = 4096;
let host = TestInput::builder(client.clone(), shape![1, n])
.normal(11, 0.0, 1.0)
.f32_host_data();
let values: Vec<f32> = (0..n).map(|i| host.get_f32(&[0, i])).collect();
let mean = values.iter().copied().sum::<f32>() / n as f32;
let var = values.iter().map(|v| (v - mean).powi(2)).sum::<f32>() / n as f32;
let std = var.sqrt();
assert!(mean.abs() < 0.1, "mean drifted: {mean}");
assert!((std - 1.0).abs() < 0.1, "std drifted: {std}");
}
#[test]
fn host_data_typed_accessors_and_iter() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let host = TestInput::builder(client.clone(), shape![2, 3])
.arange()
.f32_host_data();
assert_eq!(host.get_f32(&[1, 2]), 5.0);
assert_eq!(host.try_get_f32(&[1, 2]), Some(5.0));
assert_eq!(host.try_get_i32(&[1, 2]), None);
assert_eq!(host.try_get_bool(&[1, 2]), None);
let collected: Vec<(Vec<usize>, f32)> = host.iter_indexed_f32().collect();
assert_eq!(collected.len(), 6);
assert_eq!(collected[0], (vec![0, 0], 0.0));
assert_eq!(collected[3], (vec![1, 0], 3.0));
assert_eq!(collected.last().unwrap(), &(vec![1, 2], 5.0));
}
#[test]
fn host_data_iter_respects_strides() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let actual = TestInput::builder(client.clone(), shape![2, 3])
.stride(StrideSpec::ColMajor)
.arange()
.f32_host_data();
let collected: Vec<f32> = actual.iter_indexed_f32().map(|(_, v)| v).collect();
assert_eq!(collected, vec![0.0, 1.0, 2.0, 3.0, 4.0, 5.0]);
}
#[test]
fn playground_partial_mismatch() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let eps = 0.001f32;
let expected = TestInput::builder(client.clone(), shape![4, 4])
.custom(vec![
0.10, 0.20, 0.30, 0.40, 1.00, 1.10, 1.20, 1.30, 2.00, 2.10, 2.20, 2.30, 3.00, 3.10,
3.20, 3.30,
])
.f32_host_data();
let actual = TestInput::builder(client.clone(), shape![4, 4])
.custom(vec![
0.10,
0.20,
0.30,
0.40,
1.0001,
1.0999,
1.2001,
1.2999,
2.01,
2.10,
2.21,
2.31,
3.50,
3.60,
3.70,
f32::NAN,
])
.f32_host_data();
let result = assert_equals_approx(&actual, &expected, eps);
assert!(
matches!(result, ValidationResult::Fail(_)),
"expected partial mismatch to be flagged as Fail, got {result:?}"
);
}
#[test]
fn print_tensors_skips_shape_mismatch() {
use ruda_test_utils::print_tensors;
let client = <TestRuntime as Runtime>::client(&Default::default());
let a = TestInput::builder(client.clone(), shape![2, 3])
.arange()
.f32_host_data();
let b = TestInput::builder(client.clone(), shape![3, 2])
.arange()
.f32_host_data();
print_tensors("mismatched", &[&a, &b], Some(0.001));
}
#[test]
fn print_tensors_skips_rank_mismatch() {
use ruda_test_utils::print_tensors;
let client = <TestRuntime as Runtime>::client(&Default::default());
let r2 = TestInput::builder(client.clone(), shape![2, 3])
.arange()
.f32_host_data();
let r3 = TestInput::builder(client.clone(), shape![2, 2, 3])
.arange()
.f32_host_data();
print_tensors("mismatched_rank", &[&r2, &r3], Some(0.001));
}
#[test]
fn print_tensor_is_no_op_in_correct_mode() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let r2 = TestInput::builder(client.clone(), shape![2, 3])
.arange()
.f32_host_data();
print_tensor("rank-2 arange", &r2);
let r3 = TestInput::builder(client.clone(), shape![2, 2, 3])
.arange()
.f32_host_data();
print_tensor("rank-3 arange", &r3);
}
#[test]
fn pretty_print_handles_rank_3() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let host = TestInput::builder(client.clone(), shape![2, 2, 3])
.arange()
.f32_host_data();
let printed = host.pretty_print();
assert!(
printed.contains("[0, *, *]"),
"missing first slice label:\n{printed}"
);
assert!(
printed.contains("[1, *, *]"),
"missing second slice label:\n{printed}"
);
assert!(
printed.contains("0.000"),
"expected 0.000 in printed output:\n{printed}"
);
assert!(
printed.contains("11.000"),
"expected 11.000 in printed output:\n{printed}"
);
}
#[test]
fn pretty_print_slice_filters_rows_and_cols() {
use ruda_test_utils::DimFilter;
let client = <TestRuntime as Runtime>::client(&Default::default());
let host = TestInput::builder(client.clone(), shape![2, 4])
.arange()
.f32_host_data();
let printed = host.pretty_print_slice(vec![
DimFilter::Exact(1),
DimFilter::Range { start: 1, end: 2 },
]);
assert!(printed.contains("5.000"), "expected 5.000 in: {printed}");
assert!(printed.contains("6.000"), "expected 6.000 in: {printed}");
assert!(
!printed.contains("4.000"),
"should not contain 4.000: {printed}"
);
assert!(
!printed.contains("7.000"),
"should not contain 7.000: {printed}"
);
assert!(
!printed.contains("0.000"),
"should not contain 0.000: {printed}"
);
}
#[test]
fn pretty_print_slice_filters_leading_dims() {
use ruda_test_utils::DimFilter;
let client = <TestRuntime as Runtime>::client(&Default::default());
let host = TestInput::builder(client.clone(), shape![3, 2, 2])
.arange()
.f32_host_data();
let filter = vec![DimFilter::Exact(1), DimFilter::Any, DimFilter::Any];
let printed = host.pretty_print_slice(filter);
assert!(
printed.contains("[1, *, *]"),
"expected slice [1, *, *]:\n{printed}"
);
assert!(
!printed.contains("[0, *, *]"),
"should not include slice [0, *, *]:\n{printed}"
);
assert!(
!printed.contains("[2, *, *]"),
"should not include slice [2, *, *]:\n{printed}"
);
}
#[test]
fn builder_uniform_values_in_range() {
let client = <TestRuntime as Runtime>::client(&Default::default());
let host = TestInput::builder(client.clone(), shape![4, 4])
.uniform(7, -1.0, 1.0)
.f32_host_data();
for i in 0..4 {
for j in 0..4 {
let v = host.get_f32(&[i, j]);
assert!(
(-1.0..=1.0).contains(&v),
"uniform value out of range at [{i},{j}]: {v}"
);
}
}
}