#!/usr/bin/env python3
"""Discovery and prediction-gated held-out MLP observation, pinned CPU specimen."""
import argparse, hashlib, importlib.metadata, json, platform
from pathlib import Path
import inspect_residual_stream as specimen

PAIRS=[('x = 2 + 3\nprint(x)','Mira added two and three and wrote the total.'),
 ('items = [1, 2, 3]\nlen(items)','Mira counted three items on the desk.'),
 ("names = ['Ada', 'Lin']\nnames[0]",'Mira chose Ada from a list of two names.'),
 ("word = 'blue'\nword.upper()",'Mira rewrote the word blue in capital letters.')]
SUFFIX='\nSummary:'
TOLERANCE=1e-6

def select(records):
 contrasts=[sum(records[2*k]['blocks']['2']['target_h'][j]-records[2*k+1]['blocks']['2']['target_h'][j] for k in range(2))/2 for j in range(512)]
 indices=sorted(range(512),key=lambda j:(-abs(contrasts[j]),j))[:3]
 return [{'index':j,'mean_discovery_contrast':contrasts[j],'sign':(contrasts[j]>0)-(contrasts[j]<0)} for j in indices]

def validate_predictions(predictions,selection):
 if predictions.get('selected_indices')!=[s['index'] for s in selection]:raise ValueError('Prediction indices must match the frozen discovery selection')
 directions=predictions.get('directions')
 if not isinstance(directions,list) or len(directions)!=3 or any(type(x) is not int or x not in (-1,0,1) for x in directions):raise ValueError('Record each predicted direction as -1, 0 or 1')
 if not isinstance(predictions.get('reasoning'),str) or not predictions['reasoning'].strip():raise ValueError('Record reasoning before revealing held-out responses')
 if not isinstance(predictions.get('structural_predictions'),str) or not predictions['structural_predictions'].strip():raise ValueError('Record block0/block2 structural predictions')
 return predictions

class Observer:
 def __init__(self,directory):
  for package,version in specimen.VERSIONS.items():
   if importlib.metadata.version(package).split('+')[0]!=version:raise ValueError('Use the pinned Lab05 environment')
  import torch
  from transformers import AutoTokenizer,GPTNeoXForCausalLM
  self.torch=torch;torch.set_num_threads(1)
  self.model=GPTNeoXForCausalLM.from_pretrained(str(directory),local_files_only=True,trust_remote_code=False,use_safetensors=True,dtype=torch.float32,attn_implementation='eager').to('cpu').eval()
  c=self.model.config
  if (c.num_hidden_layers,c.hidden_size,c.intermediate_size,c.use_parallel_residual)!=(6,128,512,True):raise ValueError('Unexpected model config')
  self.tokenizer=AutoTokenizer.from_pretrained(str(directory),local_files_only=True,trust_remote_code=False,use_fast=True)
  self.prompts=[text+SUFFIX for pair in PAIRS for text in pair]
  self.inputs=[self.tokenizer(text,return_tensors='pt',add_special_tokens=False,truncation=False,padding=False) for text in self.prompts]
  ids=[int(inputs['input_ids'][0,-1]) for inputs in self.inputs]
  if len(set(ids))!=1:raise ValueError('Common suffix final token is not aligned')
  self.final_id=ids[0]

 def run(self,index):
  torch=self.torch;inputs=self.inputs[index];n=inputs['input_ids'].shape[1];captured={};handles=[]
  initial_hooks={id(m):(len(m._forward_hooks),len(m._forward_pre_hooks)) for m in self.model.modules()}
  def post(block,key):
   def hook(module,args,output):
    if not isinstance(output,torch.Tensor):raise ValueError('MLP capture must be a tensor')
    captured[block][key]=output.detach().clone()
    return None
   return hook
  def pre(block):
   def hook(module,args):
    if not isinstance(args[0],torch.Tensor):raise ValueError('Output-affine input must be a tensor')
    captured[block]['h']=args[0].detach().clone()
    return None
   return hook
  with torch.inference_mode():unhooked=self.model(**inputs,use_cache=False).logits.detach().clone()
  try:
   for block in (0,2):
    mlp=self.model.gpt_neox.layers[block].mlp;captured[block]={}
    handles.extend([mlp.dense_h_to_4h.register_forward_hook(post(block,'z')),
     mlp.dense_4h_to_h.register_forward_pre_hook(pre(block)),mlp.register_forward_hook(post(block,'delta'))])
   with torch.inference_mode():hooked=self.model(**inputs,use_cache=False).logits.detach().clone()
  finally:
   for handle in handles:handle.remove()
  clean=all(initial_hooks[id(m)]==(len(m._forward_hooks),len(m._forward_pre_hooks)) for m in self.model.modules())
  invariant=float((hooked-unhooked).abs().max());blocks={}
  with torch.inference_mode():
   for block,c in captured.items():
    mlp=self.model.gpt_neox.layers[block].mlp
    assert c['z'].shape==c['h'].shape==(1,n,512) and c['delta'].shape==(1,n,128)
    activation_error=float((mlp.act(c['z'])-c['h']).abs().max())
    reconstruction_error=float((torch.nn.functional.linear(c['h'],mlp.dense_4h_to_h.weight,mlp.dense_4h_to_h.bias)-c['delta']).abs().max())
    if max(activation_error,reconstruction_error,invariant)>TOLERANCE or not clean:raise ValueError('MLP boundary or nonmutating-hook check failed')
    blocks[str(block)]={'shape_z_h':[1,n,512],'shape_delta':[1,n,128],'activation_max_abs_error':activation_error,
      'affine_reconstruction_max_abs_error':reconstruction_error,'target_h':c['h'][0,-1].tolist(),
      'full_z':c['z'].tolist(),'full_h':c['h'].tolist(),'full_delta':c['delta'].tolist()}
  ids=inputs['input_ids'][0].tolist()
  return {'prompt_index':index,'pair':index//2+1,'kind':'code' if index%2==0 else 'prose','text':self.prompts[index],
   'tokens':[{'position':i,'id':ident,'display_piece':self.tokenizer.convert_ids_to_tokens(ident),
              'decoded_alone':self.tokenizer.decode([ident],clean_up_tokenization_spaces=False)} for i,ident in enumerate(ids)],
   'target_position':n-1,'target_id':ids[-1],'blocks':blocks,'hooked_vs_unhooked_logits_max_abs_error':invariant,'hooks_removed':clean}

