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