Files
giant/tests/test_flow.py
T
lars 9277d79dff 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>
2026-06-17 10:48:03 +02:00

60 lines
1.7 KiB
Python

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)