"""Lab 17: fixed offline Pythia observations, logit lens and optional held-out probe."""
import argparse,json,time
from pathlib import Path
from inspection_common import Output,Budget,runtime,environment,load_specimen,source_hashes

ATOL,RTOL=1e-6,1e-5
CORE=("The boat is red.\nColor:","The boat is blue.\nColor:",
      "A red flag hangs by the door.\nColor:","A blue flag hangs by the door.\nColor:",
      "The label reads red.\nColor:","The label reads blue.\nColor:",
      "Mira copied the word red.\nColor:","Mira copied the word blue.\nColor:")


def probe_rows():
    rows=[]
    for family in range(4):
        for assignment in (('red','blue'),('blue','red')):
            colors=dict(zip(('cedar','maple'),assignment))
            for order in (('cedar','maple'),('maple','cedar')):
                for query in ('cedar','maple'):
                    templates=('The {box} box is {color}.','Box {box}: {color}.',
                               'A {color} tag marks {box}.','Record {box} as {color}.')
                    questions=(' What color is the {query} box?',' Give the color of box {query}.',
                               ' Which tag color belongs to {query}?',' Retrieve the color recorded for {query}.')
                    text=' '.join(templates[family].format(box=box,color=colors[box]) for box in order)
                    text+=questions[family].format(query=query)+'\nAnswer:'
                    rows.append({'id':'F'+str(family)+'-'+str(len(rows)%8),'family':'F'+str(family),
                                 'text':text,'assignment':colors.copy(),'record_order':list(order),'query':query,
                                 'label':1 if colors[query]=='red' else -1,
                                 'split':'train' if family<2 else 'heldout'})
    return rows


def tokenize_rows(rows,tokenizer):
    for row in rows:
        encoded=tokenizer(row['text'],add_special_tokens=False,padding=False,truncation=False)
        ids=encoded['input_ids']
        if not 1<=len(ids)<=96:raise ValueError('Prompt exceeds fixed context limit')
        row.update(token_ids=ids,positions=list(range(len(ids))),length=len(ids),final_position=len(ids)-1,
                   decoded_pieces=[tokenizer.decode([i],clean_up_tokenization_spaces=False) for i in ids],
                   attention_mask=encoded['attention_mask'])
    if len({r['token_ids'][-1] for r in rows})!=1:raise ValueError('Final token alignment failed')
    return rows


def checked_logits(logits,torch):
    if tuple(logits.shape)!=(1,1,50304) or not torch.isfinite(logits).all():
        raise ValueError('Invalid final-only logits')
    return logits[0,0].detach().clone()


def run_forward(model,row,budget,torch):
    budget.forward()
    return checked_logits(model(input_ids=torch.tensor([row['token_ids']],dtype=torch.long),
        attention_mask=torch.tensor([row['attention_mask']],dtype=torch.long),use_cache=False,
        logits_to_keep=1,output_attentions=False,output_hidden_states=False).logits,torch)


def capture(model,row,budget,torch,keys=None):
    """Store cloned selected-position vectors; cleanup holds even when a callback fails."""
    wanted=set(keys) if keys is not None else {*(f'r{i}' for i in range(7)),
                    *(f'a{i}' for i in range(6)),*(f'm{i}' for i in range(6)),'h_f'}
    values={};handles=[];n=row['length']
    counts={id(m):(len(m._forward_hooks),len(m._forward_pre_hooks)) for m in model.modules()}
    def put(key,value):
        if key in values:raise ValueError('Duplicate capture key')
        if not isinstance(value,torch.Tensor) or tuple(value.shape)!=(1,n,128):
            raise ValueError('Invalid boundary tensor: '+key)
        if not torch.isfinite(value).all():raise ValueError('Nonfinite boundary: '+key)
        values[key]=value[0,n-1,:].detach().clone()
    def pre(key):
        def hook(module,args):
            if not isinstance(args,tuple) or not args:raise ValueError('Invalid pre-hook input')
            put(key,args[0])
            return None
        return hook
    def post(key,block=False):
        def hook(module,args,result):
            if block:
                if not isinstance(result,tuple) or len(result)!=1:raise ValueError('Expected block tuple')
                result=result[0]
            put(key,result)
            return None
        return hook
    try:
        if 'r0' in wanted:handles.append(model.gpt_neox.layers[0].register_forward_pre_hook(pre('r0')))
        for i,layer in enumerate(model.gpt_neox.layers):
            for key,module,block in [(f'r{i+1}',layer,True),(f'a{i}',layer.post_attention_dropout,False),
                                     (f'm{i}',layer.post_mlp_dropout,False)]:
                if key in wanted:handles.append(module.register_forward_hook(post(key,block)))
        if 'h_f' in wanted:handles.append(model.gpt_neox.final_layer_norm.register_forward_hook(post('h_f')))
        logits=run_forward(model,row,budget,torch)
        if set(values)!=wanted:raise ValueError('Capture coverage failed')
        return values,logits
    finally:
        for handle in handles:handle.remove()
        if any(counts[id(m)]!=(len(m._forward_hooks),len(m._forward_pre_hooks)) for m in model.modules()):
            raise ValueError('Observation hook cleanup failed')


