| import csv |
| import os |
| import time |
| import numpy as np |
| import torch |
| import ant_hd_ppo as a |
| DEFAULT_D = 512 |
| BETA_BASE = 0.8 |
| BETA_EFF_MIN_MULT = 0.2 |
| BETA_EFF_MAX_MULT = 4.0 |
| G_EMA_DECAY = 0.99 |
|
|
| class HDEncoderGradAdaptiveHybrid: |
|
|
| def __init__(self, feat_lo, feat_hi, D, seed, feature_fn, beta_base, phi_init=None): |
| self.lo = np.array(feat_lo, np.float32) |
| self.hi = np.array(feat_hi, np.float32) |
| self.feature_fn = feature_fn |
| self.beta_base = float(beta_base) |
| self.D = D |
| self.n_feat = len(feat_lo) |
| self.sqrtD = float(np.sqrt(D)) |
| if phi_init is not None: |
| self.Phi = np.asarray(phi_init, dtype=np.float32) |
| else: |
| rng = np.random.default_rng(seed) |
| self.Phi = rng.uniform(-np.pi, np.pi, (self.n_feat, D)).astype(np.float32) |
| scale = 2.0 / (self.hi - self.lo + 1e-08) |
| self.dtheta_ds_unit = self.Phi * scale[:, None] |
| self.W_re = None |
| self.W_im = None |
| self._torch_actor = None |
| self._log_g_ema = 0.0 |
| self.beta_eff_history = [] |
|
|
| def set_weights(self, W_re, W_im): |
| self.W_re = W_re |
| self.W_im = W_im |
|
|
| def link_torch_actor(self, actor): |
| self._torch_actor = actor |
|
|
| def _current_weights(self): |
| if self._torch_actor is not None: |
| return (self._torch_actor.W_re.detach().cpu().numpy(), self._torch_actor.W_im.detach().cpu().numpy()) |
| return (self.W_re, self.W_im) |
|
|
| @property |
| def beta_vec(self): |
| return np.full(self.D, self.beta_base, dtype=np.float32) |
|
|
| def encode(self, state): |
| W_re, W_im = self._current_weights() |
| s = self.feature_fn(state) |
| s = np.clip(s, self.lo, self.hi) |
| s_norm = (2.0 * (s - self.lo) / (self.hi - self.lo + 1e-08) - 1.0).astype(np.float32) |
| proj = s_norm @ self.Phi |
| theta_base = self.beta_base * proj |
| H_re_base, H_im_base = (np.cos(theta_base), np.sin(theta_base)) |
| dtheta_ds = self.beta_base * self.dtheta_ds_unit |
| dH_re_ds = -H_im_base[None, :] * dtheta_ds |
| dH_im_ds = H_re_base[None, :] * dtheta_ds |
| J = (dH_re_ds @ W_re + dH_im_ds @ W_im) / self.sqrtD |
| g = float(np.linalg.norm(J)) |
| log_g = np.log1p(g) |
| centered = log_g - self._log_g_ema |
| self._log_g_ema = G_EMA_DECAY * self._log_g_ema + (1.0 - G_EMA_DECAY) * log_g |
| beta_eff = self.beta_base * (1.0 + centered) |
| beta_eff = float(np.clip(beta_eff, self.beta_base * BETA_EFF_MIN_MULT, self.beta_base * BETA_EFF_MAX_MULT)) |
| self.beta_eff_history.append(beta_eff) |
| theta = beta_eff * proj |
| return (np.cos(theta).astype(np.float32), np.sin(theta).astype(np.float32)) |
|
|
| def make_worker_encoder(cfg, master_seed): |
| if cfg.get('adaptive_beta', False): |
| return HDEncoderGradAdaptiveHybrid(cfg['feat_lo'], cfg['feat_hi'], cfg['D'], master_seed, cfg['feature_fn'], cfg['beta'], phi_init=cfg.get('fpe_phi_init')) |
| else: |
| return a.HDEncoderFPE(cfg['feat_lo'], cfg['feat_hi'], cfg['D'], cfg['beta'], seed=master_seed, feature_fn=cfg['feature_fn'], phi_init=cfg.get('fpe_phi_init'), beta_vec_init=cfg.get('beta_vec_init')) |
|
|
| def rollout_worker_adaptive(worker_id, cfg, conn, master_seed): |
| encoder = make_worker_encoder(cfg, master_seed) |
| adaptive = cfg.get('adaptive_beta', False) |
| beta_log_writer = None |
| if adaptive and cfg.get('adaptive_beta_log_dir'): |
| log_dir = cfg['adaptive_beta_log_dir'] |
| os.makedirs(log_dir, exist_ok=True) |
| beta_log_file = open(os.path.join(log_dir, f'worker_{worker_id}_beta_log.csv'), 'w', newline='') |
| beta_log_writer = csv.writer(beta_log_file) |
| beta_log_writer.writerow(['rollout_idx', 'n', 'mean', 'std', 'min', 'max']) |
| rollout_idx = 0 |
| sqrtD = float(np.sqrt(cfg['D'])) |
| D = cfg['D'] |
| n_obs = cfg['n_feat'] |
| action_dim = cfg['action_dim'] |
| a_lo = float(cfg['action_low']) |
| a_hi = float(cfg['action_high']) |
| rng = np.random.default_rng(master_seed + 1 + worker_id * 1000) |
| env = __import__('gymnasium').make(cfg['env_id'], **cfg.get('env_kwargs', {})) |
| state, _ = env.reset(seed=master_seed + 1 + worker_id * 1000) |
| while True: |
| cmd = conn.recv() |
| if cmd[0] == 'exit': |
| if beta_log_writer is not None: |
| beta_log_file.close() |
| env.close() |
| return |
| _, W_re, W_im, log_std, n_steps = cmd |
| if adaptive: |
| encoder.set_weights(W_re, W_im) |
| std = np.exp(log_std).astype(np.float32) |
| log_std_sum = float(log_std.sum()) |
| H_res = np.empty((n_steps, D), np.float32) |
| H_ims = np.empty((n_steps, D), np.float32) |
| obs_arr = np.empty((n_steps, n_obs), np.float32) |
| nobs_arr = np.empty((n_steps, n_obs), np.float32) |
| a_arr = np.empty((n_steps, action_dim), np.float32) |
| r_arr = np.empty(n_steps, np.float64) |
| lp_arr = np.empty(n_steps, np.float32) |
| term_arr = np.empty(n_steps, np.float64) |
| trunc_arr = np.empty(n_steps, np.float64) |
| ep_rewards = [] |
| ep_lengths = [] |
| ep_r = 0.0 |
| ep_len = 0 |
| for t in range(n_steps): |
| H_re, H_im = encoder.encode(state) |
| mu = (H_re @ W_re + H_im @ W_im) / sqrtD |
| act = mu + std * rng.standard_normal(action_dim).astype(np.float32) |
| act = np.clip(act, a_lo, a_hi) |
| diff = act - mu |
| lp = float(-0.5 * np.sum((diff / std) ** 2) - log_std_sum - action_dim * a.LOG_2PI_HALF) |
| next_s, reward, term, trunc, _ = env.step(act.astype(np.float32)) |
| H_res[t] = H_re |
| H_ims[t] = H_im |
| obs_arr[t] = np.asarray(state, dtype=np.float32) |
| nobs_arr[t] = np.asarray(next_s, dtype=np.float32) |
| a_arr[t] = act |
| r_arr[t] = float(reward) |
| lp_arr[t] = np.float32(lp) |
| term_arr[t] = float(term) |
| trunc_arr[t] = float(trunc) |
| ep_r += float(reward) |
| ep_len += 1 |
| if term or trunc: |
| ep_rewards.append(ep_r) |
| ep_lengths.append(ep_len) |
| ep_r = 0.0 |
| ep_len = 0 |
| state, _ = env.reset() |
| else: |
| state = next_s |
| if beta_log_writer is not None: |
| beta_hist = np.asarray(encoder.beta_eff_history[-n_steps:], dtype=np.float64) |
| beta_log_writer.writerow([rollout_idx, len(beta_hist), float(beta_hist.mean()), float(beta_hist.std()), float(beta_hist.min()), float(beta_hist.max())]) |
| beta_log_file.flush() |
| rollout_idx += 1 |
| conn.send((H_res, H_ims, obs_arr, nobs_arr, a_arr, r_arr, lp_arr, term_arr, trunc_arr, ep_rewards, ep_lengths, state.copy(), bool(term or trunc))) |
| a.rollout_worker = rollout_worker_adaptive |
|
|
| class HDPPOHybridAgentAdaptive(a.HDPPOHybridAgent): |
|
|
| def __init__(self, cfg, seed=a.DEFAULT_SEED): |
| super().__init__(cfg, seed=seed) |
| if cfg.get('adaptive_beta', False): |
| self.encoder = HDEncoderGradAdaptiveHybrid(cfg['feat_lo'], cfg['feat_hi'], cfg['D'], seed, cfg['feature_fn'], cfg['beta'], phi_init=cfg.get('fpe_phi_init')) |
| self.encoder.link_torch_actor(self.actor) |
| a.HDPPOHybridAgent = HDPPOHybridAgentAdaptive |
|
|
| def prune_actor_global(checkpoint, D_prime): |
| W_re, W_im = (np.asarray(checkpoint['W_actor_re']), np.asarray(checkpoint['W_actor_im'])) |
| Phi = np.asarray(checkpoint['fpe_phi']) |
| beta_base = float(checkpoint['beta_base']) if 'beta_base' in checkpoint else float(np.atleast_1d(checkpoint['beta_bands'])[0]) |
| importance = np.sqrt((W_re ** 2).sum(axis=1) + (W_im ** 2).sum(axis=1)) |
| keep_idx = np.sort(np.argsort(-importance)[:D_prime]) |
| return dict(D=D_prime, beta_base=beta_base, beta=beta_base, beta_vec=np.full(D_prime, beta_base, dtype=np.float32), fpe_phi=Phi[:, keep_idx], W_actor_re=W_re[keep_idx], W_actor_im=W_im[keep_idx]) |
|
|
| def merge_beta_logs(log_dir, rollout_steps): |
| import glob |
| worker_files = sorted(glob.glob(os.path.join(log_dir, 'worker_*_beta_log.csv'))) |
| if not worker_files: |
| return [] |
| per_worker = [] |
| for fp in worker_files: |
| with open(fp) as f: |
| rows = list(csv.DictReader(f)) |
| per_worker.append(rows) |
| n_rollouts = min((len(rows) for rows in per_worker)) |
| merged = [] |
| for i in range(n_rollouts): |
| ns = np.array([float(rows[i]['n']) for rows in per_worker]) |
| means = np.array([float(rows[i]['mean']) for rows in per_worker]) |
| stds = np.array([float(rows[i]['std']) for rows in per_worker]) |
| mins = np.array([float(rows[i]['min']) for rows in per_worker]) |
| maxs = np.array([float(rows[i]['max']) for rows in per_worker]) |
| total_n = ns.sum() |
| pooled_mean = float((ns * means).sum() / total_n) |
| pooled_var = float((ns * (stds ** 2 + means ** 2)).sum() / total_n - pooled_mean ** 2) |
| merged.append(dict(rollout_idx=i, global_step=(i + 1) * rollout_steps, n=int(total_n), mean=pooled_mean, std=float(np.sqrt(max(pooled_var, 0.0))), min=float(mins.min()), max=float(maxs.max()))) |
| return merged |
|
|
| def run_condition(seed, adaptive, total_timesteps, verbose=False, beta_log_dir=None, warm_start=None, save_weights_path=None, D=None, log_csv_path=None, eval_csv_path=None, eval_every_n_steps=None): |
| target_D = int(warm_start['D']) if warm_start is not None else int(D) if D is not None else DEFAULT_D |
| a.ANT_CONFIG['D'] = target_D |
| a.ANT_CONFIG['beta'] = BETA_BASE |
| a.ANT_CONFIG['adaptive_beta'] = adaptive |
| a.ANT_CONFIG.pop('fpe_phi_init', None) |
| a.ANT_CONFIG.pop('beta_vec_init', None) |
| if adaptive and beta_log_dir: |
| os.makedirs(beta_log_dir, exist_ok=True) |
| a.ANT_CONFIG['adaptive_beta_log_dir'] = beta_log_dir |
| else: |
| a.ANT_CONFIG.pop('adaptive_beta_log_dir', None) |
| t0 = time.perf_counter() |
| summary = a.train(total_timesteps=total_timesteps, verbose=verbose, seed=seed, wandb_mode='disabled', save_weights_path=save_weights_path, warm_start=warm_start, log_csv_path=log_csv_path, eval_csv_path=eval_csv_path, eval_every_n_steps=eval_every_n_steps) |
| wall = time.perf_counter() - t0 |
| beta_log = None |
| if adaptive and beta_log_dir: |
| beta_log = merge_beta_logs(beta_log_dir, a.ANT_CONFIG['rollout_steps']) |
| return dict(curve_steps=summary.get('curve_steps'), curve_avg100=summary.get('curve_avg100'), beta_log=beta_log, seed=seed, adaptive=adaptive, total_time=wall, final_avg100=summary.get('final_avg100'), best_avg100=summary.get('best_avg100'), eval_mean_final=summary.get('eval_mean_final'), eval_mean_best=summary.get('eval_mean_best'), critic_state_dict=summary.get('critic_state_dict')) |
|
|