"""Lab 18 evaluation only: truth deserialization follows the frozen-run manifest."""
import json,time


def correlation(a,b):
    a=a.double()-a.double().mean();b=b.double()-b.double().mean()
    denominator=float(a.norm()*b.norm())
    if denominator==0:return {'value':None,'reason':'zero_centered_norm'}
    return {'value':float(a.dot(b)/denominator),'reason':None}


def metrics(x,reconstruction,z,decoder,torch):
    if not all(torch.isfinite(t).all() for t in (x,reconstruction,z,decoder)):
        raise ValueError('Nonfinite evaluation tensors')
    residual=(x-reconstruction).square().sum(1)
    variance=float((x-x.mean(0)).square().sum())
    norms=decoder.norm(dim=0)
    return {'reconstruction_error':float(residual.mean()),
        'fvu':float(residual.sum())/variance if variance else None,
        'fvu_reason':None if variance else 'zero_test_variance',
        'mean_l1':float(z.sum(1).mean()),'mean_l0_threshold':float((z>1e-6).float().sum(1).mean()),
        'mean_l0_exact_positive':float((z>0).float().sum(1).mean()),
        'firing_frequency':(z>1e-6).float().mean(0).tolist(),
        'dead_latents':int((~(z>1e-6).any(0)).sum()),
        'decoder_norm_min':float(norms.min()),'decoder_norm_max':float(norms.max())}


def matching(dictionary,truth,decoder,z,torch,budget,result=None):
    budget.check()
    norms=decoder.norm(dim=0,keepdim=True)
    if not torch.isfinite(norms).all() or (norms<=1e-12).any():raise ValueError('Invalid matching decoder')
    cosines=dictionary.T@(decoder/norms)
    if not torch.isfinite(cosines).all():raise ValueError('Nonfinite matching cosines')
    # torch.argmax selects the first index on an exact tie, without absolute values.
    selected=cosines.argmax(1)
    permutation=torch.randperm(len(truth),generator=torch.Generator(device='cpu').manual_seed(1805))
    shuffled=truth[permutation];records=[]
    result={} if result is None else result
    result.update(status='running',records=records,denominator=32,shuffle_seed=1805,
                  shuffle_indices=permutation.tolist())
    for k,j in enumerate(selected.tolist()):
        budget.check();corr=correlation(truth[:,k],z[:,j]);negative=correlation(shuffled[:,k],z[:,j])
        records.append({'planted_index':k,'latent_index':j,'cosine':float(cosines[k,j]),
                        'correlation':corr['value'],'correlation_reason':corr['reason'],
                        'shuffled_correlation':negative['value'],'shuffled_reason':negative['reason']})
    assert len(records)==32
    budget.check()
    result.update({'mean_best_cosine':sum(r['cosine'] for r in records)/32,
        'minimum_best_cosine':min(r['cosine'] for r in records),
        'candidate_cosine_fraction':sum(r['cosine']>=0.90 for r in records)/32,
        'candidate_joint_fraction':sum(r['cosine']>=0.90 and r['correlation'] is not None
                                       and r['correlation']>=0.80 for r in records)/32,
        'denominator':32,'collision_count':32-len(set(selected.tolist())),
        'records':records,'shuffle_seed':1805,'shuffle_indices':permutation.tolist(),'status':'completed'})
    return result


