# ============================================================================

# DEEP CFENG + MLA  (char tiny-shakespeare)  -- multi-layer version of the champion

# The 1.4648 champion is ONE layer: emb -> carry-FiLM -> MLA -> CFENG -> head.

# nanoGPT char is 6 layers. This stacks the [MLA -> CFENG] block NLAYERS times,

# transformer-style (pre-norm residual), with the EMA carry-FiLM computed ONCE at input.

#   emb -> carry-FiLM  ->  [ h += MLA(LN h);  h += CFENG(LN h) ] x NLAYERS  ->  LN -> head

# Cross-token depth was the one lever still at 1 (per-token depth K=4 is tapped out).

# Defaults: NLAYERS=2 MLA=1 CFENG=1 OPT=sf. Just run it. (pip install schedulefree)

# NLAYERS=1 == the champion block, re-expressed as a residual sublayer (~1.4648 check).

# Param-fair option: NLAYERS=2 with M=224 holds params near the 1.82M champion budget.

# ============================================================================

import os, math, time, urllib.request

import torch, torch.nn as nn, torch.nn.functional as F



# =============================== CONFIG (edit me) ===============================

D        = int(os.environ.get('D', 128))        # token-embedding dim

M        = int(os.environ.get('M', 280))        # block hidden width (our best)

BLOCK    = int(os.environ.get('BLOCK', 256))    # context length

BATCH    = int(os.environ.get('BATCH', 64))     # sequences per step (our best)

FFMUL    = int(os.environ.get('FFMUL', 4))

DROPOUT  = float(os.environ.get('DROPOUT', 0.2))

SIGMATR  = float(os.environ.get('SIGMATR', 1.0))

NLAYERS  = int(os.environ.get('NLAYERS', 3))    # <<< NEW: number of stacked [MLA->CFENG] blocks (champion=1)

NEST     = int(os.environ.get('NEST', 2))       # CHAMPION = 2 (5x2=10 streams). NEST=3 is an untuned variant.

NESTSTACK= os.environ.get('NESTSTACK', '1') == '1'

LAMS     = [float(x) for x in os.environ.get('LAMS', '0.5,0.875,0.969,0.992,0.996').split(',')]

MLA_ON   = os.environ.get('MLA', '1') == '1'        # ON: MLA attention in each block

CFENG    = os.environ.get('CFENG', '1') == '1'      # continued-fraction division engine (per-token), ON

CFENG_K  = int(os.environ.get('CFENG_K', 5))        # depth: number of gated-rational-division levels (per-token)

MLA_DC   = int(os.environ.get('MLA_DC', 64))     # KV latent dim (the cached part)

MLA_NH   = int(os.environ.get('MLA_NH', 4))      # heads

MLA_RD   = int(os.environ.get('MLA_RD', 16))     # decoupled-RoPE dim

OPT      = os.environ.get('OPT', 'sf')           # CHAMPION = 'sf' (schedule-free AdamW). 'sfans' | 'adamw' also available

LR       = float(os.environ.get('LR', 2e-3))     # used only by OPT=sf / adamw

MAXIT    = int(os.environ.get('MAXIT', 4500))

EVAL_INT = int(os.environ.get('EVAL_INT', 250))  # main eval cadence (also drives the WD ramp)

EVAL_ITERS = int(os.environ.get('EVAL_ITERS', 50))  # val/train batches per eval

FINE_EVAL= int(os.environ.get('FINE_EVAL', 50))  # extra RNG-isolated val read-outs between main evals (0=off)

WD_RAMP  = float(os.environ.get('WD_RAMP', 1.7))

WD_CAP   = float(os.environ.get('WD_CAP', 0.6))  # our best

USE_AMP  = os.environ.get('AMP', '0') == '1'     # fp32 by default (matches CPU). AMP=1 for ~3x GPU speed.

# --- SF-ANS optimizer knobs (OPT=sfans) ---------------------------------------

SFANS_LR2  = float(os.environ.get('SFANS_LR2', 0.02))

SFANS_LR1  = float(os.environ.get('SFANS_LR1', 3e-4))

SFANS_VP   = float(os.environ.get('SFANS_VP', 0.75))

SFANS_BETA = float(os.environ.get('SFANS_BETA', 0.95))

SFANS_C    = float(os.environ.get('SFANS_C', 0.05))

SFANS_NS   = int(os.environ.get('SFANS_NS', 5))

SFANS_RMS  = os.environ.get('SFANS_RMS', '0') == '1'

# ===============================================================================



DEV = 'cuda' if torch.cuda.is_available() else 'cpu'

