sweedworks

← all sources

chatcost.py

The chat billing rule, checked against the encoder

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

  1. 1"""What a chat request actually costs, in tokens.
  2. 2 
  3. 3The API bills for the whole serialised conversation, not the words you typed.
  4. 4Every message is wrapped in ChatML scaffolding — <|im_start|>, the role name,
  5. 5<|im_sep|>, <|im_end|> — and the request ends with a partial header inviting the
  6. 6assistant to speak. That works out to an exact rule, which verify_rule() checks
  7. 7against the tokenizer's own chat encoder rather than trusting it:
  8. 8 
  9. 9 billed = sum(content tokens) + 4 per message + 3
  10. 10 
  11. 11The 4 is <|im_start|>, the role, <|im_sep|> and <|im_end|>. The 3 is the trailing
  12. 12<|im_start|>assistant<|im_sep|> that primes the reply.
  13. 13 
  14. 14This is the ChatML layout used by the GPT-4/4o family. Other providers wrap
  15. 15messages differently; the fact that there IS a wrapper is universal, the exact
  16. 16count is not.
  17. 17"""
  18. 18 
  19. 19import json
  20. 20import random
  21. 21 
  22. 22from cjsload import Reference
  23. 23 
  24. 24PER_MESSAGE = 4
  25. 25PER_REQUEST = 3
  26. 26 
  27. 27MODEL = "gpt-4o"
  28. 28 
  29. 29 
  30. 30class Chat:
  31. 31 def __init__(self, encoding="o200k_base"):
  32. 32 self.ref = Reference(encoding)
  33. 33 
  34. 34 def count(self, text):
  35. 35 return len(self.ref.ids(text))
  36. 36 
  37. 37 def encode_chat(self, messages):
  38. 38 self.ref.ctx.set("_m", json.dumps(messages))
  39. 39 return json.loads(self.ref.ctx.eval(
  40. 40 f'JSON.stringify(Array.from(M.encodeChat(JSON.parse(_m), "{MODEL}")))'))
  41. 41 
  42. 42 def billed(self, messages):
  43. 43 """Apply the rule rather than the encoder — same answer, and it is the
  44. 44 rule the browser calculator uses."""
  45. 45 content = sum(self.count(m["content"]) for m in messages)
  46. 46 return content + PER_MESSAGE * len(messages) + PER_REQUEST
  47. 47 
  48. 48 def conversation_cost(self, system, turns, user_tokens, reply_tokens):
  49. 49 """Cumulative tokens sent across a whole conversation.
  50. 50 
  51. 51 Turn n re-sends everything before it, so the input side grows linearly
  52. 52 per turn and the running total grows with the square of the turn count.
  53. 53 """
  54. 54 sys_tokens = self.count(system)
  55. 55 rows = []
  56. 56 cumulative = 0
  57. 57 history = sys_tokens + PER_MESSAGE # the system message
  58. 58 messages = 1
  59. 59 for turn in range(1, turns + 1):
  60. 60 history += user_tokens + PER_MESSAGE # this turn's user message
  61. 61 messages += 1
  62. 62 sent = history + PER_REQUEST
  63. 63 cumulative += sent
  64. 64 rows.append({
  65. 65 "turn": turn,
  66. 66 "sent": sent,
  67. 67 "cumulative": cumulative,
  68. 68 "system_share": sys_tokens,
  69. 69 "resent": sent - (user_tokens + PER_MESSAGE + PER_REQUEST),
  70. 70 })
  71. 71 history += reply_tokens + PER_MESSAGE # the assistant's reply
  72. 72 messages += 1
  73. 73 return rows
  74. 74 
  75. 75 
  76. 76def verify_rule(chat, trials=200, seed=5):
  77. 77 """The arithmetic must match the tokenizer's own chat encoder exactly."""
  78. 78 rng = random.Random(seed)
  79. 79 words = ["hello", "please", "summarise", "the", "document", "café", "🍓",
  80. 80 "for me", "in Japanese", "こんにちは", "x = 1", "", " ", "\n\n",
  81. 81 "a much longer sentence that goes on for a while without stopping"]
  82. 82 roles = ["system", "user", "assistant"]
  83. 83 bad = []
  84. 84 for _ in range(trials):
  85. 85 n = rng.randint(1, 6)
  86. 86 msgs = [{"role": rng.choice(roles),
  87. 87 "content": " ".join(rng.choice(words)
  88. 88 for _ in range(rng.randint(1, 8)))}
  89. 89 for _ in range(n)]
  90. 90 try:
  91. 91 actual = len(chat.encode_chat(msgs))
  92. 92 except Exception as e:
  93. 93 bad.append((msgs, f"encoder failed: {e}"))
  94. 94 continue
  95. 95 predicted = chat.billed(msgs)
  96. 96 if actual != predicted:
  97. 97 bad.append((msgs, f"rule {predicted} != encoder {actual}"))
  98. 98 return bad
  99. 99 
  100. 100 
  101. 101if __name__ == "__main__":
  102. 102 chat = Chat()
  103. 103 bad = verify_rule(chat)
  104. 104 if bad:
  105. 105 print(f"FAIL chat rule disagrees with the encoder on {len(bad)} of 200")
  106. 106 for msgs, why in bad[:3]:
  107. 107 print(f" {why}: {msgs}")
  108. 108 raise SystemExit(1)
  109. 109 print("PASS chat cost rule matches encodeChat on 200 random conversations")
  110. 110 
  111. 111 demo = [{"role": "system", "content": "You are a helpful assistant."},
  112. 112 {"role": "user", "content": "What is the capital of France?"}]
  113. 113 content = sum(chat.count(m["content"]) for m in demo)
  114. 114 print(f" example: {content} tokens of content -> "
  115. 115 f"{chat.billed(demo)} billed")