#!/usr/bin/env python3
"""
Capitolo 4: l'eco di Loschmidt dell'operatore (OLE), la grandezza misurata
nell'esperimento di Algorithmiq e IBM, calcolata esattamente su una catena
piccola con la propagazione di Pauli senza troncamento.

Definizione (README del circuito nel Quantum Advantage Tracker):
    f_delta(O) = 2^{-n} Tr( U O U^† V^† U O U^† V ),   V = exp(-i delta G),
    G = somma di X sui siti del sottoinsieme P,  U = (U_{b-alpha}^†)^L U_b^L,
    U_b = prod_{(u,v)} exp(-i pi/4 Z_u Z_v) exp(-i pi/8 (Z_u+Z_v)) exp(-i(b_u X_u + b_v X_v)).

Qui la catena è aperta, con N = 10 qubit; i siti F (b_u = 3 pi/8) sono quelli
pari, gli S (b_u = b) i dispari; l'osservabile è O = Z_4 Z_5 (peso 2), la
perturbazione agisce sui qubit 7 e 8 con delta = 0,3. Con alpha = 0 l'eco è
perfetto (U = identità) e f = 1: il segnale nasce dalla diffusione alpha.

Il calcolo di produzione è a matrici dense (2^10 = 1024); la propagazione di
Pauli senza troncamento, che è esatta, lo verifica per N = 6.

Produce:
    c04_ole_L.dat       f_delta in funzione di L per b = 0,125 e 0,25 (alpha = 0,15)
    c04_ole_delta.dat   1 - f_delta in funzione di delta, L = 4, con la parabola OTOC
    c04_ole_alpha.dat   f_delta in funzione del potenziale di diffusione alpha, L = 2, 4, 6
"""

import math

import numpy as np

from comune import catena, colorazione_lati, coniuga_rotazione, scrivi_dat

PI = math.pi


def strati_ub(n, b, alpha=0.0, inverso=False):
    """Lista ordinata di rotazioni (generatore, angolo) che compongono U_b
    (convenzione R_P(theta) = exp(-i theta P/2)).

    exp(-i pi/4 ZZ) = R_ZZ(pi/2); exp(-i pi/8 Z) = R_Z(pi/4);
    exp(-i b X) = R_X(2b). Per l'inverso si invertono ordine e segni."""
    bu = [3 * PI / 8 if u % 2 == 0 else b - alpha for u in range(n)]
    seq = []
    for strato in colorazione_lati(catena(n)):
        for (u, v) in strato:
            # prodotto di operatori: agisce per primo il fattore più a destra
            seq.append(((1 << u, 0), 2 * bu[u]))
            seq.append(((1 << v, 0), 2 * bu[v]))
            seq.append(((0, 1 << u), PI / 4))
            seq.append(((0, 1 << v), PI / 4))
            seq.append(((0, (1 << u) | (1 << v)), PI / 2))
    if inverso:
        seq = [(g, -a) for (g, a) in reversed(seq)]
    return seq


def sequenza_u(n, b, L, alpha):
    """U = (U_{b-alpha}^†)^L U_b^L come lista di rotazioni nell'ordine di applicazione."""
    seq = []
    for _ in range(L):
        seq += strati_ub(n, b)
    for _ in range(L):
        seq += strati_ub(n, b, alpha=alpha, inverso=True)
    return seq


def u_o_udag(op, seq):
    """A = U O U^†. Con U = R_k ... R_1, A = R_k ... R_1 O R_1^† ... R_k^†:
    si coniuga con R^† (angolo opposto) partendo dalla prima rotazione."""
    for g, a in seq:
        op, _ = coniuga_rotazione(op, g, -a)
    return op


def ole(n, b, L, delta, q_oss, siti_p, alpha=0.0):
    A = u_o_udag({(0, 1 << q_oss): 1.0}, sequenza_u(n, b, L, alpha))
    B = dict(A)
    # V^† A V con V = exp(-i delta X) = R_X(2 delta) su ogni sito di P
    for s in siti_p:
        B, _ = coniuga_rotazione(B, (1 << s, 0), 2 * delta)
    # 2^{-n} Tr(A B) = somma dei prodotti dei coefficienti (Pauli ortonormali)
    return sum(c * B.get(p, 0.0) for p, c in A.items()), len(A)



