Tatopenn's picture
Sync dashboard_core/ from v8.1.33
d2486ac verified
Raw
History Blame Contribute Delete
10.9 kB
"""
VQE / QM-MM telemetry: real JAX-autodiff VQE plus the synthetic mock
fallback used when a circuit has no parametric gates.
Split out of the former monolithic dashboard_core.py (Phase 1 of the
dashboard refactor) -- pure move, no behavior change.
"""
import hashlib
from typing import Tuple
import numpy as np
import pandas as pd
import jax
import jax.numpy as jnp
import dense_evolution as de
QM_MM_HEAVY_QUBIT_THRESHOLD = 12
class QMMMForceEngine:
"""Hellmann-Feynman QM/MM force engine via JAX autodiff. Adapted from dash.py:1455."""
def __init__(self, simulator_instance):
self.sim = simulator_instance
self.dim = simulator_instance.dim
self.n_qubits = simulator_instance.n
def build_loss_function(self):
def qm_mm_energy_loss(classical_positions: jnp.ndarray,
classical_charges: jnp.ndarray,
orbital_centers: jnp.ndarray,
h_pq_core: jnp.ndarray,
statevector: jnp.ndarray) -> jnp.ndarray:
def single_orbital_v(r_orb):
r_diff = classical_positions - r_orb
distanze = jnp.linalg.norm(r_diff, axis=1)
distanze_protette = jnp.where(distanze < 0.8, 0.8, distanze)
return -jnp.sum(classical_charges / distanze_protette)
v_esterno = jax.vmap(single_orbital_v)(orbital_centers)
matrice_v = jnp.diag(v_esterno)
h_pq_perturbed = h_pq_core + matrice_v
energy_eval = jnp.real(jnp.dot(jnp.conj(statevector), jnp.dot(h_pq_perturbed, statevector)))
return energy_eval
return qm_mm_energy_loss
def compute_forces(self, classical_positions: jnp.ndarray,
classical_charges: jnp.ndarray,
orbital_centers: jnp.ndarray,
h_pq_core: jnp.ndarray,
statevector: jnp.ndarray) -> Tuple[jnp.ndarray, jnp.ndarray]:
loss_fun = self.build_loss_function()
grad_fun = jax.jit(jax.value_and_grad(loss_fun, argnums=0))
energia, gradiente_posizioni = grad_fun(
classical_positions, classical_charges, orbital_centers, h_pq_core, statevector
)
forze_mm = -gradiente_posizioni
return energia, forze_mm
def _run_vqe_mock_simulation(epochs: int, lr: float, beta1: float, beta2: float,
nome_circuito: str = "Custom", on_epoch=None) -> pd.DataFrame:
"""Hash-seeded synthetic VQE trajectory, used when a circuit has no parametric gates.
Adapted from dash.py:2234 (canonical version wired into ottimizza_vqe).
`on_epoch`, if given: see run_vqe_telemetry's docstring — same purely-additive contract."""
data = {
"Step": [], "VQE_Energy": [], "Entropy": [], "Purity": [],
"Gradient": [], "Noise_Factor": [], "Theta_Correction": []
}
hash_seed = int(hashlib.md5(nome_circuito.encode('utf-8')).hexdigest(), 16) % 10000
np.random.seed(hash_seed)
target_energy = -1.2 - np.random.uniform(0.1, 0.8)
complexity = 1.0 + np.random.uniform(0.1, 0.7)
barren_step = np.random.choice([0, 20, 35, 50], p=[0.4, 0.2, 0.2, 0.2])
plateau_len = np.random.randint(15, 30)
energy = -0.5 + np.random.uniform(-0.3, 0.3)
m_g, v_g = 0.0, 0.0
epsilon = 1e-8
for epoch in range(epochs):
data["Step"].append(epoch)
if barren_step > 0 and barren_step <= epoch <= (barren_step + plateau_len):
grad_val = np.random.uniform(-0.003, 0.003)
noise = 0.012
else:
grad_val = 0.055 * (energy - target_energy) * complexity + np.random.uniform(-0.005, 0.005)
noise = 0.003
m_g = beta1 * m_g + (1 - beta1) * grad_val
v_g = beta2 * v_g + (1 - beta2) * (grad_val ** 2)
m_hat = m_g / (1 - beta1 ** (epoch + 1))
v_hat = v_g / (1 - beta2 ** (epoch + 1))
step_update = lr * m_hat / (np.sqrt(v_hat) + epsilon)
if epoch < 5:
energy -= step_update * 2.5 + np.random.uniform(-noise * 2, noise * 2)
else:
energy -= step_update * 1.75 + np.random.uniform(-noise, noise)
data["VQE_Energy"].append(float(energy))
data["Entropy"].append(float(0.45 * np.exp(-epoch / 35) + 0.08 + np.random.uniform(-0.01, 0.01)))
data["Purity"].append(float(0.72 + 0.22 * (1 - np.exp(-epoch / 45)) + np.random.uniform(-0.01, 0.01)))
data["Gradient"].append(float(grad_val))
data["Noise_Factor"].append(float(1.0 - (epoch * 0.04 / epochs) + np.random.uniform(-0.002, 0.002)))
if epoch < 5:
data["Theta_Correction"].append(float(step_update * np.cos(epoch * 0.22 * complexity) * 2 + np.random.uniform(-0.01, 0.01)))
else:
data["Theta_Correction"].append(float(step_update * np.cos(epoch * 0.22 * complexity)))
if on_epoch is not None:
on_epoch(epoch, epochs, {col: values[-1] for col, values in data.items()})
df_vqe = pd.DataFrame(data)
df_vqe.set_index("Step", inplace=True)
return df_vqe
def run_vqe_telemetry(sim, parser, qasm_text, circuit_name, n_qubits, use_float32,
epochs, lr, beta1, beta2, seed, hamiltonian_values=None,
on_epoch=None) -> pd.DataFrame:
"""Adapted from ottimizza_vqe (dash.py:3194, canonical/later definition).
Runs real JAX-autodiff VQE (via QMMMForceEngine Hellmann-Feynman forces) if the
circuit has parametric gates, else falls back to _run_vqe_mock_simulation — this
branch is unchanged from the original. `hamiltonian_values`, if given, must be an
array of length 2**n_qubits (custom Hamiltonian); otherwise a random one is used,
replacing the original's globals()-based custom-Hamiltonian widget lookup.
`on_epoch`, if given, is called as `on_epoch(epoch: int, total_epochs: int, row: dict)`
once per completed epoch — purely additive, default None, no behavior change for
existing callers. Lets a caller (e.g. a Streamlit progress bar) observe real
per-epoch state; this function otherwise runs its whole loop internally and would
only return the final DataFrame.
"""
previous_x64_state = jax.config.jax_enable_x64
jax.config.update('jax_enable_x64', not use_float32)
try:
return _run_vqe_telemetry_body(
sim, parser, qasm_text, circuit_name, n_qubits, use_float32,
epochs, lr, beta1, beta2, seed, hamiltonian_values, on_epoch,
)
finally:
# `sim` was built under a precision fixed at its own creation (run_simulation);
# this call reuses that same sim, so jax_enable_x64 must match for its whole
# duration, not whatever was left behind by unrelated code in between.
jax.config.update('jax_enable_x64', previous_x64_state)
def _run_vqe_telemetry_body(sim, parser, qasm_text, circuit_name, n_qubits, use_float32,
epochs, lr, beta1, beta2, seed, hamiltonian_values=None,
on_epoch=None) -> pd.DataFrame:
circ_obj = parser.parse(qasm_text)
energy_fn, n_params = de.circuit_to_energy_fn(circ_obj, n_qubits)
if n_params == 0:
return _run_vqe_mock_simulation(epochs=epochs, lr=lr, beta1=beta1, beta2=beta2,
nome_circuito=circuit_name, on_epoch=on_epoch)
theta = np.random.uniform(-np.pi, np.pi, n_params)
m, v = np.zeros(n_params), np.zeros(n_params)
try:
engine = QMMMForceEngine(sim)
except Exception:
engine = None
np.random.seed(seed)
classical_dtype = jnp.float32 if use_float32 else jnp.float64
if hamiltonian_values is not None and len(hamiltonian_values) == 2 ** n_qubits:
valori_energetici = np.array(hamiltonian_values, dtype=classical_dtype)
else:
valori_energetici = np.sort(np.random.uniform(-2.5, 2.5, 2 ** n_qubits)).astype(classical_dtype)
sim.H_matrix = jnp.diag(jnp.array(valori_energetici, dtype=classical_dtype))
history = []
stato_zero_dtype = jnp.complex64 if use_float32 else jnp.complex128
stato_zero = jnp.zeros(2 ** n_qubits, dtype=stato_zero_dtype).at[0].set(1.0)
classical_positions = jnp.array([[0.0, 0.0, 0.0], [1.4, 0.0, 0.0]], dtype=classical_dtype)
classical_charges = jnp.array([1.0, -1.0], dtype=classical_dtype)
orbital_centers = jnp.array([[0.0, 0.0, 0.1]], dtype=classical_dtype)
# Built once, reused every epoch — only theta changes, so this traces/
# JIT-compiles a single time instead of rebuilding the circuit (and
# calling the trace-severing risolvi_qasm) from scratch each epoch.
energy_and_grad = jax.jit(jax.value_and_grad(energy_fn, argnums=0, has_aux=True))
for epoch in range(epochs):
(energia_jax, sv), grad_jax = energy_and_grad(
jnp.asarray(theta, dtype=jnp.float64), sim.H_matrix, stato_zero
)
energia = float(energia_jax)
prob = np.clip(np.abs(np.asarray(sv)) ** 2, 0.0, 1.0)
prob_total = prob.sum()
if prob_total > 1e-12:
prob = prob / prob_total
p_safe = prob[prob > 1e-15]
entropia = float(-np.sum(p_safe * np.log2(p_safe))) if len(p_safe) > 0 else 0.0
purita = float(np.sum(prob ** 2))
norma_forze_mm = 0.0
if engine is not None:
try:
_, forze_mm = engine.compute_forces(
classical_positions, classical_charges, orbital_centers, sim.H_matrix, sv
)
norma_forze_mm = float(jnp.linalg.norm(forze_mm))
except Exception:
pass
# Real gradient (jax.grad through _vqe_energy_fn), not a formula —
# see CHANGELOG for what used to be here.
grad_vqe_params = np.asarray(grad_jax)
norm_grad_vqe_params = float(np.linalg.norm(grad_vqe_params))
t = epoch + 1
m = beta1 * m + (1 - beta1) * grad_vqe_params
v = beta2 * v + (1 - beta2) * (grad_vqe_params ** 2)
m_hat = m / (1.0 - beta1 ** t)
v_hat = v / (1.0 - beta2 ** t)
theta_correction_step_raw = (lr / (np.sqrt(v_hat) + 1e-8)) * m_hat
theta -= theta_correction_step_raw
norm_theta_correction_step = float(np.linalg.norm(theta_correction_step_raw))
row = {
"Step": epoch,
"VQE_Energy": energia,
"Entropy": entropia,
"Purity": purita,
"Gradient": norm_grad_vqe_params,
"Noise_Factor": 0.015 * (1.0 - (purita * 0.1)),
"Theta_Correction": norm_theta_correction_step,
}
history.append(row)
if on_epoch is not None:
on_epoch(epoch, epochs, row)
return pd.DataFrame(history).set_index("Step")