Add class-balanced secondary particle-type loss (gitea #44)
CI / Lint (ruff check) (push) Successful in 36s
CI / Format (ruff format) (push) Successful in 38s
CI / Sync project version with tag (push) Has been skipped
CI / Type check (ty) (push) Successful in 38s
CI / Format (ruff format) (pull_request) Successful in 43s
CI / Lint (ruff check) (pull_request) Successful in 45s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 47s
CI / Tests (push) Successful in 5m20s
CI / Tests (pull_request) Successful in 4m50s
CI / Lint (ruff check) (push) Successful in 36s
CI / Format (ruff format) (push) Successful in 38s
CI / Sync project version with tag (push) Has been skipped
CI / Type check (ty) (push) Successful in 38s
CI / Format (ruff format) (pull_request) Successful in 43s
CI / Lint (ruff check) (pull_request) Successful in 45s
CI / Sync project version with tag (pull_request) Has been skipped
CI / Type check (ty) (pull_request) Successful in 47s
CI / Tests (push) Successful in 5m20s
CI / Tests (pull_request) Successful in 4m50s
The v0.3.0 pivot exists because the 2026-08-03 WGAN rollout benchmark produced zero photon secondaries and ~4M hallucinated antineutrinos — even with a correctly-sized top-N species vocabulary (gitea #29), plain cross-entropy over a class distribution spanning orders of magnitude still under-predicts rare-but-physical species. stage2_model.particle_type.class_weighting = "none" | "inverse_freq" (default "none", fully back-compat) weights the stage-2 type head's CE loss (FlowDDPMStageTrainer._type_loss) by inverse class frequency, normalized to mean 1 so switching it on doesn't rescale the type loss against particle_type.lambda / the generator loss it's summed with. The per-class counts the weighting needs don't already exist despite the issue's premise: _topn_plus_other_map (giant/data/loader.py) previously kept counts only for keys folded into "other", dropping the kept classes' counts on the floor. TopNMap now carries class_counts (index -> count), round-tripped through the setup-cache sidecar (format version bumped 3->4, since existing sidecars have none) and through checkpoints (tolerantly — a pre-#44 checkpoint decodes to {}, since only training-time loss weighting reads it, not inference). Decisions made during planning (with the user): dropped "effective_num" from the issue's proposed three-way enum (no beta hyperparameter to design around) — final domain is "none" | "inverse_freq". Weights are mean-1-normalized. validate_config rejects class_weighting != "none" combined with particle_type.target != "onehot" or stage2_model.generator == "wgan" (both have no class CE to weight), following the #28/#30 dead-key-must-not-go-silent convention. Branch fix/issue-44 off master. Co-Authored-By: Claude Sonnet 5 <noreply@anthropic.com>
This commit is contained in:
@@ -10,6 +10,7 @@ from unittest.mock import MagicMock, patch
|
||||
import numpy as np
|
||||
import pytest
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
|
||||
from giant.config import ParticleTypeConfig
|
||||
from giant.constants import (
|
||||
@@ -35,6 +36,7 @@ from giant.training import (
|
||||
train,
|
||||
)
|
||||
from giant.training.metrics import _wandb_run_config
|
||||
from giant.training.trainers import _type_class_weight_vector
|
||||
from giant.training.stage2_inputs import (
|
||||
_ar_has_prev,
|
||||
_assemble_stage2_ar_inputs,
|
||||
@@ -599,6 +601,151 @@ def test_flow_stage_trainer_ddpm_not_implemented_for_stage2():
|
||||
FlowDDPMStageTrainer(spec, torch.nn.Linear(1, 1), torch.device("cpu"))
|
||||
|
||||
|
||||
# --- gitea #44: class-balanced secondary particle-type loss -----------------
|
||||
|
||||
|
||||
def test_type_class_weight_vector_none_scheme_returns_none():
|
||||
assert _type_class_weight_vector({0: 100, 1: 5}, n_classes=2, scheme="none") is None
|
||||
|
||||
|
||||
def test_type_class_weight_vector_raises_without_counts():
|
||||
with pytest.raises(ValueError, match="class_counts"):
|
||||
_type_class_weight_vector({}, n_classes=4, scheme="inverse_freq")
|
||||
|
||||
|
||||
def test_type_class_weight_vector_inverse_freq_favors_rare_class_and_has_mean_one():
|
||||
weights = _type_class_weight_vector({0: 1000, 1: 10, 2: 1, 3: 1}, n_classes=4, scheme="inverse_freq")
|
||||
assert weights is not None
|
||||
assert len(weights) == 4
|
||||
assert weights[1] > weights[0] # rarer class -> larger weight
|
||||
assert math.isclose(sum(weights) / len(weights), 1.0, rel_tol=1e-9)
|
||||
|
||||
|
||||
def test_type_class_weight_vector_missing_index_clamps_to_count_one():
|
||||
# n_classes=3 but only index 0 was ever observed (e.g. a tiny dataset) —
|
||||
# indices 1/2 must not divide by zero.
|
||||
weights = _type_class_weight_vector({0: 10}, n_classes=3, scheme="inverse_freq")
|
||||
assert weights is not None
|
||||
assert all(math.isfinite(w) for w in weights)
|
||||
|
||||
|
||||
def _onehot_flow_stage2_setup():
|
||||
"""A built stage-2 model + a batch, under target='onehot' + generator='flow'
|
||||
(mirrors the 'stage2_onehot_target_flow' case in test_train_end_to_end)."""
|
||||
cfg = _base_cfg()
|
||||
cfg["stage2_model"]["generator"] = "flow"
|
||||
cfg["stage2_model"]["particle_type"] = {"target": "onehot", "lambda": 1.0}
|
||||
model_config = _model_config(cfg)
|
||||
model = build_models(model_config)["stage2"]
|
||||
assert model is not None
|
||||
batch = _fake_batches(1, 8)[0]
|
||||
device = torch.device("cpu")
|
||||
cond_cont, cond_cat, x1_s1 = batch.cond_cont, batch.cond_cat, batch.target_s1
|
||||
# Mostly class 0 (common), a few slot 1's set to class 1 (rare) —
|
||||
# PARTICLE_CFG's emb_dim=8, n_classes=0 (inherit) -> 8 type classes.
|
||||
sec_type_idx = torch.zeros(8, K_MAX, dtype=torch.long)
|
||||
sec_type_idx[:, :2] = 1
|
||||
sec_mask = torch.ones(8, K_MAX, dtype=torch.bool)
|
||||
return model, cond_cont, cond_cat, x1_s1, sec_type_idx, sec_mask, device
|
||||
|
||||
|
||||
def test_flow_ddpm_trainer_type_loss_none_leaves_weight_unset():
|
||||
model, *_ = _onehot_flow_stage2_setup()
|
||||
spec = StageSpec(
|
||||
name="stage2",
|
||||
is_stage2=True,
|
||||
generator="flow",
|
||||
particle_type=ParticleTypeConfig(target="onehot", class_weighting="none"),
|
||||
particle_type_n_classes=8,
|
||||
ema_decay=0.0,
|
||||
)
|
||||
trainer = FlowDDPMStageTrainer(spec, model, torch.device("cpu"))
|
||||
assert trainer.type_class_weights is None
|
||||
|
||||
|
||||
def test_flow_ddpm_trainer_type_loss_matches_manual_weighted_cross_entropy():
|
||||
model, cond_cont, cond_cat, x1_s1, sec_type_idx, sec_mask, device = _onehot_flow_stage2_setup()
|
||||
class_counts = {0: 1000, 1: 10, 2: 1, 3: 1, 4: 1, 5: 1, 6: 1, 7: 1}
|
||||
weights = _type_class_weight_vector(class_counts, n_classes=8, scheme="inverse_freq")
|
||||
spec = StageSpec(
|
||||
name="stage2",
|
||||
is_stage2=True,
|
||||
generator="flow",
|
||||
particle_type=ParticleTypeConfig(target="onehot", class_weighting="inverse_freq"),
|
||||
particle_type_n_classes=8,
|
||||
type_class_weights=weights,
|
||||
ema_decay=0.0,
|
||||
)
|
||||
trainer = FlowDDPMStageTrainer(spec, model, device)
|
||||
assert trainer.type_class_weights is not None
|
||||
stage1_ctx = trainer._stage1_context(x1_s1, cond_cont, cond_cat, epoch=None)
|
||||
|
||||
with torch.no_grad():
|
||||
type_out = model.predict_type(cond_cont, cond_cat, stage1_ctx)
|
||||
weight_t = torch.tensor(weights)
|
||||
ce = F.cross_entropy(type_out.transpose(1, 2), sec_type_idx, weight=weight_t, reduction="none")
|
||||
expected = (ce * sec_mask.float()).sum() / sec_mask.float().sum().clamp(min=1)
|
||||
|
||||
l_type, _ = trainer._type_loss(cond_cont, cond_cat, stage1_ctx, sec_type_idx, sec_mask, device)
|
||||
|
||||
assert torch.allclose(l_type, expected, atol=1e-6)
|
||||
|
||||
# Unweighted trainer, same model/batch — the two losses must differ
|
||||
# (the batch mixes the common and rare classes, so weighting changes the
|
||||
# per-slot contributions), confirming the weight is actually plumbed in.
|
||||
spec_none = StageSpec(
|
||||
name="stage2",
|
||||
is_stage2=True,
|
||||
generator="flow",
|
||||
particle_type=ParticleTypeConfig(target="onehot", class_weighting="none"),
|
||||
particle_type_n_classes=8,
|
||||
ema_decay=0.0,
|
||||
)
|
||||
trainer_none = FlowDDPMStageTrainer(spec_none, model, device)
|
||||
with torch.no_grad():
|
||||
l_type_none, _ = trainer_none._type_loss(cond_cont, cond_cat, stage1_ctx, sec_type_idx, sec_mask, device)
|
||||
assert not torch.allclose(l_type, l_type_none)
|
||||
|
||||
|
||||
def test_build_stage_trainers_threads_sec_type_class_counts_into_weights():
|
||||
cfg = _base_cfg()
|
||||
cfg["stage2_model"]["generator"] = "flow"
|
||||
cfg["stage2_model"]["particle_type"] = {
|
||||
"target": "onehot",
|
||||
"lambda": 1.0,
|
||||
"class_weighting": "inverse_freq",
|
||||
}
|
||||
model_config = _model_config(cfg)
|
||||
models = build_models(model_config)
|
||||
critics = build_critics(model_config)
|
||||
class_counts = {i: 100 for i in range(8)}
|
||||
class_counts[1] = 1 # one rare class
|
||||
trainers = build_stage_trainers(
|
||||
cfg, models, critics, torch.device("cpu"), total_train_batches=4, sec_type_class_counts=class_counts
|
||||
)
|
||||
stage2_trainer = trainers["stage2"]
|
||||
assert isinstance(stage2_trainer, FlowDDPMStageTrainer)
|
||||
weights = stage2_trainer.type_class_weights
|
||||
assert weights is not None
|
||||
assert weights[1] > weights[0]
|
||||
|
||||
|
||||
def test_build_stage_trainers_no_class_counts_with_none_weighting_is_fine():
|
||||
"""The overwhelmingly common case (class_weighting = 'none', the
|
||||
default): build_stage_trainers must not require sec_type_class_counts at
|
||||
all."""
|
||||
cfg = _base_cfg()
|
||||
cfg["stage2_model"]["generator"] = "flow"
|
||||
cfg["stage2_model"].__setitem__("particle_type", {"target": "onehot", "lambda": 1.0})
|
||||
model_config = _model_config(cfg)
|
||||
models = build_models(model_config)
|
||||
critics = build_critics(model_config)
|
||||
trainers = build_stage_trainers(cfg, models, critics, torch.device("cpu"), total_train_batches=4)
|
||||
stage2_trainer = trainers["stage2"]
|
||||
assert isinstance(stage2_trainer, FlowDDPMStageTrainer)
|
||||
assert stage2_trainer.type_class_weights is None
|
||||
|
||||
|
||||
# --- gitea #42: freeze / init_from -------------------------------------------
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user