#!/usr/bin/env python3
"""
Capitolo 11: catena di Markov Monte Carlo con proposta quantistica (Layden
et al., Nature 2023) applicata a modelli di default correlati, confrontata con
avversari classici completi.

Modelli (costruiti, semi dichiarati), n = 10 debitori, stato d in {0,1}^n
(1 = default), spin z_i = 1 - 2 d_i, distribuzione di Boltzmann
p(d) ∝ exp(-E(d)/T) con E(z) = - sum_{i<j} J_ij z_i z_j - sum_i h_i z_i.

  settori     due settori di 5; J_ij = 0,6 nel settore, 0,25 fra settori;
              h_i ~ U(0,8; 1,2). Ising ferromagnetico con campi concordi:
              classe classicamente trattabile (Jerrum-Sinclair 1993). Due
              bacini: quasi nessun default, default di massa.
  frustrato   J_ij ~ N(0,15; 0,6) con segni misti (dipendenze positive e
              negative fra debitori), h_i ~ U(0,2; 0,8). Seme 11.
              Le coppie a segno negativo creano frustrazione: niente
              scorciatoia di tipo «inverti tutto».

Proposte (accettazione di Metropolis classica in tutti i casi):
  locale        inversione di un debitore scelto a caso;
  uniforme      nuovo stato estratto uniformemente;
  loc+glob      90% inversione locale, 10% inversione di tutti i debitori;
  misto         80% locale, 10% inversione globale, 10% uniforme
                (il miglior avversario classico «a una riga» fra quelli provati);
  quantistica   |<d'| exp(-i H t) |d>|^2 con H = (1-g) a H_prob + g sum X_i,
                media su 12 coppie (g, t), g ~ U(0,25; 0,6), t ~ U(2; 20);
  q. rumorosa   proposta quantistica mescolata alla uniforme con peso
                1 - exp(-Lambda): modello depolarizzante globale del circuito.

Velocità: spectral gap delta = 1 - max_{k>=2} |lambda_k| della matrice di
transizione, esatto su 2^10 stati; il relaxation time è 1/delta.

Effetto del readout asimmetrico (sezione 11.x): la proposta effettiva è
Q' = R Q con R prodotto di canali binari per bit (1->0 con prob. a, 0->1 con
prob. b). Q' non è simmetrica; la catena di Metropolis (che assume simmetria)
converge a una distribuzione diversa. Se ne calcola la distribuzione
stazionaria esatta e la probabilità di coda.

Produce:
    c11_gap.dat                log10 delta in funzione di T, modello frustrato, cinque proposte
    c11_tabella_mcmc.txt       1/delta per modello, temperatura e proposta
    c11_tabella_readout.txt    probabilità di coda vera e distorta dal readout asimmetrico
    c11_perdita.dat            distribuzione esatta della perdita a T = 1, modello a settori
"""

import math

import numpy as np

from comune import scrivi_dat, scrivi_txt, it, it_sci

N = 10
DIM = 2 ** N
idx = np.arange(DIM)
D = ((idx[:, None] >> np.arange(N)) & 1).astype(float)
Zs = 1 - 2 * D


def modello_settori():
    rng = np.random.default_rng(31)
    settore = np.array([0] * 5 + [1] * 5)
    J = np.where(settore[:, None] == settore[None, :], 0.6, 0.25)
    np.fill_diagonal(J, 0.0)
    h = rng.uniform(0.8, 1.2, N)
    esp = rng.uniform(0.5, 1.5, N)
    return J, h, esp


def modello_frustrato():
    rng = np.random.default_rng(11)
    A = rng.normal(0.15, 0.6, (N, N))
    J = np.triu(A, 1)
    J = J + J.T
    h = rng.uniform(0.2, 0.8, N)
    esp = rng.uniform(0.5, 1.5, N)
    return J, h, esp


def energie(J, h):
    return -0.5 * np.einsum("si,ij,sj->s", Zs, J, Zs) - Zs @ h


def boltzmann(E, T):
    w = np.exp(-(E - E.min()) / T)
    return w / w.sum()


def metropolis(Q, E, T):
    dE = E[:, None] - E[None, :]
    A = np.minimum(1.0, np.exp(-np.clip(dE / T, -700, 700)))
    P = Q * A
    np.fill_diagonal(P, 0.0)
    P[idx, idx] = 1.0 - P.sum(axis=0)
    return P


def gap(P):
    ev = np.sort(np.abs(np.linalg.eigvals(P)))[::-1]
    return 1.0 - ev[1]


def stazionaria(P):
    w, v = np.linalg.eig(P)
    k = np.argmin(np.abs(w - 1))
    p = np.real(v[:, k])
    return p / p.sum()


# proposte classiche
Q_loc = np.zeros((DIM, DIM))
for i in range(N):
    Q_loc[idx ^ (1 << i), idx] += 1.0 / N
Q_uni = np.full((DIM, DIM), 1.0 / DIM)
Q_glob = np.zeros((DIM, DIM))
Q_glob[idx ^ (DIM - 1), idx] = 1.0
Q_lg = 0.9 * Q_loc + 0.1 * Q_glob
Q_mix = 0.8 * Q_loc + 0.1 * Q_glob + 0.1 * Q_uni

Hmix = np.zeros((DIM, DIM))
for i in range(N):
    Hmix[idx ^ (1 << i), idx] += 1.0


