Skip to main content

tenferro_tensor/
config.rs

1use smallvec::SmallVec;
2
3use crate::{Error, Result, ValidationError};
4
5const DOT_GENERAL_OP: &str = "dot_general";
6
7fn invalid_dot_general_config(message: impl Into<String>) -> Error {
8    Error::invalid_argument(DOT_GENERAL_OP, "dot_general_config", message)
9}
10
11/// DotGeneral dimension configuration.
12///
13/// Each axis list stores up to four axes inline and spills to the heap for
14/// larger contractions; this is not a tensor-rank limit.
15/// Records only the dim-numbering roles (contracting / batch; free is derived).
16/// Rank info travels with the enclosing `StdTensorOp::DotGeneral` variant at
17/// the trace/StdTensorOp layer, and with `ExecInstruction::output_shapes` at
18/// the exec layer. This separation makes it structurally impossible for
19/// stored ranks to drift from actual tensor ranks (issue #664).
20///
21/// # Output layout
22///
23/// The output shape is `[lhs_free..., rhs_free..., batch...]` (col-major
24/// batch-trailing convention): batch axes come **last**, unlike PyTorch's
25/// batch-leading `bmm`. Batch dims have the largest stride so that each batch
26/// slice occupies a contiguous block of memory. Free axes keep their input
27/// order. For attention scores, `q[d, lq, b]` against `k[d, lk, b]` with
28/// `d` contracted and `b` batched gives `[lq, lk, b]`; see the
29/// "Contraction Output Layout" section of the tensor-operations guide for a
30/// worked example. Transpose afterwards only where a consumer needs batch
31/// first.
32///
33/// # Examples
34///
35/// ```rust
36/// use tenferro_tensor::DotGeneralConfig;
37///
38/// let config = DotGeneralConfig {
39///     lhs_contracting_dims: [1].as_slice().into(),
40///     rhs_contracting_dims: [0].as_slice().into(),
41///     lhs_batch_dims: [].as_slice().into(),
42///     rhs_batch_dims: [].as_slice().into(),
43/// };
44/// ```
45#[derive(Clone, Debug, Hash, PartialEq, Eq)]
46pub struct DotGeneralConfig {
47    pub lhs_contracting_dims: SmallVec<[usize; 4]>,
48    pub rhs_contracting_dims: SmallVec<[usize; 4]>,
49    pub lhs_batch_dims: SmallVec<[usize; 4]>,
50    pub rhs_batch_dims: SmallVec<[usize; 4]>,
51}
52
53impl DotGeneralConfig {
54    fn check_no_duplicates(dims: &[usize], label: &'static str) -> Result<()> {
55        let mut seen = std::collections::HashSet::new();
56        for &d in dims {
57            if !seen.insert(d) {
58                return Err(Error::validation(
59                    DOT_GENERAL_OP,
60                    ValidationError::DuplicateAxis {
61                        axis: d,
62                        role: label,
63                    },
64                ));
65            }
66        }
67        Ok(())
68    }
69
70    /// Validate that all dimension indices are within range for the given
71    /// explicit ranks and that no axis appears in multiple roles.
72    ///
73    /// Call sites supply the actual operand ranks (from the tensor shapes they
74    /// have in hand). The config itself carries only the dim-numbering roles.
75    ///
76    /// # Examples
77    ///
78    /// ```rust
79    /// use tenferro_tensor::DotGeneralConfig;
80    ///
81    /// let config = DotGeneralConfig {
82    ///     lhs_contracting_dims: [1].as_slice().into(),
83    ///     rhs_contracting_dims: [0].as_slice().into(),
84    ///     lhs_batch_dims: [].as_slice().into(),
85    ///     rhs_batch_dims: [].as_slice().into(),
86    /// };
87    /// config.validate_dims_with_ranks(2, 2).unwrap();
88    /// ```
89    /// # Errors
90    ///
91    /// Returns [`crate::Error::Validation`] with an axis, duplicate-axis,
92    /// or configuration source when the dimension roles are invalid.
93    pub fn validate_dims_with_ranks(&self, lhs_rank: usize, rhs_rank: usize) -> Result<()> {
94        for &d in &self.lhs_contracting_dims {
95            if d >= lhs_rank {
96                return Err(Error::validation(
97                    DOT_GENERAL_OP,
98                    ValidationError::AxisOutOfBounds {
99                        axis: d,
100                        rank: lhs_rank,
101                    },
102                ));
103            }
104        }
105        for &d in &self.rhs_contracting_dims {
106            if d >= rhs_rank {
107                return Err(Error::validation(
108                    DOT_GENERAL_OP,
109                    ValidationError::AxisOutOfBounds {
110                        axis: d,
111                        rank: rhs_rank,
112                    },
113                ));
114            }
115        }
116        for &d in &self.lhs_batch_dims {
117            if d >= lhs_rank {
118                return Err(Error::validation(
119                    DOT_GENERAL_OP,
120                    ValidationError::AxisOutOfBounds {
121                        axis: d,
122                        rank: lhs_rank,
123                    },
124                ));
125            }
126        }
127        for &d in &self.rhs_batch_dims {
128            if d >= rhs_rank {
129                return Err(Error::validation(
130                    DOT_GENERAL_OP,
131                    ValidationError::AxisOutOfBounds {
132                        axis: d,
133                        rank: rhs_rank,
134                    },
135                ));
136            }
137        }
138        Self::check_no_duplicates(&self.lhs_contracting_dims, "lhs_contracting_dims")?;
139        Self::check_no_duplicates(&self.rhs_contracting_dims, "rhs_contracting_dims")?;
140        Self::check_no_duplicates(&self.lhs_batch_dims, "lhs_batch_dims")?;
141        Self::check_no_duplicates(&self.rhs_batch_dims, "rhs_batch_dims")?;
142        for &d in &self.lhs_contracting_dims {
143            if self.lhs_batch_dims.contains(&d) {
144                return Err(Error::validation(
145                    DOT_GENERAL_OP,
146                    ValidationError::AxisRoleConflict {
147                        axis: d,
148                        first_role: "lhs contracting",
149                        second_role: "lhs batch",
150                    },
151                ));
152            }
153        }
154        for &d in &self.rhs_contracting_dims {
155            if self.rhs_batch_dims.contains(&d) {
156                return Err(Error::validation(
157                    DOT_GENERAL_OP,
158                    ValidationError::AxisRoleConflict {
159                        axis: d,
160                        first_role: "rhs contracting",
161                        second_role: "rhs batch",
162                    },
163                ));
164            }
165        }
166        if self.lhs_contracting_dims.len() != self.rhs_contracting_dims.len() {
167            return Err(invalid_dot_general_config(format!(
168                "lhs/rhs contracting dim counts differ ({} vs {})",
169                self.lhs_contracting_dims.len(),
170                self.rhs_contracting_dims.len()
171            )));
172        }
173        if self.lhs_batch_dims.len() != self.rhs_batch_dims.len() {
174            return Err(invalid_dot_general_config(format!(
175                "lhs/rhs batch dim counts differ ({} vs {})",
176                self.lhs_batch_dims.len(),
177                self.rhs_batch_dims.len()
178            )));
179        }
180        Ok(())
181    }
182}
183
184/// Comparison direction.
185///
186/// # Examples
187///
188/// ```rust
189/// use tenferro_tensor::CompareDir;
190///
191/// let dir = CompareDir::Eq;
192/// ```
193#[derive(Clone, Debug, Hash, PartialEq, Eq)]
194pub enum CompareDir {
195    Eq,
196    Lt,
197    Le,
198    Gt,
199    Ge,
200}
201
202/// StableHLO gather dimension configuration.
203///
204/// # Examples
205///
206/// ```rust
207/// use tenferro_tensor::GatherConfig;
208///
209/// let config = GatherConfig {
210///     offset_dims: vec![],
211///     collapsed_slice_dims: vec![0],
212///     start_index_map: vec![0],
213///     index_vector_dim: 1,
214///     slice_sizes: vec![1],
215/// };
216/// ```
217#[derive(Clone, Debug, Hash, PartialEq, Eq)]
218pub struct GatherConfig {
219    pub offset_dims: Vec<usize>,
220    pub collapsed_slice_dims: Vec<usize>,
221    pub start_index_map: Vec<usize>,
222    pub index_vector_dim: usize,
223    pub slice_sizes: Vec<usize>,
224}
225
226/// StableHLO scatter dimension configuration.
227///
228/// # Examples
229///
230/// ```rust
231/// use tenferro_tensor::ScatterConfig;
232///
233/// let config = ScatterConfig {
234///     update_window_dims: vec![],
235///     inserted_window_dims: vec![0],
236///     scatter_dims_to_operand_dims: vec![0],
237///     index_vector_dim: 1,
238/// };
239/// ```
240#[derive(Clone, Debug, Hash, PartialEq, Eq)]
241pub struct ScatterConfig {
242    pub update_window_dims: Vec<usize>,
243    pub inserted_window_dims: Vec<usize>,
244    pub scatter_dims_to_operand_dims: Vec<usize>,
245    pub index_vector_dim: usize,
246}
247
248/// Slice configuration.
249///
250/// # Examples
251///
252/// ```rust
253/// use tenferro_tensor::SliceConfig;
254///
255/// let config = SliceConfig {
256///     starts: vec![0],
257///     limits: vec![2],
258///     strides: vec![1],
259/// };
260/// ```
261#[derive(Clone, Debug, Hash, PartialEq, Eq)]
262pub struct SliceConfig {
263    pub starts: Vec<usize>,
264    pub limits: Vec<usize>,
265    pub strides: Vec<usize>,
266}
267
268/// StableHLO pad configuration.
269///
270/// # Examples
271///
272/// ```rust
273/// use tenferro_tensor::PadConfig;
274///
275/// let config = PadConfig {
276///     edge_padding_low: vec![1, 1],
277///     edge_padding_high: vec![1, 1],
278///     interior_padding: vec![0, 0],
279/// };
280/// ```
281#[derive(Clone, Debug, Hash, PartialEq, Eq)]
282pub struct PadConfig {
283    pub edge_padding_low: Vec<i64>,
284    pub edge_padding_high: Vec<i64>,
285    pub interior_padding: Vec<i64>,
286}