diff --git a/.gitignore b/.gitignore index 83d0800..6166370 100644 --- a/.gitignore +++ b/.gitignore @@ -18,3 +18,5 @@ isolate-*.log /python/data_wild/raw/ /python/data_wild/samples/ /python/data_wild/wild_preds.jsonl +/python/data_wild/negatives.jsonl +/python/data_wild/pretrain_corpus.jsonl diff --git a/python/README.md b/python/README.md index b2f16ca..c3016ae 100644 --- a/python/README.md +++ b/python/README.md @@ -156,6 +156,94 @@ production input would — `train.py` prints a warning listing the 20 most frequent such characters (and how many extra examples they came from) so you can see whether an unexpected script slipped in. +### Real negatives (`--extra-negatives`) + +```bash +# regenerate the pool (gitignored: it is derived from CC BY-SA corpus text) +python -m sankhya.wild negatives --out data_wild/negatives.jsonl --n 20000 + +python -m sankhya.train \ + --train data/train.jsonl --val data/val.jsonl \ + --extra data_llm/hi_latn.jsonl data_llm/hi_deva.jsonl \ + data_llm/mr_deva.jsonl data_llm/gu_gujr.jsonl --extra-ratio 0.2 \ + --extra-negatives data_wild/negatives.jsonl --extra-negative-ratio 0.15 \ + --lang hi_latn,hi_deva,mr_deva,gu_gujr --out models/ +``` + +`data_wild/REPORT.md` measured a **20% false-positive rate on real sentences +with no amount in them** (8.8% of them produce a *verified* span) against 0% +on the hand-written negatives — the model has never seen real non-amount text, +so names (`अण्णा हजारे`), ethnonyms (`अरब`, `અરબ`), measurements and years look +like amounts to it. `--extra-negatives` is the data-side answer: real +sentences, labelled all-`O`. + +`--extra-negatives` takes jsonl file(s) of `{"text": ..., "lang": ...}` lines +(no `bio`/`cls`/`spans` — they are *derived*: BIO all 0, class all `O`, for the +normalized, `MAX_LEN`-truncated text). They are mixed per epoch exactly like +`--extra`: `n_neg_per_epoch = round(base * r / (1 - r))` where `base` is +`--train` plus that epoch's `--extra` draw, so `--extra-negative-ratio 0.15` +makes real negatives ~15% of the combined epoch and the two ratios compose. +The pool is redrawn each epoch (without replacement when it is big enough, +with replacement otherwise). `--extra-negative-ratio 0` disables it. + +**Building the pool** — `python -m sankhya.wild negatives`: + +- draws a uniform reservoir sample of real lines per language from the same + corpora and filters `sankhya.wild sample` uses (3–40 words, ≤ 128 chars, + deduplicated); +- keeps only lines the **lexicon scorer** gives no amount signal at all + (score 0: no unit word, prefix, cardinal, currency marker or digit); +- *and* only lines the **currently shipped int8 weights leave alone** — a line + the model spans might be a real amount the lexicon missed, and training on + it as all-`O` would teach the wrong thing. About 1–4% of no-signal lines are + dropped this way (worst in `hi_latn`); +- excludes every text in `tests/gold_wild_*.jsonl` and asserts the result is + disjoint from it, so the evaluation set cannot leak into training; +- balances the kept lines across the four languages (`--n 20000` → 5,000 each). + +The output is **derived from CC BY-SA corpus text**, so like `data_wild/raw/` +and `data_wild/samples/` it is **gitignored**; only the tool is committed. +Regenerate it with the command at the top of this section (seeded, so it is +reproducible). + +## Masked-character pretraining (`pretrain.py`) + +```bash +python -m sankhya.pretrain --out models/pretrain --n 500000 --epochs 6 \ + --arch v2 --channels 48 --lang hi_latn,hi_deva,mr_deva,gu_gujr \ + --cache data_wild/pretrain_corpus.jsonl + +python -m sankhya.train ... --init-from models/pretrain/pretrain.pt +``` + +The same "the model has never seen real text" problem as above, attacked +unsupervised instead: take the **same `SankhyaCNN` trunk** (char embedding + +conv stack), put a temporary `channels -> vocab_size` head on it, mask 15% of +the characters of real sentences and train it to reconstruct **only the masked +positions** (`CrossEntropyLoss(ignore_index=-100)`, so padding and unmasked +characters contribute nothing to the loss). The head is then thrown away. + +- Corpus: `--n` real sentences pulled through `sankhya.wild`'s loaders and + filters (3–40 words, ≤ 128 chars, deduplicated), balanced across the four + languages by water-filling (a language with fewer lines than its equal share + hands the remainder to the ones that have more — `hi_latn` only has ~12k + usable lines, `hi_deva` has 716k). Every `tests/gold_wild_*.jsonl` text is + excluded, so pretraining cannot memorise the evaluation set. `--cache PATH` + writes/reads the collected sentences so a re-run skips the corpus scan. +- Charset: the **shipped** one, built from `--lang` exactly as `train.py` + builds it. Real characters outside it map to ``, as production input + does. The mask token is `` too, which keeps the embedding table the + exact shape the tagger expects. +- Budget: ~40k parameters; on 4 CPU cores one epoch over 500k sentences is a + few minutes, so pick `--epochs` after timing the first one. + +`train.py --init-from CKPT` then loads **only the trunk** tensors +(`embed.*`, `conv*.*`) out of that checkpoint and leaves the BIO and class +heads at their fresh torch init; a shape or key mismatch is fatal rather than +silently partial. It composes with everything else (`--extra`, +`--extra-negatives`, `--arch`, `--channels`), which is how experiment B in +`data_wild/REPORT.md` was run. + ## Export ```bash @@ -796,7 +884,11 @@ strict-tier predicate and the freshness of `src/data/lexicon.json`; false-positive ceilings. `tests/test_wild_rules.py` covers the R11–R18 precision rules case by case (mirrored by `test/wild-rules.test.ts` on the JS side) and `tests/test_eval_long_lines.py` the >128-character windowing. -`tests/gold.jsonl` is the +`test_pretrain_negatives.py` covers the data plumbing of the two model-side +experiments: `--extra-negatives` lines become all-`O` examples and are mixed +at the requested ratio, `sankhya.pretrain`'s masking hits 15% of real +characters with the loss ignoring every unmasked position, and `--init-from` +loads the trunk while leaving the heads fresh. `tests/gold.jsonl` is the hand-written gold set used by `eval_gold.py`, not a generator round-trip test. diff --git a/python/data_wild/REPORT.md b/python/data_wild/REPORT.md index 9726f18..220e913 100644 --- a/python/data_wild/REPORT.md +++ b/python/data_wild/REPORT.md @@ -226,7 +226,96 @@ turns it straight into a gold line. that also means everything in this report is measured on the short end of real text. -## 5. What real text has that the generator never produces +## 5. Model-side experiments + +**Decision (0.8.0):** the shipped weights stay. Re-evaluated under the R11–R18 +rules (which landed in parallel), experiment B and the shipped weights are +within seed noise of each other on every metric — real-text FP 8 vs 10 of +240, verified FP 6 vs 4, synthetic value accuracy identical, strict value +accuracy 1.000 for both — and B also changes the charset. The rules +captured most of what B learned from real negatives. Real-negative mixing +(`--extra-negatives`) and the pretrained trunk (`--init-from`) are now the +standard recipe for the next weights release, when a fifth language or a +new gold finding gives a reason to retrain. + + +Section 3's headline is a *precision* problem: 20% of real no-amount +sentences produce a span (8.8% a verified one) against 0% on the hand-written +negatives, because the model has never seen real non-amount text. Two +model-side levers, each one CPU run of the standard recipe (`v2`, 48 ch, 20 +epochs, seed 2, `--lang hi_latn,hi_deva,mr_deva,gu_gujr --mix +0.32,0.26,0.21,0.21 --cross 0.10`, 200k train / 6k val, the four verified LLM +corpora at `--extra-ratio 0.2`): + +- **A — real negatives in training.** 20,000 real sentences with no amount + signal, labelled all-`O`, mixed in at 15% of each epoch + (`--extra-negatives ... --extra-negative-ratio 0.15`). Pool built with + `python -m sankhya.wild negatives --out data_wild/negatives.jsonl --n 20000`: + lexicon score 0, *and* not spanned by the shipped weights, balanced 5,000 per + language, asserted disjoint from `tests/gold_wild_*.jsonl`. The file is + derived from CC BY-SA corpus text, so it is gitignored — only the tool is + committed (see python/README.md, "Real negatives"). +- **B — masked-character pretraining, then A.** `python -m sankhya.pretrain` + trains the same `SankhyaCNN` trunk (embedding + conv stack) with a throwaway + per-char vocab head to reconstruct 15% randomly masked characters of 500,000 + real sentences (shipped 170-char charset, `` as the mask token, loss + only on masked positions, seed 23, wild-gold text excluded); `train.py + --init-from pretrain.pt` then runs recipe A on top with fresh heads. + +### Results (exported **int8** weights, same gates as section 3) + +| model | syn. value acc | syn. strict cov | wild value acc | wild strict cov | wild FP | wild strict FP | +| --- | ---: | ---: | ---: | ---: | ---: | ---: | +| shipped 0.7.0 | 0.962 | 0.909 | 0.899 | 0.869 | **0.200** | **0.088** | +| A (real negatives) | 0.952 | 0.891 | 0.881 | 0.869 | **0.079** | **0.063** | +| B (pretrain + A) | **0.962** | 0.906 | **0.905** | **0.893** | **0.033** | **0.025** | + +Per-language wild false-positive rate (FP / strict FP, over the 240 real +negatives in the wild gold): + +| model | hi_latn (79) | hi_deva (57) | mr_deva (52) | gu_gujr (52) | +| --- | ---: | ---: | ---: | ---: | +| shipped 0.7.0 | 0.367 / 0.177 | 0.140 / 0.070 | 0.135 / 0.019 | 0.077 / 0.038 | +| A | 0.101 / 0.076 | 0.088 / 0.053 | 0.058 / 0.058 | 0.058 / 0.058 | +| B | 0.051 / 0.038 | 0.070 / 0.053 | **0.000 / 0.000** | **0.000 / 0.000** | + +Spurious spans on the 400 wild gold lines: 63 (shipped) -> 22 (A) -> 10 (B). +Strict value accuracy stays **1.000** on every row, synthetic and wild: the +strict tier's promise is intact, and B finally makes strict mode mean what the +synthetic numbers implied it meant. + +Synthetic negatives stay at 0 false positives for all three models, and +synthetic strict value accuracy stays 1.000. + +**Did the synthetic numbers move?** A costs about a point of synthetic value +accuracy (0.962 -> 0.952) and 1.8 points of synthetic strict coverage +(0.909 -> 0.891) — real negatives are 15% of every epoch, and the model spends +some capacity on them. B gets that back: 0.962 / 0.906, i.e. synthetic value +accuracy indistinguishable from shipped, while cutting the wild FP rate by 6x. +Note the caveat the sweep section already makes: seed-to-seed spread on this +recipe is +/-1-2 points on gold, and these are single runs on a freshly +generated dataset, so the small synthetic differences (and A's small wild +value-accuracy dip) are within noise. The FP collapse — 48 -> 19 -> 8 false +positives out of 240 — is far outside it. + +### Wall time (this container, 4 CPU cores) + +| step | time | +| --- | ---: | +| `wild negatives --n 20000` (scan + shipped-model filter) | 1m 46s | +| A: train (20 epochs, 294k examples/epoch) | 22m 16s | +| B: pretrain corpus collection (500k sentences) | 1m 33s | +| B: pretrain (12 epochs x 500k, masked-char acc 0.412 -> 0.549) | 22m 10s | +| B: train (20 epochs, same as A, `--init-from`) | 20m 1s | + +The per-epoch cost of A over the shipped recipe is the extra 44k negatives +(294k vs 250k examples per epoch); pretraining adds a flat ~24 minutes once, +reusable across tagger runs (the trunk checkpoint is 40,475 parameters). + +**Not adopted.** These weights are experiments, not a release: nothing under +`models/default/` was touched. + +## 6. What real text has that the generator never produces - **Amounts are rare and unevenly distributed.** ~1% of real sentences mention one; the generator's world is mostly amounts. Precision, not recall, is what diff --git a/python/sankhya/pretrain.py b/python/sankhya/pretrain.py new file mode 100644 index 0000000..36b88de --- /dev/null +++ b/python/sankhya/pretrain.py @@ -0,0 +1,290 @@ +"""Masked-character pretraining for the SankhyaCNN trunk. + +The tagger only ever sees text the generator wrote (plus a small verified LLM +corpus). Real sentences -- names, ethnonyms, measurements, years, the ordinary +prose amounts are embedded in -- are what it has never had any signal about, +and `data_wild/REPORT.md` shows the cost: 20% of real no-amount sentences get +a span. + +This module gives the *trunk* (char embedding + conv stack, the same +`SankhyaCNN` the tagger uses) a cheap unsupervised look at that text first: +mask a fraction of the characters of a real sentence, put a temporary +per-character vocab head on the conv stack, and train it to reconstruct only +the masked positions. The head is thrown away; `sankhya.train --init-from` +loads just the trunk and re-initialises the BIO/class heads. + + python -m sankhya.pretrain --out /tmp/pretrain --n 500000 --epochs 8 + +The charset is the SHIPPED one (built from the `--lang` packs exactly as +`train.py` builds it), never rebuilt from the real text: a real character +outside it maps to ``, which is what production input does too. The mask +token is `` as well, so the embedding table keeps the exact shape the +tagger expects and `--init-from` is a straight tensor copy. + +Wild-gold text (`tests/gold_wild_*.jsonl`) is excluded, so pretraining cannot +memorise the evaluation set. +""" +from __future__ import annotations + +import argparse +import json +import math +import os +import random +import time + +import numpy as np +import torch +import torch.nn as nn + +from . import classes as C +from .charset import build_charset, build_charset_multi, normalize_text +from .langs import base as langbase +from .model import SankhyaCNN, count_params, ARCHS + +MAX_LEN = 128 +DEFAULT_MASK_RATE = 0.15 +IGNORE_INDEX = -100 + + +class PretrainModel(nn.Module): + """`SankhyaCNN` trunk + a temporary `channels -> vocab_size` head. + + The trunk is the real thing (same module, same state_dict keys), so its + weights drop straight into a tagger via `train.py --init-from`. The + tagger's own BIO/class heads exist here too but are never trained -- only + `self.head` gets gradient from the reconstruction loss. + """ + + def __init__(self, vocab_size, arch=None, channels=48, embed_dim=16): + super().__init__() + self.trunk = SankhyaCNN( + vocab_size=vocab_size, n_cls=len(C.CLASSES), arch=arch, + channels=channels, embed_dim=embed_dim, + ) + self.head = nn.Linear(channels, vocab_size) + + def features(self, chars): + t = self.trunk + x = t.embed(chars).transpose(1, 2) + for layer, conv in zip(t.arch, t.convs): + y = t.act(conv(x)) + if layer["residual"]: + y = y + x + x = y + return x.transpose(1, 2) + + def forward(self, chars): + return self.head(self.features(chars)) + + def trunk_state_dict(self): + """Only the tensors `train.py --init-from` consumes.""" + return {k: v for k, v in self.trunk.state_dict().items() + if k.startswith("embed.") or k.startswith("conv")} + + +def apply_mask(chars, mask, mask_id, rate=DEFAULT_MASK_RATE, rng=None): + """Mask `rate` of the real (unpadded) characters of each row. + + Returns `(masked_chars, targets)`. `targets` is the original id at every + masked position and `IGNORE_INDEX` everywhere else, so a plain + `CrossEntropyLoss(ignore_index=IGNORE_INDEX)` sees loss ONLY on masked + positions -- padding and unmasked characters contribute nothing. + """ + if rng is None: + rng = torch.Generator() + rng.manual_seed(0) + real = mask > 0 + draw = torch.rand(chars.shape, generator=rng) + chosen = real & (draw < rate) + targets = torch.full_like(chars, IGNORE_INDEX) + targets[chosen] = chars[chosen] + out = chars.clone() + out[chosen] = mask_id + return out, targets + + +# -------------------------------------------------------------------------- +# corpus +# -------------------------------------------------------------------------- + +def collect_sentences(n=500000, seed=23, langs=None, limit=None, verbose=True): + """Up to `n` real sentences, balanced across languages as far as each + corpus allows (water-filling: a language with fewer lines than its equal + share gives the remainder to the ones that have more). + + Reuses `sankhya.wild`'s loaders and filters (3-40 words, <= 128 chars, + deduplicated) and drops anything that appears in the wild gold. + """ + from . import wild + + langs = list(langs or wild.LANGS) + gold_keys = wild.gold_wild_keys(langs=langs) + pools = {} + for lang in langs: + rng = random.Random(f"pretrain:{seed}:{lang}") + seen = set() + res = [] + n_seen = 0 + cap = n # never keep more than the whole budget for one language + for src in wild.sources_for(lang): + for text, _score, _signals in wild.scan_source(src, limit=limit): + key = wild._dedup_key(text) + if not key or key in seen or key in gold_keys: + continue + seen.add(key) + n_seen += 1 + if len(res) < cap: + res.append(text) + else: + j = rng.randrange(n_seen) + if j < cap: + res[j] = text + rng.shuffle(res) + pools[lang] = res + if verbose: + print(f" [{lang}] {len(res)} usable lines") + + # water-fill: equal shares, shortfalls redistributed + remaining = n + quotas = {} + for lang in sorted(langs, key=lambda l: len(pools[l])): + share = remaining // max(1, (len(langs) - len(quotas))) + quotas[lang] = min(share, len(pools[lang])) + remaining -= quotas[lang] + + out = [] + for lang in langs: + out.extend({"text": t, "lang": lang} for t in pools[lang][:quotas[lang]]) + random.Random(seed).shuffle(out) + if verbose: + print(f" -> {len(out)} sentences " + + " ".join(f"{l}={quotas[l]}" for l in langs)) + return out + + +def tensorize(texts, char_to_id, max_len=MAX_LEN): + n = len(texts) + chars = np.zeros((n, max_len), dtype=np.int64) + mask = np.zeros((n, max_len), dtype=np.float32) + unk = char_to_id.get("", 1) + for i, t in enumerate(texts): + s = normalize_text(t)[:max_len] + for j, ch in enumerate(s): + chars[i, j] = char_to_id.get(ch, unk) + mask[i, :len(s)] = 1.0 + return torch.from_numpy(chars), torch.from_numpy(mask) + + +# -------------------------------------------------------------------------- + +def main(argv=None): + ap = argparse.ArgumentParser(prog="python -m sankhya.pretrain") + ap.add_argument("--out", default="models/pretrain") + ap.add_argument("--n", type=int, default=500000) + ap.add_argument("--epochs", type=int, default=6) + ap.add_argument("--batch", type=int, default=256) + ap.add_argument("--lr", type=float, default=3e-3) + ap.add_argument("--seed", type=int, default=23) + ap.add_argument("--mask-rate", type=float, default=DEFAULT_MASK_RATE) + ap.add_argument("--lang", default="hi_latn,hi_deva,mr_deva,gu_gujr") + ap.add_argument("--arch", default="v2", choices=list(ARCHS.keys())) + ap.add_argument("--channels", type=int, default=48) + ap.add_argument("--embed-dim", type=int, default=16) + ap.add_argument("--grad-clip", type=float, default=1.0) + ap.add_argument("--limit", type=int, default=None, help="cap lines read per source") + ap.add_argument("--cache", default=None, + help="jsonl of collected sentences; read if present, written if not") + args = ap.parse_args(argv) + + torch.manual_seed(args.seed) + np.random.seed(args.seed) + random.seed(args.seed) + torch.set_num_threads(4) + os.makedirs(args.out, exist_ok=True) + + lang_ids = [x.strip() for x in args.lang.split(",") if x.strip()] + packs = [langbase.get_pack(l) for l in lang_ids] + vocab = build_charset_multi(packs) if len(packs) > 1 else build_charset(packs[0]) + char_to_id = {c: i for i, c in enumerate(vocab)} + mask_id = char_to_id[""] + print(f"charset: {len(vocab)} chars (mask token = , id {mask_id})") + + t0 = time.time() + rows = None + if args.cache and os.path.exists(args.cache): + rows = [json.loads(l) for l in open(args.cache, encoding="utf-8") if l.strip()] + print(f"loaded {len(rows)} cached sentences from {args.cache}") + if rows is None: + rows = collect_sentences(n=args.n, seed=args.seed, langs=lang_ids, limit=args.limit) + if args.cache: + os.makedirs(os.path.dirname(args.cache) or ".", exist_ok=True) + with open(args.cache, "w", encoding="utf-8") as f: + for r in rows: + f.write(json.dumps(r, ensure_ascii=False) + "\n") + print(f"wrote {args.cache}") + print(f"corpus ready in {time.time()-t0:.1f}s") + + chars, cmask = tensorize([r["text"] for r in rows], char_to_id) + n = chars.shape[0] + print(f"tensorized {n} sentences") + + model = PretrainModel(vocab_size=len(vocab), arch=ARCHS[args.arch], + channels=args.channels, embed_dim=args.embed_dim) + print(f"trunk params: {count_params(model.trunk)} " + f"(+{count_params(model.head)} throwaway head)") + + opt = torch.optim.Adam(model.parameters(), lr=args.lr) + steps = math.ceil(n / args.batch) * args.epochs + sched = torch.optim.lr_scheduler.OneCycleLR(opt, max_lr=args.lr, total_steps=steps) + ce = nn.CrossEntropyLoss(ignore_index=IGNORE_INDEX) + gen = torch.Generator() + gen.manual_seed(args.seed) + + t0 = time.time() + for epoch in range(args.epochs): + model.train() + perm = torch.randperm(n) + tot, nb = 0.0, 0 + correct = masked_total = 0 + for i in range(0, n, args.batch): + idx = perm[i:i + args.batch] + cb, mb = chars[idx], cmask[idx] + inp, tgt = apply_mask(cb, mb, mask_id, rate=args.mask_rate, rng=gen) + logits = model(inp) + V = logits.shape[-1] + loss = ce(logits.reshape(-1, V), tgt.reshape(-1)) + opt.zero_grad() + loss.backward() + torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip) + opt.step() + sched.step() + tot += loss.item() + nb += 1 + with torch.no_grad(): + sel = tgt != IGNORE_INDEX + if sel.any(): + correct += (logits.argmax(-1)[sel] == tgt[sel]).sum().item() + masked_total += int(sel.sum().item()) + acc = correct / masked_total if masked_total else 0.0 + print(f"epoch {epoch+1}/{args.epochs} loss={tot/max(nb,1):.4f} " + f"masked_char_acc={acc:.4f} time={time.time()-t0:.1f}s") + + ckpt_path = os.path.join(args.out, "pretrain.pt") + torch.save({ + "state_dict": model.trunk_state_dict(), + "vocab": vocab, + "arch": model.trunk.arch, + "channels": args.channels, + "embed_dim": args.embed_dim, + "seed": args.seed, + "mask_rate": args.mask_rate, + "n_sentences": n, + "epochs": args.epochs, + "pretrain": True, + }, ckpt_path) + print(f"saved trunk to {ckpt_path} ({time.time()-t0:.1f}s total)") + + +if __name__ == "__main__": + main() diff --git a/python/sankhya/train.py b/python/sankhya/train.py index 16ac1af..0fec2cc 100644 --- a/python/sankhya/train.py +++ b/python/sankhya/train.py @@ -61,6 +61,30 @@ def tensorize(examples, char_to_id, max_len=MAX_LEN): ) +def negatives_to_examples(rows, max_len=MAX_LEN): + """Turn `{"text", "lang"}` lines into all-O training examples. + + A real sentence with no amount in it is a perfectly good labelled + example: every character is BIO `O` (0) and class `O` (index 0 in + `classes.CLASSES`). Labels are built at the normalized/truncated length + the tensorizer will actually use, so a long line cannot carry labels + past the model's window. + """ + out = [] + o_cls = C.CLASSES.index("O") + for r in rows: + text = r["text"] + L = min(len(normalize_text(text)), max_len) + out.append({ + "text": text, + "lang": r.get("lang"), + "bio": [0] * L, + "cls": [o_cls] * L, + "spans": [], + }) + return out + + def span_f1(gold_spans_list, pred_spans_list): tp = fp = fn = 0 for gold, pred in zip(gold_spans_list, pred_spans_list): @@ -200,6 +224,17 @@ def main(argv=None): ap.add_argument("--extra-ratio", type=float, default=0.2, help="target share of each epoch's examples drawn from --extra (0 disables mixing); " "--extra is oversampled (with replacement) or subsampled each epoch to hit it") + ap.add_argument("--extra-negatives", nargs="+", default=None, + help="jsonl file(s) of real no-amount lines ({\"text\", \"lang\"}); " + "mixed in as all-O examples at --extra-negative-ratio " + "(see sankhya.wild negatives)") + ap.add_argument("--extra-negative-ratio", type=float, default=0.15, + help="target share of each epoch's examples drawn from " + "--extra-negatives (0 disables mixing)") + ap.add_argument("--init-from", default=None, + help="checkpoint (e.g. sankhya.pretrain output) to initialise the " + "trunk (embedding + conv stack) from; the BIO/class heads stay " + "freshly initialised") args = ap.parse_args(argv) if os.environ.get("SANKHYA_ANOMALY"): @@ -251,12 +286,24 @@ def main(argv=None): f"they map to . Top 20: " + " ".join(f"{c!r}={n}" for c, n in top) ) + neg_ex = [] + if args.extra_negatives: + neg_rows = [] + for np_path in args.extra_negatives: + neg_rows.extend(load_jsonl(np_path)) + neg_ex = negatives_to_examples(neg_rows) + print(f"loaded {len(neg_ex)} real-negative lines from " + f"{len(args.extra_negatives)} file(s) (labelled all-O)") + t0 = time.time() tr_chars, tr_bio, tr_cls, tr_mask = tensorize(train_ex, char_to_id) va_chars, va_bio, va_cls, va_mask = tensorize(val_ex, char_to_id) ex_chars = ex_bio = ex_cls = ex_mask = None if extra_ex: ex_chars, ex_bio, ex_cls, ex_mask = tensorize(extra_ex, char_to_id) + ng_chars = ng_bio = ng_cls = ng_mask = None + if neg_ex: + ng_chars, ng_bio, ng_cls, ng_mask = tensorize(neg_ex, char_to_id) print(f"tensorized in {time.time()-t0:.1f}s") if args.device == "auto" and torch.cuda.is_available(): @@ -271,6 +318,24 @@ def main(argv=None): vocab_size=len(vocab), n_cls=len(C.CLASSES), arch=arch, dilation=args.dilation, channels=args.channels, layers=args.layers, embed_dim=args.embed_dim, crf=args.crf, ).to(device) + if args.init_from: + ck = torch.load(args.init_from, map_location="cpu") + sd = ck.get("state_dict", ck) + trunk = {k: v for k, v in sd.items() + if k.startswith("embed.") or k.startswith("conv")} + model_sd = model.state_dict() + missing = [k for k in model_sd + if (k.startswith("embed.") or k.startswith("conv")) and k not in trunk] + if missing: + raise SystemExit( + f"--init-from {args.init_from}: trunk mismatch, missing {missing}") + bad = [k for k, v in trunk.items() + if k in model_sd and model_sd[k].shape != v.shape] + if bad: + raise SystemExit(f"--init-from {args.init_from}: shape mismatch on {bad}") + model.load_state_dict(trunk, strict=False) + print(f"initialised trunk ({len(trunk)} tensors) from {args.init_from}; heads are fresh") + n_params = count_params(model) print(f"param count: {n_params} (arch={args.arch or 'legacy'} channels={args.channels} layers={len(model.arch)})") @@ -293,7 +358,20 @@ def main(argv=None): f"{n_extra_pool} (target ratio={ratio:.2f}, " f"{'oversampled' if n_extra_per_epoch > n_extra_pool else 'subsampled'})" ) - n = n_main + n_extra_per_epoch + n_neg_pool = ng_chars.shape[0] if ng_chars is not None else 0 + n_neg_per_epoch = 0 + if n_neg_pool and args.extra_negative_ratio > 0: + nratio = min(max(args.extra_negative_ratio, 0.0), 0.95) + # same "share of the combined epoch" arithmetic as --extra, computed + # against the main + extra base so the two ratios compose. + base = n_main + n_extra_per_epoch + n_neg_per_epoch = int(round(base * nratio / (1 - nratio))) + print( + f"mixing in {n_neg_per_epoch} real negatives/epoch from a pool of " + f"{n_neg_pool} (target ratio={nratio:.2f}, " + f"{'oversampled' if n_neg_per_epoch > n_neg_pool else 'subsampled'})" + ) + n = n_main + n_extra_per_epoch + n_neg_per_epoch opt = torch.optim.Adam(model.parameters(), lr=args.lr) n_batches_per_epoch = math.ceil(n / args.batch) @@ -335,6 +413,14 @@ def main(argv=None): ep_mask = torch.cat([tr_mask, ex_mask[extra_idx]], dim=0) else: ep_chars, ep_bio, ep_cls, ep_mask = tr_chars, tr_bio, tr_cls, tr_mask + if n_neg_per_epoch: + replace = n_neg_per_epoch > n_neg_pool + neg_idx = torch.from_numpy( + np.random.choice(n_neg_pool, size=n_neg_per_epoch, replace=replace)) + ep_chars = torch.cat([ep_chars, ng_chars[neg_idx]], dim=0) + ep_bio = torch.cat([ep_bio, ng_bio[neg_idx]], dim=0) + ep_cls = torch.cat([ep_cls, ng_cls[neg_idx]], dim=0) + ep_mask = torch.cat([ep_mask, ng_mask[neg_idx]], dim=0) perm = torch.randperm(n) epoch_loss = 0.0 nb = 0 diff --git a/python/sankhya/wild.py b/python/sankhya/wild.py index f5aebac..5267fbf 100644 --- a/python/sankhya/wild.py +++ b/python/sankhya/wild.py @@ -401,6 +401,123 @@ def load_samples(sample_dir, langs=LANGS): return out +# -------------------------------------------------------------------------- +# real-negative pool for training (`negatives`) +# -------------------------------------------------------------------------- + +GOLD_WILD_GLOB = "gold_wild_{lang}.jsonl" +DEFAULT_NEGATIVES = os.path.join(DATA_WILD, "negatives.jsonl") + + +def gold_wild_keys(tests_dir=None, langs=LANGS): + """Dedup keys for every hand-labelled wild-gold line, so the training + negatives can be asserted disjoint from the evaluation set.""" + tests_dir = tests_dir or os.path.join(_PY_ROOT, "tests") + keys = set() + for lang in langs: + path = os.path.join(tests_dir, GOLD_WILD_GLOB.format(lang=lang)) + if not os.path.exists(path): + continue + with open(path, encoding="utf-8") as f: + for line in f: + line = line.strip() + if not line: + continue + keys.add(_dedup_key(json.loads(line)["text"])) + return keys + + +def _reservoir_no_signal(lang, k, seed, limit=None, verbose=True): + """Uniform reservoir sample of k no-signal lines for one language. + + "No signal" is the same bar `sample` uses for its negatives: the lexicon + scorer returns score 0, i.e. not a unit word, prefix, cardinal, currency + marker or digit anywhere in the line. Reservoir sampling keeps the pass + single and the memory bounded even over the 929k-line Hindi wiki file. + """ + rng = random.Random(f"neg:{seed}:{lang}") + seen_keys = set() + res = [] + n_seen = 0 + for src in sources_for(lang): + for text, score, signals in scan_source(src, limit=limit): + if score != 0 or _is_amount_signal(signals): + continue + key = _dedup_key(text) + if not key or key in seen_keys: + continue + seen_keys.add(key) + row = {"text": text, "lang": lang, "source": src.id} + n_seen += 1 + if len(res) < k: + res.append(row) + else: + j = rng.randrange(n_seen) + if j < k: + res[j] = row + if verbose: + print(f" [{lang}] {n_seen} unique no-signal lines -> reservoir of {len(res)}") + rng.shuffle(res) + return res + + +def build_negatives(n=20000, seed=23, weights_json=None, int8=True, limit=None, + oversample=2.0, langs=LANGS, verbose=True): + """Real lines with no amount signal that the SHIPPED model also leaves + alone, balanced across languages. + + Two filters, deliberately both: the lexicon scorer (cheap, and the same + definition of "no signal" the wild sampling uses) and the shipped weights + (so a line the current model spans -- which might be a real amount the + lexicon missed -- never becomes an all-O training label). Lines that + appear in the hand-labelled wild gold are excluded, and the caller + asserts the disjointness. + """ + if weights_json is None: + weights_json = os.path.join(_PY_ROOT, "..", "models", "default", + "sankhya.weights.int8.json") + gold_keys = gold_wild_keys(langs=langs) + per_lang = int(round(n / len(langs))) + out = [] + stats = {} + for lang in langs: + k = int(round(per_lang * oversample)) + cand = _reservoir_no_signal(lang, k, seed, limit=limit, verbose=verbose) + cand = [r for r in cand if _dedup_key(r["text"]) not in gold_keys] + results = predict(cand, weights_json, int8=int8) + kept = [] + n_spanned = 0 + for r, res in zip(cand, results): + if res["spans"]: + n_spanned += 1 + continue + kept.append(r) + if len(kept) >= per_lang: + break + stats[lang] = {"candidates": len(cand), "model_spanned": n_spanned, + "kept": len(kept), "target": per_lang} + if verbose: + print(f" [{lang}] candidates={len(cand)} model-spanned={n_spanned} " + f"kept={len(kept)} (target {per_lang})") + out.extend(kept) + + keys = {_dedup_key(r["text"]) for r in out} + assert not (keys & gold_keys), ( + "negatives pool overlaps the hand-labelled wild gold -- these lines " + "would leak evaluation text into training" + ) + return out, stats + + +def write_negatives(path, rows): + os.makedirs(os.path.dirname(path) or ".", exist_ok=True) + with open(path, "w", encoding="utf-8") as f: + for r in rows: + f.write(json.dumps({"text": r["text"], "lang": r["lang"], + "source": r["source"]}, ensure_ascii=False) + "\n") + print(f"wrote {len(rows)} negatives to {path}") + + # -------------------------------------------------------------------------- # running the shipped weights, with eval_gold's gates # -------------------------------------------------------------------------- @@ -531,6 +648,19 @@ def main(argv=None): sp.add_argument("--limit", type=int, default=None, help="cap lines read per source") sp.add_argument("--stats-out", default=os.path.join(DATA_WILD, "source_stats.json")) + np_ = sub.add_parser("negatives", help="build an all-O real-negative pool for training") + np_.add_argument("--out", default=DEFAULT_NEGATIVES) + np_.add_argument("--n", type=int, default=20000) + np_.add_argument("--seed", type=int, default=23) + np_.add_argument("--oversample", type=float, default=2.0, + help="how many candidates to draw per kept negative (the shipped " + "model spans a few percent of no-signal lines)") + np_.add_argument("--limit", type=int, default=None, help="cap lines read per source") + np_.add_argument("--weights-json", + default=os.path.join(_PY_ROOT, "..", "models", "default", "sankhya.weights.int8.json")) + np_.add_argument("--int8", action="store_true", default=True) + np_.add_argument("--float32", dest="int8", action="store_false") + rp = sub.add_parser("run", help="run shipped weights over the samples") rp.add_argument("--samples", default=os.path.join(DATA_WILD, "samples")) rp.add_argument("--weights-json", @@ -563,6 +693,14 @@ def main(argv=None): print(f"wrote {args.stats_out}") return + if args.cmd == "negatives": + rows, stats = build_negatives(n=args.n, seed=args.seed, + weights_json=args.weights_json, int8=args.int8, + limit=args.limit, oversample=args.oversample) + write_negatives(args.out, rows) + print(json.dumps(stats, indent=2)) + return + if args.cmd == "run": samples = load_samples(args.samples) if not samples: diff --git a/python/tests/test_pretrain_negatives.py b/python/tests/test_pretrain_negatives.py new file mode 100644 index 0000000..5055152 --- /dev/null +++ b/python/tests/test_pretrain_negatives.py @@ -0,0 +1,265 @@ +"""Data plumbing for the two model-side experiments in data_wild/REPORT.md: + +- `train.negatives_to_examples` / `--extra-negatives` (real all-O sentences + mixed into each epoch at a ratio), +- `sankhya.pretrain`'s masking (15% of real characters, loss only there), +- `train.py --init-from` (trunk copied, heads fresh). + +Everything here is deliberately tiny (a few hundred examples, 1-2 epochs) so +the file runs in seconds on CPU. +""" +import json +import os +import subprocess +import sys +import tempfile +import unittest + +import torch + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) + +from sankhya import classes as C +from sankhya import train as T +from sankhya import pretrain as P +from sankhya.model import SankhyaCNN, ARCHS + +PY_ROOT = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + + +class TestNegativesToExamples(unittest.TestCase): + def test_all_o_labels(self): + rows = [{"text": "अरब सागर के किनारे", "lang": "hi_deva"}, + {"text": "anna hazare ka andolan", "lang": "hi_latn"}] + ex = T.negatives_to_examples(rows) + self.assertEqual(len(ex), 2) + for e, r in zip(ex, rows): + self.assertEqual(e["text"], r["text"]) + self.assertEqual(e["spans"], []) + self.assertTrue(e["bio"]) + self.assertEqual(set(e["bio"]), {0}) + self.assertEqual(set(e["cls"]), {C.CLASSES.index("O")}) + self.assertEqual(len(e["bio"]), len(e["cls"])) + + def test_labels_truncated_to_max_len(self): + long_text = "क " * 200 + ex = T.negatives_to_examples([{"text": long_text, "lang": "hi_deva"}])[0] + self.assertEqual(len(ex["bio"]), T.MAX_LEN) + + def test_tensorizes_as_all_o(self): + vocab = ["", "", "a", "b", " "] + c2i = {c: i for i, c in enumerate(vocab)} + ex = T.negatives_to_examples([{"text": "ab ba", "lang": "hi_latn"}]) + chars, bio, cls, mask = T.tensorize(ex, c2i) + self.assertEqual(int(mask.sum().item()), 5) + self.assertEqual(int(bio.sum().item()), 0) + self.assertEqual(int(cls.sum().item()), 0) + + +def _write_jsonl(path, rows): + with open(path, "w", encoding="utf-8") as f: + for r in rows: + f.write(json.dumps(r, ensure_ascii=False) + "\n") + + +def _synthetic_examples(n, lang="hi_latn"): + """Minimal train/val rows in the generator's schema.""" + out = [] + for i in range(n): + text = f"das hazaar {i}" + out.append({"text": text, "lang": lang, + "bio": [0] * len(text), "cls": [0] * len(text), "spans": []}) + return out + + +class TestExtraNegativeMixing(unittest.TestCase): + """The ratio arithmetic is what a training run actually depends on, so it + is asserted end to end through `train.main` on a tiny model.""" + + def test_ratio_reported_and_epoch_size(self): + with tempfile.TemporaryDirectory() as d: + tr = os.path.join(d, "train.jsonl") + va = os.path.join(d, "val.jsonl") + neg = os.path.join(d, "neg.jsonl") + _write_jsonl(tr, _synthetic_examples(200)) + _write_jsonl(va, _synthetic_examples(20)) + _write_jsonl(neg, [{"text": f"koi amount nahi hai {i}", "lang": "hi_latn"} + for i in range(50)]) + out = os.path.join(d, "models") + import io + import contextlib + buf = io.StringIO() + with contextlib.redirect_stdout(buf): + T.main(["--train", tr, "--val", va, "--out", out, + "--lang", "hi_latn", "--epochs", "1", "--batch", "64", + "--channels", "8", "--embed-dim", "4", + "--extra-negatives", neg, "--extra-negative-ratio", "0.15"]) + log = buf.getvalue() + self.assertIn("loaded 50 real-negative lines", log) + # 200 main examples, ratio 0.15 -> 35 negatives (35/235 = 0.149) + self.assertIn("mixing in 35 real negatives/epoch", log) + self.assertTrue(os.path.exists(os.path.join(out, "sankhya.pt"))) + + def test_ratio_zero_disables(self): + with tempfile.TemporaryDirectory() as d: + tr = os.path.join(d, "train.jsonl") + va = os.path.join(d, "val.jsonl") + neg = os.path.join(d, "neg.jsonl") + _write_jsonl(tr, _synthetic_examples(100)) + _write_jsonl(va, _synthetic_examples(10)) + _write_jsonl(neg, [{"text": "kuch bhi nahi", "lang": "hi_latn"}]) + import io + import contextlib + buf = io.StringIO() + with contextlib.redirect_stdout(buf): + T.main(["--train", tr, "--val", va, "--out", os.path.join(d, "m"), + "--lang", "hi_latn", "--epochs", "1", "--batch", "64", + "--channels", "8", "--embed-dim", "4", + "--extra-negatives", neg, "--extra-negative-ratio", "0"]) + self.assertNotIn("mixing in", buf.getvalue()) + + +class TestPretrainMasking(unittest.TestCase): + def test_mask_rate_and_loss_positions(self): + torch.manual_seed(0) + B, L, V = 64, 100, 30 + chars = torch.randint(2, V, (B, L)) + mask = torch.ones(B, L) + gen = torch.Generator() + gen.manual_seed(7) + inp, tgt = P.apply_mask(chars, mask, mask_id=1, rate=0.15, rng=gen) + sel = tgt != P.IGNORE_INDEX + rate = sel.float().mean().item() + self.assertAlmostEqual(rate, 0.15, delta=0.02) + # masked positions carry the mask id in the input and the ORIGINAL id + # as the target; every other position is untouched and ignored. + self.assertTrue((inp[sel] == 1).all()) + self.assertTrue((tgt[sel] == chars[sel]).all()) + self.assertTrue((inp[~sel] == chars[~sel]).all()) + self.assertTrue((tgt[~sel] == P.IGNORE_INDEX).all()) + + def test_padding_never_masked(self): + chars = torch.randint(2, 20, (16, 50)) + mask = torch.zeros(16, 50) + mask[:, :10] = 1.0 + gen = torch.Generator() + gen.manual_seed(1) + inp, tgt = P.apply_mask(chars, mask, mask_id=1, rate=0.5, rng=gen) + self.assertTrue((tgt[:, 10:] == P.IGNORE_INDEX).all()) + self.assertTrue((inp[:, 10:] == chars[:, 10:]).all()) + + def test_loss_ignores_unmasked_positions(self): + """Changing the logits at an UNMASKED position must not move the loss.""" + torch.manual_seed(0) + V = 12 + ce = torch.nn.CrossEntropyLoss(ignore_index=P.IGNORE_INDEX) + logits = torch.randn(1, 8, V) + tgt = torch.full((1, 8), P.IGNORE_INDEX) + tgt[0, 2] = 5 + base = ce(logits.reshape(-1, V), tgt.reshape(-1)).item() + bumped = logits.clone() + bumped[0, 6] += 100.0 # an unmasked position + self.assertAlmostEqual(ce(bumped.reshape(-1, V), tgt.reshape(-1)).item(), base, places=6) + bumped2 = logits.clone() + bumped2[0, 2, 5] += 100.0 # the masked position + self.assertLess(ce(bumped2.reshape(-1, V), tgt.reshape(-1)).item(), base) + + def test_trunk_state_dict_keys_match_tagger(self): + m = P.PretrainModel(vocab_size=40, arch=ARCHS["v2"], channels=8, embed_dim=4) + tagger = SankhyaCNN(vocab_size=40, n_cls=len(C.CLASSES), arch=ARCHS["v2"], + channels=8, embed_dim=4) + trunk = m.trunk_state_dict() + tsd = tagger.state_dict() + for k, v in trunk.items(): + self.assertIn(k, tsd) + self.assertEqual(tuple(v.shape), tuple(tsd[k].shape)) + # every trunk tensor of the tagger is covered + for k in tsd: + if k.startswith("embed.") or k.startswith("conv"): + self.assertIn(k, trunk) + # and the heads are NOT part of it + self.assertFalse(any(k.startswith(("bio_head", "cls_head")) for k in trunk)) + + +class TestInitFrom(unittest.TestCase): + def test_loads_trunk_leaves_heads_fresh(self): + with tempfile.TemporaryDirectory() as d: + tr = os.path.join(d, "train.jsonl") + va = os.path.join(d, "val.jsonl") + _write_jsonl(tr, _synthetic_examples(64)) + _write_jsonl(va, _synthetic_examples(8)) + + # build a pretrain-shaped checkpoint on the real charset + from sankhya.charset import build_charset + from sankhya.langs import base as langbase + vocab = build_charset(langbase.get_pack("hi_latn")) + pm = P.PretrainModel(vocab_size=len(vocab), arch=None, channels=8, embed_dim=4) + for p in pm.trunk.parameters(): + torch.nn.init.constant_(p, 0.25) + ck = os.path.join(d, "pretrain.pt") + torch.save({"state_dict": pm.trunk_state_dict(), "vocab": vocab}, ck) + + out = os.path.join(d, "m") + import io + import contextlib + buf = io.StringIO() + with contextlib.redirect_stdout(buf): + T.main(["--train", tr, "--val", va, "--out", out, + "--lang", "hi_latn", "--epochs", "1", "--batch", "64", + "--channels", "8", "--embed-dim", "4", "--lr", "1e-9", + "--init-from", ck]) + self.assertIn("initialised trunk", buf.getvalue()) + self.assertIn("heads are fresh", buf.getvalue()) + + got = torch.load(os.path.join(out, "sankhya.pt"), map_location="cpu")["state_dict"] + # with lr ~ 0 the trunk should still be the constant we saved + self.assertTrue(torch.allclose(got["conv1.weight"], + torch.full_like(got["conv1.weight"], 0.25), + atol=1e-4)) + # heads were never in the pretrain checkpoint, so they are the + # fresh torch init -- not the constant. + self.assertFalse(torch.allclose(got["bio_head.weight"], + torch.full_like(got["bio_head.weight"], 0.25), + atol=1e-3)) + + def test_shape_mismatch_is_fatal(self): + with tempfile.TemporaryDirectory() as d: + tr = os.path.join(d, "train.jsonl") + va = os.path.join(d, "val.jsonl") + _write_jsonl(tr, _synthetic_examples(16)) + _write_jsonl(va, _synthetic_examples(4)) + from sankhya.charset import build_charset + from sankhya.langs import base as langbase + vocab = build_charset(langbase.get_pack("hi_latn")) + pm = P.PretrainModel(vocab_size=len(vocab), arch=None, channels=16, embed_dim=4) + ck = os.path.join(d, "pretrain.pt") + torch.save({"state_dict": pm.trunk_state_dict(), "vocab": vocab}, ck) + with self.assertRaises(SystemExit): + T.main(["--train", tr, "--val", va, "--out", os.path.join(d, "m"), + "--lang", "hi_latn", "--epochs", "1", "--batch", "16", + "--channels", "8", "--embed-dim", "4", "--init-from", ck]) + + +class TestWildNegativesTool(unittest.TestCase): + def test_gold_wild_keys_nonempty(self): + from sankhya import wild + keys = wild.gold_wild_keys() + self.assertGreaterEqual(len(keys), 350) + + def test_committed_negatives_are_disjoint_from_gold(self): + """If the (gitignored) pool has been built locally, it must not share + a line with the evaluation set.""" + from sankhya import wild + path = wild.DEFAULT_NEGATIVES + if not os.path.exists(path): + self.skipTest("data_wild/negatives.jsonl not built here") + gold = wild.gold_wild_keys() + with open(path, encoding="utf-8") as f: + for line in f: + line = line.strip() + if line: + self.assertNotIn(wild._dedup_key(json.loads(line)["text"]), gold) + + +if __name__ == "__main__": + unittest.main()