Spaces:
Running on Zero
Running on Zero
File size: 11,008 Bytes
1dbc4b5 aa72ab5 1dbc4b5 aa72ab5 1dbc4b5 aa72ab5 5693c0e 1dbc4b5 aa72ab5 5e1a704 aa72ab5 5e1a704 aa72ab5 1dbc4b5 1ceaa1d 1dbc4b5 aa72ab5 1dbc4b5 aa72ab5 5e1a704 5693c0e 5e1a704 5693c0e 38e8d0e 5693c0e 38e8d0e aa72ab5 1dbc4b5 aa72ab5 1dbc4b5 aa72ab5 1dbc4b5 aa72ab5 1dbc4b5 aa72ab5 5e1a704 261dfe0 929cc98 5e1a704 1dbc4b5 929bb2f 1dbc4b5 1ceaa1d 1dbc4b5 261dfe0 327bb4f aa72ab5 5e1a704 261dfe0 5e1a704 1dbc4b5 261dfe0 929cc98 1dbc4b5 929cc98 aa72ab5 1dbc4b5 aa72ab5 1dbc4b5 5e1a704 aa72ab5 1dbc4b5 929bb2f 1dbc4b5 1210a65 1dbc4b5 5e1a704 1dbc4b5 4975726 1dbc4b5 | 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 | """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()
|