| """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.