| import csv |
| import json |
| import os |
| import time |
| import numpy as np |
| import numba |
| from numba import njit |
| import gymnasium as gym |
| import psutil |
| from collections import deque |
|
|
| @njit |
| def ppo_epoch_jit(H_re, H_im, W_re, W_im, Wc_re, Wc_im, a_vec, adv_n, ret_, old_lps, log_std, log_std_grad_out, bias, perm, scale, sqrtD, actor_lr, critic_lr, clip_eps, bias_lr, ent_coef, action_dim): |
| T = H_re.shape[0] |
| A = action_dim |
| D = H_re.shape[1] |
| LOG_2PI = numba.float32(1.8378770664093453) |
| policy_loss_sum = numba.float64(0.0) |
| value_loss_sum = numba.float64(0.0) |
| for ii in range(T): |
| idx = perm[ii] |
| Hr = H_re[idx] |
| Hi = H_im[idx] |
| mu = np.empty(A, numba.float32) |
| for k in range(A): |
| s = numba.float32(0.0) |
| for d in range(D): |
| s += W_re[k, d] * Hr[d] + W_im[k, d] * Hi[d] |
| mu[k] = s / scale |
| new_lp = numba.float32(0.0) |
| for k in range(A): |
| sigma = np.exp(log_std[k]) |
| z = (a_vec[idx, k] - mu[k]) / sigma |
| new_lp += numba.float32(-0.5) * (z * z + numba.float32(2.0) * log_std[k] + LOG_2PI) |
| ratio = np.exp(new_lp - old_lps[idx]) |
| if ratio > numba.float32(20.0): |
| ratio = numba.float32(20.0) |
| adv = adv_n[idx] |
| lo = numba.float32(1.0) - clip_eps |
| hi2 = numba.float32(1.0) + clip_eps |
| rc = ratio if ratio < hi2 else hi2 |
| rc = rc if rc > lo else lo |
| surr_u = ratio * adv |
| surr_c = rc * adv |
| surr = surr_u if surr_u < surr_c else surr_c |
| policy_loss_sum += numba.float64(-surr) |
| COEF_CLIP = numba.float32(2.0) |
| Z_CLIP = numba.float32(5.0) |
| for k in range(A): |
| sigma = np.exp(log_std[k]) |
| diff = a_vec[idx, k] - mu[k] |
| mean_score = diff / (sigma * sigma) |
| coef = actor_lr * surr * mean_score / scale |
| if coef > COEF_CLIP: |
| coef = COEF_CLIP |
| elif coef < -COEF_CLIP: |
| coef = -COEF_CLIP |
| for d in range(D): |
| W_re[k, d] += coef * Hr[d] |
| W_im[k, d] += coef * Hi[d] |
| z = diff / sigma |
| if z > Z_CLIP: |
| z = Z_CLIP |
| elif z < -Z_CLIP: |
| z = -Z_CLIP |
| log_std_grad_out[k] += surr * (z * z - numba.float32(1.0)) + ent_coef |
| target = ret_[idx] |
| bias = bias + bias_lr * (target - bias) |
| v = numba.float32(0.0) |
| for d in range(D): |
| v += Wc_re[d] * Hr[d] + Wc_im[d] * Hi[d] |
| v /= sqrtD |
| td_error = target - bias - v |
| value_loss_sum += numba.float64(td_error * td_error) |
| res = td_error * (critic_lr / sqrtD) |
| if res > COEF_CLIP: |
| res = COEF_CLIP |
| elif res < -COEF_CLIP: |
| res = -COEF_CLIP |
| for d in range(D): |
| Wc_re[d] += res * Hr[d] |
| Wc_im[d] += res * Hi[d] |
| return (bias, policy_loss_sum / T, value_loss_sum / T) |
|
|
| @njit |
| def compute_gae_jit(rewards, values, dones, last_val, gamma, lam): |
| n = rewards.shape[0] |
| adv = np.zeros(n, numba.float64) |
| gae = 0.0 |
| nv = last_val |
| for t in range(n - 1, -1, -1): |
| m = 1.0 - dones[t] |
| gae = rewards[t] + gamma * nv * m - values[t] + gamma * lam * m * gae |
| adv[t] = gae |
| nv = values[t] |
| return (adv, adv + values) |
|
|
| def warmup_jit(D, T, action_dim): |
| Hr = np.zeros((T, D), np.float32) |
| Hi = np.zeros((T, D), np.float32) |
| Wr = np.zeros((action_dim, D), np.float32) |
| Wi = np.zeros((action_dim, D), np.float32) |
| Cr = np.zeros(D, np.float32) |
| Ci = np.zeros(D, np.float32) |
| av = np.zeros((T, action_dim), np.float32) |
| an = np.zeros(T, np.float32) |
| rt = np.zeros(T, np.float32) |
| ol = np.zeros(T, np.float32) |
| ls = np.zeros(action_dim, np.float32) |
| lsg = np.zeros(action_dim, np.float32) |
| pm = np.arange(T, dtype=np.int32) |
| sc = np.float32(np.sqrt(D)) |
| sq = np.float32(np.sqrt(D)) |
| ppo_epoch_jit(Hr, Hi, Wr, Wi, Cr, Ci, av, an, rt, ol, ls, lsg, np.float32(0.0), pm, sc, sq, np.float32(0.001), np.float32(0.005), np.float32(0.2), np.float32(0.05), np.float32(0.01), action_dim) |
| compute_gae_jit(np.zeros(T, np.float64), np.zeros(T, np.float64), np.zeros(T, np.float64), 0.0, 0.99, 0.95) |
|
|
| def pendulum_features(obs): |
| cos_t, sin_t, td = (float(obs[0]), float(obs[1]), float(obs[2])) |
| return np.array([cos_t, sin_t, td, td / 8.0, td * td, sin_t * td, cos_t * td], dtype=np.float32) |
| OBS_DIM = 7 |
| ACTION_DIM = 1 |
| BASE_BETA = 2.5 |
| PENDULUM_CONFIG = dict(feat_lo=[-1.0, -1.0, -8.0, -1.0, 0.0, -8.0, -8.0], feat_hi=[1.0, 1.0, 8.0, 1.0, 64.0, 8.0, 8.0], feature_fn=pendulum_features, action_dim=ACTION_DIM, action_low=-2.0, action_high=2.0, D=512, beta=BASE_BETA, rollout_steps=1024, actor_lr=0.003, critic_lr=0.005, n_epochs=8, log_std_init=-0.5, log_std_min=-2.0, log_std_max=0.7, log_std_lr=0.01, entropy_coef=0.1, entropy_decay=1.0, entropy_min=0.1, clip_eps=0.2, gamma=0.95, lam=0.95, solve_thresh=-200.0, ema_interval=100, ema_alpha=0.15) |
| REWARD_THRESHOLDS = [-1200, -800, -500, -300, -200, -150, -120, -100] |
| DEFAULT_SEED = 123 |
| BETA_EFF_MIN_MULT = 0.2 |
| BETA_EFF_MAX_MULT = 4.0 |
| G_EMA_DECAY = 0.99 |
|
|
| class HDEncoderGradAdaptiveContinuous: |
|
|
| 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_norm = 2.0 / (self.hi - self.lo + 1e-08) |
| self.dtheta_ds_unit = self.Phi * scale_norm[:, None] |
| self._actor = None |
| self._log_g_ema = 0.0 |
| self.beta_eff_history = [] |
|
|
| def link_actor(self, actor): |
| self._actor = actor |
|
|
| @property |
| def beta_vec(self): |
| return np.full(self.D, self.beta_base, dtype=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) |
| 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)) |
| W_re, W_im = (self._actor.W_re, self._actor.W_im) |
| 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 = (W_re @ dH_re_ds.T + W_im @ dH_im_ds.T) / 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 HDActorContinuous: |
|
|
| def __init__(self, D, action_dim, log_std_init, action_low, action_high): |
| self.sqrtD = float(np.sqrt(D)) |
| self.action_dim = action_dim |
| self.W_re = np.zeros((action_dim, D), dtype=np.float32) |
| self.W_im = np.zeros((action_dim, D), dtype=np.float32) |
| self.log_std = np.full(action_dim, log_std_init, dtype=np.float32) |
| self.a_lo = float(action_low) |
| self.a_hi = float(action_high) |
|
|
| def mean(self, H_re, H_im): |
| return (self.W_re @ H_re + self.W_im @ H_im) / self.sqrtD |
|
|
| def sample(self, H_re, H_im): |
| mu = self.mean(H_re, H_im) |
| sigma = np.exp(self.log_std) |
| a_raw = mu + sigma * np.random.standard_normal(self.action_dim).astype(np.float32) |
| a_env = np.clip(a_raw, self.a_lo, self.a_hi).astype(np.float32) |
| z = (a_raw - mu) / sigma |
| lp = float(np.sum(-0.5 * (z * z + 2.0 * self.log_std + np.log(2.0 * np.pi)))) |
| return (a_raw, a_env, lp) |
|
|
| def greedy(self, H_re, H_im): |
| mu = self.mean(H_re, H_im) |
| return np.clip(mu, self.a_lo, self.a_hi).astype(np.float32) |
|
|
| def snapshot(self): |
| return (self.W_re.copy(), self.W_im.copy()) |
|
|
| def restore(self, snap, alpha): |
| wr, wi = snap |
| self.W_re = (1 - alpha) * self.W_re + alpha * wr |
| self.W_im = (1 - alpha) * self.W_im + alpha * wi |
|
|
| class HDCritic: |
|
|
| def __init__(self, D, bias_lr=0.05): |
| self.sqrtD = float(np.sqrt(D)) |
| self.W_re = np.zeros(D, dtype=np.float32) |
| self.W_im = np.zeros(D, dtype=np.float32) |
| self.bias = 0.0 |
| self.bias_lr = bias_lr |
|
|
| def value(self, H_re, H_im): |
| return float(np.dot(self.W_re, H_re) + np.dot(self.W_im, H_im)) / self.sqrtD + self.bias |
|
|
| class TrajectoryBuffer: |
|
|
| def __init__(self): |
| self.reset() |
|
|
| def reset(self): |
| self.H_res, self.H_ims = ([], []) |
| self.actions, self.rewards = ([], []) |
| self.log_probs, self.values = ([], []) |
| self.dones = [] |
|
|
| def store(self, H_re, H_im, a, r, lp, v, done): |
| self.H_res.append(H_re) |
| self.H_ims.append(H_im) |
| self.actions.append(a) |
| self.rewards.append(r) |
| self.log_probs.append(lp) |
| self.values.append(v) |
| self.dones.append(done) |
|
|
| def __len__(self): |
| return len(self.rewards) |
|
|
| def to_arrays(self): |
| return (np.array(self.H_res, dtype=np.float32), np.array(self.H_ims, dtype=np.float32), np.array(self.actions, dtype=np.float32), np.array(self.rewards, dtype=np.float64), np.array(self.log_probs, dtype=np.float32), np.array(self.values, dtype=np.float64), np.array(self.dones, dtype=np.float64)) |
|
|
| class HDPPOAgentContinuous: |
|
|
| def __init__(self, cfg, seed=DEFAULT_SEED): |
| D = cfg['D'] |
| self.encoder = HDEncoderGradAdaptiveContinuous(cfg['feat_lo'], cfg['feat_hi'], D, seed, cfg['feature_fn'], cfg['beta'], phi_init=cfg.get('fpe_phi_init')) |
| self.actor = HDActorContinuous(D, cfg['action_dim'], cfg['log_std_init'], cfg['action_low'], cfg['action_high']) |
| self.encoder.link_actor(self.actor) |
| self.critic = HDCritic(D) |
| self.buffer = TrajectoryBuffer() |
| self.cfg = cfg |
| self.entropy_coef = cfg['entropy_coef'] |
| self.best_avg100 = -np.inf |
| self.best_snap = None |
|
|
| def select_action(self, state): |
| H_re, H_im = self.encoder.encode(state) |
| a_raw, a_env, lp = self.actor.sample(H_re, H_im) |
| val = self.critic.value(H_re, H_im) |
| return (a_raw, a_env, lp, val, H_re.copy(), H_im.copy()) |
|
|
| def store(self, H_re, H_im, a_raw, r, lp, v, done): |
| self.buffer.store(H_re, H_im, a_raw, r, lp, v, done) |
|
|
| def update(self, last_state, last_done): |
| buf = self.buffer |
| if not len(buf): |
| return (None, None) |
| cfg = self.cfg |
| H_re, H_im, a_vec, rewards, old_lps, values, dones = buf.to_arrays() |
| T = len(rewards) |
| if last_done: |
| last_val = 0.0 |
| else: |
| lr, li = self.encoder.encode(last_state) |
| last_val = self.critic.value(lr, li) |
| adv, returns = compute_gae_jit(rewards, values, dones, last_val, cfg['gamma'], cfg['lam']) |
| adv_std = float(adv.std()) |
| adv_n = np.clip((adv - adv.mean()) / (adv_std + 1e-08), -3.0, 3.0) if adv_std > 0.0001 else np.zeros(T, np.float64) |
| adv_n32 = adv_n.astype(np.float32) |
| ret32 = returns.astype(np.float32) |
| scale = np.float32(self.actor.sqrtD) |
| sqrtD = np.float32(self.critic.sqrtD) |
| act_lr = np.float32(cfg['actor_lr']) |
| crit_lr = np.float32(cfg['critic_lr']) |
| clip_e = np.float32(cfg['clip_eps']) |
| ent_c = np.float32(self.entropy_coef) |
| bias_lr = np.float32(self.critic.bias_lr) |
| bias = np.float32(self.critic.bias) |
| log_std = self.actor.log_std |
| policy_losses, value_losses = ([], []) |
| for _ in range(cfg['n_epochs']): |
| perm = np.random.permutation(T).astype(np.int32) |
| ls_grad = np.zeros(cfg['action_dim'], np.float32) |
| bias, pl, vl = ppo_epoch_jit(H_re, H_im, self.actor.W_re, self.actor.W_im, self.critic.W_re, self.critic.W_im, a_vec, adv_n32, ret32, old_lps, log_std, ls_grad, bias, perm, scale, sqrtD, act_lr, crit_lr, clip_e, bias_lr, ent_c, cfg['action_dim']) |
| policy_losses.append(float(pl)) |
| value_losses.append(float(vl)) |
| ls_step = cfg['log_std_lr'] * (ls_grad / np.float32(T)) |
| log_std = log_std + ls_step.astype(np.float32) |
| np.clip(log_std, cfg['log_std_min'], cfg['log_std_max'], out=log_std) |
| self.actor.log_std = log_std |
| self.critic.bias = float(bias) |
| self.buffer.reset() |
| return (float(np.mean(policy_losses)), float(np.mean(value_losses))) |
|
|
| def maybe_snapshot(self, avg100): |
| if avg100 > self.best_avg100: |
| self.best_avg100 = avg100 |
| self.best_snap = self.actor.snapshot() |
|
|
| def ema_restore(self): |
| if self.best_snap is not None: |
| self.actor.restore(self.best_snap, self.cfg['ema_alpha']) |
|
|
| class SystemMonitor: |
|
|
| def __init__(self): |
| self.proc = psutil.Process(os.getpid()) |
|
|
| def ram_mb(self): |
| return self.proc.memory_info().rss / 1024 ** 2 |
|
|
| def evaluate_agent(agent, n_episodes=20, seed_base=10000): |
| env = gym.make('Pendulum-v1') |
| rewards = np.empty(n_episodes, dtype=np.float64) |
| for i in range(n_episodes): |
| state, _ = env.reset(seed=seed_base + i) |
| ep_r = 0.0 |
| done = False |
| while not done: |
| Hr, Hi = agent.encoder.encode(state) |
| a = agent.actor.greedy(Hr, Hi) |
| state, reward, term, trunc, _ = env.step(a) |
| ep_r += float(reward) |
| done = term or trunc |
| rewards[i] = ep_r |
| 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, n_episodes=n_episodes, rewards=rewards.tolist()) |
|
|
| def prune_actor_global(checkpoint_npz, D_prime): |
| W_re, W_im = (checkpoint_npz['W_actor_re'], checkpoint_npz['W_actor_im']) |
| Wc_re, Wc_im = (checkpoint_npz['W_critic_re'], checkpoint_npz['W_critic_im']) |
| Phi = checkpoint_npz['fpe_phi'] |
| beta_base = float(checkpoint_npz['beta_base']) |
| log_std = checkpoint_npz['log_std'] |
| importance = np.sqrt((W_re ** 2).sum(axis=0) + (W_im ** 2).sum(axis=0)) |
| keep_idx = np.sort(np.argsort(-importance)[:D_prime]) |
| return dict(D=D_prime, beta=beta_base, log_std=log_std, fpe_phi=Phi[:, keep_idx], W_actor_re=W_re[:, keep_idx], W_actor_im=W_im[:, keep_idx], W_critic_re=Wc_re[keep_idx], W_critic_im=Wc_im[keep_idx], critic_bias=float(checkpoint_npz.get('critic_bias', 0.0))) |
|
|
| def train_one_seed(seed, total_timesteps, save_weights_path=None, warm_start=None, log_csv_path=None, eval_csv_path=None, eval_every_n_steps=None, D=None, verbose=True): |
| cfg = dict(PENDULUM_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) |
| elif D is not None: |
| cfg['D'] = int(D) |
| 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', 'best_avg100', 'policy_loss', 'value_loss', 'log_std_mean', 'entropy_coef', 'fps', 'ram_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']) |
| env = gym.make('Pendulum-v1') |
| np.random.seed(seed) |
| agent = HDPPOAgentContinuous(cfg, seed=seed) |
| if warm_start is not None: |
| agent.actor.W_re = np.asarray(warm_start['W_actor_re'], dtype=np.float32).copy() |
| agent.actor.W_im = np.asarray(warm_start['W_actor_im'], dtype=np.float32).copy() |
| agent.critic.W_re = np.asarray(warm_start['W_critic_re'], dtype=np.float32).copy() |
| agent.critic.W_im = np.asarray(warm_start['W_critic_im'], dtype=np.float32).copy() |
| agent.critic.bias = float(warm_start.get('critic_bias', 0.0)) |
| agent.actor.log_std = np.asarray(warm_start['log_std'], dtype=np.float32).copy() |
| if verbose: |
| print(f' Warm-started actor from provided checkpoint: D={cfg['D']}') |
| print(' Warm-started critic from provided checkpoint (warm-start, not reset)') |
| print(f' Warm-started log_std from provided checkpoint (warm-start, not reset): {agent.actor.log_std}') |
| if eval_csv_writer is not None: |
| post_prune_eval = evaluate_agent(agent) |
| 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) |
| ep = 0 |
| ep_r = 0.0 |
| steps_roll = 0 |
| global_step = 0 |
| update_count = 0 |
| ep_batch_rewards = [] |
| state, _ = env.reset(seed=seed) |
| steps_to_thresh = {T: None for T in REWARD_THRESHOLDS} |
| episodes_to_thresh = {T: None for T in REWARD_THRESHOLDS} |
| solved = False |
| t_solve = None |
| ep_solve = None |
| if verbose: |
| print('=' * 80) |
| print(f'HD-PPO -> Pendulum-v1') |
| print(f' D={cfg['D']} beta={cfg['beta']} total_timesteps={total_timesteps:,}') |
| print('=' * 80) |
| t0 = time.perf_counter() |
| while global_step < total_timesteps: |
| a_raw, a_env, lp, val, H_re, H_im = agent.select_action(state) |
| next_s, reward, term, trunc, _ = env.step(a_env) |
| done = term or trunc |
| global_step += 1 |
| agent.store(H_re, H_im, a_raw, reward, lp, val, done) |
| ep_r += reward |
| steps_roll += 1 |
| if eval_csv_writer is not None and next_eval_at is not None: |
| while global_step >= next_eval_at: |
| periodic_eval = evaluate_agent(agent) |
| 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 |
| if done: |
| ep += 1 |
| recent.append(ep_r) |
| ep_batch_rewards.append(ep_r) |
| if len(recent) == 100: |
| avg100 = sum(recent) / 100.0 |
| agent.maybe_snapshot(avg100) |
| if ep % ema_interval == 0: |
| agent.ema_restore() |
| for T in REWARD_THRESHOLDS: |
| if steps_to_thresh[T] is None and avg100 >= T: |
| steps_to_thresh[T] = global_step |
| episodes_to_thresh[T] = ep |
| if not solved and avg100 >= cfg['solve_thresh']: |
| solved = True |
| t_solve = time.perf_counter() - t0 |
| ep_solve = ep |
| if verbose: |
| print(f' *** FIRST SOLVE at step {global_step:,} (ep {ep}, avg100={avg100:.1f}) -- continuing to full budget ***') |
| state, _ = env.reset() |
| ep_r = 0.0 |
| else: |
| state = next_s |
| if steps_roll >= rollout_steps: |
| t_update_start = time.perf_counter() |
| policy_loss, value_loss = agent.update(state, done) |
| update_s = time.perf_counter() - t_update_start |
| update_count += 1 |
| steps_roll = 0 |
| if csv_writer is not None and len(recent) > 0: |
| fps = rollout_steps / (update_s + 1e-08) |
| csv_writer.writerow([global_step, int(time.perf_counter() - t0), ep, len(ep_batch_rewards), float(np.mean(recent)), float(np.max(recent)), float(np.min(recent)), agent.best_avg100, policy_loss, value_loss, float(np.mean(agent.actor.log_std)), agent.entropy_coef, fps, sysmon.ram_mb()]) |
| csv_file.flush() |
| ep_batch_rewards = [] |
| if verbose and len(recent) > 0 and (update_count % 20 == 0): |
| print(f' [step {global_step:>9,}] ep {ep:>5} avg100={float(np.mean(recent)):>+8.1f} best={agent.best_avg100:>+8.1f} log_std={float(np.mean(agent.actor.log_std)):+.3f}') |
| total_time = time.perf_counter() - t0 |
| if csv_file is not None: |
| csv_file.close() |
| final_avg = sum(recent) / len(recent) if recent else 0.0 |
| eval_final = evaluate_agent(agent) |
| eval_best = None |
| if agent.best_snap is not None: |
| saved_wr, saved_wi = (agent.actor.W_re.copy(), agent.actor.W_im.copy()) |
| agent.actor.W_re, agent.actor.W_im = agent.best_snap |
| saved_log_g_ema = agent.encoder._log_g_ema |
| agent.encoder._log_g_ema = 0.0 |
| eval_best = evaluate_agent(agent) |
| agent.encoder._log_g_ema = saved_log_g_ema |
| agent.actor.W_re, agent.actor.W_im = (saved_wr, saved_wi) |
| if eval_csv_file is not None: |
| eval_csv_file.close() |
| if save_weights_path is not None: |
| use_best = eval_best is not None and eval_best['mean_reward'] > eval_final['mean_reward'] |
| W_re_save, W_im_save = agent.best_snap if use_best else (agent.actor.W_re, agent.actor.W_im) |
| np.savez(save_weights_path, W_actor_re=W_re_save, W_actor_im=W_im_save, W_critic_re=agent.critic.W_re, W_critic_im=agent.critic.W_im, critic_bias=np.float32(agent.critic.bias), 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(agent.encoder.n_feat), action_dim=np.int32(cfg['action_dim']), log_std=agent.actor.log_std, 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), used_best_snapshot=np.bool_(use_best)) |
| if verbose: |
| print(f' Saved actor+critic -> {save_weights_path} ({os.path.getsize(save_weights_path) / 1024:.1f} KB)') |
| if verbose: |
| eb_str = f'{eval_best['mean_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_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}') |
| env.close() |
| return dict(seed=seed, solved=solved, total_time=total_time, final_avg100=final_avg, best_avg100=float(agent.best_avg100), eval_mean_final=eval_final['mean_reward'], eval_mean_best=eval_best['mean_reward'] if eval_best is not None else None, global_steps=global_step) |
| THIS_DIR = os.path.dirname(os.path.abspath(__file__)) |
| SEED = DEFAULT_SEED |
| STAGES = [(512, 1000000), (128, 1000000), (32, 1000000)] |
| OUT_DIR = os.path.join(THIS_DIR, 'prune_finetune') |
|
|
| def weights_path(D, stage_label): |
| return os.path.join(OUT_DIR, f'pendulum_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') |
| warmup_jit(STAGES[0][0], PENDULUM_CONFIG['rollout_steps'], ACTION_DIM) |
| stage_records = [] |
| prev_path = 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) |
| 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_one_seed(seed=SEED, total_timesteps=timesteps, warm_start=warm_start, save_weights_path=path, log_csv_path=curve_csv, eval_csv_path=eval_csv, eval_every_n_steps=50000, D=D if warm_start is None else None, verbose=True) |
| 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 |
| total_time = time.perf_counter() - t_chain0 |
| print('\n' + '=' * 100) |
| 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('=' * 100) |
| 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__': |
| main() |
|
|