#[derive(Clone, Copy, PartialEq, Eq, Hash, Debug)]
pub struct SymbolId(pub u32);
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum Dim {
Static(usize),
Symbolic(SymbolId),
}
impl Dim {
pub fn as_static(self) -> Option<usize> {
match self {
Dim::Static(n) => Some(n),
Dim::Symbolic(_) => None,
}
}
pub fn is_static(self) -> bool {
matches!(self, Dim::Static(_))
}
}
impl From<usize> for Dim {
fn from(n: usize) -> Self {
Dim::Static(n)
}
}
impl From<SymbolId> for Dim {
fn from(s: SymbolId) -> Self {
Dim::Symbolic(s)
}
}
pub type Shape = Vec<Dim>;
pub fn static_shape(dims: impl IntoIterator<Item = usize>) -> Shape {
dims.into_iter().map(Dim::Static).collect()
}
pub fn as_static_shape(shape: &[Dim]) -> Option<Vec<usize>> {
shape.iter().map(|d| d.as_static()).collect()
}
pub fn is_fully_static(shape: &[Dim]) -> bool {
shape.iter().all(|d| d.is_static())
}
#[derive(Clone, Debug, PartialEq, Eq, Default)]
pub struct SymbolConstraints {
pub id: Option<SymbolId>,
pub name: Option<String>,
pub min: Option<usize>,
pub max: Option<usize>,
pub divisible_by: Option<usize>,
}
impl SymbolConstraints {
pub fn new(id: SymbolId, name: Option<String>) -> Self {
Self {
id: Some(id),
name,
..Default::default()
}
}
pub fn accepts(&self, value: usize) -> bool {
self.min.is_none_or(|lo| value >= lo)
&& self.max.is_none_or(|hi| value <= hi)
&& self
.divisible_by
.is_none_or(|m| m != 0 && value.is_multiple_of(m))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn static_shape_helpers() {
let s = static_shape([2, 3, 4]);
assert_eq!(s.len(), 3);
assert!(is_fully_static(&s));
assert_eq!(as_static_shape(&s), Some(vec![2, 3, 4]));
}
#[test]
fn symbolic_shape_is_not_static() {
let s = vec![Dim::Symbolic(SymbolId(0)), Dim::Static(768)];
assert!(!is_fully_static(&s));
assert_eq!(as_static_shape(&s), None);
assert_eq!(s[1].as_static(), Some(768));
}
#[test]
fn dim_conversions() {
assert_eq!(Dim::from(5usize), Dim::Static(5));
assert_eq!(Dim::from(SymbolId(7)), Dim::Symbolic(SymbolId(7)));
}
#[test]
fn constraints_accept() {
let c = SymbolConstraints {
id: Some(SymbolId(0)),
name: Some("seq".into()),
min: Some(1),
max: Some(2048),
divisible_by: Some(8),
};
assert!(c.accepts(8));
assert!(c.accepts(2048));
assert!(!c.accepts(0)); assert!(!c.accepts(4096)); assert!(!c.accepts(12)); }
}