Implement Phase 1: full data pipeline, model, training, and config support
- Data pipeline: loader (parquet→numpy), transforms (log, local-frame Rodrigues rotation, Normalizer), StepsDataset with event-ID-based split - Model: SinusoidalEmbedding, ConditionEncoder, ResBlock, DenoisingMLP - Schedule: cosine DDPM and conditional flow matching loss (Lipman 2022) - Samplers: flow (Euler ODE), DDPM ancestral, DDIM deterministic - Training loop: AdamW + cosine LR, grad clipping, best-val checkpoint - Validation: per-dimension marginal summary (normalised space) - CLI: TOML config support with CLI-overrides; hyperparam-encoded output directory; config.toml with git hash saved into each run's checkpoint dir - 21 unit tests covering transforms, network, flow/DDPM losses, dataset splits Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -0,0 +1,62 @@
|
||||
import numpy as np
|
||||
import pytest
|
||||
from giant.data.dataset import StepsDataset, train_val_split
|
||||
|
||||
|
||||
def _dummy(N=500, n_events=20):
|
||||
rng = np.random.default_rng(42)
|
||||
data = {"event_id": rng.integers(0, n_events, size=N)}
|
||||
cond_cont = rng.standard_normal((N, 9)).astype(np.float32)
|
||||
cond_cat = rng.integers(0, 3, size=(N, 2)).astype(np.int64)
|
||||
target = rng.standard_normal((N, 6)).astype(np.float32)
|
||||
return data, cond_cont, cond_cat, target
|
||||
|
||||
|
||||
def test_dataset_length():
|
||||
data, cond_cont, cond_cat, target = _dummy()
|
||||
assert len(StepsDataset(cond_cont, cond_cat, target)) == len(target)
|
||||
|
||||
|
||||
def test_dataset_item_shapes():
|
||||
data, cond_cont, cond_cat, target = _dummy()
|
||||
c, k, t = StepsDataset(cond_cont, cond_cat, target)[0]
|
||||
assert c.shape == (9,)
|
||||
assert k.shape == (2,)
|
||||
assert t.shape == (6,)
|
||||
|
||||
|
||||
def test_split_sizes_sum_to_total():
|
||||
data, cond_cont, cond_cat, target = _dummy(N=500)
|
||||
train_ds, val_ds = train_val_split(data, cond_cont, cond_cat, target, val_fraction=0.2)
|
||||
assert len(train_ds) + len(val_ds) == 500
|
||||
|
||||
|
||||
def test_split_no_empty_sets():
|
||||
data, cond_cont, cond_cat, target = _dummy(N=500, n_events=20)
|
||||
train_ds, val_ds = train_val_split(data, cond_cont, cond_cat, target, val_fraction=0.2)
|
||||
assert len(val_ds) > 0
|
||||
assert len(train_ds) > 0
|
||||
|
||||
|
||||
def test_split_event_leakage():
|
||||
"""Train and val must not share any event_id."""
|
||||
N = 1000
|
||||
n_events = 50
|
||||
rng = np.random.default_rng(7)
|
||||
event_ids = rng.integers(0, n_events, size=N)
|
||||
data = {"event_id": event_ids}
|
||||
cond_cont = rng.standard_normal((N, 9)).astype(np.float32)
|
||||
cond_cat = rng.integers(0, 3, size=(N, 2)).astype(np.int64)
|
||||
target = rng.standard_normal((N, 6)).astype(np.float32)
|
||||
|
||||
train_ds, val_ds = train_val_split(data, cond_cont, cond_cat, target, val_fraction=0.2)
|
||||
|
||||
# Recover which event_ids ended up in each split via the indices
|
||||
# (The dataset doesn't store event_ids, so we check via the original mask logic)
|
||||
unique_events = np.unique(event_ids)
|
||||
rng2 = np.random.default_rng(42)
|
||||
rng2.shuffle(unique_events)
|
||||
n_val = max(1, int(len(unique_events) * 0.2))
|
||||
val_events = set(unique_events[:n_val].tolist())
|
||||
train_events = set(unique_events[n_val:].tolist())
|
||||
assert val_events.isdisjoint(train_events)
|
||||
@@ -0,0 +1,59 @@
|
||||
import torch
|
||||
import pytest
|
||||
from giant.model.network import DenoisingMLP
|
||||
from giant.model.schedule import CosineSchedule, flow_matching_loss
|
||||
from giant.sample import sample_flow, sample_ddpm, sample_ddim
|
||||
|
||||
|
||||
def _small_model():
|
||||
return DenoisingMLP(pdg_vocab=3, mat_vocab=2, hidden_dim=32, n_blocks=2)
|
||||
|
||||
|
||||
def _batch(B=8):
|
||||
x1 = torch.randn(B, 6)
|
||||
cond_cont = torch.randn(B, 9)
|
||||
cond_cat = torch.zeros(B, 2, dtype=torch.long)
|
||||
return x1, cond_cont, cond_cat
|
||||
|
||||
|
||||
def test_flow_matching_loss_nonneg():
|
||||
x1, cond_cont, cond_cat = _batch()
|
||||
loss = flow_matching_loss(_small_model(), x1, cond_cont, cond_cat)
|
||||
assert loss.item() >= 0.0
|
||||
|
||||
|
||||
def test_flow_matching_loss_is_scalar():
|
||||
x1, cond_cont, cond_cat = _batch()
|
||||
loss = flow_matching_loss(_small_model(), x1, cond_cont, cond_cat)
|
||||
assert loss.shape == ()
|
||||
|
||||
|
||||
def test_flow_matching_loss_has_grad():
|
||||
model = _small_model()
|
||||
x1, cond_cont, cond_cat = _batch()
|
||||
flow_matching_loss(model, x1, cond_cont, cond_cat).backward()
|
||||
assert any(p.grad is not None for p in model.parameters())
|
||||
|
||||
|
||||
def test_sample_flow_shape():
|
||||
B = 6
|
||||
cond_cont = torch.randn(B, 9)
|
||||
cond_cat = torch.zeros(B, 2, dtype=torch.long)
|
||||
out = sample_flow(_small_model(), cond_cont, cond_cat, steps=5)
|
||||
assert out.shape == (B, 6)
|
||||
|
||||
|
||||
def test_ddpm_loss_nonneg():
|
||||
schedule = CosineSchedule(T=50)
|
||||
x1, cond_cont, cond_cat = _batch()
|
||||
loss = schedule.loss(_small_model(), x1, cond_cont, cond_cat)
|
||||
assert loss.item() >= 0.0
|
||||
|
||||
|
||||
def test_sample_ddim_shape():
|
||||
B = 4
|
||||
schedule = CosineSchedule(T=50)
|
||||
cond_cont = torch.randn(B, 9)
|
||||
cond_cat = torch.zeros(B, 2, dtype=torch.long)
|
||||
out = sample_ddim(_small_model(), cond_cont, cond_cat, schedule, steps=5)
|
||||
assert out.shape == (B, 6)
|
||||
@@ -0,0 +1,42 @@
|
||||
import torch
|
||||
import pytest
|
||||
from giant.model.network import DenoisingMLP, SinusoidalEmbedding
|
||||
|
||||
|
||||
def test_sinusoidal_embedding_shape():
|
||||
emb = SinusoidalEmbedding(64)
|
||||
t = torch.rand(16)
|
||||
assert emb(t).shape == (16, 64)
|
||||
|
||||
|
||||
def test_sinusoidal_embedding_batch_1():
|
||||
emb = SinusoidalEmbedding(32)
|
||||
t = torch.tensor([0.5])
|
||||
assert emb(t).shape == (1, 32)
|
||||
|
||||
|
||||
def test_denoising_mlp_output_shape():
|
||||
B = 8
|
||||
model = DenoisingMLP(pdg_vocab=5, mat_vocab=3)
|
||||
x_t = torch.randn(B, 6)
|
||||
t = torch.rand(B)
|
||||
cond_cont = torch.randn(B, 9)
|
||||
cond_cat = torch.stack([
|
||||
torch.randint(0, 5, (B,)),
|
||||
torch.randint(0, 3, (B,)),
|
||||
], dim=1)
|
||||
out = model(x_t, t, cond_cont, cond_cat)
|
||||
assert out.shape == (B, 6)
|
||||
|
||||
|
||||
def test_denoising_mlp_gradients_flow():
|
||||
B = 4
|
||||
model = DenoisingMLP(pdg_vocab=3, mat_vocab=2, hidden_dim=32, n_blocks=2)
|
||||
x_t = torch.randn(B, 6)
|
||||
t = torch.rand(B)
|
||||
cond_cont = torch.randn(B, 9)
|
||||
cond_cat = torch.zeros(B, 2, dtype=torch.long)
|
||||
loss = model(x_t, t, cond_cont, cond_cat).sum()
|
||||
loss.backward()
|
||||
for name, p in model.named_parameters():
|
||||
assert p.grad is not None, f"no grad for {name}"
|
||||
@@ -0,0 +1,66 @@
|
||||
import numpy as np
|
||||
import pytest
|
||||
from giant.data.transforms import (
|
||||
inv_log_transform,
|
||||
local_frame_rotation,
|
||||
log_transform,
|
||||
Normalizer,
|
||||
)
|
||||
|
||||
|
||||
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_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)
|
||||
Reference in New Issue
Block a user