torch.manual_seed(1337)



PATH = 'tinyshakespeare.txt'

if not os.path.exists(PATH):

    urllib.request.urlretrieve("https://raw.githubusercontent.com/karpathy/char-rnn/master/data/tinyshakespeare/input.txt", PATH)

text = open(PATH).read(); chars = sorted(set(text)); V = len(chars)

stoi = {c: i for i, c in enumerate(chars)}

data = torch.tensor([stoi[c] for c in text], dtype=torch.long)

n = int(0.9 * len(data)); train_data, val_data = data[:n], data[n:]



def get_batch(split):

    d = train_data if split == 'train' else val_data

    ix = torch.randint(len(d) - BLOCK - 1, (BATCH,))

    return (torch.stack([d[i:i+BLOCK] for i in ix]).to(DEV),

            torch.stack([d[i+1:i+BLOCK+1] for i in ix]).to(DEV))



def decay_matrix(T):

    lam = torch.tensor(LAMS).view(len(LAMS), 1, 1)

    t = torch.arange(T).view(1, T, 1); s = torch.arange(T).view(1, 1, T)

    mask = (s < t).float(); expo = (t - 1 - s).clamp(min=0).float()

    return (1 - lam) * (lam ** expo) * mask



def build_carry(T):

    base = decay_matrix(T)

    if NESTSTACK:

        return torch.cat([torch.linalg.matrix_power(base, p) for p in range(1, NEST + 1)], 0), len(LAMS) * NEST

    if NEST > 1:

        return torch.linalg.matrix_power(base, NEST), len(LAMS)

    return base, len(LAMS)



def _rope(x):                                            # x: (B, nh, T, hd)

    T, hd = x.shape[-2], x.shape[-1]; half = hd // 2

    fr = 1.0 / (10000.0 ** (torch.arange(half, device=x.device).float() / half))

    ang = torch.outer(torch.arange(T, device=x.device).float(), fr)

    cos, sin = ang.cos()[None, None], ang.sin()[None, None]

    x1, x2 = x[..., :half], x[..., half:]

    return torch.cat([x1*cos - x2*sin, x1*sin + x2*cos], -1)



class MLA(nn.Module):                                    # GLM-5.2 / DeepSeek-V3 multi-latent attention (low-rank KV)

    def __init__(s, M, dc=MLA_DC, nh=MLA_NH, rd=MLA_RD):

        super().__init__()

        s.nh = nh; s.hd = M // nh; s.rd = rd

        s.WDKV = nn.Linear(M, dc, bias=False)

        s.WUK  = nn.Linear(dc, nh*s.hd, bias=False)

        s.WUV  = nn.Linear(dc, nh*s.hd, bias=False)

        s.WQ   = nn.Linear(M, nh*s.hd, bias=False)

        s.WQR  = nn.Linear(M, nh*rd, bias=False)

        s.WKR  = nn.Linear(M, rd, bias=False)

        s.WO   = nn.Linear(nh*s.hd, M, bias=False); nn.init.zeros_(s.WO.weight)   # zero-init -> starts as no-op

    def forward(s, h):

        B, T, M = h.shape; c = s.WDKV(h)

        k = s.WUK(c).view(B, T, s.nh, s.hd).transpose(1, 2)

        v = s.WUV(c).view(B, T, s.nh, s.hd).transpose(1, 2)

        q = s.WQ(h).view(B, T, s.nh, s.hd).transpose(1, 2)

        qr = _rope(s.WQR(h).view(B, T, s.nh, s.rd).transpose(1, 2))

        kr = _rope(s.WKR(h).view(B, T, 1, s.rd).transpose(1, 2)).expand(B, s.nh, T, s.rd)

        qf = torch.cat([q, qr], -1); kf = torch.cat([k, kr], -1)

        o = F.scaled_dot_product_attention(qf, kf, v, is_causal=True, dropout_p=DROPOUT if s.training else 0.0)

        return s.WO(o.transpose(1, 2).reshape(B, T, s.nh*s.hd))