def max_target_difference(records,block):
 reference=records[0]['blocks'][str(block)]['target_h']
 return max(abs(a-b) for record in records[1:] for a,b in zip(reference,record['blocks'][str(block)]['target_h']))

def write(path,data):path.write_text(json.dumps(data,indent=2,allow_nan=False)+'\n')

def main():
 p=argparse.ArgumentParser(description=__doc__);p.add_argument('phase',choices=['discover','reveal','challenge'])
 p.add_argument('--artifacts',type=Path,default=Path('pythia-artifacts'));p.add_argument('--output',type=Path,required=True)
 p.add_argument('--discovery',type=Path);p.add_argument('--predictions',type=Path)
 p.add_argument('--structural-predictions',type=Path);p.add_argument('--contexts',type=Path);args=p.parse_args()
 if args.output.exists() and any(args.output.iterdir()):raise ValueError('Output directory must be empty')
 specimen.artifacts(args.artifacts,offline=True) # Download explicitly through Lab05 first.
 selection=None;predictions=None
 if args.phase!='discover':
  if args.discovery is None or args.predictions is None:raise ValueError('Reveal requires discovery directory and written predictions')
  saved=json.loads((args.discovery/'discovery.json').read_text());selection=select(saved['records'])
  if saved['selection']!=selection:raise ValueError('Discovery selection changed')
  if saved['revision']!=specimen.REVISION:raise ValueError('Wrong discovery checkpoint')
  predictions=validate_predictions(json.loads(args.predictions.read_text()),selection)
 structure=None
 if args.phase=='discover':
  if args.structural_predictions is None:raise ValueError('Save structural predictions before discovery')
  structure=args.structural_predictions.read_text().strip()
  if not structure:raise ValueError('Structural predictions must be nonempty')
 args.output.mkdir(parents=True,exist_ok=True)
 if structure is not None:(args.output/'structural-predictions.txt').write_text(structure+'\n')
 if predictions is not None:
  write(args.output/'frozen_predictions.json',predictions) # Durable before any held-out inference.
  write(args.output/'frozen_selection.json',selection)
 observer=Observer(args.artifacts)
 indices=range(4) if args.phase=='discover' else range(4,8)
 if args.phase=='challenge':
  if args.contexts is None:raise ValueError('Challenge requires a JSON file with two context strings')
  contexts=json.loads(args.contexts.read_text())
  if not isinstance(contexts,list) or len(contexts)!=2 or any(not isinstance(x,str) or not x.strip() for x in contexts):raise ValueError('Supply exactly two nonempty context strings')
  for context in contexts:
   text=context+SUFFIX;inputs=observer.tokenizer(text,return_tensors='pt',add_special_tokens=False,truncation=False,padding=False)
   if int(inputs['input_ids'][0,-1])!=observer.final_id or inputs['input_ids'].shape[1]>256:raise ValueError('Challenge target must align and contain at most 256 tokens')
   observer.prompts.append(text);observer.inputs.append(inputs)
  write(args.output/'frozen_contexts.json',contexts)
  indices=range(8,10)
 records=[observer.run(i) for i in indices]
 if args.phase=='challenge':
  for i,record in enumerate(records):record['kind']='challenge_first' if i==0 else 'challenge_second'
 block0=max_target_difference(records,0)
 checks={'absolute_tolerance':TOLERANCE,'block0_same_final_id_max_abs_difference':block0,
         'block0_within_tolerance':block0<=TOLERANCE,'block2_max_abs_difference':max_target_difference(records,2)}
 if args.phase=='discover':
  repeat=observer.run(0)
  errors=[]
  for block in ('0','2'):
   for key in ('full_z','full_h','full_delta'):
    errors += [abs(a-b) for row_a,row_b in zip(records[0]['blocks'][block][key][0],repeat['blocks'][block][key][0]) for a,b in zip(row_a,row_b)]
  checks['discovery_repeat_max_abs_error']=max(errors)
  if max(errors)>TOLERANCE:raise ValueError('Discovery repeatability failed')
  selection=select(records)
  write(args.output/'discovery.json',{'revision':specimen.REVISION,'records':records,'selection':selection,'checks':checks})
  write(args.output/'predictions-template.json',{'selected_indices':[s['index'] for s in selection],
    'directions':[None,None,None],'reasoning':'','structural_predictions':''})
  print('Discovery saved. Fill predictions-template.json in a separate file before the reveal phase.')
 else:
  outcomes=[]
  for j,item in enumerate(selection):
   index=item['index'];pairs=[]
   for k in range(len(records)//2):
    code=records[2*k]['blocks']['2']['target_h'][index];prose=records[2*k+1]['blocks']['2']['target_h'][index];diff=code-prose
    pairs.append({'pair':records[2*k]['pair'],'code':code,'prose':prose,'difference':diff,'observed_sign':(diff>0)-(diff<0),
                  'predicted_sign':predictions['directions'][j],'sign_matches':((diff>0)-(diff<0))==predictions['directions'][j]})
   outcomes.append({'index':index,'pairs':pairs})
  write(args.output/('heldout.json' if args.phase=='reveal' else 'challenge.json'),{'revision':specimen.REVISION,'records':records,'frozen_selection':selection,'outcomes':outcomes,'checks':checks})
  print(('Held-out' if args.phase=='reveal' else 'Challenge')+' observations saved, including every sign failure. No reselection occurred.')
 manifest={'revision':specimen.REVISION,'phase':args.phase,'python':platform.python_version(),
  'packages':{name:importlib.metadata.version(name) for name in specimen.VERSIONS},'download_client':'CPython urllib.request '+platform.python_version(),
  'attention_backend':'eager','device':'cpu','dtype':'float32','eval':True,'use_cache':False,'add_special_tokens':False,
  'truncation':False,'padding':False,'target_token_id':observer.final_id,'artifacts':specimen.FILES}
 if args.phase!='discover':manifest['discovery_sha256']=hashlib.sha256((args.discovery/'discovery.json').read_bytes()).hexdigest()
 write(args.output/'manifest.json',manifest)
 print(json.dumps(checks,indent=2))

if __name__=='__main__':main()
