tenferro_linalg/householder.rs
1use std::fmt;
2use std::ops::Range;
3
4use tenferro_tensor::{BackendSession, Tensor};
5
6use crate::backend::CompactQrResult;
7use crate::QrOptions;
8
9/// Opaque compact Householder QR state.
10///
11/// The packed reflector tensors remain private so callers cannot accidentally
12/// treat provider state as ordinary tensor values.
13///
14/// # Examples
15///
16/// ```rust
17/// use tenferro_cpu::CpuBackend;
18/// use tenferro_linalg::{QrOptions, TensorLinalgExt};
19/// use tenferro_tensor::{BackendSessionHost, Tensor};
20///
21/// let a = Tensor::from_vec_col_major(
22/// vec![3, 2],
23/// vec![1.0_f64, 0.0, 1.0, 0.0, 1.0, 1.0],
24/// )?;
25/// let mut host = CpuBackend::new();
26/// let r = host.with_backend_session(|session| {
27/// let qr = a.householder_qr(session)?;
28/// qr.r(QrOptions::default(), session)
29/// })??;
30/// assert_eq!(r.shape(), &[2, 2]);
31/// # Ok::<(), tenferro_tensor::Error>(())
32/// ```
33#[derive(Clone)]
34pub struct HouseholderQr<T> {
35 pub(crate) packed: T,
36 pub(crate) coeff: T,
37}
38
39impl<T> fmt::Debug for HouseholderQr<T> {
40 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
41 formatter
42 .debug_struct("HouseholderQr")
43 .finish_non_exhaustive()
44 }
45}
46
47impl HouseholderQr<Tensor> {
48 pub(crate) fn from_backend(state: CompactQrResult) -> Self {
49 Self {
50 packed: state.packed,
51 coeff: state.coeff,
52 }
53 }
54
55 /// Construct compact state for the product of compatible factors `Q * R`.
56 ///
57 /// # Errors
58 ///
59 /// Returns `tenferro_tensor::Error::Validation` for incompatible rank,
60 /// shape, dtype, placement, or a non-trapezoidal R factor;
61 /// `tenferro_tensor::Error::Unsupported` when the provider lacks compact
62 /// QR; or `tenferro_tensor::Error::BackendSource` for provider failures.
63 pub fn from_factors(
64 q: &Tensor,
65 r: &Tensor,
66 session: &mut dyn BackendSession,
67 ) -> tenferro_tensor::Result<Self> {
68 crate::tensor_ext::with_linalg_backend(session, "householder_qr_from_factors", |backend| {
69 backend
70 .householder_qr_from_factors(q, r)
71 .map(Self::from_backend)
72 })
73 }
74
75 /// Append a column block without refactorizing existing columns.
76 ///
77 /// # Errors
78 ///
79 /// Returns `tenferro_tensor::Error::Validation` for incompatible shape,
80 /// dtype, placement, or malformed state; `tenferro_tensor::Error::Unsupported`
81 /// when append is unavailable; or `tenferro_tensor::Error::BackendSource`
82 /// for reflector or factorization failures.
83 pub fn append_columns(
84 &self,
85 block: &Tensor,
86 session: &mut dyn BackendSession,
87 ) -> tenferro_tensor::Result<Self> {
88 crate::tensor_ext::with_linalg_backend(session, "householder_qr_append", |backend| {
89 backend
90 .householder_qr_append(&self.packed, &self.coeff, block)
91 .map(Self::from_backend)
92 })
93 }
94
95 /// Extract the thin upper-trapezoidal factor.
96 ///
97 /// # Errors
98 ///
99 /// Returns `tenferro_tensor::Error::Validation` for malformed state,
100 /// `tenferro_tensor::Error::Unsupported` for an unavailable provider path,
101 /// or `tenferro_tensor::Error::BackendSource` for extraction failures.
102 pub fn r(
103 &self,
104 options: QrOptions,
105 session: &mut dyn BackendSession,
106 ) -> tenferro_tensor::Result<Tensor> {
107 crate::tensor_ext::with_linalg_backend(session, "householder_qr_r", |backend| {
108 backend.householder_qr_r(&self.packed, &self.coeff, options)
109 })
110 }
111
112 /// Materialize a contiguous range of Q columns, up to the full-Q width.
113 ///
114 /// For an `m x n` input with `k = min(m, n)`, columns `0..k` are the thin-Q
115 /// factor and columns `k..m` span the orthogonal complement of the input's
116 /// column space — the cheaper route to a nullspace basis than a full SVD.
117 /// `QrOptions::gauge` is defined by R's diagonal, so `PositiveDiagonal`
118 /// fixes the first `k` columns only; a complement column has no diagonal to
119 /// fix and is returned as the reflector product produced it.
120 ///
121 /// # Errors
122 ///
123 /// Returns `tenferro_tensor::Error::Validation` when the range is outside
124 /// the full-Q width `0..m` or state metadata is malformed,
125 /// `tenferro_tensor::Error::Unsupported` for an unavailable provider path,
126 /// or `tenferro_tensor::Error::BackendSource` for execution failures.
127 pub fn q_columns(
128 &self,
129 columns: Range<usize>,
130 options: QrOptions,
131 session: &mut dyn BackendSession,
132 ) -> tenferro_tensor::Result<Tensor> {
133 crate::tensor_ext::with_linalg_backend(session, "householder_qr_q_columns", |backend| {
134 backend.householder_qr_q_columns(&self.packed, &self.coeff, columns, options)
135 })
136 }
137}
138
139#[cfg(feature = "autodiff")]
140impl HouseholderQr<tenferro_ad::EagerTensor> {
141 pub(crate) fn from_eager_outputs(
142 packed: tenferro_ad::EagerTensor,
143 coeff: tenferro_ad::EagerTensor,
144 ) -> Self {
145 Self { packed, coeff }
146 }
147
148 /// Construct compact state from compatible eager factors.
149 ///
150 /// # Examples
151 /// ```rust
152 /// use tenferro_ad::{EagerRuntime, EagerTensor, Tensor};
153 /// use tenferro_linalg::HouseholderQr;
154 /// let ctx = EagerRuntime::new()?;
155 /// let state = ctx.with_eager_session(|s| {
156 /// let q = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![1.0_f64])?)?;
157 /// let r = s.constant_from(Tensor::from_vec_col_major([1, 1], vec![2.0_f64])?)?;
158 /// HouseholderQr::<EagerTensor>::from_factors(&q, &r, s)
159 /// })?;
160 /// assert!(format!("{state:?}").starts_with("HouseholderQr"));
161 /// # Ok::<(), tenferro_ad::Error>(())
162 /// ```
163 /// # Errors
164 ///
165 /// Returns `tenferro_ad::Error::Validation` for known invalid metadata,
166 /// `tenferro_ad::Error::Extension` for unsupported or provider failures,
167 /// or `tenferro_ad::Error::RuntimeState` when eager execution is unavailable.
168 pub fn from_factors(
169 q: &tenferro_ad::EagerTensor,
170 r: &tenferro_ad::EagerTensor,
171 session: &mut tenferro_ad::EagerSession<'_>,
172 ) -> tenferro_ad::Result<Self> {
173 eager_state(crate::eager_ext::apply_linalg_eager_in_session(
174 session,
175 crate::extension::LinalgOp::HouseholderQrFromFactors,
176 &[q, r],
177 )?)
178 }
179
180 /// Append an eager column block functionally.
181 ///
182 /// # Examples
183 /// ```rust
184 /// use tenferro_ad::{EagerRuntime, Tensor};
185 /// use tenferro_linalg::EagerSessionLinalgExt;
186 /// let ctx = EagerRuntime::new()?;
187 /// let state = ctx.with_eager_session(|s| {
188 /// let a = s.constant_from(Tensor::from_vec_col_major([2, 1], vec![1.0_f64, 0.0])?)?;
189 /// let b = s.constant_from(Tensor::from_vec_col_major([2, 1], vec![0.0_f64, 1.0])?)?;
190 /// s.householder_qr(&a)?.append_columns(&b, s)
191 /// })?;
192 /// assert!(format!("{state:?}").starts_with("HouseholderQr"));
193 /// # Ok::<(), tenferro_ad::Error>(())
194 /// ```
195 /// # Errors
196 ///
197 /// Returns `tenferro_ad::Error::Validation` for known invalid metadata,
198 /// `tenferro_ad::Error::Extension` for unsupported or provider failures,
199 /// or `tenferro_ad::Error::RuntimeState` when eager execution is unavailable.
200 pub fn append_columns(
201 &self,
202 block: &tenferro_ad::EagerTensor,
203 session: &mut tenferro_ad::EagerSession<'_>,
204 ) -> tenferro_ad::Result<Self> {
205 eager_state(crate::eager_ext::apply_linalg_eager_in_session(
206 session,
207 crate::extension::LinalgOp::HouseholderQrAppend,
208 &[&self.packed, &self.coeff, block],
209 )?)
210 }
211
212 /// Extract eager R.
213 ///
214 /// # Examples
215 /// ```rust
216 /// use tenferro_ad::{EagerRuntime, Tensor};
217 /// use tenferro_linalg::{EagerSessionLinalgExt, QrOptions};
218 /// let ctx = EagerRuntime::new()?;
219 /// let r = ctx.with_eager_session(|s| {
220 /// let a = s.constant_from(Tensor::from_vec_col_major([2, 1], vec![1.0_f64, 2.0])?)?;
221 /// s.householder_qr(&a)?.r(QrOptions::default(), s)
222 /// })?;
223 /// assert_eq!(r.shape(), &[1, 1]);
224 /// # Ok::<(), tenferro_ad::Error>(())
225 /// ```
226 /// # Errors
227 ///
228 /// Returns `tenferro_ad::Error::Validation` for known invalid metadata,
229 /// `tenferro_ad::Error::Extension` for unsupported or provider failures,
230 /// or `tenferro_ad::Error::RuntimeState` when eager execution is unavailable.
231 pub fn r(
232 &self,
233 options: QrOptions,
234 session: &mut tenferro_ad::EagerSession<'_>,
235 ) -> tenferro_ad::Result<tenferro_ad::EagerTensor> {
236 eager_one(
237 crate::eager_ext::apply_linalg_eager_in_session(
238 session,
239 crate::extension::LinalgOp::HouseholderQrR {
240 gauge: options.gauge,
241 },
242 &[&self.packed, &self.coeff],
243 )?,
244 "householder_qr_r",
245 )
246 }
247
248 /// Materialize eager Q columns, up to the full-Q width.
249 ///
250 /// # Examples
251 /// ```rust
252 /// use tenferro_ad::{EagerRuntime, Tensor};
253 /// use tenferro_linalg::{EagerSessionLinalgExt, QrOptions};
254 /// let ctx = EagerRuntime::new()?;
255 /// let q = ctx.with_eager_session(|s| {
256 /// let a = s.constant_from(Tensor::from_vec_col_major([2, 1], vec![1.0_f64, 2.0])?)?;
257 /// s.householder_qr(&a)?.q_columns(0..1, QrOptions::default(), s)
258 /// })?;
259 /// assert_eq!(q.shape(), &[2, 1]);
260 /// # Ok::<(), tenferro_ad::Error>(())
261 /// ```
262 /// Columns `0..k` are the thin-Q factor; columns `k..m` span the orthogonal
263 /// complement of the input's column space. `PositiveDiagonal` fixes only
264 /// the first `k` columns, because the gauge comes from R's diagonal.
265 /// Differentiating through a range that reaches past `k` is unsupported and
266 /// surfaces a typed AD error rather than a silently wrong derivative.
267 ///
268 /// # Errors
269 ///
270 /// Returns `tenferro_ad::Error::Validation` for known invalid metadata,
271 /// `tenferro_ad::Error::Extension` for unsupported or provider failures,
272 /// or `tenferro_ad::Error::RuntimeState` when eager execution is unavailable.
273 pub fn q_columns(
274 &self,
275 columns: Range<usize>,
276 options: QrOptions,
277 session: &mut tenferro_ad::EagerSession<'_>,
278 ) -> tenferro_ad::Result<tenferro_ad::EagerTensor> {
279 eager_one(
280 crate::eager_ext::apply_linalg_eager_in_session(
281 session,
282 crate::extension::LinalgOp::HouseholderQrQColumns {
283 start: columns.start,
284 end: columns.end,
285 gauge: options.gauge,
286 },
287 &[&self.packed, &self.coeff],
288 )?,
289 "householder_qr_q_columns",
290 )
291 }
292}
293
294#[cfg(feature = "autodiff")]
295pub(crate) fn eager_state(
296 outputs: Vec<tenferro_ad::EagerTensor>,
297) -> tenferro_ad::Result<HouseholderQr<tenferro_ad::EagerTensor>> {
298 let mut outputs = outputs.into_iter();
299 match (outputs.next(), outputs.next(), outputs.next()) {
300 (Some(packed), Some(coeff), None) => Ok(HouseholderQr::from_eager_outputs(packed, coeff)),
301 _ => Err(tenferro_ad::Error::Internal(
302 "compact Householder QR returned an unexpected output count".into(),
303 )),
304 }
305}
306
307#[cfg(feature = "autodiff")]
308fn eager_one(
309 outputs: Vec<tenferro_ad::EagerTensor>,
310 op: &'static str,
311) -> tenferro_ad::Result<tenferro_ad::EagerTensor> {
312 let mut outputs = outputs.into_iter();
313 match (outputs.next(), outputs.next()) {
314 (Some(output), None) => Ok(output),
315 _ => Err(tenferro_ad::Error::Internal(format!(
316 "{op} returned an unexpected output count"
317 ))),
318 }
319}
320
321impl HouseholderQr<tenferro_runtime::TracedTensor> {
322 pub(crate) fn from_traced_outputs(
323 packed: tenferro_runtime::TracedTensor,
324 coeff: tenferro_runtime::TracedTensor,
325 ) -> Self {
326 Self { packed, coeff }
327 }
328
329 /// Construct compact state from compatible traced factors.
330 ///
331 /// # Errors
332 ///
333 /// Returns `tenferro_runtime::Error::Validation` for known invalid metadata
334 /// or `tenferro_runtime::Error::Extension` for unsupported operation state.
335 ///
336 /// # Deferred errors
337 ///
338 /// Symbolic shape constraints and backend provider failures may be reported
339 /// during compile or execution.
340 pub fn from_factors(
341 q: &tenferro_runtime::TracedTensor,
342 r: &tenferro_runtime::TracedTensor,
343 ) -> tenferro_runtime::Result<Self> {
344 crate::validation::ensure_float_or_complex("householder_qr_from_factors", q.dtype())?;
345 crate::validation::ensure_float_or_complex("householder_qr_from_factors", r.dtype())?;
346 traced_state(
347 crate::extension::LinalgOp::HouseholderQrFromFactors,
348 &[q, r],
349 )
350 }
351
352 /// Append a traced column block functionally.
353 ///
354 /// # Errors
355 ///
356 /// Returns `tenferro_runtime::Error::Validation` for known invalid metadata
357 /// or `tenferro_runtime::Error::Extension` for unsupported operation state.
358 ///
359 /// # Deferred errors
360 ///
361 /// Symbolic shape constraints and backend provider failures may be reported
362 /// during compile or execution.
363 pub fn append_columns(
364 &self,
365 block: &tenferro_runtime::TracedTensor,
366 ) -> tenferro_runtime::Result<Self> {
367 crate::validation::ensure_float_or_complex("householder_qr_append", block.dtype())?;
368 traced_state(
369 crate::extension::LinalgOp::HouseholderQrAppend,
370 &[&self.packed, &self.coeff, block],
371 )
372 }
373
374 /// Extract traced R.
375 ///
376 /// # Errors
377 ///
378 /// Returns `tenferro_runtime::Error::Validation` for known invalid metadata
379 /// or `tenferro_runtime::Error::Extension` for unsupported operation state.
380 ///
381 /// # Deferred errors
382 ///
383 /// Symbolic shape constraints and backend provider failures may be reported
384 /// during compile or execution.
385 pub fn r(
386 &self,
387 options: QrOptions,
388 ) -> tenferro_runtime::Result<tenferro_runtime::TracedTensor> {
389 traced_one(
390 crate::extension::LinalgOp::HouseholderQrR {
391 gauge: options.gauge,
392 },
393 &[&self.packed, &self.coeff],
394 "householder_qr_r",
395 )
396 }
397
398 /// Materialize traced Q columns, up to the full-Q width.
399 ///
400 /// Columns `0..k` are the thin-Q factor; columns `k..m` span the orthogonal
401 /// complement of the input's column space. `PositiveDiagonal` fixes only
402 /// the first `k` columns, because the gauge comes from R's diagonal.
403 /// Differentiating through a range that reaches past `k` is unsupported and
404 /// surfaces a typed AD error rather than a silently wrong derivative.
405 ///
406 /// # Errors
407 ///
408 /// Returns `tenferro_runtime::Error::Validation` for known invalid metadata
409 /// or `tenferro_runtime::Error::Extension` for unsupported operation state.
410 ///
411 /// # Deferred errors
412 ///
413 /// Symbolic shape constraints and backend provider failures may be reported
414 /// during compile or execution.
415 pub fn q_columns(
416 &self,
417 columns: Range<usize>,
418 options: QrOptions,
419 ) -> tenferro_runtime::Result<tenferro_runtime::TracedTensor> {
420 traced_one(
421 crate::extension::LinalgOp::HouseholderQrQColumns {
422 start: columns.start,
423 end: columns.end,
424 gauge: options.gauge,
425 },
426 &[&self.packed, &self.coeff],
427 "householder_qr_q_columns",
428 )
429 }
430}
431
432fn traced_outputs(
433 op: crate::extension::LinalgOp,
434 inputs: &[&tenferro_runtime::TracedTensor],
435) -> tenferro_runtime::Result<Vec<tenferro_runtime::TracedTensor>> {
436 tenferro_runtime::extension::apply(
437 std::sync::Arc::new(crate::extension::LinalgExtensionOp::new(op)),
438 inputs,
439 )
440}
441
442fn traced_state(
443 op: crate::extension::LinalgOp,
444 inputs: &[&tenferro_runtime::TracedTensor],
445) -> tenferro_runtime::Result<HouseholderQr<tenferro_runtime::TracedTensor>> {
446 let mut outputs = traced_outputs(op, inputs)?.into_iter();
447 match (outputs.next(), outputs.next(), outputs.next()) {
448 (Some(packed), Some(coeff), None) => Ok(HouseholderQr::from_traced_outputs(packed, coeff)),
449 _ => Err(tenferro_runtime::Error::Internal(
450 "compact Householder QR returned an unexpected output count".into(),
451 )),
452 }
453}
454
455fn traced_one(
456 op: crate::extension::LinalgOp,
457 inputs: &[&tenferro_runtime::TracedTensor],
458 name: &'static str,
459) -> tenferro_runtime::Result<tenferro_runtime::TracedTensor> {
460 let mut outputs = traced_outputs(op, inputs)?.into_iter();
461 match (outputs.next(), outputs.next()) {
462 (Some(output), None) => Ok(output),
463 _ => Err(tenferro_runtime::Error::Internal(format!(
464 "{name} returned an unexpected output count"
465 ))),
466 }
467}