Skip to main content

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}