"""Overlap stable LayerNorm moments with RF-resident decode projections.

For each projection, cache -gamma@W and beta@W (plus its output bias).
Then evaluate ((x*gamma)@W + mean*(-gamma@W))*invstd + beta@W.
This changes FP32 association, so complete correctness checks are required.
"""

import json
from project.compiler import rf,hbm,imm


def prepare_constants(b):
    b.projection_constants={}
    for bi,ids in b.bin_tasks.items():
        b.wg=b.worker_names[bi];cursor=1024
        for ti in ids:
            task=b.tasks[ti];layer,kind=task['stage'].split('/')
            if kind not in ('wqkv','w1'):continue
            n=task['n']
            if task['inner']!=128 or task.get('parts',1)!=1 or n!=16:
                raise ValueError('Projected normalization currently requires complete K128/N16 tasks')
            if cursor+2*n>2048:raise ValueError('Projected constants exceed the reserved RF lane')
            norm='ln1' if kind=='wqkv' else 'ln2'
            gamma=b.symbol(layer+'/'+norm+'_g')
            g,beta=b.worker_norm_cache[bi,gamma.offset]
            delta=rf(7,n,cursor);offset=rf(7,n,cursor+n)
            if b.allocation.get('lazy_norm_constants',False):
                b.projection_constants[ti]=(delta,offset);cursor+=2*n
                continue
            b.vec('add',[imm(0),imm(0)],rf(7,2*n,cursor))
            combined=b.allocation.get('combined_norm_constants',False)
            if combined:b.vec('add',[imm(0),imm(0)],rf(0,2*n,1024))
            for index,inner in enumerate(range(0,128,b.k_chunk)):
                lane,start=b.slots[ti,0];weight=rf(lane,b.k_chunk*n,start+inner*n)
                if combined:
                    part=rf(g['lane'],2*b.k_chunk,g['offset']+inner)|{'shape':[2,b.k_chunk],'strides':[128,1]}
                    acc=rf(7,2*n,cursor) if index%2==0 else rf(0,2*n,1024)
                    b.emit('MMA.ACC',a=part,b=weight,acc=acc,m=2,n=n,k=b.k_chunk,event=None)
                else:
                    for source,target in ((g,delta),(beta,offset)):
                        part=rf(source['lane'],b.k_chunk,source['offset']+inner)
                        b.emit('MMA.ACC',a=part,b=weight,acc=target,m=1,n=n,k=b.k_chunk,event=None)
            if combined:b.vec('add',[rf(7,2*n,cursor),rf(0,2*n,1024)],rf(7,2*n,cursor))
            b.vec('mul',[delta,imm(-1)],delta)
            if kind=='w1':
                if not b.cache_bias:raise ValueError('Projected normalization requires cached output bias')
                b.vec('add',[offset,b.bias_view(ti,0,n)],offset)
            b.projection_constants[ti]=(delta,offset);cursor+=2*n
    b.wg=b.controller
    if not b.allocation.get('lazy_norm_constants',False):
        b.emit('BARRIER',wgs=list(b.participants),events=[])


def bootstrap_projection(b,tasks,k_extent,gelu):
    """Compute the first projection and its two constants in one M=3 GEMM."""
    if k_extent!=128 or any(b.tasks[ti]['n']!=16 for ti in tasks) or len(tasks)>2:
        raise ValueError('Joint constant bootstrap requires at most two K128/N16 tasks')
    cells=3*16*len(tasks)
    for base in (1024,1280):b.vec('add',[imm(0),imm(0)],rf(0,cells,base))
    for index,inner in enumerate(range(0,k_extent,b.k_chunk)):
        for pos,ti in enumerate(tasks):
            lane,offset=b.slots[ti,0]
            a=rf(0,3*b.k_chunk,inner)|{'shape':[3,b.k_chunk],'strides':[128,1]}
            b.emit('MMA.ACC',a=a,b=rf(lane,b.k_chunk*16,offset+inner*16),
                   acc=rf(0,48,(1024 if index%2==0 else 1280)+pos*48),
                   m=3,n=16,k=b.k_chunk,event=None)
    for pos,ti in enumerate(tasks):
        base=1024+pos*48
        b.vec('add',[rf(0,48,base),rf(0,48,1280+pos*48)],rf(0,48,base))
        delta,offset=b.projection_constants[ti]
        b.vec('mul',[rf(0,16,base+16),imm(-1)],delta)
        if gelu is not None:
            b.vec('add',[rf(0,16,base+32),b.bias_view(ti,pos,16)],offset)
        else:b.vec('add',[rf(0,16,base+32),imm(0)],offset)
    # Save all constants before compacting outputs over their scratch rows.
    for pos,ti in enumerate(tasks):
        if pos:b.vec('add',[rf(0,16,1024+pos*48),imm(0)],rf(0,16,1024+pos*16))


