import csv import json import os import sys if sys.platform.startswith('linux'): os.environ.setdefault('MUJOCO_GL', 'egl') os.environ.setdefault('PYOPENGL_PLATFORM', 'egl') import time import multiprocessing as mp from contextlib import contextmanager from collections import deque import numpy as np import gymnasium as gym import psutil import torch import torch.nn as nn import torch.nn.functional as F ENV_ID = 'InvertedPendulum-v5' N_EVAL_EPISODES = 20 EVAL_SEED_BASE = 10000 LOG_2PI_HALF = 0.5 * np.log(2.0 * np.pi) DEVICE = torch.device('cpu') torch.set_num_threads(max(1, psutil.cpu_count(logical=False) or 2)) def state_features(obs): return np.asarray(obs, dtype=np.float32) N_WORKERS = max(1, min(4, os.cpu_count() or 2)) ROLLOUT_STEPS = 2048 OBS_DIM = 4 ACTION_DIM = 1 ENV_KWARGS = dict() CONFIG = dict(env_id=ENV_ID, env_kwargs=ENV_KWARGS, feat_lo=[-10.0] * OBS_DIM, feat_hi=[10.0] * OBS_DIM, feature_fn=state_features, n_feat=OBS_DIM, action_dim=ACTION_DIM, action_low=-3.0, action_high=3.0, D=512, beta=2.0 ** (-0.5), adaptive_beta=True, rollout_steps=ROLLOUT_STEPS, n_workers=N_WORKERS, clip_eps=0.2, vf_coef=0.5, entropy_coef=0.001, entropy_decay=1.0, entropy_min=0.0, gamma=0.99, lam=0.95, adv_clip=5.0, actor_lr=0.0003, critic_lr=0.001, log_std_lr=0.0003, n_epochs=10, minibatch_size=256, max_grad_norm=0.5, target_kl=None, log_std_init=-0.5, log_std_min=-2.0, log_std_max=2.0, critic_hidden=(64, 64), ema_interval=200, ema_alpha=0.15) DEFAULT_SEED = 42 BETA_EFF_MIN_MULT = 0.2 BETA_EFF_MAX_MULT = 4.0 G_EMA_DECAY = 0.99 class HDEncoderGradAdaptive: 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)) class HDEncoderFixed: def __init__(self, feat_lo, feat_hi, D, beta, seed, feature_fn, 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.D = D self.beta_vec = np.full(D, float(beta), dtype=np.float32) self.n_feat = len(feat_lo) 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) def encode(self, state): 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) theta = s_norm @ self.Phi * self.beta_vec return (np.cos(theta), np.sin(theta)) def make_encoder(cfg, seed): if cfg.get('adaptive_beta', True): return HDEncoderGradAdaptive(cfg['feat_lo'], cfg['feat_hi'], cfg['D'], seed, cfg['feature_fn'], cfg['beta'], phi_init=cfg.get('fpe_phi_init')) return HDEncoderFixed(cfg['feat_lo'], cfg['feat_hi'], cfg['D'], cfg['beta'], seed, cfg['feature_fn'], phi_init=cfg.get('fpe_phi_init')) class HDLinearActor(nn.Module): def __init__(self, D, action_dim, log_std_init, action_low, action_high): super().__init__() self.D = D self.action_dim = action_dim self.sqrtD = float(np.sqrt(D)) self.W_re = nn.Parameter(torch.zeros(D, action_dim, dtype=torch.float32)) self.W_im = nn.Parameter(torch.zeros(D, action_dim, dtype=torch.float32)) self.log_std = nn.Parameter(torch.full((action_dim,), float(log_std_init), dtype=torch.float32)) self.a_lo = float(action_low) self.a_hi = float(action_high) def mean_from_hd(self, H_re, H_im): return (H_re @ self.W_re + H_im @ self.W_im) / self.sqrtD def dist(self, H_re, H_im): mu = self.mean_from_hd(H_re, H_im) std = self.log_std.exp().unsqueeze(0).expand_as(mu) return (mu, std) @torch.no_grad() def greedy_action_np(self, H_re_np, H_im_np): Hr = torch.from_numpy(H_re_np).unsqueeze(0) Hi = torch.from_numpy(H_im_np).unsqueeze(0) mu = self.mean_from_hd(Hr, Hi).squeeze(0).cpu().numpy() return np.clip(mu, self.a_lo, self.a_hi).astype(np.float32) class MLPCritic(nn.Module): def __init__(self, n_obs, hidden=(64, 64)): super().__init__() layers = [] last = n_obs for h in hidden: lin = nn.Linear(last, h) nn.init.orthogonal_(lin.weight, gain=np.sqrt(2)) nn.init.zeros_(lin.bias) layers += [lin, nn.Tanh()] last = h head = nn.Linear(last, 1) nn.init.orthogonal_(head.weight, gain=1.0) nn.init.zeros_(head.bias) layers.append(head) self.net = nn.Sequential(*layers) def forward(self, obs): return self.net(obs).squeeze(-1) def compute_gae_with_trunc(rewards, values, next_values, terminated, truncated, last_value, gamma, lam): n = rewards.shape[0] adv = np.zeros(n, dtype=np.float64) gae = 0.0 for t in range(n - 1, -1, -1): if terminated[t]: next_v = 0.0 elif truncated[t]: next_v = next_values[t] elif t == n - 1: next_v = last_value else: next_v = values[t + 1] delta = rewards[t] + gamma * next_v - values[t] gae = delta if terminated[t] or truncated[t] else delta + gamma * lam * gae adv[t] = gae return (adv, adv + values) class SystemMonitor: def __init__(self): self.proc = psutil.Process(os.getpid()) def _tree_memory_mb(self): total = self.proc.memory_info().rss for child in self.proc.children(recursive=True): try: total += child.memory_info().rss except (psutil.NoSuchProcess, psutil.AccessDenied): pass return total / 1024 ** 2 def snapshot(self): return {'ram_process_tree_mb': self._tree_memory_mb()} def rollout_worker(worker_id, cfg, conn, master_seed): encoder = make_encoder(cfg, master_seed) adaptive = cfg.get('adaptive_beta', True) 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 = gym.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 * 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))) class WorkerPool: def __init__(self, n_workers, cfg, master_seed): self.n_workers = n_workers self.conns = [] self.procs = [] ctx = mp.get_context('fork') for i in range(n_workers): parent_conn, child_conn = ctx.Pipe() p = ctx.Process(target=rollout_worker, args=(i, cfg, child_conn, master_seed), daemon=True) p.start() child_conn.close() self.conns.append(parent_conn) self.procs.append(p) def collect(self, W_re, W_im, log_std, total_steps): steps_each = total_steps // self.n_workers for conn in self.conns: conn.send(('rollout', W_re.copy(), W_im.copy(), log_std.copy(), steps_each)) return [conn.recv() for conn in self.conns] def close(self): for conn in self.conns: try: conn.send(('exit',)) except Exception: pass for p in self.procs: p.join(timeout=3) if p.is_alive(): p.terminate() def actor_snapshot(actor): return {k: v.detach().clone() for k, v in actor.state_dict().items()} def actor_restore(actor, snap, alpha): sd = actor.state_dict() with torch.no_grad(): for k in sd: sd[k].mul_(1.0 - alpha).add_(snap[k], alpha=alpha) @contextmanager def actor_weights_swapped(actor, snap): saved = actor_snapshot(actor) actor_restore(actor, snap, alpha=1.0) try: yield finally: actor_restore(actor, saved, alpha=1.0) def evaluate_agent(encoder, actor, env_id, env_kwargs, n_episodes=N_EVAL_EPISODES, seed_base=EVAL_SEED_BASE): env = gym.make(env_id, **env_kwargs) rewards = np.empty(n_episodes, dtype=np.float64) actor.eval() for i in range(n_episodes): state, _ = env.reset(seed=seed_base + i) ep_r = 0.0 done = False while not done: H_re, H_im = encoder.encode(state) a = actor.greedy_action_np(H_re, H_im) state, reward, term, trunc, _ = env.step(a) ep_r += float(reward) done = term or trunc rewards[i] = ep_r actor.train() env.close() n = len(rewards) sem = float(rewards.std(ddof=1) / np.sqrt(n)) if n > 1 else 0.0 return dict(mean_reward=float(rewards.mean()), ci95_reward=1.96 * sem, rewards=rewards.tolist()) class HDPPOHybridAgent: def __init__(self, cfg, seed=DEFAULT_SEED): D = cfg['D'] torch.manual_seed(seed) np.random.seed(seed) self.actor = HDLinearActor(D, cfg['action_dim'], cfg['log_std_init'], cfg['action_low'], cfg['action_high']).to(DEVICE) self.encoder = make_encoder(cfg, seed) if cfg.get('adaptive_beta', True): self.encoder.link_torch_actor(self.actor) self.critic = MLPCritic(cfg['n_feat'], cfg['critic_hidden']).to(DEVICE) self.opt_actor = torch.optim.Adam([{'params': [self.actor.W_re, self.actor.W_im], 'lr': cfg['actor_lr']}, {'params': [self.actor.log_std], 'lr': cfg['log_std_lr']}]) self.opt_critic = torch.optim.Adam(self.critic.parameters(), lr=cfg['critic_lr']) self.cfg = cfg self.entropy_coef = cfg['entropy_coef'] self.best_avg100 = -np.inf self.best_snap = None self.pool = WorkerPool(cfg['n_workers'], cfg, seed) def collect_and_update(self, total_steps): cfg = self.cfg t_rollout_start = time.time() with torch.no_grad(): W_re_np = self.actor.W_re.detach().cpu().numpy().astype(np.float32) W_im_np = self.actor.W_im.detach().cpu().numpy().astype(np.float32) log_std_np = self.actor.log_std.detach().cpu().numpy().astype(np.float32) results = self.pool.collect(W_re_np, W_im_np, log_std_np, total_steps) t_rollout = time.time() - t_rollout_start t_update_start = time.time() all_H_re, all_H_im, all_obs, all_nobs = ([], [], [], []) all_a, all_lp, all_r, all_term, all_trunc = ([], [], [], [], []) ep_rews, ep_lens = ([], []) for H_re, H_im, obs, nobs, a, r, lp, term, trunc, ep_r, ep_l, ls, ld in results: all_H_re.append(H_re) all_H_im.append(H_im) all_obs.append(obs) all_nobs.append(nobs) all_a.append(a) all_lp.append(lp) all_r.append(r) all_term.append(term) all_trunc.append(trunc) ep_rews.extend(ep_r) ep_lens.extend(ep_l) H_re = np.concatenate(all_H_re) H_im = np.concatenate(all_H_im) obs_np = np.concatenate(all_obs) nobs_np = np.concatenate(all_nobs) a_np = np.concatenate(all_a).astype(np.float32) old_lps = np.concatenate(all_lp).astype(np.float32) r_np = np.concatenate(all_r) term_np = np.concatenate(all_term) trunc_np = np.concatenate(all_trunc) self.critic.eval() with torch.no_grad(): obs_t = torch.from_numpy(obs_np).to(DEVICE) nobs_t = torch.from_numpy(nobs_np).to(DEVICE) v_pred = self.critic(obs_t).cpu().numpy().astype(np.float64) v_next = self.critic(nobs_t).cpu().numpy().astype(np.float64) self.critic.train() adv_list, ret_list = ([], []) cursor = 0 for (_, _, _, _, _, _, _, _, _, _, _, ls, ld), worker_r in zip(results, all_r): n = worker_r.shape[0] r_seg = r_np[cursor:cursor + n] v_seg = v_pred[cursor:cursor + n] nv_seg = v_next[cursor:cursor + n] term_seg = term_np[cursor:cursor + n] trunc_seg = trunc_np[cursor:cursor + n] if ld: last_val = 0.0 else: with torch.no_grad(): ls_t = torch.from_numpy(np.asarray(ls, dtype=np.float32)).unsqueeze(0).to(DEVICE) last_val = float(self.critic(ls_t).item()) adv, ret = compute_gae_with_trunc(r_seg.astype(np.float64), v_seg.astype(np.float64), nv_seg.astype(np.float64), term_seg.astype(np.float64), trunc_seg.astype(np.float64), last_val, cfg['gamma'], cfg['lam']) adv_list.append(adv) ret_list.append(ret) cursor += n adv_np = np.concatenate(adv_list) ret_np = np.concatenate(ret_list) T = len(adv_np) adv_std = float(adv_np.std()) adv_clip = float(cfg['adv_clip']) adv_n_np = np.clip((adv_np - adv_np.mean()) / (adv_std + 1e-08), -adv_clip, adv_clip) if adv_std > 0.0001 else np.zeros(T, np.float64) var_y = float(np.var(ret_np)) explained_var = float(1 - np.var(ret_np - v_pred) / var_y) if var_y > 1e-08 else 0.0 H_re_t = torch.from_numpy(H_re).to(DEVICE) H_im_t = torch.from_numpy(H_im).to(DEVICE) obs_t = torch.from_numpy(obs_np).to(DEVICE) a_t = torch.from_numpy(a_np).to(DEVICE) old_lp_t = torch.from_numpy(old_lps).to(DEVICE) adv_t = torch.from_numpy(adv_n_np.astype(np.float32)).to(DEVICE) ret_t = torch.from_numpy(ret_np.astype(np.float32)).to(DEVICE) clip_eps = float(cfg['clip_eps']) vf_coef = float(cfg['vf_coef']) ent_coef = float(self.entropy_coef) mb_size = int(cfg['minibatch_size']) max_gn = float(cfg['max_grad_norm']) log_std_min = float(cfg['log_std_min']) log_std_max = float(cfg['log_std_max']) action_dim = cfg['action_dim'] rng = np.random.default_rng() kls, clip_fracs, actor_losses, critic_losses = ([], [], [], []) for epoch in range(cfg['n_epochs']): perm = rng.permutation(T) for start in range(0, T, mb_size): idx = perm[start:start + mb_size] idx_t = torch.from_numpy(idx).to(DEVICE) Hr_b = H_re_t.index_select(0, idx_t) Hi_b = H_im_t.index_select(0, idx_t) obs_b = obs_t.index_select(0, idx_t) a_b = a_t.index_select(0, idx_t) oldlp_b = old_lp_t.index_select(0, idx_t) adv_b = adv_t.index_select(0, idx_t) ret_b = ret_t.index_select(0, idx_t) mu, std = self.actor.dist(Hr_b, Hi_b) var = std * std new_lp = (-0.5 * (a_b - mu).pow(2) / var - self.actor.log_std.unsqueeze(0) - LOG_2PI_HALF).sum(dim=-1) ratio = (new_lp - oldlp_b).exp() surr1 = ratio * adv_b surr2 = torch.clamp(ratio, 1.0 - clip_eps, 1.0 + clip_eps) * adv_b policy_loss = -torch.min(surr1, surr2).mean() entropy = 0.5 * np.log(2 * np.pi * np.e) * action_dim + self.actor.log_std.sum() actor_loss = policy_loss - ent_coef * entropy self.opt_actor.zero_grad(set_to_none=True) actor_loss.backward() nn.utils.clip_grad_norm_([self.actor.W_re, self.actor.W_im, self.actor.log_std], max_gn) self.opt_actor.step() with torch.no_grad(): self.actor.log_std.clamp_(log_std_min, log_std_max) v_b = self.critic(obs_b) value_loss = F.mse_loss(v_b, ret_b) self.opt_critic.zero_grad(set_to_none=True) (vf_coef * value_loss).backward() nn.utils.clip_grad_norm_(self.critic.parameters(), max_gn) self.opt_critic.step() with torch.no_grad(): log_ratio = new_lp - oldlp_b approx_kl = (log_ratio.exp() - 1 - log_ratio).mean().item() clip_frac = ((ratio > 1 + clip_eps) | (ratio < 1 - clip_eps)).float().mean().item() kls.append(approx_kl) clip_fracs.append(clip_frac) actor_losses.append(policy_loss.item()) critic_losses.append(value_loss.item()) self.entropy_coef = max(self.entropy_coef * cfg['entropy_decay'], cfg['entropy_min']) t_update = time.time() - t_update_start with torch.no_grad(): log_std_now = self.actor.log_std.detach().cpu().numpy() diag = dict(rollout_s=t_rollout, update_s=t_update, fps=total_steps / (t_rollout + t_update + 1e-08), explained_var=explained_var, policy_loss=float(np.mean(actor_losses)), value_loss=float(np.mean(critic_losses)), approx_kl=float(np.mean(kls)), clip_fraction=float(np.mean(clip_fracs)), entropy_coef=float(self.entropy_coef), log_std_mean=float(log_std_now.mean()), std_mean=float(np.exp(log_std_now).mean()), mean_value=float(v_pred.mean())) return (ep_rews, ep_lens, diag) def maybe_snapshot(self, avg100): if avg100 > self.best_avg100: self.best_avg100 = avg100 self.best_snap = actor_snapshot(self.actor) def ema_restore(self): if self.best_snap is not None: actor_restore(self.actor, self.best_snap, self.cfg['ema_alpha']) def evaluate_final(self, n_episodes=N_EVAL_EPISODES, seed_base=EVAL_SEED_BASE): return evaluate_agent(self.encoder, self.actor, self.cfg['env_id'], self.cfg.get('env_kwargs', {}), n_episodes, seed_base) def evaluate_best(self, n_episodes=N_EVAL_EPISODES, seed_base=EVAL_SEED_BASE): if self.best_snap is None: return None with actor_weights_swapped(self.actor, self.best_snap): return evaluate_agent(self.encoder, self.actor, self.cfg['env_id'], self.cfg.get('env_kwargs', {}), n_episodes, seed_base) def close(self): self.pool.close() def train(total_timesteps, seed=DEFAULT_SEED, verbose=True, warm_start=None, save_weights_path=None, log_csv_path=None, eval_csv_path=None, eval_every_n_steps=None): csv_file = csv_writer = None if log_csv_path is not None: csv_file = open(log_csv_path, 'w', newline='') csv_writer = csv.writer(csv_file) csv_writer.writerow(['global_step', 'wall_time_sec', 'episode', 'episodes_this_update', 'ep_rew_mean', 'ep_rew_max', 'ep_rew_min', 'ep_len_mean', 'best_avg100', 'policy_loss', 'value_loss', 'explained_var', 'approx_kl', 'clip_fraction', 'entropy_coef', 'log_std_mean', 'std_mean', 'mean_value', 'rollout_s', 'update_s', 'fps', 'ram_process_tree_mb']) eval_csv_file = eval_csv_writer = None if eval_csv_path is not None: eval_csv_file = open(eval_csv_path, 'w', newline='') eval_csv_writer = csv.writer(eval_csv_file) eval_csv_writer.writerow(['global_step', 'eval_mean', 'eval_ci95', 'tag']) cfg = dict(CONFIG) if warm_start is not None else CONFIG if warm_start is not None: cfg['D'] = int(warm_start['D']) cfg['beta'] = warm_start.get('beta', cfg['beta']) cfg['fpe_phi_init'] = np.asarray(warm_start['fpe_phi'], dtype=np.float32) agent = HDPPOHybridAgent(cfg, seed=seed) if warm_start is not None: with torch.no_grad(): agent.actor.W_re.copy_(torch.from_numpy(np.asarray(warm_start['W_actor_re'], dtype=np.float32))) agent.actor.W_im.copy_(torch.from_numpy(np.asarray(warm_start['W_actor_im'], dtype=np.float32))) if 'log_std' in warm_start: agent.actor.log_std.copy_(torch.from_numpy(np.asarray(warm_start['log_std'], dtype=np.float32))) if verbose: print(f' Warm-started from provided checkpoint: D={cfg['D']}') if warm_start.get('critic_state_dict') is not None: agent.critic.load_state_dict(warm_start['critic_state_dict']) if verbose: print(' Warm-started critic from provided checkpoint (warm-start, not reset)') if eval_csv_writer is not None: post_prune_eval = agent.evaluate_final() eval_csv_writer.writerow([0, post_prune_eval['mean_reward'], post_prune_eval['ci95_reward'], 'post_prune']) eval_csv_file.flush() if verbose: print(f' Post-prune eval (before fine-tuning): {post_prune_eval['mean_reward']:+.1f} +/- {post_prune_eval['ci95_reward']:.1f}') next_eval_at = eval_every_n_steps rollout_steps = cfg['rollout_steps'] ema_interval = cfg['ema_interval'] sysmon = SystemMonitor() recent = deque(maxlen=100) len_recent = deque(maxlen=100) ep = 0 total_steps = 0 update_count = 0 if verbose: print('=' * 80) print(f'HD-PPO -> {cfg['env_id']}') print(f' D={cfg['D']} beta={cfg['beta']} n_workers={cfg['n_workers']} total_timesteps={total_timesteps:,}') print('=' * 80) t0 = time.perf_counter() while total_steps < total_timesteps: ep_batch, ep_len_batch, diag = agent.collect_and_update(rollout_steps) total_steps += rollout_steps update_count += 1 wall_sec = int(time.perf_counter() - t0) if eval_csv_writer is not None and next_eval_at is not None: while total_steps >= next_eval_at: periodic_eval = agent.evaluate_final() eval_csv_writer.writerow([next_eval_at, periodic_eval['mean_reward'], periodic_eval['ci95_reward'], 'periodic']) eval_csv_file.flush() if verbose: print(f' [eval @ step {next_eval_at:>9,}] {periodic_eval['mean_reward']:+.1f} +/- {periodic_eval['ci95_reward']:.1f}') next_eval_at += eval_every_n_steps for r, L in zip(ep_batch, ep_len_batch): ep += 1 recent.append(r) len_recent.append(L) if len(recent) > 0: avg100 = float(np.mean(recent)) agent.maybe_snapshot(avg100) if ep % ema_interval == 0 and ep > 0: agent.ema_restore() if csv_writer is not None: sysm = sysmon.snapshot() csv_writer.writerow([total_steps, wall_sec, ep, len(ep_batch), avg100, float(np.max(recent)), float(np.min(recent)), float(np.mean(len_recent)) if len_recent else 0.0, agent.best_avg100, diag['policy_loss'], diag['value_loss'], diag['explained_var'], diag['approx_kl'], diag['clip_fraction'], diag['entropy_coef'], diag['log_std_mean'], diag['std_mean'], diag['mean_value'], diag['rollout_s'], diag['update_s'], diag['fps'], sysm['ram_process_tree_mb']]) csv_file.flush() if verbose and update_count % 5 == 0: sysm = sysmon.snapshot() print(f' [step {total_steps:>9,}] ep {ep:>5} avg100={avg100:>+8.1f} best={agent.best_avg100:>+8.1f} pol_loss={diag['policy_loss']:>+.3f} v_loss={diag['value_loss']:>.2f} fps={diag['fps']:>6.0f} ram={sysm['ram_process_tree_mb']:>5.0f}MB') if csv_file is not None: csv_file.close() total_time = time.perf_counter() - t0 final_train_avg = float(np.mean(list(recent))) if recent else 0.0 if verbose: print(f'\n Running held-out eval ({N_EVAL_EPISODES} eps)...', flush=True) eval_final = agent.evaluate_final() eval_best = agent.evaluate_best() if eval_csv_file is not None: eval_csv_file.close() critic_state_dict = {k: v.detach().clone() for k, v in agent.critic.state_dict().items()} if save_weights_path is not None: if agent.best_snap is not None: best_sd = agent.best_snap W_re_save = best_sd['W_re'].cpu().numpy() W_im_save = best_sd['W_im'].cpu().numpy() log_std_save = best_sd['log_std'].cpu().numpy() else: W_re_save = agent.actor.W_re.detach().cpu().numpy() W_im_save = agent.actor.W_im.detach().cpu().numpy() log_std_save = agent.actor.log_std.detach().cpu().numpy() np.savez(save_weights_path, W_actor_re=W_re_save, W_actor_im=W_im_save, log_std=log_std_save, fpe_phi=agent.encoder.Phi, beta_base=np.float32(cfg['beta']), feat_lo=np.array(cfg['feat_lo'], dtype=np.float32), feat_hi=np.array(cfg['feat_hi'], dtype=np.float32), D=np.int32(cfg['D']), n_feat=np.int32(cfg['n_feat']), action_dim=np.int32(cfg['action_dim']), action_low=np.float32(cfg['action_low']), action_high=np.float32(cfg['action_high']), eval_mean_final=np.float32(eval_final['mean_reward']), eval_mean_best=np.float32(eval_best['mean_reward'] if eval_best is not None else np.nan)) if verbose: print(f' Saved actor -> {save_weights_path} ({os.path.getsize(save_weights_path) / 1024:.1f} KB)') if verbose: eb_str = f'{eval_best['mean_reward']:.1f}+/-{eval_best['ci95_reward']:.1f}' if eval_best is not None else '-' print(f'\n Training time: {total_time:.1f}s') print(f' Episodes: {ep:,}') print(f' Final train avg100: {final_train_avg:+.1f}') print(f' Best train avg100: {agent.best_avg100:+.1f}') print(f' Eval (final wts): {eval_final['mean_reward']:+.1f} +/- {eval_final['ci95_reward']:.1f}') print(f' Eval (best wts): {eb_str}') agent.close() return dict(eval_mean_final=eval_final['mean_reward'], eval_ci95_final=eval_final['ci95_reward'], eval_mean_best=eval_best['mean_reward'] if eval_best is not None else None, final_avg100=final_train_avg, best_avg100=float(agent.best_avg100), training_time_sec=total_time, critic_state_dict=critic_state_dict) 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']) 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=beta_base, fpe_phi=Phi[:, keep_idx], W_actor_re=W_re[keep_idx], W_actor_im=W_im[keep_idx]) THIS_DIR = os.path.dirname(os.path.abspath(__file__)) SEED = DEFAULT_SEED STAGES = [(512, 1000000), (128, 1000000), (64, 1000000)] OUT_DIR = os.path.join(THIS_DIR, 'prune_finetune') def weights_path(D, stage_label): return os.path.join(OUT_DIR, f'hdppo_D{D}_{stage_label}.npz') def curve_csv_path(D, stage_label): return os.path.join(OUT_DIR, f'training_curve_D{D}_{stage_label}.csv') def eval_csv_path_for(D, stage_label): return os.path.join(OUT_DIR, f'eval_curve_D{D}_{stage_label}.csv') def main(): os.makedirs(OUT_DIR, exist_ok=True) results_json = os.path.join(OUT_DIR, 'results.json') table_txt = os.path.join(OUT_DIR, 'results_table.txt') stage_records = [] prev_path = None prev_critic_sd = None t_chain0 = time.perf_counter() for i, (D, timesteps) in enumerate(STAGES): if i == 0: stage_label = 'teacher_fresh' warm_start = None print(f'\n{'#' * 90}\nSTAGE {i + 1}/{len(STAGES)}: D={D} FRESH, {timesteps:,} steps\n{'#' * 90}', flush=True) else: stage_label = 'finetuned' prev_D = STAGES[i - 1][0] print(f'\n{'#' * 90}\nSTAGE {i + 1}/{len(STAGES)}: prune D={prev_D} -> D={D}, then fine-tune {timesteps:,} steps (critic warm-started)\n{'#' * 90}', flush=True) prev_ckpt = np.load(prev_path) warm_start = prune_actor_global(prev_ckpt, D_prime=D) warm_start['critic_state_dict'] = prev_critic_sd print(f' Pruned: kept top-{D}/{prev_D} dimensions by weight importance ({prev_D / D:.1f}x cut)') path = weights_path(D, stage_label) curve_csv = curve_csv_path(D, stage_label) eval_csv = eval_csv_path_for(D, stage_label) t0 = time.time() result = train(total_timesteps=timesteps, seed=SEED, verbose=True, warm_start=warm_start, save_weights_path=path, log_csv_path=curve_csv, eval_csv_path=eval_csv, eval_every_n_steps=50000) wall = time.time() - t0 print(f' STAGE {i + 1} done: D={D} eval_final={result['eval_mean_final']:+.1f} eval_best={result['eval_mean_best']:+.1f} wall={wall:.0f}s', flush=True) stage_records.append(dict(stage=stage_label, D=D, configured_timesteps=timesteps, weights_path=path, curve_csv=curve_csv, eval_csv=eval_csv, eval_mean_final=result['eval_mean_final'], eval_mean_best=result['eval_mean_best'], final_avg100=result['final_avg100'], best_avg100=result['best_avg100'], wall_time_sec=wall)) with open(results_json, 'w') as f: json.dump(dict(seed=SEED, stages=stage_records, complete=False), f, indent=2) prev_path = path prev_critic_sd = result['critic_state_dict'] total_time = time.perf_counter() - t_chain0 print('\n' + '=' * 90) header = f'{'stage':<16} {'D':>6} {'steps':>10} {'eval_final':>12} {'eval_best':>12} {'wall (min)':>11}' print(header) lines_txt = [header] for r in stage_records: line = f'{r['stage']:<16} {r['D']:>6} {r['configured_timesteps']:>10,} {r['eval_mean_final']:>12.1f} {r['eval_mean_best']:>12.1f} {r['wall_time_sec'] / 60:>11.1f}' print(line) lines_txt.append(line) print(f'\nTotal chain wall time: {total_time / 60:.1f} min') print('=' * 90) lines_txt.append(f'\nTotal chain wall time: {total_time / 60:.1f} min') with open(table_txt, 'w') as f: f.write('\n'.join(lines_txt) + '\n') print(f'\nWrote {table_txt}') with open(results_json, 'w') as f: json.dump(dict(seed=SEED, stages=stage_records, complete=True, total_chain_time_sec=total_time), f, indent=2) print(f'Wrote {results_json}') if __name__ == '__main__': mp.set_start_method('fork', force=True) main()