#!/usr/bin/env python3
"""Lab 12: bounded offline CPU sampling, matched cache checks, and local timings.

Requires the sibling inspect_residual_stream.py and its pinned environment.
Ordinary EOS-stopped completions and deliberately post-EOS profiling are separate.
"""
import argparse, contextlib, hashlib, importlib.metadata, inspect, json, logging
import platform, random, signal, statistics, sys, time, warnings
from pathlib import Path
from inspect_residual_stream import artifacts, FILES, REPOSITORY, REVISION, VERSIONS

PROMPTS = {'P1': 'The small red boat crossed the lake.',
           'P2': 'A notebook lay beside the window.'}
ATOL, RTOL = 1e-5, 1e-4
MAX_TRIALS, MAX_SELECTIONS, MAX_FORWARDS = 50, 800, 816

def safe_path(path):
    if any(p.is_symlink() for p in (path, *path.parents)):
        raise ValueError('Paths must not traverse symlinks')

def save(path, data):
    path.write_text(json.dumps(data, indent=2, allow_nan=False) + '\n', encoding='utf-8')

class Budget:
    def __init__(self):
        self.deadline = time.perf_counter() + 900
        self.forwards = self.trials = self.selections = 0
    def forward(self, *args):
        if time.perf_counter() >= self.deadline: raise TimeoutError('900-second execution limit')
        if self.forwards >= MAX_FORWARDS: raise ValueError('Forward budget exhausted')
        self.forwards += 1
    @contextlib.contextmanager
    def execution(self):
        if not hasattr(signal, 'setitimer'):
            raise RuntimeError('This runner requires POSIX wall-clock timers')
        def expired(signum, frame): raise TimeoutError('900-second execution limit')
        self.global_handler = expired
        previous = signal.signal(signal.SIGALRM, expired)
        signal.setitimer(signal.ITIMER_REAL, max(0.001, self.deadline-time.perf_counter()))
        try: yield
        finally:
            signal.setitimer(signal.ITIMER_REAL, 0)
            signal.signal(signal.SIGALRM, previous)
    @contextlib.contextmanager
    def trial(self):
        remaining = self.deadline - time.perf_counter()
        if remaining <= 0: raise TimeoutError('900-second execution limit')
        self.trials += 1
        if self.trials > MAX_TRIALS: raise ValueError('Trial budget exhausted')
        if not hasattr(signal, 'setitimer'):
            raise RuntimeError('This runner requires POSIX wall-clock timers; no unbounded fallback')
        def expired(signum, frame): raise TimeoutError('Trial or cumulative deadline exceeded')
        previous = signal.signal(signal.SIGALRM, expired)
        signal.setitimer(signal.ITIMER_REAL, min(30, remaining))
        try: yield
        finally:
            signal.setitimer(signal.ITIMER_REAL, 0)
            signal.signal(signal.SIGALRM, previous)
            if getattr(self, 'global_handler', None) is previous:
                signal.setitimer(signal.ITIMER_REAL, max(0.001, self.deadline-time.perf_counter()))
    def selected(self, count):
        self.selections += count
        if count > 16 or self.selections > MAX_SELECTIONS: raise ValueError('Selection budget exhausted')

def storage_bytes(tensors):
    storages = {}
    for x in tensors:
        storage = x.untyped_storage()
        storages[(str(x.device), storage.data_ptr())] = storage.nbytes()
    return sum(storages.values())

def cache_record(cache):
    if cache is None: return None
    tensors = [x for pair in cache for x in pair]
    return {'logical_length': cache.get_seq_length(), 'storage_bytes': storage_bytes(tensors),
            'tensors': [{'shape': list(x.shape), 'dtype': str(x.dtype),
                         'device': str(x.device)} for x in tensors]}

