etomoscow's picture
download
raw
38.7 kB
"""Broad multi-model, multi-task classification sweep for Fisher-LoRA.
This runner covers the broad-sweep tasks for XNLI, BoolQ, and RTE using a
shared sequence-classification setup:
- qwen3 / llama / ministral use their native sequence-classification heads.
- gemma4 is wrapped with a lightweight pooled hidden-state score head.
LoRA initializations:
- fpeft, fpeft_low, fpeft_bi_low, fpeft_bi_high use Fisher-informed init helpers.
- filet uses Fisher-energy selection from FILet factors.
- peft_default uses native PEFT initialization.
- random_orthogonal and peft_scaled_random are zero-preserving controls.
- mff_top is supported for the Llama BoolQ gate only.
The script writes per-run JSON under the requested output directory and a
summary.json with per-method aggregate statistics plus collapse counts.
"""
from __future__ import annotations
import argparse
import gc
import json
import sys
import time
from pathlib import Path
from typing import Any
import numpy as np
import torch
from torch import nn
from transformers.modeling_outputs import SequenceClassifierOutput
REPO = Path(__file__).resolve().parents[2]
SRC = REPO / "src"
if str(REPO) not in sys.path:
sys.path.insert(0, str(REPO))
if str(SRC) not in sys.path:
sys.path.insert(0, str(SRC))
import mfflora._compat # noqa: E402,F401
from mfflora.init import ( # noqa: E402
fisher_peft_bi_init,
fisher_peft_init,
fisher_peft_low_init,
filet_init,
mff_lora_init,
)
from mfflora.train.lora_runner import inject_lora_init # noqa: E402
from mfflora.utils import set_seed # noqa: E402
MODEL_REGISTRY = {
"llama": {
"hf_id": "unsloth/Llama-3.1-8B",
"factor_root": REPO / "outputs" / "e1_0_additional_factors",
"kind": "auto_seqcls",
},
"qwen3": {
"hf_id": "Qwen/Qwen3-8B",
"factor_root": REPO / "outputs" / "factors" / "qwen3",
"kind": "auto_seqcls",
},
"ministral": {
"hf_id": "mistralai/Ministral-3-8B-Base-2512",
"factor_root": REPO / "outputs" / "factors" / "ministral",
"kind": "ministral_seqcls",
},
"gemma4": {
"hf_id": "google/gemma-4-E4B",
"factor_root": REPO / "outputs" / "factors" / "gemma4",
"kind": "gemma4_wrapper",
},
}
DATASETS = {"hellaswag", "xnli", "boolq", "rte"}
ALL_METHODS = [
"fpeft_low", "fpeft_bi_low", "fpeft_bi_high", "peft_default", "filet",
"mff_top", "random_orthogonal", "peft_scaled_random",
]
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser()
p.add_argument("--model-id", choices=sorted(MODEL_REGISTRY), required=True)
p.add_argument("--dataset", choices=sorted(DATASETS), required=True)
p.add_argument("--model-path", default=None, help="Override HF model id or local path.")
p.add_argument("--factor-root", default=None, help="Override factor root directory.")
p.add_argument("--methods", nargs="+", default=None, help="Methods to run.")
p.add_argument("--ranks", nargs="+", type=int, default=[8, 16, 32])
p.add_argument("--seeds", nargs="+", type=int, default=[42, 1337, 2024, 7, 123])
p.add_argument("--train-size", type=int, default=None, help="Subsample train split size; defaults vary by dataset.")
p.add_argument("--eval-size", type=int, default=None, help="Optional eval subsample size.")
p.add_argument("--epochs", type=int, default=None, help="Override training epochs (defaults vary by dataset).")
p.add_argument("--lr", type=float, default=1e-4)
p.add_argument("--batch-size", type=int, default=8)
p.add_argument("--grad-accum", type=int, default=2)
p.add_argument("--max-length", type=int, default=256)
p.add_argument("--gradient-checkpointing", action="store_true")
p.add_argument("--hellaswag-eval-batch-size", type=int, default=32)
p.add_argument("--alpha-fpeft", type=float, default=0.5)
p.add_argument("--alpha-baseline", type=float, default=1.0)
p.add_argument("--collapse-threshold", type=float, default=70.0)
p.add_argument("--output-dir", required=True)
p.add_argument("--dataset-cache-dir", default=None)
p.add_argument("--model-cache-dir", default=None)
p.add_argument("--no-wandb", action="store_true")
return p.parse_args()
class PooledScoreWrapper(nn.Module):
def __init__(
self,
backbone: nn.Module,
hidden_size: int,
num_labels: int,
*,
use_base_model: bool = False,
):
super().__init__()
self.backbone = backbone
self.config = backbone.config
self.score = nn.Linear(hidden_size, num_labels, bias=False)
self.use_base_model = use_base_model
def forward(self, input_ids=None, attention_mask=None, labels=None, **kwargs):
kwargs.pop("output_hidden_states", None)
kwargs.pop("return_dict", None)
kwargs.pop("use_cache", None)
target = (
self.backbone.model
if self.use_base_model and hasattr(self.backbone, "model")
else self.backbone
)
out = target(
input_ids=input_ids,
attention_mask=attention_mask,
return_dict=True,
use_cache=False,
**kwargs,
)
hidden = out.last_hidden_state if hasattr(out, "last_hidden_state") else out[0]
if attention_mask is None:
pooled = hidden[:, -1]
else:
idx = attention_mask.to(torch.long).sum(dim=1).clamp_min(1) - 1
pooled = hidden[torch.arange(hidden.size(0), device=hidden.device), idx]
logits = self.score(pooled)
loss = None
if labels is not None:
loss = nn.CrossEntropyLoss()(logits, labels)
return SequenceClassifierOutput(loss=loss, logits=logits)
def _model_spec(model_id: str, model_path: str | None, factor_root: str | None):
spec = dict(MODEL_REGISTRY[model_id])
if model_path is not None:
spec["hf_id"] = model_path
if factor_root is not None:
spec["factor_root"] = Path(factor_root)
else:
spec["factor_root"] = Path(spec["factor_root"])
return spec
def _load_tokenizer(hf_id: str, cache_dir: str | None):
import json
from transformers import AutoTokenizer
kwargs = {"trust_remote_code": True, "cache_dir": cache_dir}
tokenizer_config = Path(hf_id) / "tokenizer_config.json"
if tokenizer_config.is_file():
raw = json.loads(tokenizer_config.read_text())
if isinstance(raw.get("extra_special_tokens"), list):
# Older Gemma snapshots store this as a list; newer Transformers
# expects a token-to-id mapping. Text classification does not use
# the media token, so an empty mapping is the correct text-only
# compatibility value.
kwargs["extra_special_tokens"] = {}
tok = AutoTokenizer.from_pretrained(hf_id, **kwargs)
if tok.pad_token is None:
tok.pad_token = tok.eos_token
tok.padding_side = "right"
return tok
def _dataset_spec(dataset: str):
if dataset == "hellaswag":
return {
"name": "hellaswag",
"hf_args": ("Rowan/hellaswag", None),
"num_labels": 4,
"train_split": "train",
"eval_split": "validation",
"label_name": "label",
"train_default": 8192,
"epochs_default": 3,
}
if dataset == "xnli":
return {
"name": "xnli",
"hf_args": ("xnli", "en"),
"num_labels": 3,
"train_split": "train",
"eval_split": "validation",
"label_name": "label",
"train_default": 8192,
"epochs_default": 3,
}
if dataset == "boolq":
return {
"name": "boolq",
"hf_args": ("super_glue", "boolq"),
"num_labels": 2,
"train_split": "train",
"eval_split": "validation",
"label_name": "label",
"train_default": 8192,
"epochs_default": 3,
}
if dataset == "rte":
return {
"name": "rte",
"hf_args": ("glue", "rte"),
"num_labels": 2,
"train_split": "train",
"eval_split": "validation",
"label_name": "label",
"train_default": None,
"epochs_default": 5,
}
raise ValueError(dataset)
def _load_dataset(
dataset: str,
tok,
train_size: int | None,
eval_size: int | None,
cache_dir: str | None,
max_length: int,
):
from datasets import load_dataset
spec = _dataset_spec(dataset)
ds = load_dataset(spec["hf_args"][0], spec["hf_args"][1], cache_dir=cache_dir)
train = ds[spec["train_split"]]
eval_ds = ds[spec["eval_split"]]
if train_size is not None:
train = train.shuffle(seed=42).select(range(min(train_size, len(train))))
elif spec["train_default"] is not None and len(train) > spec["train_default"]:
train = train.shuffle(seed=42).select(range(spec["train_default"]))
if eval_size is not None:
eval_ds = eval_ds.select(range(min(eval_size, len(eval_ds))))
def tok_fn(batch):
if dataset == "xnli":
enc = tok(batch["premise"], batch["hypothesis"], truncation=True, max_length=max_length)
elif dataset == "boolq":
prompts = [f"question: {q} passage: {p}" for q, p in zip(batch["question"], batch["passage"])]
enc = tok(prompts, truncation=True, max_length=max_length)
else: # rte
prompts = [f"sentence1: {s1}" for s1 in batch["sentence1"]]
hyps = [f"sentence2: {s2}" for s2 in batch["sentence2"]]
enc = tok(prompts, hyps, truncation=True, max_length=max_length)
enc["labels"] = batch[spec["label_name"]]
return enc
remove_cols = [c for c in train.column_names if c != spec["label_name"]]
train_tok = train.map(tok_fn, batched=True, remove_columns=remove_cols)
eval_tok = eval_ds.map(tok_fn, batched=True, remove_columns=remove_cols)
return train_tok, eval_tok
def _load_hellaswag_dataset(tok, train_size: int | None, eval_size: int | None, cache_dir: str | None, max_length: int):
from datasets import load_dataset
ds = load_dataset("Rowan/hellaswag", cache_dir=cache_dir)
train = ds["train"]
eval_ds = ds["validation"]
if train_size is not None:
train = train.shuffle(seed=42).select(range(min(train_size, len(train))))
elif len(train) > 8192:
train = train.shuffle(seed=42).select(range(8192))
if eval_size is not None:
eval_ds = eval_ds.select(range(min(eval_size, len(eval_ds))))
def tok_train(batch):
input_ids = []
attention_mask = []
labels = []
for ctx, endings, label in zip(batch["ctx"], batch["endings"], batch["label"]):
prompt = ctx + " "
answer = endings[int(label)]
prompt_ids = tok(prompt, truncation=True, max_length=max_length, add_special_tokens=True)["input_ids"]
enc = tok(prompt + answer, truncation=True, max_length=max_length, padding="max_length")
labs = list(enc["input_ids"])
prompt_len = min(len(prompt_ids), len(labs))
for i in range(prompt_len):
labs[i] = -100
for i, m in enumerate(enc["attention_mask"]):
if m == 0:
labs[i] = -100
input_ids.append(enc["input_ids"])
attention_mask.append(enc["attention_mask"])
labels.append(labs)
return {"input_ids": input_ids, "attention_mask": attention_mask, "labels": labels}
train_tok = train.map(tok_train, batched=True, remove_columns=train.column_names)
return train_tok, eval_ds
def _examples_from_dataset_slice(batch: Any) -> list[dict[str, Any]]:
if isinstance(batch, dict):
if not batch:
return []
n = len(next(iter(batch.values())))
return [{key: value[i] for key, value in batch.items()} for i in range(n)]
return list(batch)
def _score_hellaswag(model: nn.Module, tok, eval_ds, max_length: int, batch_size: int = 8) -> dict[str, float]:
model.eval()
total = 0
correct = 0
losses = []
def seq_loss(input_ids: torch.Tensor, attention_mask: torch.Tensor, labels: torch.Tensor) -> torch.Tensor:
out = model(input_ids=input_ids, attention_mask=attention_mask, return_dict=True)
logits = out.logits[:, :-1, :]
targets = labels[:, 1:]
mask = targets.ne(-100)
safe_targets = targets.clamp_min(0)
log_probs = torch.log_softmax(logits, dim=-1)
tok_nll = -log_probs.gather(-1, safe_targets.unsqueeze(-1)).squeeze(-1)
tok_nll = tok_nll * mask
denom = mask.sum(dim=1).clamp_min(1)
return tok_nll.sum(dim=1) / denom
for start in range(0, len(eval_ds), batch_size):
batch = _examples_from_dataset_slice(eval_ds[start : start + batch_size])
seqs = []
golds = []
for ex in batch:
ctx = ex["ctx"] + " "
prompt_ids = tok(ctx, truncation=True, max_length=max_length, add_special_tokens=True)["input_ids"]
for ending in ex["endings"]:
enc = tok(ctx + ending, truncation=True, max_length=max_length, padding="max_length")
labs = list(enc["input_ids"])
prompt_len = min(len(prompt_ids), len(labs))
for i in range(prompt_len):
labs[i] = -100
for i, m in enumerate(enc["attention_mask"]):
if m == 0:
labs[i] = -100
seqs.append((enc["input_ids"], enc["attention_mask"], labs))
golds.append(int(ex["label"]))
input_ids = torch.tensor([s[0] for s in seqs], device=next(model.parameters()).device)
attention_mask = torch.tensor([s[1] for s in seqs], device=input_ids.device)
labels = torch.tensor([s[2] for s in seqs], device=input_ids.device)
with torch.no_grad():
seq_losses = seq_loss(input_ids, attention_mask, labels).detach().cpu().tolist()
for i, gold in enumerate(golds):
row = seq_losses[i * 4 : (i + 1) * 4]
pred = int(np.argmin(row))
correct += int(pred == gold)
total += 1
losses.append(float(row[pred]))
return {"accuracy": correct / max(total, 1), "mean_loss": float(np.mean(losses)) if losses else 0.0, "num_examples": total}
def _load_backbone(model_id: str, hf_id: str, num_labels: int, tokenizer, dtype: torch.dtype):
spec = MODEL_REGISTRY[model_id]["kind"]
if spec == "auto_seqcls":
from transformers import AutoModelForSequenceClassification
model = AutoModelForSequenceClassification.from_pretrained(
hf_id,
num_labels=num_labels,
trust_remote_code=True,
torch_dtype=dtype,
cache_dir=MODEL_REGISTRY[model_id].get("model_cache_dir"),
)
model.config.pad_token_id = tokenizer.pad_token_id
return model
if spec == "ministral_seqcls":
backbone = _load_ministral_text_backbone(hf_id, tokenizer, dtype)
return PooledScoreWrapper(backbone, backbone.config.hidden_size, num_labels)
if spec == "gemma4_wrapper":
raise RuntimeError("gemma4 wrapper loads via a separate path")
def _load_gemma4_backbone(hf_id: str, tokenizer, dtype: torch.dtype):
from huggingface_hub import hf_hub_download, scan_cache_dir
from huggingface_hub.utils import LocalEntryNotFoundError
from safetensors import safe_open
from transformers import AutoConfig
from transformers.models.gemma4 import Gemma4TextModel
cfg = AutoConfig.from_pretrained(hf_id, trust_remote_code=True)
text_cfg = cfg.text_config
backbone = Gemma4TextModel(text_cfg).to(dtype=dtype)
expected = set(backbone.state_dict().keys())
try:
weights_path = Path(hf_id) / "model.safetensors" if Path(hf_id).is_dir() else hf_hub_download(
hf_id, "model.safetensors", local_files_only=True
)
except LocalEntryNotFoundError:
revisions = [
revision
for repo in scan_cache_dir().repos
if repo.repo_id == hf_id
for revision in repo.revisions
]
weights_path = _cached_file_from_revisions(revisions, "model.safetensors")
if weights_path is None:
raise
state_dict = {}
prefix = "model.language_model."
with safe_open(weights_path, framework="pt", device="cpu") as handle:
for key in handle.keys():
if not key.startswith(prefix):
continue
target_key = key[len(prefix) :]
if target_key not in expected:
continue
state_dict[target_key] = handle.get_tensor(key).to(dtype=dtype)
missing, unexpected = backbone.load_state_dict(state_dict, strict=False)
if missing or unexpected:
print(
"[gemma4-text-load] "
f"loaded={len(state_dict)} missing={len(missing)} unexpected={len(unexpected)}"
)
if missing:
print(f"[gemma4-text-load] missing sample: {missing[:8]}")
if unexpected:
print(f"[gemma4-text-load] unexpected sample: {unexpected[:8]}")
backbone.config.pad_token_id = tokenizer.pad_token_id
backbone.config.use_cache = False
return backbone
def _load_ministral_text_backbone(hf_id: str, tokenizer, dtype: torch.dtype):
from huggingface_hub import hf_hub_download
from safetensors import safe_open
from transformers import MistralConfig, MistralModel
config_path = Path(hf_id) / "config.json" if Path(hf_id).is_dir() else Path(
hf_hub_download(hf_id, "config.json", local_files_only=True)
)
raw_config = json.loads(config_path.read_text())
text_config = dict(raw_config["text_config"])
# This checkpoint uses a stale ``ministral3`` text model label. The
# installed Transformers implementation exposes the decoder as Mistral.
text_config["model_type"] = "mistral"
backbone = MistralModel(MistralConfig(**text_config)).to(dtype=dtype)
expected = set(backbone.state_dict().keys())
index_path = Path(hf_id) / "model.safetensors.index.json" if Path(hf_id).is_dir() else hf_hub_download(
hf_id, "model.safetensors.index.json", local_files_only=True
)
index = json.loads(Path(index_path).read_text())
state_dict = {}
prefix = "language_model.model."
for key, shard in index["weight_map"].items():
if not key.startswith(prefix):
continue
target_key = key[len(prefix) :]
if target_key not in expected:
continue
shard_path = Path(hf_id) / shard if Path(hf_id).is_dir() else hf_hub_download(
hf_id, shard, local_files_only=True
)
with safe_open(shard_path, framework="pt", device="cpu") as handle:
state_dict[target_key] = handle.get_tensor(key).to(dtype=dtype)
missing, unexpected = backbone.load_state_dict(state_dict, strict=False)
print(
"[ministral-text-load] "
f"loaded={len(state_dict)} missing={len(missing)} unexpected={len(unexpected)}"
)
if missing:
print(f"[ministral-text-load] missing sample: {missing[:8]}")
if unexpected:
print(f"[ministral-text-load] unexpected sample: {unexpected[:8]}")
backbone.config.pad_token_id = tokenizer.pad_token_id
backbone.config.use_cache = False
return backbone
def _load_ministral_causal_lm(hf_id: str, tokenizer, dtype: torch.dtype):
from huggingface_hub import hf_hub_download
from safetensors import safe_open
from transformers import AutoConfig
from transformers.models.ministral3 import Ministral3ForCausalLM
cfg = AutoConfig.from_pretrained(hf_id, trust_remote_code=True)
model = Ministral3ForCausalLM(cfg.text_config).to(dtype=dtype)
expected = set(model.state_dict().keys())
index_path = hf_hub_download(hf_id, "model.safetensors.index.json", local_files_only=True)
index = json.loads(Path(index_path).read_text())
state_dict = {}
prefixes = {
"language_model.model.": "model.",
"language_model.lm_head.": "lm_head.",
}
for key, shard in index["weight_map"].items():
target_key = None
for src_prefix, dst_prefix in prefixes.items():
if key.startswith(src_prefix):
target_key = dst_prefix + key[len(src_prefix) :]
break
if target_key is None or target_key not in expected:
continue
shard_path = hf_hub_download(hf_id, shard, local_files_only=True)
with safe_open(shard_path, framework="pt", device="cpu") as handle:
state_dict[target_key] = handle.get_tensor(key).to(dtype=dtype)
missing, unexpected = model.load_state_dict(state_dict, strict=False)
print(
"[ministral-causal-load] "
f"loaded={len(state_dict)} missing={len(missing)} unexpected={len(unexpected)}"
)
if missing or unexpected:
raise RuntimeError(
"Ministral causal loader did not load cleanly: "
f"missing={missing[:8]} unexpected={unexpected[:8]}"
)
model.config.pad_token_id = tokenizer.pad_token_id
model.config.use_cache = False
return model
def _target_layer_names(model: nn.Module, target_modules: list[str]) -> list[str]:
return [
name
for name, module in model.named_modules()
if isinstance(module, nn.Linear) and name.rsplit(".", 1)[-1] in target_modules
]
def _factor_path(factor_root: Path, language: str, base_layer_name: str) -> Path:
module = base_layer_name.rsplit(".", 1)[-1]
filename = base_layer_name.replace(".", "__") + ".safetensors"
return factor_root / language / module / filename
def _normalize_base_name(peft_name: str) -> str:
base = peft_name.removeprefix("base_model.model.")
if base.startswith("backbone."):
base = base.removeprefix("backbone.")
if base.startswith("layers.") or base.startswith("embed_tokens") or base.startswith("norm"):
base = f"model.{base}"
return base
def _make_init_fn(method: str, rank: int, seed: int, factor_root: Path, language: str, scale: float = 1.0):
def fn(peft_name: str, base_linear: nn.Linear) -> dict | None:
base_ln = _normalize_base_name(peft_name)
W = base_linear.weight.detach()
if method == "peft_default":
return {"U": torch.zeros(W.shape[0], rank, dtype=W.dtype, device=W.device), "V": torch.zeros(rank, W.shape[1], dtype=W.dtype, device=W.device), "residual": None}
if method in {"random_orthogonal", "peft_scaled_random"}:
from experiments.e1_single_language.run import (
kaiming_random_init,
random_orthogonal_init,
)
init = random_orthogonal_init if method == "random_orthogonal" else kaiming_random_init
U, V = init(W, rank, seed=seed)
return {"U": U, "V": V, "residual": None}
from safetensors.torch import load_file
p = _factor_path(factor_root, language, base_ln)
data = load_file(str(p))
# Factor files store K-FAC moments: A = SY (output-side), B = SX (input-side).
SY = data["A"].to(device=W.device, dtype=torch.float32)
SX = data["B"].to(device=W.device, dtype=torch.float32)
if method == "mff_top":
r = mff_lora_init(W, SY, SX, rank, selection="top", strategy="alpha", seed=seed)
return {"U": r.U, "V": r.V, "residual": r.residual}
if method == "filet":
U, V = filet_init(W, rank, sx=SX, sy=SY, candidate_pool=512 if max(W.shape) > 8192 else None)
residual = (W - scale * (U @ V)).to(W.dtype)
return {"U": U, "V": V, "residual": residual}
if method == "fpeft":
U, V = fisher_peft_init(W, SX, rank, seed=seed)
return {"U": U, "V": V, "residual": None}
if method == "fpeft_low":
U, V = fisher_peft_low_init(W, SX, rank, seed=seed)
return {"U": U, "V": V, "residual": None}
if method == "fpeft_bi_low":
U, V = fisher_peft_bi_init(W, SX, SY, rank, select="low", seed=seed)
return {"U": U, "V": V, "residual": None}
if method == "fpeft_bi_high":
U, V = fisher_peft_bi_init(W, SX, SY, rank, select="high", seed=seed)
return {"U": U, "V": V, "residual": None}
raise ValueError(f"Unknown method: {method}")
return fn
@torch.no_grad()
def _collapse_flag(acc: float | None, threshold: float) -> bool | None:
if acc is None:
return None
return bool(acc < threshold / 100.0)
def _prediction_distribution(logits: Any, num_labels: int) -> list[int]:
if isinstance(logits, (tuple, list)):
logits = logits[0]
predictions = np.argmax(np.asarray(logits), axis=-1)
return np.bincount(predictions, minlength=num_labels).astype(int).tolist()
def _cached_file_from_revisions(revisions: Any, filename: str) -> str | None:
candidates = [
Path(file.file_path)
for revision in revisions
for file in revision.files
if file.file_name == filename and Path(file.file_path).is_file()
]
if not candidates:
return None
return str(max(candidates, key=lambda path: path.stat().st_size))
def _is_retryable_cell_error(exc: Exception) -> bool:
text = repr(exc).lower()
return (
isinstance(exc, BrokenPipeError)
or "brokenpipeerror" in text
or "broken pipe" in text
or "outofmemoryerror" in text
or "cuda out of memory" in text
)
def _cleanup_after_failed_cell() -> None:
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
torch.cuda.ipc_collect()
def _load_factor_cache_summary(cache_dir: Path, language: str, target_modules: list[str], base_model: nn.Module):
# The broad sweep reuses the existing FILet factor cache layout. If the
# cache already exists, we do nothing here. For freshly collected factors,
# this function is not used.
return None
def run_cell(
*,
model_id: str,
dataset: str,
method: str,
rank: int,
seed: int,
args: argparse.Namespace,
tokenizer,
train_tok,
eval_tok,
num_labels: int,
factor_root: Path,
model_path: str | None = None,
) -> dict[str, Any]:
if dataset == "hellaswag" and model_id == "gemma4":
raise ValueError("Gemma4 HellaSwag is not supported until a verified causal-LM loader exists.")
set_seed(seed)
from peft import LoraConfig, get_peft_model
from transformers import Trainer, TrainingArguments, DataCollatorWithPadding, default_data_collator
import evaluate
dtype = torch.bfloat16 if torch.cuda.is_available() else torch.float32
hf_id = model_path or MODEL_REGISTRY[model_id]["hf_id"]
if dataset == "hellaswag":
if model_id == "gemma4":
model = _load_gemma4_backbone(hf_id, tokenizer, dtype)
elif model_id == "ministral":
model = _load_ministral_causal_lm(hf_id, tokenizer, dtype)
else:
from transformers import AutoModelForCausalLM
model = AutoModelForCausalLM.from_pretrained(
hf_id,
trust_remote_code=True,
torch_dtype=dtype,
)
model.config.pad_token_id = tokenizer.pad_token_id
model.config.use_cache = False
elif model_id == "gemma4":
backbone = _load_gemma4_backbone(hf_id, tokenizer, dtype)
model = PooledScoreWrapper(
backbone, backbone.config.hidden_size, num_labels, use_base_model=True
)
if hasattr(backbone, "gradient_checkpointing_enable"):
backbone.gradient_checkpointing_enable()
else:
model = _load_backbone(model_id, hf_id, num_labels, tokenizer, dtype)
if getattr(args, "gradient_checkpointing", False) and hasattr(model, "gradient_checkpointing_enable"):
model.gradient_checkpointing_enable()
model.config.use_cache = False
if torch.cuda.is_available():
model = model.to("cuda")
alpha = args.alpha_fpeft if method.startswith("fpeft") else args.alpha_baseline
if method in {"filet", "peft_default", "mff_top"}:
alpha = args.alpha_baseline
lora_cfg = LoraConfig(
r=rank,
lora_alpha=rank * alpha,
target_modules=["q_proj", "v_proj", "down_proj"],
bias="none",
task_type="CAUSAL_LM" if dataset == "hellaswag" else "SEQ_CLS",
init_lora_weights=(method == "peft_default"),
modules_to_save=[] if dataset == "hellaswag" else ["score"],
)
peft_model = get_peft_model(model, lora_cfg)
if method != "peft_default":
init_fn = _make_init_fn(method, rank, seed, factor_root, "en", scale=alpha)
inits: dict[str, dict[str, torch.Tensor | None]] = {}
for peft_name, mod in peft_model.named_modules():
if not hasattr(mod, "lora_A"):
continue
if peft_name.rsplit(".", 1)[-1] not in ["q_proj", "v_proj", "down_proj"]:
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 = p.predictions
if isinstance(preds, (tuple, list)):
preds = preds[0]
preds = np.argmax(np.asarray(preds), 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)
t0 = time.perf_counter()
eval_batch_size = args.batch_size if model_id == "gemma4" else args.batch_size * 2
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=eval_batch_size,
gradient_accumulation_steps=args.grad_accum,
learning_rate=args.lr,
eval_strategy="no" if dataset == "hellaswag" else "epoch",
save_strategy="no",
logging_steps=50,
report_to=[] if args.no_wandb else ["wandb"],
seed=seed,
bf16=torch.cuda.is_available(),
remove_unused_columns=False,
dataloader_num_workers=0,
)
if dataset == "hellaswag":
trainer = Trainer(
model=peft_model,
args=targs,
train_dataset=train_tok,
data_collator=default_data_collator,
)
trainer.train()
final = _score_hellaswag(
peft_model,
tokenizer,
eval_tok,
args.max_length,
batch_size=args.hellaswag_eval_batch_size,
)
elapsed = time.perf_counter() - t0
result = {k: float(v) for k, v in final.items() if isinstance(v, (int, float))}
result["eval_accuracy"] = float(final["accuracy"])
result["eval_loss"] = float(final["mean_loss"])
result["elapsed_seconds"] = float(elapsed)
result["status"] = "complete"
if torch.cuda.is_available():
torch.cuda.empty_cache()
return result
trainer = Trainer(
model=peft_model,
args=targs,
train_dataset=train_tok,
eval_dataset=eval_tok,
data_collator=DataCollatorWithPadding(tokenizer),
compute_metrics=compute_metrics,
)
trainer.train()
final = trainer.predict(eval_tok)
logits = final.predictions[0] if isinstance(final.predictions, (tuple, list)) else final.predictions
prediction_distribution = _prediction_distribution(logits, num_labels)
eval_accuracy = float(metric.compute(
predictions=np.argmax(np.asarray(logits), axis=-1),
references=np.asarray(eval_tok["labels"]),
)["accuracy"])
elapsed = time.perf_counter() - t0
result = {k: float(v) for k, v in final.metrics.items() if isinstance(v, (int, float))}
result["eval_accuracy"] = eval_accuracy
result["prediction_distribution"] = prediction_distribution
result["elapsed_seconds"] = float(elapsed)
result["status"] = "complete"
if torch.cuda.is_available():
torch.cuda.empty_cache()
return result
def main() -> None:
args = parse_args()
spec = _model_spec(args.model_id, args.model_path, args.factor_root)
out_dir = Path(args.output_dir)
out_dir.mkdir(parents=True, exist_ok=True)
results_path = out_dir / "results.json"
summary_path = out_dir / "summary.json"
results: dict[str, Any] = json.loads(results_path.read_text()) if results_path.exists() else {}
dspec = _dataset_spec(args.dataset)
epochs = args.epochs if args.epochs is not None else dspec["epochs_default"]
train_size = args.train_size if args.train_size is not None else dspec["train_default"]
tokenizer = _load_tokenizer(spec["hf_id"], args.model_cache_dir)
if args.dataset == "hellaswag":
train_tok, eval_tok = _load_hellaswag_dataset(tokenizer, train_size, args.eval_size, args.dataset_cache_dir, args.max_length)
else:
train_tok, eval_tok = _load_dataset(args.dataset, tokenizer, train_size, args.eval_size, args.dataset_cache_dir, args.max_length)
num_labels = dspec["num_labels"]
methods = args.methods or (["fpeft_low", "fpeft_bi_low", "fpeft_bi_high", "peft_default", "filet"] if args.dataset != "boolq" or args.model_id != "llama" else ["fpeft_low", "fpeft_bi_low", "fpeft_bi_high", "peft_default", "filet", "mff_top"])
if args.dataset == "boolq" and args.model_id != "llama":
methods = [m for m in methods if m != "mff_top"]
factor_root = spec["factor_root"]
for method in methods:
for rank in args.ranks:
for seed in args.seeds:
key = f"{method}__rank{rank}__seed{seed}"
if key in results and isinstance(results[key], dict) and results[key].get("status") == "complete":
print(f"[skip] {key}")
continue
print(f"==> {key}")
attempts = 0
while True:
attempts += 1
try:
cell = run_cell(
model_id=args.model_id,
model_path=spec["hf_id"],
dataset=args.dataset,
method=method,
rank=rank,
seed=seed,
args=argparse.Namespace(
output_dir=str(out_dir),
lr=args.lr,
batch_size=args.batch_size,
grad_accum=args.grad_accum,
no_wandb=args.no_wandb,
alpha_fpeft=args.alpha_fpeft,
alpha_baseline=args.alpha_baseline,
epochs=epochs,
max_length=args.max_length,
gradient_checkpointing=args.gradient_checkpointing,
hellaswag_eval_batch_size=args.hellaswag_eval_batch_size,
),
tokenizer=tokenizer,
train_tok=train_tok,
eval_tok=eval_tok,
num_labels=num_labels,
factor_root=factor_root,
)
if attempts > 1:
cell["retry_attempts"] = attempts - 1
results[key] = cell
break
except Exception as exc:
retryable = _is_retryable_cell_error(exc)
if retryable and attempts < 4:
delay = 2 ** (attempts - 1)
print(f"[retry] {key} attempt {attempts}/3 failed with {repr(exc)}; sleeping {delay}s")
_cleanup_after_failed_cell()
time.sleep(delay)
continue
results[key] = {
"status": "error",
"error": repr(exc),
"attempts": attempts,
"retryable": retryable,
}
break
results_path.write_text(json.dumps(results, indent=2))
summary: dict[str, Any] = {
"model_id": args.model_id,
"dataset": args.dataset,
"methods": methods,
"collapse_threshold": args.collapse_threshold,
"gate_decisions": {},
"per_method": {},
"tasks_complete": [],
}
for method in methods:
accs = []
healthy = 0
collapses = 0
for rank in args.ranks:
for seed in args.seeds:
cell = results.get(f"{method}__rank{rank}__seed{seed}")
if not isinstance(cell, dict) or cell.get("status") != "complete":
continue
acc = cell.get("eval_accuracy")
if acc is None:
continue
accs.append(float(acc))
if float(acc) < args.collapse_threshold / 100.0:
collapses += 1
else:
healthy += 1
if accs:
summary["per_method"][method] = {
"mean": float(np.mean(accs)),
"std": float(np.std(accs)),
"n": len(accs),
"healthy_cells": healthy,
"collapse_cells": collapses,
}
if args.dataset == "boolq" and args.model_id == "llama":
fpl = summary["per_method"].get("fpeft_low", {}).get("mean")
fpb = summary["per_method"].get("fpeft_bi_low", {}).get("mean")
pft = summary["per_method"].get("peft_default", {}).get("mean")
gate = False
if fpl is not None and fpb is not None and pft is not None:
gate = bool(
(fpl >= 0.80 or fpb >= 0.80)
and (abs(fpl - pft) <= 0.03 or abs(fpb - pft) <= 0.03)
)
summary["gate_decisions"]["boolq_llama_proceed"] = gate
summary_path.write_text(json.dumps(summary, indent=2))
print(json.dumps(summary, indent=2))
if __name__ == "__main__":
main()

Xet Storage Details

Size:
38.7 kB
·
Xet hash:
61a7c3d6bace3f53177135075b184d84a01a2fef628af09272170d3a87ef7bb0

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