Add rotation-free S-space adapter cores (antipasto family)

Replace antipasto's rotation/Cayley with a bounded 1+ELU gain and split the
S-space idea into four interpretable PiSSA-style cores (frozen U/S/Vh, small
trainable core):

- antipasto: S_eff = S*(1+ELU(coeff*g)). exp-bounded attenuation, linear
  amplification (constant gradient, no runaway). g=0 -> exact identity.
- antipasto_rot: keeps the block-Cayley rotation as a separate variant for
  cost comparison (its per-forward solve is the 72ms vs 36ms gap).
- antipasto_ablate: contractive (I - a c c^T) diag(S), eigenvalues in [0,1],
  cannot blow up. Optional cov_orient (CorDA) basis.
- antipasto_corda: covariance-oriented oblique projector P = Vh C^{-1/2}, the
  data-energy basis rather than the weight-gain basis. 1+ELU gain.

Add scripts/_cost.py + scripts/cost_report.py: one-row-per-variant cost table
(trainable params, peak GPU mem, fwd/bwd ms, added MACs/tok, group_init ms).
Wire all four into the benchmark, smoke test, and __init__ exports.

External review (DeepSeek-v4-pro, docs/reviews/) verified the math; acted on
its one real point (corda g now inits to zeros for exact identity).