def proposta_quantistica(J, h, E, seme=7):
    rng = np.random.default_rng(seme)
    alfa = math.sqrt(N) / math.sqrt(np.sum(np.triu(J, 1) ** 2) + np.sum(h ** 2))
    Q = np.zeros((DIM, DIM))
    for _ in range(12):
        g = rng.uniform(0.25, 0.6)
        t = rng.uniform(2.0, 20.0)
        H = (1 - g) * alfa * np.diag(E) + g * Hmix
        w, V = np.linalg.eigh(H)
        U = (V * np.exp(-1j * w * t)) @ V.T
        Q += np.abs(U) ** 2 / 12
    assert np.allclose(Q, Q.T, atol=1e-10)
    return Q


def readout(a, b):
    """Matrice R (colonne: stringa letta come partenza ideale) per errori indipendenti per bit."""
    r1 = np.array([[1 - b, a], [b, 1 - a]])   # colonne: bit vero 0, 1; righe: bit letto
    R = np.ones((1, 1))
    for _ in range(N):
        R = np.kron(r1, R)
    return R


modelli = {"settori": modello_settori(), "frustrato": modello_frustrato()}
dati = {}
for nome, (J, h, esp) in modelli.items():
    E = energie(J, h)
    Qq = proposta_quantistica(J, h, E)
    dati[nome] = (J, h, esp, E, Qq)

# ------------------------------------------------ spectral gap in funzione di T, modello frustrato
J, h, esp, E, Qq = dati["frustrato"]
righe = []
for k in range(0, 17):
    T = 10 ** (math.log10(0.3) + k * (math.log10(5.0) - math.log10(0.3)) / 16)
    g = [gap(metropolis(Q, E, T)) for Q in (Q_loc, Q_uni, Q_mix, Qq,
                                              math.exp(-2) * Qq + (1 - math.exp(-2)) * Q_uni)]
    righe.append([T] + [math.log10(max(x, 1e-300)) for x in g])
    print(f"frustrato T={T:.3f}: 1/delta locale {1/g[0]:.3g}, uniforme {1/g[1]:.3g}, "
          f"misto {1/g[2]:.3g}, quantistica {1/g[3]:.3g}, q. rumorosa (Lambda=2) {1/g[4]:.3g}", flush=True)
scrivi_dat("c11_gap.dat",
           "T log10_delta_locale log10_delta_uniforme log10_delta_misto log10_delta_quantistica "
           "log10_delta_q_rumorosa_L2  (modello frustrato, n=10)", righe)

# ------------------------------------------------ tabella riassuntiva
tab = []
for nome in ("settori", "frustrato"):
    J, h, esp, E, Qq = dati[nome]
    for T in (0.5, 1.0):
        val = []
        for Q in (Q_loc, Q_lg, Q_mix, Qq, math.exp(-2) * Qq + (1 - math.exp(-2)) * Q_uni):
            val.append(1 / gap(metropolis(Q, E, T)))
        tab.append([nome, it(T, 1)] + [it_sci(v, 1) if v >= 1000 else it(v, 1) for v in val])
        print(f"{nome} T={T}: 1/delta locale {val[0]:.4g}, loc+glob {val[1]:.4g}, misto {val[2]:.4g}, "
              f"quantistica {val[3]:.4g}, q. rumorosa {val[4]:.4g}", flush=True)
scrivi_txt("c11_tabella_mcmc.txt",
           ["modello", "T", "locale", "locale + globale", "misto classico", "quantistica ideale",
            "quantistica rumorosa (Lambda = 2)"], tab,
           "1/delta (relaxation time in passi), n = 10; misto = 80% locale, 10% globale, 10% uniforme")

# ------------------------------------------------ readout asimmetrico: distorsione della coda
tab = []
for nome in ("settori", "frustrato"):
    J, h, esp, E, Qq = dati[nome]
    perdita = D @ (esp * 0.6)
    soglia = 0.5 * perdita.max()
    T = 1.0
    vera = float(boltzmann(E, T)[perdita > soglia].sum())
    for a, b in ((0.02, 0.01), (0.05, 0.01)):
        Qn = readout(a, b) @ Qq
        P = metropolis(Qn, E, T)
        pis = stazionaria(P)
        dist = float(pis[perdita > soglia].sum())
        tv = 0.5 * float(np.abs(pis - boltzmann(E, T)).sum())
        tab.append([nome, f"{it(100*a, 0)}% / {it(100*b, 0)}%", it_sci(vera, 2), it_sci(dist, 2),
                    it(dist / vera, 2), it_sci(tv, 1)])
        print(f"readout {nome} a={a} b={b}: coda vera {vera:.3e}, distorta {dist:.3e}, rapporto {dist/vera:.3f}, TV {tv:.2e}")
scrivi_txt("c11_tabella_readout.txt",
           ["modello", "errore di readout 1->0 / 0->1", "P(coda) vera", "P(coda) della catena",
            "rapporto", "distanza in variazione totale"], tab,
           "T = 1; coda = perdita oltre metà del massimo; proposta quantistica ideale seguita da readout asimmetrico per bit")

# ------------------------------------------------ distribuzione di perdita, modello a settori
J, h, esp, E, Qq = dati["settori"]
perdita = D @ (esp * 0.6)
p = boltzmann(E, 1.0)
bordi = np.linspace(0, perdita.max() + 1e-9, 21)
hist = np.histogram(perdita, bins=bordi, weights=p)[0]
scrivi_dat("c11_perdita.dat", "perdita probabilita  (modello a settori, T=1; classi di ampiezza costante)",
           [[0.5 * (bordi[i] + bordi[i + 1]), max(hist[i], 1e-12)] for i in range(20)])
print(f"settori: perdita massima {perdita.max():.3f}, P(coda) a T=1 {float(p[perdita > 0.5*perdita.max()].sum()):.3e}")
J, h, esp, E, Qq = dati["frustrato"]
print("frustrato: coppie a segno negativo", int(np.sum(np.triu(J, 1) < 0)), "su", N * (N - 1) // 2)
