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'))