def emit_moments(b,source,d,name):
    if d!=128:raise ValueError('Projected decode normalization requires width 128')
    if b.allocation.get('distributed_moments'):
        return distributed_moments(b,source,d,name)
    pending=getattr(b,'pending_store_events',[])
    if pending:b.wait_for_store_events(pending,(b.controller,))
    b.wg=b.controller
    output=b.alloc(name+'_moments',64)  # Preserve 256-byte alignment of later tensors.
    mean=rf(0,1,256);inv=rf(0,1,257)
    b.ld(hbm(source.offset,d),rf(0,d))
    b.reduce('sum',rf(0,d),mean);b.vec('mul',[mean,imm(1/d)],mean)
    b.vec('sub',[rf(0,d),mean],rf(0,d))
    b.vec('mul',[rf(0,d),rf(0,d)],rf(0,d,128))
    b.reduce('sum',rf(0,d,128),inv)
    b.vec('fma',[inv,imm(1/d),imm(1e-5)],inv);b.sfu('rsqrt',inv,inv)
    b.st(rf(0,2,256),hbm(output.offset,2))
    return output,json.loads(b.lines[-1].split(' ',1)[1])['event']


def load_moments(b):
    output,event,*_=b.projection_normalizer
    b.emit('WAIT',wg=b.wg,events=[event])
    b.ld(hbm(output.offset,2),rf(0,2,672))


def finish_projection(b,tasks):
    for pos,ti in enumerate(tasks):
        n=b.tasks[ti]['n'];x=rf(0,n,1024+pos*n);delta,_=b.projection_constants[ti]
        b.vec('fma',[delta,rf(0,1,672),x],x)
    for pos,ti in enumerate(tasks):
        n=b.tasks[ti]['n'];x=rf(0,n,1024+pos*n);_,offset=b.projection_constants[ti]
        b.vec('fma',[x,rf(0,1,673),offset],x)


def distributed_moments(b,source,d,name):
    """Merge local centered sums of squares, avoiding cancellation in E[x²]."""
    parts=b.allocation['distributed_moments'];stride=b.allocation.get('moment_stride',64)
    if parts not in (2,4,8,16) or len(b.history_workers)!=16 or not b.loose_store or stride not in (16,64):
        raise ValueError('Distributed moments require 8/16 groups, four-way history, and tracked stores')
    light=b.allocation.get('moment_worker_layout')=='query_light' and name.endswith('ln1')
    indices=([1,5] if light else [1,9]) if parts==2 else ([1,3,5,7] if light else [1,5,9,13]) if parts==4 else list(range(8) if light else range(1,16,2)) if parts==8 else list(range(16))
    workers=tuple(b.history_workers[i] for i in indices)
    pending=getattr(b,'pending_store_events',[])
    if pending:b.wait_for_store_events(pending,(b.controller,*workers))
    partial=b.alloc(name+'_partial_moments',parts*stride)
    output=b.alloc(name+'_moments',64)
    chunk=d//parts;events=[]
    for index,worker in enumerate(workers):
        b.wg=worker;mean=rf(0,1,256);m2=rf(0,1,257)
        b.ld(hbm(source.offset+index*chunk,chunk),rf(0,chunk))
        b.reduce('sum',rf(0,chunk),mean);b.vec('mul',[mean,imm(1/chunk)],mean)
        b.vec('sub',[rf(0,chunk),mean],rf(0,chunk))
        b.vec('mul',[rf(0,chunk),rf(0,chunk)],rf(0,chunk,128))
        b.reduce('sum',rf(0,chunk,128),m2)
        b.st(rf(0,2,256),hbm(partial.offset+index*stride,2))
        events.append(json.loads(b.lines[-1].split(' ',1)[1])['event'])
    b.wg=b.controller;b.emit('WAIT',wg=b.wg,events=events)
    b.ld(hbm(partial.offset,2*parts,[parts,2],[stride,1]),rf(0,2*parts,128))
    means=rf(0,parts,128)|{'shape':[parts],'strides':[2]}
    m2=rf(0,parts,129)|{'shape':[parts],'strides':[2]}
    mean=rf(0,1,256);inv=rf(0,1,257);delta=rf(0,parts,512)
    b.reduce('sum',means,mean);b.vec('mul',[mean,imm(1/parts)],mean)
    b.vec('sub',[means,mean],delta);b.vec('mul',[delta,delta],delta)
    b.vec('fma',[delta,imm(chunk),m2],delta);b.reduce('sum',delta,inv)
    b.vec('fma',[inv,imm(1/d),imm(1e-5)],inv);b.sfu('rsqrt',inv,inv)
    b.st(rf(0,2,256),hbm(output.offset,2))
    return output,json.loads(b.lines[-1].split(' ',1)[1])['event'],workers
