Spaces:
Sleeping
Sleeping
Download app.py from jaydeepb/predicting_memorization: direct link, hf CLI and curl.
- Browser
- Download file 11 kB
-
https://huggingface.co/spaces/jaydeepb/predicting_memorization/resolve/main/app.py
- Command line
-
hf download hf://spaces/jaydeepb/predicting_memorization/app.py
-
curl -L -o app.py https://huggingface.co/spaces/jaydeepb/predicting_memorization/resolve/main/app.py
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) | |
| 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() | |