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, bias, perm, scale, sqrtD, actor_lr, critic_lr, clip_eps, ent_coef, bias_lr, n_actions): T = H_re.shape[0] K = n_actions D = H_re.shape[1] 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] z = np.empty(K, numba.float32) for k in range(K): s = numba.float32(0.0) for d in range(D): s += W_re[k, d] * Hr[d] + W_im[k, d] * Hi[d] z[k] = s / scale z_max = z[0] for k in range(1, K): if z[k] > z_max: z_max = z[k] ex = np.empty(K, numba.float32) ex_sum = numba.float32(0.0) for k in range(K): ex[k] = np.exp(z[k] - z_max) ex_sum += ex[k] pi = ex / ex_sum a = a_vec[idx] new_lp = np.log(pi[a] + numba.float32(1e-08)) ratio = np.exp(new_lp - old_lps[idx]) if ratio > numba.float32(10.0): ratio = numba.float32(10.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) mean_ne = numba.float32(0.0) for k in range(K): mean_ne += pi[k] * np.log(pi[k] + numba.float32(1e-08)) for k in range(K): delta = (numba.float32(1.0) if k == a else numba.float32(0.0)) - pi[k] pg = actor_lr * surr * delta / scale ent = actor_lr * ent_coef * pi[k] * (mean_ne - np.log(pi[k] + numba.float32(1e-08))) / scale coef = pg + ent for d in range(D): W_re[k, d] += coef * Hr[d] W_im[k, d] += coef * Hi[d] 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) 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, n_actions): Hr = np.zeros((T, D), np.float32) Hi = np.zeros((T, D), np.float32) Wr = np.zeros((n_actions, D), np.float32) Wi = np.zeros((n_actions, D), np.float32) Cr = np.zeros(D, np.float32) Ci = np.zeros(D, np.float32) av = np.zeros(T, np.int32) an = np.zeros(T, np.float32) rt = np.zeros(T, np.float32) ol = np.zeros(T, np.float32) pm = np.arange(T, dtype=np.int32) sc = np.float32(np.sqrt(D) * 0.5) sq = np.float32(np.sqrt(D)) ppo_epoch_jit(Hr, Hi, Wr, Wi, Cr, Ci, av, an, rt, ol, np.float32(0.0), pm, sc, sq, np.float32(0.001), np.float32(0.005), np.float32(0.2), np.float32(0.02), np.float32(0.05), n_actions) 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 lunarlander_features(obs): x, y, vx, vy, ang, av, leg1, leg2 = (float(obs[0]), float(obs[1]), float(obs[2]), float(obs[3]), float(obs[4]), float(obs[5]), float(obs[6]), float(obs[7])) return np.array([x, y, vx, vy, ang, av, leg1, leg2, np.sin(ang), np.cos(ang), x * x + y * y, vx * vx + vy * vy], dtype=np.float32) OBS_DIM = 12 N_ACTIONS = 4 BASE_BETA = 1.0 LUNARLANDER_CONFIG = dict(feat_lo=[-1.5, -0.5, -2.0, -2.0, -3.1416, -2.0, 0.0, 0.0, -1.0, -1.0, 0.0, 0.0], feat_hi=[1.5, 1.5, 2.0, 2.0, 3.1416, 2.0, 1.0, 1.0, 1.0, 1.0, 4.5, 8.0], feature_fn=lunarlander_features, n_actions=N_ACTIONS, D=512, beta=BASE_BETA, rollout_steps=1024, actor_lr=0.001, critic_lr=0.005, n_epochs=6, temperature=0.5, entropy_coef=0.02, entropy_decay=0.997, entropy_min=0.001, clip_eps=0.2, gamma=0.99, lam=0.95, solve_thresh=200.0, ema_interval=100, ema_alpha=0.15) REWARD_THRESHOLDS = [-200, -100, 0, 100, 150, 200] DEFAULT_SEED = 123 BETA_EFF_MIN_MULT = 0.2 BETA_EFF_MAX_MULT = 4.0 G_EMA_DECAY = 0.99 class HDEncoderGradAdaptiveDiscrete: 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) 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): actor = self._actor 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 @ actor.W_re.T + dH_im_ds @ actor.W_im.T) / actor.sqrtD_tau 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 HDActorDiscrete: def __init__(self, D, n_actions, temperature): self.sqrtD_tau = float(np.sqrt(D)) * temperature self.n_actions = n_actions self.W_re = np.zeros((n_actions, D), dtype=np.float32) self.W_im = np.zeros((n_actions, D), dtype=np.float32) def probs(self, H_re, H_im): z = (self.W_re @ H_re + self.W_im @ H_im) / self.sqrtD_tau z -= z.max() ex = np.exp(z) return ex / ex.sum() def sample(self, H_re, H_im): pi = self.probs(H_re, H_im) a = int(np.random.choice(self.n_actions, p=pi)) return (a, float(np.log(pi[a] + 1e-08)), pi) def argmax(self, H_re, H_im): z = self.W_re @ H_re + self.W_im @ H_im return int(np.argmax(z)) 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, self.pis = ([], []) def store(self, H_re, H_im, a, r, lp, v, done, pi): 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) self.pis.append(pi) 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.int32), 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), np.array(self.pis, dtype=np.float32)) class HDPPOAgentDiscrete: def __init__(self, cfg, seed=DEFAULT_SEED): D = cfg['D'] self.encoder = HDEncoderGradAdaptiveDiscrete(cfg['feat_lo'], cfg['feat_hi'], D, seed, cfg['feature_fn'], cfg['beta'], phi_init=cfg.get('fpe_phi_init')) self.actor = HDActorDiscrete(D, cfg['n_actions'], cfg['temperature']) 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, lp, pi = self.actor.sample(H_re, H_im) val = self.critic.value(H_re, H_im) return (a, lp, val, H_re.copy(), H_im.copy(), pi) def store(self, H_re, H_im, a, r, lp, v, done, pi): self.buffer.store(H_re, H_im, a, r, lp, v, done, pi) 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_tau) 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) policy_losses, value_losses = ([], []) for _ in range(cfg['n_epochs']): perm = np.random.permutation(T).astype(np.int32) 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, bias, perm, scale, sqrtD, act_lr, crit_lr, clip_e, ent_c, bias_lr, cfg['n_actions']) policy_losses.append(float(pl)) value_losses.append(float(vl)) self.critic.bias = float(bias) self.entropy_coef = max(self.entropy_coef * cfg['entropy_decay'], cfg['entropy_min']) 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('LunarLander-v3') 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.argmax(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']) 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, 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(LUNARLANDER_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', '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('LunarLander-v3') np.random.seed(seed) agent = HDPPOAgentDiscrete(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)) if verbose: print(f' Warm-started actor from provided checkpoint: D={cfg['D']}') print(' Warm-started critic from provided checkpoint (warm-start, not reset)') 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 ep_len = 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 -> LunarLander-v3') print(f' D={cfg['D']} beta={cfg['beta']} total_timesteps={total_timesteps:,}') print('=' * 80) t0 = time.perf_counter() while global_step < total_timesteps: a, lp, val, H_re, H_im, pi = agent.select_action(state) next_s, reward, term, trunc, _ = env.step(a) done = term or trunc global_step += 1 agent.store(H_re, H_im, a, reward, lp, val, done, pi) ep_r += reward ep_len += 1 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, ep_len = (0.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, 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} ent={agent.entropy_coef:.4f}') 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), n_actions=np.int32(cfg['n_actions']), 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'lunarlander_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], LUNARLANDER_CONFIG['rollout_steps'], N_ACTIONS) 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()