def agreement(a,b,torch):
    return {'max_abs_error':float((a-b).abs().max()),
            'passed':bool(torch.allclose(a,b,atol=ATOL,rtol=RTOL))}


def geometry(a,b=None):
    # Reporting metrics use float64 copies; native captures and model stay float32.
    a=a.double();norm=float(a.norm());result={'norm':norm}
    if b is not None:
        b=b.double();bn=float(b.norm())
        result.update(difference_norm=float((a-b).norm()),
                      cosine=float(a.dot(b)/(norm*bn)) if norm and bn else None,
                      cosine_status='defined' if norm and bn else 'undefined_zero_norm')
    return result


def candidates(logits,tokenizer,torch):
    # Stable descending sort preserves ascending token IDs on exact ties.
    ids=torch.argsort(logits,descending=True,stable=True)[:5].tolist()
    size=len(tokenizer)
    return [{'id':i,'logit':float(logits[i]),
             'decoded_piece':tokenizer.decode([i],clean_up_tokenization_spaces=False) if i<size else None,
             'piece_status':'present' if i<size else 'missing_tokenizer_entry'} for i in ids]


def fit_probe(train,labels,test,torch):
    train=train.double();test=test.double();labels=labels.double()
    mean=train.mean(0);scale=train.std(0,correction=0);scale=torch.where(scale==0,1.,scale)
    a=(train-mean)/scale;b=(test-mean)/scale
    a=torch.cat((a,torch.ones((len(a),1),dtype=torch.float64)),dim=1)
    b=torch.cat((b,torch.ones((len(b),1),dtype=torch.float64)),dim=1)
    penalty=torch.eye(a.shape[1],dtype=torch.float64);penalty[-1,-1]=0
    weights=torch.linalg.solve(a.T@a+penalty,a.T@labels)
    return {'mean':mean.tolist(),'scale':scale.tolist(),'weights':weights[:-1].tolist(),
            'intercept':float(weights[-1]),'train_scores':(a@weights).tolist(),'test_scores':(b@weights).tolist()}


def score_probe(record,train_labels,test_labels,rows):
    predicted=[1 if x>=0 else -1 for x in record['test_scores']]
    train_predicted=[1 if x>=0 else -1 for x in record['train_scores']]
    record.update(predictions=predicted,heldout_correct=sum(a==b for a,b in zip(predicted,test_labels)),
                  heldout_total=16,training_correct=sum(a==b for a,b in zip(train_predicted,train_labels)),
                  per_family_errors={family:[r['id'] for r,p,y in zip(rows,predicted,test_labels)
                                              if r['family']==family and p!=y] for family in ('F2','F3')})
    return record


def interrupted_probe(probe,reason):
    probe.update(status='not_completed',reason=reason)
    probe.setdefault('features',{'r0':[],'r3':[]})
    probe.setdefault('feature_rows',[])
    probe.setdefault('readouts',{key:{'status':'not_started'} for key in ('r0','r3')})
    probe.setdefault('shuffles',[])
    if probe.get('feature_capture_status')!='completed':
        probe['feature_capture_status']='partial' if probe['feature_rows'] else 'not_started'
    for record in probe['readouts'].values():
        if record['status']=='running':record.update(status='interrupted',reason=reason)
    if probe.get('shuffles_status')!='completed':
        probe['shuffles_status']='partial' if probe['shuffles'] else 'not_started'
    if probe.get('pending_shuffle'):
        probe['pending_shuffle'].update(status='interrupted',reason=reason)


