Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
107 changes: 99 additions & 8 deletions cutile-rs/cutile/src/api.rs
Original file line number Diff line number Diff line change
Expand Up @@ -530,6 +530,12 @@ pub fn ones<T: DType>(shape: &[usize]) -> impl DeviceOp<Output = Tensor<T>> {
/// Allocates GPU memory and fills it with the specified value. This uses a GPU kernel
/// to initialize the memory efficiently.
///
/// ## Errors
///
/// Returns an error on execution if a dimension is zero, the element count
/// overflows `usize`, or the flattened element count exceeds `i32::MAX`.
/// Invalid shapes are rejected before GPU allocation or kernel execution.
///
/// ## Examples
///
/// ```rust,ignore
Expand All @@ -540,15 +546,40 @@ pub fn ones<T: DType>(shape: &[usize]) -> impl DeviceOp<Output = Tensor<T>> {
/// let matrix = api::full(-1, &[128, 128]).await;
/// ```
pub fn full<T: DType>(val: T, shape: &[usize]) -> impl DeviceOp<Output = Tensor<T>> {
let len = match checked_full_len(shape) {
Ok(len) => len,
Err(message) => return fail::<Tensor<T>>(message).boxed(),
};
let shape = shape.to_vec();
let len = shape.iter().product::<usize>();
Tensor::<T>::uninitialized(len).then(move |t| {
// TODO (hme): It's awkward to assume_init this before actually initializing it.
let partition_size = 128;
let result = unsafe { t.assume_init() }.partition([partition_size]);
let (_, res) = value((val, result)).then(full_apply).unzip();
res.unpartition().reshape(&shape)
})
Tensor::<T>::uninitialized(len)
.then(move |t| {
// TODO (hme): It's awkward to assume_init this before actually initializing it.
let partition_size = 128;
let result = unsafe { t.assume_init() }.partition([partition_size]);
let (_, res) = value((val, result)).then(full_apply).unzip();
res.unpartition().reshape(&shape)
})
.boxed()
}

// `full` initializes a flat tensor before reshaping it. Its element count
// must fit that tensor's i32 dimension, which also bounds every positive
// target dimension and contiguous stride. An empty shape remains a scalar.
fn checked_full_len(shape: &[usize]) -> Result<usize, String> {
let len = shape
.iter()
.try_fold(1usize, |len, &dim| len.checked_mul(dim))
.ok_or_else(|| format!("full: shape {shape:?} element count overflows usize"))?;
if len == 0 {
return Err(format!("full: shape {shape:?} contains a zero dimension"));
}
if i32::try_from(len).is_err() {
return Err(format!(
"full: shape {shape:?} element count {len} exceeds i32::MAX ({})",
i32::MAX
));
}
Ok(len)
}

pub fn fill<T: DType>(tensor: Tensor<T>, val: T) -> impl DeviceOp<Output = Tensor<T>> {
Expand Down Expand Up @@ -963,3 +994,63 @@ pub trait DeviceOpReshapeShared<T: DType + Send>:
}

impl<T: DType + Send, DI: DeviceOp<Output = Arc<Tensor<T>>>> DeviceOpReshapeShared<T> for DI {}

#[cfg(test)]
mod full_shape_tests {
use super::{checked_full_len, full, ones, zeros};

#[test]
fn accepts_valid_shapes_and_i32_boundary() {
assert_eq!(checked_full_len(&[]), Ok(1));
assert_eq!(checked_full_len(&[1024]), Ok(1024));
assert_eq!(checked_full_len(&[3, 5]), Ok(15));
assert_eq!(checked_full_len(&[2, 3, 4]), Ok(24));
assert_eq!(
checked_full_len(&[i32::MAX as usize]),
Ok(i32::MAX as usize)
);
}

#[test]
fn rejects_zero_dimensions() {
for shape in [&[0][..], &[3, 0][..], &[0, usize::MAX][..]] {
let error = checked_full_len(shape).expect_err("zero dimension must fail");
assert!(error.contains("zero dimension"), "{error}");
}
}

#[test]
fn rejects_usize_product_overflow() {
// The first product wraps to 1 without overflow checks. The second
// overflows even though every individual dimension fits in i32.
for shape in [&[usize::MAX, usize::MAX][..], &[i32::MAX as usize; 3][..]] {
let error = checked_full_len(shape).expect_err("shape product must fail");
assert!(error.contains("overflows usize"), "{error}");
}
}

#[test]
fn rejects_flattened_length_above_i32_max() {
// Both products fit usize on 32-bit and 64-bit hosts.
for shape in [&[i32::MAX as usize + 1][..], &[32768, 65536][..]] {
let error = checked_full_len(shape).expect_err("flat dimension must fail");
assert!(error.contains("exceeds i32::MAX"), "{error}");
}
}

#[test]
fn invalid_shapes_do_not_panic_during_construction() {
// Only construct the lazy operations: this must not need a CUDA
// context, panic on arithmetic, or reach the zero-length assertion.
for shape in [
&[usize::MAX, usize::MAX][..],
&[i32::MAX as usize + 1][..],
&[32768, 65536][..],
&[0][..],
] {
let _ = full(7.0f32, shape);
let _ = zeros::<f32>(shape);
let _ = ones::<f32>(shape);
}
}
}
67 changes: 67 additions & 0 deletions cutile-rs/cutile/tests/gpu/tensor_guards.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@

use std::panic::{self, AssertUnwindSafe};

use cuda_async::error::DeviceError;
use cutile::api;
use cutile::prelude::*;

Expand Down Expand Up @@ -104,3 +105,69 @@ fn oversized_allocation_is_an_error_not_a_panic() {
.expect("allocation after OOM");
assert_eq!(t.shape(), &[16]);
}

const INVALID_FULL_SHAPES: &[(&[usize], &str)] = &[
(&[usize::MAX, usize::MAX], "overflows usize"),
(&[i32::MAX as usize; 3], "overflows usize"),
(&[i32::MAX as usize + 1], "exceeds i32::MAX"),
(&[32768, 65536], "exceeds i32::MAX"),
(&[0], "zero dimension"),
(&[3, 0], "zero dimension"),
];

fn assert_full_shape_error(error: DeviceError, expected: &str) {
match error {
DeviceError::Internal(message) => {
assert!(message.starts_with("full: "), "{message}");
assert!(message.contains(expected), "{message}");
}
other => panic!("expected full shape validation error, got: {other:?}"),
}
}

#[test]
fn full_reports_invalid_shapes_as_errors() {
// Check the specific error, not just is_err(): in release builds the
// old product could wrap, allocate and fill, then fail during reshape.
for &(shape, expected) in INVALID_FULL_SHAPES {
let error = api::full(7.0f32, shape)
.sync()
.expect_err("invalid full shape must fail");
assert_full_shape_error(error, expected);
}
}

#[test]
fn full_wrappers_report_invalid_shapes_as_errors() {
for &(shape, expected) in INVALID_FULL_SHAPES {
let zeros_error = api::zeros::<f32>(shape)
.sync()
.expect_err("invalid zeros shape must fail");
assert_full_shape_error(zeros_error, expected);

let ones_error = api::ones::<f32>(shape)
.sync()
.expect_err("invalid ones shape must fail");
assert_full_shape_error(ones_error, expected);
}
}

#[test]
fn full_and_wrappers_preserve_shape_and_values() {
// Exercise a matrix reshape and a partial edge tile of the fill kernel.
for (shape, expected_shape) in [
(&[2usize, 3][..], &[2i32, 3][..]),
(&[129usize][..], &[129i32][..]),
] {
let tensors = [
(api::full(3.25f32, shape).sync().expect("valid full"), 3.25),
(api::zeros::<f32>(shape).sync().expect("valid zeros"), 0.0),
(api::ones::<f32>(shape).sync().expect("valid ones"), 1.0),
];
for (tensor, expected) in tensors {
assert_eq!(tensor.shape(), expected_shape);
let host: Vec<f32> = tensor.to_host_vec().sync().expect("copy");
assert_eq!(host, vec![expected; shape.iter().product()]);
}
}
}