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