def optional_probe(model,rows,budget,torch,output,probe):
    """Persist each completed feature row and readout before advancing stages."""
    probe.update(status='running',stage='capture',feature_capture_status='running',
        feature_dtype='float32',regression_dtype='float64',**{'lambda':1},
        prompt_manifest='prompts.json:optional',features={'r0':[],'r3':[]},feature_rows=[],
        feature_files=[],readouts={key:{'status':'not_started'} for key in ('r0','r3')},
        shuffles=[],shuffle_files=[],shuffles_status='not_started',pending_shuffle=None)
    try:
        for index,row in enumerate(rows):
            values,_=capture(model,row,budget,torch,('r0','r3'))
            vectors={key:values[key].tolist() for key in ('r0','r3')}
            address={'row_index':index,'prompt_id':row['id'],'family':row['family'],'split':row['split'],
                     'token_id':row['token_ids'][-1],'position':row['final_position'],
                     'boundaries':{'r0':'gpt_neox.layers.0/input','r3':'gpt_neox.layers.2/output'},
                     'status':'completed'}
            for key in vectors:probe['features'][key].append(vectors[key])
            probe['feature_rows'].append(address)
            filename='probe-feature-'+str(index+1).zfill(3)+'.json'
            output.json(filename,{'address':address,'vectors':vectors})
            probe['feature_files'].append({'path':filename,'sha256':output.hash(filename)})
        probe['feature_capture_status']='completed';probe['stage']='fit'
        labels=torch.tensor([r['label'] for r in rows],dtype=torch.float64)
        features={key:torch.tensor(vectors,dtype=torch.float32) for key,vectors in probe['features'].items()}
        train_ids=[r['id'] for r in rows[:16]];heldout_ids=[r['id'] for r in rows[16:]]
        probe.update(training_labels=labels[:16].tolist(),heldout_labels=labels[16:].tolist(),
                     training_prompt_ids=train_ids,heldout_prompt_ids=heldout_ids,total=16)
        def majority(values):return 1 if sum(values)>=0 else -1
        train_labels=labels[:16].tolist();test_labels=labels[16:].tolist()
        majority_predictions=[majority(train_labels)]*16
        query_predictions=[majority([r['label'] for r in rows[:16] if r['query']==test['query']]) for test in rows[16:]]
        id_predictions=[majority([r['label'] for r in rows[:16] if r['token_ids'][-1]==test['token_ids'][-1]]) for test in rows[16:]]
        for name,predictions in [('majority',majority_predictions),('query_name',query_predictions),('final_id',id_predictions)]:
            probe[name+'_correct']=sum(a==b for a,b in zip(predictions,test_labels))
            probe[name+'_predictions']=predictions
        for key in ('r0','r3'):
            budget.check();probe['readouts'][key]={'status':'running'}
            result=score_probe(fit_probe(features[key][:16],labels[:16],features[key][16:],torch),
                               train_labels,test_labels,rows[16:])
            result.update(status='completed',training_prompt_ids=train_ids,heldout_prompt_ids=heldout_ids)
            probe['readouts'][key]=result
            output.json('probe-readout-'+key+'.json',result)
        probe['stage']='shuffle';probe['shuffles_status']='running'
        rng=torch.Generator(device='cpu').manual_seed(20261007)
        for trial in range(20):
            budget.check();permutation=torch.randperm(16,generator=rng,device='cpu')
            probe['pending_shuffle']={'trial':trial,'permutation':permutation.tolist(),'status':'running'}
            shuffled=labels[:16][permutation]
            result=fit_probe(features['r3'][:16],shuffled,features['r3'][16:],torch)
            record={'trial':trial,'permutation':permutation.tolist(),'status':'completed',
                    **score_probe(result,shuffled.tolist(),test_labels,rows[16:])}
            probe['shuffles'].append(record);probe['pending_shuffle']=None
            filename='probe-shuffle-'+str(trial+1).zfill(3)+'.json'
            output.json(filename,record)
            probe['shuffle_files'].append({'path':filename,'sha256':output.hash(filename)})
        probe.update(status='completed',stage='completed',shuffles_status='completed')
    except Exception as exc:
        interrupted_probe(probe,type(exc).__name__+': '+str(exc))
        raise


