Sor0ush's picture
download
raw
6.64 kB
#!/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.