Files
giant/tests/test_transforms.py
T
lars e5bf7c51cb Fix silent failure modes surfaced by extensive code review
- energy_simplex_encode: warn when clipping post_E to pre_E discards
  recorded edep/e_sec instead of silently zeroing them
- local/inv_local_frame_rotation: validate and normalize pre_dir instead
  of silently assuming unit norm; raise on near-zero-norm rows
- train(): make --lr authoritative on resume instead of being silently
  overwritten by the checkpoint's optimizer/scheduler state; print and
  exit cleanly instead of silently training zero epochs when the
  checkpoint already meets --epochs; truncate metrics.csv on a fresh
  run instead of always appending
- dwarf update-manifest: check file existence for every manifest line,
  not just ones whose gen/schema actually changed
- pyproject.toml: dev extra now pulls in convert+analysis so the
  documented `uv sync --extra cpu --extra dev` + `pytest` actually
  passes collection

Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
2026-07-02 16:54:02 +02:00

206 lines
8.4 KiB
Python

import numpy as np
import pytest
from giant.data.transforms import (
energy_simplex_decode,
energy_simplex_encode,
inv_local_frame_rotation,
inv_log_transform,
local_frame_rotation,
log_transform,
Normalizer,
reconstruct_post_pos,
travel_direction,
)
def test_log_transform_invertible():
x = np.array([0.1, 1.0, 10.0, 1000.0], dtype=np.float32)
np.testing.assert_allclose(inv_log_transform(log_transform(x)), x, rtol=1e-5)
def test_local_frame_rotation_noop_when_aligned():
N = 8
pre_dir = np.tile([0.0, 0.0, 1.0], (N, 1)).astype(np.float32)
rng = np.random.default_rng(0)
post_dir = rng.standard_normal((N, 3)).astype(np.float32)
post_dir /= np.linalg.norm(post_dir, axis=1, keepdims=True)
result = local_frame_rotation(pre_dir, post_dir)
np.testing.assert_allclose(result, post_dir, atol=1e-5)
def test_local_frame_rotation_preserves_angle():
"""Angle between pre_dir and post_dir must equal angle between ẑ and rotated."""
rng = np.random.default_rng(1)
N = 200
pre_dir = rng.standard_normal((N, 3)).astype(np.float32)
pre_dir /= np.linalg.norm(pre_dir, axis=1, keepdims=True)
post_dir = rng.standard_normal((N, 3)).astype(np.float32)
post_dir /= np.linalg.norm(post_dir, axis=1, keepdims=True)
rotated = local_frame_rotation(pre_dir, post_dir)
cos_before = (pre_dir * post_dir).sum(axis=1)
cos_after = rotated[:, 2] # dot with ẑ = z-component (unit vectors)
np.testing.assert_allclose(cos_after, cos_before, atol=1e-5)
def test_local_frame_rotation_preserves_norm():
rng = np.random.default_rng(2)
N = 100
pre_dir = rng.standard_normal((N, 3)).astype(np.float32)
pre_dir /= np.linalg.norm(pre_dir, axis=1, keepdims=True)
post_dir = rng.standard_normal((N, 3)).astype(np.float32)
post_dir /= np.linalg.norm(post_dir, axis=1, keepdims=True)
result = local_frame_rotation(pre_dir, post_dir)
np.testing.assert_allclose(np.linalg.norm(result, axis=1), 1.0, atol=1e-5)
def test_local_frame_rotation_rejects_near_zero_pre_dir():
"""A degenerate (near-zero-norm) pre_dir has no well-defined frame — must
raise instead of silently falling back to an arbitrary rotation axis."""
pre_dir = np.array([[0.0, 0.0, 0.0], [0.0, 0.0, 1.0]], dtype=np.float32)
post_dir = np.array([[1.0, 0.0, 0.0], [0.0, 1.0, 0.0]], dtype=np.float32)
with pytest.raises(ValueError, match="near-zero norm"):
local_frame_rotation(pre_dir, post_dir)
with pytest.raises(ValueError, match="near-zero norm"):
inv_local_frame_rotation(pre_dir, post_dir)
def test_local_frame_rotation_normalizes_non_unit_pre_dir():
"""A pre_dir with float32-drift norm (not exactly 1) must still produce the
same result as its exactly-normalized counterpart, not a skewed frame."""
rng = np.random.default_rng(9)
N = 50
pre_dir_unit = rng.standard_normal((N, 3)).astype(np.float32)
pre_dir_unit /= np.linalg.norm(pre_dir_unit, axis=1, keepdims=True)
post_dir = rng.standard_normal((N, 3)).astype(np.float32)
post_dir /= np.linalg.norm(post_dir, axis=1, keepdims=True)
pre_dir_scaled = pre_dir_unit * rng.uniform(0.9, 1.1, size=(N, 1)).astype(
np.float32
)
expected = local_frame_rotation(pre_dir_unit, post_dir)
result = local_frame_rotation(pre_dir_scaled, post_dir)
np.testing.assert_allclose(result, expected, atol=1e-4)
def test_travel_direction_is_unit_norm():
rng = np.random.default_rng(5)
N = 50
pre_pos = rng.standard_normal((N, 3)).astype(np.float32)
post_pos = pre_pos + rng.standard_normal((N, 3)).astype(np.float32)
result = travel_direction(pre_pos, post_pos)
np.testing.assert_allclose(np.linalg.norm(result, axis=1), 1.0, atol=1e-5)
def test_travel_direction_matches_normalized_displacement():
rng = np.random.default_rng(6)
N = 50
pre_pos = rng.standard_normal((N, 3)).astype(np.float32)
disp = rng.standard_normal((N, 3)).astype(np.float32)
post_pos = pre_pos + disp
expected = disp / np.linalg.norm(disp, axis=1, keepdims=True)
np.testing.assert_allclose(travel_direction(pre_pos, post_pos), expected, atol=1e-5)
def test_reconstruct_post_pos_straight_line():
"""When post_pos = pre_pos + L * pre_dir, travel_dir equals pre_dir, so its
local-frame encoding is ẑ — reconstruction must recover post_pos exactly."""
rng = np.random.default_rng(7)
N = 20
pre_pos = rng.standard_normal((N, 3)).astype(np.float32)
pre_dir = rng.standard_normal((N, 3)).astype(np.float32)
pre_dir /= np.linalg.norm(pre_dir, axis=1, keepdims=True)
step_length = rng.uniform(0.1, 5.0, size=N).astype(np.float32)
post_pos = pre_pos + step_length[:, None] * pre_dir
travel_dir_local = local_frame_rotation(
pre_dir, travel_direction(pre_pos, post_pos)
)
np.testing.assert_allclose(travel_dir_local, np.tile([0, 0, 1], (N, 1)), atol=1e-4)
reconstructed = reconstruct_post_pos(
pre_pos, pre_dir, step_length, travel_dir_local
)
np.testing.assert_allclose(reconstructed, post_pos, atol=1e-4)
def test_reconstruct_post_pos_general_roundtrip():
"""Full encode (build_features-style) -> decode (cli.py predict-style) path."""
rng = np.random.default_rng(8)
N = 100
pre_pos = rng.standard_normal((N, 3)).astype(np.float32)
pre_dir = rng.standard_normal((N, 3)).astype(np.float32)
pre_dir /= np.linalg.norm(pre_dir, axis=1, keepdims=True)
post_pos = pre_pos + rng.standard_normal((N, 3)).astype(np.float32)
step_length = np.linalg.norm(post_pos - pre_pos, axis=1).astype(np.float32)
travel_dir_local = local_frame_rotation(
pre_dir, travel_direction(pre_pos, post_pos)
)
reconstructed = reconstruct_post_pos(
pre_pos, pre_dir, step_length, travel_dir_local
)
np.testing.assert_allclose(reconstructed, post_pos, atol=1e-4)
def test_energy_simplex_conservation():
"""Decoding any ALR coords yields energies that sum to pre_E exactly."""
rng = np.random.default_rng(11)
N = 500
z = rng.standard_normal((N, 2)).astype(np.float32) * 3.0
pre_E = rng.uniform(1.0, 100.0, N).astype(np.float32)
edep, e_sec, post_E, delta_e = energy_simplex_decode(z, pre_E)
np.testing.assert_allclose(edep + e_sec + post_E, pre_E, rtol=1e-5, atol=1e-4)
np.testing.assert_allclose(delta_e, edep + e_sec, rtol=1e-5, atol=1e-4)
assert np.all(edep >= 0) and np.all(e_sec >= 0) and np.all(post_E >= 0)
def test_energy_simplex_roundtrip():
"""Encode → decode recovers energies whose lost part already sums to delta_e."""
rng = np.random.default_rng(12)
N = 500000
pre_E = rng.uniform(1.0, 100.0, N).astype(np.float32)
post_E = (pre_E * rng.uniform(0.0, 1.0, N)).astype(np.float32)
delta_e = pre_E - post_E
g = rng.uniform(0.0, 1.0, N).astype(np.float32)
edep = (g * delta_e).astype(np.float32)
e_sec = ((1.0 - g) * delta_e).astype(np.float32)
z = energy_simplex_encode(edep, e_sec, post_E, pre_E)
edep_r, e_sec_r, post_E_r, _ = energy_simplex_decode(z, pre_E)
# Tolerance reflects the tiny simplex floor (~1e-5 of pre_E).
np.testing.assert_allclose(edep_r, edep, atol=5e-3)
np.testing.assert_allclose(e_sec_r, e_sec, atol=5e-3)
np.testing.assert_allclose(post_E_r, post_E, atol=5e-3)
def test_energy_simplex_handles_boundary_zeros():
"""e_sec=0 (no secondaries) and post_E=0 (track end) stay finite and decode near 0."""
pre_E = np.array([10.0, 50.0, 100.0], dtype=np.float32)
edep = np.array([4.0, 50.0, 0.0], dtype=np.float32)
e_sec = np.array([0.0, 0.0, 0.0], dtype=np.float32) # no secondaries
post_E = np.array([6.0, 0.0, 100.0], dtype=np.float32) # row 1: track ends
z = energy_simplex_encode(edep, e_sec, post_E, pre_E)
assert np.all(np.isfinite(z))
_, e_sec_r, post_E_r, _ = energy_simplex_decode(z, pre_E)
np.testing.assert_allclose(e_sec_r, 0.0, atol=1e-2)
assert post_E_r[1] < 1e-2 # the absorbed track decodes to ~0 post energy
def test_normalizer_roundtrip():
rng = np.random.default_rng(3)
X = rng.standard_normal((200, 9)).astype(np.float32)
norm = Normalizer().fit(X)
np.testing.assert_allclose(norm.inverse_transform(norm.transform(X)), X, atol=1e-5)
def test_normalizer_serialization():
rng = np.random.default_rng(4)
X = rng.standard_normal((50, 6)).astype(np.float32)
norm = Normalizer().fit(X)
norm2 = Normalizer.from_dict(norm.to_dict())
np.testing.assert_allclose(norm2.mean, norm.mean)
np.testing.assert_allclose(norm2.std, norm.std)