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()