72bd65ff9f
step_length already encodes |post_pos - pre_pos| by definition, so a raw post_pos target would duplicate that magnitude and could drift inconsistent with step_length during sampling. Instead add travel_dir, a unit vector (local frame) giving only the direction of pre_pos->post_pos; post_pos is reconstructed at inference as pre_pos + step_length * travel_dir, keeping the two self-consistent. Target grows from 6D to 9D. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
121 lines
4.8 KiB
Python
121 lines
4.8 KiB
Python
import numpy as np
|
|
import pytest
|
|
from giant.data.transforms import (
|
|
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_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_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)
|