Sor0ush's picture
download
raw
5.56 kB
#!/usr/bin/env python3
"""Evaluate Claim 2 adaptation curve from an existing MAML checkpoint."""
from __future__ import annotations
import json
import sys
from pathlib import Path
import numpy as np
import torch
from scipy.optimize import curve_fit
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
sys.path.insert(0, str(ROOT / "experiments"))
from fig_appendix_meta_training.train_meta_gamma import ( # noqa: E402
create_gamma_loss_function,
create_gamma_task_distribution,
gamma_data_generator,
gamma_task_sampler,
)
from metaqctrl.meta_rl.maml import MAML # noqa: E402
from metaqctrl.meta_rl.policy_gamma import GammaPulsePolicy # noqa: E402
from metaqctrl.quantum.gates import TargetGates # noqa: E402
def exponential_saturation(K, c, beta):
return c * (1 - np.exp(-beta * K))
def load_policy(ckpt_path: Path, device: torch.device) -> GammaPulsePolicy:
ckpt = torch.load(ckpt_path, map_location=device, weights_only=False)
policy = GammaPulsePolicy(
task_feature_dim=3, hidden_dim=128, n_hidden_layers=2, n_segments=20, n_controls=2
).to(device)
state = ckpt.get("policy_state_dict") or ckpt
policy.load_state_dict(state)
policy.eval()
return policy
def main():
device = torch.device("cpu")
out_dir = ROOT / "outputs" / "claim2"
out_dir.mkdir(parents=True, exist_ok=True)
ckpt_candidates = [
ROOT / "temp_maml_claim2_best.pt",
ROOT / "temp_maml_claim2.pt",
]
ckpt_path = next((p for p in ckpt_candidates if p.exists()), None)
if ckpt_path is None:
raise FileNotFoundError("No Claim 2 checkpoint found (temp_maml_claim2*.pt)")
U_target = TargetGates.pauli_x()
ket_0 = np.array([1, 0], dtype=complex)
target_state = np.outer(U_target @ ket_0, (U_target @ ket_0).conj())
config = {
"gamma_deph_range": [0.001, 0.01],
"gamma_relax_range": [0.0005, 0.005],
"inner_lr": 0.01,
"inner_steps": 5,
"meta_lr": 0.001,
"first_order": True,
"gate_time": 1.0,
}
task_dist = create_gamma_task_distribution(config)
loss_fn = create_gamma_loss_function(target_state, device, config)
policy = load_policy(ckpt_path, device)
maml = MAML(
policy=policy,
inner_lr=0.01,
inner_steps=5,
meta_lr=0.001,
first_order=True,
device=device,
)
def make_task_data(task):
support = gamma_data_generator(task, 1, "test", device)
query = gamma_data_generator(task, 1, "test", device)
return {"support": support, "query": query}
test_tasks = gamma_task_sampler(30, "test", task_dist, np.random.default_rng(123))
K_values = list(range(0, 31, 2))
G_K_means = []
for K in K_values:
gaps = []
for task in test_tasks:
task_data = make_task_data(task)
with torch.no_grad():
pre_loss = loss_fn(maml.policy, task_data["query"]).item()
if K == 0:
gaps.append(0.0)
else:
adapted, _ = maml.inner_loop(task_data, loss_fn, num_steps=K)
with torch.no_grad():
post_loss = loss_fn(adapted, task_data["query"]).item()
gaps.append(pre_loss - post_loss) # fidelity gain = loss drop
G_K_means.append(float(np.mean(gaps)))
print(f"K={K}, Mean Gap={np.mean(gaps):.6f}", flush=True)
K_arr = np.array(K_values, dtype=float)
G_arr = np.array(G_K_means, dtype=float)
popt, _ = curve_fit(
exponential_saturation,
K_arr,
G_arr,
p0=[max(float(G_arr.max()), 1e-4), 0.1],
bounds=([0, 0], [1, 5]),
)
c_fit, beta_fit = map(float, popt)
G_fit = exponential_saturation(K_arr, c_fit, beta_fit)
ss_res = float(np.sum((G_arr - G_fit) ** 2))
ss_tot = float(np.sum((G_arr - np.mean(G_arr)) ** 2))
R2 = 1 - ss_res / ss_tot if ss_tot > 0 else 0.0
# Low-variance regime
low_cfg = {
**config,
"gamma_deph_range": [0.0045, 0.0055],
"gamma_relax_range": [0.00225, 0.00275],
}
low_dist = create_gamma_task_distribution(low_cfg)
sigma_tau2 = float(low_dist.compute_variance())
low_tasks = gamma_task_sampler(30, "test", low_dist, np.random.default_rng(7))
low_gaps_k10 = []
for task in low_tasks:
task_data = make_task_data(task)
with torch.no_grad():
pre_loss = loss_fn(maml.policy, task_data["query"]).item()
adapted, _ = maml.inner_loop(task_data, loss_fn, num_steps=10)
with torch.no_grad():
post_loss = loss_fn(adapted, task_data["query"]).item()
low_gaps_k10.append(pre_loss - post_loss)
low_gap = float(np.mean(low_gaps_k10))
result = {
"checkpoint": str(ckpt_path),
"paper_target": {"R2": 0.99, "beta": 0.083, "low_variance_threshold": 0.002},
"fit": {"c": c_fit, "beta": beta_fit, "R2": R2},
"task_variance_train": float(task_dist.compute_variance()),
"K_values": K_values,
"G_K_means": G_K_means,
"low_variance": {
"sigma_tau2": sigma_tau2,
"mean_gap_K10": low_gap,
"negligible": abs(low_gap) < 0.01 or sigma_tau2 < 0.002,
},
}
out_path = out_dir / "claim2_results.json"
out_path.write_text(json.dumps(result, indent=2))
print(json.dumps(result["fit"], indent=2))
print(json.dumps(result["low_variance"], indent=2))
print(f"Wrote {out_path}")
if __name__ == "__main__":
main()

Xet Storage Details

Size:
5.56 kB
·
Xet hash:
faf9312a7cca4dc4c5fff9ae96b82f591f8a60d214ebdb128b778e9864eb0ceb

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.