def evaluate(output,budget,torch,model_class,train_mean,runs,stage,results,checks):
    """Load truth here only, once every scheduled training status is frozen."""
    frozen_path=output.path/'completed-runs.json'
    if not frozen_path.is_file():raise ValueError('Training manifest must precede evaluation')
    frozen=json.loads(frozen_path.read_text())
    if len(frozen['runs'])!=6 or any(r['training_status']=='running' for r in frozen['runs']):
        raise ValueError('Training is not frozen')
    for name,digest in frozen['fixture_hashes'].items():
        if output.hash(name)!=digest:raise ValueError('Fixture changed before evaluation')
    for run in frozen['runs']:
        if run['training_status']=='completed' and output.hash(run['checkpoint'])!=run['checkpoint_sha256']:
            raise ValueError('Frozen checkpoint changed')
    for control in frozen['controls']:
        if output.hash(control['checkpoint'])!=control['sha256']:raise ValueError('Untrained checkpoint changed')
    budget.check()
    stage('evaluation_truth_deserialize',{'frozen_manifest_sha256':output.hash('completed-runs.json')})
    truth=torch.load(output.path/'evaluation-truth.pt',map_location='cpu',weights_only=True)
    budget.check()
    x=torch.load(output.path/'test-input.pt',map_location='cpu',weights_only=True)
    dictionary,z_true=truth['dictionary'],truth['coefficients']
    if x.shape!=(2048,16) or dictionary.shape!=(16,32) or z_true.shape!=(2048,32):
        raise ValueError('Evaluation fixture shapes differ')
    checks['truth_loaded_after_frozen_manifest']={'status':'passed'}
    jobs=[('untrained-'+str(seed),'untrained-'+str(seed)+'.pt',seed) for seed in (101,202)]
    jobs += [(r['id'],r['checkpoint'],r['seed']) for r in runs if r['training_status']=='completed']
    for name,checkpoint,seed in jobs:
        results[name]={'status':'not_started'}
    results['training_mean']={'status':'not_started'};results['signed_identity']={'status':'not_started'}
    output.json('evaluation-plan.json',results)
    with torch.inference_mode():
        for name,checkpoint,seed in jobs:
            if time.perf_counter()>=budget.deadline:
                results[name]={'status':'not_started','reason':'wall_time_limit'};continue
            partial={'status':'running'}
            try:
                budget.check();model=model_class(seed,train_mean)
                model.load_state_dict(torch.load(output.path/checkpoint,map_location='cpu',weights_only=True))
                model.eval();budget.check();reconstruction,z=model(x)
                partial['metrics']=metrics(x,reconstruction,z,model.decoder,torch)
                if not torch.allclose(model.decoder.norm(dim=0),torch.ones(64),atol=1e-5,rtol=0):
                    raise ValueError('Decoder normalization failed')
                partial['matching']={}
                matching(dictionary,z_true,model.decoder,z,torch,budget,partial['matching'])
                budget.check()
                # Permute decoder columns and matching encoder rows/biases together.
                order=torch.arange(63,-1,-1)
                permuted=model_class(seed,train_mean)
                permuted.load_state_dict({'decoder':model.decoder[:,order].clone(),
                    'encoder':model.encoder[order].clone(),'encoder_bias':model.encoder_bias[order].clone(),
                    'decoder_bias':model.decoder_bias.clone()})
                permuted.eval();other,_=permuted(x)
                error=float((other-reconstruction).abs().max())
                partial['permutation_check']={'max_abs_error':error,'passed':bool(torch.allclose(other,reconstruction,atol=1e-5,rtol=0))}
                partial['status']='completed';results[name]=partial
                output.json('evaluation-'+name+'.json',partial)
            except Exception as exc:
                partial.update(status='interrupted' if isinstance(exc,TimeoutError) else 'failed',
                               reason=type(exc).__name__+': '+str(exc));results[name]=partial
                output.json('evaluation-'+name+'.json',partial)
        for name in ('training_mean','signed_identity'):
            if time.perf_counter()>=budget.deadline:
                results[name]={'status':'not_started','reason':'wall_time_limit'};continue
            try:
                budget.check()
                if name=='training_mean':
                    reconstruction=train_mean.expand_as(x)
                    variance=float((x-x.mean(0)).square().sum());error=float((x-reconstruction).square().sum())
                    results[name]={'status':'completed','kind':'training_mean_reconstruction',
                        'reconstruction_error':error/len(x),'fvu':error/variance if variance else None,
                        'fvu_reason':None if variance else 'zero_test_variance'}
                else:
                    decoder=torch.cat((torch.eye(16),-torch.eye(16)),dim=1)
                    z=torch.cat((torch.relu(x),torch.relu(-x)),dim=1);reconstruction=z@decoder.T
                    results[name]={'status':'completed','kind':'analytical_hand_constructed_control',
                        'metrics':metrics(x,reconstruction,z,decoder,torch),
                        'identity_max_abs_error':float((x-reconstruction).abs().max()),
                        'identity_passed':bool(torch.allclose(x,reconstruction,atol=1e-6,rtol=0))}
                output.json('evaluation-'+name+'.json',results[name])
            except Exception as exc:
                results[name]={'status':'interrupted' if isinstance(exc,TimeoutError) else 'failed','reason':str(exc)}
    return results,checks
