#!/usr/bin/env python3
"""Lab 09 completed reference. Attempt the learner scaffold before running."""
import json
import math
import sys
from pathlib import Path
SPEC_ID = 'tiny-moe-routing-v1'
K = 2
INPUTS = [{'position': 0, 'token_id': 11, 'x': [1.0, 1.0]}, {'position': 1, 'token_id': 11, 'x': [1.0, 0.0]}, {'position': 2, 'token_id': 17, 'x': [0.0, 1.0]}]
LAYERS = [{'router_W': [[math.log(4), 0.0], [0.0, math.log(3)], [math.log(2), 0.0], [0.0, 0.0]], 'router_b': [0.0, 0.0, 0.0, 0.0], 'experts': [[[2.0, 0.0], [0.0, 0.0]], [[0.0, 0.0], [0.0, 3.0]], [[-1.0, 0.0], [0.0, 1.0]], [[0.0, 1.0], [1.0, 0.0]]]}, {'router_W': [[0.0, 0.0] for _ in range(4)], 'router_b': [0.0, math.log(2), math.log(4), 0.0], 'experts': [[[0.0, 1.0], [0.0, 0.0]], [[1.0, 0.0], [0.0, 0.0]], [[0.0, 0.0], [0.0, 2.0]], [[-1.0, 0.0], [0.0, -1.0]]]}]

def matvec(matrix, vector):
    return [sum((a * b for a, b in zip(row, vector))) for row in matrix]

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

def choose(logits, eligible, k):
    if len(set(eligible)) != len(eligible) or len(eligible) < k:
        raise ValueError('Need at least k distinct eligible experts')
    return sorted(eligible, key=lambda e: (-logits[e], e))[:k]

def trace_one(x, layer_id, position, token_id, scenario):
    local = scenario if layer_id == 0 and position == 0 else 'baseline'
    layer = LAYERS[layer_id]
    projection = matvec(layer['router_W'], x)
    raw = [a + b for a, b in zip(projection, layer['router_b'])]
    logits = raw[:]
    if local == 'boost_2':
        logits[2] += math.log(4)
    eligible = list(range(4))
    if local == 'prohibit_0':
        eligible = [1, 2, 3]
    elif local == 'force_23':
        eligible = [2, 3]
    selected = choose(logits, eligible, K)
    weights = softmax([logits[e] for e in selected])
    original = [matvec(layer['experts'][e], x) for e in selected]
    used = [vector[:] for vector in original]
    if local == 'ablate_0':
        used[selected.index(0)] = [0.0, 0.0]
    update = [sum((w * vector[j] for w, vector in zip(weights, used))) for j in range(2)]
    output = [a + b for a, b in zip(x, update)]
    return {'scenario': scenario, 'local_intervention': local, 'position': position, 'token_id': token_id, 'layer': layer_id, 'input': x[:], 'raw_logits': raw, 'effective_logits': logits, 'full_probabilities': softmax(logits), 'eligible': eligible, 'selected': selected, 'combine_weights': weights, 'expert_outputs_original': original, 'expert_outputs_used': used, 'update': update, 'output': output}
SCENARIOS = ['baseline', 'prohibit_0', 'force_23', 'boost_2', 'ablate_0']


def run():
    records = []
    for scenario in SCENARIOS:
        for item in INPUTS:
            x = item['x'][:]
            for layer in range(len(LAYERS)):
                row = trace_one(x, layer, item['position'], item['token_id'], scenario)
                records.append(row)
                x = row['output'][:]
    return records


def main():
    import argparse
    import platform
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output', type=Path, default=Path('moe-results'))
    parser.add_argument('--overwrite', action='store_true', help='Explicitly replace the two named result files')
    args = parser.parse_args()
    if args.output.is_symlink():
        raise ValueError('Output directory must not be a symlink')
    args.output.mkdir(parents=True, exist_ok=True)
    files = [args.output / 'run-spec.json', args.output / 'routing-log.jsonl']
    if any(p.is_symlink() for p in files):
        raise ValueError('Output files must not be symlinks')
    if list(args.output.iterdir()) and not args.overwrite:
        raise ValueError('Output directory is nonempty; choose a new directory or --overwrite')
    records = run()
    spec = {'spec_id': SPEC_ID, 'python': platform.python_version(),
            'kind': 'original two-layer arithmetic toy; not trained-model tracing',
            'top_k': K, 'layers': LAYERS, 'inputs': INPUTS,
            'tie_rule': 'descending logit then ascending expert ID',
            'combine_rule': 'softmax over selected effective logits',
            'normalization': 'none', 'attention': 'none', 'capacity_limit': None,
            'scenarios': SCENARIOS, 'intervention_site': {'layer': 0, 'position': 0}}
    files[0].write_text(json.dumps(spec, indent=2) + '\n', encoding='utf-8')
    files[1].write_text(''.join(json.dumps(row) + '\n' for row in records), encoding='utf-8')
    print(f'Saved {len(records)} routing records in {args.output}')


if __name__ == '__main__':
    main()
