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: 9,994 Bytes
481fbb6 930bb27 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 | #!/usr/bin/env python3
"""Evaluate a trained Memory-LoRA head: recall EM/EditSim on held-out QA,
split into in-corpus (ir_test) and cross-corpus (cr_val/cr_test), plus a
manual spot-check comparing the LoRA-adapted model against the bare base
model on hand-picked Code2LoRA-paper questions (proof the adapter, not
general pretraining, is doing the recall).
Forked from Code2LoRA's evaluation metrics (EM after whitespace collapsing
+ trailing-punctuation removal with relaxed prefix matching; EditSim via
difflib.SequenceMatcher).
Usage:
python scripts/eval_memory_lora.py --checkpoint runs/full1/head.best.pt
"""
from __future__ import annotations
import argparse
import difflib
import json
import re
import sys
from pathlib import Path
from typing import Any, Dict, List
import torch
from transformers import AutoModelForImageTextToText, AutoTokenizer
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 # noqa: E402
from memory_lora.core import ( # noqa: E402
MemoryLoRAHead,
DEFAULT_ROOT_PREFIX,
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",
]
SPOT_CHECK_QUESTIONS = [
("What LoRA rank does Code2LoRA's static hypernetwork use?", "16"),
("How many trainable parameters does Code2LoRA-Static have?", "720 million"),
("What is the cross-repo exact match of Code2LoRA-Static on the static track?", "63.8%"),
("How many Python repositories are in RepoPeftBench?", "604"),
("What is the base LLM used in Code2LoRA's experiments?", "Qwen2.5-Coder-1.5B"),
]
def normalize(s: str) -> str:
s = s.strip().rstrip(".,;:!?")
s = re.sub(r"\s+", " ", s)
return s.lower()
def exact_match(pred: str, target: str) -> bool:
p, t = normalize(pred), normalize(target)
return p == t or p.startswith(t) or t.startswith(p)
def edit_sim(pred: str, target: str) -> float:
return difflib.SequenceMatcher(None, normalize(pred), normalize(target)).ratio()
@torch.no_grad()
def generate(base_model, tokenizer, prefix: str, device, max_new_tokens: int = 12) -> str:
"""Greedy-decode the answer, then truncate at the first newline.
Without this the model frequently keeps going past the answer into a
hallucinated ``\\nQ: <next question>`` continuation (base models without
an EOS-triggering chat template rarely stop cleanly on a bare
completion prompt); comparing the *untruncated* string against the gold
answer would mark an otherwise-correct short answer wrong just because
of what it rambled into afterward.
"""
enc = tokenizer(prefix, return_tensors="pt").to(device)
out = base_model.generate(
**enc, max_new_tokens=max_new_tokens, do_sample=False,
pad_token_id=tokenizer.pad_token_id or tokenizer.eos_token_id,
)
gen_ids = out[0][enc["input_ids"].shape[1]:]
text = tokenizer.decode(gen_ids, skip_special_tokens=True)
return text.split("\n")[0]
def load_head_and_model(checkpoint: Path, model_name: str, target_modules: List[str],
root_prefix: str, device: torch.device, dtype: torch.dtype,
attn_implementation: str):
# 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(checkpoint, map_location="cpu", weights_only=False)
tokenizer = AutoTokenizer.from_pretrained(model_name)
if tokenizer.pad_token is None:
tokenizer.pad_token = tokenizer.eos_token
base_model = AutoModelForImageTextToText.from_pretrained(
model_name, torch_dtype=dtype, attn_implementation=attn_implementation,
).to(device)
base_model.eval()
for p in base_model.parameters():
p.requires_grad = False
specs = get_module_specs(base_model, target_modules, root_prefix=root_prefix)
rank = ckpt["config"]["rank"]
alpha = ckpt["args"].get("alpha", 32.0)
replace_with_lora(base_model, specs, rank=rank, alpha=alpha)
head = MemoryLoRAHead(
input_dim=ckpt["config"]["input_dim"],
type_dims={k: tuple(v) for k, v in ckpt["config"]["type_dims"].items()},
hidden_dim=ckpt["config"]["hidden_dim"],
rank=rank,
).to(device)
head.load_state_dict(ckpt["state_dict"])
head.eval()
return base_model, head, specs, tokenizer
def eval_suite(base_model, head, specs, tokenizer, doc_rows, qnas_by_doc, device,
max_qna_per_doc: int = 20) -> Dict[str, float]:
n_em, n_total, sum_editsim = 0, 0, 0.0
for dr in doc_rows:
pairs = qnas_by_doc.get(dr.doc_id, [])[:max_qna_per_doc]
if not pairs:
continue
ctx = torch.from_numpy(dr.doc_embedding).to(device).unsqueeze(0)
head_out = head(ctx)
inject_lora_weights(base_model, specs, head_out, batch_index=0)
for p in pairs:
pred = generate(base_model, tokenizer, p["prefix"], device)
target = p["target"]
if exact_match(pred, target):
n_em += 1
sum_editsim += edit_sim(pred, target)
n_total += 1
return {
"em": n_em / max(n_total, 1),
"editsim": sum_editsim / max(n_total, 1),
"n": n_total,
}
def main() -> None:
ap = argparse.ArgumentParser()
ap.add_argument("--checkpoint", required=True)
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("--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("--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("--suites", nargs="+", default=["cr_val", "cr_test", "ir_test"])
ap.add_argument("--max-qna-per-doc", type=int, default=20)
ap.add_argument("--limit-docs", type=int, default=0,
help="Random-sample at most N docs per suite (fixed seed) "
"for a fast estimate -- cr_test has thousands of real "
"repos, far too many to greedy-generate on CPU.")
ap.add_argument("--skip-spot-check", action="store_true")
args = ap.parse_args()
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]
base_model, head, specs, tokenizer = load_head_and_model(
Path(args.checkpoint), args.model_name, args.target_modules,
args.root_prefix, device, dtype, args.attn_implementation,
)
all_docs = load_doc_rows(Path(args.embeddings_path))
all_qnas = load_qna_rows(Path(args.qna_path))
qnas_by_doc_all = {}
qnas_held_out_by_doc = {}
for q in all_qnas:
qnas_by_doc_all.setdefault(q.doc_id, []).append({"prefix": q.prefix, "target": q.target})
if q.qna_split == "held_out":
qnas_held_out_by_doc.setdefault(q.doc_id, []).append({"prefix": q.prefix, "target": q.target})
train_docs = [d for d in all_docs if d.split == "train"]
results: Dict[str, Any] = {}
for suite in args.suites:
if suite in ("cr_val", "cr_test"):
rows, q_by_doc = [d for d in all_docs if d.split == suite], qnas_by_doc_all
elif suite == "ir_test":
rows, q_by_doc = train_docs, qnas_held_out_by_doc
else:
continue
if args.limit_docs and len(rows) > args.limit_docs:
import random as _r
rows = _r.Random(3407).sample(rows, args.limit_docs)
print(f"Evaluating {suite} ({len(rows)} docs) ...", flush=True)
m = eval_suite(base_model, head, specs, tokenizer, rows, q_by_doc, device,
max_qna_per_doc=args.max_qna_per_doc)
results[suite] = m
print(f" {suite}: EM={m['em']:.3f} EditSim={m['editsim']:.3f} n={m['n']}", flush=True)
print(json.dumps(results, indent=2))
if not args.skip_spot_check:
print("\n=== Spot check: base model vs. LoRA-adapted, Code2LoRA paper facts ===", flush=True)
paper_doc = next((d for d in all_docs if d.doc_id == "code2lora_paper"), None)
if paper_doc is None:
print(" [skip] code2lora_paper doc not found in embeddings", flush=True)
else:
ctx = torch.from_numpy(paper_doc.doc_embedding).to(device).unsqueeze(0)
head_out = head(ctx)
for q, gold in SPOT_CHECK_QUESTIONS:
prefix = f"Q: {q}\nA:"
# base: zero out LoRA (A=B=None) by re-wrapping without injection
for sp in specs:
named = dict(base_model.named_modules())
named[sp.full_name].A = None
named[sp.full_name].B = None
base_pred = generate(base_model, tokenizer, prefix, device)
inject_lora_weights(base_model, specs, head_out, batch_index=0)
adapted_pred = generate(base_model, tokenizer, prefix, device)
print(f"Q: {q}")
print(f" gold: {gold}")
print(f" base: {base_pred!r}")
print(f" adapted: {adapted_pred!r}")
print()
if __name__ == "__main__":
main()
|