| """Composed multiply+reduce CoT — the DEPLOYABLE surface. |
| |
| divcot validates the reduction in isolation by feeding N=x*y as input; a real |
| submission only gets (x, y, p) with x,y < p. This encoding chains the two |
| validated scaffolds in one generated trace: |
| |
| 1. schoolbook multiply as a running accumulator: after consuming each y-limb |
| y_j the model emits acc_j = acc_{j-1} + x*y_j*B^j (2W limbs, LSB-first so |
| the product+add carry chain is local). acc_{W-1} = N. |
| 2. copy-reverse: re-emit N MSB-first (attention reversal — staging the |
| dividend so the division block looks exactly like validated divcot). |
| 3. long division by p: quotient limbs MSB-first with running-remainder |
| install targets, then the answer (= final remainder) LSB-first. |
| |
| Surface: |
| "A" x(W,LSB) "B" y(W,LSB) "M" p(W,MSB) "=" |
| [ "a" acc_j(2W,LSB) ] * W |
| "N" N(2W,MSB) Q(2W) "R" ans(W,LSB) "\n" |
| |
| Tokens/example = 2W^2 + 9W + 7 (W=5: 102; base-100 tier 5 W=10: 297). |
| |
| Install targets (NTP-aligned): |
| acarry carry into each emitted acc limb (the multiply-accumulate state) |
| rem{j} MSB limb j of the running remainder at each quotient position |
| qhat leading-limbs quotient estimate at each quotient position |
| ans limb k of (x*y) mod p at each answer position |
| |
| Compliance: identical status to divcot/karatsuba — fixed trace SHAPE, every |
| value generated by trained weights; randomising weights collapses accuracy. |
| """ |
| from __future__ import annotations |
|
|
| from limbs import limb_char, limb_str, parse_limbs, to_limbs |
| from divcot2_encoding import long_division, qhat_estimate |
|
|
| AMARK, BMARK, DIVMARK, EQ, ACC, NMARK, REVMARK, NL = "A", "B", "M", "=", "a", "N", "R", "\n" |
| PLSB = "m" |
|
|
|
|
| def prompt_str(x: int, y: int, p: int, W: int, base: int, subpad: bool = False, |
| mulonly: bool = False) -> str: |
| """subpad adds p LSB-first to the prompt: the MSB field serves quotient |
| estimation, the LSB field aligns with the LSB-emitted qd*p / remainder |
| chains (per-step diagnostic: subtraction digits need REVERSED access into |
| the MSB-only field — the aligned acc stage learned, the division didn't).""" |
| if mulonly: |
| return AMARK + limb_str(x, W, base) + BMARK + limb_str(y, W, base) + EQ |
| s = (AMARK + limb_str(x, W, base) + BMARK + limb_str(y, W, base) |
| + DIVMARK + limb_str(p, W, base, msb_first=True)) |
| if subpad: |
| s += PLSB + limb_str(p, W, base) |
| return s + EQ |
|
|
|
|
| def gen_len(W: int, scratch: bool = False, cursor: bool = False, |
| subpad: bool = False, stepidx: bool = False, |
| skiptriv: bool = False, subnum: bool = False, |
| bemit: bool = False, srt: bool = False, |
| mulonly: bool = False) -> int: |
| """Tokens generated after '=': W acc blocks, N restage, quotient, answer. |
| scratch=True: each quotient limb is followed by the running remainder |
| (W limbs, LSB-first) — externalizes the division state, which the model |
| cannot hold internally past W~3 (composed tier-3 diagnostic: first |
| quotient digit right, then collapse). |
| cursor=True (requires scratch): each step block starts with the consumed |
| dividend limb, copied from the restaged N — an explicit progress anchor. |
| Diagnostic at scratch step-4k: mid-sequence remainder drift + the model |
| losing count of steps (9 blocks emitted instead of 2W); copying gives the |
| same position mechanism that put restage at 1.00. |
| subpad=True (requires cursor): each step also emits qd*p (W+1 limbs, LSB) |
| before the remainder — externalizes the multiply-subtract that the |
| per-step diagnostic showed failing (steps 0-4 = copies, perfect; real |
| division steps 5+ at 0.04-0.15 even teacher-forced). Tier-2's compact |
| success could memorize qd*p over 48 primes; tier 3's 400 primes need the |
| generic circuit, so generate it like the (perfectly learned) acc rows. |
| stepidx=True (requires cursor): each step opens with its 2-digit ASCII |
| index, and the restage pairs every dividend limb with its index — a |
| content-addressable name per step. Tier-4 diagnostic (W=10, 20 steps): |
| steps 0-9 at ~1.00, cliff to ~0.45 at 10+; rem fails worst (0.36 late) |
| despite all-local operands, so the wall is step MISBINDING under |
| relative-position (ALiBi) attention — the number of identical-looking |
| distractor blocks grows with step index. Indexed steps turn 'previous |
| step's remainder' and 'my dividend limb' into content lookups. |
| skiptriv=True (requires stepidx): x,y < p guarantees Q = N//p < p, so the |
| first W division steps ALWAYS emit qd=0 — provably trivial filler that |
| doubles the distractor count, dilutes gradient (250 of 828 tokens at |
| W=10), and pushes real steps ~250 tokens further from the prompt's p |
| field. Skip them: division = W real steps keeping their TRUE indices |
| (so restage lookups still match), starting from r_prev = N's top half. |
| Also re-emits p LSB-first ('m' + W limbs) right after the restage, so |
| qd*p rows read a local p copy instead of a ~700-token fetch (tier-4 |
| stepidx diagnostic: qdp 0.78-0.84 was the residual fetch weakness). |
| subnum=True (requires subpad): each step copies its numerator |
| num = r_prev*base + cursor limb (W+1 limbs, LSB) between cursor and qd, |
| so qd / qd*p / rem are all WITHIN-step local arithmetic. skiptriv |
| diagnostic: qdp hit 0.9949 once p became a local read while rem |
| plateaued at 0.43 fetching the previous step's rem block -- every pure |
| copy in this stack reaches ~1.00, every fetch+arithmetic fusion stalls; |
| decouple them (the subpad lesson, applied to the subtraction). |
| bemit=True (requires subnum): the remainder block becomes (borrow, limb) |
| pairs -- the borrow bit is EMITTED as an ASCII 0/1 token before each rem |
| limb, so b_k sits in context when rem[k] is predicted and b_{k+1} is a |
| local function of (num[k], qdp[k], b_k), all adjacent. v4 diagnostic: |
| rem errors were 58% unstructured noise, 0.54 even at b=0 -- the |
| install-only borrow supervision never built a circuit; emit the chain |
| (the scratch lesson, applied to the borrow state).""" |
| if mulonly: |
| return W * (1 + 2 * W) + 1 + 2 * W |
| restage = 1 + (3 if stepidx else 1) * 2 * W |
| pcopy = (1 + W) if skiptriv else 0 |
| if srt: |
| |
| per_step = 2 + 1 + (W + 2) + 1 + (W + 1) + 2 * (W + 2) |
| final = 1 + 2 * W |
| return (W * (1 + 2 * W) + restage + pcopy + W * per_step |
| + final + 1 + W) |
| per_step = ((2 if stepidx else 0) + (1 if cursor else 0) |
| + ((W + 1) if subnum else 0) + 1 |
| + ((W + 1) if subpad else 0) |
| + ((2 * W) if bemit else (W if scratch else 0))) |
| n_steps = W if skiptriv else 2 * W |
| return W * (1 + 2 * W) + restage + pcopy + n_steps * per_step + 1 + W |
|
|
|
|
| def _srt_qd(num: int, p: int, base: int) -> int: |
| """Signed SRT digit from |num|'s top-4 / p's top-3 limbs (exact integer |
| rounding). Guarantees |num - qd*p| < p given |num| < base*p (verified: |
| 0 violations / 500k tier-4 selections, worst |s|/p = 0.54).""" |
| if num == 0: |
| return 0 |
| sgn = 1 if num > 0 else -1 |
| an = abs(num) |
| nw = 1 |
| while base ** nw <= an: |
| nw += 1 |
| pw = 1 |
| while base ** pw <= p: |
| pw += 1 |
| n4 = an // base ** max(0, nw - 4) |
| p3 = p // base ** max(0, pw - 3) |
| sh = max(0, nw - 4) - max(0, pw - 3) |
| if sh >= 0: |
| a, b = n4 * base ** sh, p3 |
| else: |
| a, b = n4, p3 * base ** -sh |
| q = (2 * a + b) // (2 * b) |
| return sgn * min(base - 1, q) |
|
|
|
|
| def build_example(x: int, y: int, p: int, W: int, base: int, scratch: bool = False, |
| cursor: bool = False, subpad: bool = False, stepidx: bool = False, |
| skiptriv: bool = False, subnum: bool = False, |
| bemit: bool = False, srt: bool = False, mulonly: bool = False, |
| srt_dither: float = 0.0): |
| """Return (text, ann). x,y < p < base**W (mulonly: x,y < base**W).""" |
| assert not subnum or subpad, "subnum requires subpad" |
| assert not bemit or subnum, "bemit requires subnum" |
| assert not srt or (skiptriv and bemit), "srt requires skiptriv+bemit" |
| assert not srt or base == 10, "srt: signed-digit tokens are base-10 only" |
| assert not cursor or scratch, "cursor requires scratch" |
| assert not subpad or cursor, "subpad requires cursor" |
| assert not stepidx or cursor, "stepidx requires cursor" |
| assert not skiptriv or stepidx, "skiptriv requires stepidx" |
| assert not stepidx or 2 * W <= 100, "stepidx: 2-digit indices cap at 100 steps" |
| assert not skiptriv or (x < p and y < p), "skiptriv: Q < p needs x,y < p" |
| N = x * y |
| y_limbs = to_limbs(y, W, base) |
| x_limbs = to_limbs(x, W, base) |
| parts, ann_g = [], [] |
| pos = 0 |
|
|
| acc = 0 |
| for j, yj in enumerate(y_limbs): |
| prev = acc |
| acc = prev + x * yj * base ** j |
| parts.append(ACC) |
| pos += 1 |
| prev_l = to_limbs(prev, 2 * W, base) |
| c = 0 |
| for i in range(2 * W): |
| ann_g.append((pos + i, "acarry", c)) |
| xi = x_limbs[i - j] if 0 <= i - j < W else 0 |
| c = (prev_l[i] + xi * yj + c) // base |
| parts.append(limb_str(acc, 2 * W, base)) |
| pos += 2 * W |
| assert acc == N |
|
|
| if mulonly: |
| parts.append(NMARK) |
| pos += 1 |
| n_l = to_limbs(N, 2 * W, base) |
| for k in range(2 * W): |
| ann_g.append((pos + k, "ans", n_l[k])) |
| parts.append(limb_str(N, 2 * W, base) + NL) |
| pos += 2 * W + 1 |
| prompt = prompt_str(x, y, p, W, base, mulonly=True) |
| text = prompt + "".join(parts) |
| assert len(text) == len(prompt) + gen_len(W, mulonly=True) + 1 |
| ann = [dict() for _ in range(len(text))] |
| for gi, var, val in ann_g: |
| ann[len(prompt) + gi - 1][var] = int(val) |
| return text, ann |
|
|
| n_msb = to_limbs(N, 2 * W, base, msb_first=True) |
| if stepidx: |
| parts.append(NMARK + "".join(f"{i:02d}" + limb_char(n_msb[i]) |
| for i in range(2 * W))) |
| pos += 1 + 3 * 2 * W |
| else: |
| parts.append(NMARK + limb_str(N, 2 * W, base, msb_first=True)) |
| pos += 1 + 2 * W |
|
|
| if skiptriv: |
| parts.append(PLSB + limb_str(p, W, base)) |
| pos += 1 + W |
|
|
| if srt: |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| K = W + 2 |
| r = N // base ** W |
| for i in range(W, 2 * W): |
| d = n_msb[i] |
| num = r * base + d |
| parts.append(f"{i:02d}" + limb_char(d)) |
| pos += 3 |
| num_c = to_limbs((num + base ** K) % base ** K, K, base) |
| parts.append("".join(limb_char(v) for v in num_c)) |
| pos += K |
| qd = _srt_qd(num, p, base) |
| if srt_dither: |
| |
| |
| |
| |
| |
| |
| h = hash((x, y, p, i)) & 0xFFFF |
| if h < int(srt_dither * 0x10000): |
| alt = qd + (1 if h & 1 else -1) |
| if abs(alt) <= base - 1 and abs(num - alt * p) < p: |
| qd = alt |
| assert abs(num - qd * p) < p, (x, y, p, i, num, qd) |
| ann_g.append((pos, "qhat", qd + base - 1)) |
| parts.append(chr(0x90 + base - 1 + qd)) |
| pos += 1 |
| aq = abs(qd) * p |
| p_l = to_limbs(p, W, base) + [0] |
| c = 0 |
| for k in range(W + 1): |
| ann_g.append((pos + k, "spcarry", c)) |
| c = (abs(qd) * p_l[k] + c) // base |
| parts.append(limb_str(aq, W + 1, base)) |
| pos += W + 1 |
| s = num - qd * p |
| s_c = to_limbs((s + base ** K) % base ** K, K, base) |
| aq_l = to_limbs(aq, K, base) |
| ch = 0 |
| for k in range(K): |
| parts.append(str(ch) + limb_char(s_c[k])) |
| if qd >= 0: |
| ch = 1 if num_c[k] - aq_l[k] - ch < 0 else 0 |
| else: |
| ch = (num_c[k] + aq_l[k] + ch) // base |
| pos += 2 * K |
| r = s |
| neg = 1 if r < 0 else 0 |
| parts.append(str(neg)) |
| pos += 1 |
| r_c = to_limbs((r + base ** (W + 1)) % base ** (W + 1), W + 1, base) |
| pn_l = to_limbs(p * neg, W + 1, base) |
| answer = r + p * neg |
| a_l = to_limbs(answer, W, base) |
| ch = 0 |
| for k in range(W): |
| parts.append(str(ch) + limb_char(a_l[k])) |
| ch = (r_c[k] + pn_l[k] + ch) // base |
| pos += 2 * W |
| assert answer == N % p, (x, y, p) |
| q_limbs = [] |
|
|
| if not srt: |
| q_limbs, rems, answer = long_division(N, p, W, base) |
| r_prev = 0 |
| for i, (qd, r) in enumerate(zip(q_limbs, rems) if not srt else ()): |
| if skiptriv and i < W: |
| assert qd == 0 |
| r_prev = r |
| continue |
| if stepidx: |
| parts.append(f"{i:02d}") |
| pos += 2 |
| if cursor: |
| parts.append(limb_char(n_msb[i])) |
| pos += 1 |
| if subnum: |
| num = r_prev * base + n_msb[i] |
| parts.append(limb_str(num, W + 1, base)) |
| pos += W + 1 |
| |
| |
| |
| |
| |
| ann_g.append((pos, "qhat", qd)) |
| if not scratch: |
| rstr = to_limbs(r, W, base, msb_first=True) |
| for j in range(W): |
| ann_g.append((pos, f"rem{j}", rstr[j])) |
| parts.append(limb_char(qd)) |
| pos += 1 |
| if subpad: |
| qdp = qd * p |
| p_l = to_limbs(p, W, base) + [0] |
| c = 0 |
| for k in range(W + 1): |
| ann_g.append((pos + k, "spcarry", c)) |
| c = (qd * p_l[k] + c) // base |
| parts.append(limb_str(qdp, W + 1, base)) |
| pos += W + 1 |
| if scratch: |
| num_l = [n_msb[i]] + to_limbs(r_prev, W, base) |
| qdp_l = to_limbs(qd * p, W + 1, base) |
| if bemit: |
| r_l = to_limbs(r, W, base) |
| b = 0 |
| for k in range(W): |
| parts.append(str(b) + limb_char(r_l[k])) |
| b = 1 if num_l[k] - qdp_l[k] - b < 0 else 0 |
| pos += 2 * W |
| else: |
| if subpad: |
| b = 0 |
| for k in range(W): |
| ann_g.append((pos + k, "sborrow", b)) |
| b = 1 if num_l[k] - qdp_l[k] - b < 0 else 0 |
| parts.append(limb_str(r, W, base)) |
| pos += W |
| r_prev = r |
| parts.append(REVMARK) |
| pos += 1 |
| ans_limbs = to_limbs(answer, W, base) |
| for k in range(W): |
| ann_g.append((pos + k, "ans", ans_limbs[k])) |
| parts.append(limb_str(answer, W, base) + NL) |
| pos += W + 1 |
|
|
| prompt = prompt_str(x, y, p, W, base, subpad) |
| text = prompt + "".join(parts) |
| assert len(text) == len(prompt) + gen_len(W, scratch, cursor, subpad, stepidx, |
| skiptriv, subnum, bemit, srt) + 1, \ |
| (len(text), len(prompt), gen_len(W, scratch, cursor, subpad, stepidx, |
| skiptriv, subnum, bemit, srt)) |
| ann = [dict() for _ in range(len(text))] |
| base_i = len(prompt) |
| for gi, var, val in ann_g: |
| ann[base_i + gi - 1][var] = int(val) |
| return text, ann |
|
|
|
|
| def var_specs(W: int, base: int, scratch: bool = False, subpad: bool = False, |
| srt: bool = False, mulonly: bool = False): |
| if mulonly: |
| return [("ans", base), ("acarry", base)] |
| specs = [("ans", base), ("qhat", 2 * base - 1 if srt else base), |
| ("acarry", base)] |
| if subpad: |
| specs += [("spcarry", base), ("sborrow", 2)] |
| 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, |
| cursor: bool = False, subpad: bool = False, |
| stepidx: bool = False, skiptriv: bool = False, |
| subnum: bool = False, bemit: bool = False, |
| srt: bool = False, mulonly: bool = False) -> int: |
| if mulonly: |
| off = gen_len(W, mulonly=True) - 2 * W |
| return parse_limbs(gen_chars[off:off + 2 * W], base) |
| |
| off = gen_len(W, scratch, cursor, subpad, stepidx, skiptriv, subnum, bemit, |
| srt) - W |
| return parse_limbs(gen_chars[off:off + W], base) |
|
|
|
|
| if __name__ == "__main__": |
| import random |
| rng = random.Random(3) |
| for base in (10, 100): |
| for scratch, cursor, subpad, stepidx, skiptriv, subnum, bemit, srt in ( |
| (False, False, False, False, False, False, False, False), |
| (True, False, False, False, False, False, False, False), |
| (True, True, False, False, False, False, False, False), |
| (True, True, True, False, False, False, False, False), |
| (True, True, True, True, False, False, False, False), |
| (True, True, True, True, True, False, False, False), |
| (True, True, True, True, True, True, False, False), |
| (True, True, True, True, True, True, True, False), |
| (True, True, True, True, True, True, True, True)): |
| if srt and base != 10: |
| continue |
| for _ in range(4000): |
| W = rng.randint(1, 12) |
| p = rng.randrange(max(2, base ** (W - 1)), base ** W) |
| x, y = rng.randrange(p), rng.randrange(p) |
| text, ann = build_example(x, y, p, W, base, scratch, cursor, |
| subpad, stepidx, skiptriv, subnum, bemit, |
| srt) |
| plen = len(prompt_str(x, y, p, W, base, subpad)) |
| assert decode_answer(text[plen:], W, base, scratch, cursor, |
| subpad, stepidx, skiptriv, subnum, bemit, |
| srt) \ |
| == (x * y) % p, (x, y, p, W, base, scratch, cursor, subpad, |
| stepidx, skiptriv, subnum, bemit, srt) |
| assert len(ann) == len(text) |
| print(f"composed base={base} scratch={scratch} cursor={cursor} " |
| f"subpad={subpad} stepidx={stepidx} skiptriv={skiptriv} " |
| f"subnum={subnum} bemit={bemit} srt={srt}: 4000/4000 decode OK") |
| print("\ntokens/example (prompt+gen+NL):") |
| for tier, Wd in [(3, 5), (4, 10), (5, 20), (6, 39), (7, 78)]: |
| for base, W in ((10, Wd), (100, (Wd + 1) // 2)): |
| print(f" tier {tier} base {base:>3} W={W:>3} compact {3 * W + 4 + gen_len(W) + 1:>6}" |
| f" scratch {3 * W + 4 + gen_len(W, True) + 1:>6}") |
|
|