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}