# ------------------------------------------------ calcolo a matrici dense
def _applica(M, n, gen, theta):
    """R_G(theta) M per una rotazione a uno o due qubit (X, Z, ZZ) su M di forma (2^n, m)."""
    x, z = gen
    idx = np.arange(2 ** n)
    if x == 0:   # rotazione diagonale
        par = np.zeros(2 ** n, dtype=int)
        for k in range(n):
            if (z >> k) & 1:
                par ^= (idx >> k) & 1
        fase = np.exp(-1j * theta / 2 * (1 - 2 * par))
        return fase[:, None] * M
    q = x.bit_length() - 1
    c, s = math.cos(theta / 2), math.sin(theta / 2)
    flip = idx ^ (1 << q)
    return c * M - 1j * s * M[flip, :]


def matrice_u(n, seq):
    U = np.eye(2 ** n, dtype=complex)
    for g, a in seq:
        U = _applica(U, n, g, a)
    return U


def ole_denso(n, b, L, delta, q_oss, siti_p, alpha=0.0, U=None):
    if U is None:
        U = matrice_u(n, sequenza_u(n, b, L, alpha))
    idx = np.arange(2 ** n)
    zq = (1 - 2 * ((idx >> q_oss) & 1)).astype(float)
    A = (U * zq[None, :]) @ U.conj().T
    V = np.eye(2 ** n, dtype=complex)
    for s_ in siti_p:
        V = _applica(V, n, (1 << s_, 0), 2 * delta)
    return float(np.real(np.trace(A @ V.conj().T @ A @ V))) / 2 ** n


for (bb, LL, dd, aa) in ((0.25, 2, 0.3, 0.1), (0.125, 1, 0.5, 0.0)):
    fp, _ = ole(6, bb, LL, dd, 2, (5,), aa)
    fd = ole_denso(6, bb, LL, dd, 2, (5,), aa)
    print(f"verifica N=6, b={bb}, L={LL}: Pauli {fp:.12f}  denso {fd:.12f}  diff {abs(fp-fd):.1e}", flush=True)
    assert abs(fp - fd) < 1e-10

N = 10
O = (4, 5)        # osservabile Z_4 Z_5
P = (7, 8)        # siti della perturbazione
DELTA = 0.3       # come nell'istanza 56x1488


def ole_n(b, L, delta, alpha, U=None):
    if U is None:
        U = matrice_u(N, sequenza_u(N, b, L, alpha))
    idx = np.arange(2 ** N)
    par = np.zeros(2 ** N, dtype=int)
    for q in O:
        par ^= (idx >> q) & 1
    zq = (1 - 2 * par).astype(float)
    A = (U * zq[None, :]) @ U.conj().T
    V = np.eye(2 ** N, dtype=complex)
    for s_ in P:
        V = _applica(V, N, (1 << s_, 0), 2 * delta)
    return float(np.real(np.trace(A @ V.conj().T @ A @ V))) / 2 ** N


righe = []
for k in range(0, 13):
    a = k * 0.025
    righe.append([a] + [ole_n(0.125, L, DELTA, a) for L in (2, 4, 6)])
scrivi_dat("c04_ole_alpha.dat",
           "alpha f_L2 f_L4 f_L6  (N=10, b=0,125, delta=0,3, O=Z4Z5, P={7,8})", righe)
for r in righe[::2]:
    print("alpha", round(r[0], 3), [round(v, 4) for v in r[1:]], flush=True)

righe = []
for L in range(0, 9):
    righe.append([L] + [ole_n(b, L, DELTA, 0.15) for b in (0.125, 0.25)])
scrivi_dat("c04_ole_L.dat", "L f_b0125 f_b025  (N=10, delta=0,3, alpha=0,15)", righe)
print("L:", [[r[0]] + [round(v, 4) for v in r[1:]] for r in righe], flush=True)

U4 = matrice_u(N, sequenza_u(N, 0.125, 4, 0.15))
eps = 1e-3
curv = (1 - ole_n(0.125, 4, eps, 0.15, U=U4)) / eps ** 2
righe = []
for k in range(0, 21):
    d = k * 0.05
    righe.append([d, 1 - ole_n(0.125, 4, d, 0.15, U=U4), curv * d * d])
scrivi_dat("c04_ole_delta.dat", "delta uno_meno_f parabola_otoc  (N=10, b=0,125, L=4, alpha=0,15)", righe)
print(f"coefficiente (1-f)/delta^2 per delta->0: {curv:.4f}; a delta=0,3: 1-f = {righe[6][1]:.4f}, "
      f"parabola {righe[6][2]:.4f}; a delta=1: 1-f = {righe[20][1]:.4f}, parabola {righe[20][2]:.4f}")
