"""Codex comparison harness: identical training settings, unchanged lesson files.

Run: python3 compare_models.py
Save the measurements: python3 compare_models.py --json results.json
Uses Python's standard library only.
"""
import argparse
import contextlib
import hashlib
import io
import json
import math
from pathlib import Path
import random
import runpy

ROOT = Path(__file__).resolve().parent
EPOCHS = 100
LEARNING_RATE = 0.1
SAMPLE_COUNT = 1000
MAX_TOKENS = 30


def load_lesson(filename, generator_name):
    path = ROOT / filename
    with contextlib.redirect_stdout(io.StringIO()):
        lesson = runpy.run_path(str(path))
    # The lessons run their demonstrations on load. Reset every weight in place
    # so both comparisons start from zero, including the bigram's extra update.
    weights = lesson['weights']
    rows = weights.values() if isinstance(weights, dict) else weights
    for row in rows:
        for token_id in range(len(row)):
            row[token_id] = 0.0
    lesson['comparison_generator'] = lesson[generator_name]
    lesson['source_sha256'] = hashlib.sha256(path.read_bytes()).hexdigest()
    return lesson


def average_loss(lesson):
    total = 0.0
    for context, target_id in zip(lesson['inputs'], lesson['targets']):
        probabilities = lesson['softmax'](lesson['weights'][context])
        total += -math.log(probabilities[target_id])
    return total / len(lesson['targets'])


def train(lesson):
    history = [{'epoch': 0, 'loss': average_loss(lesson)}]
    for epoch in range(1, EPOCHS + 1):
        for context, target_id in zip(lesson['inputs'], lesson['targets']):
            row = lesson['weights'][context]
            predicted = lesson['softmax'](row)
            for candidate_id in range(len(lesson['vocab'])):
                target = 1.0 if candidate_id == target_id else 0.0
                row[candidate_id] -= LEARNING_RATE * (predicted[candidate_id] - target)
        history.append({'epoch': epoch, 'loss': average_loss(lesson)})
    return history


def counting_reference(lesson):
    counts = {}
    for context, target_id in zip(lesson['inputs'], lesson['targets']):
        following = counts.setdefault(context, {})
        following[target_id] = following.get(target_id, 0) + 1
    total_loss = 0.0
    for context, target_id in zip(lesson['inputs'], lesson['targets']):
        following = counts[context]
        total_loss -= math.log(following[target_id] / sum(following.values()))
    return total_loss / len(lesson['targets'])


def summarize(lesson):
    history = train(lesson)
    rows = lesson['weights']
    parameter_count = sum(len(row) for row in (rows.values() if isinstance(rows, dict) else rows))
    samples = []
    matches = empty = at_limit = 0
    # Reset the seed for each output in each model: reproducible samples without
    # selecting only the attractive outputs. Different models may draw different
    # numbers of tokens; identical seeds do not imply identical token choices.
    saved_random_state = random.getstate()
    try:
        for seed in range(SAMPLE_COUNT):
            random.seed(seed)
            output = lesson['comparison_generator'](MAX_TOKENS)
            if seed < 5:
                samples.append({'seed': seed, 'text': output})
            matches += output in lesson['training_texts']
            empty += output == ''
            at_limit += len(output) == MAX_TOKENS
    finally:
        random.setstate(saved_random_state)
    return {
        'source_sha256': lesson['source_sha256'],
        'allocated_contexts': len(rows),
        'observed_contexts': len(set(lesson['inputs'])),
        'parameters': parameter_count,
        'training_pairs': len(lesson['targets']),
        'history': history,
        'counting_reference_loss': counting_reference(lesson),
        'samples': samples,
        'sample_count': SAMPLE_COUNT,
        'training_string_matches': matches,
        'empty_outputs': empty,
        'outputs_at_character_limit': at_limit,
    }


def compare():
    one = load_lesson('bigram_checkpoint.py', 'generate_learned')
    two = load_lesson('context_two.py', 'generate')
    assert one['training_texts'] == two['training_texts']
    assert one['vocab'] == two['vocab']
    assert one['targets'] == two['targets']
    assert one['inputs'] == [context[-1] for context in two['inputs']]
    result = {
        'training_texts': one['training_texts'],
        'vocab': one['vocab'],
        'epochs': EPOCHS,
        'learning_rate': LEARNING_RATE,
        'updates_per_model': EPOCHS * len(one['targets']),
        'max_generated_tokens': MAX_TOKENS,
        'sample_seeds': [0, SAMPLE_COUNT - 1],
        'loss_scope': 'Full softmax, average over the training pairs; not held-out evaluation.',
        'generation_rule': 'Exclude START from sampling, stop at END or the 30-token cap.',
        'one_token': summarize(one),
        'two_tokens': summarize(two),
    }
    return result


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--json', type=Path, help='Write the full measurements to this file.')
    args = parser.parse_args()
    result = compare()
    print('Same three examples, 24 pairs, zero initialization, 100 epochs, learning rate 0.1.')
    print('All losses below are training losses, not held-out scores.\n')
    for key in ('one_token', 'two_tokens'):
        item = result[key]
        print(key)
        print('  Weights:', item['parameters'])
        print('  Loss:', round(item['history'][0]['loss'], 6), '->', round(item['history'][-1]['loss'], 6))
        print('  Counting reference:', round(item['counting_reference_loss'], 6))
        print('  Training-string matches:', item['training_string_matches'], '/', SAMPLE_COUNT)
        print('  Empty / at length limit:', item['empty_outputs'], '/', item['outputs_at_character_limit'])
        for sample in item['samples']:
            print('  Seed', sample['seed'], ':', repr(sample['text']))
    if args.json:
        args.json.write_text(json.dumps(result, indent=2) + '\n')
        print('\nSaved measurements to', args.json)


if __name__ == '__main__':
    main()
