add dW magnitude/direction ablation eval

Constructs four variants of a trained dW and evaluates each on daily
dilemmas at coeffs {-1, 0, +1}:
  full         original (control)
  dir_only     elementwise direction preserved, all tensors rescaled
               to a common Frobenius norm (flattens per-tensor magnitude)
  mag_only     random direction per tensor, original per-tensor norm
               (preserves which layers/modules carry the load)
  random_norm  random direction + common norm (control)

Tests whether the trained behavior is carried by element direction or
by the per-tensor magnitude pattern. Default adapter is delora since
it has the largest raw dd_delta and the worst SI -- which factor is
load-bearing?

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
This commit is contained in:
wassname
2026-04-28 08:31:24 +08:00
co-authored by Claude Opus 4.7
parent 64adf9267d
commit 19bc3edb2e
+159
View File
@@ -0,0 +1,159 @@
"""DeLoRA magnitude vs direction ablation.
Question: is the trained dW useful because of (a) its element-wise direction
or (b) the per-tensor magnitude pattern (which layers / modules get bigger
updates)? Constructs three variants:
full original dW (control)
dir_only dW with all tensors rescaled to a common Frobenius norm
(preserves elementwise direction; flattens the per-tensor
magnitude pattern)
mag_only each tensor replaced by a Gaussian random tensor scaled to
the original tensor's norm (preserves the per-tensor
magnitude pattern; randomises the direction)
random_norm Gaussian random tensors all rescaled to a common norm
(control: neither direction nor magnitude pattern)
Eval all four on daily-dilemmas (full 219 split) at coeffs {-1, 0, +1}
and dump dilemmas_per_row.csv so SI can be recomputed offline.
"""
from __future__ import annotations
from dataclasses import dataclass
from pathlib import Path
import polars as pl
import torch
import tyro
from loguru import logger
from tabulate import tabulate
from transformers import AutoModelForCausalLM, AutoTokenizer
from ws._log import final_summary, get_argv, setup_logging
from ws.diff import DIFF_FILENAME, load_diff
from ws.eval.dilemmas import DilemmasCfg, compute_full_metrics, evaluate
@dataclass
class DWDecompCfg:
model: str = "Qwen/Qwen3-0.6B"
behavior: str = "honesty"
adapter: str = "delora"
coeffs: tuple[float, ...] = (-1.0, 0.0, 1.0)
n_dilemmas: int = 219
batch_size: int = 8
out: Path = Path("out")
seed: int = 0
def _frob_norms(w: dict[str, torch.Tensor]) -> dict[str, float]:
return {k: float(t.float().norm().item()) for k, t in w.items()}
def make_variants(w: dict[str, torch.Tensor], seed: int = 0) -> dict[str, dict[str, torch.Tensor]]:
g = torch.Generator(device="cpu").manual_seed(seed)
norms = _frob_norms(w)
n_tensors = len(w)
total_sq = sum(n ** 2 for n in norms.values())
common_norm = (total_sq / n_tensors) ** 0.5 # equal-norm such that sum-of-squares matches
full = {k: t.clone() for k, t in w.items()}
# dir_only: same elementwise direction, all tensors rescaled to common_norm
dir_only = {}
for k, t in w.items():
n = norms[k]
if n == 0:
dir_only[k] = t.clone()
else:
dir_only[k] = (t.float() * (common_norm / n)).to(t.dtype)
# mag_only: random direction, original per-tensor norm
mag_only = {}
for k, t in w.items():
r = torch.randn(t.shape, generator=g, dtype=torch.float32)
rn = float(r.norm().item())
mag_only[k] = (r * (norms[k] / rn)).to(t.dtype)
# random_norm: random direction, common per-tensor norm
random_norm = {}
for k, t in w.items():
r = torch.randn(t.shape, generator=g, dtype=torch.float32)
rn = float(r.norm().item())
random_norm[k] = (r * (common_norm / rn)).to(t.dtype)
return {"full": full, "dir_only": dir_only, "mag_only": mag_only, "random_norm": random_norm}
def main(cfg: DWDecompCfg) -> None:
setup_logging("dw_decomp_ablation")
out_dir = cfg.out / cfg.behavior / "dw_decomp_ablation" / cfg.adapter
out_dir.mkdir(parents=True, exist_ok=True)
tok = AutoTokenizer.from_pretrained(cfg.model)
if tok.pad_token is None:
tok.pad_token = tok.eos_token
model = AutoModelForCausalLM.from_pretrained(cfg.model, torch_dtype=torch.bfloat16, device_map="auto")
model.eval()
w_path = cfg.out / cfg.behavior / cfg.adapter / DIFF_FILENAME
w = load_diff(w_path)
logger.info(f"loaded {len(w)} tensors from {w_path}; total ||dW||_F={sum(t.float().norm().item()**2 for t in w.values())**0.5:.3f}")
variants = make_variants(w, seed=cfg.seed)
parts = []
dcfg = DilemmasCfg(
model_id=cfg.model,
coeffs=cfg.coeffs,
n_dilemmas=cfg.n_dilemmas,
batch_size=cfg.batch_size,
)
for name, ww in variants.items():
logger.info(f"variant={name}: total ||W||_F={sum(t.float().norm().item()**2 for t in ww.values())**0.5:.3f}")
df = evaluate(dcfg, ww, model=model, tok=tok).with_columns(pl.lit(name).alias("variant"))
parts.append(df)
per_row = pl.concat(parts)
per_row_path = out_dir / "dilemmas_per_row.csv"
per_row.write_csv(per_row_path)
rows = []
for name in variants:
sub = per_row.filter(pl.col("variant") == name)
m = compute_full_metrics(sub)
zero = float(sub.filter(pl.col("coeff") == 0.0)["logratio_honesty"].mean())
pos_mean = float(sub.filter(pl.col("coeff") == 1.0)["logratio_honesty"].mean())
neg_mean = float(sub.filter(pl.col("coeff") == -1.0)["logratio_honesty"].mean())
rows.append({
"variant": name,
"SI": m.get("surgical_informedness", float("nan")),
"si_fwd": m.get("si_fwd", float("nan")),
"si_rev": m.get("si_rev", float("nan")),
"fix_fwd": m.get("fix_fwd", -1),
"broke_fwd": m.get("broke_fwd", -1),
"flip_rev": m.get("flip_rev", -1),
"counter_rev": m.get("counter_rev", -1),
"lr_at_zero": zero,
"delta_pos": pos_mean - zero,
"delta_neg": neg_mean - zero,
})
summary = pl.DataFrame(rows).sort("SI", descending=True, nulls_last=True)
summary_path = out_dir / "summary.csv"
summary.write_csv(summary_path)
print("\nDeLoRA dW magnitude/direction ablation (honesty axis, daily dilemmas)")
print(tabulate(summary.to_pandas(), headers="keys", tablefmt="tsv", floatfmt="+.3f", showindex=False))
final_summary(
out=summary_path,
argv=get_argv(),
main_metric=f"best_variant={summary['variant'][0]} SI={float(summary['SI'][0] or 0):+.3f}",
cue="🟢",
table_rows=summary.select("variant", "SI", "si_fwd", "si_rev", "delta_pos", "delta_neg").rows(),
headers=["variant", "SI", "si_fwd", "si_rev", "delta_pos", "delta_neg"],
floatfmt="",
)
if __name__ == "__main__":
main(tyro.cli(DWDecompCfg))