hdppo-Ant-v5 / gradient_adaptive_fpe_experiment.py
ChirathD's picture
Add hdppo-Ant-v5 package (weights, code, model card)
f7dd012 verified
Raw
History Blame Contribute Delete
10.7 kB
import csv
import os
import time
import numpy as np
import torch
import ant_hd_ppo as a
DEFAULT_D = 512
BETA_BASE = 0.8
BETA_EFF_MIN_MULT = 0.2
BETA_EFF_MAX_MULT = 4.0
G_EMA_DECAY = 0.99
class HDEncoderGradAdaptiveHybrid:
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))
def make_worker_encoder(cfg, master_seed):
if cfg.get('adaptive_beta', False):
return HDEncoderGradAdaptiveHybrid(cfg['feat_lo'], cfg['feat_hi'], cfg['D'], master_seed, cfg['feature_fn'], cfg['beta'], phi_init=cfg.get('fpe_phi_init'))
else:
return a.HDEncoderFPE(cfg['feat_lo'], cfg['feat_hi'], cfg['D'], cfg['beta'], seed=master_seed, feature_fn=cfg['feature_fn'], phi_init=cfg.get('fpe_phi_init'), beta_vec_init=cfg.get('beta_vec_init'))
def rollout_worker_adaptive(worker_id, cfg, conn, master_seed):
encoder = make_worker_encoder(cfg, master_seed)
adaptive = cfg.get('adaptive_beta', False)
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 = __import__('gymnasium').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 * a.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)))
a.rollout_worker = rollout_worker_adaptive
class HDPPOHybridAgentAdaptive(a.HDPPOHybridAgent):
def __init__(self, cfg, seed=a.DEFAULT_SEED):
super().__init__(cfg, seed=seed)
if cfg.get('adaptive_beta', False):
self.encoder = HDEncoderGradAdaptiveHybrid(cfg['feat_lo'], cfg['feat_hi'], cfg['D'], seed, cfg['feature_fn'], cfg['beta'], phi_init=cfg.get('fpe_phi_init'))
self.encoder.link_torch_actor(self.actor)
a.HDPPOHybridAgent = HDPPOHybridAgentAdaptive
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']) if 'beta_base' in checkpoint else float(np.atleast_1d(checkpoint['beta_bands'])[0])
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_base=beta_base, beta=beta_base, beta_vec=np.full(D_prime, beta_base, dtype=np.float32), fpe_phi=Phi[:, keep_idx], W_actor_re=W_re[keep_idx], W_actor_im=W_im[keep_idx])
def merge_beta_logs(log_dir, rollout_steps):
import glob
worker_files = sorted(glob.glob(os.path.join(log_dir, 'worker_*_beta_log.csv')))
if not worker_files:
return []
per_worker = []
for fp in worker_files:
with open(fp) as f:
rows = list(csv.DictReader(f))
per_worker.append(rows)
n_rollouts = min((len(rows) for rows in per_worker))
merged = []
for i in range(n_rollouts):
ns = np.array([float(rows[i]['n']) for rows in per_worker])
means = np.array([float(rows[i]['mean']) for rows in per_worker])
stds = np.array([float(rows[i]['std']) for rows in per_worker])
mins = np.array([float(rows[i]['min']) for rows in per_worker])
maxs = np.array([float(rows[i]['max']) for rows in per_worker])
total_n = ns.sum()
pooled_mean = float((ns * means).sum() / total_n)
pooled_var = float((ns * (stds ** 2 + means ** 2)).sum() / total_n - pooled_mean ** 2)
merged.append(dict(rollout_idx=i, global_step=(i + 1) * rollout_steps, n=int(total_n), mean=pooled_mean, std=float(np.sqrt(max(pooled_var, 0.0))), min=float(mins.min()), max=float(maxs.max())))
return merged
def run_condition(seed, adaptive, total_timesteps, verbose=False, beta_log_dir=None, warm_start=None, save_weights_path=None, D=None, log_csv_path=None, eval_csv_path=None, eval_every_n_steps=None):
target_D = int(warm_start['D']) if warm_start is not None else int(D) if D is not None else DEFAULT_D
a.ANT_CONFIG['D'] = target_D
a.ANT_CONFIG['beta'] = BETA_BASE
a.ANT_CONFIG['adaptive_beta'] = adaptive
a.ANT_CONFIG.pop('fpe_phi_init', None)
a.ANT_CONFIG.pop('beta_vec_init', None)
if adaptive and beta_log_dir:
os.makedirs(beta_log_dir, exist_ok=True)
a.ANT_CONFIG['adaptive_beta_log_dir'] = beta_log_dir
else:
a.ANT_CONFIG.pop('adaptive_beta_log_dir', None)
t0 = time.perf_counter()
summary = a.train(total_timesteps=total_timesteps, verbose=verbose, seed=seed, wandb_mode='disabled', save_weights_path=save_weights_path, warm_start=warm_start, log_csv_path=log_csv_path, eval_csv_path=eval_csv_path, eval_every_n_steps=eval_every_n_steps)
wall = time.perf_counter() - t0
beta_log = None
if adaptive and beta_log_dir:
beta_log = merge_beta_logs(beta_log_dir, a.ANT_CONFIG['rollout_steps'])
return dict(curve_steps=summary.get('curve_steps'), curve_avg100=summary.get('curve_avg100'), beta_log=beta_log, seed=seed, adaptive=adaptive, total_time=wall, final_avg100=summary.get('final_avg100'), best_avg100=summary.get('best_avg100'), eval_mean_final=summary.get('eval_mean_final'), eval_mean_best=summary.get('eval_mean_best'), critic_state_dict=summary.get('critic_state_dict'))