Apply ruff format and document lint/type tooling in CLAUDE.md
First repo-wide ruff format pass, plus a note in CLAUDE.md to run ruff and ty periodically. Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
+65
-35
@@ -39,27 +39,36 @@ def _unit_vectors(rng, n):
|
||||
|
||||
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)
|
||||
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)
|
||||
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,
|
||||
real_raw=real,
|
||||
gen_raw=gen,
|
||||
real_norm=real,
|
||||
gen_norm=gen,
|
||||
)
|
||||
|
||||
|
||||
@@ -150,25 +159,35 @@ 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))
|
||||
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)},
|
||||
})
|
||||
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)
|
||||
@@ -263,14 +282,21 @@ def test_marginal_table_pl_matches_numpy_version(tmp_path, group_by):
|
||||
path = _predicted_local_path(tmp_path)
|
||||
collection = load_predicted_local(path)
|
||||
|
||||
expected = marginal_table(collection, group_by=group_by).sort_values(["group", "dim"])
|
||||
actual = marginal_table_pl(path, group_by=group_by).sort(["group", "dim"]).to_pandas()
|
||||
expected = marginal_table(collection, group_by=group_by).sort_values(
|
||||
["group", "dim"]
|
||||
)
|
||||
actual = (
|
||||
marginal_table_pl(path, group_by=group_by).sort(["group", "dim"]).to_pandas()
|
||||
)
|
||||
|
||||
assert list(expected["group"]) == list(actual["group"])
|
||||
assert list(expected["n"]) == list(actual["n"])
|
||||
for col in ["real_mean", "gen_mean", "real_std", "gen_std", "kl_real_gen"]:
|
||||
np.testing.assert_allclose(
|
||||
expected[col].to_numpy(), actual[col].to_numpy(), atol=1e-4, rtol=1e-4,
|
||||
expected[col].to_numpy(),
|
||||
actual[col].to_numpy(),
|
||||
atol=1e-4,
|
||||
rtol=1e-4,
|
||||
)
|
||||
|
||||
|
||||
@@ -290,8 +316,12 @@ def test_constraint_report_pl_matches_numpy_version(tmp_path):
|
||||
|
||||
assert list(expected["check"]) == list(actual["check"])
|
||||
np.testing.assert_allclose(
|
||||
expected["violation_rate"].to_numpy(), actual["violation_rate"].to_numpy(), atol=1e-6,
|
||||
expected["violation_rate"].to_numpy(),
|
||||
actual["violation_rate"].to_numpy(),
|
||||
atol=1e-6,
|
||||
)
|
||||
np.testing.assert_allclose(
|
||||
expected["mean_abs_error"].to_numpy(), actual["mean_abs_error"].to_numpy(), atol=1e-4,
|
||||
expected["mean_abs_error"].to_numpy(),
|
||||
actual["mean_abs_error"].to_numpy(),
|
||||
atol=1e-4,
|
||||
)
|
||||
|
||||
+22
-7
@@ -22,7 +22,10 @@ def test_merge_cli_overrides_applies_file_then_cli(tmp_path, monkeypatch):
|
||||
_write_config(path, "abc123")
|
||||
|
||||
cfg = gconfig.merge_cli_overrides(
|
||||
gconfig.DEFAULT_CONFIG, path, train_overrides={}, model_overrides={"hidden_dim": 128},
|
||||
gconfig.DEFAULT_CONFIG,
|
||||
path,
|
||||
train_overrides={},
|
||||
model_overrides={"hidden_dim": 128},
|
||||
)
|
||||
assert cfg["train"]["epochs"] == 5 # from file
|
||||
assert cfg["model"]["hidden_dim"] == 128 # CLI override wins over file
|
||||
@@ -42,7 +45,9 @@ def test_merge_cli_overrides_warns_on_git_hash_mismatch(tmp_path, monkeypatch, c
|
||||
assert "current999" in captured.err
|
||||
|
||||
|
||||
def test_merge_cli_overrides_no_warning_on_matching_git_hash(tmp_path, monkeypatch, capsys):
|
||||
def test_merge_cli_overrides_no_warning_on_matching_git_hash(
|
||||
tmp_path, monkeypatch, capsys
|
||||
):
|
||||
monkeypatch.setattr(gconfig, "git_hash", lambda: "same123")
|
||||
path = tmp_path / "config.toml"
|
||||
_write_config(path, "same123")
|
||||
@@ -51,7 +56,9 @@ def test_merge_cli_overrides_no_warning_on_matching_git_hash(tmp_path, monkeypat
|
||||
assert capsys.readouterr().err == ""
|
||||
|
||||
|
||||
def test_merge_cli_overrides_no_warning_when_git_hash_unknown(tmp_path, monkeypatch, capsys):
|
||||
def test_merge_cli_overrides_no_warning_when_git_hash_unknown(
|
||||
tmp_path, monkeypatch, capsys
|
||||
):
|
||||
monkeypatch.setattr(gconfig, "git_hash", lambda: "unknown")
|
||||
path = tmp_path / "config.toml"
|
||||
_write_config(path, "abc123")
|
||||
@@ -60,7 +67,9 @@ def test_merge_cli_overrides_no_warning_when_git_hash_unknown(tmp_path, monkeypa
|
||||
assert capsys.readouterr().err == ""
|
||||
|
||||
|
||||
def test_merge_cli_overrides_no_warning_when_meta_section_absent(tmp_path, monkeypatch, capsys):
|
||||
def test_merge_cli_overrides_no_warning_when_meta_section_absent(
|
||||
tmp_path, monkeypatch, capsys
|
||||
):
|
||||
monkeypatch.setattr(gconfig, "git_hash", lambda: "current999")
|
||||
path = tmp_path / "config.toml"
|
||||
path.write_text("[train]\nepochs = 5\n")
|
||||
@@ -69,7 +78,9 @@ def test_merge_cli_overrides_no_warning_when_meta_section_absent(tmp_path, monke
|
||||
assert capsys.readouterr().err == ""
|
||||
|
||||
|
||||
def test_warn_if_checkpoint_config_mismatch_finds_sibling_toml(tmp_path, monkeypatch, capsys):
|
||||
def test_warn_if_checkpoint_config_mismatch_finds_sibling_toml(
|
||||
tmp_path, monkeypatch, capsys
|
||||
):
|
||||
monkeypatch.setattr(gconfig, "git_hash", lambda: "current999")
|
||||
ckpt_path = tmp_path / "best.pt"
|
||||
ckpt_path.write_bytes(b"") # contents irrelevant, only its directory is used
|
||||
@@ -83,7 +94,9 @@ def test_warn_if_checkpoint_config_mismatch_finds_sibling_toml(tmp_path, monkeyp
|
||||
assert "current999" in captured.err
|
||||
|
||||
|
||||
def test_warn_if_checkpoint_config_mismatch_no_warning_when_toml_absent(tmp_path, monkeypatch, capsys):
|
||||
def test_warn_if_checkpoint_config_mismatch_no_warning_when_toml_absent(
|
||||
tmp_path, monkeypatch, capsys
|
||||
):
|
||||
monkeypatch.setattr(gconfig, "git_hash", lambda: "current999")
|
||||
ckpt_path = tmp_path / "best.pt"
|
||||
ckpt_path.write_bytes(b"")
|
||||
@@ -92,7 +105,9 @@ def test_warn_if_checkpoint_config_mismatch_no_warning_when_toml_absent(tmp_path
|
||||
assert capsys.readouterr().err == ""
|
||||
|
||||
|
||||
def test_warn_if_checkpoint_config_mismatch_no_warning_when_hashes_match(tmp_path, monkeypatch, capsys):
|
||||
def test_warn_if_checkpoint_config_mismatch_no_warning_when_hashes_match(
|
||||
tmp_path, monkeypatch, capsys
|
||||
):
|
||||
monkeypatch.setattr(gconfig, "git_hash", lambda: "same123")
|
||||
ckpt_path = tmp_path / "best.pt"
|
||||
ckpt_path.write_bytes(b"")
|
||||
|
||||
@@ -26,13 +26,17 @@ def test_dataset_item_shapes():
|
||||
|
||||
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)
|
||||
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)
|
||||
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
|
||||
|
||||
@@ -48,7 +52,9 @@ def test_split_event_leakage():
|
||||
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)
|
||||
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)
|
||||
|
||||
@@ -20,10 +20,13 @@ def test_denoising_mlp_output_shape():
|
||||
x_t = torch.randn(B, 9)
|
||||
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)
|
||||
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, 9)
|
||||
|
||||
|
||||
@@ -81,10 +81,14 @@ def test_reconstruct_post_pos_straight_line():
|
||||
step_length = rng.uniform(0.1, 5.0, size=N).astype(np.float32)
|
||||
post_pos = pre_pos + step_length[:, None] * pre_dir
|
||||
|
||||
travel_dir_local = local_frame_rotation(pre_dir, travel_direction(pre_pos, post_pos))
|
||||
travel_dir_local = local_frame_rotation(
|
||||
pre_dir, travel_direction(pre_pos, post_pos)
|
||||
)
|
||||
np.testing.assert_allclose(travel_dir_local, np.tile([0, 0, 1], (N, 1)), atol=1e-4)
|
||||
|
||||
reconstructed = reconstruct_post_pos(pre_pos, pre_dir, step_length, travel_dir_local)
|
||||
reconstructed = reconstruct_post_pos(
|
||||
pre_pos, pre_dir, step_length, travel_dir_local
|
||||
)
|
||||
np.testing.assert_allclose(reconstructed, post_pos, atol=1e-4)
|
||||
|
||||
|
||||
@@ -98,8 +102,12 @@ def test_reconstruct_post_pos_general_roundtrip():
|
||||
post_pos = pre_pos + rng.standard_normal((N, 3)).astype(np.float32)
|
||||
step_length = np.linalg.norm(post_pos - pre_pos, axis=1).astype(np.float32)
|
||||
|
||||
travel_dir_local = local_frame_rotation(pre_dir, travel_direction(pre_pos, post_pos))
|
||||
reconstructed = reconstruct_post_pos(pre_pos, pre_dir, step_length, travel_dir_local)
|
||||
travel_dir_local = local_frame_rotation(
|
||||
pre_dir, travel_direction(pre_pos, post_pos)
|
||||
)
|
||||
reconstructed = reconstruct_post_pos(
|
||||
pre_pos, pre_dir, step_length, travel_dir_local
|
||||
)
|
||||
np.testing.assert_allclose(reconstructed, post_pos, atol=1e-4)
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user