def config_for(model, tokenizer, condition):
    from transformers import GenerationConfig
    # Fresh library defaults avoid inheriting checkpoint-specific penalties.
    kwargs = dict(max_new_tokens=16, num_beams=1, num_return_sequences=1,
                  use_cache=True, do_sample=condition != 'G', repetition_penalty=1.0,
                  eos_token_id=model.config.eos_token_id, pad_token_id=tokenizer.eos_token_id,
                  bos_token_id=model.config.bos_token_id, min_length=0, min_new_tokens=None,
                  forced_bos_token_id=None, forced_eos_token_id=None, bad_words_ids=None,
                  sequence_bias=None, suppress_tokens=None, begin_suppress_tokens=None,
                  watermarking_config=None, output_scores=False, output_logits=False,
                  output_attentions=False, output_hidden_states=False, return_dict_in_generate=False)
    if condition != 'G':
        kwargs.update(temperature=0.5 if condition == 'S2' else 1.0,
                      top_p=0.8 if condition == 'S3' else 1.0, top_k=0)
    return GenerationConfig(**kwargs)

REFERENCE_TRIAL = 'ordinary-P1-G-0'

def reference_trial(records):
    designated=next((r for r in records if r.get('trial_id')==REFERENCE_TRIAL),None)
    if designated is None or designated.get('status')!='complete':return None
    return designated

@contextlib.contextmanager
def resolved_config(model):
    """Observe the configuration object returned by the pinned generation resolver."""
    original=model._prepare_generation_config; captured=[]
    def prepare(*args,**kwargs):
        result=original(*args,**kwargs);captured.append(result[0]);return result
    model._prepare_generation_config=prepare
    try:yield captured
    finally:del model._prepare_generation_config

def configuration_record(config):
    public={k:v for k,v in config.to_dict().items() if not k.startswith('_')}
    # The resolver's special-token tensors are derived fields outside public to_dict.
    public['resolved_special_token_ids']={key:getattr(config,key).tolist() for key in
        ('_eos_token_tensor','_pad_token_tensor','_bos_token_tensor') if getattr(config,key,None) is not None}
    return public

class SelectionRecorder:
    """Transformers sends the prompt first, then each selected token."""
    def __init__(self,budget):self.budget=budget;self.prompt=None;self.ids=[]
    def put(self,value):
        values=value.reshape(-1).tolist()
        if self.prompt is None:self.prompt=values;return
        self.budget.selected(len(values));self.ids.extend(values)
    def end(self):pass

@contextlib.contextmanager
def library_warnings():
    messages=[]
    class Collector(logging.Handler):
        def emit(self,record):messages.append(record.getMessage())
    logger=logging.getLogger('transformers');handler=Collector(logging.WARNING)
    logger.addHandler(handler)
    try:yield messages
    finally:logger.removeHandler(handler)

