#[inline]
pub(crate) fn validate_gemm_nn(
a_len: usize,
b_len: usize,
c_len: usize,
m: usize,
k: usize,
n: usize,
op: &'static str,
) {
assert!(m.checked_mul(k).is_some(), "{op}: shape overflow: m*k");
assert!(k.checked_mul(n).is_some(), "{op}: shape overflow: k*n");
assert!(m.checked_mul(n).is_some(), "{op}: shape overflow: m*n");
assert!(a_len >= m * k, "{op}: a too short for m*k");
assert!(b_len >= k * n, "{op}: b too short for k*n");
assert!(c_len >= m * n, "{op}: c too short for m*n");
}
#[inline]
pub(crate) fn validate_gemm_bt(
a_len: usize,
b_len: usize,
c_len: usize,
m: usize,
k: usize,
n: usize,
op: &'static str,
) {
assert!(m.checked_mul(k).is_some(), "{op}: shape overflow: m*k");
assert!(n.checked_mul(k).is_some(), "{op}: shape overflow: n*k");
assert!(m.checked_mul(n).is_some(), "{op}: shape overflow: m*n");
assert!(a_len >= m * k, "{op}: a too short for m*k");
assert!(b_len >= n * k, "{op}: b too short for n*k");
assert!(c_len >= m * n, "{op}: c too short for m*n");
}
#[inline]
pub(crate) fn validate_gemm_strided_shape(
m: usize,
k: usize,
n: usize,
lda: usize,
ldb: usize,
ldc: usize,
transposed_b: bool,
op: &'static str,
) {
assert!(m.checked_mul(k).is_some(), "{op}: shape overflow: m*k");
assert!(n.checked_mul(k).is_some(), "{op}: shape overflow: n*k");
assert!(m.checked_mul(n).is_some(), "{op}: shape overflow: m*n");
assert!(lda >= k, "{op}: lda too small for row extent k");
if transposed_b {
assert!(ldb >= k, "{op}: ldb too small for row extent k");
} else {
assert!(ldb >= n, "{op}: ldb too small for row extent n");
}
assert!(ldc >= n, "{op}: ldc too small for row extent n");
}
#[inline]
pub(crate) fn validate_ternary_matvec_args(
x_q_len: usize,
alphas_len: usize,
packed_w_len: usize,
output_len: usize,
n: usize,
k: usize,
packed_row_bytes: usize,
op: &'static str,
) {
assert!(
n.checked_mul(packed_row_bytes).is_some(),
"{op}: shape overflow: n*packed_row_bytes"
);
assert!(x_q_len >= k, "{op}: x_q too short for k");
assert!(alphas_len >= n, "{op}: alphas too short for n");
assert!(
packed_w_len >= n * packed_row_bytes,
"{op}: packed_w too short for n*packed_row_bytes"
);
assert!(output_len >= n, "{op}: output too short for n");
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn validate_gemm_nn_accepts_oversized_buffers() {
validate_gemm_nn(3, 5, 4, 1, 2, 2, "test");
}
#[test]
#[should_panic(expected = "a too short for m*k")]
fn validate_gemm_nn_rejects_short_a() {
validate_gemm_nn(1, 4, 2, 1, 2, 2, "test");
}
#[test]
#[should_panic(expected = "b too short for k*n")]
fn validate_gemm_nn_rejects_short_b() {
validate_gemm_nn(2, 3, 2, 1, 2, 2, "test");
}
#[test]
#[should_panic(expected = "c too short for m*n")]
fn validate_gemm_nn_rejects_short_c() {
validate_gemm_nn(2, 4, 1, 1, 2, 2, "test");
}
#[test]
#[should_panic(expected = "shape overflow: m*k")]
fn validate_gemm_nn_rejects_overflow() {
validate_gemm_nn(2, 2, 2, 2, usize::MAX, 2, "test");
}
#[test]
fn validate_gemm_bt_accepts_oversized_buffers() {
validate_gemm_bt(3, 5, 4, 1, 2, 2, "test");
}
#[test]
#[should_panic(expected = "b too short for n*k")]
fn validate_gemm_bt_rejects_short_b() {
validate_gemm_bt(2, 3, 2, 1, 2, 2, "test");
}
#[test]
#[should_panic(expected = "shape overflow: n*k")]
fn validate_gemm_bt_rejects_overflow() {
validate_gemm_bt(2, 2, 2, 2, 2, usize::MAX, "test");
}
#[test]
fn validate_gemm_strided_shape_accepts_exact() {
validate_gemm_strided_shape(1, 2, 2, 2, 2, 2, false, "test");
validate_gemm_strided_shape(1, 2, 2, 2, 2, 2, true, "test");
}
#[test]
#[should_panic(expected = "ldb too small for row extent n")]
fn validate_gemm_strided_shape_rejects_short_ldb_nn() {
validate_gemm_strided_shape(1, 2, 4, 2, 2, 4, false, "test");
}
#[test]
#[should_panic(expected = "ldb too small for row extent k")]
fn validate_gemm_strided_shape_rejects_short_ldb_bt() {
validate_gemm_strided_shape(1, 4, 2, 4, 2, 2, true, "test");
}
#[test]
#[should_panic(expected = "shape overflow: n*k")]
fn validate_gemm_strided_shape_rejects_overflow() {
validate_gemm_strided_shape(2, 2, usize::MAX, 2, 2, 2, true, "test");
}
#[test]
fn validate_ternary_matvec_args_accepts_oversized() {
validate_ternary_matvec_args(5, 4, 20, 4, 3, 4, 1, "test");
}
#[test]
#[should_panic(expected = "x_q too short for k")]
fn validate_ternary_matvec_args_rejects_short_x_q() {
validate_ternary_matvec_args(3, 4, 20, 4, 3, 4, 1, "test");
}
#[test]
#[should_panic(expected = "alphas too short for n")]
fn validate_ternary_matvec_args_rejects_short_alphas() {
validate_ternary_matvec_args(5, 2, 20, 4, 3, 4, 1, "test");
}
#[test]
#[should_panic(expected = "packed_w too short for n*packed_row_bytes")]
fn validate_ternary_matvec_args_rejects_short_packed_w() {
validate_ternary_matvec_args(5, 4, 2, 4, 3, 4, 1, "test");
}
#[test]
#[should_panic(expected = "output too short for n")]
fn validate_ternary_matvec_args_rejects_short_output() {
validate_ternary_matvec_args(5, 4, 20, 2, 3, 4, 1, "test");
}
#[test]
#[should_panic(expected = "shape overflow: n*packed_row_bytes")]
fn validate_ternary_matvec_args_rejects_overflow() {
validate_ternary_matvec_args(5, 4, 20, 4, usize::MAX, 4, usize::MAX, "test");
}
}