"""CPU-only teaching experiments. All input scores are synthetic, not assay results.

Run: python3 experiments.py
The fixed grouping and seeded selection match labs.mjs in the web edition.
"""
from pathlib import Path
import csv
import json
import math


def load_rows():
    groups = {}
    path = Path(__file__).resolve().parent / 'data' / 'simulated-variants.csv'
    with path.open() as f:
        for row in csv.DictReader(f):
            assert row['data_kind'] == 'synthetic_not_experimental'
            key = row['variant_id']
            groups.setdefault(key, dict(id=key, x=int(row['mutations']), split=row['split'], values=[]))
            assert groups[key]['split'] == row['split']
            groups[key]['values'].append(float(row['score']))
    return [dict(id=r['id'], x=r['x'], split=r['split'], y=sum(r['values']) / len(r['values'])) for r in groups.values()]


def fit(rows):
    mx = sum(r['x'] for r in rows) / len(rows)
    my = sum(r['y'] for r in rows) / len(rows)
    variance = sum((r['x']-mx)**2 for r in rows)
    slope = sum((r['x']-mx)*(r['y']-my) for r in rows) / variance if variance else 0
    return dict(intercept=my-slope*mx, slope=slope, mean=my)


def predict(model, x):
    return model['intercept'] + model['slope'] * x


def mae(rows, prediction):
    return sum(abs(r['y']-prediction(r)) for r in rows) / len(rows)


def random_generator(seed):
    def random():
        nonlocal seed
        seed = (1664525*seed + 1013904223) % (2**32)
        return seed / (2**32)
    return random


def choose(candidates, observed, strategy, random):
    # Candidate records deliberately have no unrevealed labels.
    assert all('y' not in row for row in candidates)
    if strategy == 'random':
        return candidates[math.floor(random()*len(candidates))]['id']
    model = fit(observed)
    return max(candidates, key=lambda row: predict(model, row['x']) + 3*min(abs(row['x']-r['x']) for r in observed))['id']


def replay(rows, strategy, seed):
    random = random_generator(seed)
    test = [r for r in rows if r['split'] == 'test']
    pool = [r for r in rows if r['split'] == 'train']
    shuffled = list(pool)
    for i in range(len(shuffled)-1, 0, -1):
        j = math.floor(random()*(i+1))
        shuffled[i], shuffled[j] = shuffled[j], shuffled[i]
    observed = shuffled[:4]
    initial = [r['id'] for r in observed]
    candidates = [dict(id=r['id'], x=r['x']) for r in pool if r['id'] not in initial]
    oracle = {r['id']: r for r in pool}
    def record():
        model = fit(observed)
        return dict(budget=len(observed)-4, best=max(r['y'] for r in observed), mae=mae(test, lambda r: predict(model,r['x'])))
    history = [record()]
    selected = []
    for _ in range(6):
        id = choose(candidates, observed, strategy, random)
        selected.append(id)
        observed.append(oracle[id])
        candidates = [r for r in candidates if r['id'] != id]
        history.append(record())
    return dict(initial=initial, selected=selected, history=history)


def run():
    rows = load_rows()
    train = [r for r in rows if r['split'] == 'train']
    test = [r for r in rows if r['split'] == 'test']
    assert len(train) == 18 and len(test) == 6
    assert not ({r['id'] for r in train} & {r['id'] for r in test})
    exact = fit([dict(x=0,y=2),dict(x=1,y=5),dict(x=2,y=8)])
    assert abs(predict(exact,3)-11) < 1e-12
    model = fit(train)
    result = dict(data_kind='synthetic_not_experimental', train=len(train), test=len(test), model=model,
                  baseline_mae=mae(test, lambda r:model['mean']), linear_mae=mae(test, lambda r:predict(model,r['x'])),
                  predictions=[dict(id=r['id'], observed=r['y'], predicted=predict(model,r['x']), residual=r['y']-predict(model,r['x'])) for r in test],
                  replays={})
    for seed in [11,29,47]:
        result['replays'][str(seed)] = {strategy:replay(rows,strategy,seed) for strategy in ['random','model']}
        a,b = result['replays'][str(seed)].values()
        assert a['initial']==b['initial']
        for replay_result in [a,b]:
            assert len(set(replay_result['selected']))==6
            assert not set(replay_result['selected']) & {r['id'] for r in test}
    return result


if __name__ == '__main__':
    print(json.dumps(run(),ensure_ascii=False,indent=2))
