diff --git a/src/wwgpt/ww.py b/src/wwgpt/ww.py index 95418d7..d1e1098 100644 --- a/src/wwgpt/ww.py +++ b/src/wwgpt/ww.py @@ -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}) diff --git a/tests/test_external_wwpgd_pass3.py b/tests/test_external_wwpgd_pass3.py index 176bb46..92f8bd7 100644 --- a/tests/test_external_wwpgd_pass3.py +++ b/tests/test_external_wwpgd_pass3.py @@ -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