"""Lab 18: fixed six-run synthetic CPU SAE experiment; no language-model execution."""
import argparse,ast,hashlib,inspect,json,platform,time
from datetime import datetime,timezone
from pathlib import Path
from inspection_common import Output,Budget,source_hashes,host_record
from sparse_feature_evaluation import evaluate

SEEDS=(101,202);PENALTIES=(0.,0.03,0.1)
CONFIG={'dimensions':16,'planted_directions':32,'learned_latents':64,
    'dictionary_seed':1801,'splits':{'train':[4096,1802],'validation':[1024,1803],'test':[2048,1804]},
    'support_probability':1/16,'amplitude_uniform':[0.5,1.5],'noise_std':0.01,
    'seeds':SEEDS,'penalties':PENALTIES,'epochs':80,'batch_size':256,'updates_per_epoch':16,
    'optimizer':{'name':'Adam','lr':0.003,'betas':[0.9,0.999],'eps':1e-8,'weight_decay':0,'foreach':False,'fused':False},
    'permutation_seed_rule':'100000+initialization_seed+zero_based_epoch',
    'activity_threshold':1e-6,'candidate_cosine':0.90,'candidate_correlation':0.80,
    'limits':{'runs':6,'optimizer_updates':7680,'training_and_evaluation_seconds':600}}


def tensor_hash(tensor,torch):
    return hashlib.sha256(bytes(tensor.detach().contiguous().view(torch.uint8).flatten().tolist())).hexdigest()


def state_hash(model,torch):
    return hashlib.sha256(''.join(name+tensor_hash(value,torch) for name,value in sorted(model.state_dict().items())).encode()).hexdigest()


def model_class(torch):
    class ToySAE(torch.nn.Module):
        def __init__(self,seed,train_mean):
            super().__init__();g=torch.Generator(device='cpu').manual_seed(seed)
            d=torch.randn((16,64),generator=g,dtype=torch.float32);d=d/d.norm(dim=0,keepdim=True)
            self.decoder=torch.nn.Parameter(d.clone());self.encoder=torch.nn.Parameter(d.T.clone())
            self.encoder_bias=torch.nn.Parameter(torch.full((64,),-0.05,dtype=torch.float32))
            self.decoder_bias=torch.nn.Parameter(train_mean.clone())
        def forward(self,x):
            z=torch.relu((x-self.decoder_bias)@self.encoder.T+self.encoder_bias)
            return z@self.decoder.T+self.decoder_bias,z
    return ToySAE


def generate(output,torch,budget,stage):
    budget.check();g=torch.Generator(device='cpu').manual_seed(1801)
    dictionary=torch.randn((16,32),generator=g,dtype=torch.float32)
    dictionary=dictionary/dictionary.norm(dim=0,keepdim=True)
    if not torch.allclose(dictionary.norm(dim=0),torch.ones(32),atol=1e-6,rtol=0):raise ValueError('Planted norms')
    hashes={}
    for name,(count,seed) in CONFIG['splits'].items():
        budget.check();g=torch.Generator(device='cpu').manual_seed(seed)
        support=torch.rand((count,32),generator=g,dtype=torch.float32)<1/16
        amplitude=0.5+torch.rand((count,32),generator=g,dtype=torch.float32)
        coefficients=support.float()*amplitude
        noise=0.01*torch.randn((count,16),generator=g,dtype=torch.float32)
        x=coefficients@dictionary.T+noise
        if x.shape!=(count,16) or x.dtype!=torch.float32 or not torch.isfinite(x).all():raise ValueError('Fixture validation')
        filename={'train':'training-input.pt','validation':'validation-input.pt','test':'test-input.pt'}[name]
        torch.save(x,output.path/filename);hashes[filename]=output.hash(filename)
        if name=='test':
            torch.save({'dictionary':dictionary,'coefficients':coefficients},output.path/'evaluation-truth.pt')
            hashes['evaluation-truth.pt']=output.hash('evaluation-truth.pt')
    stage('fixture_generation_complete',{'hashes':hashes,'training_and_validation_coefficients_discarded':True})
    return hashes


def objective(x,x_hat,z,penalty):
    reconstruction=(x-x_hat).square().sum(dim=1).mean()
    mean_l1=z.sum(dim=1).mean()
    return reconstruction+penalty*mean_l1,reconstruction,mean_l1


