Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions src/wwgpt/ww.py
Original file line number Diff line number Diff line change
Expand Up @@ -708,6 +708,13 @@ def apply_external_wwpgd(
"combined_hardness_requested": req_h, "combined_hardness_applied": req_h * scale, "trust_region_limit": limit, "trust_region_scale": scale,
"relative_frobenius_change_requested": req_rel, "relative_frobenius_change_applied": app_rel, "relative_frobenius_change": app_rel,
"relative_frobenius_weight_change": app_rel, "changed": changed, "projection_attempted": req_h > 0.0, "projected": changed})
# Dose identity for cross-run / cross-actuator comparison (logging only).
# Definition: Frobenius ||W_after - W_before|| / ||W_before|| of the applied update.
rows[-1].update({
"dose_definition": "relative_frobenius_change_applied",
"dose_relative_frobenius": app_rel,
"is_first_projection_event": int(event_index) == 1,
})
rows[-1].update({"displacement_kind": displacement_kind,
"real_candidate_displacement_cosine": displacement_cosine,
"sham_seed": int(sham_seed) if sham_seed is not None else None})
Expand Down
13 changes: 13 additions & 0 deletions tests/test_external_wwpgd_pass3.py
Original file line number Diff line number Diff line change
Expand Up @@ -106,6 +106,19 @@ def test_no_strength_multiplier_external_blend_eta(monkeypatch):
assert all(row["blend_eta"] == 0.5 for row in rows)


def test_projection_rows_include_dose_identity_fields(monkeypatch):
"""Every projection row names the dose metric used for strength comparison."""
install_fake_ww_pgd(monkeypatch, [])
rows = apply_external_wwpgd(tiny_model(), actual_step=1, event_index=1)
assert rows
for row in rows:
assert row["dose_definition"] == "relative_frobenius_change_applied"
assert row["dose_relative_frobenius"] == row["relative_frobenius_change_applied"]
assert row["is_first_projection_event"] is True
rows2 = apply_external_wwpgd(tiny_model(), actual_step=2, event_index=2)
assert all(row["is_first_projection_event"] is False for row in rows2)


def test_projection_interval_accepts_positive_integers():
ext = WWPGDExtension(cfg=WWPGDConfig(), interval=3)
assert ext.interval == 3
Expand Down