Buckets:
| 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.