jaydeepb's picture
Fix short-input error: check token length in main process before @spaces.GPU call
1ceaa1d verified
Raw History Blame Contribute Delete
11 kB
"""Predicting Memorization Before Fine-Tuning — interactive demo.
Two modes:
(1) Explore the held-out 1M — pick a real training sequence, see the fine-tuned
model's actual 50-token generation (green where it matches the ground truth),
the memorized/not verdict, and a per-example explanation of the classifier's
pre-fine-tuning risk score.
(2) Score your own text — paste a passage, get the classifier's predicted
memorization risk from base-model features, with the same per-example explanation.
A collapsible panel shows the headline result: base-model gradient features alone
predict memorization about as well as the full feature set.
"""
import os, glob, json, joblib, numpy as np, pandas as pd, torch, gradio as gr
import spaces # Hugging Face ZeroGPU; @spaces.GPU is a no-op off ZeroGPU
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import matplotlib.patches as mpatches
HERE = os.path.dirname(os.path.abspath(__file__))
BASE_MODEL = os.environ.get("BASE_MODEL", "EleutherAI/pythia-1.4b")
LOOKUP = pd.read_parquet(os.path.join(HERE, "data", "lookup.parquet"))
DATASETS = sorted(LOOKUP["dataset"].unique())
CLFS = {}
for f in glob.glob(os.path.join(HERE, "models", "clf_*.joblib")):
b = joblib.load(f)
CLFS[b["dataset"]] = b
# Feature distributions per dataset, split by memorization outcome (for the "where does
# this example fall" plots). Model-free: reads straight from the lookup table.
FEAT_ORDER = ["base_ppl", "base_loss_variance", "one_step_rel_change",
"one_step_abs_change", "gradient_norm", "zlib_entropy"]
FEAT_NAMES = {"zlib_entropy": "zlib entropy", "base_ppl": "base perplexity",
"base_loss_variance": "base loss variance", "gradient_norm": "gradient norm",
"one_step_abs_change": "one-step loss change", "one_step_rel_change": "one-step rel. loss change"}
DIST = {ds: {"mem": LOOKUP[(LOOKUP.dataset == ds) & LOOKUP.memorized],
"non": LOOKUP[(LOOKUP.dataset == ds) & ~LOOKUP.memorized]} for ds in DATASETS}
GREEN, RED = "#2e7d32", "#c62828" # generation highlight
MEMCLR, NONCLR = "#f58518", "#4c78a8" # memorized (orange) vs not-memorized (blue)
# ZeroGPU: place the base model on cuda at module level (a CUDA-emulation mode makes
# this valid at startup; a real GPU is attached only inside @spaces.GPU functions).
_ON_ZEROGPU = bool(os.environ.get("SPACES_ZERO_GPU"))
DEVICE = "cuda" if (_ON_ZEROGPU or torch.cuda.is_available()) else "cpu"
try:
from byo_features import BYOExtractor, WINDOW_LEN
EXTRACTOR = BYOExtractor(BASE_MODEL, device=DEVICE)
_LOAD_ERR = None
except Exception as e:
EXTRACTOR, _LOAD_ERR = None, str(e)
@spaces.GPU(duration=120)
def _compute_features(text):
return EXTRACTOR.features(text)
# ------------------------------- graphics ----------------------------------
def dist_figure(feats, dataset):
"""Six panels: where this example's feature values fall relative to the memorized
and not-memorized populations of the chosen dataset. Model-free."""
d = DIST.get(dataset) or DIST[next(iter(DIST))]
mem, non = d["mem"], d["non"]
fig, axes = plt.subplots(2, 3, figsize=(9.6, 5.4))
for ax, f in zip(axes.ravel(), FEAT_ORDER):
mv = mem[f].to_numpy(float); nv = non[f].to_numpy(float)
both = np.concatenate([mv, nv])
lo, hi = np.nanpercentile(both, [1, 99])
if not np.isfinite(lo) or not np.isfinite(hi) or lo == hi:
lo, hi = float(np.nanmin(both)), float(np.nanmax(both)) + 1e-9
use_log = lo > 0 and hi / lo > 50
bins = (np.logspace(np.log10(lo), np.log10(hi), 40) if use_log
else np.linspace(lo, hi, 40))
centers = 0.5 * (bins[:-1] + bins[1:])
for arr, color, label in [(nv, NONCLR, "not memorized"), (mv, MEMCLR, "memorized")]:
c, _ = np.histogram(np.clip(arr, lo, hi), bins=bins)
c = c / c.max() if c.max() > 0 else c # normalize each class to its own peak
ax.fill_between(centers, c, step="mid", alpha=0.5, color=color, label=label)
if use_log:
ax.set_xscale("log")
v = float(feats[f]); vc = min(max(v, lo), hi)
ax.axvline(vc, color="k", lw=1.8)
ax.plot([vc], [ax.get_ylim()[1] * 0.9], marker="v", color="k", ms=9, clip_on=False)
off = "" if lo <= v <= hi else (" (off-scale ←)" if v < lo else " (off-scale →)")
ax.set_title(f"{FEAT_NAMES[f]}\nthis example = {v:.3g}{off}", fontsize=8.5)
ax.set_yticks([]); ax.tick_params(labelsize=7)
# top strip: "this example" note on the left, 2-column class legend on the right
handles = [mpatches.Patch(color=MEMCLR, alpha=0.5, label="memorized"),
mpatches.Patch(color=NONCLR, alpha=0.5, label="not memorized")]
fig.text(0.5, 0.965, "black ▼ = this example", ha="center", va="center", fontsize=9)
fig.legend(handles=handles, ncol=2, loc="upper center", bbox_to_anchor=(0.5, 0.94),
fontsize=9, frameon=False, columnspacing=1.4, handlelength=1.2)
fig.tight_layout(rect=[0, 0, 1, 0.88])
return fig
def _segments(gen_segments_json):
segs = json.loads(gen_segments_json)
return [(t, lab) for t, lab in segs]
# ----------------------------- Explore (lookup) -----------------------------
def draw(dataset, which):
df = LOOKUP[LOOKUP.dataset == dataset]
if which == "Memorized only":
df = df[df.memorized]
elif which == "Not memorized":
df = df[~df.memorized]
if len(df) == 0:
return "No examples for this filter.", [], "", None
r = df.sample(1).iloc[0]
gen = _segments(r.gen_segments)
if r.memorized:
verdict = "### 🔴 Memorized\nThe fine-tuned model reproduced the ground-truth suffix exactly (all tokens green)."
else:
verdict = ("### 🟢 Not memorized\nThe fine-tuned model's greedy generation diverges from the ground truth; "
"green tokens happen to match, red tokens do not.")
clf_name = dataset if dataset in CLFS else next(iter(CLFS))
b = CLFS[clf_name]
feats = {k: float(r[k]) for k in FEAT_ORDER}
X = np.nan_to_num(np.array([[feats[k] for k in b["feats"]]]), posinf=1e10, neginf=-1e10)
risk = float(b["clf"].predict_proba(b["scaler"].transform(X))[0, 1])
verdict += (f"\n\n### Classifier risk score: {risk:.3f}\n"
f"Computed by the memorization classifier from base-model features measured before "
f"fine-tuning. The plots below show where "
f"this example's feature values fall relative to memorized and not-memorized examples.")
return r.prefix, gen, verdict, dist_figure(feats, dataset)
# ----------------------------- Predict (BYO) --------------------------------
def predict(text, clf_name):
text = (text or "").strip()
if not text:
return "Please enter some input text first.", None
if EXTRACTOR is None:
return f"⚠️ Base model unavailable: {_LOAD_ERR}", None
# length check in the main process (CPU tokenization); a ValueError raised inside the
# @spaces.GPU worker would not survive the process boundary as a catchable ValueError.
n_tokens = len(EXTRACTOR.tok(text, truncation=True, max_length=WINDOW_LEN)["input_ids"])
if n_tokens < WINDOW_LEN:
return (f"⚠️ The input should be at least {WINDOW_LEN} tokens (or about 75 words); "
f"you entered {n_tokens} tokens."), None
try:
feats = _compute_features(text)
except ValueError as e:
return f"⚠️ {e}", None
b = CLFS[clf_name]
X = np.nan_to_num(np.array([[feats[k] for k in b["feats"]]]), posinf=1e10, neginf=-1e10)
risk = float(b["clf"].predict_proba(b["scaler"].transform(X))[0, 1])
flag = "high" if risk >= 0.5 else "low"
md = (f"### Predicted memorization risk: {risk:.3f} ({flag})\n"
f"Forecast from base-model features using the classifier for Pythia fine-tuned on {clf_name}. "
f"There is no ground truth for arbitrary text, so treat it as a relative risk score. "
f"The plots below show where your text's features fall relative to memorized and not-memorized "
f"examples in {clf_name}.")
return md, dist_figure(feats, clf_name)
INTRO = """# Predicting Memorization Before Fine-Tuning
A classifier trained only on pre-fine-tuning base-model features predicts which sequences a
fine-tuned model will memorize. The "Explore held-out examples" tab verifies this against ground truth on
held-out examples from a fresh fine-tuning run; the "Score your own text" tab applies the same classifier
to a passage you provide.
"""
with gr.Blocks(title="Predicting Memorization Before Fine-Tuning") as demo:
gr.Markdown(INTRO)
with gr.Tab("Explore held-out examples"):
gr.Markdown("Pick a real Pythia-1.4B fine-tuning sequence and compare the classifier's "
"pre-fine-tuning risk score with what the fine-tuned model actually generated.")
with gr.Row():
d_in = gr.Dropdown(DATASETS, value=DATASETS[0], label="Dataset")
f_in = gr.Radio(["All", "Memorized only", "Not memorized"], value="All", label="Filter")
go = gr.Button("Show a random example", variant="primary")
pref = gr.Textbox(label="Prefix (50 tokens given to the model)", lines=3)
gen = gr.HighlightedText(
label="Model's actual 50-token generation (green = matches ground truth, red = differs)",
color_map={"match": GREEN, "diff": RED}, combine_adjacent=True, show_legend=False)
verdict = gr.Markdown()
feat = gr.Plot(label="Where this example's features fall (memorized vs not-memorized)")
go.click(draw, [d_in, f_in], [pref, gen, verdict, feat])
with gr.Tab("Score your own text"):
gr.Markdown("Please enter an input text that is at least 100 tokens (or 75 words). Using its "
"base-model features from Pythia-1.4B, our learned classifiers (one per fine-tuning "
"dataset) predict how likely Pythia-1.4B would be to memorize it if you fine-tuned on "
"it. The plots below show where your text's features fall relative to memorized and "
"not-memorized examples.")
txt = gr.Textbox(label="Your text", lines=6)
clf_in = gr.Dropdown(sorted(CLFS.keys()), value=sorted(CLFS.keys())[0],
label="Choose a classifier")
pgo = gr.Button("Predict memorization risk", variant="primary")
pmd = gr.Markdown()
pfeat = gr.Plot(label="Where your text's features fall (memorized vs not-memorized)")
pgo.click(predict, [txt, clf_in], [pmd, pfeat])
gr.Markdown("_Datasets shown are public (FineWeb, PG-19, The Stack, OpenWebMath). "
"The classifier uses only base-model features and was trained on a separate fine-tuning run._")
if __name__ == "__main__":
demo.launch()