use crate::bounds::{bounds_from_bop, Bounds, TileBinaryOp};
use crate::compiler::_value::TileRustValue;
use cuda_async::predicate::{Atom, Term};
pub(crate) struct ScalarFacts {
pub(crate) bounds: Option<Bounds<i64>>,
pub(crate) term: Option<Term>,
pub(crate) floor_div: Option<FloorDiv>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) struct FloorDiv {
pub(crate) numerator: Term,
pub(crate) divisor: i64,
}
pub(crate) fn int_value_domain(elem_ty: &str) -> Option<Bounds<i64>> {
match elem_ty {
"bool" | "i1" => Some(Bounds::new(0, 1)),
"i32" => Some(Bounds::new(i32::MIN as i64, i32::MAX as i64)),
"u32" => Some(Bounds::new(0, u32::MAX as i64)),
_ => None,
}
}
pub(crate) fn transfer(
op: &TileBinaryOp,
lhs: &TileRustValue,
rhs: &TileRustValue,
result_domain: Option<Bounds<i64>>,
) -> ScalarFacts {
let bounds = propagate_bounds(op, lhs, rhs).filter(|b| {
result_domain.is_some_and(|domain| domain.start <= b.start && b.end <= domain.end)
});
ScalarFacts {
bounds,
term: propagate_term(op, lhs, rhs),
floor_div: propagate_floor_div(op, lhs, rhs),
}
}
pub(crate) fn propagate_floor_div(
op: &TileBinaryOp,
lhs: &TileRustValue,
rhs: &TileRustValue,
) -> Option<FloorDiv> {
if !matches!(op, TileBinaryOp::Div) {
return None;
}
let divisor = rhs.bounds.filter(|b| b.is_exact()).map(|b| b.start)?;
if divisor <= 0 {
return None;
}
Some(FloorDiv {
numerator: lhs.term.clone()?,
divisor,
})
}
pub(crate) fn propagate_bounds(
op: &TileBinaryOp,
lhs: &TileRustValue,
rhs: &TileRustValue,
) -> Option<Bounds<i64>> {
match (lhs.bounds, rhs.bounds) {
(Some(a), Some(b)) => bounds_from_bop(op, &a, &b),
_ => None,
}
}
pub(crate) fn propagate_term(
op: &TileBinaryOp,
lhs: &TileRustValue,
rhs: &TileRustValue,
) -> Option<Term> {
let term_of = |v: &TileRustValue| -> Option<Term> {
v.term.clone().or_else(|| {
v.bounds
.filter(|b| b.is_exact())
.map(|b| Term::constant(b.start))
})
};
let lt = term_of(lhs)?;
let rt = term_of(rhs)?;
match op {
TileBinaryOp::Add => lt.add(&rt),
TileBinaryOp::Sub => lt.sub(&rt),
TileBinaryOp::Mul => {
if let Some(c) = rt.as_constant() {
lt.mul_const(c)
} else if let Some(c) = lt.as_constant() {
rt.mul_const(c)
} else {
None
}
}
_ => None,
}
}
pub(crate) fn term_range(
term: &Term,
atom_range: &impl Fn(&Atom) -> Option<Bounds<i64>>,
) -> Option<Bounds<i64>> {
let mut lo = term.constant_part();
let mut hi = term.constant_part();
for (atom, &coeff) in term.coeffs() {
let r = atom_range(atom)?;
let (a, b) = (coeff.checked_mul(r.start)?, coeff.checked_mul(r.end)?);
let (add_lo, add_hi) = if a <= b { (a, b) } else { (b, a) };
lo = lo.checked_add(add_lo)?;
hi = hi.checked_add(add_hi)?;
}
Some(Bounds { start: lo, end: hi })
}
#[cfg(test)]
mod tests {
use super::*;
fn dim(param: usize, axis: usize) -> Atom {
Atom::Dim { param, axis }
}
#[test]
fn term_range_of_a_constant_is_the_exact_range() {
let range = term_range(&Term::constant(7), &|_| None).unwrap();
assert_eq!(range, Bounds { start: 7, end: 7 });
assert!(range.is_exact());
}
#[test]
fn term_range_of_affine_uses_atom_ranges() {
let term = Term::affine(dim(0, 0), 2, 3);
let env = |_: &Atom| Some(Bounds { start: 0, end: 4 });
assert_eq!(term_range(&term, &env), Some(Bounds { start: 3, end: 11 }));
}
#[test]
fn term_range_negative_coefficient_swaps_endpoints() {
let term = Term::affine(dim(0, 0), -1, 10);
let env = |_: &Atom| Some(Bounds { start: 0, end: 4 });
assert_eq!(term_range(&term, &env), Some(Bounds { start: 6, end: 10 }));
}
#[test]
fn term_range_is_none_when_an_atom_has_no_range() {
let term = Term::atom(dim(9, 9));
assert_eq!(term_range(&term, &|_| None), None);
}
#[test]
fn runtime_domains_cover_boolean_and_32_bit_integer_results() {
assert_eq!(int_value_domain("bool"), Some(Bounds::new(0, 1)));
assert_eq!(
int_value_domain("i32"),
Some(Bounds::new(i32::MIN as i64, i32::MAX as i64))
);
assert_eq!(
int_value_domain("u32"),
Some(Bounds::new(0, u32::MAX as i64))
);
assert_eq!(int_value_domain("i64"), None);
assert_eq!(int_value_domain("usize"), None);
}
}