Sillage / app.py
riscoss's picture
sillage 1.0.2; regime table: measured rank-16 figure (16.8 at 7.4 MB)
7482550 verified
Raw
History Blame Contribute Delete
20.3 kB
"""Sillage — Hugging Face Space.
Three screens, in the order that convinces:
1. "It has already read this paper." A memory built at deploy time is
already loaded, so the very first thing a visitor sees is a frozen model
and a remembering model completing the same sentence side by side. No
waiting, no setup, no account.
2. "Give it your own text." Half a minute, then the same comparison on
THEIR document, plus the list of predictions the memory actually
corrected -- the part that turns a number into a fact.
3. "Where it does not work." The regime boundary, stated by the authors
before anyone else has to point it out.
Everything here runs on `pip install sillage`; the read loop below is the
library's own, instrumented to report which tokens changed.
On ZeroGPU the visitor pays quota only for time spent inside a @spaces.GPU
function, so the split matters: the frozen forward passes are batched into
one decorated call, the Hebbian mechanism stays in numpy on the CPU, and the
suggested prompts of screen 1 are served from runs recorded at build time.
Looking at the page costs nothing; a click costs a second or two.
"""
import collections
import io
import json
import os
import re
import shutil
import tempfile
import gradio as gr
import numpy as np
import torch
from sillage import Sillage
from sillage.core import CAP
try: # ZeroGPU: a real GPU is attached only
import spaces # inside a decorated function
GPU = spaces.GPU
except ImportError: # anywhere else the decorator is a no-op
def GPU(*args, **kwargs):
if args and callable(args[0]):
return args[0]
return lambda fn: fn
import sample
STATE = os.path.join(os.path.dirname(os.path.abspath(__file__)), "state")
MAX_TOKENS = 1600 # 160 MB of logits, ~30 s end to end
MIN_TOKENS = 900 # below this, the split has nothing to say
WINDOW, STRIDE = 1024, 512
PAPER = "Sillage (paper 1), 8 969 tokens"
# Chosen by running them: each one makes the two columns visibly diverge,
# and the right-hand column recites the paper rather than inventing.
EXAMPLE_PROMPTS = [
"On a 36k-token stream of novel technical text, the memory",
"At 500k tokens the fixed matrix",
"surprise gating quadruples the gain of",
"Every write is gated by the model's own",
]
# --------------------------------------------------------------- backend ---
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
memory = Sillage(model="gpt2", state=STATE, device=DEVICE, quiet=True)
frozen = Sillage(model="gpt2", state=tempfile.mkdtemp(), device=DEVICE,
quiet=True)
# Each visitor who reads a document leaves a 7.4 MB state behind. Keep the
# most recent few so the follow-up completion still works, and delete the
# rest: a public Space should not fill its disk with strangers' memories.
SESSIONS = collections.deque(maxlen=8)
def new_session_dir():
if len(SESSIONS) == SESSIONS.maxlen:
shutil.rmtree(SESSIONS[0], ignore_errors=True)
d = tempfile.mkdtemp(prefix="sillage-session-")
SESSIONS.append(d)
return d
def models():
"""Load GPT-2 once and let both assistants share the frozen weights."""
tok, model = memory.load_model()
frozen._tok, frozen._model = tok, model
frozen.device = memory.device
return tok, model
# ZeroGPU asks for the weights to be placed at startup rather than inside a
# decorated call, where a cold load would eat into the visitor's quota and
# could outlast the function's duration budget.
models()
# Screen 1's suggested prompts were run once, at build time, into data.json.
# Greedy decoding over a fixed memory is deterministic, so a cached answer is
# the same answer -- and serving it costs the visitor no GPU quota at all.
# Anything a visitor types instead is computed live.
CACHED = {}
try:
for _c in json.load(io.open(os.path.join(
os.path.dirname(os.path.abspath(__file__)), "data.json"),
encoding="utf-8"))["completions"]:
CACHED[_c["prompt"].strip()] = (
_c["prompt"] + _c["frozen"], _c["prompt"] + _c["memory"],
_c["same"])
except Exception: # no capture yet: everything runs live
CACHED = {}
def compare(prompt, n_tokens=24):
"""Screen 1. Recorded prompts are served from cache, the rest is live."""
hit = CACHED.get((prompt or "").strip())
if hit:
return hit[0], hit[1], _verdict(hit[2])
return compare_live(prompt, n_tokens)
@GPU(duration=30)
def compare_live(prompt, n_tokens=24):
"""The decorated path: generation is the only part that needs the model.
Twenty-four tokens twice is a second or two of GPU, which is what a
visitor's daily ZeroGPU quota can afford many times over.
"""
return _compare(prompt, n_tokens)
def _verdict(same):
return ("identical here — the memory abstains when it is not confident, "
"which is the point of the abstention threshold" if same else
"the two columns diverge: everything after the prompt on the "
"right comes from what was read, not from GPT-2's weights")
@GPU(duration=60)
@torch.no_grad()
def forward_all(ids):
"""Every frozen forward pass for a document, in one GPU visit.
The Hebbian loop that follows is numpy on the CPU: keeping it outside
this function is what stops a 30-second read from costing 30 seconds of
the visitor's GPU quota. Returned in fp16 -- 1600 positions of GPT-2
logits are 160 MB, and the loop only ever reads them.
"""
tok, model = models()
x = torch.tensor(ids, device=memory.device)
n = len(ids) - 1
logits = np.empty((n + 1, model.config.vocab_size), dtype=np.float16)
hidden = np.empty((n + 1, model.config.n_embd), dtype=np.float16)
a = 0
while a < n:
w = min(WINDOW, len(ids) - a)
out = model(x[a:a + w].unsqueeze(0), output_hidden_states=True)
lo = 0 if a == 0 else WINDOW - STRIDE
logits[a + lo:a + w] = (
out.logits[0, lo:w].to(torch.float16).cpu().numpy())
hidden[a + lo:a + w] = (
out.hidden_states[-1][0, lo:w].to(torch.float16).cpu().numpy())
if a + w >= len(ids):
break
a += STRIDE
return logits, hidden
def _compare(prompt, n_tokens=24):
"""The same prompt, decoded greedily, with and without the memory."""
if not prompt.strip():
return "", "", ""
models()
a = frozen.complete(prompt, n=n_tokens)
b = memory.complete(prompt, n=n_tokens)
return prompt + a, prompt + b, _verdict(a.strip() == b.strip())
def suggest_prompt(text, tail):
"""Half of a sentence whose opening recurs — so the memory knows the rest.
A prompt taken at random usually lands where the memory has nothing to
say, and the demo falls flat through no fault of the mechanism. Picking a
sentence whose first words already appeared earlier is the honest way to
show the memory doing what it is for.
"""
best = ""
for sentence in re.split(r"(?<=[.:])\s+|\n+", tail):
words = sentence.split()
if len(words) < 9:
continue
if text.count(" ".join(words[:5])) >= 2:
cut = " ".join(words[:max(5, len(words) // 2)])
if len(cut) > len(best):
best = cut
return best or " ".join(tail.split()[:10])
def read_and_report(text, progress=gr.Progress()):
"""Stream a document through a fresh memory, reporting what it changed.
This is `sillage.runtime.Sillage.read_text`, unrolled for two reasons:
the demo collects the positions where the memory turned a wrong
prediction into a right one, and the frozen forward passes are done in
one batch beforehand (on the GPU when there is one) so that this loop --
plain numpy, no model -- costs the visitor nothing. Strictly prequential
either way: every token is scored before it is written.
"""
text = (text or "").strip()
tok, lm = models()
state_dir = new_session_dir()
fresh = Sillage(model="gpt2", state=state_dir, device=memory.device,
quiet=True)
fresh._tok, fresh._model = tok, lm
mem = fresh.mem
ids = np.array(tok.encode(text), dtype=np.int64)[:MAX_TOKENS]
n = len(ids) - 1
if n < MIN_TOKENS:
return (None, f"That is {max(n, 0)} tokens. Below about "
f"{MIN_TOKENS} there is nothing to measure honestly: the "
f"first half of a document is what builds the memory, and "
f"the abstention threshold needs a few hundred observations "
f"before it will let the memory speak at all. Paste a longer "
f"one — or click an example below.", None, "", "")
dev_end = n // 2 # first half: the memory reads and calibrates
mem.new_stream()
thrG = thrS = np.inf # silent until the dev half is behind us
nll_b = nll_f = nll_m = 0.0
cnt = active = 0
fixes = []
logits, hidden = forward_all(ids) # the only GPU visit
for j in range(n):
if j == dev_end: # the papers' protocol: dev decides, test tells
thrG, thrS = mem.thresholds()
truth = int(ids[j + 1])
lb = logits[j].astype(np.float32)
mx = lb.max()
lpb = lb - (mx + np.log(np.exp(lb - mx).sum()))
lp = float(lpb[truth])
la, phi = mem.adapt(lb, hidden[j].astype(np.float32))
m2 = la.max()
p_ad = np.exp(la - m2)
p_ad /= p_ad.sum()
lp_f = float(np.log(max(p_ad[truth], 1e-30)))
qG = mem.step_key(int(ids[j]))
uG, sG = mem.scores(mem.M, qG)
mem.res_G.append(float(sG.max()))
pc = mem.cold_lookup(truth)
p_true = mem.mix_true(np.exp(lp_f), sG, truth, None, pc, thrG, thrS)
speaks = float(sG.max()) >= thrG
if j >= dev_end:
active += speaks
nll_b += -lp
nll_f += -lp_f
nll_m += -np.log(max(p_true, 1e-30))
cnt += 1
if speaks and len(fixes) < 12:
# mix_true takes the probability of one token, mix_full takes the
# whole distribution -- so the cold tier is asked twice, once for
# each shape
full = mem.mix_full(p_ad.copy(), sG, None, mem.cold_lookup(),
thrG, thrS)
said, now = int(np.argmax(p_ad)), int(np.argmax(full))
if now == truth != said:
ctx = tok.decode(ids[max(0, j - 9):j + 1])
fixes.append([" ".join(ctx.split())[-58:],
repr(tok.decode([truth])),
repr(tok.decode([said]))])
g = min(CAP, max(0.0, -lp))
mem.write_all(qG, uG, None, None, truth, g, phi, p_ad)
if j % 128 == 0:
progress(j / n, desc=f"reading {j}/{n} tokens")
fresh.index.add(text, "your document")
fresh.save()
ppl = [float(np.exp(v / cnt)) for v in (nll_b, nll_f, nll_m)]
plot = {"frozen GPT-2": ppl[0], "+ fast weights": ppl[1],
"+ memory": ppl[2]}
knn = n * 1536 * 4 / 1e6 # one 1536-float key+value per token
summary = (
f"**Read once, left to right: {dev_end} tokens built the memory, the "
f"next {cnt} were measured with it.** That split is the papers' own "
f"protocol — the numbers below are on text the memory had not seen "
f"when it scored them, and every token was scored *before* being "
f"written.\n\n"
f"Perplexity on that second half: {ppl[0]:.2f} frozen → "
f"{ppl[1]:.2f} with the rank-16 adapter → **{ppl[2]:.2f}** with the "
f"memory on top — **{100 * (1 - ppl[2] / ppl[0]):.0f}% lower**. The "
f"memory spoke on {100 * active / cnt:.0f}% of those positions and "
f"kept quiet on the rest; abstaining when it has nothing to say is "
f"what keeps it from doing harm.\n\n"
f"State on disk: **7.4 MB**, and it would still be 7.4 MB after a "
f"million tokens. A kNN-LM datastore over the same text would "
f"already hold about {knn:.1f} MB, and would keep growing.")
return (plot, summary, fixes or None, state_dir,
suggest_prompt(text, tok.decode(ids[dev_end:])))
def read_then_compare(text, progress=gr.Progress()):
"""Read, report, and hand back a prompt worth trying plus the new state."""
plot, summary, fixes, state_dir, suggestion = read_and_report(text,
progress)
return (plot, summary, fixes, state_dir,
gr.update(value=suggestion, visible=bool(state_dir)),
gr.update(visible=bool(state_dir)))
@GPU(duration=30)
def compare_user(prompt, state_dir, n_tokens=12):
"""The same comparison, against the memory the visitor just built."""
if not state_dir or not os.path.exists(os.path.join(state_dir,
"state.npz")):
return "", "", ("That session's memory has been cleared — read a "
"document again.")
if not (prompt or "").strip():
return "", "", ""
tok, model = models()
theirs = Sillage(model="gpt2", state=state_dir, device=memory.device,
quiet=True)
theirs._tok, theirs._model = tok, model
a = frozen.complete(prompt, n=n_tokens)
b = theirs.complete(prompt, n=n_tokens)
note = ("identical here — try a phrase that recurs in your document, "
"the memory only speaks where it is confident"
if a.strip() == b.strip() else
"the right-hand column comes from your document, not from "
"GPT-2's weights")
return prompt + a, prompt + b, note
# ------------------------------------------------------------------ ui -----
CSS = """
.hero {text-align:center}
.small {font-size:0.9em; opacity:0.75}
footer {visibility:hidden}
"""
with gr.Blocks(title="Sillage: a frozen LM that remembers "
"what it reads") as demo:
gr.Markdown(
"# Sillage\n"
"### A frozen language model that remembers what it reads — "
"4.2 MB, no gradients, no fine-tuning, no vector database.\n"
"[![PyPI](https://img.shields.io/pypi/v/sillage.svg)]"
"(https://pypi.org/project/sillage/) "
"[![GitHub](https://img.shields.io/badge/code-GitHub-181717?logo=github)]"
"(https://github.com/riscoss63/sillage) "
"[![DOI](https://zenodo.org/badge/DOI/10.5281/zenodo.22079016.svg)]"
"(https://doi.org/10.5281/zenodo.22079016)",
elem_classes="hero")
gr.Markdown(
"> This Space runs **GPT-2 124M on a free CPU**. Its prose is weak in "
"absolute terms — that is not what is on display. What is on display "
"is the **difference** between the two columns, and it comes from a "
"4.2 MB matrix written while reading, with no gradient anywhere.",
elem_classes="small")
with gr.Tab("1 · It has already read a paper"):
gr.Markdown(
f"The memory loaded here has read **{PAPER}** — the paper that "
"describes this very mechanism. GPT-2 has never seen that text. "
"Complete a sentence from it and watch the right-hand column "
"recall what the left one cannot know.")
prompt = gr.Textbox(label="Beginning of a sentence",
value=EXAMPLE_PROMPTS[0], lines=2)
gr.Examples(examples=[[p] for p in EXAMPLE_PROMPTS], inputs=prompt,
label="Try one of these")
go = gr.Button("Complete both", variant="primary")
with gr.Row():
out_frozen = gr.Textbox(label="GPT-2, frozen", lines=6)
out_memory = gr.Textbox(label="GPT-2 + Sillage memory", lines=6)
verdict = gr.Markdown(elem_classes="small")
go.click(compare, prompt, [out_frozen, out_memory, verdict])
demo.load(compare, prompt, [out_frozen, out_memory, verdict])
with gr.Tab("2 · Give it your own text"):
gr.Markdown(
f"Paste something GPT-2 has never seen: your notes, an internal "
f"document, a README, a specification. It is read **once**, left "
f"to right, and every token is scored *before* it is memorised — "
f"so the numbers below are honest online measurements, not a "
f"replay. Capped at {MAX_TOKENS} tokens (about 90 seconds on this "
f"free CPU).")
user_text = gr.Textbox(label="Your document", lines=12,
value=sample.MANUAL)
gr.Examples(
examples=[[sample.MANUAL], [sample.paper_excerpt()]],
inputs=user_text, example_labels=[
"An invented operations manual (GPT-2 cannot know it)",
"The opening of paper 1 (novel technical prose)"],
label="Or start from one of these")
read_btn = gr.Button("Read it", variant="primary")
chart = gr.Label(label="Perplexity, lower is better")
report = gr.Markdown()
fixed = gr.Dataframe(
headers=["context", "what came next", "what frozen GPT-2 said"],
label="Predictions the memory corrected — the frozen model had "
"no way of knowing these", wrap=True)
gr.Markdown("### Now ask it to continue a sentence from *your* text")
session = gr.State("")
your_prompt = gr.Textbox(label="Beginning of a sentence", lines=2,
visible=False)
your_btn = gr.Button("Complete both", variant="primary",
visible=False)
with gr.Row():
your_frozen = gr.Textbox(label="GPT-2, frozen", lines=5)
your_memory = gr.Textbox(label="GPT-2 + the memory you just "
"built", lines=5)
your_note = gr.Markdown(elem_classes="small")
read_btn.click(read_then_compare, user_text,
[chart, report, fixed, session, your_prompt, your_btn])
your_btn.click(compare_user, [your_prompt, session],
[your_frozen, your_memory, your_note])
with gr.Tab("3 · Where it does not work"):
gr.Markdown(
"**The regime matters, and it is measured.**\n\n"
"| system | perplexity | memory used |\n|---|---|---|\n"
"| frozen GPT-2 | 31.2 | 0 |\n"
"| + RAG-style retrieve & rescore | 29.9 | corpus + index |\n"
"| + kNN-LM, unbounded store | 23.6 | 55 MB, grows forever |\n"
"| + **this memory** (fixed) | **19.2** | **4.2 MB, constant** |\n"
"| + memory and fast weights | **16.8** | 7.4 MB, constant |\n\n"
"That is on 36k tokens of novel technical text. **On long, "
"low-repetition narrative, an unbounded kNN-LM still wins** "
"(+0.048 vs +0.007 nats): this memory captures verbatim "
"recurrence, not paraphrase. The boundary is mapped in the "
"papers rather than hidden.\n\n"
"Three results that did *not* work are published too: hidden "
"states make poor Hebbian keys; surprise gating helps the memory "
"and *hurts* the fast-weight adapter; and calibrating the readout "
"on your own stream loses to a proper tuning, because the "
"calibration window is read by a colder memory than the one it "
"will govern.\n\n"
"---\n\n"
"**Run it yourself, on any causal LM:**\n\n"
"```bash\npip install sillage\nsillage read notes.md\n"
"sillage ask \"what did the report say?\"\n"
"sillage complete \"The report said\"\n```\n\n"
"Four preprints with DOIs, the full reproduction pipeline and "
"every number as committed JSON: "
"[github.com/riscoss63/sillage](https://github.com/riscoss63/sillage)")
if __name__ == "__main__":
# one reading at a time: the whole point is a single CPU
demo.queue(default_concurrency_limit=1).launch(
theme=gr.themes.Soft(), css=CSS)