Buckets:
| #!/usr/bin/env python3 | |
| """Claim 5 / Corollary 4.10: adaptation benefit vs task variance (reduced-scale GPU run).""" | |
| from __future__ import annotations | |
| import json | |
| import os | |
| import sys | |
| from pathlib import Path | |
| import numpy as np | |
| import torch | |
| import yaml | |
| 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, MAMLTrainer # 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 main(): | |
| device = torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| out_env = os.environ.get("OUT_ROOT") | |
| out_root = Path(out_env) / "claim5" if out_env else ( | |
| Path("/mnt/outputs/claim5") if Path("/mnt").exists() else (ROOT / "outputs" / "claim5") | |
| ) | |
| out_root.mkdir(parents=True, exist_ok=True) | |
| print(f"device={device}") | |
| cfg_path = ROOT / "configs" / "experiment_config_gamma.yaml" | |
| with open(cfg_path) as f: | |
| config = yaml.safe_load(f) | |
| n_iterations = int(os.environ.get("CLAIM5_ITERS", "300")) | |
| tasks_per_batch = int(os.environ.get("CLAIM5_BATCH", "8")) | |
| n_segments = int(config.get("n_segments", 60)) | |
| hidden_dim = int(config.get("hidden_dim", 128)) | |
| 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()) | |
| task_dist = create_gamma_task_distribution(config) | |
| print(f"train sigma_tau^2 = {task_dist.compute_variance():.6f}") | |
| policy = GammaPulsePolicy( | |
| task_feature_dim=3, | |
| hidden_dim=hidden_dim, | |
| n_hidden_layers=2, | |
| n_segments=n_segments, | |
| n_controls=2, | |
| ).to(device) | |
| maml = MAML( | |
| policy=policy, | |
| inner_lr=float(config.get("inner_lr", 0.01)), | |
| inner_steps=int(config.get("inner_steps", 5)), | |
| meta_lr=float(config.get("meta_lr", 0.001)), | |
| first_order=True, | |
| device=device, | |
| ) | |
| loss_fn = create_gamma_loss_function(target_state, device, config) | |
| def data_generator_wrapper(task_params, n_trajectories, split): | |
| return gamma_data_generator(task_params, n_trajectories, split, device) | |
| rng = np.random.default_rng(42) | |
| trainer = MAMLTrainer( | |
| maml=maml, | |
| task_sampler=lambda n, split: gamma_task_sampler(n, split, task_dist, rng), | |
| data_generator=data_generator_wrapper, | |
| loss_fn=loss_fn, | |
| n_support=min(4, int(config.get("n_support", 10))), | |
| n_query=min(4, int(config.get("n_query", 10))), | |
| log_interval=50, | |
| val_interval=50, | |
| ) | |
| ckpt = out_root / "maml_gamma_pauli_x_claim5.pt" | |
| print(f"Training {n_iterations} iters, batch={tasks_per_batch}, segments={n_segments}") | |
| trainer.train( | |
| n_iterations=n_iterations, | |
| tasks_per_batch=tasks_per_batch, | |
| val_tasks=5, | |
| save_path=str(ckpt), | |
| ) | |
| diversity_scales = [0.05, 0.1, 0.25, 0.5, 1.0] | |
| K_budget = 10 | |
| rows = [] | |
| for scale in diversity_scales: | |
| cfg = dict(config) | |
| cfg["diversity_scale"] = scale | |
| dist = create_gamma_task_distribution(cfg) | |
| sigma2 = float(dist.compute_variance()) | |
| tasks = gamma_task_sampler(20, "test", dist, np.random.default_rng(1000 + int(scale * 100))) | |
| gaps = [] | |
| for task in tasks: | |
| support = gamma_data_generator(task, 1, "test", device) | |
| query = gamma_data_generator(task, 1, "test", device) | |
| task_data = {"support": support, "query": query} | |
| with torch.no_grad(): | |
| pre = loss_fn(maml.policy, query).item() | |
| adapted, _ = maml.inner_loop(task_data, loss_fn, num_steps=K_budget) | |
| with torch.no_grad(): | |
| post = loss_fn(adapted, query).item() | |
| gaps.append(pre - post) | |
| mean_gap = float(np.mean(gaps)) | |
| rows.append( | |
| { | |
| "diversity_scale": scale, | |
| "sigma_tau2": sigma2, | |
| "mean_gap_K10": mean_gap, | |
| "negligible": abs(mean_gap) < 0.01 or sigma2 < 0.002, | |
| } | |
| ) | |
| print( | |
| f"scale={scale:.2f} sigma2={sigma2:.6f} gap@K10={mean_gap:.6f} " | |
| f"negligible={rows[-1]['negligible']}", | |
| flush=True, | |
| ) | |
| full_tasks = gamma_task_sampler(20, "test", task_dist, np.random.default_rng(999)) | |
| K_values = list(range(0, 21, 2)) | |
| G = [] | |
| for K in K_values: | |
| gaps = [] | |
| for task in full_tasks: | |
| support = gamma_data_generator(task, 1, "test", device) | |
| query = gamma_data_generator(task, 1, "test", device) | |
| task_data = {"support": support, "query": query} | |
| with torch.no_grad(): | |
| pre = loss_fn(maml.policy, query).item() | |
| if K == 0: | |
| gaps.append(0.0) | |
| continue | |
| adapted, _ = maml.inner_loop(task_data, loss_fn, num_steps=K) | |
| with torch.no_grad(): | |
| post = loss_fn(adapted, query).item() | |
| gaps.append(pre - post) | |
| G.append(float(np.mean(gaps))) | |
| print(f"full K={K} gap={G[-1]:.6f}", flush=True) | |
| popt, _ = curve_fit( | |
| exponential_saturation, | |
| np.array(K_values, float), | |
| np.array(G, float), | |
| p0=[max(max(G), 1e-4), 0.1], | |
| bounds=([0, 0], [1, 5]), | |
| ) | |
| c_fit, beta_fit = map(float, popt) | |
| result = { | |
| "device": str(device), | |
| "n_iterations": n_iterations, | |
| "tasks_per_batch": tasks_per_batch, | |
| "checkpoint": str(ckpt), | |
| "variance_sweep": rows, | |
| "fit_full_diversity": {"c": c_fit, "beta": beta_fit, "K_values": K_values, "G_K": G}, | |
| "corollary_check": { | |
| "low_variance_negligible": any( | |
| r["negligible"] and r["diversity_scale"] <= 0.25 for r in rows | |
| ), | |
| "beta_vs_1_over_K": { | |
| "beta": beta_fit, | |
| "one_over_K10": 0.1, | |
| "beta_much_less": beta_fit < 0.05, | |
| }, | |
| }, | |
| } | |
| out_path = out_root / "claim5_results.json" | |
| out_path.write_text(json.dumps(result, indent=2)) | |
| print(json.dumps(result["corollary_check"], indent=2)) | |
| print(f"Wrote {out_path}") | |
| if __name__ == "__main__": | |
| main() | |
Xet Storage Details
- Size:
- 6.64 kB
- Xet hash:
- f3f4e6a58b5d2127e50da4346376a73847fed32aabd7591bf2cda8c8c3d786e9
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.