Co-Authored-By: Claudypoo <noreply@anthropic.com>
This commit is contained in:
wassname
2026-06-14 19:12:27 +08:00
co-authored by Claudypoo
parent e5048fcaff
commit b80d7778af
11 changed files with 1059 additions and 107 deletions
+131
View File
@@ -0,0 +1,131 @@
"""Measure the cost of an attached adapter: params, FLOPs/MACs, time, GPU mem.
Which metric is "best" for comparing adapters? They answer different questions:
- trainable_params -- deterministic "size" number. The headline.
- macs_per_token -- deterministic, hardware-INDEPENDENT compute. Best for an
apples-to-apples comparison: wall-time is noisy and the old
rotation adapter paid a per-forward Cayley solve the new ones
do not. "adds" (additions) ~= MACs; FLOPs ~= 2 * MACs.
- fwd_ms / bwd_ms -- felt cost, but noisy: warmup + median over `iters`, never one run.
- peak_gpu_mb -- resident + activation peak around fwd(+bwd).
FLOPs come from torch.utils.flop_counter.FlopCounterMode (built in, no new dep). Its
convention is MACs (a (m,k)@(k,n) matmul counts as m*n*k); we expose both `flops`
(as returned) and `macs_per_token = flops / n_tokens` -- calibrate once on a known
matmul if you need to be sure of the factor of 2.
"""
from __future__ import annotations
import statistics
import time
import torch
from torch.utils.flop_counter import FlopCounterMode
def _time_call(fn, warmup: int, iters: int, cuda: bool) -> float:
"""Median wall-time of fn() in milliseconds (warmup excluded)."""
for _ in range(warmup):
fn()
if cuda:
torch.cuda.synchronize()
samples = []
for _ in range(iters):
if cuda:
start = torch.cuda.Event(enable_timing=True)
end = torch.cuda.Event(enable_timing=True)
start.record()
fn()
end.record()
torch.cuda.synchronize()
samples.append(start.elapsed_time(end))
else:
t0 = time.perf_counter()
fn()
samples.append((time.perf_counter() - t0) * 1e3)
return statistics.median(samples)
def measure_cost(
model: torch.nn.Module,
fwd_fn,
*,
bwd_step_fn=None,
n_tokens: int | None = None,
adapter_filter: str = "lora_",
warmup: int = 3,
iters: int = 10,
) -> dict:
"""Cost of the currently-attached adapter.
fwd_fn(): run one forward (no grad). Used for FLOPs + fwd timing.
bwd_step_fn(): zero_grad + forward + loss.backward(). Used for bwd timing.
n_tokens: tokens in the fwd_fn batch, for macs_per_token.
adapter_filter: substring marking adapter params/buffers (default 'lora_').
"""
dev = next(model.parameters()).device
cuda = dev.type == "cuda"
trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)
named = list(model.named_parameters()) + list(model.named_buffers())
adapter_bytes = sum(t.numel() * t.element_size() for n, t in named if adapter_filter in n)
# FLOPs: one forward under the counter (no grad so we count inference cost).
# FlopCounterMode can assert on some fused attention shapes; degrade to None.
try:
fc = FlopCounterMode(display=False)
with torch.no_grad(), fc:
fwd_fn()
flops = fc.get_total_flops()
except Exception as e:
print(f" [warn] FLOP count failed ({type(e).__name__}: {e}); flops=None")
flops = None
if cuda:
torch.cuda.synchronize()
torch.cuda.reset_peak_memory_stats()
fwd_ms = _time_call(lambda: _no_grad(fwd_fn), warmup, iters, cuda)
bwd_ms = _time_call(bwd_step_fn, warmup, iters, cuda) if bwd_step_fn is not None else None
peak_gpu_mb = (torch.cuda.max_memory_allocated() / 1e6) if cuda else None
return dict(
trainable_params=trainable_params,
adapter_resident_mb=adapter_bytes / 1e6,
flops=flops,
macs_per_token=(flops / n_tokens) if (flops and n_tokens) else None,
fwd_ms=fwd_ms,
bwd_ms=bwd_ms,
peak_gpu_mb=peak_gpu_mb,
)
def _no_grad(fn):
with torch.no_grad():
return fn()
class group_init_meter:
"""Context manager: wall-time + peak CPU RAM of a group_init / attach-with-calib.
CorDA accumulates C = E[xx^T] on CPU and runs eigh(d_in^3) -- the expensive corner.
Use around ll.attach(model, cfg, calibration_data=...) to log that asymmetry.
"""
def __init__(self):
self.ms = None
self.peak_cpu_mb = None
def __enter__(self):
import tracemalloc
self._tm = tracemalloc
tracemalloc.start()
self._t0 = time.perf_counter()
return self
def __exit__(self, *exc):
self.ms = (time.perf_counter() - self._t0) * 1e3
_, peak = self._tm.get_traced_memory()
self._tm.stop()
self.peak_cpu_mb = peak / 1e6
return False
+142
View File
@@ -0,0 +1,142 @@
"""One-row-per-variant cost table: params, MACs/token, fwd/bwd ms, peak GPU, group_init.
Answers "which is best -- time / flops / adds / params?": MACs/token is the
deterministic apples-to-apples compute number; trainable_params is the size headline;
wall-time is the felt-but-noisy number; group_init is where CorDA's eigh(d_in^3) bites.
Usage:
uv run --extra benchmark python scripts/cost_report.py \
--model Qwen/Qwen3-0.6B-Base --variants antipasto antipasto_corda antipasto_ablate lora \
--target-name 'q_proj$' 'v_proj$' --r 32 --out logs/cost_qwen0.6b.log
Point --target-name at down_proj to see the CorDA covariance corner (large d_in).
"""
from __future__ import annotations
import argparse
import importlib.util
import sys
from pathlib import Path
import torch
from tabulate import tabulate
import lora_lite as ll
_HERE = Path(__file__).resolve().parent
_BENCH = importlib.util.spec_from_file_location("metamath_benchmark", _HERE / "metamath_gsm8k_benchmark.py")
benchmark = importlib.util.module_from_spec(_BENCH)
sys.modules[_BENCH.name] = benchmark
_BENCH.loader.exec_module(benchmark)
_COST = importlib.util.spec_from_file_location("_cost", _HERE / "_cost.py")
cost = importlib.util.module_from_spec(_COST)
sys.modules[_COST.name] = cost
_COST.loader.exec_module(cost)
def build_cfg(variant: str, args, dtype) -> ll.AdapterConfig:
"""Reuse the benchmark's variant->config map; only need r/targets/dtype here."""
bcfg = benchmark.BenchmarkConfig(
model=args.model, variant=variant, r=args.r, alpha=float(args.r),
target_name=list(args.target_name), layers=args.layers, torch_dtype=args.dtype,
antipasto_cov_orient=args.cov_orient,
)
return benchmark.cfg_for_variant(bcfg, dtype)
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--model", default="Qwen/Qwen3-0.6B-Base")
ap.add_argument("--variants", nargs="+",
default=["lora", "antipasto", "antipasto_rot", "antipasto_corda", "antipasto_ablate"])
ap.add_argument("--target-name", nargs="+", default=[r"q_proj$", r"v_proj$"])
ap.add_argument("--r", type=int, default=32)
ap.add_argument("--layers", default="all",
help="'all' or comma list e.g. '0,1' -- limit layers (CorDA down_proj eigh is slow).")
ap.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
ap.add_argument("--dtype", default="bfloat16")
ap.add_argument("--seq-len", type=int, default=256)
ap.add_argument("--batch", type=int, default=2)
ap.add_argument("--calib-batches", type=int, default=4)
ap.add_argument("--cov-orient", action="store_true",
help="CorDA-orient antipasto_ablate (measure the eigh corner).")
ap.add_argument("--out", default="logs/cost.log")
args = ap.parse_args()
dtype = getattr(torch, args.dtype)
# eager attention: FlopCounterMode's sdpa_flop_count asserts on GQA (Qwen3) SDPA
# shapes (q heads != kv heads). eager uses explicit matmuls it can count.
from transformers import AutoModelForCausalLM, AutoTokenizer
tok = AutoTokenizer.from_pretrained(args.model)
model = AutoModelForCausalLM.from_pretrained(
args.model, dtype=dtype, attn_implementation="eager"
).to(args.device)
model.eval()
n_tokens = args.batch * args.seq_len
ids = torch.randint(0, model.config.vocab_size, (args.batch, args.seq_len), device=args.device)
calib = [{"input_ids": torch.randint(0, model.config.vocab_size,
(args.batch, args.seq_len), device=args.device)}
for _ in range(args.calib_batches)]
def fwd():
model(input_ids=ids)
def bwd_step():
model.zero_grad(set_to_none=True)
loss = model(input_ids=ids).logits.float().pow(2).mean()
loss.backward()
# base (no-adapter) cost, so each row can report the adapter's ADDED MACs/token.
base = cost.measure_cost(model, fwd, bwd_step_fn=bwd_step, n_tokens=n_tokens)
base_macs = base["macs_per_token"]
print(f"base (no adapter): MACs/tok={int(base_macs) if base_macs else None} "
f"fwd_ms={round(base['fwd_ms'],2)} bwd_ms={round(base['bwd_ms'],2)}")
# base = no adapter; model params left trainable, so this is the full-finetune
# GPU-mem reference (its backward stores grads for every weight).
total_params = sum(p.numel() for p in model.parameters())
rows = [{
"variant": "base(full-FT)", "train_params": total_params,
"fwd_ms": round(base["fwd_ms"], 2), "bwd_ms": round(base["bwd_ms"], 2),
"peak_GPU_MB": round(base["peak_gpu_mb"], 1) if base["peak_gpu_mb"] else None,
"added_MACs/tok": 0 if base_macs else None,
"ginit_ms": 0.0, "ginit_CPU_MB": 0.0,
}]
for variant in args.variants:
cfg = build_cfg(variant, args, dtype)
# group_init / attach cost (CorDA's eigh + C live here).
with cost.group_init_meter() as gi:
ll.attach(model, cfg, calibration_data=calib)
c = cost.measure_cost(model, fwd, bwd_step_fn=bwd_step, n_tokens=n_tokens)
ll.detach(model)
rows.append({
"variant": variant,
"train_params": c["trainable_params"],
"fwd_ms": round(c["fwd_ms"], 2),
"bwd_ms": round(c["bwd_ms"], 2) if c["bwd_ms"] else None,
"peak_GPU_MB": round(c["peak_gpu_mb"], 1) if c["peak_gpu_mb"] else None,
# flat across same-r adapters; kept only as a sanity check, not a comparator.
"added_MACs/tok": int(c["macs_per_token"] - base_macs) if (c["macs_per_token"] and base_macs) else None,
"ginit_ms": round(gi.ms, 1),
"ginit_CPU_MB": round(gi.peak_cpu_mb, 1),
})
print(f" {variant}: params={rows[-1]['train_params']} "
f"peak_GPU_MB={rows[-1]['peak_GPU_MB']} bwd_ms={rows[-1]['bwd_ms']} ginit_ms={rows[-1]['ginit_ms']}")
table = tabulate(rows, headers="keys", tablefmt="pipe")
header = (f"# cost report: {args.model} targets={args.target_name} r={args.r} "
f"seq={args.seq_len} batch={args.batch} dtype={args.dtype}\n"
f"# COMPARATORS: train_params, peak_GPU_MB (fwd+bwd, process-local max), bwd_ms, ginit_ms.\n"
f"# added_MACs/tok is flat across same-r adapters (sanity check only).\n"
f"# ginit_CPU_MB undercounts: tracemalloc misses torch C++ tensor allocs (the CorDA C matrix).\n")
out_path = Path(args.out)
out_path.parent.mkdir(parents=True, exist_ok=True)
out_path.write_text(header + table + "\n")
print("\n" + header + table)
print(f"\nsaved -> {out_path}")
if __name__ == "__main__":
main()
+20 -3
View File
@@ -34,6 +34,9 @@ CFG_BY_VARIANT = {
"hra": ll.HRAConfig,
"eva": ll.EVAConfig,
"antipasto": ll.AntiPaSTOConfig,
"antipasto_rot": ll.AntiPaSTORotConfig,
"antipasto_ablate": ll.AntiPaSTOAblateConfig,
"antipasto_corda": ll.AntiPaSTOCorDAConfig,
"road": ll.RoadConfig,
}
@@ -43,7 +46,7 @@ class BenchmarkConfig:
"""MetaMathQA -> GSM8K benchmark config. Tyro turns this into the CLI."""
model: str = "Qwen/Qwen3-0.6B-Base"
variant: Literal["lora", "pissa", "delora", "ia3", "ia3_ff", "dora", "hra", "eva", "antipasto", "road"] = "lora"
variant: Literal["lora", "pissa", "delora", "ia3", "ia3_ff", "dora", "hra", "eva", "antipasto", "antipasto_rot", "antipasto_ablate", "antipasto_corda", "road"] = "lora"
mode: Literal["benchmark", "probe"] = "benchmark"
device: str = "cuda"
torch_dtype: str = "bfloat16"
@@ -52,6 +55,13 @@ class BenchmarkConfig:
alpha: float = 64.0
delora_lambda0: float = 0.1
road_group_size: int = 64
# AntiPaSTO family (gain / corda) runtime knobs.
antipasto_coeff: float = 1.0
antipasto_suppress_only: bool = False
# AntiPaSTO-ablate.
antipasto_ablate_k: int = 1
antipasto_cov_orient: bool = False
# AntiPaSTO-rot (legacy rotation variant) basis to rotate.
antipasto_rotate_basis: Literal["V", "U", "none"] = "V"
target_name: list[str] = field(default_factory=lambda: list(DEFAULT_TARGETS))
layers: str = "all"
@@ -124,8 +134,15 @@ def cfg_for_variant(args: BenchmarkConfig, dtype: torch.dtype) -> ll.AdapterConf
extra = {"lambda0": args.delora_lambda0} if args.variant == "delora" else {}
if args.variant == "road":
extra = {"group_size": args.road_group_size}
if args.variant == "antipasto":
if args.variant == "antipasto_rot":
extra = {"rotate_basis": args.antipasto_rotate_basis}
if args.variant == "antipasto":
extra = {"coeff": args.antipasto_coeff, "suppress_only": args.antipasto_suppress_only}
if args.variant == "antipasto_corda":
extra = {"coeff": args.antipasto_coeff, "suppress_only": args.antipasto_suppress_only}
if args.variant == "antipasto_ablate":
extra = {"coeff": args.antipasto_coeff, "k": args.antipasto_ablate_k,
"cov_orient": args.antipasto_cov_orient}
return CFG_BY_VARIANT[args.variant](
r=args.r,
alpha=args.r if args.variant == "pissa" else args.alpha,
@@ -155,7 +172,7 @@ def count_base_grad_leaks(model: torch.nn.Module) -> int:
def perturb_first_adapter(model: torch.nn.Module) -> None:
priority = ("lora_B", "lora_g", "lora_U", "lora_A", "lora_lambda", "lora_gate", "lora_delta_s", "lora_rot_T", "lora_m", "lora_road_theta", "lora_road_alpha")
priority = ("lora_B", "lora_g", "lora_c", "lora_alpha", "lora_U", "lora_A", "lora_lambda", "lora_gate", "lora_delta_s", "lora_m", "lora_road_theta", "lora_road_alpha")
for key in priority:
for _, p in model.named_parameters():
if p.requires_grad and key in _: