"""Reproduce the two public computational notes by Maurizio Viviani.

Usage: python reproduce_research.py --output-dir results
Dependencies: numpy==2.3.5 matplotlib==3.10.8
All data are synthetic. No quantum hardware or particle-detector data are used.
"""
from pathlib import Path
import argparse
import csv
import json
import platform

import numpy as np
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt

AUTHOR = 'Maurizio Viviani'
SEED = 20261011
EDGES = [(0, 1), (1, 2), (2, 3), (3, 0), (0, 2)]
COSTS = np.array([sum(((z >> u) & 1) != ((z >> v) & 1)
                     for u, v in EDGES) for z in range(16)], dtype=float)


def qaoa_probabilities(gamma, beta):
    state = np.exp(-1j * gamma * COSTS) / 4.0
    for qubit in range(4):
        partner = np.arange(16) ^ (1 << qubit)
        state = np.cos(beta) * state - 1j * np.sin(beta) * state[partner]
    probabilities = np.abs(state) ** 2
    if not np.isclose(probabilities.sum(), 1.0, atol=1e-12):
        raise RuntimeError('State normalization failed')
    return probabilities


def shot_study():
    best = (-np.inf, 0.0, 0.0, None)
    for gamma in np.linspace(0, np.pi, 51):
        for beta in np.linspace(-np.pi / 2, np.pi / 2, 51):
            probabilities = qaoa_probabilities(gamma, beta)
            expectation = float(probabilities @ COSTS)
            if expectation > best[0]:
                best = (expectation, float(gamma), float(beta), probabilities)
    expectation, gamma, beta, probabilities = best
    variance = float(probabilities @ ((COSTS - expectation) ** 2))
    cost_probability = np.bincount(COSTS.astype(int), weights=probabilities, minlength=5)
    cost_probability /= cost_probability.sum()
    rng = np.random.default_rng(SEED)
    repetitions = 2000
    rows = []
    for shots in [64, 256, 1024, 4096]:
        counts = rng.multinomial(shots, cost_probability, size=repetitions)
        means = counts @ np.arange(5) / shots
        rows.append({
            'shots': shots,
            'repetitions': repetitions,
            'bias': float(np.mean(means) - expectation),
            'rmse': float(np.sqrt(np.mean((means - expectation) ** 2))),
            'predicted_standard_error': float(np.sqrt(variance / shots)),
            'empirical_standard_deviation': float(np.std(means, ddof=1)),
        })
    return {
        'description': 'Fixed-angle ideal four-qubit p=1 QAOA sampling study',
        'seed': SEED,
        'edges': EDGES,
        'grid_points_per_angle': 51,
        'gamma': gamma,
        'beta': beta,
        'exact_expectation': expectation,
        'exact_variance': variance,
        'classical_optimum': int(COSTS.max()),
        'optimal_cut_probability': float(probabilities[COSTS == COSTS.max()].sum()),
        'probabilities': probabilities.tolist(),
        'costs': COSTS.astype(int).tolist(),
        'rows': rows,
    }


def sigmoid(logits):
    return 1.0 / (1.0 + np.exp(-logits))


def synthetic_classifier_data(rng, count):
    x = rng.normal(size=(count, 2))
    true_logits = 1.2 * x[:, 0] - 0.9 * x[:, 1] + 0.7 * x[:, 0] * x[:, 1]
    y = rng.binomial(1, sigmoid(true_logits))
    return 2.5 * true_logits, y


def negative_log_likelihood(logits, y):
    return float(np.mean(np.logaddexp(0.0, logits) - y * logits))


def auc_from_scores(scores, y):
    order = np.argsort(scores, kind='stable')
    sorted_scores = scores[order]
    ranks = np.empty(len(scores), dtype=float)
    start = 0
    while start < len(scores):
        end = start + 1
        while end < len(scores) and sorted_scores[end] == sorted_scores[start]:
            end += 1
        ranks[order[start:end]] = (start + 1 + end) / 2.0
        start = end
    positive = int(y.sum())
    negative = len(y) - positive
    return float((ranks[y == 1].sum() - positive * (positive + 1) / 2) /
                 (positive * negative))


def reliability_bins(probabilities, y):
    ids = np.minimum((15 * probabilities).astype(int), 14)
    rows = []
    for bin_id in range(15):
        selected = ids == bin_id
        count = int(selected.sum())
        if count:
            rows.append({'bin': bin_id, 'count': count,
                         'mean_probability': float(probabilities[selected].mean()),
                         'positive_fraction': float(y[selected].mean())})
    return rows


def classification_metrics(logits, y):
    probabilities = sigmoid(logits)
    bins = reliability_bins(probabilities, y)
    return {
        'auc': auc_from_scores(logits, y),
        'negative_log_likelihood': negative_log_likelihood(logits, y),
        'brier_score': float(np.mean((probabilities - y) ** 2)),
        'binary_ece_15_bins': float(sum(row['count'] / len(y) *
                                      abs(row['mean_probability'] - row['positive_fraction'])
                                      for row in bins)),
        'accuracy_at_0_5': float(np.mean((probabilities >= 0.5) == y)),
        'reliability_bins': bins,
    }


