use crate::{Tensor, check, check::TensorCheck, linalg::l2_norm, s};
use alloc::vec;
use burn_std::Slice;
pub fn qr<const D: usize>(tensor: Tensor<D>, reduced: bool) -> (Tensor<D>, Tensor<D>) {
let dims = tensor.dims();
let original_dtype = tensor.dtype();
check!(TensorCheck::qr_input_tensor::<D>(
"linalg::qr",
&dims,
original_dtype
));
let device = tensor.device();
let (n_rows, n_cols) = (dims[D - 2], dims[D - 1]);
let max_iters = n_rows.min(n_cols);
let mut r = tensor.clone();
let identity: Tensor<2> = Tensor::eye(n_rows, &device);
let mut reshape_dims = [1; D];
reshape_dims[D - 2] = n_rows;
reshape_dims[D - 1] = n_rows;
let reshaped_identity = identity.reshape(reshape_dims);
let mut expand_dims = [n_rows; D];
expand_dims[..(D - 2)].copy_from_slice(&dims[..(D - 2)]);
let mut q = reshaped_identity.expand(expand_dims);
let mut slices = vec![Slice::full(); D];
for i in 0..max_iters {
let sub_tensor = r
.clone()
.slice_dim(D - 2, s![i..])
.slice_dim(D - 1, s![i..]);
let v = sub_tensor.clone().slice_dim(D - 1, 0..1);
let v0 = v.clone().slice_dim(D - 2, s![0]);
let zeros = v0.clone().zeros_like();
let norm_v = l2_norm(v.clone().slice_dim(D - 2, s![..]), D - 2);
let sign = -v0.clone().sign();
let mask = sign.clone().is_close(zeros.clone(), None, None);
let sign = sign.mask_fill(mask, -1.0);
let u0 = v0.clone().sub(norm_v.clone().mul(sign.clone()));
let mask = norm_v.clone().is_close(zeros.clone(), None, None);
let mut tau = -u0.clone().div(norm_v.clone()).mul(sign.clone());
tau = tau.clone().mask_fill(mask.clone(), 0.0);
let e0 = v0.clone().mul_scalar(0.0).add_scalar(1.0);
let mut w = v.clone().div(u0.clone());
slices[D - 2] = s![0];
w = w.slice_assign(&slices, e0.clone());
w = w.clone().mask_fill(mask, 0.0);
let f = |a: Tensor<D>| -> Tensor<D> {
let aw_out = a.matmul(w.clone());
let aw = aw_out.clone().mul(tau.clone());
w.clone().matmul(aw.transpose())
};
slices[D - 2] = s![i..];
slices[D - 1] = s![i..];
r = r.slice_assign(&slices, sub_tensor.clone() - f(sub_tensor.transpose()));
slices[D - 2] = Slice::full();
let q_sub_tensor = q.clone().slice(&slices);
q = q.slice_assign(&slices, q_sub_tensor.clone() - f(q_sub_tensor).transpose());
slices[D - 1] = Slice::full();
}
if reduced & (n_rows > n_cols) {
slices[D - 1] = s![0..n_cols];
let result_q = q.clone().slice(&slices);
slices[D - 2] = s![0..n_cols];
let result_r = r.clone().slice(&slices);
return (result_q, result_r);
}
(q, r)
}