Buckets:
| #!/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.