Instructions to use moncefem/memory-lora-gemma4 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- PEFT
How to use moncefem/memory-lora-gemma4 with PEFT:
Task type is invalid.
- Notebooks
- Google Colab
- Kaggle
File size: 38,995 Bytes
481fbb6 f17edea 481fbb6 f17edea 481fbb6 f17edea 481fbb6 f17edea 481fbb6 f17edea 481fbb6 f17edea 481fbb6 f17edea 481fbb6 930bb27 f17edea 481fbb6 f17edea 481fbb6 f17edea 481fbb6 f17edea 481fbb6 f17edea 481fbb6 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 | #!/usr/bin/env python3
"""Train the Memory-LoRA hypernetwork on google/gemma-4-E2B.
Forked from Code2LoRA's ``hypernetwork/train_code2lora_static_v2.py``
(direct-projection trainer), retargeted:
* repo embedding -> doc embedding (memory_lora/encoder.py output)
* Qwen2.5-Coder -> google/gemma-4-E2B (memory_lora/core.py target modules)
* cuda + flash_attn2 -> mps + sdpa (falls back to eager)
* no wandb/TRL -> plain PyTorch loop + TensorBoard (SummaryWriter)
Same core trick as the paper: only the hypernetwork head is trained; the
base LLM is frozen (gradient-checkpointed for memory headroom); LoRA A/B
tensors are non-detached so the causal-LM loss's backward graph flows
straight into the head's parameters.
Usage:
python scripts/train_memory_lora.py --output-dir runs/pilot1 \\
--limit-train-docs 5 --epochs 1 # smoke test
python scripts/train_memory_lora.py --output-dir runs/full1
"""
from __future__ import annotations
import argparse
import json
import random
import sys
import time
from pathlib import Path
from typing import Any, Dict, List, Optional, Tuple
import numpy as np
import psutil
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch.utils.tensorboard import SummaryWriter
from transformers import (
AutoModelForImageTextToText,
AutoTokenizer,
get_cosine_schedule_with_warmup,
)
HERE = Path(__file__).resolve().parent
REPO_ROOT = HERE.parent
sys.path.insert(0, str(REPO_ROOT))
from memory_lora.data_paths import EMBEDDINGS_DIR, QNA_DIR, RUNS_DIR, ensure_dirs # noqa: E402
from memory_lora.core import ( # noqa: E402
MemoryLoRAHead,
DEFAULT_ROOT_PREFIX,
discover_module_types_and_dims,
get_module_specs,
inject_lora_weights,
load_doc_rows,
load_qna_rows,
replace_with_lora,
)
DEFAULT_MODEL = "google/gemma-4-E2B"
DEFAULT_TARGET_MODULES = [
"q_proj", "k_proj", "v_proj", "o_proj",
"up_proj", "gate_proj", "down_proj",
]
# ---------------------------------------------------------------------------
# Dataset & batching
# ---------------------------------------------------------------------------
class DocDataset:
"""One example = one document with its train-split QnAs."""
def __init__(
self,
docs_by_id: Dict[str, Dict[str, Any]],
qnas_by_doc: Dict[str, List[Dict[str, str]]],
doc_ids: List[str],
max_qna_per_doc: int = 32,
seed: int = 3407,
):
self.doc_ids = list(doc_ids)
self.docs = docs_by_id
self.qnas = qnas_by_doc
self.max_qna = max_qna_per_doc
self.rng = random.Random(seed)
def __len__(self) -> int:
return len(self.doc_ids)
def __getitem__(self, idx: int) -> Optional[Dict[str, Any]]:
d = self.doc_ids[idx]
pairs = list(self.qnas.get(d, []))
if not pairs:
return None
if len(pairs) > self.max_qna:
pairs = self.rng.sample(pairs, self.max_qna)
return {"doc_id": d, "embedding": self.docs[d]["emb"], "qnas": pairs}
def _tokenize_lm_batch(tokenizer, prefixes: List[str], targets: List[str],
max_seq_len: int = 2048,
fixed_len: bool = False) -> Dict[str, torch.Tensor]:
"""Causal-LM batch with the loss masked on prefix tokens. Keeps the
rightmost prefix tokens on overflow; targets are never truncated.
fixed_len: pad every batch to EXACTLY max_seq_len instead of the
batch's own local max length. Real code prefixes vary widely (100-1024+
tokens across different repos), so per-batch padding produces a new
tensor shape almost every document. On MPS this repeatedly triggered
unbounded memory growth (observed: a run that stayed under 13GB on
homogeneous-length synthetic docs hit 70+GB and got OS-killed within
~10 documents of real, variable-length code) -- MPS's caching allocator
does not appear to reliably reclaim/reuse blocks across many distinct
shapes the way CUDA's does. Fixing every batch to one shape avoids the
allocator ever seeing a new size after the first batch. Slightly wastes
compute on padding for short sequences; that trade is worth it for
system stability.
"""
eos = tokenizer.eos_token or ""
input_ids_list: List[torch.Tensor] = []
labels_list: List[torch.Tensor] = []
for p, t in zip(prefixes, targets):
t_ids = tokenizer(t + eos, add_special_tokens=False)["input_ids"]
if not t_ids:
continue
prefix_budget = max(8, max_seq_len - len(t_ids))
p_ids_full = tokenizer(p, add_special_tokens=False)["input_ids"]
p_ids = p_ids_full[-prefix_budget:] if len(p_ids_full) > prefix_budget else p_ids_full
ids = p_ids + t_ids
labels = ([-100] * len(p_ids)) + list(t_ids)
input_ids_list.append(torch.tensor(ids, dtype=torch.long))
labels_list.append(torch.tensor(labels, dtype=torch.long))
if not input_ids_list:
return {}
local_max = max(t.size(0) for t in input_ids_list)
L = max(max_seq_len, local_max) if fixed_len else local_max
pad_id = tokenizer.pad_token_id or 0
def _lpad(x, val):
return F.pad(x, (L - x.size(0), 0), value=val)
input_ids = torch.stack([_lpad(t, pad_id) for t in input_ids_list], 0)
labels = torch.stack([_lpad(t, -100) for t in labels_list], 0)
attn_list = [torch.ones(t.size(0), dtype=torch.long) for t in input_ids_list]
attn = torch.stack([_lpad(t, 0) for t in attn_list], 0)
return {"input_ids": input_ids, "labels": labels, "attention_mask": attn}
# ---------------------------------------------------------------------------
# Eval
# ---------------------------------------------------------------------------
@torch.no_grad()
def evaluate_suite(
base_model: nn.Module, head: MemoryLoRAHead, specs, tokenizer,
doc_rows: List[Any], qnas_by_doc: Dict[str, List[Dict[str, str]]],
*, device: torch.device, max_seq_len: int = 512,
lm_micro_batch: int = 4, max_qna_per_doc: int = 32,
fixed_len: bool = False, with_baseline: bool = False,
) -> Dict[str, float]:
"""Evaluate the adapted model, and (with ``with_baseline``) the SAME model
with no adapter injected.
The baseline is the metric that actually matters: an eval loss of 2.6 says
nothing on its own, because it does not reveal whether the generated
adapter is helping, doing nothing, or actively hurting. Tracking only the
adapted loss is how a head that had collapsed to emitting one constant,
worse-than-random adapter for every repo went unnoticed. ``delta`` below is
the number to watch: it must go NEGATIVE and stay there.
"""
base_model.eval()
head.eval()
total_loss = 0.0
total_tokens = 0
n_docs = 0
base_loss_total = 0.0
base_tokens = 0
for dr in doc_rows:
pairs = qnas_by_doc.get(dr.doc_id)
if not pairs:
continue
if len(pairs) > max_qna_per_doc:
pairs = pairs[:max_qna_per_doc]
ctx = torch.from_numpy(dr.doc_embedding).to(device).unsqueeze(0)
head_out = head(ctx)
if with_baseline:
# Detach the adapter (A=B=None -> LoRA.forward returns base output)
# and score the identical batches through the frozen model.
for sp in specs:
m = dict(base_model.named_modules())[sp.full_name]
m.A, m.B = None, None
prefixes_b = [p["prefix"] for p in pairs]
targets_b = [p["target"] for p in pairs]
for i in range(0, len(prefixes_b), lm_micro_batch):
j = min(i + lm_micro_batch, len(prefixes_b))
b = _tokenize_lm_batch(tokenizer, prefixes_b[i:j], targets_b[i:j],
max_seq_len=max_seq_len, fixed_len=fixed_len)
if not b:
continue
b = {k: v.to(device) for k, v in b.items()}
o = base_model(**b)
nt = (b["labels"] != -100).sum().item()
base_loss_total += o.loss.item() * nt
base_tokens += nt
inject_lora_weights(base_model, specs, head_out, batch_index=0)
prefixes = [p["prefix"] for p in pairs]
targets = [p["target"] for p in pairs]
for i in range(0, len(prefixes), lm_micro_batch):
j = min(i + lm_micro_batch, len(prefixes))
batch = _tokenize_lm_batch(tokenizer, prefixes[i:j], targets[i:j],
max_seq_len=max_seq_len, fixed_len=fixed_len)
if not batch:
continue
batch = {k: v.to(device) for k, v in batch.items()}
out = base_model(**batch)
loss = out.loss
ntok = (batch["labels"] != -100).sum().item()
total_loss += loss.item() * ntok
total_tokens += ntok
n_docs += 1
avg = total_loss / max(total_tokens, 1)
out: Dict[str, float] = {"eval_loss": avg, "n_docs": n_docs,
"n_tokens": total_tokens}
if with_baseline and base_tokens:
base_avg = base_loss_total / base_tokens
out["baseline_loss"] = base_avg
out["delta_vs_baseline"] = avg - base_avg # negative == adapter helps
return out
@torch.no_grad()
def adapter_input_sensitivity(head: MemoryLoRAHead, doc_rows, device,
n: int = 16) -> Dict[str, float]:
"""Mean pairwise cosine between the adapters generated for different repos.
~1.0 means the head ignores its input and emits one constant adapter (the
failure mode that made a trained head score worse than random noise); low
values mean the emitted adapter is genuinely repo-conditional.
"""
rows = doc_rows[:n]
if len(rows) < 2:
return {}
ctx = torch.from_numpy(
np.stack([r.doc_embedding for r in rows])).to(device)
o = head(ctx)
t = sorted(o["A"].keys())[0]
D = torch.einsum("nor,nri->noi", o["B"][t].float(), o["A"][t].float()).flatten(1)
Dn = F.normalize(D, dim=1)
C = Dn @ Dn.T
iu = torch.triu_indices(len(D), len(D), 1)
return {"adapter_cosine": float(C[iu[0], iu[1]].mean()),
"adapter_delta_fro": float(D.norm(dim=1).mean())}
# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------
def _docs_by_id(doc_rows) -> Dict[str, Dict[str, Any]]:
return {dr.doc_id: {"emb": dr.doc_embedding} for dr in doc_rows}
def _group_qnas_by_doc(rows) -> Dict[str, List[Dict[str, str]]]:
out: Dict[str, List[Dict[str, str]]] = {}
for qr in rows:
out.setdefault(qr.doc_id, []).append({"prefix": qr.prefix, "target": qr.target})
return out
def main() -> None:
ap = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("--embeddings-path", default=str(EMBEDDINGS_DIR / "doc_embeddings.parquet"))
ap.add_argument("--qna-path", default=str(QNA_DIR / "qna.jsonl"))
ap.add_argument("--output-dir", required=True)
ap.add_argument("--model-name", default=DEFAULT_MODEL)
ap.add_argument("--target-modules", nargs="+", default=DEFAULT_TARGET_MODULES)
ap.add_argument("--root-prefix", default=DEFAULT_ROOT_PREFIX)
ap.add_argument("--rank", type=int, default=16)
ap.add_argument("--alpha", type=float, default=32.0)
ap.add_argument("--head-hidden-dim", type=int, default=128,
help="Kept small deliberately -- with ~165 training "
"docs, the paper's 512-1024 hidden dim overfits "
"within ~2 epochs (see memory_lora/core.py docstring).")
ap.add_argument("--head-dropout", type=float, default=0.1)
ap.add_argument("--epochs", type=int, default=3)
ap.add_argument("--lr", type=float, default=1e-4)
ap.add_argument("--weight-decay", type=float, default=0.05)
ap.add_argument("--warmup-ratio", type=float, default=0.03)
ap.add_argument("--lr-total-steps", type=int, default=0,
help="Override the cosine LR schedule's total-step "
"target with a realistic estimate of what "
"--max-hours will actually cover, instead of "
"steps_per_epoch * epochs (which assumes the run "
"finishes a full epoch -- unrealistic at "
"tens-of-thousands-of-docs scale). 0 = use the "
"epoch-based calculation.")
ap.add_argument("--max-grad-norm", type=float, default=1.0)
ap.add_argument("--early-stop-patience", type=int, default=8,
help="Stop after this many consecutive evals with no "
"improvement on --primary-eval-suite. 0 = disabled.")
ap.add_argument("--max-qna-per-doc", type=int, default=32)
ap.add_argument("--lm-micro-batch", type=int, default=4)
ap.add_argument("--max-seq-len", type=int, default=512)
ap.add_argument("--fixed-seq-len", action="store_true", default=True,
help="Pad every batch to exactly --max-seq-len instead "
"of each batch's own local max length. See "
"_tokenize_lm_batch docstring: on MPS, varying "
"tensor shapes across many real-code documents "
"of wildly different lengths caused unbounded "
"memory growth (a run went from healthy to "
"OS-killed within ~10 documents). Costs some "
"wasted padding compute; worth it for stability.")
ap.add_argument("--no-fixed-seq-len", dest="fixed_seq_len", action="store_false")
ap.add_argument("--eval-every-steps", type=int, default=50)
ap.add_argument("--eval-suites", nargs="+", default=["cr_val", "cr_test", "ir_test"])
ap.add_argument("--limit-eval-docs", type=int, default=200,
help="Cap docs per eval suite (random sample, fixed "
"seed) for speed at real-corpus scale -- e.g. "
"cr_val alone can be 8,600+ real repo-commit "
"docs; evaluating all of them every eval cycle "
"would dominate wall-clock time. 0 = no cap "
"(matches the paper's own --limit-eval-snapshots).")
ap.add_argument("--primary-eval-suite", default="cr_val")
ap.add_argument("--log-every-iters", type=int, default=10)
ap.add_argument("--seed", type=int, default=3407)
ap.add_argument("--device", default="mps")
ap.add_argument("--dtype", default="bfloat16", choices=["bfloat16", "float16", "float32"])
ap.add_argument("--attn-implementation", default="sdpa", choices=["sdpa", "eager"])
ap.add_argument("--limit-train-docs", type=int, default=0)
ap.add_argument("--priority-doc-ids", nargs="+", default=[],
help="doc_ids to oversample within the shared "
"multi-document head (--priority-oversample "
"extra passes per epoch), instead of raising "
"global head capacity -- raising rank/hidden_dim "
"fixes within-document fact interference but "
"makes the whole corpus overfit faster (see "
"runs/full2 diagnosis); oversampling gives "
"specific documents more gradient signal without "
"changing capacity or hurting the rest.")
ap.add_argument("--priority-oversample", type=int, default=5,
help="How many times to repeat each --priority-doc-ids "
"entry per epoch's shuffled training order.")
ap.add_argument("--only-doc-ids", nargs="+", default=[],
help="Restrict training (and, if present in this set, "
"ir_test eval) to exactly these doc_ids. Used for "
"single-document capacity diagnostics -- e.g. can "
"the architecture memorize ONE document's facts "
"when not sharing hypernetwork capacity across "
"165 others?")
ap.add_argument("--gradient-checkpointing", action="store_true", default=True)
ap.add_argument("--no-gradient-checkpointing", dest="gradient_checkpointing", action="store_false")
ap.add_argument("--max-hours", type=float, default=0.0,
help="Wall-clock training budget in hours. 0 = unlimited "
"(stop only after --epochs). Checked once per doc "
"iteration; when exceeded, saves a final checkpoint "
"and stops cleanly (does not just get killed mid-write).")
ap.add_argument("--min-available-gb", type=float, default=15.0,
help="Hard safety floor: stop (with a final "
"checkpoint) if SYSTEM-WIDE available memory "
"(psutil.virtual_memory().available -- NOT this "
"process's own RSS, which undercounts MPS "
"memory on Apple Silicon) drops below this many "
"GB. Checked every 2 iterations. 0 = disabled.")
ap.add_argument("--checkpoint-every-steps", type=int, default=50,
help="Overwrite head.latest.pt every N optimizer steps "
"so a crash/kill never loses more than N steps of "
"progress. Overwrites (doesn't accumulate files), "
"so it's disk-safe even for a 3GB head. 0=disabled.")
ap.add_argument("--checkpoint-every-minutes", type=float, default=30.0,
help="Save a timestamped checkpoint every N minutes of "
"wall-clock time, independent of eval/epoch "
"boundaries. 0 = disabled (epoch-end saves only).")
ap.add_argument("--epoch-ckpt-every", type=int, default=10,
help="Only write a NEW numbered head.epN.pt every N "
"epochs (head.latest.pt still updates every "
"epoch). Each checkpoint is a full head save "
"(hundreds of MB) -- with many small/fast epochs "
"(e.g. a tiny single-document run), saving one "
"per epoch can fill the disk in minutes.")
ap.add_argument("--no-eval-baseline", action="store_true",
help="skip the no-adapter baseline during eval (faster, but "
"you lose the only signal that says whether the "
"adapter is actually helping)")
ap.add_argument("--resume-from", default="",
help="Path to a head.*.pt checkpoint to load weights "
"from before training starts (optimizer/scheduler "
"restart fresh; only head weights are resumed).")
args = ap.parse_args()
out_dir = Path(args.output_dir)
if not out_dir.is_absolute():
out_dir = RUNS_DIR / out_dir
out_dir.mkdir(parents=True, exist_ok=True)
ensure_dirs()
device = torch.device(args.device if (args.device != "mps" or torch.backends.mps.is_available()) else "cpu")
dtype = {"bfloat16": torch.bfloat16, "float16": torch.float16, "float32": torch.float32}[args.dtype]
random.seed(args.seed)
np.random.seed(args.seed)
torch.manual_seed(args.seed)
tb = SummaryWriter(log_dir=str(out_dir / "tb"))
# ---- Load embeddings + QnAs ----
print("Loading document embeddings ...", flush=True)
all_docs = load_doc_rows(Path(args.embeddings_path))
only_ids = set(args.only_doc_ids) if args.only_doc_ids else None
train_docs = [d for d in all_docs if d.split == "train"]
if only_ids:
train_docs = [d for d in train_docs if d.doc_id in only_ids]
if args.limit_train_docs:
train_docs = train_docs[: args.limit_train_docs]
print(f" {len(train_docs)} train docs (of {len(all_docs)} total)", flush=True)
print("Loading QnAs ...", flush=True)
all_qnas = load_qna_rows(Path(args.qna_path))
if only_ids:
all_qnas = [q for q in all_qnas if q.doc_id in only_ids]
train_qnas = [q for q in all_qnas if q.qna_split == "train"]
qnas_train = _group_qnas_by_doc(train_qnas)
docs_by_id = _docs_by_id(train_docs)
doc_ids = [d for d in docs_by_id if d in qnas_train]
print(f" {sum(len(v) for v in qnas_train.values())} train QA pairs across {len(doc_ids)} docs", flush=True)
if args.priority_doc_ids and args.priority_oversample > 1:
extra = []
for pid in args.priority_doc_ids:
if pid in doc_ids:
extra.extend([pid] * (args.priority_oversample - 1))
else:
print(f" [warn] --priority-doc-ids {pid!r} not in training set, skipping", flush=True)
doc_ids = doc_ids + extra
print(f" oversampled {args.priority_doc_ids} x{args.priority_oversample} "
f"-> {len(doc_ids)} entries/epoch", flush=True)
ds = DocDataset(docs_by_id, qnas_train, doc_ids,
max_qna_per_doc=args.max_qna_per_doc, seed=args.seed)
# ---- Build LLM, discover modules, wrap them ----
print(f"Loading {args.model_name} ...", flush=True)
tokenizer = AutoTokenizer.from_pretrained(args.model_name)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
base_model = AutoModelForImageTextToText.from_pretrained(
args.model_name, torch_dtype=dtype,
attn_implementation=args.attn_implementation,
).to(device)
base_model.eval()
for p in base_model.parameters():
p.requires_grad = False
if args.gradient_checkpointing:
base_model.config.use_cache = False
try:
base_model.gradient_checkpointing_enable(
gradient_checkpointing_kwargs={"use_reentrant": False})
print(" gradient checkpointing: ON", flush=True)
except Exception as e: # noqa: BLE001
print(f" [warn] gradient checkpointing unavailable: {e}", flush=True)
specs = get_module_specs(base_model, args.target_modules, root_prefix=args.root_prefix)
type_dims = discover_module_types_and_dims(specs)
print(f" discovered {len(specs)} target modules, {len(type_dims)} types: {sorted(type_dims)}", flush=True)
if not specs:
raise SystemExit(
f"No modules matched root_prefix={args.root_prefix!r} + "
f"{args.target_modules}. Inspect base_model.named_modules() and "
f"pass --root-prefix explicitly."
)
replace_with_lora(base_model, specs, rank=args.rank, alpha=args.alpha)
head = MemoryLoRAHead(
input_dim=train_docs[0].doc_embedding.shape[0],
type_dims=type_dims,
hidden_dim=args.head_hidden_dim,
rank=args.rank,
dropout=args.head_dropout,
).to(device)
# Standardize the conditioning input using TRAIN docs only. Without this the
# ~64% DC component shared by all repo embeddings dominates the trunk and it
# emits a near-identical adapter for every repo (see MemoryLoRAHead docs).
head.fit_input_stats(torch.from_numpy(
np.stack([d.doc_embedding for d in train_docs])).to(device))
print(f" fitted input standardization over {len(train_docs)} train docs",
flush=True)
if args.resume_from:
# weights_only=False: these checkpoints carry the run's config/args dicts,
# not just tensors, and torch>=2.6 defaults the strict unpickler on --
# which rejects them ("Unsupported operand"). They are produced by this
# project's own training script, so loading them fully is intended.
ckpt = torch.load(args.resume_from, map_location=device, weights_only=False)
# strict=False: checkpoints written before input standardization
# existed carry no input_mean/input_std, so the stats fitted above are
# kept. A checkpoint that does carry them overwrites the fresh fit,
# which is what a resumed run wants -- the transform must not change
# mid-training.
missing, _ = head.load_state_dict(ckpt["state_dict"], strict=False)
if missing:
print(f" (new buffers not in checkpoint: {missing})", flush=True)
print(f" resumed head weights from {args.resume_from}", flush=True)
n_head_params = sum(p.numel() for p in head.parameters())
print(f" head params: {n_head_params / 1e6:.1f}M", flush=True)
optim = torch.optim.AdamW(head.parameters(), lr=args.lr, weight_decay=args.weight_decay)
steps_per_epoch = max(1, len(ds))
if args.lr_total_steps:
# At real-corpus scale (tens of thousands of docs), --max-hours will
# cut training off long before steps_per_epoch * epochs is reached,
# so a schedule calibrated to full-epoch coverage would barely start
# annealing from its LR peak. Calibrate to the realistically
# achievable step count instead.
total_steps = args.lr_total_steps
else:
total_steps = steps_per_epoch * args.epochs
warmup_steps = max(1, int(total_steps * args.warmup_ratio))
sched = get_cosine_schedule_with_warmup(optim, warmup_steps, total_steps)
# ---- Eval suites ----
eval_suites: Dict[str, Dict[str, Any]] = {}
print("Loading eval suites ...", flush=True)
qnas_by_doc_all = _group_qnas_by_doc(all_qnas)
qnas_held_out_by_doc = _group_qnas_by_doc([q for q in all_qnas if q.qna_split == "held_out"])
eval_rng = random.Random(args.seed)
for suite in args.eval_suites:
if suite in ("cr_val", "cr_test"):
rows = [d for d in all_docs if d.split == suite]
q_by_doc = qnas_by_doc_all
elif suite == "ir_test":
rows = train_docs
q_by_doc = qnas_held_out_by_doc
else:
continue
if args.limit_eval_docs and len(rows) > args.limit_eval_docs:
rows = eval_rng.sample(rows, args.limit_eval_docs)
eval_suites[suite] = {"doc_rows": rows, "qnas_by_doc": q_by_doc}
n_q = sum(len(q_by_doc.get(d.doc_id, [])) for d in rows)
print(f" {suite}: {len(rows)} docs, {n_q} qnas", flush=True)
# ---- Train ----
metrics_log: List[Dict[str, Any]] = []
best_eval = float("inf")
global_step = 0
t0 = time.time()
last_ckpt_wall = t0
budget_seconds = args.max_hours * 3600.0 if args.max_hours > 0 else float("inf")
ckpt_interval_seconds = args.checkpoint_every_minutes * 60.0 if args.checkpoint_every_minutes > 0 else float("inf")
stop_training = False
patience_counter = 0
for epoch in range(args.epochs):
if stop_training:
break
order = list(range(len(ds)))
random.shuffle(order)
head.train()
running_loss, running_n = 0.0, 0
for it, di in enumerate(order):
now = time.time()
if now - t0 >= budget_seconds:
print(f" [budget] {args.max_hours:.2f}h training budget reached "
f"(epoch {epoch}, it {it}/{len(order)}) -- stopping.", flush=True)
stop_training = True
break
if now - last_ckpt_wall >= ckpt_interval_seconds:
mins = int((now - t0) / 60)
p = _save_ckpt(out_dir, head, type_dims, args, name=f"t{mins:04d}m")
_save_ckpt(out_dir, head, type_dims, args, name="latest")
print(f" [ckpt] periodic ({args.checkpoint_every_minutes:.0f}min interval) -> {p}", flush=True)
last_ckpt_wall = now
if args.min_available_gb > 0 and it % 2 == 0:
# IMPORTANT: this checks SYSTEM-WIDE available memory
# (psutil.virtual_memory), not this process's own RSS.
# psutil.Process().memory_info().rss -- like `ps -o rss` --
# does NOT reliably capture MPS/GPU-resident allocations
# on Apple Silicon: a run was observed at 55-83GB actual
# usage (per `top`'s MEM column, corroborated by system
# vm_stat showing genuine memory exhaustion) while RSS
# reported under 1GB the whole time. System-wide available
# memory is the metric that's actually reliable here.
available_gb = psutil.virtual_memory().available / 1e9
if available_gb < args.min_available_gb:
print(f" [safety] system available memory {available_gb:.1f}GB "
f"below --min-available-gb {args.min_available_gb:.1f}GB "
f"(epoch {epoch}, it {it}) -- saving and stopping to "
f"protect system stability.", flush=True)
_save_ckpt(out_dir, head, type_dims, args, name="latest")
stop_training = True
break
sample = ds[di]
if sample is None:
continue
ctx = torch.from_numpy(sample["embedding"]).to(device).unsqueeze(0)
qnas = sample["qnas"]
prefixes = [q["prefix"] for q in qnas]
targets = [q["target"] for q in qnas]
micro_batches = []
for i in range(0, len(prefixes), args.lm_micro_batch):
j = min(i + args.lm_micro_batch, len(prefixes))
b = _tokenize_lm_batch(tokenizer, prefixes[i:j], targets[i:j],
max_seq_len=args.max_seq_len, fixed_len=args.fixed_seq_len)
if b:
micro_batches.append({k: v.to(device) for k, v in b.items()})
if not micro_batches:
continue
n_tok_seen, loss_acc = 0, 0.0
for mb_idx, batch in enumerate(micro_batches):
if args.min_available_gb > 0 and mb_idx % 3 == 0:
# Same system-wide check as the per-document one below,
# but INSIDE the micro-batch loop too: observed runaway
# growth can blow past a safe threshold within a
# single document's micro-batches, before the
# per-document check would ever fire.
available_gb = psutil.virtual_memory().available / 1e9
if available_gb < args.min_available_gb:
print(f" [safety] system available memory {available_gb:.1f}GB "
f"below --min-available-gb {args.min_available_gb:.1f}GB "
f"mid-document (epoch {epoch}, it {it}, micro-batch "
f"{mb_idx}) -- saving and stopping immediately.", flush=True)
_save_ckpt(out_dir, head, type_dims, args, name="latest")
stop_training = True
break
head_out = head(ctx)
inject_lora_weights(base_model, specs, head_out, batch_index=0)
out = base_model(**batch)
ntok = (batch["labels"] != -100).sum().item()
loss = out.loss * ntok
loss.backward()
loss_acc += loss.detach().item()
n_tok_seen += ntok
del head_out, out, loss
if stop_training:
break
if n_tok_seen == 0:
continue
if device.type == "mps" and it % 5 == 0:
# MPS's caching allocator is markedly less aggressive about
# returning freed blocks to the OS than CUDA's -- on a
# unified-memory Mac (CPU and GPU share physical RAM,
# unlike a discrete-GPU box with isolated VRAM) that cache
# growth directly threatens the whole system, not just this
# process. Without this, a real-corpus run OOM'd the OS
# itself (83GB RSS, process state "stuck", heavy swapping)
# within the first ~10 minutes.
torch.mps.empty_cache()
torch.nn.utils.clip_grad_norm_(head.parameters(), args.max_grad_norm)
optim.step()
sched.step()
optim.zero_grad(set_to_none=True)
global_step += 1
if args.checkpoint_every_steps > 0 and global_step % args.checkpoint_every_steps == 0:
_save_ckpt(out_dir, head, type_dims, args, name="latest")
running_loss += loss_acc
running_n += n_tok_seen
if it % max(1, args.log_every_iters) == 0:
avg = running_loss / max(running_n, 1)
elapsed = (time.time() - t0) / 60
print(f"[ep{epoch} it{it}/{len(order)} step{global_step}] "
f"loss={avg:.4f} lr={sched.get_last_lr()[0]:.2e} elapsed={elapsed:.1f}m", flush=True)
tb.add_scalar("train/loss", avg, global_step)
tb.add_scalar("train/lr", sched.get_last_lr()[0], global_step)
running_loss, running_n = 0.0, 0
if (args.eval_every_steps > 0 and global_step > 0
and global_step % args.eval_every_steps == 0
and it + 1 != len(order)):
# The `it + 1 != len(order)` guard skips a redundant eval when
# eval_every_steps happens to equal (a multiple of) steps-per-
# epoch -- the unconditional end-of-epoch eval below would
# otherwise double-count this exact step in the early-stop
# patience counter every epoch.
prev_best = best_eval
best_eval = _do_eval(args, base_model, head, specs, tokenizer, eval_suites,
device, out_dir, metrics_log, best_eval, global_step, epoch, tb)
patience_counter = 0 if best_eval < prev_best else patience_counter + 1
if args.early_stop_patience > 0 and patience_counter >= args.early_stop_patience:
print(f" [early-stop] no improvement on {args.primary_eval_suite} for "
f"{patience_counter} evals -- stopping.", flush=True)
stop_training = True
break
_save_ckpt(out_dir, head, type_dims, args, name="latest")
if epoch % max(1, args.epoch_ckpt_every) == 0:
ep_path = _save_ckpt(out_dir, head, type_dims, args, name=f"ep{epoch}")
print(f" [ckpt] end-of-epoch ep{epoch} -> {ep_path}", flush=True)
else:
print(f" [ckpt] end-of-epoch ep{epoch} -> (latest.pt only)", flush=True)
prev_best = best_eval
best_eval = _do_eval(args, base_model, head, specs, tokenizer, eval_suites,
device, out_dir, metrics_log, best_eval, global_step, epoch, tb,
end_of_epoch=True)
patience_counter = 0 if best_eval < prev_best else patience_counter + 1
if args.early_stop_patience > 0 and patience_counter >= args.early_stop_patience:
print(f" [early-stop] no improvement on {args.primary_eval_suite} for "
f"{patience_counter} evals -- stopping.", flush=True)
stop_training = True
tb.close()
print(f"\nTraining done. Best primary eval = {best_eval:.4f}", flush=True)
def _save_ckpt(out_dir: Path, head: MemoryLoRAHead, type_dims, args, name: str = "latest") -> Path:
out = out_dir / f"head.{name}.pt"
torch.save({
"state_dict": head.state_dict(),
"config": head.config_dict(),
"type_dims": type_dims,
"args": vars(args),
}, out)
return out
def _do_eval(args, base_model, head, specs, tokenizer, eval_suites, device, out_dir,
metrics_log, best_eval, global_step, epoch, tb, end_of_epoch: bool = False) -> float:
suite_metrics: Dict[str, Dict[str, float]] = {}
for name, suite in eval_suites.items():
m = evaluate_suite(base_model, head, specs, tokenizer,
suite["doc_rows"], suite["qnas_by_doc"],
device=device, max_seq_len=args.max_seq_len,
lm_micro_batch=args.lm_micro_batch,
max_qna_per_doc=args.max_qna_per_doc,
fixed_len=args.fixed_seq_len,
with_baseline=not args.no_eval_baseline)
suite_metrics[name] = m
delta_s = ""
if "delta_vs_baseline" in m:
verdict = "HELPS" if m["delta_vs_baseline"] < 0 else "HURTS"
delta_s = (f" | base={m['baseline_loss']:.4f} "
f"delta={m['delta_vs_baseline']:+.4f} {verdict}")
print(f" [eval {name}] step={global_step} loss={m['eval_loss']:.4f} "
f"docs={m['n_docs']} tok={m['n_tokens']}{delta_s}", flush=True)
tb.add_scalar(f"eval/{name}_loss", m["eval_loss"], global_step)
if "baseline_loss" in m:
tb.add_scalar(f"eval/{name}_baseline_loss", m["baseline_loss"], global_step)
# THE metric to watch: must be negative for the adapter to be useful.
tb.add_scalar(f"eval/{name}_delta_vs_baseline",
m["delta_vs_baseline"], global_step)
# Input-sensitivity diagnostic: is the head emitting repo-conditional
# adapters, or one constant adapter regardless of input?
any_suite = next(iter(eval_suites.values()), None)
if any_suite:
diag = adapter_input_sensitivity(head, any_suite["doc_rows"], device)
for k, v in diag.items():
tb.add_scalar(f"diag/{k}", v, global_step)
if diag:
print(f" [diag] adapter_cosine={diag['adapter_cosine']:.4f} "
f"(1.0 = same adapter for every repo) "
f"delta_fro={diag['adapter_delta_fro']:.3f}", flush=True)
suite_metrics["_diag"] = diag
primary = suite_metrics.get(args.primary_eval_suite)
primary_loss = primary["eval_loss"] if primary else float("inf")
row = {"step": global_step, "epoch": epoch, "end_of_epoch": end_of_epoch,
"eval_loss": primary_loss, "suites": suite_metrics}
metrics_log.append(row)
(out_dir / "metrics.jsonl").open("a").write(json.dumps(row) + "\n")
if primary_loss < best_eval:
best_eval = primary_loss
p = _save_ckpt(out_dir, head, head.type_dims, args, name="best")
print(f" [ckpt] best updated -> {p} (loss={primary_loss:.4f})", flush=True)
_save_ckpt(out_dir, head, head.type_dims, args, name="latest")
head.train()
return best_eval
if __name__ == "__main__":
main()
|