"""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 # single limb 0..base-1 (r < base*p) r -= qd * p q_limbs.append(qd) rems.append(r) return q_limbs, rems, r # r == N % p 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) # top 2 limbs of numerator p_top = p // base ** (pw - 1) # top 1 limb of divisor 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 = [], [] # ann_q: (char_index, var, val) 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: # compact: remainder is internal state 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: # emit remainder LSB-first (local borrow) parts.append(limb_str(r, W, base)) pos += W r_prev = r ans_limbs = to_limbs(answer, W, base) # LSB-first 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) # NTP alignment ans_base = pos + 1 # after 'R' 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 # skip quotient(+rems) and 'R' 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)) # +1: NL 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") # tier-relevant lengths 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)}")