Sor0ush's picture
download
raw
3.96 kB
import sys
import os
import torch
import numpy as np
from pathlib import Path
from scipy.optimize import curve_fit
# Add experiments folder to path
sys.path.insert(0, str(Path(__file__).parent / "experiments"))
from fig_appendix_meta_training.train_meta_gamma import MAMLTrainer, create_gamma_loss_function, gamma_data_generator, gamma_task_sampler, create_gamma_task_distribution
from metaqctrl.meta_rl.maml import MAML
from metaqctrl.meta_rl.policy_gamma import GammaPulsePolicy
from metaqctrl.quantum.gates import TargetGates
def exponential_saturation(K, c, beta):
return c * (1 - np.exp(-beta * K))
def main():
device = torch.device('cpu')
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': False,
}
task_dist = create_gamma_task_distribution(config)
policy = GammaPulsePolicy(
task_feature_dim=3, hidden_dim=128, n_hidden_layers=2, n_segments=20, n_controls=2
).to(device)
maml = MAML(policy=policy, inner_lr=0.01, inner_steps=5, meta_lr=0.001, first_order=False, 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)
trainer = MAMLTrainer(
maml=maml,
task_sampler=lambda n, split: gamma_task_sampler(n, split, task_dist, np.random.default_rng(42)),
data_generator=data_generator_wrapper,
loss_fn=loss_fn,
n_support=1, n_query=1, log_interval=100, val_interval=100
)
print("Training MAML for 300 iterations (fast version)...")
trainer.train(n_iterations=300, tasks_per_batch=4, val_tasks=5, save_path="temp_maml_claim2.pt")
print("Testing adaptation curve...")
maml.policy.eval()
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:
data = data_generator_wrapper(task, 1, 'test')
with torch.no_grad():
pre_loss = loss_fn(maml.policy(data['inputs']), data['targets'], data['masks']).item()
pre_fid = 1.0 - pre_loss
if K == 0:
gaps.append(0.0)
else:
adapted_policy = maml.clone_policy()
# manual adaptation
opt = torch.optim.SGD(adapted_policy.parameters(), lr=0.01)
for _ in range(K):
opt.zero_grad()
loss = loss_fn(adapted_policy(data['inputs']), data['targets'], data['masks'])
loss.backward()
opt.step()
with torch.no_grad():
post_loss = loss_fn(adapted_policy(data['inputs']), data['targets'], data['masks']).item()
post_fid = 1.0 - post_loss
gaps.append(post_fid - pre_fid)
G_K_means.append(np.mean(gaps))
print(f"K={K}, Mean Gap={np.mean(gaps):.4f}")
K_arr = np.array(K_values)
G_arr = np.array(G_K_means)
popt, _ = curve_fit(exponential_saturation, K_arr, G_arr, p0=[0.05, 0.1], bounds=([0, 0], [1, 5]))
c_fit, beta_fit = popt
G_fit = exponential_saturation(K_arr, c_fit, beta_fit)
ss_res = np.sum((G_arr - G_fit) ** 2)
ss_tot = np.sum((G_arr - np.mean(G_arr)) ** 2)
R2 = 1 - ss_res / ss_tot if ss_tot > 0 else 0
print(f"Fit Result: beta = {beta_fit:.4f}, R^2 = {R2:.4f}")
if __name__ == '__main__':
main()

Xet Storage Details

Size:
3.96 kB
·
Xet hash:
b0b06c6fd0e7d557eb54558c9ead751c9b33cee1ed8b68e4c63fe071613c89d8

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