"""CPU CLIP-style objective on synthetic RGB pixels and six-word captions.

No pretrained CLIP, vision transformer, BPE, or natural-image benchmark.
"""
import argparse
import hashlib
import json
import platform
from pathlib import Path
import numpy as np

COLORS = ['red', 'green', 'blue']
SHAPES = ['square', 'cross', 'ring']
CAPTIONS = [f'{c} {s}' for c in COLORS for s in SHAPES]


def dataset(seed, per_class):
    rng = np.random.default_rng(seed)
    yy, xx = np.mgrid[:16, :16]
    images, labels = [], []
    for label in range(9):
        for _ in range(per_class):
            cx, cy = rng.integers(6, 10, size=2)
            dx, dy = np.abs(xx-cx), np.abs(yy-cy)
            masks = [(dx <= 3) & (dy <= 3),
                     ((dx <= 1) & (dy <= 4)) | ((dy <= 1) & (dx <= 4)),
                     ((dx*dx+dy*dy) <= 20) & ((dx*dx+dy*dy) >= 8)]
            im = rng.uniform(0, .03, (16, 16, 3))
            im[:, :, label//3] += masks[label % 3] * rng.uniform(.7, 1.)
            images.append(im.clip(0, 1).reshape(-1))
            labels.append(label)
    return np.array(images), np.array(labels)


def text_features():
    # Explicit six-word bag of words: no prelearned linguistic representation.
    vocab = COLORS + SHAPES
    return np.array([[float(w in caption.split()) for w in vocab]
                     for caption in CAPTIONS])


def normalize(x):
    n = np.linalg.norm(x, axis=1, keepdims=True)
    assert n.min() > 1e-12
    return x/n, n


def softmax(x):
    z = x-x.max(axis=1, keepdims=True)
    return np.exp(z)/np.exp(z).sum(axis=1, keepdims=True)


def objective(x, t, wi, wt, scale=10.):
    u, un = normalize(x @ wi)
    v, vn = normalize(t @ wt)
    s = scale * u @ v.T
    n = len(x)
    p, q = softmax(s), softmax(s.T)
    loss = -.5 * (np.log(np.diag(p)).mean()+np.log(np.diag(q)).mean())
    ds = (p + q.T - 2*np.eye(n))/(2*n)
    du, dv = scale*ds @ v, scale*ds.T @ u
    gu = (du-u*(du*u).sum(axis=1, keepdims=True))/un
    gv = (dv-v*(dv*v).sum(axis=1, keepdims=True))/vn
    return float(loss), [x.T @ gu, t.T @ gv]


def retrieval(scores, query_ids, gallery_ids):
    # Stable tie rule: smaller gallery index wins. Multiple relevant ids allowed.
    order = np.argsort(-scores, axis=1, kind='stable')
    hits = gallery_ids[order] == query_ids[:, None]
    positives = (query_ids[:, None] == gallery_ids[None, :]).sum(axis=1)
    assert np.all(positives > 0)
    out = {}
    for k in [1, 3, 5]:
        h = hits[:, :k].sum(axis=1)
        out[f'hit@{k}'] = float(np.mean(h > 0))
        out[f'relevant_recall@{k}'] = float(np.mean(h/positives))
    out['mrr'] = float(np.mean(1/(np.argmax(hits, axis=1)+1)))
    return out


def evaluate(x, labels, t, wi, wt):
    u, _ = normalize(x @ wi)
    v, _ = normalize(t @ wt)
    s = u @ v.T
    return s, {'i2t': retrieval(s, labels, np.arange(9)),
               't2i': retrieval(s.T, np.arange(9), labels)}


def checks():
    rng = np.random.default_rng(7100)
    x, t = rng.normal(size=(4, 5)), rng.normal(size=(4, 3))
    wi, wt = rng.normal(size=(5, 4)), rng.normal(size=(3, 4))
    _, grad = objective(x, t, wi, wt)
    errors = []
    for w, g in zip([wi, wt], grad):
        for idx in np.ndindex(w.shape):
            old = w[idx]
            w[idx] = old+1e-6
            plus = objective(x, t, wi, wt)[0]
            w[idx] = old-1e-6
            minus = objective(x, t, wi, wt)[0]
            w[idx] = old
            errors.append(abs((plus-minus)/2e-6-g[idx]))
    assert max(errors) < 1e-7
    # Perfect semantic fixture, 3 identical members per group.
    ids = np.repeat(np.arange(9), 3)
    s = (ids[:, None] == ids[None, :]).astype(float)
    semantic = retrieval(s, ids, ids)
    identity_hit = float(np.mean(np.argmax(s, axis=1) == np.arange(27)))
    assert semantic['hit@1'] == 1 and abs(identity_hit-1/3) < 1e-12
    # Same semantic class: diagonal CE cannot distinguish duplicate captions.
    p = softmax(s*10)
    diagonal_ce = float(-np.log(np.diag(p)).mean())
    positive_mass_ce = float(-np.log((p*(ids[:,None]==ids[None,:])).sum(1)).mean())
    assert abs(diagonal_ce-positive_mass_ce-np.log(3)) < 1e-12
    # Positive scale preserves ranking but not normalized confidence.
    assert np.array_equal(np.argmax(s, 1), np.argmax(10*s, 1))
    return {'gradient_parameters': 32, 'gradient_max_abs_error': max(errors),
            'duplicate_semantic_hit@1': semantic['hit@1'],
            'duplicate_identity_hit@1': identity_hit,
            'duplicate_diagonal_ce': diagonal_ce,
            'duplicate_positive_mass_ce': positive_mass_ce,
            'duplicate_ce_gap': diagonal_ce-positive_mass_ce}


def run(seed, out):
    rng = np.random.default_rng(seed)
    train_x, train_y = dataset(seed+1000, 12)
    test_x, test_y = dataset(seed+2000, 5)
    hashes = lambda a: {hashlib.sha256(row.tobytes()).hexdigest() for row in a}
    assert not hashes(train_x) & hashes(test_x)
    t = text_features()
    init = [rng.normal(0, .05, (768, 16)), rng.normal(0, .2, (6, 16))]
    _, baseline = evaluate(test_x, test_y, t, *init)
    batches = [np.arange(9)*12 + rng.integers(0, 12, 9) for _ in range(400)]
    arms = {}
    for arm, mapping in [('aligned', np.arange(9)), ('mismatched', np.roll(np.arange(9), 1))]:
        w = [a.copy() for a in init]
        m, v = [np.zeros_like(a) for a in w], [np.zeros_like(a) for a in w]
        losses = []
        for step, batch in enumerate(batches, 1):
            loss, grads = objective(train_x[batch], t[mapping], *w)
            losses.append(loss)
            for j in range(2):
                m[j] = .9*m[j]+.1*grads[j]
                v[j] = .999*v[j]+.001*grads[j]**2
                w[j] -= .01*(m[j]/(1-.9**step))/(np.sqrt(v[j]/(1-.999**step))+1e-8)
        s, metrics = evaluate(test_x, test_y, t, *w)
        assert np.isfinite(losses).all()
        # Global temperature leaves ranks invariant for the saved score matrix.
        assert np.array_equal(np.argsort(s, axis=1), np.argsort(10*s, axis=1))
        np.savez_compressed(out/f'{seed}_{arm}.npz', wi=w[0], wt=w[1],
                            scores=s, losses=losses, mapping=mapping)
        arms[arm] = {'last_batch_loss': losses[-1], **metrics}
    np.savez_compressed(out/f'{seed}_data.npz', train_x=train_x, train_y=train_y,
                        test_x=test_x, test_y=test_y, text=t, batches=batches,
                        init_wi=init[0], init_wt=init[1])
    return {'seed': seed, 'baseline': baseline, **arms}


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument('--output', type=Path, default=Path(__file__).parent/'results')
    args = ap.parse_args()
    args.output.mkdir(parents=True, exist_ok=True)
    result = {'python': platform.python_version(), 'numpy': np.__version__,
              'config': {'seeds': [71,72,73], 'steps':400, 'batch':9, 'scale':10,
                         'adam_lr':.01, 'dtype':'float64', 'parameters':12384,
                         'train_images':108, 'test_images':45, 'captions':9},
              'checks': checks(), 'runs': []}
    print('train pixels [108,768]; text [9,6]; embedding [B,16]; logits [9,9]')
    print('test scores [45,9]; text-to-image scores [9,45]')
    print('checks', json.dumps(result['checks']))
    for seed in [71,72,73]:
        r = run(seed, args.output)
        result['runs'].append(r)
        print(json.dumps(r))
    result['aggregate'] = {}
    for arm in ['baseline', 'aligned', 'mismatched']:
        result['aggregate'][arm] = {}
        for direction in ['i2t', 't2i']:
            vals = np.array([r[arm][direction]['hit@1'] for r in result['runs']])
            result['aggregate'][arm][direction] = {'hit@1_mean':float(vals.mean()),
                                                  'hit@1_sample_sd':float(vals.std(ddof=1))}
    (args.output/'summary.json').write_text(json.dumps(result, indent=2)+'\n')
    print('aggregate', json.dumps(result['aggregate']))


if __name__ == '__main__':
    main()
