#!/usr/bin/env python3
"""Original, standard-library implementation of the fixed Lab 10 bigram protocol.

All output stays local. No downloads, packages, accounts or arbitrary serializers.
"""
import argparse
from collections import Counter
import hashlib
import json
import math
from pathlib import Path
import platform
import random
import time

VOCAB = ['<bos>', '<eos>', 'red', 'blue', 'green', 'gold', 'orb', 'cube',
         'cone', 'ring', 'glows', 'spins', 'rests', 'rolls']
VERSION = 'course-bigram-softmax-v1'
SCHEDULE = [0, 1, 5, 20, 100, 300]
CONFIG = {'epochs': 300, 'batch_size': 20, 'seed': 17,
          'optimizer': 'SGD', 'optimizer_state': {},
          'shuffle': 'fresh canonical pairs; random.Random(17 + epoch)'}


def encoded(value):
    return (json.dumps(value, sort_keys=True, indent=2, allow_nan=False) + '\n').encode()


def digest(data):
    return hashlib.sha256(data).hexdigest()


def unique_keys(items):
    result = {}
    for key, value in items:
        if key in result:
            raise ValueError('Duplicate JSON key')
        result[key] = value
    return result


def load(path):
    return json.loads(Path(path).read_text(), object_pairs_hook=unique_keys,
                      parse_constant=lambda value: (_ for _ in ()).throw(ValueError(value)))


def destination(path):
    path = Path(path)
    if any(p.is_symlink() for p in [path, *path.parents]):
        raise ValueError('Output path must not traverse symlinks')
    if path.exists() and (not path.is_dir() or any(path.iterdir())):
        raise ValueError('Use a fresh empty output directory')
    path.mkdir(parents=True, exist_ok=True)
    return path


def write(path, value):
    with Path(path).open('xb') as stream:
        stream.write(encoded(value))


def records():
    groups = {name: [] for name in ['train', 'validation', 'test']}
    for c in range(4):
        for s in range(4):
            for a in range(4):
                r = (c + s + a) % 4
                split = 'test' if r == 0 else 'validation' if r == 1 else 'train'
                groups[split].append({'id': f'c{c}-s{s}-a{a}', 'indices': [c, s, a],
                                      'tokens': [0, 2+c, 6+s, 10+a, 2+c, 1]})
    return groups


def prepare(output):
    output = destination(output)
    write(output/'vocabulary.json', VOCAB)
    files = {}
    for name, rows in records().items():
        content = ''.join(json.dumps(r, sort_keys=True) + '\n' for r in rows).encode()
        with (output/f'{name}.jsonl').open('xb') as stream:
            stream.write(content)
        files[name] = {'file': f'{name}.jsonl', 'records': len(rows), 'sha256': digest(content)}
    manifest = {'generator': VERSION, 'vocabulary': VOCAB,
                'vocabulary_sha256': digest((output/'vocabulary.json').read_bytes()),
                'assignment': '(c+s+a)%4: 0 test, 1 validation, 2/3 train', 'files': files}
    write(output/'manifest.json', manifest)
    return manifest


def read_records(path, split=None):
    rows = [json.loads(line, object_pairs_hook=unique_keys) for line in Path(path).read_text().splitlines()]
    expected = records()
    if split is None:
        matches = [name for name, values in expected.items() if rows == values]
        if len(matches) != 1:
            raise ValueError('File is not an exact canonical Lab 10 partition')
        split = matches[0]
    if rows != expected[split] or any(type(t) is not int for row in rows for t in row['tokens']+row['indices']):
        raise ValueError(f'Expected canonical {split} records in nested generator order')
    return rows


def pairs(rows):
    return [(i, y) for row in rows for i, y in zip(row['tokens'], row['tokens'][1:])]


