use crate::{Error, Result};
const DOT_GENERAL_OP: &str = "dot_general";
fn invalid_dot_general_config(message: impl Into<String>) -> Error {
Error::InvalidConfig {
op: DOT_GENERAL_OP,
message: message.into(),
}
}
#[derive(Clone, Debug, Hash, PartialEq, Eq)]
pub struct DotGeneralConfig {
pub lhs_contracting_dims: Vec<usize>,
pub rhs_contracting_dims: Vec<usize>,
pub lhs_batch_dims: Vec<usize>,
pub rhs_batch_dims: Vec<usize>,
}
impl DotGeneralConfig {
fn check_no_duplicates(dims: &[usize], label: &str) -> Result<()> {
let mut seen = std::collections::HashSet::new();
for &d in dims {
if !seen.insert(d) {
return Err(invalid_dot_general_config(format!(
"{} contains duplicate dim {}",
label, d
)));
}
}
Ok(())
}
pub fn validate_dims_with_ranks(&self, lhs_rank: usize, rhs_rank: usize) -> Result<()> {
for &d in &self.lhs_contracting_dims {
if d >= lhs_rank {
return Err(invalid_dot_general_config(format!(
"lhs_contracting_dim {} out of bounds for lhs_rank {}",
d, lhs_rank
)));
}
}
for &d in &self.rhs_contracting_dims {
if d >= rhs_rank {
return Err(invalid_dot_general_config(format!(
"rhs_contracting_dim {} out of bounds for rhs_rank {}",
d, rhs_rank
)));
}
}
for &d in &self.lhs_batch_dims {
if d >= lhs_rank {
return Err(invalid_dot_general_config(format!(
"lhs_batch_dim {} out of bounds for lhs_rank {}",
d, lhs_rank
)));
}
}
for &d in &self.rhs_batch_dims {
if d >= rhs_rank {
return Err(invalid_dot_general_config(format!(
"rhs_batch_dim {} out of bounds for rhs_rank {}",
d, rhs_rank
)));
}
}
Self::check_no_duplicates(&self.lhs_contracting_dims, "lhs_contracting_dims")?;
Self::check_no_duplicates(&self.rhs_contracting_dims, "rhs_contracting_dims")?;
Self::check_no_duplicates(&self.lhs_batch_dims, "lhs_batch_dims")?;
Self::check_no_duplicates(&self.rhs_batch_dims, "rhs_batch_dims")?;
for &d in &self.lhs_contracting_dims {
if self.lhs_batch_dims.contains(&d) {
return Err(invalid_dot_general_config(format!(
"lhs dim {} appears in both contracting and batch dims",
d
)));
}
}
for &d in &self.rhs_contracting_dims {
if self.rhs_batch_dims.contains(&d) {
return Err(invalid_dot_general_config(format!(
"rhs dim {} appears in both contracting and batch dims",
d
)));
}
}
if self.lhs_contracting_dims.len() != self.rhs_contracting_dims.len() {
return Err(invalid_dot_general_config(format!(
"lhs/rhs contracting dim counts differ ({} vs {})",
self.lhs_contracting_dims.len(),
self.rhs_contracting_dims.len()
)));
}
if self.lhs_batch_dims.len() != self.rhs_batch_dims.len() {
return Err(invalid_dot_general_config(format!(
"lhs/rhs batch dim counts differ ({} vs {})",
self.lhs_batch_dims.len(),
self.rhs_batch_dims.len()
)));
}
Ok(())
}
}
#[derive(Clone, Debug, Hash, PartialEq, Eq)]
pub enum CompareDir {
Eq,
Lt,
Le,
Gt,
Ge,
}
#[derive(Clone, Debug, Hash, PartialEq, Eq)]
pub struct GatherConfig {
pub offset_dims: Vec<usize>,
pub collapsed_slice_dims: Vec<usize>,
pub start_index_map: Vec<usize>,
pub index_vector_dim: usize,
pub slice_sizes: Vec<usize>,
}
#[derive(Clone, Debug, Hash, PartialEq, Eq)]
pub struct ScatterConfig {
pub update_window_dims: Vec<usize>,
pub inserted_window_dims: Vec<usize>,
pub scatter_dims_to_operand_dims: Vec<usize>,
pub index_vector_dim: usize,
}
#[derive(Clone, Debug, Hash, PartialEq, Eq)]
pub struct SliceConfig {
pub starts: Vec<usize>,
pub limits: Vec<usize>,
pub strides: Vec<usize>,
}
#[derive(Clone, Debug, Hash, PartialEq, Eq)]
pub struct PadConfig {
pub edge_padding_low: Vec<i64>,
pub edge_padding_high: Vec<i64>,
pub interior_padding: Vec<i64>,
}