a867fc4aae
Provides stratified marginal comparisons, joint-structure checks (correlation matrices, physically-coupled pairwise plots, direction alignment), and physical-constraint validation (unit-norm directions, non-negative raw targets) for a trained model's generated samples, building on the aggregate marginal/KL check already in giant.validate. Supports two entry points: live sampling against a checkpoint + val data (load_model_bundle/collect_samples), or loading a precomputed `giant predict --coord local` parquet directly (load_predicted_local) without needing the checkpoint at all. Predict output is now tagged with parquet schema metadata so the loader can verify a file's format and reject coord=global or untagged files with a clear error instead of guessing from column names. Also extends the config git-hash mismatch warning (added for --config loading) to checkpoint loading: both `giant predict` and analysis.load_model_bundle now look for a config.toml next to the checkpoint and warn (without failing) if it was generated from a different git commit. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
244 lines
8.4 KiB
Python
244 lines
8.4 KiB
Python
import matplotlib
|
|
|
|
matplotlib.use("Agg") # no display needed for plot smoke tests
|
|
|
|
import numpy as np
|
|
import pandas as pd
|
|
import pyarrow as pa
|
|
import pyarrow.parquet as pq
|
|
import pytest
|
|
|
|
from giant.analysis import (
|
|
RAW_TARGET_NAMES,
|
|
SampleCollection,
|
|
constraint_report,
|
|
correlation_matrices,
|
|
direction_alignment,
|
|
load_predicted_local,
|
|
marginal_table,
|
|
plot_constraint_violations,
|
|
plot_correlation_matrices,
|
|
plot_direction_alignment,
|
|
plot_marginals,
|
|
plot_pairwise,
|
|
)
|
|
from giant.constants import (
|
|
LOCAL_TARGET_NAMES,
|
|
PREDICT_COORD_METADATA_KEY,
|
|
PREDICT_SCHEMA_VERSION,
|
|
PREDICT_SCHEMA_VERSION_KEY,
|
|
)
|
|
from giant.data.transforms import log_transform
|
|
|
|
|
|
def _unit_vectors(rng, n):
|
|
v = rng.standard_normal((n, 3)).astype(np.float32)
|
|
return v / np.linalg.norm(v, axis=1, keepdims=True)
|
|
|
|
|
|
def _make_collection(n=200, seed=0, gen_offset=0.0) -> SampleCollection:
|
|
rng = np.random.default_rng(seed)
|
|
real = np.column_stack([
|
|
rng.uniform(0.1, 5.0, n), # step_length
|
|
rng.uniform(0.1, 5.0, n), # delta_e
|
|
rng.uniform(0.1, 5.0, n), # edep
|
|
_unit_vectors(rng, n), # post_dir
|
|
_unit_vectors(rng, n), # travel_dir
|
|
]).astype(np.float32)
|
|
gen = real + gen_offset
|
|
|
|
pre_E = rng.uniform(1.0, 100.0, n).astype(np.float32)
|
|
cond_cont_raw = np.column_stack([
|
|
rng.standard_normal((n, 3)), pre_E, rng.standard_normal((n, 3)),
|
|
rng.integers(0, 5, n), rng.integers(0, 3, n),
|
|
]).astype(np.float32)
|
|
|
|
return SampleCollection(
|
|
cond_cont_raw=cond_cont_raw,
|
|
pdg=rng.choice([11, -11, 22], size=n),
|
|
material=rng.choice(["W", "Pb"], size=n),
|
|
real_raw=real, gen_raw=gen,
|
|
real_norm=real, gen_norm=gen,
|
|
)
|
|
|
|
|
|
def test_marginal_table_aggregate_has_all_dims():
|
|
table = marginal_table(_make_collection())
|
|
assert set(table["dim"]) == set(RAW_TARGET_NAMES)
|
|
assert (table["group"] == "all").all()
|
|
|
|
|
|
@pytest.mark.parametrize("group_by", ["pdg", "material", "energy"])
|
|
def test_marginal_table_grouped_covers_all_rows(group_by):
|
|
collection = _make_collection()
|
|
table = marginal_table(collection, group_by=group_by)
|
|
assert table["n"].groupby(table["group"]).first().sum() == len(collection.pdg)
|
|
|
|
|
|
def test_marginal_table_identical_distributions_have_zero_kl():
|
|
collection = _make_collection(gen_offset=0.0)
|
|
table = marginal_table(collection)
|
|
np.testing.assert_allclose(table["kl_real_gen"], 0.0, atol=1e-6)
|
|
|
|
|
|
def test_marginal_table_shifted_distribution_has_positive_kl():
|
|
collection = _make_collection(gen_offset=3.0)
|
|
table = marginal_table(collection)
|
|
assert (table["kl_real_gen"] > 0).all()
|
|
|
|
|
|
def test_correlation_matrices_are_symmetric_unit_diagonal():
|
|
real_corr, gen_corr = correlation_matrices(_make_collection())
|
|
for corr in (real_corr, gen_corr):
|
|
np.testing.assert_allclose(np.diag(corr), 1.0, atol=1e-5)
|
|
np.testing.assert_allclose(corr, corr.T, atol=1e-5)
|
|
|
|
|
|
def test_direction_alignment_real_data_is_unit_norm_dot_product():
|
|
real_cos, gen_cos = direction_alignment(_make_collection())
|
|
assert np.all(real_cos >= -1.0 - 1e-5) and np.all(real_cos <= 1.0 + 1e-5)
|
|
assert np.all(gen_cos >= -1.0 - 1e-5) and np.all(gen_cos <= 1.0 + 1e-5)
|
|
|
|
|
|
def test_constraint_report_clean_data_has_no_violations():
|
|
report = constraint_report(_make_collection(gen_offset=0.0))
|
|
assert (report["violation_rate"] == 0.0).all()
|
|
|
|
|
|
def test_constraint_report_flags_negative_log_dims_and_bad_norms():
|
|
collection = _make_collection(gen_offset=0.0)
|
|
collection.gen_raw[:, 0] = -1.0 # negative step_length
|
|
collection.gen_raw[:, 3:6] *= 2.0 # post_dir no longer unit norm
|
|
report = constraint_report(collection)
|
|
violations = dict(zip(report["check"], report["violation_rate"]))
|
|
assert violations["step_length >= 0"] == 1.0
|
|
assert violations["post_dir unit norm"] == 1.0
|
|
|
|
|
|
def test_plot_marginals_runs_without_error():
|
|
fig = plot_marginals(_make_collection())
|
|
assert fig is not None
|
|
|
|
|
|
def test_plot_marginals_grouped_runs_without_error():
|
|
fig = plot_marginals(_make_collection(), group_by="material")
|
|
assert fig is not None
|
|
|
|
|
|
def test_plot_correlation_matrices_runs_without_error():
|
|
fig = plot_correlation_matrices(_make_collection())
|
|
assert fig is not None
|
|
|
|
|
|
def test_plot_pairwise_runs_without_error():
|
|
fig = plot_pairwise(_make_collection())
|
|
assert fig is not None
|
|
|
|
|
|
def test_plot_direction_alignment_runs_without_error():
|
|
fig = plot_direction_alignment(_make_collection())
|
|
assert fig is not None
|
|
|
|
|
|
def test_plot_constraint_violations_runs_without_error():
|
|
fig = plot_constraint_violations(_make_collection())
|
|
assert fig is not None
|
|
|
|
|
|
def _write_predicted_local_parquet(path, n=50, metadata=None, rng=None):
|
|
"""Mimic `giant predict --coord local`'s output schema for the loader tests."""
|
|
rng = rng or np.random.default_rng(0)
|
|
true_log_local = rng.standard_normal((n, 9)).astype(np.float32)
|
|
true_log_local[:, :3] = log_transform(rng.uniform(0.1, 5.0, (n, 3)).astype(np.float32))
|
|
pred_log_local = true_log_local + rng.normal(0, 0.01, (n, 9)).astype(np.float32)
|
|
|
|
table = pa.table({
|
|
"event_id": rng.integers(0, 10, n),
|
|
"pdg": rng.choice([11, -11, 22], n),
|
|
"pre_x": rng.standard_normal(n).astype(np.float32),
|
|
"pre_y": rng.standard_normal(n).astype(np.float32),
|
|
"pre_z": rng.standard_normal(n).astype(np.float32),
|
|
"pre_E": rng.uniform(1.0, 100.0, n).astype(np.float32),
|
|
"pre_dx": rng.standard_normal(n).astype(np.float32),
|
|
"pre_dy": rng.standard_normal(n).astype(np.float32),
|
|
"pre_dz": rng.standard_normal(n).astype(np.float32),
|
|
"material": rng.choice(["W", "Pb"], n),
|
|
"layer_id": rng.integers(0, 10, n).astype(np.int32),
|
|
"n_sec": rng.integers(0, 3, n).astype(np.int32),
|
|
**{f"pred_{name}": pred_log_local[:, j] for j, name in enumerate(LOCAL_TARGET_NAMES)},
|
|
**{f"true_{name}": true_log_local[:, j] for j, name in enumerate(LOCAL_TARGET_NAMES)},
|
|
})
|
|
if metadata is not None:
|
|
table = table.replace_schema_metadata(metadata)
|
|
pq.write_table(table, path)
|
|
return true_log_local, pred_log_local
|
|
|
|
|
|
def test_load_predicted_local_round_trips_values(tmp_path):
|
|
path = tmp_path / "predicted_local.parquet"
|
|
true_log_local, pred_log_local = _write_predicted_local_parquet(
|
|
path,
|
|
metadata={
|
|
PREDICT_COORD_METADATA_KEY: "local",
|
|
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
|
|
},
|
|
)
|
|
|
|
collection = load_predicted_local(path)
|
|
|
|
expected_real = true_log_local.copy()
|
|
expected_real[:, :3] = np.exp(expected_real[:, :3]) - 1e-8
|
|
expected_gen = pred_log_local.copy()
|
|
expected_gen[:, :3] = np.exp(expected_gen[:, :3]) - 1e-8
|
|
np.testing.assert_allclose(collection.real_raw, expected_real, atol=1e-4)
|
|
np.testing.assert_allclose(collection.gen_raw, expected_gen, atol=1e-4)
|
|
assert collection.real_norm is None
|
|
assert collection.gen_norm is None
|
|
|
|
|
|
def test_load_predicted_local_usable_by_downstream_plots(tmp_path):
|
|
path = tmp_path / "predicted_local.parquet"
|
|
_write_predicted_local_parquet(
|
|
path,
|
|
metadata={
|
|
PREDICT_COORD_METADATA_KEY: "local",
|
|
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
|
|
},
|
|
)
|
|
collection = load_predicted_local(path)
|
|
assert marginal_table(collection) is not None
|
|
assert plot_marginals(collection) is not None
|
|
|
|
|
|
def test_load_predicted_local_rejects_missing_metadata(tmp_path):
|
|
path = tmp_path / "no_metadata.parquet"
|
|
_write_predicted_local_parquet(path, metadata=None)
|
|
with pytest.raises(ValueError, match="no '.*' parquet metadata"):
|
|
load_predicted_local(path)
|
|
|
|
|
|
def test_load_predicted_local_rejects_global_coord(tmp_path):
|
|
path = tmp_path / "global.parquet"
|
|
_write_predicted_local_parquet(
|
|
path,
|
|
metadata={
|
|
PREDICT_COORD_METADATA_KEY: "global",
|
|
PREDICT_SCHEMA_VERSION_KEY: PREDICT_SCHEMA_VERSION,
|
|
},
|
|
)
|
|
with pytest.raises(ValueError, match="coord=local"):
|
|
load_predicted_local(path)
|
|
|
|
|
|
def test_load_predicted_local_rejects_mismatched_schema_version(tmp_path):
|
|
path = tmp_path / "old_version.parquet"
|
|
_write_predicted_local_parquet(
|
|
path,
|
|
metadata={
|
|
PREDICT_COORD_METADATA_KEY: "local",
|
|
PREDICT_SCHEMA_VERSION_KEY: "999",
|
|
},
|
|
)
|
|
with pytest.raises(ValueError, match="schema version"):
|
|
load_predicted_local(path)
|