precompute_predict.py
Computes the figures on /predict/
80 lines. This is the file the build actually runs, copied verbatim at build time.
- 1
"""Generate predict/data.json — every figure on /predict/.""" - 2
- 3
import json - 4
import os - 5
- 6
import ngram - 7
from corpora import TALE - 8
- 9
OUT = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", - 10
"predict", "data.json") - 11
- 12
ORDER = 2 - 13
CONTEXT = "it was the" - 14
SEED = 11 - 15
LENGTH = 22 - 16
- 17
- 18
def marked(full, kept): - 19
"""Full distribution with a flag for whether each token survived the cut.""" - 20
keep = {t for t, _ in kept} - 21
reshaped = dict(kept) - 22
return [{"token": t, "p": p, "kept": t in keep, "after": reshaped.get(t, 0.0)} - 23
for t, p in full] - 24
- 25
- 26
def main(): - 27
trained = ngram.train(TALE, ORDER) - 28
context = ngram.tokenize(CONTEXT) - 29
items, used = ngram.distribution(trained, context) - 30
- 31
data = { - 32
"corpus": TALE, - 33
"corpus_tokens": len(trained["tokens"]), - 34
"corpus_vocab": len(trained["vocab"]), - 35
"order": ORDER, - 36
"context": CONTEXT, - 37
"context_order_used": used, - 38
"distribution": [{"token": t, "p": p} for t, p in items], - 39
"temperatures": [ - 40
{"temp": temp, - 41
"items": [{"token": t, "p": p} - 42
for t, p in ngram.apply_temperature(items, temp)]} - 43
for temp in [0.0, 0.5, 1.0, 2.0] - 44
], - 45
"top_k": {"k": 3, "items": marked(items, ngram.apply_top_k(items, 3))}, - 46
"top_p": {"p": 0.5, "items": marked(items, ngram.apply_top_p(items, 0.5))}, - 47
"samples": [ - 48
{"label": label, "temp": temp, "top_k": k, "top_p": p, - 49
"text": ngram.generate(trained, CONTEXT, LENGTH, temp, k, p, - 50
SEED)["text"]} - 51
for label, temp, k, p in [ - 52
("Greedy (T = 0)", 0.0, 0, 0.0), - 53
("T = 0.7", 0.7, 0, 0.0), - 54
("T = 1.0", 1.0, 0, 0.0), - 55
("T = 2.5", 2.5, 0, 0.0), - 56
] - 57
], - 58
"orders": [ - 59
{"order": o, - 60
"text": ngram.generate(ngram.train(TALE, o), "it was", 18, 1.0, - 61
0, 0.0, 5)["text"]} - 62
for o in [2, 1, 0] - 63
], - 64
"seed": SEED, - 65
} - 66
- 67
with open(OUT, "w", encoding="utf-8") as fh: - 68
json.dump(data, fh, ensure_ascii=False, indent=1) - 69
- 70
print(f"wrote predict/data.json") - 71
print(f" corpus {data['corpus_tokens']} tokens, vocab {data['corpus_vocab']}") - 72
print(f" after {CONTEXT!r} (order {used}): " - 73
f"{len(data['distribution'])} candidates") - 74
for s in data["samples"]: - 75
print(f" {s['label']:<16} {s['text'][:56]!r}") - 76
- 77
- 78
if __name__ == "__main__": - 79
main()