class Block(nn.Module):

    """One stacked layer: residual MLA (token mixing) + residual CFENG (per-token division engine).

    Both sublayers are pre-norm and end in a zero-init projection (WO / cf_out) so the block

    starts as identity and 'unlocks' after step 1 -- stable to stack to depth."""

    def __init__(s):

        super().__init__()

        if MLA_ON:

            s.ln_mla = nn.LayerNorm(M); s.mla = MLA(M)

        if CFENG:

            s.ln_cf = nn.LayerNorm(M)

            s.cf_in = nn.Linear(M, M)

            s.cf_N  = nn.ModuleList([nn.Linear(M, M) for _ in range(CFENG_K)])   # numerators

            s.cf_D  = nn.ModuleList([nn.Linear(M, M) for _ in range(CFENG_K)])   # denominators

            s.cf_T  = nn.ModuleList([nn.Linear(M, M) for _ in range(CFENG_K)])   # per-level contraction

            s.cf_ln = nn.ModuleList([nn.LayerNorm(M) for _ in range(CFENG_K)])

            s.cf_out = nn.Linear(M, M)

            for T in s.cf_T: nn.init.normal_(T.weight, 0.0, 0.02 / (CFENG_K ** 0.5)); nn.init.zeros_(T.bias)

            nn.init.zeros_(s.cf_out.weight); nn.init.zeros_(s.cf_out.bias)       # residual sublayer starts as no-op

        else:

            s.ln_ffn = nn.LayerNorm(M)

            s.ffn = nn.Sequential(nn.Linear(M, FFMUL*M), nn.GELU(), nn.Dropout(DROPOUT), nn.Linear(FFMUL*M, M))

            nn.init.zeros_(s.ffn[-1].weight); nn.init.zeros_(s.ffn[-1].bias)

    def cfeng(s, x):                                          # depth-K continued fraction of gated rational divisions

        h = s.cf_in(x)

        for k in range(CFENG_K):

            hn = s.cf_ln[k](h)

            r = s.cf_N[k](hn) / (0.5 + F.softplus(s.cf_D[k](hn)))   # scale-inv division, denom >= 0.5

            h = h + s.cf_T[k](torch.sigmoid(r) * r)                 # gate IS the ratio: sigma(r)*r, contracted, residual

        return s.cf_out(h)

    def forward(s, h):

        if MLA_ON: h = h + s.mla(s.ln_mla(h))                       # token mixing: sharp content retrieval

        if CFENG:  h = h + s.cfeng(s.ln_cf(h))                      # per-token division engine (residual sublayer)

        else:      h = h + s.ffn(s.ln_ffn(h))

        return h



class ChampionLM(nn.Module):

    def __init__(s):

        super().__init__()

        Dm, s.LC = build_carry(BLOCK)

        s.register_buffer('Dm', Dm)

        s.emb = nn.Embedding(V, D); s.Vin = nn.Linear(D, M)

        s.Fg = nn.Linear(8, M); s.Fd = nn.Linear(8, M)

        s.Cbot = nn.Linear(s.LC * D, M); s.Cg = nn.Linear(M, M); s.Cd = nn.Linear(M, M)

        for l in (s.Fg, s.Fd, s.Cg, s.Cd): nn.init.zeros_(l.weight); nn.init.zeros_(l.bias)

        s.drop = nn.Dropout(DROPOUT)

        s.blocks = nn.ModuleList([Block() for _ in range(NLAYERS)])   # <<< stacked depth

        s.ln_f = nn.LayerNorm(M)                                      # final pre-head norm

        s.head = nn.Linear(M, V)

    def carry(s, e):

        B, T, Dd = e.shape

        return torch.einsum('lts,bsd->btld', s.Dm.to(e.dtype), e).reshape(B, T, s.LC * Dd)

    def forward(s, idx, sigma=0.0):

        e = s.emb(idx)

        if sigma > 0: e = e + sigma * torch.randn_like(e)

        c = s.carry(e); h = s.Vin(e)

        a = float(sigma) * torch.exp(torch.linspace(0, 3, 4, device=idx.device))

        es = torch.cat([torch.sin(a), torch.cos(a)])

        h = h * (1 + s.Fg(es)) + s.Fd(es)                            # sigma-FiLM

        cb = F.gelu(s.Cbot(c)); h = h * (1 + s.Cg(cb)) + s.Cd(cb)    # carry-FiLM (once, input conditioning)

        for blk in s.blocks: h = blk(h)                              # stacked [MLA -> CFENG] depth

        return s.head(s.ln_f(h))



# ======================== SF-ANS: custom optimizer =============================

def newton_schulz(G, steps=5):

    a, b, c = 3.4445, -4.7750, 2.0315

    X = G.float(); X = X / (X.norm() + 1e-7)

    transposed = X.shape[-2] > X.shape[-1]

    if transposed: X = X.T

    for _ in range(steps):

        A = X @ X.T

        X = a * X + b * (A @ X) + c * (A @ A @ X)

    if transposed: X = X.T

    return X.to(G.dtype)



