"""COPT assignment of complete GEMV output tiles to persistent RF workgroups."""

import json
import math
from pathlib import Path

import coptpy as cp

from codesign.challenge.service import mma_service


def solve_packing(hardware, tile_n=16, seconds=20, io_aware=False, bin_slots=None, coalesce_w1=False,w2_partitions=1,reserve_norm_lane=False,qkv_partitions=1,stage_tiles=None,query_early=False):
    sm_count = hardware.sm_count
    if query_early and (sm_count!=16 or tile_n!=16 or w2_partitions!=2 or qkv_partitions!=1 or stage_tiles or not reserve_norm_lane):
        raise ValueError('Early queries require the reserved-lane 16-SM N16 layout')
    if qkv_partitions not in (1,2):raise ValueError('QKV supports one or two K partitions')
    unit=64 if qkv_partitions==2 else 128
    unit_words=1024 if stage_tiles else unit*tile_n
    cap = (6 if reserve_norm_lane else 7)*2048//unit_words
    if w2_partitions not in (1,2,4):raise ValueError('W2 supports one, two or four K partitions')
    tasks=[]
    for layer in range(2):
        for kind, inner, width in (("wqkv",128,384),("wo",128,128),("w1",128,512),("w2",512,128)):
            task_n=(stage_tiles or {}).get(kind,tile_n)
            if task_n not in (8,16,32) or width%task_n:raise ValueError('Unsupported stage tile width')
            service=mma_service(hardware,1,task_n,128)
            ports={"8R4W":(128,64),"4R2W":(64,32),"2R1W":(32,16)}[hardware.rf_ports]
            partitions=w2_partitions if kind=='w2' else qkv_partitions if kind=='wqkv' else 1
            part_inner=inner//partitions
            cost=(service.compute_cycles+service.rf_read_bytes/ports[0]+service.rf_write_bytes/ports[1])*(part_inner/128)
            for col in range(0,width,task_n):
                for part in range(partitions):
                    task=dict(stage=f"layer{layer}/{kind}",inner=part_inner,width=width,col=col,n=task_n,size=part_inner*task_n//unit_words,cost=cost)
                    if partitions>1:task.update(k_start=part*part_inner,full_inner=inner,part=part,parts=partitions)
                    tasks.append(task)
    count=math.ceil(sum(t['size'] for t in tasks)/(sm_count*cap))
    if count>4: raise ValueError("Weights exceed available workgroup RF")
    if bin_slots is None:
        bin_slots={sm:min(count,3 if sm==0 else 4) for sm in range(sm_count)}
    if sum(bin_slots.values())*cap < sum(t['size'] for t in tasks):
        raise ValueError('Requested RF bin layout lacks capacity')
    bins=[dict(sm=sm,slot=j) for sm in range(sm_count) for j in range(bin_slots.get(sm,0))]
    used=[0]*len(bins);warm={};stage_load={}
    for ti in sorted(range(len(tasks)),key=lambda i:-tasks[i]['size']):
        t=tasks[ti]
        legal=[bi for bi,b in enumerate(bins) if used[bi]+t['size']<=cap]
        if not legal:raise ValueError("Greedy capacity initialization failed")
        bi=min(legal,key=lambda i:(stage_load.get((t['stage'],bins[i]['sm']),0),sum(used[j] for j,b in enumerate(bins) if b['sm']==bins[i]['sm']),i))
        used[bi]+=t['size'];warm[ti]=bi
        key=t['stage'],bins[bi]['sm'];stage_load[key]=stage_load.get(key,0)+t['cost']
    env=cp.Envr();model=env.createModel('persistent_rf_packing')
    model.setParam(cp.COPT.Param.Logging,0);model.setParam(cp.COPT.Param.Threads,4);model.setParam(cp.COPT.Param.TimeLimit,seconds)
    x={}
    for ti,t in enumerate(tasks):
        for bi,b in enumerate(bins):
            if query_early:
                layer,kind=t['stage'].split('/');layer_index=int(layer.removeprefix('layer'))
                if b['slot']!=layer_index:continue
                if kind=='wqkv':owner=t['col']//16 if t['col']<128 else 8+(t['col']-128)//32
                elif kind=='wo':owner=t['col']//16
                elif kind=='w1':owner=t['col']//32
                else:owner=(t['col']//16)*2+t['part']
                if b['sm']!=owner:continue
            # Keep balanced W2 ownership as a bounded-search symmetry anchor.
            if t['stage'].endswith('/w2') and b['sm']!=bins[warm[ti]]['sm']:continue
            x[ti,bi]=model.addVar(vtype=cp.COPT.BINARY,name=f'x_{ti}_{bi}')
        model.addConstr(cp.quicksum(v for (i,_),v in x.items() if i==ti)==1)
    for bi in range(len(bins)):
        model.addConstr(cp.quicksum(tasks[ti]['size']*v for (ti,j),v in x.items() if j==bi)<=cap)
    stages=sorted({t['stage'] for t in tasks});tau=model.addVars(stages,lb=0,nameprefix='load')
    active={}
    if io_aware or coalesce_w1:
        for stage in stages:
            for bi in range(len(bins)):
                expr=cp.quicksum(v for (ti,j),v in x.items() if j==bi and tasks[ti]['stage']==stage)
                active[stage,bi]=model.addVar(vtype=cp.COPT.BINARY,name=f'active_{stage}_{bi}')
                model.addConstr(expr<=len(tasks)*active[stage,bi]);model.addConstr(active[stage,bi]<=expr)
        if coalesce_w1:
            for stage in stages:
                if stage.endswith('/w1') or (qkv_partitions>1 and stage.endswith('/wqkv')):
                    for sm in range(sm_count):
                        model.addConstr(cp.quicksum(active[stage,bi] for bi,b in enumerate(bins) if b['sm']==sm)<=1)
    for stage in stages:
        for sm in range(sm_count):
            expr=cp.quicksum(tasks[ti]['cost']*v for (ti,bi),v in x.items() if tasks[ti]['stage']==stage and bins[bi]['sm']==sm)
            if io_aware:
                k=next(t['inner'] for t in tasks if t['stage']==stage)
                io=4+4*k/hardware.sm_noc_bytes_per_cycle+4*k/ports[1]+math.ceil(4*k/64)
                expr=expr+cp.quicksum(io*active[stage,bi] for bi,b in enumerate(bins) if b['sm']==sm)
            model.addConstr(tau[stage]>=expr)
    model.setObjective(cp.quicksum(tau.values()),cp.COPT.MINIMIZE)
    model.setMipStart(list(x.values()),[int(warm[ti]==bi) for ti,bi in x])
    model.loadMipStart();model.solve()
    if not model.hasmipsol:raise ValueError('COPT did not produce a feasible RF assignment')
    assignment=[next(bi for (i,bi),v in x.items() if i==ti and v.x>.5) for ti in range(len(tasks))]
    for bi in range(len(bins)):
        if sum(t['size'] for t,j in zip(tasks,assignment) if j==bi)>cap:raise AssertionError('RF assignment exceeds capacity')
    return {'tasks':tasks,'bins':bins,'assignment':assignment,'tile_n':tile_n,'capacity_units':cap,'slot_words':unit_words,
            'objective':model.objval,'bound':model.bestbnd,'gap':model.objval-model.bestbnd,
            'io_aware':io_aware,'coalesce_w1':coalesce_w1,'w2_partitions':w2_partitions,'reserve_norm_lane':reserve_norm_lane,'qkv_partitions':qkv_partitions,'stage_tiles':stage_tiles,'query_early':query_early,
            'stage_k_chunks':({kind:min(128,(16384//hardware.vector_lanes)//(stage_tiles or {}).get(kind,tile_n)) for kind in ('wqkv','wo','w1','w2')} if stage_tiles else {}),
            'constraint_scope':'Exact RF/occupancy limits with fixed balanced W2 SM ownership'}