def scale_control(torch,decoder_scale=4.,coefficient_scale=.25):
    d=torch.tensor([1.,0.],dtype=torch.float32)
    z=torch.tensor([2.],dtype=torch.float32)
    scaled_d=d*decoder_scale;scaled_z=z*coefficient_scale
    before=d*z;after=scaled_d*scaled_z
    l0_before=int((z!=0).sum());l0_after=int((scaled_z!=0).sum())
    l1_before=float(z.abs().sum());l1_after=float(scaled_z.abs().sum())
    passed=bool(torch.allclose(before,after,atol=1e-6,rtol=0)) and l0_before==l0_after and abs(l1_after-l1_before/4)<1e-6
    return {'decoder_before':d.tolist(),'decoder_after':scaled_d.tolist(),
            'coefficient_before':z.tolist(),'coefficient_after':scaled_z.tolist(),
            'reconstruction_before':before.tolist(),'reconstruction_after':after.tolist(),
            'l0_before':l0_before,'l0_after':l0_after,'l1_before':l1_before,'l1_after':l1_after,
            'passed':passed}


def arithmetic_checks(torch):
    root2=2**0.5
    x=torch.tensor([[1.,1.]])
    dictionary=torch.tensor([[1.,0.,1/root2],[0.,1.,1/root2]])
    a=torch.tensor([[1.,1.,0.]]);b=torch.tensor([[0.,0.,root2]])
    xp=torch.tensor([[1.,0.8]])
    exact=torch.tensor([[1.,0.8,0.]]);diagonal=torch.tensor([[0.,0.,1.8/root2]])
    obj_exact=float(objective(xp,exact@dictionary.T,exact,0.1)[0])
    obj_diagonal=float(objective(xp,diagonal@dictionary.T,diagonal,0.1)[0])
    small=torch.tensor([[1.,2.],[3.,4.]])
    loss,reconstruction,l1=objective(small,torch.zeros_like(small),torch.tensor([[1.,2.,3.],[4.,5.,6.]]),0.1)
    return {'exact_reconstructions':{'passed':bool(torch.allclose(a@dictionary.T,x,atol=1e-6,rtol=0) and torch.allclose(b@dictionary.T,x,atol=1e-6,rtol=0)),
                'l0':[int((a>0).sum()),int((b>0).sum())],'l1':[float(a.sum()),float(b.sum())]},
            'fixed_objectives':{'passed':abs(obj_exact-.18)<1e-6 and abs(obj_diagonal-(.02+.1*1.8/root2))<1e-6,'exact':obj_exact,'diagonal':obj_diagonal},
            'small_batch_reductions':{'passed':abs(float(reconstruction)-15)<1e-6 and abs(float(l1)-10.5)<1e-6 and abs(float(loss)-16.05)<2e-6,
                                      'reconstruction':float(reconstruction),'mean_l1':float(l1),'objective':float(loss)},
            'regularized_diagonal':{'value':max(0,1.8/root2-.05),'kind':'analytical'},
            'scale_control':scale_control(torch)}


def summaries(model,x,torch):
    model.eval()
    with torch.inference_mode():
        reconstruction,z=model(x)
        if not torch.isfinite(reconstruction).all() or not torch.isfinite(z).all():raise ValueError('Nonfinite history evaluation')
        result={'reconstruction':float((x-reconstruction).square().sum(1).mean()),
                'mean_l1':float(z.sum(1).mean()),'mean_l0':float((z>1e-6).float().sum(1).mean())}
    model.train();return result


