hdppo-InvertedPendulum-v5 / train_hdppo.py
ChirathD's picture
Add hdppo-InvertedPendulum-v5 package (weights, code, model card)
0c0f8ca verified
Raw
History Blame Contribute Delete
35.3 kB
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()