pad_safe_svd#

t3toolbox.backend.linalg.pad_safe_svd(A, row_mask, col_mask)#
def pad_safe_svd(
        A:        NDArray, # shape=(...,N,M); padded rows/columns identically zero
        row_mask: NDArray, # bool, shape=broadcastable to (...,N); True = REAL row, False = padded
        col_mask: NDArray, # bool, shape=broadcastable to (...,M); True = REAL column, False = padded
) -> typ.Tuple[
    NDArray, # U,  shape=(...,N,K), K=minimum(N,M); first q=minimum(n,m) columns bitwise zero on padded rows
    NDArray, # ss, shape=(...,K);   ss[q:] == 0 exactly
    NDArray, # Vt, shape=(...,K,M); first q rows bitwise zero on padded columns
]:

SVD of a zero-padded matrix whose null-space vectors are forced OFF the padding.

A black-box SVD of a padded matrix is wrong when the real block is numerically rank-deficient: the sigma ~= 0 left singular vectors are an arbitrary basis of a degenerate subspace that contains the pad coordinates, so they generically land on padded rows – and a downstream mask then erases them (a lost direction: sliced to its real rows, U is non-orthonormal and rank-deficient). pad_safe_svd takes the real/pad partition as data and realizes the contract

pad -> svd -> unpad  ==  svd of the unpadded n x m real block

with n = row_mask.sum(), m = col_mask.sum() and q = minimum(n, m):

  • A == U @ diag(ss) @ Vt – a genuine economy SVD of the padded matrix;

  • the first q triplets are a valid economy SVD of the real block: its positive singular values, then real-supported null completions, with U[..., :q] bitwise zero on padded rows and Vt[..., :q, :] bitwise zero on padded columns (exact == 0.0, not small);

  • the remaining K - q don’t-care triplets carry ss == 0 exactly;

  • no rank tolerance is used anywhere – every count comes from the masks or from bitwise {0, 1} indicators, so exact and roundoff-level zeros need no distinction.

Any real (n, m) is supported – tall, wide, or mixed across a batch – and pads may sit at arbitrary (interior) positions. Masks are runtime data of static length: under jax they may be traced (or host-numpy constants), and one jit compile covers every mask pattern. The only branch is on the static padded shape (N < M transposes internally).

Algorithm (Method D, “sketch-project”, from the pad-safe SVD design record: docs/pad_safe_svd.tex). Load-bearing details – do not “simplify”:

  • pad rows are permuted to the TRAILING pivot positions before the QR (Householder reflectors then never place mass on a padded row; the surplus columns come out as exact pad coordinate vectors, flagged by the bitwise indicator t);

  • the separation constant is c = 4 * ||A||_F (Frobenius, per batch element). The filter sigma > c/2 then has ||A||-sized margins on both sides. c = 2 * sigma_max is fragile: the threshold sits exactly at sigma_max and one-ulp rounding deletes the largest triplet (~38%% of generic rank-1 matrices, measured);

  • the augmented SVD’s right factor is discarded (zero columns where t = 0 pollute only it);

  • Vt is rebuilt from A.T @ U == V @ diag(ss) – exactly its own QR up to column signs.

The bitwise-zero guarantees rest on Householder-QR semantics (true for LAPACK, cuSOLVER, and jax’s qr lowering on all backends; the big matrix never sees an SVD – only the small augmented core does). They hold in float32 as well; only orthonormality/sigma accuracy scales with precision.

Cost O(N M^2 + M^3), independent of the pad counts; no (N, N) intermediate exists.

The complete derivation – the problem and contract, why each step works, the two-sided separation-constant measurements, and every alternative considered (augmentation/GSVD, post-hoc completion, masked noise, …) – is docs/pad_safe_svd.tex (+pdf).

Examples

The failure this exists for: a zero-padded warm start (interior pad rows, and a real column that is exactly zero – the padding of a rank-continuation restart). A black-box SVD puts null-space mass on the padded rows, so the masked real block is no longer orthonormal:

>>> import numpy as np
>>> import t3toolbox.backend.linalg as linalg
>>> np.random.seed(0)
>>> row_mask = np.array([True, False, True, True, False, True])   # pads at rows 1, 4 (interior)
>>> col_mask = np.array([True, True, True, False])                # 3 real columns of 4
>>> A = np.zeros((6, 4))
>>> A[np.ix_(row_mask, col_mask)] = np.hstack([np.random.randn(4, 2), np.zeros((4, 1))])
>>> U0, ss0, _ = np.linalg.svd(A, full_matrices=False)            # black-box SVD:
>>> print(bool(np.all(U0[~row_mask][:, :3] == 0.0)))              #   pad rows contaminated
False
>>> Ur0 = U0[row_mask][:, :3]
>>> print(float(np.round(np.linalg.norm(Ur0.T @ Ur0 - np.eye(3)), 2)))   # real block skewed
0.26

pad_safe_svd takes the masks (True = real, the library polarity) and returns clean factors – bitwise-zero pads, the real block orthonormal at full mask rank, singular values exactly those of the unpadded block:

>>> U, ss, Vt = linalg.pad_safe_svd(A, row_mask, col_mask)
>>> print(bool(np.all(U[~row_mask][:, :3] == 0.0)), bool(np.all(Vt[:3, ~col_mask] == 0.0)))
True True
>>> Ur = U[row_mask][:, :3]
>>> print(np.allclose(Ur.T @ Ur, np.eye(3)))                      # no lost directions
True
>>> print(np.allclose(ss[:3], np.linalg.svd(A[np.ix_(row_mask, col_mask)], compute_uv=False)))
True
>>> print(np.allclose(np.einsum('ix,x,xj->ij', U, ss, Vt), A), float(ss[3]))
True 0.0

A wide real block (n < m) needs no transpose by the caller – the contract is symmetric in min(n, m) (here 2) – and a statically wide matrix (N < M) transposes internally:

>>> B = np.zeros((3, 5)); B[:2, :4] = np.random.randn(2, 4)
>>> Uw, sw, Vtw = linalg.pad_safe_svd(B, np.array([1, 1, 0], bool), np.array([1, 1, 1, 1, 0], bool))
>>> print(Uw.shape, sw.shape, Vtw.shape)                          # K = min(N, M) = 3 triplets
(3, 3) (3,) (3, 5)
>>> print(np.allclose(sw[:2], np.linalg.svd(B[:2, :4], compute_uv=False)), bool(np.all(sw[2:] == 0.0)))
True True
Parameters:
Return type:

t3toolbox.backend.common.typ.Tuple[NDArray, NDArray, NDArray]