"""Lower the exact fixed-batch ReLU learner to affine-loop v4 primitives.

No learned values enter the program: initial weights are seed-only raw literals.
Gradient entries are streamed through one accumulator once hidden deltas exist,
so all updates use the same pre-update dependency values without gradient arrays.

optimizer='sgd' reproduces the accepted scorer byte for byte. optimizer='adam'
replaces each two-instruction SGD update with the ordered FP32 Adam rule used by
the research code, using zero-initialized m and v arrays per parameter, a
per-step bias-correction table, the div primitive, and a fixed Newton-Raphson
square root built only from add/mul/div.
"""
from __future__ import annotations
import math
import numpy as np

from affine import ref as R, ins as I, loop as L, make_program

def initial_parameters(features, width, seed):
    rng = np.random.Generator(np.random.PCG64(seed))
    return [rng.uniform(-1 / math.sqrt(features), 1 / math.sqrt(features), (features, width)).astype(np.float32),
            np.zeros(width, dtype=np.float32),
            rng.uniform(-1 / math.sqrt(width), 1 / math.sqrt(width), (width, 10)).astype(np.float32),
            np.zeros(10, dtype=np.float32)]


def raw(value):
    return int(np.float32(value).view(np.uint32))


NR_ITERATIONS = 12
NR_SEED_OFFSET = 1e-6


def newton_sqrt(x, y, q, half, tiny, iterations=NR_ITERATIONS):
    """Newton-Raphson sqrt leaves y ~ sqrt(x) using only add/mul/div.

    Deterministic seed y0 = x + 1e-6 (the offset caps the low end of the
    validated domain, x in [1e-12, 1e3]) and a fixed iteration count
    y <- 0.5*(y + x/y). The first iteration is evaluated q = x/y,
    q = y + q, y = 0.5*q in that order.
    """
    leaves = [I('add', y, x, tiny)]
    for _ in range(iterations):
        leaves += [I('div', q, x, y), I('add', q, y, q), I('mul', y, half, q)]
    return leaves


def adam_update(w, g, m, v, c1, c2, step, z1, z2, z3, nr_iterations):
    """Ordered FP32 Adam leaves for one parameter element.

    m = 0.9*m + 0.1*g; v = 0.999*v + 0.001*g*g; then
    w = w - step*(m/c1)/(approx_sqrt(v/c2)+1e-8), matching Python's
    left-to-right evaluation of step * mhat / (sqrt(vhat) + eps) with the
    sqrt replaced by the fixed Newton-Raphson sequence above.
    """
    beta1, omb1, beta2, omb2, eps = [R('k', i) for i in range(15, 20)]
    tiny, half = R('k', 20), R('k', 3)
    return [
        I('mul',z1,beta1,m), I('mul',z2,omb1,g), I('add',m,z1,z2),
        I('mul',z1,beta2,v), I('mul',z2,omb2,g), I('mul',z2,z2,g), I('add',v,z1,z2),
        I('div',z2,v,c2), *newton_sqrt(z2,z3,z1,half,tiny,nr_iterations),
        I('div',z1,m,c1), I('add',z3,z3,eps),
        I('mul',z1,step,z1), I('div',z1,z1,z3), I('sub',w,w,z1),
    ]


