sweedworks

← all sources

precompute_predict.py

Computes the figures on /predict/

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

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