def main():
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--artifacts',type=Path,required=True)
    parser.add_argument('--output',type=Path,required=True)
    parser.add_argument('--probe',action='store_true')
    args=parser.parse_args()
    output=Output(args.output,10_000_000);torch=runtime()
    model,tokenizer,identity=load_specimen(args.artifacts,torch)
    core=tokenize_rows([{'id':'P'+str(i+1),'text':text} for i,text in enumerate(CORE)],tokenizer)
    optional=tokenize_rows(probe_rows(),tokenizer) if args.probe else []
    output.json('prompts.json',{'core':core,'optional':optional})
    output.json('predictions.json',{'recorded_before_forward':True,'r0_equal':'expected_same_final_embedding',
        'm0_equal':'expected_parallel_MLP_reads_same_embedding','attention_later_equal':'not_required',
        'monotonic_norm_or_lens':'not_required','tolerances':{'atol':ATOL,'rtol':RTOL},
        'metrics_dtype':'float64_reporting_copies','claim_limit':'lexical_context_observation_not_causal_semantics'})
    output.json('manifest.json',identity|environment(torch)|{'probe_requested':args.probe,
        'sources':source_hashes(['inspect_activations.py','inspection_common.py','inspect_residual_stream.py']),
        'limits':{'seconds':900,'distinct_prompts':40,'context_tokens':96,'saved_bytes':10_000_000,
                  'forwards':56 if args.probe else 24},'seed':0,'attn_backend':'eager','use_cache':False})
    budget=Budget(900,56 if args.probe else 24);checks=[];records=[];captures={};probe={'status':'not_run'}
    status='completed';failure=None
    try:
        with budget.execution(),torch.inference_mode():
            for row in core:
                before=run_forward(model,row,budget,torch)
                values,observed=capture(model,row,budget,torch)
                after=run_forward(model,row,budget,torch)
                local={'prompt':row['id'],'hooked_vs_baseline':agreement(before,observed,torch),
                       'after_removal_vs_baseline':agreement(before,after,torch),
                       'hooks_removed':True,'addition':{}}
                for i in range(6):
                    local['addition'][str(i)]=agreement((values[f'm{i}']+values[f'a{i}'])+values[f'r{i}'],values[f'r{i+1}'],torch)
                normalized=model.gpt_neox.final_layer_norm(values['r6'])
                local['final_normalization']=agreement(normalized,values['h_f'],torch)
                local['normalized_logits']=agreement(model.embed_out(values['h_f']),observed,torch)
                local['raw_logits']=agreement(model.embed_out(normalized),observed,torch)
                checks.append(local)
                captures[row['id']]={key:{'vector':value.tolist(),'shape':[128],'dtype':'float32',
                    'source_shape':[1,row['length'],128],
                    'module_path':('gpt_neox.layers.0/input' if key=='r0' else
                        'gpt_neox.layers.'+str(int(key[1:])-1)+'/output' if key.startswith('r') else
                        'gpt_neox.layers.'+key[1:]+'.post_attention_dropout' if key.startswith('a') else
                        'gpt_neox.layers.'+key[1:]+'.post_mlp_dropout' if key.startswith('m') else
                        'gpt_neox.final_layer_norm'),
                    'prompt':row['id'],'token_id':row['token_ids'][-1],'position':row['final_position'],
                    'boundary':key,'block_index':int(key[1:])-1 if key.startswith('r') and key!='r0' else
                       (int(key[1:]) if key.startswith(('a','m')) else (0 if key=='r0' else None))}
                                     for key,value in values.items()}
                required=[local[k] for k in ('hooked_vs_baseline','after_removal_vs_baseline','final_normalization','normalized_logits','raw_logits')]+list(local['addition'].values())
                if not all(check['passed'] for check in required):raise ValueError('Numerical invariant failed; interpretation blocked')
                record={'prompt':row['id'],'geometry':{},'lens':{}}
                for i in range(7):
                    key='r'+str(i);record['geometry'][key]=geometry(values[key],values['r'+str(i-1)] if i else None)
                    lens=model.embed_out(model.gpt_neox.final_layer_norm(values[key]))
                    if not torch.isfinite(lens).all():raise ValueError('Nonfinite diagnostic lens')
                    record['lens'][key]=candidates(lens,tokenizer,torch)
                record['normalization_separate']=geometry(values['h_f'],values['r6']);records.append(record)
            structural={key:[agreement(torch.tensor(captures['P1'][key]['vector']),torch.tensor(captures[r['id']][key]['vector']),torch) for r in core]
                        for key in ('r0','m0')}
            checks.append({'same_final_token':structural})
            if not all(c['passed'] for series in structural.values() for c in series):raise ValueError('Structural invariant failed')
            pairs=[]
            for i in range(0,8,2):
                for k in range(7):
                    key='r'+str(k)
                    pairs.append({'red':core[i]['id'],'blue':core[i+1]['id'],'boundary':key,
                        **geometry(torch.tensor(captures[core[i]['id']][key]['vector']),torch.tensor(captures[core[i+1]['id']][key]['vector']))})
            output.json('paired-geometry.json',pairs)
            if args.probe:
                optional_probe(model,optional,budget,torch,output,probe)
    except Exception as exc:
        status='interrupted' if isinstance(exc,TimeoutError) else 'failed';failure=type(exc).__name__+': '+str(exc)
        if args.probe and probe['status']!='completed':interrupted_probe(probe,failure)
    output.json('captures.json',captures);output.json('checks.json',checks)
    output.json('observations.json',records);output.json('probe.json',probe)
    output.json('outcome.json',{'status':status,'failure':failure,'forwards':budget.forwards,
        'elapsed_seconds':time.perf_counter()-budget.started,'core_prompts_completed':len(records),
        'interpretation_status':'available_for_review' if status=='completed' else 'blocked',
        'no_generated_continuations':True,'no_causal_or_semantic_claim':True})
    output.check();print(status,':',budget.forwards,'forwards; saved-data limit checked')
    if status!='completed':raise SystemExit(1)


if __name__=='__main__':main()