def fixed_loop(model, inputs, cached, budget):
    """Exactly sixteen greedy selections; EOS remains a possible unsuppressed choice."""
    import torch
    ids = inputs['input_ids']; mask = inputs['attention_mask']; past = None
    positions = torch.arange(ids.shape[1]); timestamps = []; prefill_end = None; logits_refs = []
    started = time.perf_counter()
    selected_refs=[]
    try:
        for step in range(16):
            result = model(input_ids=ids if not cached or step == 0 else ids[:, -1:],
                           attention_mask=mask, past_key_values=past if cached else None,
                           position_ids=positions.unsqueeze(0), cache_position=positions,
                           use_cache=cached, logits_to_keep=1,
                           output_attentions=False, output_hidden_states=False)
            if step == 0: prefill_end = time.perf_counter()
            logits_refs.append(result.logits)
            selected = result.logits[:, -1].argmax(dim=-1, keepdim=True)
            budget.selected(1)
            timestamps.append(time.perf_counter())
            selected_refs.append(selected)
            if step < 15:
                ids = torch.cat((ids, selected), dim=1)
                mask = torch.cat((mask, torch.ones_like(selected)), dim=1)
                positions = torch.arange(ids.shape[1]) if not cached else torch.tensor([ids.shape[1]-1])
                past = result.past_key_values if cached else None
    except Exception as exc:
        # Materialize only after the interrupted timing region has ended.
        generated=[int(x.item()) for x in selected_refs]
        exc.partial_trial={'generated_ids':generated,'output_length':len(generated),
            'full_ids':inputs['input_ids'][0].tolist()+generated,
            'selection_timestamps_relative':[t-started for t in timestamps],
            'prefill_forward_seconds':prefill_end-started if prefill_end is not None else None,
            'ttft_seconds':timestamps[0]-started if timestamps else None,
            'generation_seconds':None,'final_cache':None}
        raise
    # Transfers, finite checks, cache inspection and decoding happen after the clock.
    full = torch.cat((ids, selected), dim=1)[0].tolist()
    if any(x.shape != (1, 1, model.config.vocab_size) or not torch.isfinite(x).all() for x in logits_refs):
        raise ValueError('Invalid final-only timing logits')
    generated = full[inputs['input_ids'].shape[1]:]
    return {'generated_ids': generated, 'full_ids': full, 'output_length': len(generated),
            'prefill_forward_seconds': prefill_end-started, 'ttft_seconds': timestamps[0]-started,
            'selection_timestamps_relative': [t-started for t in timestamps],
            'generation_seconds': timestamps[-1]-started,
            'mean_post_first_interval_seconds': (timestamps[-1]-timestamps[0])/15,
            'throughput_tokens_per_second_including_prefill': 16/(timestamps[-1]-started),
            'final_cache': cache_record(result.past_key_values) if cached else None}

def matched_prefixes(model, inputs, targets, budget):
    import torch
    rows = []; past = None; prompt = inputs['input_ids']; n = prompt.shape[1]
    targets = targets[:8]
    for step in range(len(targets)):
        ids = torch.cat((prompt, torch.tensor([targets[:step]], dtype=torch.long)), dim=1)
        mask = torch.ones_like(ids); positions = torch.arange(ids.shape[1])
        a = model(input_ids=ids, attention_mask=mask, position_ids=positions.unsqueeze(0),
                  cache_position=positions, use_cache=False, logits_to_keep=1,
                  output_attentions=False, output_hidden_states=False).logits[0, -1]
        cp = positions if step == 0 else positions[-1:]
        result = model(input_ids=ids if step == 0 else ids[:, -1:], attention_mask=mask,
                       position_ids=cp.unsqueeze(0), cache_position=cp, past_key_values=past,
                       use_cache=True, logits_to_keep=1, output_attentions=False, output_hidden_states=False)
        b = result.logits[0, -1]; past = result.past_key_values
        if a.shape != b.shape or a.numel() != model.config.vocab_size:
            raise ValueError('Matched-prefix vocabulary shape mismatch')
        if not torch.isfinite(a).all() or not torch.isfinite(b).all():
            raise ValueError('Nonfinite matched-prefix logits')
        error = (a-b).abs(); failed = error > ATOL + RTOL*a.abs()
        av, ai = a.topk(2); bv, bi = b.topk(2)
        cache = cache_record(past)
        if cache['logical_length'] != n+step: raise ValueError('Cache length mismatch')
        rows.append({'step': step, 'prefix_ids': ids[0].tolist(), 'next_fixed_id': targets[step],
                     'attention_mask': mask[0].tolist(), 'uncached_position_ids': positions.tolist(),
                     'cached_position_ids': cp.tolist(), 'uncached_logits': a.tolist(), 'cached_logits': b.tolist(),
                     'max_abs_difference': float(error.max()), 'failure_fraction': float(failed.float().mean()),
                     'criterion_passed': not bool(failed.any()), 'uncached_argmax': int(ai[0]),
                     'cached_argmax': int(bi[0]), 'uncached_top_two_margin': float(av[0]-av[1]),
                     'cached_top_two_margin': float(bv[0]-bv[1]), 'cache': cache})
    return {'absolute_tolerance': ATOL, 'relative_tolerance': RTOL, 'reference': 'uncached',
            'rows': rows, 'forwards': 2*len(targets), 'cache_reuse_exercised': len(targets)>1,
            'status': 'complete' if rows else 'unavailable: no successful greedy continuation'}

