palimpseste-max / examples /train_1b.py
thefinalboss's picture
Upload examples/train_1b.py with huggingface_hub
f57cdb9 verified
Raw
History Blame Contribute Delete
5.65 kB
#!/usr/bin/env python
"""PALIMPSESTE — Train the 1B-scale model on real data.
Combines the built-in conversational corpus with a real trivia dataset
(TriviaQA, 3000 Q/A pairs) and trains at D=100,000 (the '1b' preset —
theoretical capacity ~2^25000 ≈ 10^7525 associations).
Usage:
python examples/train_1b.py --output ./palimpseste-1b --n-trivia 2000
python examples/train_1b.py --output ./palimpseste-1b --push-to-hub thefinalboss/palimpseste-1b
"""
from __future__ import annotations
import argparse
import json
import sys
import time
from pathlib import Path
import numpy as np
sys.path.insert(0, str(Path(__file__).resolve().parent))
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
from palimseste.lm import PalimpsesteConfig, PRESETS
from palimseste.hf import HFPalimpsesteLM
from corpus_chat import get_corpus as get_chat_corpus
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(
description="Train the PALIMPSESTE 1B-scale model on real data.",
formatter_class=argparse.RawDescriptionHelpFormatter,
)
p.add_argument("--output", "-o", type=str, required=True)
p.add_argument("--n-trivia", type=int, default=2000,
help="number of TriviaQA pairs to use (default 2000)")
p.add_argument("--trivia-file", type=str, default="trivia_qa_pairs.json",
help="path to extracted TriviaQA pairs JSON")
p.add_argument("--D", type=int, default=100000,
help="hypervector dimension (default 100000 = 1b preset)")
p.add_argument("--context-window", type=int, default=80)
p.add_argument("--kernel-radius", type=int, default=800)
p.add_argument("--push-to-hub", type=str, default=None)
p.add_argument("--seed", type=int, default=42)
return p.parse_args()
def load_trivia(path: str, n: int) -> list[tuple[str, str]]:
p = Path(path)
if not p.exists():
print(f"warning: {path} not found, downloading TriviaQA...", file=sys.stderr)
from datasets import load_dataset
ds = load_dataset('mandarjoshi/trivia_qa', 'rc.nocontext',
split='train', streaming=True)
pairs = []
for i, ex in enumerate(ds):
q = ex['question'].strip().lower()
ans = ex['answer']['value'].strip().lower()
if q and ans and len(ans) < 100:
pairs.append((q, ans))
if i >= n * 2:
break
with open(path, 'w', encoding='utf-8') as f:
json.dump([{'q': q, 'a': a} for q, a in pairs], f, ensure_ascii=False)
return pairs[:n]
with open(path, 'r', encoding='utf-8') as f:
data = json.load(f)
return [(d['q'], d['a']) for d in data[:n]]
def main() -> None:
args = parse_args()
# load data
chat_pairs = get_chat_corpus()
trivia_pairs = load_trivia(args.trivia_file, args.n_trivia)
all_pairs = chat_pairs + trivia_pairs
print("=" * 70)
print("PALIMPSESTE — 1B-SCALE TRAINING ON REAL DATA")
print("=" * 70)
cap_log2 = float(np.log2(0.14 * args.D) + args.D / 4.0)
print(f"D = {args.D:,} (theoretical capacity: 2^{cap_log2:.0f} ≈ 10^{cap_log2*0.30103:.0f})")
print(f"context_window = {args.context_window} kernel_radius = {args.kernel_radius}")
print(f"chat corpus: {len(chat_pairs)} pairs")
print(f"trivia corpus: {len(trivia_pairs)} pairs")
print(f"total: {len(all_pairs)} pairs")
print()
config = PalimpsesteConfig(
D=args.D,
context_window=args.context_window,
kernel_radius=args.kernel_radius,
temperature=0.3,
)
rng = np.random.default_rng(args.seed)
lm = HFPalimpsesteLM(config=config, rng=rng)
# build vocab from all text
full_text = "".join(q + a for q, a in all_pairs)
lm.build_tokenizer(full_text)
print(f"vocab: {lm.tokenizer.vocab_size} chars")
# train
t0 = time.perf_counter()
n_tokens = lm.train_on_qa_pairs(all_pairs, verbose=True)
dt = time.perf_counter() - t0
print(f"\ntrained: {n_tokens:,} tokens in {dt:.1f}s ({n_tokens/dt:.0f} tok/s)")
print(f"|M| = {len(lm.mem):,} traces")
# test on chat questions + a few trivia
print("\n--- chat quality test ---")
chat_tests = ["bonjour", "qui es-tu", "tu utilises un gpu", "who are you",
"what is python", "merci"]
for q in chat_tests:
r = lm.respond(q, temperature=0.0, seed=0)
ok = "✓" if r and "desole" not in r.lower() else "✗"
print(f" {ok} {q:30s} -> {r[:70]}")
print("\n--- trivia recall test (exact match required) ---")
trivia_tests = trivia_pairs[:10]
correct = 0
for q, expected in trivia_tests:
r = lm.respond(q, temperature=0.0, seed=0)
ok = "✓" if expected in r else "✗"
if expected in r:
correct += 1
print(f" {ok} {q[:50]:50s} -> {r[:40]} (expected: {expected})")
print(f" trivia recall: {correct}/{len(trivia_tests)}")
# save
out = Path(args.output)
print(f"\nsaving to {out} ...")
lm.save_pretrained(out)
sz = sum(f.stat().st_size for f in out.rglob('*') if f.is_file())
print(f"saved: {sz / 1e6:.1f} MB")
# push
if args.push_to_hub:
print(f"\npushing to HF Hub: {args.push_to_hub} ...")
url = lm.push_to_hub(args.push_to_hub)
print(f"pushed: {url}")
print("\n" + "=" * 70)
print("DONE. Load with:")
print(f" from palimseste.hf import HFPalimpsesteLM")
print(f" lm = HFPalimpsesteLM.from_pretrained('{args.push_to_hub or out}')")
print("=" * 70)
if __name__ == "__main__":
main()