#!/usr/bin/env python3
"""Bounded Pythia-14M CPU observation; download scope and boundaries are fixed."""
import argparse, hashlib, importlib.metadata, json, math, platform, sys
from datetime import datetime, timezone
from pathlib import Path
from urllib.request import urlopen

REPOSITORY = 'EleutherAI/pythia-14m'
REVISION = '94f7c35d5e9f2e9bac8ca839329f505b4d007d5d'
# Publisher sizes, then SHA-256 measured and checked at this immutable revision.
FILES = {
    'config.json': (698, 'f97f966a66c444890ed461fff2a51eefb15d74303df05b948124719f199b0b17'),
    'generation_config.json': (111, '494cbba887cc299273f0c54d7612596918e74c7881a27cbd95e4e7e2464ea090'),
    'special_tokens_map.json': (441, '10b8c8852c1e1f70b54d9aff61728408c28971c0e97a6c5a7b2debbd1d3e9c0c'),
    'tokenizer.json': (2114042, '870f4e2baa6b683221fa52004d5d6f40ab8c9d31961617304b78c910c2c3caf2'),
    'tokenizer_config.json': (4834, 'eee017c5bd133137f45907bd0a6e781e2ccd1a533734b7ed2a2f2f4446659809'),
    'model.safetensors': (28143920, '116a02532db461f91386a5b20f942ff2c8d4de7341e21b55caafc3d7b25f49a1'),
}
VERSIONS = {'torch': '2.9.0', 'transformers': '4.57.1', 'tokenizers': '0.22.1',
            'safetensors': '0.6.2', 'huggingface-hub': '0.35.3'}
TEXT = 'The small red boat crossed the lake.'
TOLERANCE = 1e-6


def artifacts(directory, offline=False):
    directory.mkdir(parents=True, exist_ok=True)
    if any(p.name not in FILES for p in directory.iterdir()):
        raise ValueError('Artifact directory contains an unapproved file; use a dedicated directory')
    for name, (size, expected) in FILES.items():
        path = directory / name
        if path.is_symlink():
            raise ValueError('Artifact symlinks are not supported')
        if not path.exists():
            if offline:
                raise ValueError(f'Missing offline artifact: {name}')
            with urlopen(f'https://huggingface.co/{REPOSITORY}/resolve/{REVISION}/{name}', timeout=90) as response:
                data = response.read(size + 1)
            if len(data) != size or hashlib.sha256(data).hexdigest() != expected:
                raise ValueError(f'Download integrity failure: {name}')
            path.write_bytes(data)
        data = path.read_bytes()
        if len(data) != size or hashlib.sha256(data).hexdigest() != expected:
            raise ValueError(f'Cached artifact integrity failure: {name}')
    return {name: {'bytes': size, 'sha256': sha} for name, (size, sha) in FILES.items()}


