File size: 12,567 Bytes
ecc81b3 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 | """Phase 6 acceptance for the kernel family. See PLAN.md.
The load-bearing test builds the joint operator explicitly as a Kronecker
product and checks the factorized contraction equals it. That is only possible
while the lattice is small, which is exactly why it happens now rather than
after the attention modules are layered on top.
"""
import pytest
import torch
from torch_dimensions import Lattice
from torch_dimensions.compose.kernel import axial_contract, kron_operator
RANKS = [1, 2, 3, 4]
def _lat(rank, **kw):
return Lattice(shape=tuple(range(2, 2 + rank)), **kw)
def _kernels(lat, seed=0):
g = torch.Generator().manual_seed(seed)
return [torch.randn(s, s, generator=g, dtype=torch.float64) for s in lat.shape]
def _contract_all(x, lat, kernels, valid=None):
for axis, k in enumerate(kernels):
x = axial_contract(x, lat, axis, k, valid=valid)
return x
# -- the identity the whole family rests on ----------------------------------
@pytest.mark.parametrize("rank", RANKS)
def test_sequential_contraction_equals_the_kronecker_product(rank):
lat = _lat(rank)
kernels = _kernels(lat)
x = torch.randn(2, *lat.shape, 3, dtype=torch.float64)
got = _contract_all(x, lat, kernels)
# Independent reference: flatten the lattice and apply the joint operator.
joint = kron_operator(kernels)
flat = x.reshape(2, lat.n_cells, 3)
want = (joint @ flat).reshape(x.shape)
assert torch.allclose(got, want, atol=1e-10), (got - want).abs().max()
def test_the_joint_operator_is_as_large_as_advertised():
"""The reason the factorization exists: the explicit operator is quadratic
in cells, the factorized one only in axial size."""
lat = _lat(3) # (2, 3, 4) -> 24 cells
joint = kron_operator(_kernels(lat))
assert joint.shape == (24, 24)
assert sum(k.numel() for k in _kernels(lat)) == 4 + 9 + 16 < 24 * 24
@pytest.mark.parametrize("rank", RANKS)
def test_contraction_order_does_not_matter_on_a_dense_lattice(rank):
"""Kronecker factors commute across distinct axes; if ours do not, the
contraction is entangling axes it should not."""
lat = _lat(rank)
kernels = _kernels(lat)
x = torch.randn(2, *lat.shape, 3, dtype=torch.float64)
forward = _contract_all(x, lat, kernels)
backward = x
for axis in reversed(range(rank)):
backward = axial_contract(backward, lat, axis, kernels[axis])
assert torch.allclose(forward, backward, atol=1e-10)
def test_identity_kernels_leave_the_input_alone():
lat = _lat(3)
eye = [torch.eye(s, dtype=torch.float64) for s in lat.shape]
x = torch.randn(2, *lat.shape, 3, dtype=torch.float64)
assert torch.allclose(_contract_all(x, lat, eye), x, atol=1e-12)
def test_a_single_axis_contraction_is_a_plain_matmul():
lat = _lat(1)
k = _kernels(lat)[0]
x = torch.randn(2, 2, 3, dtype=torch.float64)
assert torch.allclose(axial_contract(x, lat, 0, k), k @ x, atol=1e-12)
def test_contraction_works_with_a_time_axis():
lat = _lat(2, time=True)
kernels = _kernels(lat)
x = torch.randn(2, 4, *lat.shape, 3, dtype=torch.float64)
out = x
for axis, k in enumerate(kernels):
out = axial_contract(out, lat, lat.axis_names[axis + 1], k)
assert out.shape == x.shape
def test_axes_can_be_named():
lat = Lattice(shape=(3, 4), names=("h", "w"))
k = torch.randn(4, 4, dtype=torch.float64)
x = torch.randn(2, 3, 4, 5, dtype=torch.float64)
assert torch.equal(axial_contract(x, lat, "w", k), axial_contract(x, lat, 1, k))
def test_a_batched_kernel_broadcasts_over_the_folded_batch():
lat = _lat(2)
x = torch.randn(2, *lat.shape, 3, dtype=torch.float64)
m = x.shape[0] * lat.shape[1] # folded rows when sweeping axis 0
k = torch.randn(m, 2, 2, dtype=torch.float64)
assert axial_contract(x, lat, 0, k).shape == x.shape
# -- sparse renormalization --------------------------------------------------
def _sparse(rank=2, seed=0):
shape = tuple(range(2, 2 + rank))
g = torch.Generator().manual_seed(seed)
valid = torch.rand(shape, generator=g) > 0.4
valid.reshape(-1)[0] = True
valid.reshape(-1)[-1] = True
return Lattice(shape=shape, valid=valid)
def test_renormalization_makes_a_uniform_kernel_average_only_present_cells():
"""With a uniform kernel the contraction is a mean; renormalized, it must
be the mean over cells that exist, not over all of them."""
valid = torch.tensor([[True, True, True], [True, False, False]])
lat = Lattice(shape=(2, 3), valid=valid)
x = torch.ones(1, 2, 3, 1, dtype=torch.float64) * lat.mask().to(torch.float64)
ones = torch.ones(3, 3, dtype=torch.float64)
out = axial_contract(x, lat, 1, ones, valid=lat.mask().to(torch.float64))
# Row 0 has three present cells all equal to 1 -> mean 1.
assert torch.allclose(out[0, 0], torch.ones(3, 1, dtype=torch.float64))
# Row 1 has one present cell equal to 1 -> still 1, not 1/3.
assert torch.allclose(out[0, 1], torch.ones(3, 1, dtype=torch.float64))
def test_without_renormalization_structural_zeros_dilute_the_result():
"""The control that gives the test above its meaning.
Needs a *row-stochastic* kernel to say anything: with an unnormalized
all-ones kernel the contraction is a sum rather than a mean, and a sum has
no dilution to show.
"""
valid = torch.tensor([[True, True, True], [True, False, False]])
lat = Lattice(shape=(2, 3), valid=valid)
mask = lat.mask().to(torch.float64)
x = torch.ones(1, 2, 3, 1, dtype=torch.float64) * mask
uniform = torch.full((3, 3), 1 / 3, dtype=torch.float64) # rows sum to 1
plain = axial_contract(x, lat, 1, uniform)
renormed = axial_contract(x, lat, 1, uniform, valid=mask)
one = torch.ones(3, 1, dtype=torch.float64)
# Row 0: all three cells present, so both agree on the true mean of 1.
assert torch.allclose(plain[0, 0], one)
assert torch.allclose(renormed[0, 0], one)
# Row 1: only one cell present. Unrenormalized it is averaged over three
# slots, two of which are structural zeros -> 1/3. That is the dilution.
assert torch.allclose(plain[0, 1], one / 3)
assert torch.allclose(renormed[0, 1], one)
@pytest.mark.parametrize("rank", [2, 3])
def test_absent_cell_values_cannot_influence_present_outputs(rank):
lat = _sparse(rank)
kernels = _kernels(lat)
mask = lat.mask().to(torch.float64)
x = torch.randn(2, *lat.shape, 3, dtype=torch.float64) * mask
noise = torch.randn_like(x) * 1e3 * (1 - mask)
a = _contract_all(x, lat, kernels, valid=mask) * mask
b = _contract_all(x + noise, lat, kernels, valid=mask) * mask
assert torch.equal(a, b), "absent cells leaked into present outputs"
def test_a_line_with_no_present_cells_stays_finite():
"""Dead lines divide by clamped zero; they must not produce NaN."""
valid = torch.tensor([[True, True], [False, False]])
lat = Lattice(shape=(2, 2), valid=valid)
mask = lat.mask().to(torch.float64)
x = torch.randn(1, 2, 2, 3, dtype=torch.float64) * mask
out = axial_contract(x, lat, 1, torch.randn(2, 2, dtype=torch.float64), valid=mask)
assert torch.isfinite(out).all()
def test_renormalization_is_a_no_op_on_a_dense_lattice_with_a_stochastic_kernel():
"""When every cell is present and the kernel rows sum to one, the
denominator is one everywhere and nothing changes."""
lat = _lat(2)
ones = torch.ones(*lat.shape, 1, dtype=torch.float64)
kernels = [torch.softmax(k, dim=-1) for k in _kernels(lat)]
x = torch.randn(2, *lat.shape, 3, dtype=torch.float64)
plain = _contract_all(x, lat, kernels)
renorm = _contract_all(x, lat, kernels, valid=ones)
assert torch.allclose(plain, renorm, atol=1e-10)
# -- autograd ----------------------------------------------------------------
def test_contraction_is_differentiable_through_both_arguments():
lat = _lat(2)
x = torch.randn(1, *lat.shape, 2, dtype=torch.float64, requires_grad=True)
kernels = [k.clone().requires_grad_(True) for k in _kernels(lat)]
_contract_all(x, lat, kernels).pow(2).sum().backward()
assert x.grad is not None
assert all(k.grad is not None for k in kernels)
def test_gradcheck_passes_through_the_contraction():
lat = _lat(2)
kernels = _kernels(lat)
def fn(x):
return _contract_all(x, lat, kernels)
x = torch.randn(1, *lat.shape, 2, dtype=torch.float64, requires_grad=True)
assert torch.autograd.gradcheck(fn, (x,), fast_mode=True)
def test_a_signed_kernel_does_not_explode_when_the_mass_cancels():
"""`clamp_min` assumes a non-negative mass. A signed kernel — LeakyReLU
scores, as upstream CaFA uses by default — can cancel to zero while the
numerator stays nonzero, and clamping to +eps then divides by ~0."""
lat = Lattice(shape=(2, 4), valid=torch.tensor([[1, 1, 0, 0], [1, 1, 1, 1]]).bool())
mask = lat.mask().to(torch.float64)
x = torch.randn(1, 2, 4, 3, dtype=torch.float64) * mask
signed = torch.tensor(
[
[1.0, -1.0, 0.5, 0.5],
[-1.0, 1.0, 0.5, 0.5],
[0.5, 0.5, 1.0, -1.0],
[0.5, 0.5, -1.0, 1.0],
],
dtype=torch.float64,
)
out = axial_contract(x, lat, 1, signed, valid=mask)
assert torch.isfinite(out).all()
# Row 0's mass cancels exactly; the output must stay the same order of
# magnitude as the input rather than blowing up by ~1e6.
assert out.abs().max() < 10 * x.abs().max(), out.abs().max().item()
def test_a_genuinely_dead_line_is_still_zero_under_the_guard():
"""Leaving degenerate lines unscaled must not resurrect them: with no
present cells the numerator is zero, so the output stays zero."""
lat = Lattice(shape=(2, 2), valid=torch.tensor([[True, True], [False, False]]))
mask = lat.mask().to(torch.float64)
x = torch.randn(1, 2, 2, 3, dtype=torch.float64) * mask
out = axial_contract(x, lat, 1, torch.rand(2, 2, dtype=torch.float64), valid=mask)
assert torch.isfinite(out).all()
assert out[0, 1].abs().max() == 0.0
def test_a_nan_in_the_input_is_not_silently_laundered():
"""A `nan_to_num` after the division zeroed NaNs arriving in `x`, hiding a
diverging model mid-network behind finite numbers. The magnitude guard
already makes the division itself safe, so the only NaNs reaching that
point are real upstream failures — and a NaN that arrives must leave."""
lat = Lattice(shape=(4,), valid=torch.tensor([True, True, True, False]))
mask = lat.mask().to(torch.float64)
x = torch.randn(2, 4, 3, dtype=torch.float64) * mask
x[0, 1, 2] = float("nan") # a present cell diverged upstream
out = axial_contract(x, lat, 0, torch.randn(4, 4, dtype=torch.float64), valid=mask)
assert bool(out.isnan().any()), "an input NaN vanished into finite output"
def test_float32_near_cancellation_does_not_explode():
"""The absolute-epsilon guard waved through a denominator of ~1e-4 —
small enough to amplify by 1e4, large enough to pass any tiny fixed
threshold — and float32 outputs blew up ~7000x. Degeneracy is
cancellation, and cancellation is *relative* to the absolute mass."""
lat = Lattice(shape=(2, 4), valid=torch.tensor([[1, 1, 0, 0], [1, 1, 1, 1]]).bool())
mask = lat.mask().to(torch.float32)
x = (torch.randn(1, 2, 4, 3) * 100) * mask
near_cancel = torch.tensor(
[
[1.0, -0.9999, 0.5, 0.5],
[-1.0, 1.0001, 0.5, 0.5],
[0.5, 0.5, 1.0, -1.0],
[0.5, 0.5, -1.0, 1.0],
]
)
out = axial_contract(x, lat, 1, near_cancel, valid=mask)
assert torch.isfinite(out).all()
assert out.abs().max() < 10 * x.abs().max(), out.abs().max().item()
def test_a_genuinely_small_mass_still_renormalizes_exactly():
"""The relative guard must not overreach: a tiny but uncancelled mass
divides out exactly, because the numerator carries the same factor."""
lat = Lattice(shape=(3,), valid=torch.tensor([True, False, False]))
mask = lat.mask().to(torch.float64)
x = torch.randn(2, 3, 4, dtype=torch.float64) * mask
tiny = torch.full((3, 3), 1e-6, dtype=torch.float64) # small, all-positive
out = axial_contract(x, lat, 0, tiny, valid=mask)
# one present cell, mass 1e-6, numerator 1e-6 * x -> renormalizes to x
assert torch.allclose(out[:, 0], x[:, 0], atol=1e-9)
def test_kron_operator_refuses_an_empty_kernel_list():
with pytest.raises(ValueError, match="at least one kernel"):
kron_operator([])
|