# householder.jl
#
# A basic implementation of Householder factorization to obtain a QR
# decomposition of a matrix. We are not too carefuly---this should not
# be taken as anything more than a reference implementation---but it should
# work

using LinearAlgebra;

# u = find_annihilator(A::Matrix; tol = 1e-10)
#
# Finds an "annihilator" vector for the first column of A. By this,
# we mean a vector u so that if a = A[:, 1] is the first column
# of A, then
#
#   z = a - 2 * u * u' * a = (I - 2 * u * u') * a
#
# is a multiple of the first standard basis vector, where the sign
# may change depending on the sign of the first entry of a.
# The value |z[1]| = norm(a).
#
# If norm(a) < tol, returns the all-zeros vector, where tol is a specified
# tolerance for numerical accuracy.
function find_annihilator(A::Matrix; tol = 1e-10)
  n = size(A, 1);
  a = A[:, 1];
  if (norm(a) < tol)
    return zeros(n)
  end
  v = a / norm(a);
  u = v / norm(v);  # Just make sure we are good an normalized.

  # Find the annihilator vector by adding e_1 or subtracting e_1 to it.
  # We choose the sign of this addition/subtraction to make sure that
  # norm(u +/- e_1) is maximized and improve numerical conditioning.
  if (u[1] >= 0)
    u[1] += 1;
  else
    u[1] -= 1;
  end
  return u / norm(u);
end

# U = populate_annihilators(A::Matrix; tol = 1e-10)
#
# Recursively applies Householder transformations to an n-by-k matrix
# A, populating an n-by-m matrix U (with height n) and number m =
# min(n, k) - 1 of columns. The ith column of U annihilates the
# recursively constructed lower-right corner of A in the
# Householder-based construction of the QR decomposition.
#
# When this method terminates, we can construct a series of matrices
# H_1, ..., H_m of the form
#
#   H_i = I - 2 * U[:, i] * U[:, i]',
#
# where these H_i provide the guarantee that
#
#   Q = H_1 * H_2 * ... * H_m
#
# has orthogonal columns, and R = Q' * A is upper triangular, that is,
#
#   A = Q * R.
#
# The parameter tol is a numerical precision parameter
function populate_annihilators(A::Matrix; tol = 1e-10)
  (n, k) = size(A);
  m = min(k, n - 1);
  # U will be a matrix of the annihilation vectors
  U = zeros(n, m);
  curr_A = A;
  for ii = 1:m
    u = find_annihilator(curr_A, tol = tol);
    U[ii:end, ii] = u;
    # Now recurse down on the lower right corner of the annihilated matrix:
    B = curr_A - 2 * u * u' * curr_A;
    curr_A = B[2:end, 2:end];
  end
  return U;
end

# (Q, R) = compute_QR(A::Matrix; tol = 1e-10)
#
# Computes a full QR factorization of the matrix A using Householder
# transformations. When this method returns, at least to within
# high precision, the result satisfies
#
#   A = Q * R
#
# where R is an upper triangluar matrix and Q is n-by-n, where A is
# assumed to be n-by-k. Note that we have made the choice that if
# k < n, then we still return an n-by-n matrix Q.
# If k > n, then the matrix R will be of the form
#
#   R = [T *]
#
# where T is an n-by-n upper triangular matrix, and * is a potentially
# full matrix.
#
# If A does not have linearly independent columns, then while R will
# still be upper triangular, it will have zeros on diagonal entries
# for columns a_i of A that are linearly dependent on a_1, ...,
# a_{i-1}.
#
# The tolerance tol is a numerical precision value, where norms
# less than tol are treated as 0 in finding the Householder transformations.
function compute_QR(A::Matrix; tol = 1e-10, positive_diagonal::Bool = false)
  (n, k) = size(A);
  U = populate_annihilators(A, tol = tol);
  m = size(U, 2);

  # Now, we compute the matrices Q and R, but try to do it a little
  # bit efficiently by leveraging the structure of the annihilator
  # vectors and Householder-type transformations.
  Q = zeros(n, n);
  for ii = 1:n
    # Populate by multiplying sequence of householder transformations
    # by e_i, the i-th standard basis vector.  To do this, we
    # initialize the i-th column of Q to be e_i, then apply
    #
    #  q_i = (I - 2 * u_1 * u_1') * ... * (I - 2 * u_m * u_m') * q_i
    #
    # by iterating backwards
    Q[ii, ii] = 1;
    for jj = m:-1:1
      Q[:, ii] -= 2 * U[:, jj] * U[:, jj]' * Q[:, ii];
    end
  end
  # Since A = Q * R, we have R = Q' * A.
  R = Q' * A;

  # The below code fixes the signs of the diagonal of R to be nonnegative.
  # It is not necessary.
  if (positive_diagonal)
    for ii = 1:min(n, k)
      if (R[ii, ii] < 0)
        Q[:, ii] *= -1;
        R[ii, :] *= -1;
      end
    end
  end
  # For completeness, zero out the rows of R below the diagonal:
  for ii = 1:min(n, k)
    R[(ii+1):end, ii] .= 0;
  end
  # Uncomment this line to zero out entries of R that are too small.
  # R[abs.(R) .< tol / n] .= 0;
  return (Q, R);
end
