Support fitting lenses on any dataset and equal-parts mixes
The fit corpus was limited to three hardcoded choices (wikitext, Semantic-Harmless, mixed), and any other id was rejected. Now any HuggingFace dataset id works, and any number of them can be ticked to fit on an equal-parts mix, shuffled. n_prompts now counts training SEQUENCES (what the fit iterates over) instead of source rows: each dataset is packed up to its quota, so the number entered is exactly what runs, regardless of the dataset. The fixed dropdown becomes a checkable dataset library persisted in localStorage. Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Opus 4.8
parent
daeb127651
commit
d9773394e3
@@ -211,9 +211,11 @@ Gemma — the error shows up in the UI if so.)
|
||||
|
||||
### 6. Fit your own lens
|
||||
|
||||
In **Fit** (model unloaded, VRAM free), fit a lens on streamed WikiText across one
|
||||
or more GPUs, with per-prompt checkpoints (stop/resume without loss) and weighted
|
||||
merging. Metadata is written to `lenses/<name>/meta.json`.
|
||||
In **Fit** (model unloaded, VRAM free), fit a lens across one or more GPUs on the
|
||||
corpus of your choice: tick any HuggingFace dataset by id (WikiText by default),
|
||||
or tick several to fit on an equal-parts mix. Per-prompt checkpoints give
|
||||
stop/resume without loss, and multi-GPU slices are merged by weighted average.
|
||||
Metadata is written to `lenses/<name>/meta.json`.
|
||||
|
||||
### CLI (no UI)
|
||||
|
||||
|
||||
+2
-2
@@ -175,7 +175,7 @@ class FitRequest(BaseModel):
|
||||
dtype: str = "bf16"
|
||||
quant: str | None = None
|
||||
n_prompts: int = 100
|
||||
dataset: str = fitting.DATASET_WIKITEXT # or DATASET_HARMLESS, or "mixed"
|
||||
datasets: list[str] = [fitting.DATASET_WIKITEXT] # any HF ids; several = equal-parts mix
|
||||
devices: list[str] = ["cuda:0"]
|
||||
name: str | None = None
|
||||
dim_batch: int | None = None
|
||||
@@ -811,7 +811,7 @@ def api_fit(req: FitRequest):
|
||||
n_prompts=req.n_prompts,
|
||||
dtype=req.dtype,
|
||||
quant=req.quant,
|
||||
dataset=req.dataset,
|
||||
datasets=req.datasets,
|
||||
devices=req.devices,
|
||||
name=req.name,
|
||||
dim_batch=req.dim_batch,
|
||||
|
||||
+124
-62
@@ -1,5 +1,6 @@
|
||||
import hashlib
|
||||
import json
|
||||
import re
|
||||
import subprocess
|
||||
import sys
|
||||
import threading
|
||||
@@ -14,65 +15,125 @@ from core.gpus import gpu_stats
|
||||
FITS_DIR = config.DATA_DIR / "fits"
|
||||
WORKER = config.ROOT / "scripts" / "fit_worker.py"
|
||||
|
||||
# Fit corpora. "mixed" = both, equal parts (rounded to the nearest prompt).
|
||||
# Fit corpus. Any HuggingFace dataset id works as-is: wikitext is the default
|
||||
# and keeps a dedicated streamed path, every other id goes through the generic
|
||||
# loader below. Several ids = an equal-parts mix.
|
||||
DATASET_WIKITEXT = "Salesforce/wikitext-103-raw-v1"
|
||||
DATASET_HARMLESS = "heretic-org/Semantic-Harmless"
|
||||
FIT_DATASETS = (DATASET_WIKITEXT, DATASET_HARMLESS, "mixed")
|
||||
|
||||
# Seed shared by the row sampling and the mix shuffle: a continued fit that asks
|
||||
# for skip+n rows deterministically extends the sequence it drew the first n
|
||||
# from (sample(skip+n) then drop the head).
|
||||
_SAMPLE_SEED = 1729
|
||||
|
||||
|
||||
def _load_corpus(dataset, n, skip=0):
|
||||
"""``n`` prompts from ``dataset``, skipping the first ``skip`` picks
|
||||
(continue-from: the new prompts must not overlap the base lens's).
|
||||
def _slug(dataset):
|
||||
"""Short, filename-safe tag from a dataset id's last path segment, e.g.
|
||||
``heretic-org/Semantic-Harmless`` -> ``semantic-harmless``."""
|
||||
tail = dataset.rstrip("/").split("/")[-1].lower()
|
||||
return re.sub(r"[^a-z0-9]+", "-", tail).strip("-")[:24] or "dataset"
|
||||
|
||||
wikitext keeps the historical behavior (first records ≥600 chars, streamed).
|
||||
Semantic-Harmless is a small instruct set (~416 one-line prompts): we draw a
|
||||
seeded random sample — sample(skip+n) then drop the head, so a continued fit
|
||||
extends the same sequence — and PACK the picks into ~350-char sequences
|
||||
(median prompt ≈ 10 tokens, and jlens skips the first 16 positions of every
|
||||
sequence as attention sinks: unpacked, almost every pick would be dropped as
|
||||
"too short"). ``n``/``skip`` count SOURCE prompts, not packs. "mixed" takes
|
||||
equal parts of both (n odd: the extra prompt goes to wikitext) and shuffles
|
||||
the union so multi-GPU slices stay mixed."""
|
||||
if dataset == "mixed":
|
||||
import random
|
||||
|
||||
n_wiki = (n + 1) // 2
|
||||
s_wiki = (skip + 1) // 2
|
||||
prompts = _load_corpus(DATASET_WIKITEXT, n_wiki, s_wiki)
|
||||
prompts += _load_corpus(DATASET_HARMLESS, n - n_wiki, skip // 2)
|
||||
random.Random(1729).shuffle(prompts)
|
||||
return prompts
|
||||
if dataset == DATASET_HARMLESS:
|
||||
import random
|
||||
def _text_column(features):
|
||||
"""Column to fit on: prefer ``text``, else the first string-valued column."""
|
||||
from datasets import Value
|
||||
|
||||
from datasets import load_dataset
|
||||
if "text" in features:
|
||||
return "text"
|
||||
for name, feat in features.items():
|
||||
if isinstance(feat, Value) and feat.dtype == "string":
|
||||
return name
|
||||
raise ValueError(
|
||||
"dataset exposes no text column to fit on "
|
||||
f"(columns: {', '.join(features) or 'none'})"
|
||||
)
|
||||
|
||||
texts = [r["text"] for r in load_dataset(DATASET_HARMLESS, split="train")]
|
||||
if skip + n > len(texts):
|
||||
raise ValueError(
|
||||
f"{DATASET_HARMLESS} has {len(texts)} prompts, "
|
||||
f"{skip + n} requested (continue included) — lower n_prompts"
|
||||
)
|
||||
picks = random.Random(1729).sample(texts, skip + n)[skip:]
|
||||
packs, cur = [], ""
|
||||
for text in picks:
|
||||
cur = f"{cur}\n\n{text}" if cur else text
|
||||
if len(cur) >= 350:
|
||||
packs.append(cur)
|
||||
cur = ""
|
||||
if cur:
|
||||
# a lone sub-16-token tail would be skipped by jlens anyway: fold it
|
||||
# into the previous pack instead of losing it
|
||||
if packs and len(cur) < 120:
|
||||
packs[-1] += "\n\n" + cur
|
||||
else:
|
||||
packs.append(cur)
|
||||
return packs
|
||||
from jlens.examples import load_wikitext_prompts
|
||||
|
||||
# load skip + n then keep the tail: the new prompts don't overlap
|
||||
# those of the base lens
|
||||
return load_wikitext_prompts(skip + n)[skip:]
|
||||
def _pack(texts, count, target=350):
|
||||
"""Pack ``texts`` into ~``target``-char sequences, stopping as soon as
|
||||
``count`` sequences are ready. jlens skips the first 16 positions of every
|
||||
sequence as attention sinks, so short unpacked prompts would almost all be
|
||||
dropped as too short. Returns fewer than ``count`` only if ``texts`` runs
|
||||
out (the caller decides whether that is an error)."""
|
||||
packs, cur = [], ""
|
||||
for text in texts:
|
||||
text = (text or "").strip()
|
||||
if not text:
|
||||
continue
|
||||
cur = f"{cur}\n\n{text}" if cur else text
|
||||
if len(cur) >= target:
|
||||
packs.append(cur)
|
||||
cur = ""
|
||||
if len(packs) >= count:
|
||||
return packs
|
||||
if cur and len(packs) < count:
|
||||
# trailing remainder: keep it so a just-large-enough dataset still fills
|
||||
# its quota (jlens tolerates a slightly-short final sequence)
|
||||
packs.append(cur)
|
||||
return packs
|
||||
|
||||
|
||||
def _load_split(dataset):
|
||||
"""``dataset``'s ``train`` split, or its first split if it has no ``train``."""
|
||||
from datasets import load_dataset
|
||||
|
||||
try:
|
||||
return load_dataset(dataset, split="train")
|
||||
except ValueError:
|
||||
dd = load_dataset(dataset)
|
||||
return dd[next(iter(dd))]
|
||||
|
||||
|
||||
def _load_one(dataset, n, skip):
|
||||
"""``n`` training SEQUENCES from a single ``dataset`` id, skipping the first
|
||||
``skip`` (continue-from: the new sequences must not overlap the base lens's).
|
||||
|
||||
wikitext keeps its historical path (first records >=600 chars, streamed —
|
||||
one record already is one sequence). Any other HF dataset is loaded whole,
|
||||
its text column shuffled with a fixed seed, then PACKED into ~350-char
|
||||
sequences until skip+n are ready (median instruct prompt ~10 tokens, so
|
||||
several rows per sequence). ``n``/``skip`` count OUTPUT sequences, so the
|
||||
number the user asks for is exactly what the fit iterates over — not source
|
||||
rows, whose count varies per dataset."""
|
||||
if n <= 0:
|
||||
return []
|
||||
if dataset == DATASET_WIKITEXT:
|
||||
from jlens.examples import load_wikitext_prompts
|
||||
|
||||
# load skip + n then keep the tail: the new sequences don't overlap the
|
||||
# base lens's
|
||||
return load_wikitext_prompts(skip + n)[skip:]
|
||||
import random
|
||||
|
||||
ds = _load_split(dataset)
|
||||
col = _text_column(ds.features)
|
||||
texts = [r[col] for r in ds]
|
||||
random.Random(_SAMPLE_SEED).shuffle(texts)
|
||||
packs = _pack(texts, skip + n)
|
||||
if len(packs) < skip + n:
|
||||
raise ValueError(
|
||||
f"{dataset}: {len(texts)} rows pack into only {len(packs)} sequences, "
|
||||
f"{skip + n} requested — lower n_prompts"
|
||||
)
|
||||
return packs[skip:skip + n]
|
||||
|
||||
|
||||
def _load_corpus(datasets, n, skip=0):
|
||||
"""``n`` training SEQUENCES drawn from ``datasets`` (a list of HF dataset
|
||||
ids). A single id loads that dataset; several are mixed in EQUAL parts — n
|
||||
and skip are each split across them (the first datasets take the rounding
|
||||
remainder) and the union is shuffled so multi-GPU slices stay mixed. Because
|
||||
the count is in sequences, ``n`` is exactly what the fit iterates over."""
|
||||
datasets = list(datasets)
|
||||
if len(datasets) == 1:
|
||||
return _load_one(datasets[0], n, skip)
|
||||
import random
|
||||
|
||||
k = len(datasets)
|
||||
prompts = []
|
||||
for i, ds in enumerate(datasets):
|
||||
prompts += _load_one(ds, n // k + int(i < n % k), skip // k + int(i < skip % k))
|
||||
random.Random(_SAMPLE_SEED).shuffle(prompts)
|
||||
return prompts
|
||||
|
||||
def _default_dim_batch(device):
|
||||
"""Default dim_batch scaled to the device's VRAM.
|
||||
@@ -102,14 +163,15 @@ class FitManager:
|
||||
def start(self, *, model_id, source, n_prompts=100, dtype="bf16", quant=None,
|
||||
devices=("cuda:0",), name=None, dim_batch=None,
|
||||
max_seq_len=128, source_layers=None, model_revision=None,
|
||||
continue_from=None, dataset=DATASET_WIKITEXT):
|
||||
continue_from=None, datasets=(DATASET_WIKITEXT,)):
|
||||
with self._lock:
|
||||
if self.state.get("state") == "running":
|
||||
raise ValueError("a fitting is already in progress")
|
||||
if not devices:
|
||||
raise ValueError("at least one device required")
|
||||
if dataset not in FIT_DATASETS:
|
||||
raise ValueError(f"unknown dataset: {dataset} (choices: {', '.join(FIT_DATASETS)})")
|
||||
datasets = [d.strip() for d in datasets if d and d.strip()]
|
||||
if not datasets:
|
||||
raise ValueError("at least one dataset required")
|
||||
skip_prompts = 0
|
||||
base_lens = None
|
||||
if continue_from:
|
||||
@@ -120,10 +182,10 @@ class FitManager:
|
||||
source_layers = list(base_lens.source_layers)
|
||||
if name is None:
|
||||
base = model_id.split("/")[-1]
|
||||
if dataset == "mixed":
|
||||
if len(datasets) > 1:
|
||||
base += "_mixed"
|
||||
elif dataset == DATASET_HARMLESS:
|
||||
base += "_harmless"
|
||||
elif datasets[0] != DATASET_WIKITEXT:
|
||||
base += "_" + _slug(datasets[0])
|
||||
total = n_prompts + skip_prompts
|
||||
name = f"{base}_n{total}" if continue_from else f"{base}_n{n_prompts}"
|
||||
params = {
|
||||
@@ -133,7 +195,7 @@ class FitManager:
|
||||
"dtype": dtype,
|
||||
"quant": quant,
|
||||
"n_prompts": n_prompts,
|
||||
"dataset": dataset,
|
||||
"datasets": datasets,
|
||||
"devices": list(devices),
|
||||
"dim_batch": dim_batch,
|
||||
"max_seq_len": max_seq_len,
|
||||
@@ -176,7 +238,7 @@ class FitManager:
|
||||
prompts = json.loads(corpus_path.read_text(encoding="utf-8"))
|
||||
else:
|
||||
prompts = _load_corpus(
|
||||
params.get("dataset", DATASET_WIKITEXT),
|
||||
params.get("datasets", [DATASET_WIKITEXT]),
|
||||
params["n_prompts"],
|
||||
params.get("skip_prompts", 0),
|
||||
)
|
||||
@@ -299,9 +361,9 @@ class FitManager:
|
||||
"quant": params["quant"],
|
||||
"n_prompts": merged.n_prompts,
|
||||
"corpus": (
|
||||
f"mixed: {DATASET_WIKITEXT} + {DATASET_HARMLESS} (equal parts)"
|
||||
if params.get("dataset") == "mixed"
|
||||
else params.get("dataset", DATASET_WIKITEXT)
|
||||
"mixed: " + " + ".join(params["datasets"]) + " (equal parts)"
|
||||
if len(params["datasets"]) > 1
|
||||
else params["datasets"][0]
|
||||
),
|
||||
"max_seq_len": params["max_seq_len"],
|
||||
"devices": params["devices"],
|
||||
|
||||
+49
-10
@@ -195,7 +195,34 @@ export default function App() {
|
||||
|
||||
const [fitModel, setFitModel] = useState('')
|
||||
const [fitN, setFitN] = useState(100)
|
||||
const [fitDataset, setFitDataset] = useState('Salesforce/wikitext-103-raw-v1')
|
||||
// fit corpus library: tick one or several; several = mixed in equal parts.
|
||||
// persisted so the user's added datasets survive a refresh.
|
||||
const [fitDatasets, setFitDatasets] = useState(() => {
|
||||
try {
|
||||
const saved = JSON.parse(localStorage.getItem('jlens_fit_datasets') || 'null')
|
||||
const restored = Array.isArray(saved)
|
||||
? saved.filter((d) => d && d.id).map((d) => ({ id: String(d.id), on: !!d.on }))
|
||||
: []
|
||||
if (restored.length) return restored
|
||||
} catch { /* ignore corrupt storage */ }
|
||||
return [
|
||||
{ id: 'Salesforce/wikitext-103-raw-v1', on: true },
|
||||
{ id: 'heretic-org/Semantic-Harmless', on: false },
|
||||
]
|
||||
})
|
||||
useEffect(() => {
|
||||
localStorage.setItem('jlens_fit_datasets', JSON.stringify(fitDatasets))
|
||||
}, [fitDatasets])
|
||||
const [fitDatasetInput, setFitDatasetInput] = useState('')
|
||||
const addFitDataset = () => {
|
||||
const id = fitDatasetInput.trim()
|
||||
if (!id) return
|
||||
setFitDatasets((prev) =>
|
||||
prev.some((d) => d.id === id)
|
||||
? prev.map((d) => (d.id === id ? { ...d, on: true } : d))
|
||||
: [...prev, { id, on: true }])
|
||||
setFitDatasetInput('')
|
||||
}
|
||||
const [fitQuant, setFitQuant] = useState('')
|
||||
const [fitDevices, setFitDevices] = useState([])
|
||||
const [fitDimBatch, setFitDimBatch] = useState('')
|
||||
@@ -1287,15 +1314,27 @@ export default function App() {
|
||||
<div className="row"><label>name</label>
|
||||
<input type="text" placeholder="(auto: model_nN)" value={fitName} onChange={(e) => setFitName(e.target.value)} />
|
||||
</div>
|
||||
<div className="row"><label title="number of corpus prompts (existing lenses were made with n=100 unless marked _nNNN)">prompts</label>
|
||||
<div className="row"><label title="number of training sequences the fit iterates over — what you set is exactly what runs (existing lenses were made with n=100 unless marked _nNNN)">sequences</label>
|
||||
<input type="number" min="4" step="1" value={fitN} onChange={(e) => setFitN(e.target.value)} />
|
||||
</div>
|
||||
<div className="row"><label title="fit corpus. mixed = both datasets in equal parts (rounded to the nearest prompt), shuffled">dataset</label>
|
||||
<select value={fitDataset} onChange={(e) => setFitDataset(e.target.value)}>
|
||||
<option value="Salesforce/wikitext-103-raw-v1">Salesforce/wikitext-103-raw-v1</option>
|
||||
<option value="heretic-org/Semantic-Harmless">heretic-org/Semantic-Harmless</option>
|
||||
<option value="mixed">mixed (50/50)</option>
|
||||
</select>
|
||||
<div className="row"><label title="fit corpus. Tick one or several HuggingFace datasets; several ticked = mixed in equal parts (rounded to the nearest sequence), shuffled">datasets</label>
|
||||
<span style={{ display: 'flex', flexDirection: 'column', gap: 6, flex: 1 }}>
|
||||
{fitDatasets.map((d, i) => (
|
||||
<label key={d.id} style={{ width: 'auto', display: 'flex', alignItems: 'center', gap: 6 }} title={d.id}>
|
||||
<input type="checkbox" checked={d.on}
|
||||
onChange={(e) => setFitDatasets(fitDatasets.map((x, j) => j === i ? { ...x, on: e.target.checked } : x))} />
|
||||
<span style={{ flex: 1, overflow: 'hidden', textOverflow: 'ellipsis', whiteSpace: 'nowrap' }}>{d.id}</span>
|
||||
<button className="linkbtn" title="remove from the list (does not delete anything on disk)"
|
||||
onClick={() => setFitDatasets(fitDatasets.filter((_, j) => j !== i))}>✕</button>
|
||||
</label>
|
||||
))}
|
||||
<span style={{ display: 'flex', gap: 6 }}>
|
||||
<input type="text" placeholder="org/dataset (HuggingFace id)" value={fitDatasetInput}
|
||||
onChange={(e) => setFitDatasetInput(e.target.value)}
|
||||
onKeyDown={(e) => { if (e.key === 'Enter') { e.preventDefault(); addFitDataset() } }} />
|
||||
<button onClick={addFitDataset} title="add this dataset to the list">+</button>
|
||||
</span>
|
||||
</span>
|
||||
</div>
|
||||
<div className="row"><label>quant</label>
|
||||
<select value={fitQuant} onChange={(e) => setFitQuant(e.target.value)} disabled={!!fitContinue}>
|
||||
@@ -1347,7 +1386,7 @@ export default function App() {
|
||||
)}
|
||||
<button
|
||||
className="primary"
|
||||
disabled={(!fitModel && !fitContinue) || !fitDevices.length || !!loadedId || !!busy}
|
||||
disabled={(!fitModel && !fitContinue) || !fitDevices.length || !fitDatasets.some((d) => d.on) || !!loadedId || !!busy}
|
||||
onClick={async () => {
|
||||
try {
|
||||
const layers = []
|
||||
@@ -1361,7 +1400,7 @@ export default function App() {
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({
|
||||
model_id: fitModel, n_prompts: +fitN, quant: fitQuant || null,
|
||||
dataset: fitDataset,
|
||||
datasets: fitDatasets.filter((d) => d.on).map((d) => d.id),
|
||||
name: fitName.trim() || null, devices: fitDevices,
|
||||
dim_batch: fitDimBatch ? +fitDimBatch : null,
|
||||
max_seq_len: +fitMaxSeq,
|
||||
|
||||
Reference in New Issue
Block a user