def calibration_study():
    rng = np.random.default_rng(SEED + 1)
    validation_logits, validation_y = synthetic_classifier_data(rng, 4000)
    test_logits, test_y = synthetic_classifier_data(rng, 20000)
    temperatures = np.linspace(0.5, 5.0, 901)
    losses = [negative_log_likelihood(validation_logits / temperature, validation_y)
              for temperature in temperatures]
    temperature = float(temperatures[int(np.argmin(losses))])
    raw = classification_metrics(test_logits, test_y)
    calibrated = classification_metrics(test_logits / temperature, test_y)
    if abs(raw['auc'] - calibrated['auc']) > 1e-14:
        raise RuntimeError('Positive temperature scaling must preserve score ordering')
    return {
        'description': 'Synthetic binary classification with deliberately overconfident scores',
        'seed': SEED + 1,
        'validation_count': 4000,
        'test_count': 20000,
        'true_logit_formula': '1.2*x0 - 0.9*x1 + 0.7*x0*x1; x0,x1 independent N(0,1)',
        'raw_score_multiplier': 2.5,
        'temperature_grid': {'minimum': 0.5, 'maximum': 5.0, 'points': 901},
        'fitted_temperature': temperature,
        'raw': raw,
        'calibrated': calibrated,
    }


def write_csv(path, rows):
    with path.open('w', newline='', encoding='utf-8') as file:
        writer = csv.DictWriter(file, fieldnames=list(rows[0]))
        writer.writeheader()
        writer.writerows(rows)


def figures(output, shots, calibration):
    plt.rcParams.update({'font.family': 'DejaVu Sans', 'font.size': 11,
                         'axes.spines.top': False, 'axes.spines.right': False,
                         'svg.hashsalt': 'Maurizio-Viviani-public-research'})
    metadata = {'Creator': AUTHOR, 'Date': '2026-10-11'}
    fig, ax = plt.subplots(figsize=(8.4, 4.7), layout='constrained')
    ns = [row['shots'] for row in shots['rows']]
    ax.loglog(ns, [row['predicted_standard_error'] for row in shots['rows']],
              color='#173f37', marker='o', label='Exact standard error: sqrt(variance / shots)')
    ax.loglog(ns, [row['rmse'] for row in shots['rows']], color='#7f9565',
              marker='s', linestyle='--', label='Empirical RMSE: 2,000 repeated estimates')
    ax.set(xlabel='Measurement shots per estimate', ylabel='Error in expected cut score',
           title='Fixed-angle QAOA: precision follows the sampling budget')
    ax.set_xticks(ns, [f'{n:,}' for n in ns])
    ax.grid(alpha=0.18, which='both')
    ax.legend(fontsize=9, frameon=False)
    fig.savefig(output / 'qaoa-shot-budget.svg', metadata=metadata)
    fig.savefig(output / 'qaoa-shot-budget.png', dpi=160)
    plt.close(fig)
    fig, ax = plt.subplots(figsize=(8.4, 4.7), layout='constrained')
    ax.plot([0, 1], [0, 1], color='#b2beb4', linestyle=':', label='Perfect calibration')
    for key, label, color, marker in [('raw', 'Raw overconfident scores', '#9b785f', 's'),
                                      ('calibrated', 'Validation-fitted temperature', '#173f37', 'o')]:
        rows = calibration[key]['reliability_bins']
        ax.plot([r['mean_probability'] for r in rows], [r['positive_fraction'] for r in rows],
                color=color, marker=marker, label=label)
    ax.set(xlabel='Mean predicted positive-class probability',
           ylabel='Observed positive-class fraction',
           title='Same ranking, different probability reliability', xlim=(0, 1), ylim=(0, 1))
    ax.grid(alpha=0.18)
    ax.legend(fontsize=9, frameon=False)
    fig.savefig(output / 'ai-calibration.svg', metadata=metadata)
    fig.savefig(output / 'ai-calibration.png', dpi=160)
    plt.close(fig)


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output-dir', type=Path, default=Path('results'))
    args = parser.parse_args()
    args.output_dir.mkdir(parents=True, exist_ok=True)
    shots = shot_study()
    calibration = calibration_study()
    report = {'author': AUTHOR, 'date': '2026-10-11',
              'environment': {'python': platform.python_version(), 'numpy': np.__version__,
                              'matplotlib': matplotlib.__version__},
              'qaoa': shots, 'calibration': calibration}
    (args.output_dir / 'research-results.json').write_text(json.dumps(report, indent=2) + '\n')
    write_csv(args.output_dir / 'qaoa-shot-budget.csv', shots['rows'])
    write_csv(args.output_dir / 'ai-calibration.csv', [
        {'model': key, **{k: v for k, v in calibration[key].items() if k != 'reliability_bins'}}
        for key in ['raw', 'calibrated']])
    figures(args.output_dir, shots, calibration)
    print(json.dumps({'qaoa_expectation': shots['exact_expectation'],
                      'qaoa_variance': shots['exact_variance'],
                      'shot_rows': shots['rows'],
                      'fitted_temperature': calibration['fitted_temperature'],
                      'classification_metrics': {key: {k: v for k, v in calibration[key].items()
                                                       if k != 'reliability_bins'}
                                                 for key in ['raw', 'calibrated']}}, indent=2))


if __name__ == '__main__':
    main()