def audit(data):
    data = Path(data)
    manifest = load(data/'manifest.json')
    if load(data/'vocabulary.json') != VOCAB or manifest['vocabulary_sha256'] != digest((data/'vocabulary.json').read_bytes()):
        raise ValueError('Vocabulary mismatch')
    partitions = {name: read_records(data/f'{name}.jsonl', name) for name in records()}
    ids, sequences = set(), set()
    for name, rows in partitions.items():
        if digest((data/f'{name}.jsonl').read_bytes()) != manifest['files'][name]['sha256']:
            raise ValueError('Partition fingerprint mismatch')
        for row in rows:
            t = row['tokens']
            assert len(t) == 6 and t[0] == 0 and t[-1] == 1 and t[1] == t[4]
            assert row['id'] not in ids and tuple(t) not in sequences
            ids.add(row['id']); sequences.add(tuple(t))
    counts = {name: Counter(pairs(rows)) for name, rows in partitions.items()}
    assert {t for row in partitions['train'] for t in row['tokens']} == set(range(14))
    for pair in set().union(*counts.values()):
        assert counts['train'][pair] == 2*counts['validation'][pair] == 2*counts['test'][pair]
    assert [len(partitions[n]) for n in ['train', 'validation', 'test']] == [32, 16, 16]
    return {'records': {n: len(r) for n, r in partitions.items()},
            'targets': {n: len(pairs(r)) for n, r in partitions.items()},
            'unique_ids': len(ids), 'all_ordinary_tokens_in_training': True,
            'pair_count_balance': True, 'cross_record_pairs': False,
            'scope': 'held-out combinations from a shared generator; no performance evaluation'}


def softmax(row):
    peak = max(row)
    exp = [math.exp(x-peak) for x in row]
    total = sum(exp)
    return [x/total for x in exp]


def nll(row, target):
    peak = max(row)
    return peak + math.log(sum(math.exp(x-peak) for x in row)) - row[target]


def gradient(weights, batch):
    if not batch:
        raise ValueError('Empty batch')
    g = [[0.0]*len(row) for row in weights]
    for i, target in batch:
        probabilities = softmax(weights[i])
        for j, probability in enumerate(probabilities):
            g[i][j] += probability - (j == target)
    return [[value/len(batch) for value in row] for row in g]


def update(weights, batch, learning_rate):
    g = gradient(weights, batch)
    return [[w-learning_rate*d for w, d in zip(row, derivatives)]
            for row, derivatives in zip(weights, g)]


def metrics(weights, examples):
    total = sum(nll(weights[i], target) for i, target in examples)
    mean = total/len(examples)
    return {'total_nll': total, 'targets': len(examples), 'mean_nats': mean,
            'bits_per_token': mean/math.log(2), 'perplexity': math.exp(mean)}


def baselines(examples):
    targets = Counter(target for _, target in examples)
    assert [targets[j] for j in range(14)] == [0, 32]+[16]*4+[8]*8
    return {'uniform': [1/14]*14, 'training_unigram': [(targets[j]+1)/174 for j in range(14)]}


def baseline_metrics(probabilities, examples):
    total = sum(-math.log(probabilities[y]) for _, y in examples)
    return {'total_nll': total, 'targets': len(examples), 'mean_nats': total/len(examples),
            'bits_per_token': total/len(examples)/math.log(2), 'perplexity': math.exp(total/len(examples))}


def readouts(weights):
    return {VOCAB[i]: softmax(weights[i]) for i in [0, 2, 6, 10]}


def checkpoint_load(path):
    cp = load(path)
    if cp['format'] != VERSION or cp['vocabulary'] != VOCAB or cp['vocabulary_sha256'] != digest(encoded(VOCAB)):
        raise ValueError('Checkpoint format or vocabulary mismatch')
    if cp['program_sha256'] != digest(Path(__file__).read_bytes()):
        raise ValueError('Checkpoint belongs to a different implementation')
    w = cp['weights']
    if not isinstance(w, list) or len(w) != 14 or any(not isinstance(row, list) or len(row) != 14 for row in w):
        raise ValueError('Checkpoint matrix shape mismatch')
    if any(type(x) not in (float, int) or not math.isfinite(x) for row in w for x in row):
        raise ValueError('Checkpoint contains nonfinite or nonnumeric weights')
    c = cp['config']
    if any(c.get(k) != v for k, v in CONFIG.items()) or type(c.get('learning_rate')) not in (float, int) or c['learning_rate'] not in [1.0, 0.1]:
        raise ValueError('Unsupported training protocol')
    if type(cp['epoch']) is not int or type(cp['updates']) is not int or not 0 <= cp['epoch'] <= 300 or cp['updates'] != 8*cp['epoch']:
        raise ValueError('Checkpoint step accounting mismatch')
    if cp['weights_sha256'] != digest(encoded(w)):
        raise ValueError('Checkpoint weight fingerprint mismatch')
    if cp['baselines'] != baselines(pairs(records()['train'])):
        raise ValueError('Frozen baseline does not match training-only target counts')
    for distribution in cp['baselines'].values():
        if len(distribution) != 14 or any(type(x) not in (float, int) or not math.isfinite(x) or x <= 0 for x in distribution) or abs(sum(distribution)-1) > 1e-12:
            raise ValueError('Invalid frozen baseline')
    return cp


