Encode edep/secondary/post energy as a conservation-constrained simplex

Replaces the independent log_delta_e/log_edep targets with 2 additive-log-ratio
coordinates over the deposit/secondary/post-energy simplex (fractions of pre_E
summing to 1), so edep + e_sec + post_E == pre_E holds by construction after
decoding (softmax) rather than being learned approximately. Requires e_sec
(secondary energy) as a new conditioning input and a steps_to_parquet.py pass
to derive it from child track first-step energies.
This commit is contained in:
2026-06-25 16:01:13 +02:00
parent 64c6bd1cef
commit 8475199609
13 changed files with 415 additions and 83 deletions
+4 -3
View File
@@ -1,4 +1,5 @@
import torch
from giant.constants import COND_DIM
from giant.model.network import DenoisingMLP
from giant.model.schedule import CosineSchedule, flow_matching_loss
from giant.sample import sample_flow, sample_ddim
@@ -10,7 +11,7 @@ def _small_model():
def _batch(B=8):
x1 = torch.randn(B, 9)
cond_cont = torch.randn(B, 9)
cond_cont = torch.randn(B, COND_DIM)
cond_cat = torch.zeros(B, 2, dtype=torch.long)
return x1, cond_cont, cond_cat
@@ -36,7 +37,7 @@ def test_flow_matching_loss_has_grad():
def test_sample_flow_shape():
B = 6
cond_cont = torch.randn(B, 9)
cond_cont = torch.randn(B, COND_DIM)
cond_cat = torch.zeros(B, 2, dtype=torch.long)
out = sample_flow(_small_model(), cond_cont, cond_cat, steps=5)
assert out.shape == (B, 9)
@@ -52,7 +53,7 @@ def test_ddpm_loss_nonneg():
def test_sample_ddim_shape():
B = 4
schedule = CosineSchedule(T=50)
cond_cont = torch.randn(B, 9)
cond_cont = torch.randn(B, COND_DIM)
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, 9)