def train_run(train,validation,train_mean,seed,penalty,budget,output,torch,ToySAE,stage,updates):
    """Trainer receives input tensors only; it never loads files or evaluation truth."""
    identifier='seed-'+str(seed)+'-penalty-'+str(penalty)
    record={'id':identifier,'seed':seed,'penalty':penalty,'training_status':'running',
            'checkpoint_status':'not_started','evaluation_status':'not_started','epoch':0,
            'updates':0,'permutation_hashes':[],'history':[],'checkpoints':[]}
    try:
        stage('training_run_started',{'id':identifier})
        budget.check();model=ToySAE(seed,train_mean)
        if any(not torch.isfinite(p).all() for p in model.parameters()):raise ValueError('Nonfinite initialization')
        optimizer=torch.optim.Adam(model.parameters(),lr=.003,betas=(.9,.999),eps=1e-8,weight_decay=0,foreach=False,fused=False)
        record['initial_state_hash']=state_hash(model,torch)
        for epoch in range(81):
            if epoch in (0,20,40,60,80):
                budget.check()
                record['history'].append({'epoch':epoch,'training':summaries(model,train,torch),
                                          'validation':summaries(model,validation,torch)})
                filename=identifier+'-epoch-'+str(epoch)+'.pt'
                torch.save(model.state_dict(),output.path/filename)
                record['checkpoints'].append({'epoch':epoch,'path':filename,'sha256':output.hash(filename)})
                record['checkpoint_status']='completed';record['epoch']=epoch
            if epoch==80:break
            budget.check();permutation=torch.randperm(4096,generator=torch.Generator(device='cpu').manual_seed(100000+seed+epoch))
            record['permutation_hashes'].append(tensor_hash(permutation,torch))
            for offset in range(0,4096,256):
                budget.check()
                if updates['count']>=7680:raise ValueError('Optimizer update limit')
                batch=train[permutation[offset:offset+256]]
                optimizer.zero_grad(set_to_none=True);x_hat,z=model(batch)
                loss,_,_=objective(batch,x_hat,z,penalty)
                if not torch.isfinite(loss) or not torch.isfinite(z).all():raise ValueError('Nonfinite objective or coefficients')
                loss.backward()
                if any(p.grad is None or not torch.isfinite(p.grad).all() for p in model.parameters()):raise ValueError('Nonfinite gradients')
                with torch.no_grad():
                    d=model.decoder;grad=d.grad;grad.sub_(d*(grad*d).sum(dim=0,keepdim=True))
                budget.check();updates['count']+=1;record['updates']+=1;optimizer.step()
                with torch.no_grad():
                    norms=model.decoder.norm(dim=0,keepdim=True)
                    if not torch.isfinite(norms).all() or (norms<=1e-12).any():raise ValueError('Invalid decoder norms')
                    model.decoder.div_(norms)
                    if any(not torch.isfinite(p).all() for p in model.parameters()):raise ValueError('Nonfinite parameters')
            record['epoch']=epoch+1
        record.update(training_status='completed',checkpoint=record['checkpoints'][-1]['path'],
                      checkpoint_sha256=record['checkpoints'][-1]['sha256'])
    except Exception as exc:
        record.update(training_status='interrupted' if isinstance(exc,TimeoutError) else 'failed',
                      reason=type(exc).__name__+': '+str(exc))
    stage('training_run_finalized',{'id':identifier,'status':record['training_status'],'updates':record['updates']})
    output.json('training-'+identifier+'.json',record)
    return record


