etomoscow's picture
download
raw
31.6 kB
"""E1 — Single-language XNLI head-to-head on Llama-3.1-8B-base.
Methods (9 baseline + 5 diagnostic; KaSA and LoRA-Dash skipped — no open
implementation available):
Baseline: random, pissa, milora, lora_ga, eva, dora, filet, mff_top, mff_energy
Diagnostic: overscaled_random, kaiming_random, peft_scaled_random, unit_random,
random_orthogonal, peft_default
Ranks: 4, 8, 16, 32
Seeds: 42, 1337, 2024
Target modules: q_proj, v_proj, down_proj
Train: 8192 XNLI English examples (MultiNLI English train subset), 3 epochs
Eval: full XNLI English validation (2490 examples)
Online factors (FILet, EVA, LoRA-GA) are collected once from the base model
and cached under {output_dir}/.factor_cache/ — reused across rank and seed
sweeps.
MFF factors are loaded lazily from:
{mff_factor_dir}/en/{module}/{layer_name_escaped}.safetensors
Supports resume: existing (method, rank, seed) triples are skipped.
Usage:
# Full sweep on one GPU:
CUDA_VISIBLE_DEVICES=0 \\
TRITON_CACHE_DIR=/tmp/triton_cache \\
python experiments/e1_single_language/run.py
# Parallel: split by method across GPUs:
CUDA_VISIBLE_DEVICES=0 python experiments/e1_single_language/run.py \\
--methods random pissa milora lora_ga
CUDA_VISIBLE_DEVICES=1 python experiments/e1_single_language/run.py \\
--methods eva dora filet mff_top mff_energy
"""
from __future__ import annotations
import argparse
import json
import os
import sys
from pathlib import Path
import numpy as np
import torch
from torch import nn
from tqdm.auto import tqdm
REPO = Path(__file__).resolve().parents[2]
SRC = REPO / "src"
if str(SRC) not in sys.path:
sys.path.insert(0, str(SRC))
KRONLINGUA = Path(os.environ.get("MFFLORA_KRONLINGUA_ROOT", "external/kronlingua"))
if str(KRONLINGUA) not in sys.path:
sys.path.insert(0, str(KRONLINGUA))
import mfflora._compat # noqa: F401,E402
from mfflora.factors import topk_eigvecs # noqa: E402
from mfflora.init import ( # noqa: E402
kaiming_random_init,
milora_init,
mff_lora_init,
overscaled_random_init,
pissa_init,
fisher_peft_init,
fisher_peft_low_init,
fisher_peft_bi_init,
fisher_orthogonal_init,
fws_safe_init,
)
from mfflora.init.random import random_init # backward compat alias noqa: E402
from mfflora.init.eva import eva_init # noqa: E402
from mfflora.init.filet import filet_init # noqa: E402
from mfflora.init.lora_ga import lora_ga_init # noqa: E402
from mfflora.train.lora_runner import inject_lora_init # noqa: E402
from mfflora.utils import set_seed # noqa: E402
ALL_METHODS = (
"peft_default",
"overscaled_random",
"kaiming_random",
"peft_scaled_random",
"unit_random",
"random_orthogonal",
"random",
"pissa",
"milora",
"lora_ga",
"eva",
"dora",
"filet",
"mff_top",
"mff_energy",
"fpeft",
"fpeft_low",
"fpeft_bi_high",
"fpeft_bi_low",
"foi",
"fws_safe",
)
ALL_RANKS = (4, 8, 16, 32)
DEFAULT_SEEDS = (42, 1337, 2024)
MODEL_PATH = os.environ.get("MFFLORA_MODEL_PATH", "unsloth/Llama-3.1-8B")
# Pre-computed K-FAC factors for en (Task 3.0 output)
MFF_FACTOR_DIR_DEFAULT = os.environ.get(
"MFFLORA_FACTOR_DIR", "outputs/e1_0_additional_factors"
)
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser()
p.add_argument("--model-path", default=MODEL_PATH)
p.add_argument("--language", default="en")
p.add_argument("--train-size", type=int, default=8192)
p.add_argument("--ranks", nargs="+", type=int, default=list(ALL_RANKS))
p.add_argument("--seeds", nargs="+", type=int, default=list(DEFAULT_SEEDS))
p.add_argument("--methods", nargs="+", default=list(ALL_METHODS))
p.add_argument("--epochs", type=int, default=3)
p.add_argument("--lr", type=float, default=1e-4)
p.add_argument(
"--lora-alpha-mults",
nargs="+",
type=float,
default=[2.0],
help="One or more global multipliers for LoRA alpha (lora_alpha = rank * multiplier).",
)
p.add_argument("--batch-size", type=int, default=8)
p.add_argument("--grad-accum", type=int, default=2)
p.add_argument(
"--precision",
choices=["bf16", "fp16", "fp32"],
default="bf16",
help="Model/training precision. Use fp16 on devices with limited VRAM.",
)
p.add_argument(
"--gradient-checkpointing",
action="store_true",
help="Enable gradient checkpointing for lower VRAM use.",
)
p.add_argument("--max-length", type=int, default=256)
p.add_argument("--target-modules", nargs="+", default=["q_proj", "v_proj", "down_proj"])
p.add_argument("--n-factor-batches", type=int, default=32,
help="Batches used for online factor collection (FILet/EVA/LoRA-GA).")
p.add_argument("--factor-batch-size", type=int, default=4)
p.add_argument("--mff-factor-dir", default=MFF_FACTOR_DIR_DEFAULT)
p.add_argument("--output-dir", default="outputs/e1_single_language")
p.add_argument("--no-wandb", action="store_true")
return p.parse_args()
# ---------------------------------------------------------------------------
# Data
# ---------------------------------------------------------------------------
def _load_xnli(train_size: int, max_length: int, model_path: str):
from datasets import load_dataset
from transformers import AutoTokenizer
tok = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
if tok.pad_token is None:
tok.pad_token = tok.eos_token
tok.padding_side = "right"
train = (
load_dataset("xnli", "en", split="train")
.shuffle(seed=42)
.select(range(train_size))
)
val = load_dataset("xnli", "en", split="validation")
def _tok(b):
return tok(
b["premise"],
b["hypothesis"],
truncation=True,
max_length=max_length,
padding=False,
)
# Drop everything except the label column; tokenizer adds input_ids etc.
keep = {"label"}
train_tok = train.map(_tok, batched=True, remove_columns=[c for c in train.column_names if c not in keep])
val_tok = val.map(_tok, batched=True, remove_columns=[c for c in val.column_names if c not in keep])
train_tok = train_tok.rename_column("label", "labels")
val_tok = val_tok.rename_column("label", "labels")
return tok, train_tok, val_tok
def _make_factor_loader(train_tok, tok, batch_size: int):
from torch.utils.data import DataLoader
from transformers import DataCollatorWithPadding
collator = DataCollatorWithPadding(tok)
return DataLoader(train_tok, batch_size=batch_size, shuffle=False, collate_fn=collator)
# ---------------------------------------------------------------------------
# Online factor collection (FILet / EVA / LoRA-GA)
# ---------------------------------------------------------------------------
def _cache_path(cache_dir: Path, factor_type: str, layer_name: str) -> Path:
return cache_dir / factor_type / (layer_name.replace(".", "__") + ".safetensors")
def _collect_and_cache_online_factors(
model: nn.Module,
target_layer_names: list[str],
loader,
n_batches: int,
cache_dir: Path,
factor_types: tuple[str, ...],
) -> None:
"""Collect FILet (SX,SY), EVA (B), LoRA-GA (grad) for all target layers.
Uses a single set of forward+backward passes with multi-layer hooks for
efficiency. Results saved to cache_dir; existing files are skipped.
"""
from safetensors.torch import save_file
need_by_type: dict[str, list[str]] = {ft: [] for ft in factor_types}
for layer_name in target_layer_names:
for ft in factor_types:
if not _cache_path(cache_dir, ft, layer_name).exists():
need_by_type[ft].append(layer_name)
all_needed = sorted({ln for lns in need_by_type.values() for ln in lns})
if not all_needed:
print(" All online factor caches already exist — skipping collection.", flush=True)
return
print(f" Collecting online factors for {len(all_needed)} layers "
f"({factor_types}) over {n_batches} batches …", flush=True)
named_modules = dict(model.named_modules())
device = next(model.parameters()).device
# Initialise accumulators
sx: dict[str, torch.Tensor] = {}
sy: dict[str, torch.Tensor] = {}
grads: dict[str, torch.Tensor] = {}
counts_x: dict[str, int] = {}
counts_y: dict[str, int] = {}
for ln in all_needed:
layer = named_modules[ln]
assert isinstance(layer, nn.Linear), f"{ln} is not nn.Linear"
# FILet needs both sx and sy; EVA needs only sx.
# Use independent ifs so both caches can be written when requested together.
needs_sx = (
("filet" in factor_types and ln in need_by_type.get("filet", []))
or ("eva" in factor_types and ln in need_by_type.get("eva", []))
)
needs_sy = "filet" in factor_types and ln in need_by_type.get("filet", [])
if needs_sx and ln not in sx:
sx[ln] = torch.zeros(layer.in_features, layer.in_features, dtype=torch.float32)
counts_x[ln] = 0
if needs_sy and ln not in sy:
sy[ln] = torch.zeros(layer.out_features, layer.out_features, dtype=torch.float32)
counts_y[ln] = 0
if "lora_ga" in factor_types and ln in need_by_type.get("lora_ga", []):
grads[ln] = torch.zeros_like(layer.weight, dtype=torch.float32)
# Forward hooks (SX for FILet/EVA)
fwd_handles = []
bwd_handles = []
def make_fwd_hook(ln: str):
def hook(_mod, inputs, _out):
x = inputs[0].detach().reshape(-1, inputs[0].shape[-1]).to(torch.float32).cpu()
sx[ln].add_(x.T @ x)
counts_x[ln] = counts_x.get(ln, 0) + x.shape[0]
return hook
def make_bwd_hook(ln: str):
def hook(_mod, _gin, grad_output):
g = grad_output[0].detach().reshape(-1, grad_output[0].shape[-1]).to(torch.float32).cpu()
sy[ln].add_(g.T @ g)
counts_y[ln] = counts_y.get(ln, 0) + g.shape[0]
return hook
for ln in all_needed:
layer = named_modules[ln]
if ln in sx:
fwd_handles.append(layer.register_forward_hook(make_fwd_hook(ln)))
if ln in sy:
bwd_handles.append(layer.register_full_backward_hook(make_bwd_hook(ln)))
if ln in grads:
layer.weight.requires_grad_(True)
# Determine if we need backward at all
need_backward = bool(sy) or bool(grads)
was_training = model.training
if need_backward:
model.train()
else:
model.eval()
try:
for i, batch in enumerate(tqdm(loader, total=n_batches, desc=" factor batches", leave=False)):
if i >= n_batches:
break
batch = {k: (v.to(device) if torch.is_tensor(v) else v) for k, v in batch.items()}
if need_backward:
model.zero_grad(set_to_none=True)
out = model(**batch)
out.loss.backward()
for ln in grads:
layer = named_modules[ln]
if layer.weight.grad is not None:
grads[ln].add_(layer.weight.grad.detach().to(torch.float32))
else:
with torch.no_grad():
model(**batch)
finally:
for h in fwd_handles + bwd_handles:
h.remove()
for ln in grads:
named_modules[ln].weight.requires_grad_(False)
model.train(was_training)
model.zero_grad(set_to_none=True)
# Normalise and save. FILet and EVA share the sx accumulation (both need
# input activation cov), so save both caches when FILet is collected.
for ln in all_needed:
if "filet" in factor_types and ln in need_by_type.get("filet", []):
nx = max(counts_x.get(ln, 1), 1)
ny = max(counts_y.get(ln, 1), 1)
sx_norm = (sx[ln] / nx).to(torch.float32)
sy_norm = (sy[ln] / ny).to(torch.float32)
p = _cache_path(cache_dir, "filet", ln)
p.parent.mkdir(parents=True, exist_ok=True)
save_file({"SX": sx_norm, "SY": sy_norm}, str(p))
# EVA's B = input activation cov = FILet's SX → save for free.
if "eva" in factor_types:
p_eva = _cache_path(cache_dir, "eva", ln)
p_eva.parent.mkdir(parents=True, exist_ok=True)
save_file({"B": sx_norm}, str(p_eva))
elif "eva" in factor_types and ln in need_by_type.get("eva", []):
nx = max(counts_x.get(ln, 1), 1)
out = {"B": (sx[ln] / nx).to(torch.float32)}
p = _cache_path(cache_dir, "eva", ln)
p.parent.mkdir(parents=True, exist_ok=True)
save_file(out, str(p))
if "lora_ga" in factor_types and ln in need_by_type.get("lora_ga", []):
out = {"grad": (grads[ln] / n_batches).to(torch.float32)}
p = _cache_path(cache_dir, "lora_ga", ln)
p.parent.mkdir(parents=True, exist_ok=True)
save_file(out, str(p))
# ---------------------------------------------------------------------------
# MFF factor path resolver
# ---------------------------------------------------------------------------
def _mff_factor_path(mff_factor_dir: str, language: str, base_layer_name: str) -> Path:
module = base_layer_name.rsplit(".", 1)[-1]
filename = base_layer_name.replace(".", "__") + ".safetensors"
return Path(mff_factor_dir) / language / module / filename
def random_orthogonal_init(
layer_weight: torch.Tensor,
rank: int,
*,
seed: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Zero-preserving LoRA init with an orthonormal random V basis.
This is a PLANv2 control for separating "stable basis" effects from
Fisher-specific alignment. As with standard LoRA, U=0 so ΔW=0 at init.
"""
if rank <= 0:
raise ValueError("rank must be positive")
n, m = layer_weight.shape
dtype = layer_weight.dtype
device = layer_weight.device
gen = torch.Generator(device=device).manual_seed(seed)
basis = torch.randn(m, rank, dtype=torch.float32, device=device, generator=gen)
q, _ = torch.linalg.qr(basis, mode="reduced")
V = q.T.to(dtype=dtype)
U = torch.zeros(n, rank, dtype=dtype, device=device)
return U, V
def unit_random_init(
layer_weight: torch.Tensor,
rank: int,
*,
seed: int,
) -> tuple[torch.Tensor, torch.Tensor]:
"""Zero-preserving random init with each V row normalized to unit norm."""
if rank <= 0:
raise ValueError("rank must be positive")
n, m = layer_weight.shape
dtype = layer_weight.dtype
device = layer_weight.device
gen = torch.Generator(device=device).manual_seed(seed)
V = torch.randn(rank, m, dtype=torch.float32, device=device, generator=gen)
V = V / V.norm(dim=1, keepdim=True).clamp_min(1e-12)
U = torch.zeros(n, rank, dtype=dtype, device=device)
return U, V.to(dtype=dtype)
# ---------------------------------------------------------------------------
# Per-method init function factory
# ---------------------------------------------------------------------------
def _make_init_fn(
method: str,
rank: int,
seed: int,
mff_factor_dir: str,
language: str,
cache_dir: Path,
):
"""Return fn(peft_module_name, base_linear) → {U, V, residual} or None for DoRA."""
def fn(peft_name: str, base_linear: nn.Linear) -> dict | None:
# Strip 'base_model.model.' prefix to get the base model layer name
base_ln = peft_name.removeprefix("base_model.model.")
W = base_linear.weight.detach()
if method in ("random", "custom_random", "overscaled_random"):
U, V = overscaled_random_init(W, rank, seed=seed)
return {"U": U, "V": V, "residual": None}
if method == "kaiming_random":
U, V = kaiming_random_init(W, rank, seed=seed)
return {"U": U, "V": V, "residual": None}
if method == "random_orthogonal":
U, V = random_orthogonal_init(W, rank, seed=seed)
return {"U": U, "V": V, "residual": None}
if method == "peft_scaled_random":
U, V = kaiming_random_init(W, rank, seed=seed)
return {"U": U, "V": V, "residual": None}
if method == "unit_random":
U, V = unit_random_init(W, rank, seed=seed)
return {"U": U, "V": V, "residual": None}
if method == "pissa":
r = pissa_init(W, rank)
return {"U": r.U, "V": r.V, "residual": r.residual}
if method == "milora":
r = milora_init(W, rank)
return {"U": r.U, "V": r.V, "residual": r.residual}
if method == "lora_ga":
from safetensors.torch import load_file
p = _cache_path(cache_dir, "lora_ga", base_ln)
data = load_file(str(p))
# Move gradient to W.device for GPU-accelerated SVD (CPU SVD of
# [4096,14336] takes ~minutes; GPU SVD is ~0.1s per layer).
grad = data["grad"].to(device=W.device, dtype=W.dtype)
U, V = lora_ga_init(W, rank, gradient=grad)
return {"U": U, "V": V, "residual": None}
if method == "eva":
from safetensors.torch import load_file
p_eva = _cache_path(cache_dir, "eva", base_ln)
p_filet = _cache_path(cache_dir, "filet", base_ln)
if p_eva.exists():
B = load_file(str(p_eva))["B"].to(torch.float32)
elif p_filet.exists():
# EVA's B = input activation cov = FILet's SX.
B = load_file(str(p_filet))["SX"].to(torch.float32)
else:
raise FileNotFoundError(f"No EVA or FILet cache for {base_ln}")
m = W.shape[1]
if m > 4096:
eigvecs, _ = topk_eigvecs(B, k=rank, ascending=False, method="randomized")
V = eigvecs.T.to(W.dtype)
U = torch.zeros(W.shape[0], rank, dtype=W.dtype, device=W.device)
else:
U, V = eva_init(W, rank, activation_cov=B)
return {"U": U, "V": V, "residual": None}
if method == "dora":
return None # peft handles DoRA natively
if method == "filet":
from safetensors.torch import load_file
p = _cache_path(cache_dir, "filet", base_ln)
data = load_file(str(p))
# Move to W.device so filet_init can mix them in Fisher Energy scoring.
# SX for down_proj is 820 MB; H200 has ample VRAM headroom.
dev = W.device
SX = data["SX"].to(device=dev, dtype=torch.float32)
SY = data["SY"].to(device=dev, dtype=torch.float32)
# For large matrices, cap pool to avoid prohibitive SVD cost
pool = min(min(W.shape), 512) if max(W.shape) > 8192 else None
U, V = filet_init(W, rank, sx=SX, sy=SY, candidate_pool=pool)
return {"U": U, "V": V, "residual": None}
if method in ("mff_top", "mff_energy"):
from safetensors import safe_open
sel = method.split("_", 1)[1] # "top" or "energy"
p = _mff_factor_path(mff_factor_dir, language, base_ln)
need_B = (sel == "energy")
with safe_open(str(p), framework="pt", device="cpu") as f:
A = f.get_tensor("A").to(torch.float32)
B = f.get_tensor("B").to(torch.float32) if need_B else None
r = mff_lora_init(W, A, B, rank, selection=sel, strategy="alpha", seed=seed)
return {"U": r.U, "V": r.V, "residual": r.residual}
if method in ("fpeft", "fpeft_low", "fpeft_bi_high", "fpeft_bi_low", "foi", "fws_safe"):
# All Fisher-informed methods use the same online factors as FILet
from safetensors.torch import load_file as _load_st
p = _cache_path(cache_dir, "filet", base_ln)
data = _load_st(str(p))
dev = W.device
SX = data["SX"].to(device=dev, dtype=torch.float32) # [m, m] input cov
SY = data["SY"].to(device=dev, dtype=torch.float32) # [n, n] output grad cov
if method == "fpeft":
B_lo, A_lo = fisher_peft_init(W, SX, rank, seed=seed)
return {"U": B_lo, "V": A_lo, "residual": None}
if method == "fpeft_low":
B_lo, A_lo = fisher_peft_low_init(W, SX, rank, seed=seed)
return {"U": B_lo, "V": A_lo, "residual": None}
if method == "fpeft_bi_high":
B_lo, A_lo = fisher_peft_bi_init(W, SX, SY, rank, select="high", seed=seed)
return {"U": B_lo, "V": A_lo, "residual": None}
if method == "fpeft_bi_low":
B_lo, A_lo = fisher_peft_bi_init(W, SX, SY, rank, select="low", seed=seed)
return {"U": B_lo, "V": A_lo, "residual": None}
if method == "foi":
B_lo, A_lo = fisher_orthogonal_init(W, SX, rank, seed=seed)
return {"U": B_lo, "V": A_lo, "residual": None}
if method == "fws_safe":
U, V = fws_safe_init(W, SX, SY, rank, seed=seed)
return {"U": U, "V": V, "residual": None}
raise ValueError(f"Unknown method: {method!r}")
return fn
# ---------------------------------------------------------------------------
# Single run
# ---------------------------------------------------------------------------
def run_single(
*,
method: str,
rank: int,
seed: int,
alpha_mult: float,
args: argparse.Namespace,
tok,
train_tok,
val_tok,
cache_dir: Path,
) -> dict:
set_seed(seed)
from peft import LoraConfig, get_peft_model
from transformers import (
AutoModelForSequenceClassification,
DataCollatorWithPadding,
Trainer,
TrainingArguments,
)
import evaluate
use_dora = method == "dora"
use_peft_default = method == "peft_default"
dtype = {
"bf16": torch.bfloat16,
"fp16": torch.float16,
"fp32": torch.float32,
}[args.precision]
model = AutoModelForSequenceClassification.from_pretrained(
args.model_path,
num_labels=3,
torch_dtype=dtype,
trust_remote_code=True,
)
model.config.pad_token_id = tok.pad_token_id
if args.gradient_checkpointing:
model.config.use_cache = False
model.gradient_checkpointing_enable()
if torch.cuda.is_available():
model = model.to("cuda")
lora_cfg = LoraConfig(
r=rank,
lora_alpha=rank * alpha_mult,
target_modules=args.target_modules,
bias="none",
task_type="SEQ_CLS",
# DoRA and peft_default use peft's native initialization; all others
# are injected manually for controlled comparisons.
init_lora_weights=True if (use_dora or use_peft_default) else False,
use_dora=use_dora,
modules_to_save=["score"],
)
peft_model = get_peft_model(model, lora_cfg)
if not use_dora and not use_peft_default:
init_fn = _make_init_fn(method, rank, seed, args.mff_factor_dir, args.language, cache_dir)
inits: dict[str, dict | None] = {}
for peft_name, mod in peft_model.named_modules():
if not hasattr(mod, "lora_A"):
continue
suffix = peft_name.rsplit(".", 1)[-1]
if suffix not in args.target_modules:
continue
inits[peft_name] = init_fn(peft_name, mod.base_layer)
inject_lora_init(peft_model, inits)
metric = evaluate.load("accuracy")
def compute_metrics(p):
preds = np.argmax(p.predictions, axis=-1)
return metric.compute(predictions=preds, references=p.label_ids)
out_dir = Path(args.output_dir) / f"{method}__rank{rank}__seed{seed}"
out_dir.mkdir(parents=True, exist_ok=True)
targs = TrainingArguments(
output_dir=str(out_dir),
num_train_epochs=args.epochs,
per_device_train_batch_size=args.batch_size,
per_device_eval_batch_size=args.batch_size * 2,
gradient_accumulation_steps=args.grad_accum,
learning_rate=args.lr,
eval_strategy="epoch",
save_strategy="no",
logging_steps=50,
report_to=[] if args.no_wandb else ["wandb"],
seed=seed,
bf16=torch.cuda.is_available() and args.precision == "bf16",
fp16=torch.cuda.is_available() and args.precision == "fp16",
gradient_checkpointing=args.gradient_checkpointing,
remove_unused_columns=False,
dataloader_num_workers=0,
)
trainer = Trainer(
model=peft_model,
args=targs,
train_dataset=train_tok,
eval_dataset=val_tok,
processing_class=tok,
data_collator=DataCollatorWithPadding(tok),
compute_metrics=compute_metrics,
)
trainer.train()
final = trainer.evaluate()
del peft_model, model
if torch.cuda.is_available():
torch.cuda.empty_cache()
return {k: float(v) for k, v in final.items() if isinstance(v, (int, float))}
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def main():
args = parse_args()
out_dir = Path(args.output_dir)
out_dir.mkdir(parents=True, exist_ok=True)
cache_dir = out_dir / ".factor_cache"
results_path = out_dir / "results.json"
summary_path = out_dir / "summary.json"
# Load existing results for resume support
if results_path.exists():
results: dict = json.loads(results_path.read_text())
else:
results = {}
print("Loading XNLI …", flush=True)
tok, train_tok, val_tok = _load_xnli(args.train_size, args.max_length, args.model_path)
# -----------------------------------------------------------------------
# Online factor collection (done once, before any training run)
# -----------------------------------------------------------------------
online_methods = {
m for m in args.methods
if m in ("filet", "eva", "lora_ga", "fpeft", "fpeft_low", "fpeft_bi_high", "fpeft_bi_low", "foi", "fws_safe")
}
# FWS reuses FILet's SX/SY factors, so map FWS -> "filet" for collection.
factor_types = tuple({
"filet" if m in ("fpeft", "fpeft_low", "fpeft_bi_high", "fpeft_bi_low", "foi", "fws_safe") else m
for m in online_methods
})
if online_methods:
print(f"Ensuring online factors are cached for {factor_types} …", flush=True)
from transformers import AutoModelForSequenceClassification
base_model = AutoModelForSequenceClassification.from_pretrained(
args.model_path,
num_labels=3,
torch_dtype={
"bf16": torch.bfloat16,
"fp16": torch.float16,
"fp32": torch.float32,
}[args.precision],
trust_remote_code=True,
)
base_model.config.pad_token_id = tok.pad_token_id
if args.gradient_checkpointing:
base_model.config.use_cache = False
base_model.gradient_checkpointing_enable()
if torch.cuda.is_available():
base_model = base_model.to("cuda")
# Target layer names in the base model (before peft wrapping)
target_layer_names = [
name for name, mod in base_model.named_modules()
if isinstance(mod, nn.Linear)
and name.rsplit(".", 1)[-1] in args.target_modules
]
print(f" Target layers: {len(target_layer_names)}", flush=True)
factor_loader = _make_factor_loader(train_tok, tok, args.factor_batch_size)
_collect_and_cache_online_factors(
base_model,
target_layer_names,
factor_loader,
args.n_factor_batches,
cache_dir,
factor_types,
)
del base_model
if torch.cuda.is_available():
torch.cuda.empty_cache()
# -----------------------------------------------------------------------
# Training sweep
# -----------------------------------------------------------------------
for alpha_mult in args.lora_alpha_mults:
for method in args.methods:
for rank in args.ranks:
for seed in args.seeds:
key = f"{method}__alpha{alpha_mult:g}__rank{rank}__seed{seed}"
if key in results:
print(f" [skip] {key} already in results", flush=True)
continue
print(f"==> {key}", flush=True)
try:
metrics = run_single(
method=method,
rank=rank,
seed=seed,
alpha_mult=alpha_mult,
args=args,
tok=tok,
train_tok=train_tok,
val_tok=val_tok,
cache_dir=cache_dir,
)
results[key] = metrics
except Exception as exc:
print(f" ERROR: {exc}", flush=True)
results[key] = {"error": str(exc)}
results_path.write_text(json.dumps(results, indent=2))
# -----------------------------------------------------------------------
# Summary: mean ± std per (method, rank) across seeds
# -----------------------------------------------------------------------
summary: dict[str, dict[str, dict]] = {}
for alpha_mult in args.lora_alpha_mults:
summary[f"alpha{alpha_mult:g}"] = {}
for method in args.methods:
summary[f"alpha{alpha_mult:g}"][method] = {}
for rank in args.ranks:
accs = []
for seed in args.seeds:
k = f"{method}__alpha{alpha_mult:g}__rank{rank}__seed{seed}"
v = results.get(k, {})
if "eval_accuracy" in v:
accs.append(v["eval_accuracy"])
if accs:
summary[f"alpha{alpha_mult:g}"][method][f"rank{rank}"] = {
"mean": float(np.mean(accs)),
"std": float(np.std(accs)),
"n": len(accs),
}
summary_path.write_text(json.dumps(summary, indent=2))
print("\n=== E1 Summary ===")
for alpha_key, method_results in summary.items():
if not isinstance(method_results, dict) or not method_results:
continue
print(f" {alpha_key}:")
for method, rank_results in method_results.items():
if not isinstance(rank_results, dict) or not rank_results:
continue
if "mean" in rank_results and "std" in rank_results:
vals = f"mean: {rank_results['mean']:.4f}±{rank_results['std']:.4f}"
else:
vals = " | ".join(
f"{r}: {d['mean']:.4f}±{d['std']:.4f}"
for r, d in rank_results.items()
if isinstance(d, dict) and "mean" in d and "std" in d
)
if vals:
print(f" {method:12s} {vals}")
if __name__ == "__main__":
main()

Xet Storage Details

Size:
31.6 kB
·
Xet hash:
a679ed02cd0f9664642e4ca0cf5426a5391351c3a299cc467e4beadc0f3724a5

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.