sweedworks

← all sources

precompute_merges.py

Computes the figures on /vocabulary/

66 lines. This is the file the build actually runs, copied verbatim at build time.

  1. 1"""Generate merges/data.json — every figure on /merges/, computed not written."""
  2. 2 
  3. 3import json
  4. 4import os
  5. 5 
  6. 6import bpe
  7. 7from cjsload import Reference
  8. 8from corpora import BERRIES
  9. 9 
  10. 10OUT = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..",
  11. 11 "vocabulary", "data.json")
  12. 12 
  13. 13N_MERGES = 30
  14. 14 
  15. 15# Words that show the three outcomes: learned whole, learned in pieces, unseen.
  16. 16PROBES = [" strawberry", "strawberry", " blackberry", "strawberries", " kiwi"]
  17. 17 
  18. 18# Where the assembly of " strawberry" visibly changes.
  19. 19CHECKPOINTS = [0, 1, 2, 3, 4, 16, 17]
  20. 20 
  21. 21 
  22. 22def main():
  23. 23 trained = bpe.train(BERRIES, N_MERGES)
  24. 24 merges = trained["merges"]
  25. 25 real = Reference("o200k_base")
  26. 26 
  27. 27 data = {
  28. 28 "corpus": BERRIES,
  29. 29 "corpus_chars": len(BERRIES),
  30. 30 "corpus_words": len(bpe.pretokenize(BERRIES)),
  31. 31 "corpus_unique": len(set(bpe.pretokenize(BERRIES))),
  32. 32 "alphabet": trained["alphabet"],
  33. 33 "steps": trained["steps"],
  34. 34 "vocab_size": trained["vocab_size"],
  35. 35 "n_merges": len(merges),
  36. 36 "checkpoints": [
  37. 37 {"after": n, "symbols": bpe.encode(" strawberry", merges[:n])}
  38. 38 for n in CHECKPOINTS
  39. 39 ],
  40. 40 "probes": [
  41. 41 {
  42. 42 "text": p,
  43. 43 "toy": bpe.encode(p, merges),
  44. 44 "real": [t["text"] for t in real.tokens(p)],
  45. 45 }
  46. 46 for p in PROBES
  47. 47 ],
  48. 48 }
  49. 49 
  50. 50 with open(OUT, "w", encoding="utf-8") as fh:
  51. 51 json.dump(data, fh, ensure_ascii=False, indent=1)
  52. 52 
  53. 53 print(f"wrote {os.path.relpath(OUT, os.path.dirname(OUT))}")
  54. 54 print(f" corpus: {data['corpus_chars']} chars, {data['corpus_words']} words, "
  55. 55 f"{data['corpus_unique']} unique")
  56. 56 print(f" {data['n_merges']} merges, alphabet {len(data['alphabet'])}, "
  57. 57 f"vocab {data['vocab_size']}")
  58. 58 print(" first merges: " +
  59. 59 " ".join(f"{s['token']!r}" for s in data["steps"][:5]))
  60. 60 for p in data["probes"]:
  61. 61 print(f" {p['text']!r:>16} toy {len(p['toy']):>2} real {len(p['real']):>2}")
  62. 62 
  63. 63 
  64. 64if __name__ == "__main__":
  65. 65 main()