class SFOptimizer:

    """SF-ANS: Schedule-Free Adam-Newton-Schulz. 2D: Adam-precondition (m_hat/v_hat^vp) -> Newton-Schulz.

    1D: plain Adam. Schedule-free z/x average; gradient at extrapolation y=(1-b)x+bz. wd from external ramp."""

    def __init__(s, params, lr_2d=0.02, lr_1d=3e-4, betas=(0.9, 0.999), v_power=0.75,

                 eps=1e-8, weight_decay=0.0, sf_beta=0.95, sf_c=0.05, use_ns=True, ns_steps=5, use_rms=False):

        s.ps = [p for p in params if p.requires_grad]

        s.lr2, s.lr1 = lr_2d, lr_1d; s.b1, s.b2 = betas

        s.eps = eps; s.ns = ns_steps

        s.vp = v_power; s.sf_beta = sf_beta; s.sf_c = sf_c; s.use_ns = use_ns; s.rms = use_rms

        s.z = [p.data.clone() for p in s.ps]

        s.x = [p.data.clone() for p in s.ps]

        s.st = [{} for _ in s.ps]; s.t = 0

        s.param_groups = [{'weight_decay': weight_decay}]

    @torch.no_grad()

    def extrapolate(s):

        b = s.sf_beta

        for i, p in enumerate(s.ps): p.data.copy_((1 - b) * s.x[i] + b * s.z[i])

    @torch.no_grad()

    def eval_point(s):

        for i, p in enumerate(s.ps): p.data.copy_(s.x[i])

    def train_point(s): s.extrapolate()

    def zero_grad(s, set_to_none=True):

        for p in s.ps:

            if p.grad is not None: p.grad = None if set_to_none else p.grad.zero_()

    @torch.no_grad()

    def step(s):

        s.t += 1; c = s.sf_c

        wd = s.param_groups[0]['weight_decay']

        for i, p in enumerate(s.ps):

            if p.grad is None: continue

            g = p.grad; st = s.st[i]

            if not st:

                st['step'] = 0; st['m'] = torch.zeros_like(p); st['v'] = torch.zeros_like(p)

            st['step'] += 1; t = st['step']; m, v = st['m'], st['v']

            is_2d = g.ndim >= 2

            lr = s.lr2 if is_2d else s.lr1

            m.mul_(s.b1).add_(g, alpha=1 - s.b1)

            v.mul_(s.b2).addcmul_(g, g, value=1 - s.b2)

            mh = m / (1 - s.b1 ** t); vh = v / (1 - s.b2 ** t)

            if is_2d:

                g_adam = (mh / (vh ** s.vp + s.eps)).reshape(mh.shape[0], -1)

                if s.use_ns:

                    go = newton_schulz(g_adam, s.ns)

                    if s.rms: go = go * max(1.0, go.shape[-2] / go.shape[-1]) ** 0.5

                    g_update = go.reshape(p.shape)

                else:

                    g_update = g_adam.reshape(p.shape)

            else:

                g_update = mh / (vh.sqrt() + s.eps)

            if wd: s.z[i].mul_(1 - lr * wd)

            s.z[i].add_(g_update, alpha=-lr)

            s.x[i].mul_(1 - c).add_(s.z[i], alpha=c)

            b = s.sf_beta; p.data.copy_((1 - b) * s.x[i] + b * s.z[i])

# ===============================================================================



def make_opt(model):

    if OPT == 'sfans':

        opt = SFOptimizer(model.parameters(), lr_2d=SFANS_LR2, lr_1d=SFANS_LR1, v_power=SFANS_VP,

                          sf_beta=SFANS_BETA, sf_c=SFANS_C, ns_steps=SFANS_NS, use_rms=SFANS_RMS,

                          weight_decay=0.0)

        print(f"  SF-ANS: {len(opt.ps)} tensors | lr2d={SFANS_LR2} lr1d={SFANS_LR1} vp={SFANS_VP} "

              f"beta={SFANS_BETA} c={SFANS_C} ns={SFANS_NS} rms={SFANS_RMS} (wd=external ramp)", flush=True)

        return opt, 'sfans'

    if OPT == 'sf':

        try:

            from schedulefree import AdamWScheduleFree

            return AdamWScheduleFree(model.parameters(), lr=LR, warmup_steps=100, weight_decay=0.0), 'sf'

        except ImportError:

            print("  (schedulefree missing -> AdamW fallback; pip install schedulefree)", flush=True)

    return torch.optim.AdamW(model.parameters(), lr=LR, betas=(0.9, 0.99), weight_decay=0.1), 'plain'