def main():
    parser=argparse.ArgumentParser(description=__doc__);parser.add_argument('--output',type=Path,required=True)
    args=parser.parse_args();output=Output(args.output)
    import torch
    torch.set_num_threads(1);torch.set_num_interop_threads(1);torch.use_deterministic_algorithms(True)
    ToySAE=model_class(torch)
    started_utc=datetime.now(timezone.utc).isoformat();budget=Budget(600);updates={'count':0};runs=[];evaluations={};validation={}
    def stage(name,data=None):
        with (output.path/'stages.jsonl').open('a') as f:
            f.write(json.dumps({'stage':name,'monotonic':time.perf_counter(),'data':data},allow_nan=False)+'\n')
    output.json('config.json',CONFIG|{'source_hashes':source_hashes(['inspect_sparse_features.py','sparse_feature_evaluation.py','inspection_common.py'])})
    output.json('predictions.json',{'before_training':True,'penalty_effects':'may_increase_error_reduce_activity_change_dead_latents_or_matching_not_guaranteed',
        'seeds':'different_local_solutions_possible','no_semantic_names':True})
    scheduled=[{'id':'seed-'+str(seed)+'-penalty-'+str(penalty),'seed':seed,'penalty':penalty,
                'training_status':'not_started','checkpoint_status':'not_started','evaluation_status':'not_started'} for seed in SEEDS for penalty in PENALTIES]
    failure=None;fixture_hashes={};train_mean=None;controls=[]
    try:
        with budget.execution():
            fixture_hashes=generate(output,torch,budget,stage)
            train=torch.load(output.path/'training-input.pt',map_location='cpu',weights_only=True)
            validation_inputs=torch.load(output.path/'validation-input.pt',map_location='cpu',weights_only=True)
            budget.check();train_mean=train.mean(0)
            validation['fixture_shapes_finiteness_norms']={'status':'passed'}
            validation['arithmetic']=arithmetic_checks(torch)
            if not all(check['passed'] for check in validation['arithmetic'].values() if 'passed' in check):
                raise ValueError('Analytical validation failed before training')
            for seed in SEEDS:
                budget.check();model=ToySAE(seed,train_mean)
                filename='untrained-'+str(seed)+'.pt'
                torch.save(model.state_dict(),output.path/filename)
                controls.append({'seed':seed,'checkpoint':filename,'sha256':output.hash(filename)})
            for spec in scheduled:
                if time.perf_counter()>=budget.deadline:
                    runs.append(spec|{'reason':'wall_time_limit'});continue
                run=train_run(train,validation_inputs,train_mean,spec['seed'],spec['penalty'],budget,output,torch,ToySAE,stage,updates)
                runs.append(run)
            output.json('completed-runs.json',{'runs':runs,'controls':controls,'fixture_hashes':fixture_hashes,
                                              'optimizer_updates':updates['count']})
            stage('completed_run_manifest_written',{'sha256':output.hash('completed-runs.json')})
            by_seed={seed:[r for r in runs if r['seed']==seed and 'initial_state_hash' in r] for seed in SEEDS}
            validation['identical_initialization_and_batch_order']={'status':'passed' if all(
                len({r['initial_state_hash'] for r in group})<=1 and all(
                r['permutation_hashes'][:min(len(r['permutation_hashes']),len(group[0]['permutation_hashes']))]==
                group[0]['permutation_hashes'][:min(len(r['permutation_hashes']),len(group[0]['permutation_hashes']))]
                for r in group) for group in by_seed.values()) else 'failed',
                'runs_started_per_seed':{str(seed):len(group) for seed,group in by_seed.items()}}
            trainer_tree=ast.parse(inspect.getsource(train_run))
            forbidden_load=any(isinstance(node,ast.Call) and isinstance(node.func,ast.Attribute)
                               and node.func.attr in ('load','load_state_dict','read_bytes','read_text') for node in ast.walk(trainer_tree))
            validation['trainer_inputs_only']={'status':'failed' if forbidden_load else 'passed',
                'parameters':list(inspect.signature(train_run).parameters),
                'source_check':'AST checked for file-loading calls; trainer arguments contain inputs, never planted truth'}
            if time.perf_counter()<budget.deadline:
                evaluate(output,budget,torch,ToySAE,train_mean,runs,stage,evaluations,validation)
    except Exception as exc:
        failure=type(exc).__name__+': '+str(exc)
    # Metadata persistence is allowed after a timeout; no numerical work is resumed.
    known={r['id'] for r in runs}
    runs.extend(spec|{'reason':failure or 'not_started'} for spec in scheduled if spec['id'] not in known)
    if not (output.path/'completed-runs.json').exists():
        output.json('completed-runs.json',{'runs':runs,'optimizer_updates':updates['count'],'failure':failure})
        stage('completed_run_manifest_written',{'sha256':output.hash('completed-runs.json')})
    for spec in runs:
        spec['evaluation_status']=evaluations.get(spec['id'],{}).get('status','not_started')
        if spec['evaluation_status']=='not_started':spec['evaluation_reason']=failure or 'no_completed_checkpoint_or_wall_time_limit'
    for name in ['untrained-101','untrained-202','training_mean','signed_identity']:
        evaluations.setdefault(name,{'status':'not_started','reason':failure or 'wall_time_limit'})
    for check in ('fixture_shapes_finiteness_norms','arithmetic','identical_initialization_and_batch_order','trainer_inputs_only','truth_loaded_after_frozen_manifest'):
        validation.setdefault(check,{'status':'not_performed','reason':failure or 'wall_time_limit'})
    validation['identity_reconstruction']=evaluations['signed_identity'].get('identity_passed',None)
    validation['permutation_invariance']={name:r.get('permutation_check') for name,r in evaluations.items() if 'permutation_check' in r}
    validation['all_runs_accounted_for']={'status':'passed' if len(runs)==6 else 'failed'}
    output.json('outcome.json',{'runs':runs,'evaluations':evaluations,'failure':failure,'optimizer_updates':updates['count'],
                             'elapsed_seconds':time.perf_counter()-budget.started,'test_results_used_for_selection':False})
    output.json('validation.json',validation)
    output.json('environment.json',{'python':platform.python_version(),'torch':torch.__version__,**host_record(),
        'dtype':'float32','device':'cpu',
        'threads':{'intra_op':torch.get_num_threads(),'inter_op':torch.get_num_interop_threads()},
        'deterministic_algorithms':torch.are_deterministic_algorithms_enabled(),'start_utc':started_utc,
        'end_utc':datetime.now(timezone.utc).isoformat(),'elapsed_seconds':time.perf_counter()-budget.started})
    output.json('artifact-hashes.json',{p.name:output.hash(p.name) for p in output.path.iterdir() if p.is_file()})
    print('Recorded',len(runs),'scheduled runs;',updates['count'],'optimizer updates; failures retained.')
    if failure:raise SystemExit(1)


if __name__=='__main__':main()