def run(directory, out, budget):
    import torch
    from transformers import AutoTokenizer, GPTNeoXForCausalLM
    torch.set_num_threads(1); torch.set_num_interop_threads(1)
    loading_start = time.perf_counter()
    with warnings.catch_warnings(record=True) as caught, library_warnings() as loading_messages:
        warnings.simplefilter('always')
        tokenizer = AutoTokenizer.from_pretrained(str(directory), local_files_only=True,
                                                  trust_remote_code=False, use_fast=True)
        tokenizer_loaded = time.perf_counter()
        model, load_info = GPTNeoXForCausalLM.from_pretrained(
            str(directory), local_files_only=True, trust_remote_code=False, use_safetensors=True,
            dtype=torch.float32, attn_implementation='eager', output_loading_info=True)
        model.to('cpu').eval()
    loaded = time.perf_counter()
    if (model.config.num_hidden_layers, model.config.hidden_size, model.config.num_attention_heads) != (6,128,4):
        raise ValueError('Unexpected pinned architecture')
    if any(load_info.get(k) for k in ('missing_keys','unexpected_keys','mismatched_keys','error_msgs')):
        save(out/'loading-failure.json', load_info); raise ValueError('Unexpected checkpoint loading report')
    if 'logits_to_keep' not in inspect.signature(model.forward).parameters:
        raise ValueError('Installed model lacks final-position projection contract')
    if any(p.device.type != 'cpu' or p.dtype != torch.float32 for p in model.parameters()):
        raise ValueError('Required CPU FP32 parameters not observed')
    env = {'python': platform.python_version(), 'platform': platform.platform(),
           'architecture': platform.machine(), 'packages': {k:importlib.metadata.version(k) for k in VERSIONS},
           'all_installed_packages': sorted(f'{x.metadata["Name"]}=={x.version}' for x in importlib.metadata.distributions()),
           'cpu_threads': torch.get_num_threads(), 'interop_threads': torch.get_num_interop_threads(),
           'model_class': type(model).__name__, 'parameter_count': sum(p.numel() for p in model.parameters()),
           'unique_parameter_storage_bytes': storage_bytes(model.parameters()),
           'model_dtype': str(model.dtype), 'device': str(model.device), 'training': model.training,
           'attention_backend': model.config._attn_implementation, 'loading_info': load_info,
           'warnings': [str(w.message) for w in caught], 'library_warnings':loading_messages,
           'float32_matmul_precision':torch.get_float32_matmul_precision(),
           'autocast_cpu_enabled':torch.is_autocast_enabled('cpu'),
           'mkldnn_available':torch.backends.mkldnn.is_available(),'mkldnn_enabled':torch.backends.mkldnn.enabled,
           'torch_build_configuration':torch.__config__.show(),
           'tokenizer_loading_seconds': tokenizer_loaded-loading_start,
           'model_loading_seconds': loaded-tokenizer_loaded, 'logits_to_keep': 1,
           'timing_diagnostic_outputs': False, 'timing_overhead': 'budget hook, timestamps, list append and selection bookkeeping; same code for both cache paths', 'unmeasured_memory': ['RSS', 'peak RSS', 'activations', 'workspace', 'allocator retention', 'runtime libraries', 'loading copies']}
    save(out/'environment.json', env)
    encoded = {}; prompt_records = {}
    for name, text in {**PROMPTS, 'P1x8': ' '.join([PROMPTS['P1']]*8)}.items():
        started = time.perf_counter()
        inputs = tokenizer(text, return_tensors='pt', add_special_tokens=False)
        duration = time.perf_counter()-started; ids = inputs['input_ids'][0].tolist()
        encoded[name] = inputs
        prompt_records[name] = {'source_text': text, 'serialized_text': text, 'utf8_hex': text.encode().hex(),
                                'ids': ids, 'attention_mask': inputs['attention_mask'][0].tolist(),
                                'pieces': tokenizer.convert_ids_to_tokens(ids), 'input_length': len(ids),
                                'tokenization_seconds': duration, 'add_special_tokens': False}
    save(out/'inputs.json', {'prompts': prompt_records, 'special_tokens': tokenizer.special_tokens_map,
                            'eos_id': model.config.eos_token_id, 'existing_pad_id': tokenizer.eos_token_id,
                            'position_rule': 'absolute arange(full prefix); cached decode receives final position only'})
    handles = [model.register_forward_pre_hook(budget.forward)]
    records = []; first_logits = {}; configs = {}; ordinary = {}
    def log(record):
        records.append(record)
        with (out/'trials.jsonl').open('a', encoding='utf-8') as f:
            f.write(json.dumps(record, allow_nan=False)+'\n')
    try:
        with torch.inference_mode():
            for name in PROMPTS:
                for condition in ('G', 'S1', 'S2', 'S3'):
                    cfg = config_for(model, tokenizer, condition); configs[condition] = cfg.to_dict()
                    seeds = [None]*3 if condition == 'G' else [17,23,31]
                    if condition == 'S1': seeds.append(17)
                    for rep, seed in enumerate(seeds):
                        row = {'trial_id':f'ordinary-{name}-{condition}-{rep}', 'phase':'ordinary', 'prompt': name, 'condition': condition, 'repeat': rep,
                               'seed': seed, 'stopping_rule':'normal EOS or sixteen new selections',
                               'requested_generation_configuration': cfg.to_dict()}
                        capture = []; ws=[]; resolved=[]; library_messages=[]
                        recorder=SelectionRecorder(budget)
                        def capture_first(module, args, output):
                            if not capture: capture.append(output.logits[0,-1].detach())
                        hook = model.register_forward_hook(capture_first) if rep == 0 else None
                        before = budget.forwards
                        try:
                            if seed is not None: torch.manual_seed(seed); random.seed(seed)
                            with warnings.catch_warnings(record=True) as ws, library_warnings() as library_messages, budget.trial(), resolved_config(model) as resolved:
                                warnings.simplefilter('always')
                                output = model.generate(**encoded[name], generation_config=cfg, logits_to_keep=1, streamer=recorder)
                            if len(resolved)!=1:raise ValueError('Unexpected generation configuration resolution count')
                            row['effective_generation_configuration']=configuration_record(resolved[0])
                            ids = output[0].tolist(); generated = ids[len(prompt_records[name]['ids']):]
                            if generated!=recorder.ids:raise ValueError("Streamed selections differ from returned IDs")
                            if not 1 <= len(generated) <= 16: raise ValueError('Unexpected generation length')
                            if capture:
                                if capture[0].numel()!=model.config.vocab_size or not torch.isfinite(capture[0]).all():
                                    raise ValueError('Invalid raw first-step logits')
                                first_logits[name+'-'+condition] = capture[0].tolist()
                            began = time.perf_counter()
                            text = tokenizer.decode(generated, skip_special_tokens=False, clean_up_tokenization_spaces=False)
                            hidden = tokenizer.decode(generated, skip_special_tokens=True, clean_up_tokenization_spaces=False)
                            row.update(status='complete', prompt_ids=prompt_records[name]['ids'], generated_ids=generated,
                                       full_ids=ids, output_length=len(generated), continuation_with_special_tokens=text,
                                       continuation_without_special_tokens=hidden, decoding_seconds=time.perf_counter()-began,
                                       ended_on_eos=generated[-1]==model.config.eos_token_id,
                                       warnings=[str(w.message) for w in ws])
                            ordinary.setdefault((name,condition), []).append(generated)
                        except Exception as exc:
                            row.update(status='failed', error=f'{type(exc).__name__}: {exc}')
                        finally:
                            if hook: hook.remove()
                            row['warnings']=[str(w.message) for w in ws]
                            row['library_warnings']=library_messages
                            if resolved:row['effective_generation_configuration']=configuration_record(resolved[0])
                            if row.get('status')!='complete':
                                row.update(prompt_ids=prompt_records[name]['ids'],generated_ids=recorder.ids,
                                    full_ids=prompt_records[name]['ids']+recorder.ids,output_length=len(recorder.ids),
                                    continuation_with_special_tokens=None,continuation_without_special_tokens=None,
                                    ended_on_eos=None,decoding_seconds=None)
                            if capture:
                                if capture[0].numel()==model.config.vocab_size and torch.isfinite(capture[0]).all():
                                    first_logits[name+'-'+condition]=capture[0].tolist()
                                    save(out/'raw-first-step-logits.json',first_logits)
                            row['forward_calls'] = budget.forwards-before; log(row)
            save(out/'raw-first-step-logits.json', first_logits)
            save(out/'generation-configurations.json', {'requested':configs,'effective_by_trial':{r['trial_id']:r.get('effective_generation_configuration') for r in records}})
            designated=reference_trial(records)
            if designated is None:
                comparison={'status':'unavailable: designated first P1 greedy trial failed or missing',
                            'reference_trial_id':REFERENCE_TRIAL,'rows':[],'forwards':0,'cache_reuse_exercised':False,
                            'absolute_tolerance':ATOL,'relative_tolerance':RTOL}
            else:
                with budget.trial():
                    # Forward-only comparison, not an additional generation trial.
                    budget.trials -= 1
                    comparison = matched_prefixes(model, encoded['P1'], designated['generated_ids'], budget)
                    comparison['reference_trial_id']=REFERENCE_TRIAL
            save(out/'cache-comparison.json', comparison)
            schedule = []
            for name in ('P1','P1x8'):
                schedule.extend((name,cached,True,0) for cached in (True,False))
                for rep in range(5):
                    schedule.extend((name,cached,False,rep) for cached in ((True,False) if rep%2==0 else (False,True)))
            save(out/'timing-schedule.json', schedule)
            for name,cached,warmup,rep in schedule:
                row = {'phase':'timing','prompt':name,'use_cache':cached,'warmup':warmup,'repeat':rep,
                       'prompt_length':prompt_records[name]['input_length'],
                       'stopping_rule':'exactly sixteen greedy selections; EOS termination disabled without suppressing EOS',
                       'logits_to_keep':1,'diagnostic_outputs':False}
                before = budget.forwards
                try:
                    with budget.trial(): row.update(fixed_loop(model, encoded[name], cached, budget))
                    row['status'] = 'complete'
                except Exception as exc:
                    row.update(getattr(exc,'partial_trial',{}))
                    row.update(status='failed', error=f'{type(exc).__name__}: {exc}')
                row['forward_calls'] = budget.forwards-before; log(row)
    finally:
        for handle in handles: handle.remove()
    summary = {'ordinary_repeatability': {}, 'first_step_raw_logits_max_difference': {}, 'timings': {},
               'timing_limits':'local CPU, tiny model; timestamps/dispatch and retention of sixteen logits tensors add overhead; no service benchmark'}
    for name in PROMPTS:
        for condition in ('G','S1','S2','S3'):
            values = ordinary.get((name,condition), [])
            summary['ordinary_repeatability'][name+'-'+condition] = {
                'successful_trials':len(values), 'distinct_id_continuations':len(set(map(tuple,values))),
                'repeat_identical': all(v==values[0] for v in values) if condition=='G' and values else None,
                'same_seed_17_repeat_identical': values[0]==values[-1] if condition=='S1' and len(values)==4 else None}
        logits = [first_logits.get(name+'-'+c) for c in ('G','S1','S2','S3')]
        summary['first_step_raw_logits_max_difference'][name] = max(abs(a-b) for x in logits[1:] for a,b in zip(logits[0],x)) if all(logits) else None
    for name in ('P1','P1x8'):
        for cached in (True,False):
            rows = [r for r in records if r['phase']=='timing' and r['prompt']==name and r['use_cache']==cached and not r['warmup'] and r['status']=='complete']
            summary['timings'][name+'-'+str(cached)] = {key:{'all':[r[key] for r in rows],
                'median':statistics.median(r[key] for r in rows), 'min':min(r[key] for r in rows),
                'max':max(r[key] for r in rows)} for key in ('generation_seconds','ttft_seconds','mean_post_first_interval_seconds','throughput_tokens_per_second_including_prefill')} if rows else {'status':'no successful measured trials'}
        on = [r for r in records if r['phase']=='timing' and r['prompt']==name and r['use_cache'] and not r['warmup'] and r['status']=='complete']
        off = [r for r in records if r['phase']=='timing' and r['prompt']==name and not r['use_cache'] and not r['warmup'] and r['status']=='complete']
        summary.setdefault('timing_first_divergence',{})[name] = [next((i for i,(a,b) in enumerate(zip(x['generated_ids'],y['generated_ids'])) if a!=b),None) for x,y in zip(on,off)]
    comparison_rows=comparison.get('rows',[])
    summary['cache_comparison']={'status':comparison.get('status','complete'),
        'passed_prefixes':sum(r['criterion_passed'] for r in comparison_rows),
        'failed_prefixes':sum(not r['criterion_passed'] for r in comparison_rows),
        'cause':'unresolved when a tolerance failure is observed'}
    summary['trial_count'] = len(records); summary['failed_trials'] = sum(r['status']!='complete' for r in records)
    summary['forward_calls'] = budget.forwards; summary['output_selections'] = budget.selections
    save(out/'summary.json', summary)
    return summary