def observe(directory):
    for package, expected in VERSIONS.items():
        if importlib.metadata.version(package).split('+')[0] != expected:
            raise ValueError(f'Use the pinned environment: {package}=={expected}')
    import torch
    from transformers import AutoTokenizer, GPTNeoXForCausalLM
    torch.set_num_threads(1)
    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()
    config = model.config
    if (config.num_hidden_layers, config.hidden_size, config.use_parallel_residual) != (6, 128, True):
        raise ValueError('Unexpected model architecture')
    tokenizer = AutoTokenizer.from_pretrained(str(directory), local_files_only=True,
                                              trust_remote_code=False, use_fast=True)
    inputs = tokenizer(TEXT, return_tensors='pt', add_special_tokens=False)
    n = inputs['input_ids'].shape[1]
    expected_shape = (1, n, 128)
    snapshots = []; handles = []
    hook_counts = {id(m): (len(m._forward_hooks), len(m._forward_pre_hooks)) for m in model.modules()}

    def capture():
        states = {}; attention = {}; mlp = {}; block_inputs = {}
        def pre(index):
            def hook(module, args):
                block_inputs[index] = args[0].detach().clone()
            return hook
        def post(index):
            def hook(module, args, output):
                if not isinstance(output, tuple):
                    raise ValueError('GPTNeoXLayer boundary must return a tuple')
                states[index] = output[0].detach().clone()
            return hook
        def tensor(target, index):
            def hook(module, args, output):
                if not isinstance(output, torch.Tensor):
                    raise ValueError('Branch boundary must return a tensor')
                target[index] = output.detach().clone()
            return hook
        try:
            for i, layer in enumerate(model.gpt_neox.layers):
                handles.extend([layer.register_forward_pre_hook(pre(i)),
                                layer.register_forward_hook(post(i)),
                                layer.post_attention_dropout.register_forward_hook(tensor(attention, i)),
                                layer.post_mlp_dropout.register_forward_hook(tensor(mlp, i))])
            handles.append(model.gpt_neox.final_layer_norm.register_forward_hook(tensor(states, 'final_norm')))
            with torch.inference_mode():
                result = model(**inputs, use_cache=False, output_hidden_states=True)
            raw = [block_inputs[0]] + [states[i] for i in range(6)]
            captured = raw + [states['final_norm']] + list(attention.values()) + list(mlp.values())
            if any(tuple(x.shape) != expected_shape for x in captured + list(block_inputs.values())):
                raise ValueError('Unexpected activation shape')
            errors = [float(((mlp[i] + attention[i]) + block_inputs[i] - states[i]).abs().max()) for i in range(6)]
            continuity = [float((block_inputs[i] - raw[i]).abs().max()) for i in range(6)]
            returned_expected = raw[:6] + [states['final_norm']]
            if len(result.hidden_states) != 7:
                raise ValueError('Unexpected hidden_states count')
            returned_errors = [float((a - b).abs().max()) for a, b in zip(result.hidden_states, returned_expected)]
            if max(errors + continuity + returned_errors) > TOLERANCE:
                raise ValueError('Captured boundary identity exceeds declared tolerance')
            return raw, states['final_norm'], attention, mlp, result.logits.detach().clone(), {
                'expected_shape': list(expected_shape), 'parallel_sum_max_abs_errors': errors,
                'input_continuity_max_abs_errors': continuity, 'returned_hidden_states_max_abs_errors': returned_errors,
                'returned_hidden_states_labels': ['initial'] + [f'raw_block_{i}' for i in range(1,6)] + ['final_norm'],
                'raw_block_6_is_captured_separately': True, 'absolute_tolerance': TOLERANCE}
        finally:
            for handle in handles: handle.remove()
            handles.clear()

    # Unhooked run establishes observation hooks did not alter model output.
    with torch.inference_mode():
        unhooked = model(**inputs, use_cache=False).logits.detach().clone()
    first = capture(); second = capture()
    repeat_errors = [float((a - b).abs().max()) for a,b in zip(first[0]+[first[1]],second[0]+[second[1]])]
    repeat_errors += [float((first[j][i]-second[j][i]).abs().max()) for j in (2,3) for i in range(6)]
    repeat_errors.append(float((first[4]-second[4]).abs().max()))
    mutation_error = float((first[4] - unhooked).abs().max())
    hooks_removed = all(hook_counts[id(m)] == (len(m._forward_hooks),len(m._forward_pre_hooks)) for m in model.modules())
    if max(repeat_errors + [mutation_error]) > TOLERANCE or not hooks_removed:
        raise ValueError('Repeatability or nonmutating-hook check failed')
    raw, normalized = first[:2]
    selected = n - 1  # A recorded actual token position, not a presumed word.
    def metric(x, previous=None):
        v=x[0,selected]; norm=float(torch.linalg.vector_norm(v)); result={'norm':norm}
        if previous is not None:
            p=previous[0,selected]; pn=float(torch.linalg.vector_norm(p))
            result['difference_norm']=float(torch.linalg.vector_norm(v-p))
            result['cosine']=float(torch.dot(v,p)/(norm*pn)) if norm and pn else None
        return result
    ids=inputs['input_ids'][0].tolist()
    measurements={'initial':metric(raw[0])}
    for i in range(1,7):measurements[f'raw_block_{i}']=metric(raw[i],raw[i-1])
    measurements['final_normalization_separate']=metric(normalized,raw[-1])
    checks=first[5] | {'repeat_max_abs_error':max(repeat_errors), 'hooked_vs_unhooked_logits_max_abs_error':mutation_error,
                       'hooks_removed':hooks_removed, 'training':model.training, 'dtype':str(model.dtype),
                       'device':str(model.device), 'attention_backend':config._attn_implementation, 'use_cache':False}
    # Preserve actual full captures; output stays in the learner's local directory.
    tensors={'initial':raw[0].tolist(),'raw_blocks':[x.tolist() for x in raw[1:]],
             'final_normalized':normalized.tolist(), 'attention_updates':[first[2][i].tolist() for i in range(6)],
             'mlp_updates':[first[3][i].tolist() for i in range(6)]}
    tokens=[{'position':i,'id':ident,'display_piece':tokenizer.convert_ids_to_tokens(ident),
             'decoded_alone':tokenizer.decode([ident],clean_up_tokenization_spaces=False)} for i,ident in enumerate(ids)]
    return {'text':TEXT,'tokens':tokens,'selected_position':selected,'measurements':measurements,'checks':checks}, tensors


def main():
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--artifacts',type=Path,default=Path('pythia-artifacts'))
    parser.add_argument('--output',type=Path,default=Path('residual-run'))
    parser.add_argument('--offline',action='store_true')
    parser.add_argument('--plan',action='store_true',help='Print exact file scope/bytes; do not download or run')
    args=parser.parse_args()
    if args.plan:
        print(json.dumps({'repository':REPOSITORY,'revision':REVISION,'files':FILES,'total_bytes':sum(x[0] for x in FILES.values())},indent=2));return
    if args.output.exists() and any(args.output.iterdir()):raise ValueError('Output directory must be empty')
    hashes=artifacts(args.artifacts,args.offline)
    result,tensors=observe(args.artifacts)
    manifest={'repository':REPOSITORY,'revision':REVISION,'artifacts':hashes,
              'python':platform.python_version(),'packages':{k:importlib.metadata.version(k) for k in VERSIONS},
              'download_client':'CPython urllib.request '+platform.python_version(),
              'platform':platform.system()+' '+platform.machine(),'timestamp_utc':datetime.now(timezone.utc).isoformat(),
              'source_reference':'https://github.com/huggingface/transformers/blob/v4.57.1/src/transformers/models/gpt_neox/modeling_gpt_neox.py',
              'license':'Apache-2.0 (publisher repository metadata)'}
    args.output.mkdir(parents=True,exist_ok=True)
    for name,data in [('manifest',manifest),('observations',result),('captures',tensors)]:
        (args.output/(name+'.json')).write_text(json.dumps(data,indent=2,allow_nan=False)+'\n')
    print(json.dumps(result['checks'],indent=2));print(f'Saved actual CPU observations for {len(result["tokens"])} positions.')

if __name__=='__main__':main()