def build_mlp(width, epochs, learning_rate, seed=101, n_train=6000, n_test=6000, batch=30, features=81, stream_queries=True, optimizer='sgd', nr_iterations=NR_ITERATIONS):
    if optimizer not in ('sgd', 'adam'):
        raise ValueError('Optimizer must be sgd or adam')
    if nr_iterations < 1:
        raise ValueError('Newton-Raphson iterations must be positive')
    if min(width, epochs, n_train, n_test, batch, features) <= 0:
        raise ValueError('All dimensions and epochs must be positive')
    if n_train % batch:
        raise ValueError('Only complete fixed-size minibatches are supported')
    if not np.isfinite(learning_rate) or learning_rate <= 0:
        raise ValueError('Learning rate must be finite and positive')
    H, D, C, B, N, Q = width, features, 10, batch, n_train, n_test
    NB, T = N // B, epochs * (N // B)
    regions = [('s',5),('k',21 if optimizer=='adam' else 15),('w1',D*H),('b1',H),('w2',H*C),('b2',C),
               ('h',B*H),('d1',B*H),('d2',B*C),('x',N*D),('labels',N),
               ('target',N*C),('q',Q*D)]
    if stream_queries:
        # Reuse one query beside the hot scalars. Training data and targets stay
        # resident; query tape words are consumed only after training finishes.
        regions = regions[:2] + [('q', D)] + regions[2:-1]
    if optimizer == 'adam':
        # Zero-initialized moments per parameter array, the per-global-step
        # bias-correction divisors, and three FP32 scalar temporaries.
        regions += [('mw1',D*H),('vw1',D*H),('mb1',H),('vb1',H),
                    ('mw2',H*C),('vw2',H*C),('mb2',C),('vb2',C),
                    ('bc',2*T),('z',3)]
    t, a, cond, best, label = [R('s',i) for i in range(5)]
    zero, one, four, half, step = [R('k',i) for i in range(5)]
    if optimizer == 'adam':
        beta1, omb1, beta2, omb2, eps = [R('k',i) for i in range(15,20)]
        z1, z2, z3 = R('z',0), R('z',1), R('z',2)
        # Global minibatch step t under epoch and batch is epoch*NB + batch + 1.
        bc1 = R('bc', epoch=2*NB, batch=2)
        bc2 = R('bc', offset=1, epoch=2*NB, batch=2)
    body = []
    # Explicit allocation initialization is modeled work, not a free assumption.
    for region, words in regions:
        body.append(L('init',words,[I('set',R(region,init=1),0)]))
    k_values = [raw(0),raw(1),raw(4),raw(.5),raw(learning_rate/B)]+list(range(C))
    if optimizer == 'adam':
        k_values += [raw(.9),raw(.1),raw(.999),raw(.001),raw(1e-8),raw(NR_SEED_OFFSET)]
    for index, value in enumerate(k_values):
        body.append(I('set',R('k',index),value))
    initial = initial_parameters(D,H,seed)
    for region, values in [('w1',initial[0]),('w2',initial[2])]:
        body += [I('set',R(region,index),int(value))
                 for index,value in enumerate(values.reshape(-1).view(np.uint32))]
    if optimizer == 'adam':
        # Literal FP32 divisors for every global minibatch step, computed like
        # the research code: 1 - f32(beta)**t, with t starting at 1.
        for step_index in range(1, T+1):
            body.append(I('set',R('bc',2*step_index-2),raw(np.float32(1)-np.float32(.9)**step_index)))
            body.append(I('set',R('bc',2*step_index-1),raw(np.float32(1)-np.float32(.999)**step_index)))
    # Fixed tape: all train pixels, raw uint32 train labels, all test pixels.
    input_regions = [('x',N*D),('labels',N)] + ([] if stream_queries else [('q',Q*D)])
    for region, words in input_regions:
        body.append(L('recv_i',words,[I('recv',R(region,recv_i=1))]))
    normalization_regions = [('x',N*D)] + ([] if stream_queries else [('q',Q*D)])
    for region, words in normalization_regions:
        v=R(region,norm_i=1)
        body.append(L('norm_i',words,[I('mul',v,v,four),I('sub',v,v,half)]))
    # Raw integer label words compare equal to raw integer class literals;
    # select converts the boolean to the FP32 one-hot representation.
    for c in range(C):
        body.append(L('target_i',N,[I('cmp',cond,R('labels',target_i=1),R('k',5+c),predicate='eq'),
                      I('select',R('target',c,target_i=C),cond,one,zero)]))

    batch_body=[]
    hidden=R('h',b=H,h=1)
    # X_batch @ W1 + b1, ascending feature reduction; ReLU.
    batch_body.append(L('b',B,[L('h',H,[I('set',a,0),L('f',D,[
        I('mul',t,R('x',batch=B*D,b=D,f=1),R('w1',f=H,h=1)),I('add',a,a,t)]),
        I('add',a,a,R('b1',h=1)),I('cmp',cond,zero,a),I('select',hidden,cond,a,zero)])]))
    # H @ W2 + b2, immediately subtract target into delta2.
    delta2=R('d2',b=C,c=1)
    batch_body.append(L('b',B,[L('c',C,[I('set',a,0),L('h',H,[
        I('mul',t,R('h',b=H,h=1),R('w2',h=C,c=1)),I('add',a,a,t)]),
        I('add',a,a,R('b2',c=1)),I('sub',delta2,a,R('target',batch=B*C,b=C,c=1))])]))
    # delta2 @ old W2.T, masked by the ReLU activation. Strict >0 matches z>0.
    batch_body.append(L('b',B,[L('h',H,[I('set',a,0),L('c',C,[
        I('mul',t,R('d2',b=C,c=1),R('w2',h=C,c=1)),I('add',a,a,t)]),
        I('cmp',cond,zero,R('h',b=H,h=1)),I('select',R('d1',b=H,h=1),cond,a,zero)])]))
    # Each gradient uses saved activations/deltas; streaming it avoids a
    # materialized gradient matrix and preserves all pre-update dependencies.
    w=R('w1',f=H,h=1)
    if optimizer == 'adam':
        batch_body.append(L('f',D,[L('h',H,[I('set',a,0),L('b',B,[
            I('mul',t,R('x',batch=B*D,b=D,f=1),R('d1',b=H,h=1)),I('add',a,a,t)]),
            *adam_update(w,a,R('mw1',f=H,h=1),R('vw1',f=H,h=1),bc1,bc2,step,z1,z2,z3,nr_iterations)])]))
    else:
        batch_body.append(L('f',D,[L('h',H,[I('set',a,0),L('b',B,[
            I('mul',t,R('x',batch=B*D,b=D,f=1),R('d1',b=H,h=1)),I('add',a,a,t)]),
            I('mul',t,step,a),I('sub',w,w,t)])]))
    w=R('b1',h=1)
    if optimizer == 'adam':
        batch_body.append(L('h',H,[I('set',a,0),L('b',B,[I('add',a,a,R('d1',b=H,h=1))]),
                                  *adam_update(w,a,R('mb1',h=1),R('vb1',h=1),bc1,bc2,step,z1,z2,z3,nr_iterations)]))
    else:
        batch_body.append(L('h',H,[I('set',a,0),L('b',B,[I('add',a,a,R('d1',b=H,h=1))]),
                                  I('mul',t,step,a),I('sub',w,w,t)]))
    w=R('w2',h=C,c=1)
    if optimizer == 'adam':
        batch_body.append(L('h',H,[L('c',C,[I('set',a,0),L('b',B,[
            I('mul',t,R('h',b=H,h=1),R('d2',b=C,c=1)),I('add',a,a,t)]),
            *adam_update(w,a,R('mw2',h=C,c=1),R('vw2',h=C,c=1),bc1,bc2,step,z1,z2,z3,nr_iterations)])]))
    else:
        batch_body.append(L('h',H,[L('c',C,[I('set',a,0),L('b',B,[
            I('mul',t,R('h',b=H,h=1),R('d2',b=C,c=1)),I('add',a,a,t)]),
            I('mul',t,step,a),I('sub',w,w,t)])]))
    w=R('b2',c=1)
    if optimizer == 'adam':
        batch_body.append(L('c',C,[I('set',a,0),L('b',B,[I('add',a,a,R('d2',b=C,c=1))]),
                                  *adam_update(w,a,R('mb2',c=1),R('vb2',c=1),bc1,bc2,step,z1,z2,z3,nr_iterations)]))
    else:
        batch_body.append(L('c',C,[I('set',a,0),L('b',B,[I('add',a,a,R('d2',b=C,c=1))]),
                                  I('mul',t,step,a),I('sub',w,w,t)]))
    body.append(L('epoch',epochs,[L('batch',N//B,batch_body)]))

    # Inference streams queries using the first activation-buffer row.
    query_pixel = R('q',f=1) if stream_queries else R('q',query=D,f=1)
    query_body = []
    if stream_queries:
        query_body.append(L('recv_query', D, [I('recv', R('q', recv_query=1))]))
        v = R('q', norm_query=1)
        query_body.append(L('norm_query', D, [I('mul', v, v, four), I('sub', v, v, half)]))
    query_body += [L('h',H,[I('set',a,0),L('f',D,[
        I('mul',t,query_pixel,R('w1',f=H,h=1)),I('add',a,a,t)]),
        I('add',a,a,R('b1',h=1)),I('cmp',cond,zero,a),I('select',R('h',h=1),cond,a,zero)])]
    for c in range(C):
        query_body += [I('set',a,0),L('h',H,[I('mul',t,R('h',h=1),R('w2',c,h=C)),I('add',a,a,t)]),
                       I('add',a,a,R('b2',c))]
        if c==0:
            query_body += [I('copy',best,a),I('copy',label,R('k',5))]
        else:
            query_body += [I('cmp',cond,best,a),I('select',best,cond,a,best),
                           I('select',label,cond,R('k',5+c),label)]
    query_body.append(I('send',label))
    body.append(L('query',Q,query_body))
    metadata = {
        'algorithm':f'ordered-FP32 {D}-H-10 ReLU squared-error minibatch SGD',
        'width':H,'epochs':epochs,'learning_rate':learning_rate,'seed':seed,
        'train_examples':N,'test_examples':Q,'batch_size':B,
        'training_order':'fixed cyclic supplied-row order, no shuffle',
        'tape':'train pixels FP32 bits, train labels uint32, test pixels FP32 bits',
        'normalization':'x*float32(4)-float32(0.5)',
        'gradient_lowering':'stream gradient entries after all delta1 values computed; pre-update dependencies preserved',
        'initialization':'all scratch explicitly zeroed (charged); seeded initial weights set as raw FP32 literals',
        **({'features': D, 'stream_queries': True,
            'query_layout': 'one query immediately after hot scalars/constants; receive, normalize, infer, and send after training'} if stream_queries else {}),
    }
    if optimizer == 'adam':
        metadata['algorithm'] = f'ordered-FP32 {D}-H-10 ReLU squared-error minibatch Adam'
        metadata['optimizer'] = {
            'name': 'adam', 'beta1': 0.9, 'beta2': 0.999, 'epsilon': 1e-8,
            'step': 'learning_rate/batch as one FP32 scalar',
            'bias_correction': 'global minibatch step t >= 1; per-step divisors 1-beta1**t and 1-beta2**t',
            'state': 'zero-initialized m and v arrays per parameter; div FP32 primitive',
            'sqrt': f'Newton-Raphson y0 = x + 1e-6, K = {nr_iterations}, y = 0.5*(y + x/y); add/mul/div only',
        }
    return make_program(regions,body,metadata)