def train(training_file, validation_file, learning_rate, output, resume=None):
    if type(learning_rate) not in (float, int) or learning_rate not in [1.0, 0.1]:
        raise ValueError('Predeclared learning rates are 1.0 and 0.1')
    # Deliberately no data-directory or manifest argument: fitting opens only these two partitions.
    train_pairs = pairs(read_records(training_file, 'train'))
    validation_pairs = pairs(read_records(validation_file, 'validation'))
    fingerprints = {'train': digest(Path(training_file).read_bytes()),
                    'validation': digest(Path(validation_file).read_bytes())}
    frozen = baselines(train_pairs)
    config = dict(CONFIG, learning_rate=learning_rate)
    weights = [[0.0]*14 for _ in range(14)]
    first_epoch = 0
    if resume:
        cp = checkpoint_load(resume)
        if cp['config'] != config or cp['data_sha256'] != fingerprints or cp['baselines'] != frozen:
            raise ValueError('Resume configuration/data/baseline mismatch')
        weights, first_epoch = cp['weights'], cp['epoch']
    output = destination(output)
    (output/'checkpoints').mkdir()
    write(output/'config.json', config)
    rows, start = [], time.perf_counter()
    for epoch in range(first_epoch, 301):
        if epoch > first_epoch:
            shuffled = train_pairs[:]
            random.Random(17+epoch).shuffle(shuffled)
            for offset in range(0, len(shuffled), 20):
                weights = update(weights, shuffled[offset:offset+20], learning_rate)
        snapshot = {'epoch': epoch, 'updates': 8*epoch, 'learning_rate': learning_rate,
                    'elapsed_seconds': time.perf_counter()-start,
                    'train': metrics(weights, train_pairs),
                    'validation': metrics(weights, validation_pairs)}
        rows.append(snapshot)
        if epoch in SCHEDULE or epoch == first_epoch:
            cp = {'format': VERSION, 'weights': weights, 'weights_sha256': digest(encoded(weights)),
                  'vocabulary': VOCAB, 'vocabulary_sha256': digest(encoded(VOCAB)),
                  'program_sha256': digest(Path(__file__).read_bytes()),
                  'python': platform.python_version(), 'epoch': epoch, 'updates': 8*epoch,
                  'config': config, 'data_sha256': fingerprints, 'baselines': frozen,
                  'metrics_at_save': {'train': snapshot['train'], 'validation': snapshot['validation']},
                  'probabilities_at_save': readouts(weights)}
            write(output/'checkpoints'/f'epoch-{epoch:03d}.json', cp)
    with (output/'metrics.jsonl').open('x') as stream:
        for row in rows:
            stream.write(json.dumps(row, allow_nan=False)+'\n')
    return rows


def evaluate(checkpoint, evaluation_file):
    before = Path(checkpoint).read_bytes()
    cp = checkpoint_load(checkpoint)
    examples = pairs(read_records(evaluation_file))
    result = {'checkpoint_sha256': digest(before), 'evaluation_sha256': digest(Path(evaluation_file).read_bytes()),
              'model': metrics(cp['weights'], examples),
              'baselines': {name: baseline_metrics(p, examples) for name, p in cp['baselines'].items()}}
    assert Path(checkpoint).read_bytes() == before
    return result