def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--artifacts', type=Path, default=Path('pythia-artifacts'))
    parser.add_argument('--output', type=Path, required=True)
    parser.add_argument('--offline', action='store_true'); parser.add_argument('--plan', action='store_true')
    args = parser.parse_args()
    if args.plan:
        print(json.dumps({'repository':REPOSITORY,'revision':REVISION,'files':FILES,
                          'artifact_bytes':sum(v[0] for v in FILES.values()),'generation_trials':50,
                          'maximum_output_selections':800,'maximum_matched_forwards':16},indent=2)); return
    safe_path(args.artifacts); safe_path(args.output)
    if args.output.exists(): raise ValueError('Use a new output directory')
    args.output.mkdir(parents=True)
    started = time.perf_counter(); budget = None
    try:
        hashes = artifacts(args.artifacts, args.offline)
        setup_seconds = time.perf_counter()-started
        for package, expected in VERSIONS.items():
            if importlib.metadata.version(package).split('+')[0] != expected:
                raise ValueError(f'Use pinned environment: {package}=={expected}')
        budget = Budget()
        save(args.output/'manifest.json', {'repository':REPOSITORY,'revision':REVISION,'artifacts':hashes,
            'source_sha256':hashlib.sha256(Path(__file__).read_bytes()).hexdigest(),
            'download_or_reuse':'verified local reuse' if args.offline else 'fetch missing; verify all files',
            'artifact_setup_seconds':setup_seconds,'execution_limit_seconds':900,'trial_limit_seconds':30,
            'maximum_trials':50,'maximum_selections':800,'maximum_forwards':816,
            'extensions':'batching and accelerators not run','timings':'local instrumented CPU only'})
        with budget.execution():
            summary = run(args.artifacts, args.output, budget)
        print(json.dumps(summary,indent=2))
        if summary['failed_trials']: raise RuntimeError('Trial failures retained; inspect records')
    except Exception as exc:
        save(args.output/'failure.json', {'error':f'{type(exc).__name__}: {exc}',
             'elapsed_seconds':time.perf_counter()-started,
             'budget':{'forwards':budget.forwards,'trials':budget.trials,'selections':budget.selections} if budget else None})
        raise

if __name__ == '__main__': main()
