| """Generalized long-division reduction CoT — base-B limbs, estimated quotient, |
| optional scratchpad. Scales divcot_encoding.py (validated at tier 2, base 10) |
| to higher tiers. |
| |
| Three levers over the validated v1: |
| |
| 1. **Base-B limbs** (B=100 halves W): each limb is one byte (limbs.py), so |
| tier 5's W=20 decimal becomes 10 limbs, tier 7's W=78 becomes 39. |
| 2. **Estimated-quotient install** (`qhat`): the classic scaling trick. The |
| one-shot quotient digit qd = floor((r*B+d)/p) needs a W-limb comparison — |
| fine at W=3, strained at W>=5. Humans estimate qd from the LEADING limbs |
| of numerator and divisor (a constant-size table independent of W), then |
| correct by ±1. We install that estimate at an early block so the circuit |
| the model learns is estimate-then-correct, not a W-limb divide. |
| 3. **Scratchpad mode** (`scratch=True`): for large W the model can't hold the |
| running remainder internally — EMIT it (LSB-first, so the subtraction |
| borrow chain is local). Cost: O(W^2) tokens vs O(W) compact; use compact |
| while it trains, scratch when it stops. |
| |
| Surface (compact): "N" N(2W,MSB) "M" p(W,MSB) "=" Q(2W) "R" ans(W,LSB) "\n" |
| Surface (scratch): ... "=" [qd r(W,LSB)]*2W "R" ans(W,LSB) "\n" |
| |
| Install targets (NTP-aligned: label at position t supervises the residual |
| predicting byte t+1): |
| rem{j} compact only — MSB limb j of the running remainder at each |
| quotient-emission position (the division carry/state) |
| qhat leading-limbs quotient estimate at each quotient position |
| ans limb k of N mod p at each answer position |
| """ |
| from __future__ import annotations |
|
|
| from limbs import limb_char, limb_str, parse_limbs, to_limbs |
|
|
| NMARK, DIVMARK, EQ, REVMARK, NL = "N", "M", "=", "R", "\n" |
|
|
|
|
| def long_division(N: int, p: int, W: int, base: int): |
| """(quotient_limbs[2W] MSB-first, remainders[2W], answer).""" |
| q_limbs, rems, r = [], [], 0 |
| for d in to_limbs(N, 2 * W, base, msb_first=True): |
| r = r * base + d |
| qd = r // p |
| r -= qd * p |
| q_limbs.append(qd) |
| rems.append(r) |
| return q_limbs, rems, r |
|
|
|
|
| def qhat_estimate(r_prev: int, d: int, p: int, base: int) -> int: |
| """Leading-limbs estimate of floor((r*B+d)/p): numerator's top two limbs |
| over divisor's top limb — constant-size lookup regardless of W.""" |
| num = r_prev * base + d |
| if num < p: |
| return 0 |
| nw = 1 |
| while base ** nw <= num: |
| nw += 1 |
| pw = 1 |
| while base ** pw <= p: |
| pw += 1 |
| n_top = num // base ** max(0, nw - 2) |
| p_top = p // base ** (pw - 1) |
| shift = (nw - 2) - (pw - 1) |
| est = (n_top // p_top) * base ** shift if shift >= 0 else n_top // (p_top * base ** -shift) |
| return min(base - 1, max(0, est)) |
|
|
|
|
| def prompt_str(N: int, p: int, W: int, base: int) -> str: |
| return NMARK + limb_str(N, 2 * W, base, msb_first=True) \ |
| + DIVMARK + limb_str(p, W, base, msb_first=True) + EQ |
|
|
|
|
| def gen_len(W: int, scratch: bool = False) -> int: |
| steps = 2 * W * (1 + W) if scratch else 2 * W |
| return steps + 1 + W |
|
|
|
|
| def build_example(N: int, p: int, W: int, base: int, scratch: bool = False): |
| q_limbs, rems, answer = long_division(N, p, W, base) |
| n_msb = to_limbs(N, 2 * W, base, msb_first=True) |
| prompt = prompt_str(N, p, W, base) |
| parts, ann_q = [], [] |
| pos = len(prompt) |
| r_prev = 0 |
| for i, (qd, r) in enumerate(zip(q_limbs, rems)): |
| ann_q.append((pos, "qhat", qhat_estimate(r_prev, n_msb[i], p, base))) |
| if not scratch: |
| rstr = to_limbs(r, W, base, msb_first=True) |
| for j in range(W): |
| ann_q.append((pos, f"rem{j}", rstr[j])) |
| parts.append(limb_char(qd)) |
| pos += 1 |
| if scratch: |
| parts.append(limb_str(r, W, base)) |
| pos += W |
| r_prev = r |
| ans_limbs = to_limbs(answer, W, base) |
| parts.append(REVMARK + limb_str(answer, W, base) + NL) |
| text = prompt + "".join(parts) |
| ann = [dict() for _ in range(len(text))] |
| for idx, var, val in ann_q: |
| ann[idx - 1][var] = int(val) |
| ans_base = pos + 1 |
| for k in range(W): |
| ann[ans_base + k - 1]["ans"] = ans_limbs[k] |
| return text, ann |
|
|
|
|
| def var_specs(W: int, base: int, scratch: bool = False): |
| """(name, n_classes) install vars; assign blocks in the trainer.""" |
| specs = [("ans", base), ("qhat", base)] |
| if not scratch: |
| specs += [(f"rem{j}", base) for j in range(W)] |
| return specs |
|
|
|
|
| def decode_answer(gen_chars: str, W: int, base: int, scratch: bool = False) -> int: |
| off = (2 * W * (1 + W) if scratch else 2 * W) + 1 |
| return parse_limbs(gen_chars[off:off + W], base) |
|
|
|
|
| if __name__ == "__main__": |
| import random |
| rng = random.Random(1) |
| for base in (10, 100): |
| for scratch in (False, True): |
| for _ in range(4000): |
| W = rng.randint(1, 8) |
| p = rng.randrange(max(2, base ** (W - 1)), base ** W) |
| N = rng.randrange(p * p) if p > 1 else 0 |
| text, ann = build_example(N, p, W, base, scratch) |
| plen = len(prompt_str(N, p, W, base)) |
| assert len(text) == plen + gen_len(W, scratch) + 1, (len(text), plen, gen_len(W, scratch)) |
| assert decode_answer(text[plen:], W, base, scratch) == N % p |
| assert len(ann) == len(text) |
| print(f"divcot2 base={base} scratch={scratch}: 4000/4000 decode OK") |
| |
| print("\ntokens/example (prompt+gen):") |
| for tier, Wd in [(3, 5), (4, 10), (5, 20), (6, 39), (7, 78), (8, 155)]: |
| for base, W in ((10, Wd), (100, (Wd + 1) // 2)): |
| plen = 3 * W + 3 |
| print(f" tier {tier} base {base:>3} W={W:>3} compact {plen + gen_len(W)}" |
| f" scratch {plen + gen_len(W, True)}") |
|
|