def sample(checkpoint):
    cp = checkpoint_load(checkpoint)
    output = []
    for index in range(8):
        rng, ids = random.Random(23+index), [0]
        for _ in range(20):
            probs = softmax(cp['weights'][ids[-1]])
            draw, cumulative, chosen = rng.random(), 0.0, 13
            for j, p in enumerate(probs):
                cumulative += p
                if draw < cumulative:
                    chosen = j
                    break
            ids.append(chosen)
            if chosen == 1:
                break
        structure = len(ids) == 6 and ids[0] == 0 and 2 <= ids[1] <= 5 and 6 <= ids[2] <= 9 and 10 <= ids[3] <= 13 and 2 <= ids[4] <= 5 and ids[5] == 1
        output.append({'epoch': cp['epoch'], 'sample_index': index, 'seed': 23+index,
                       'ids': ids, 'tokens': [VOCAB[i] for i in ids],
                       'truncated': ids[-1] != 1, 'valid_structure': structure,
                       'matching_colors': bool(structure and ids[1] == ids[4])})
    return {'protocol': {'temperature': 1, 'max_generated_tokens': 20, 'filter': None},
            'checkpoint_sha256': digest(Path(checkpoint).read_bytes()), 'samples': output,
            'next_token_readouts': readouts(cp['weights']),
            'prefixes': [[0, 2, 6, 10], [0, 3, 6, 10]],
            'prefix_distributions': [softmax(cp['weights'][10]), softmax(cp['weights'][10])]}


def selftest():
    z, target, eps = [0.0]*3, 1, 1e-5
    analytic = [p-(j == target) for j, p in enumerate(softmax(z))]
    numeric = []
    for j in range(3):
        plus, minus = z[:], z[:]
        plus[j] += eps; minus[j] -= eps
        numeric.append((nll(plus, target)-nll(minus, target))/(2*eps))
    errors = [abs(a-b) for a, b in zip(analytic, numeric)]
    assert max(errors) < 1e-6
    new = [a-0.3*g for a, g in zip(z, analytic)]
    assert max(abs(a-b) for a, b in zip(new, [-0.1, 0.2, -0.1])) < 1e-14
    assert abs(softmax(new)[1]-0.4029599111828766) < 1e-12
    assert abs(nll(new, 1)-0.9089181979565276) < 1e-12
    w = [[0.0]*14 for _ in range(14)]
    examples = pairs(records()['train'])
    assert abs(metrics(w, examples)['mean_nats']-math.log(14)) < 1e-12
    batch = [(0, 2), (0, 3), (2, 6)]
    full = gradient(w, batch)
    singles = [gradient(w, [pair]) for pair in batch]
    error = max(abs(full[i][j]-sum(g[i][j] for g in singles)/3) for i in range(14) for j in range(14))
    assert error < 1e-14
    toy = {'tokens': [0, 2, 6, 10, 2, 1]}
    assert pairs([toy, toy]) == [(0, 2), (2, 6), (6, 10), (10, 2), (2, 1)]*2
    assert (1, 0) not in examples
    uni = baseline_metrics(baselines(examples)['training_unigram'], examples)
    assert abs(uni['mean_nats']-2.447578618) < 1e-9
    return {'finite_difference_errors': errors, 'post_update_logits': new,
            'target_probability': softmax(new)[1], 'post_update_loss': nll(new, 1),
            'repeated_input_gradient_error': error, 'uniform': metrics(w, examples),
            'training_unigram': uni, 'boundary_checks': True}


def curves(run_metrics, output):
    lines = ['<svg xmlns="http://www.w3.org/2000/svg" width="800" height="440" viewBox="0 0 800 440">',
             '<rect width="800" height="440" fill="white"/>',
             '<text x="60" y="25">Lab 10 actual loss: train/validation coincide within each run</text>',
             '<path d="M60 50V365H760" fill="none" stroke="black"/>',
             '<text x="340" y="420">Optimizer updates (0 to 2400)</text>',
             '<text x="65" y="45">Mean nats per target (1.3 to 2.7)</text>']
    for index, (lr, rows) in enumerate(run_metrics.items()):
        for split, dash in [('train', ''), ('validation', '5 5')]:
            points = ' '.join(f'{60+r["updates"]/2400*700:.2f},{365-(r[split]["mean_nats"]-1.3)/1.4*300:.2f}' for r in rows)
            color = ['#2459a0', '#a7491b'][index]
            lines.append(f'<polyline points="{points}" fill="none" stroke="{color}" stroke-width="2" stroke-dasharray="{dash}"/>')
            lines.append(f'<text x="{80+index*350}" y="{385+(split=="validation")*18}" fill="{color}">lr {lr} {split} ({"dashed" if dash else "solid"})</text>')
    lines.append('</svg>')
    Path(output).write_text('\n'.join(lines))


