diff options
Diffstat (limited to 'tests/test_core.py')
| -rw-r--r-- | tests/test_core.py | 89 |
1 files changed, 89 insertions, 0 deletions
diff --git a/tests/test_core.py b/tests/test_core.py index 50d3aad..09e58f1 100644 --- a/tests/test_core.py +++ b/tests/test_core.py @@ -234,3 +234,92 @@ def test_caption_encoding_matches_declared_factors(): assert states.shape == (2, 9) assert states[0][0] > 0 and states[0][2 + 2] > 0 and states[0][2 + 4 + 0] > 0 assert states[1][1] > 0 and states[1][2 + 0] > 0 and states[1][2 + 4 + 1] > 0 + + +def test_shared_rank_separates_aligned_from_shuffled(): + """The shared-direction count must sit at the null when nothing is shared.""" + import numpy as np + from worldalign.shared_rank import spectrum + + size, width = 192, 16 + generator = np.random.default_rng(0) + common = generator.normal(size=(size, 8)) + shared_field = common @ common.T + visual = shared_field + 0.05 * generator.normal(size=(size, size)) + text = shared_field + 0.05 * generator.normal(size=(size, size)) + visual, text = (visual + visual.T) / 2, (text + text.T) / 2 + + def count(first, second): + _, first_vectors = spectrum(first) + _, second_vectors = spectrum(second) + cosines = np.linalg.svd( + first_vectors[:, :width].T @ second_vectors[:, :width], compute_uv=False + ) + return int((np.clip(cosines, 0, 1) > 0.7).sum()) + + order = generator.permutation(size) + aligned = count(visual, text) + shuffled = count(visual, text[np.ix_(order, order)]) + assert aligned >= 8, aligned + assert shuffled <= 2, shuffled + + +def test_degree_decomposition_is_exact_and_orthogonal(): + """Fitted plus residual must reconstruct the field off the diagonal.""" + import numpy as np + from worldalign.field_anatomy import degree_part + + generator = np.random.default_rng(1) + size = 32 + rows = generator.normal(size=(size, 1)) + field = rows + rows.T + 0.1 * generator.normal(size=(size, size)) + np.fill_diagonal(field, 0.0) + fitted, residual = degree_part(field) + mask = ~np.eye(size, dtype=bool) + assert np.allclose(fitted[mask] + residual[mask], field[mask]) + # a field that is purely additive leaves almost nothing in the residual + pure = rows + rows.T + np.fill_diagonal(pure, 0.0) + _, pure_residual = degree_part(pure) + assert np.abs(pure_residual[mask]).max() < 1e-9 + + +def test_third_moment_kernel_raises_field_rank(): + """A degree-3 set kernel must span more directions than a degree-2 one.""" + import numpy as np + from worldalign.natural_pipeline import moment_field_degree + + generator = torch.Generator().manual_seed(0) + sets = [torch.randn(4, 8, generator=generator) for _ in range(24)] + second = moment_field_degree(sets, 2, 8).double().numpy() + third = moment_field_degree(sets, 3, 8).double().numpy() + + def effective_rank(matrix): + values = np.abs(np.linalg.eigvalsh((matrix + matrix.T) / 2)) + weights = values / values.sum() + weights = weights[weights > 1e-15] + return float(np.exp(-(weights * np.log(weights)).sum())) + + assert effective_rank(third) > effective_rank(second) + + +def test_structured_phrase_encoding_separates_head_from_modifier(): + """Word order must change the encoding, which averaging cannot express.""" + import numpy as np + + vectors = {"red": np.array([1.0, 0.0]), "bus": np.array([0.0, 1.0])} + + def structured(tokens): + head = vectors[tokens[-1]] + rest = ( + np.mean([vectors[t] for t in tokens[:-1]], axis=0) + if len(tokens) > 1 + else np.zeros(2) + ) + return np.concatenate([head, rest]) + + assert not np.allclose(structured(["red", "bus"]), structured(["bus", "red"])) + assert np.allclose( + np.mean([vectors[t] for t in ["red", "bus"]], axis=0), + np.mean([vectors[t] for t in ["bus", "red"]], axis=0), + ) |