@torch.no_grad()

def evaluate(model, opt, kind, iters=EVAL_ITERS):

    model.eval()

    if kind == 'sf': opt.eval()

    elif kind == 'sfans': opt.eval_point()

    losses, hit, tot = {}, 0, 0

    for split in ('train', 'val'):

        acc = 0.0

        for _ in range(iters):

            x, y = get_batch(split)

            with torch.autocast(device_type=DEV, dtype=torch.bfloat16, enabled=(USE_AMP and DEV == 'cuda')):

                logits = model(x)

                acc += F.cross_entropy(logits.reshape(-1, V), y.reshape(-1)).item()

            if split == 'val':

                hit += (logits.argmax(-1) == y).sum().item(); tot += y.numel()

        losses[split] = acc / iters

    model.train()

    if kind == 'sf': opt.train()

    elif kind == 'sfans': opt.train_point()

    return losses, 100.0 * hit / tot



@torch.no_grad()

def fine_val(model, opt, kind, iters=50):

    rng = torch.get_rng_state(); torch.manual_seed(20240601)

    model.eval()

    if kind == 'sf': opt.eval()

    elif kind == 'sfans': opt.eval_point()

    acc, hit, tot = 0.0, 0, 0

    for _ in range(iters):

        x, y = get_batch('val')

        with torch.autocast(device_type=DEV, dtype=torch.bfloat16, enabled=(USE_AMP and DEV == 'cuda')):

            logits = model(x); acc += F.cross_entropy(logits.reshape(-1, V), y.reshape(-1)).item()

        hit += (logits.argmax(-1) == y).sum().item(); tot += y.numel()

    model.train()

    if kind == 'sf': opt.train()

    elif kind == 'sfans': opt.train_point()

    torch.set_rng_state(rng)

    return acc / iters, 100.0 * hit / tot



def main():

    model = ChampionLM().to(DEV)

    npar = sum(p.numel() for p in model.parameters())

    opt, kind = make_opt(model)

    print(f"DEEP champion | D={D} M={M} ctx={BLOCK} batch={BATCH} EMAs={len(LAMS)}->{model.LC}streams "

          f"NLAYERS={NLAYERS} NEST={NEST} stack={NESTSTACK} MLA={MLA_ON} CFENG={CFENG}(K={CFENG_K}) "

          f"OPT={OPT} | {npar/1e6:.2f}M params on {DEV}", flush=True)

    print(f"[targets: CFENG+MLA 1-layer champion 1.4648 ; nanoGPT 6-layer 10.8M = 1.4697]", flush=True)

    wd = 0.1; stale = 0; best = 9.9; fine_best = 9.9; t0 = time.time(); model.train()

    if kind == 'sf': opt.train()

    elif kind == 'sfans': opt.train_point()

    for it in range(MAXIT + 1):

        if it % EVAL_INT == 0:

            lo, acc = evaluate(model, opt, kind)

            if wd < WD_CAP:

                if lo['val'] < best - 0.02: stale = 0

                else:

                    stale += 1

                    if stale >= 1:

                        wd = min(WD_CAP, wd * WD_RAMP)

                        for g in opt.param_groups: g['weight_decay'] = wd

                        stale = 0

            best = min(best, lo['val'])

            print(f"  it {it:5d} | train {lo['train']:.4f} | val {lo['val']:.4f} | best {best:.4f} | acc {acc:.1f}% | wd {wd:.3f} | {time.time()-t0:.0f}s", flush=True)

        elif FINE_EVAL and it % FINE_EVAL == 0:

            fv, fa = fine_val(model, opt, kind)

            star = ""

            if fv < fine_best: fine_best = fv; star = " *fine-best*"

            print(f"  it {it:5d} |   .fine val {fv:.4f} | fine_best {fine_best:.4f} | acc {fa:.1f}% | wd {wd:.3f}{star}", flush=True)

        if it == MAXIT: break

        if kind == 'sfans': opt.extrapolate()

        x, y = get_batch('train')

        with torch.autocast(device_type=DEV, dtype=torch.bfloat16, enabled=(USE_AMP and DEV == 'cuda')):

            loss = F.cross_entropy(model(x, sigma=SIGMATR).reshape(-1, V), y.reshape(-1))

        opt.zero_grad(set_to_none=True); loss.backward()

        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0); opt.step()

    print(f"\n===== FINAL deep champion NLAYERS={NLAYERS} OPT={OPT}: best val {best:.4f}  (1-layer 1.4648, nanoGPT 1.4697) =====", flush=True)



main()