def experiment(output):
    root = destination(output)
    started = time.perf_counter()
    prepare(root/'data')
    write(root/'audit.json', audit(root/'data'))
    write(root/'selftest.json', selftest())
    data = root/'data'
    runs = {}
    # Both predefined final checkpoints are frozen before any test metric is read.
    for lr in [1.0, 0.1]:
        runs[str(lr)] = train(data/'train.jsonl', data/'validation.jsonl', lr, root/f'lr-{lr}')
    results = {'python': platform.python_version(), 'program_sha256': digest(Path(__file__).read_bytes()), 'runs': {}}
    for lr in [1.0, 0.1]:
        run = root/f'lr-{lr}'
        cp100 = run/'checkpoints/epoch-100.json'
        cp = checkpoint_load(cp100)
        recomputed = {'train': metrics(cp['weights'], pairs(read_records(data/'train.jsonl', 'train'))),
                      'validation': metrics(cp['weights'], pairs(read_records(data/'validation.jsonl', 'validation')))}
        assert recomputed == cp['metrics_at_save'] and readouts(cp['weights']) == cp['probabilities_at_save']
        train(data/'train.jsonl', data/'validation.jsonl', lr, root/f'resume-{lr}', cp100)
        final = checkpoint_load(run/'checkpoints/epoch-300.json')
        resumed = checkpoint_load(root/f'resume-{lr}/checkpoints/epoch-300.json')
        difference = max(abs(a-b) for ar, br in zip(final['weights'], resumed['weights']) for a, b in zip(ar, br))
        assert difference == 0
        write(run/'samples.json', [sample(run/'checkpoints'/f'epoch-{epoch:03d}.json') for epoch in SCHEDULE])
        write(run/'readouts.json', {str(e): checkpoint_load(run/'checkpoints'/f'epoch-{e:03d}.json')['probabilities_at_save'] for e in SCHEDULE})
        results['runs'][str(lr)] = {'test': evaluate(run/'checkpoints/epoch-300.json', data/'test.jsonl'),
                                  'reload_metrics_equal': True, 'reload_probabilities_equal': True,
                                  'resume_max_abs_difference': difference,
                                  'updates': final['updates'], 'sample_count': 48,
                                  'max_train_validation_difference': max(abs(r['train']['mean_nats']-r['validation']['mean_nats']) for r in runs[str(lr)])}
    curves(runs, root/'loss-curves.svg')
    results['elapsed_seconds'] = time.perf_counter()-started
    results['scope'] = 'fixed bigram arithmetic experiment; no generalization or assistant capability claim'
    write(root/'experiment.json', results)
    return results


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    commands = parser.add_subparsers(dest='operation', required=True)
    for name in ['prepare', 'experiment', 'selftest']:
        p = commands.add_parser(name); p.add_argument('--output', type=Path, required=True)
    p = commands.add_parser('audit'); p.add_argument('--data', type=Path, required=True); p.add_argument('--output', type=Path, required=True)
    p = commands.add_parser('train')
    for name in ['train', 'validation', 'output']:
        p.add_argument('--'+name, type=Path, required=True)
    p.add_argument('--learning-rate', type=float, required=True); p.add_argument('--resume', type=Path)
    for name in ['evaluate', 'sample']:
        p = commands.add_parser(name); p.add_argument('--checkpoint', type=Path, required=True); p.add_argument('--output', type=Path, required=True)
        if name == 'evaluate': p.add_argument('--data', type=Path, required=True)
    args = parser.parse_args()
    if args.operation in ['prepare', 'experiment']:
        result = globals()[args.operation](args.output)
    elif args.operation == 'train':
        selftest()
        result = train(args.train, args.validation, args.learning_rate, args.output, args.resume)[-1]
    else:
        if args.operation == 'selftest': result = selftest()
        elif args.operation == 'audit': result = audit(args.data)
        elif args.operation == 'evaluate': result = evaluate(args.checkpoint, args.data)
        else: result = sample(args.checkpoint)
        directory = destination(args.output); write(directory/(args.operation+'.json'), result)
    print(json.dumps(result, indent=2, allow_nan=False))


if __name__ == '__main__':
    main()
