ec1117

Untitled

Jan 16th, 2026
59
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Python 9.74 KB | None | 0 0
  1. import os
  2. from functools import partial
  3. from typing import Tuple
  4.  
  5. import flax
  6. import jax
  7. import jax.numpy as jnp
  8. import optax
  9.  
  10. import flowrl.module.initialization as init
  11. from flowrl.agent.base import BaseAgent
  12. from flowrl.agent.online.diffsr.network import FactorizedDDPM, update_factorized_ddpm
  13. from flowrl.config.online.algo.diffsr import DiffSRLDConfig
  14. from flowrl.flow.langevin_dynamics import IBCLangevinDynamics
  15. from flowrl.functional.ema import ema_update
  16. from flowrl.module.model import Model
  17. from flowrl.module.rff import RffEnsembleCritic
  18. from flowrl.types import Batch, Metric, Param, PRNGKey
  19.  
  20.  
  21. @partial(jax.jit, static_argnames=("training", "ld_temp", "num_samples"))
  22. def jit_sample_actions_ld(
  23.     rng: PRNGKey,
  24.     ld: Model,
  25.     ddpm_target: Model,
  26.     critic_target: Model,
  27.     scaler: jnp.ndarray,
  28.     obs: jnp.ndarray,
  29.     training: bool,
  30.     ld_temp: float,
  31.     num_samples: int,
  32. ) -> Tuple[PRNGKey, jnp.ndarray, jnp.ndarray]:
  33.     assert len(obs.shape) == 2
  34.     B = obs.shape[0]
  35.     rng, xT_rng = jax.random.split(rng)
  36.  
  37.     # sample
  38.     obs_repeat = obs[..., jnp.newaxis, :].repeat(num_samples, axis=-2)
  39.     xT = jax.random.normal(xT_rng, (*obs_repeat.shape[:-1], ld.x_dim))
  40.  
  41.     def model_fn(xt, input_t, condition):
  42.         original_shape = xt.shape[:-1]
  43.         xt = xt.reshape(-1, xt.shape[-1])
  44.         input_t = input_t.reshape(-1, 1)
  45.         condition = condition.reshape(-1, condition.shape[-1])
  46.         energy_and_grad_fn = jax.vmap(jax.value_and_grad(
  47.             lambda xt, t, condition: critic_target(ddpm_target(condition, xt, method="forward_phi")).mean()
  48.         ))
  49.         energy, grad = energy_and_grad_fn(xt, input_t, condition)
  50.         energy = energy.reshape(*original_shape, 1)
  51.         grad = grad.reshape(*original_shape, -1)
  52.         grad = grad / ld_temp / (scaler + 1e-8)
  53.         return energy, grad
  54.  
  55.     rng, actions, history = ld.sample(
  56.         rng,
  57.         model_fn,
  58.         xT,
  59.         obs_repeat,
  60.         training,
  61.     )
  62.     if num_samples == 1:
  63.         actions = actions[:, 0]
  64.     else:
  65.         feature = ddpm_target(obs_repeat, actions, method="forward_phi")
  66.         qs = critic_target(feature)
  67.         qs = qs.mean(axis=0).reshape(B, num_samples)
  68.         best_idx = qs.argmax(axis=-1)
  69.         actions = actions.reshape(B, num_samples, -1)[jnp.arange(B), best_idx]
  70.  
  71.     return rng, actions, history
  72.  
  73. @partial(jax.jit, static_argnames=("discount", "ld_temp"))
  74. def update_critic(
  75.     rng: PRNGKey,
  76.     ld: Model,
  77.     critic: Model,
  78.     critic_target: Model,
  79.     ddpm_target: Model,
  80.     scaler: jnp.ndarray,
  81.     batch: Batch,
  82.     discount: float,
  83.     ld_temp: float,
  84. ) -> Tuple[PRNGKey, Model, jnp.ndarray, Metric]:
  85.     rng, sample_rng, q_rng = jax.random.split(rng, 3)
  86.     rng, next_action, history = jit_sample_actions_ld(
  87.         rng,
  88.         ld,
  89.         ddpm_target,
  90.         critic_target,
  91.         scaler,
  92.         batch.next_obs,
  93.         training=False,
  94.         ld_temp=ld_temp,
  95.         num_samples=1,
  96.     )
  97.     next_feature = ddpm_target(batch.next_obs, next_action, method="forward_phi")
  98.     q_target = critic_target(next_feature)
  99.     q_index = jax.random.choice(q_rng, q_target.shape[0], shape=(2,), replace=False)
  100.     q_target = q_target[q_index].min(0)
  101.     q_target = batch.reward + discount * (1 - batch.terminal) * q_target
  102.  
  103.     feature = ddpm_target(batch.obs, batch.action, method="forward_phi")
  104.  
  105.     def critic_loss_fn(critic_params: Param, dropout_rng: PRNGKey) -> Tuple[jnp.ndarray, Metric]:
  106.         q_pred = critic.apply(
  107.             {"params": critic_params},
  108.             feature,
  109.             rngs={"dropout": dropout_rng},
  110.         )
  111.         critic_loss = ((q_pred - q_target[jnp.newaxis, :])**2).sum(0).mean()
  112.         return critic_loss, {
  113.             "loss/critic_loss": critic_loss,
  114.             "misc/q_mean": q_pred.mean(),
  115.             "misc/reward": batch.reward.mean(),
  116.         }
  117.  
  118.     # update scaler
  119.     q_grad = history[1] * ld_temp * scaler
  120.     new_scaler = 0.995 * scaler + 0.005 * jnp.abs(q_grad).mean()
  121.  
  122.     new_critic, metrics = critic.apply_gradient(critic_loss_fn)
  123.     metrics.update({
  124.         "misc_ld/scaler": new_scaler.mean(),
  125.         "misc_ld/q_grad_l1": jnp.abs(q_grad).mean(),
  126.     })
  127.     return rng, new_critic, new_scaler, metrics
  128.  
  129.  
  130. class DiffSRLDAgent(BaseAgent):
  131.     """
  132.    Diff-SR with Langevin Dynamics Agent.
  133.    """
  134.  
  135.     name = "DiffSRLDAgent"
  136.     model_names = ["ddpm", "ddpm_target", "actor", "critic", "critic_target"]
  137.  
  138.     def __init__(self, obs_dim: int, act_dim: int, cfg: DiffSRLDConfig, seed: int):
  139.         super().__init__(obs_dim, act_dim, cfg, seed)
  140.         self.cfg = cfg
  141.  
  142.         self.ddpm_coef = cfg.ddpm_coef
  143.         self.critic_coef = cfg.critic_coef
  144.         self.reward_coef = cfg.reward_coef
  145.         self.num_noises = cfg.num_noises
  146.         self.feature_dim = cfg.feature_dim
  147.         self.rff_dim = cfg.rff_dim
  148.         self.actor_update_freq = cfg.actor_update_freq
  149.         self.target_update_freq = cfg.target_update_freq
  150.  
  151.         # networks
  152.         self.rng, ddpm_rng, ddpm_init_rng, ld_rng, critic_rng = jax.random.split(self.rng, 5)
  153.         ddpm_def = FactorizedDDPM(
  154.             self.obs_dim,
  155.             self.act_dim,
  156.             self.feature_dim,
  157.             cfg.embed_dim,
  158.             cfg.phi_hidden_dims,
  159.             cfg.mu_hidden_dims,
  160.             cfg.reward_hidden_dims,
  161.             cfg.rff_dim,
  162.             cfg.num_noises,
  163.         )
  164.         self.ddpm = Model.create(
  165.             ddpm_def,
  166.             ddpm_rng,
  167.             inputs=(
  168.                 ddpm_init_rng,
  169.                 jnp.ones((1, self.obs_dim)),
  170.                 jnp.ones((1, self.act_dim)),
  171.                 jnp.ones((1, self.obs_dim)),
  172.             ),
  173.             optimizer=optax.adamw(learning_rate=cfg.feature_lr, weight_decay=cfg.wd),
  174.             clip_grad_norm=cfg.clip_grad_norm,
  175.         )
  176.         self.ddpm_target = Model.create(
  177.             ddpm_def,
  178.             ddpm_rng,
  179.             inputs=(
  180.                 ddpm_init_rng,
  181.                 jnp.ones((1, self.obs_dim)),
  182.                 jnp.ones((1, self.act_dim)),
  183.                 jnp.ones((1, self.obs_dim)),
  184.             ),
  185.         )
  186.  
  187.         critic_def = RffEnsembleCritic(
  188.             feature_dim=self.feature_dim,
  189.             hidden_dims=cfg.critic_hidden_dims,
  190.             rff_dim=cfg.rff_dim,
  191.             ensemble_size=cfg.critic_ensemble_size,
  192.             kernel_init=init.pytorch_kernel_init,
  193.             bias_init=init.pytorch_bias_init,
  194.         )
  195.         self.critic = Model.create(
  196.             critic_def,
  197.             critic_rng,
  198.             inputs=(jnp.ones((1, self.feature_dim)),),
  199.             optimizer=optax.adamw(learning_rate=cfg.critic_lr, weight_decay=cfg.wd),
  200.         )
  201.         self.critic_target = Model.create(
  202.             critic_def,
  203.             critic_rng,
  204.             inputs=(jnp.ones((1, self.feature_dim)),),
  205.         )
  206.  
  207.         self.ld = IBCLangevinDynamics.create(
  208.             network=flax.linen.Dense(1),
  209.             rng=ld_rng,
  210.             inputs=(jnp.ones((1, 1))),
  211.             x_dim=self.act_dim,
  212.             steps=self.cfg.ld.steps,
  213.             schedule=self.cfg.ld.schedule,
  214.             stepsize_init=self.cfg.ld.stepsize_init,
  215.             stepsize_final=self.cfg.ld.stepsize_final,
  216.             stepsize_decay=self.cfg.ld.stepsize_decay,
  217.             stepsize_power=self.cfg.ld.stepsize_power,
  218.             noise_scale=self.cfg.ld.noise_scale,
  219.             grad_clip=self.cfg.ld.grad_clip,
  220.             drift_clip=self.cfg.ld.drift_clip,
  221.             margin_clip=self.cfg.ld.margin_clip,
  222.         )
  223.         self.scaler = jnp.ones((1, ), dtype=jnp.float32)
  224.  
  225.         self._n_training_steps = 0
  226.  
  227.     def train_step(self, batch: Batch, step: int) -> Metric:
  228.         metrics = {}
  229.         self.rng, self.ddpm, ddpm_metrics = update_factorized_ddpm(
  230.             self.rng,
  231.             self.ddpm,
  232.             batch,
  233.             self.reward_coef,
  234.         )
  235.         metrics.update(ddpm_metrics)
  236.  
  237.         self.rng, self.critic, self.scaler, critic_metrics = update_critic(
  238.             self.rng,
  239.             self.ld,
  240.             self.critic,
  241.             self.critic_target,
  242.             self.ddpm_target,
  243.             self.scaler,
  244.             batch,
  245.             discount=self.cfg.discount,
  246.             ld_temp=self.cfg.ld_temp,
  247.         )
  248.         metrics.update(critic_metrics)
  249.  
  250.         if self._n_training_steps % self.target_update_freq == 0:
  251.             self.sync_target()
  252.  
  253.         self._n_training_steps += 1
  254.         return metrics
  255.  
  256.     def sample_actions(
  257.         self,
  258.         obs: jnp.ndarray,
  259.         deterministic: bool = True,
  260.         num_samples: int = 1,
  261.     ) -> Tuple[jnp.ndarray, Metric]:
  262.         if deterministic:
  263.             num_samples = self.cfg.num_samples
  264.         else:
  265.             num_samples = 1
  266.         self.rng, action, history = jit_sample_actions_ld(
  267.             self.rng,
  268.             self.ld,
  269.             self.ddpm_target,
  270.             self.critic_target,
  271.             self.scaler,
  272.             obs,
  273.             training=False,
  274.             ld_temp=self.cfg.ld_temp,
  275.             num_samples=num_samples,
  276.         )
  277.         if not deterministic:
  278.             action = action + self.cfg.exploration_noise * jax.random.normal(self.rng, action.shape)
  279.             action = jnp.clip(action, -1.0, 1.0)
  280.         return action, {}
  281.  
  282.     def sync_target(self):
  283.         self.critic_target = ema_update(self.critic, self.critic_target, self.cfg.ema)
  284.         self.ddpm_target = ema_update(self.ddpm, self.ddpm_target, self.cfg.feature_ema)
  285.  
  286.     def save(self, path: str):
  287.             super().save(path)
  288.             jnp.save(os.path.join(os.getcwd(), path, "scaler.npy"), self.scaler)
  289.  
  290.     def load(self, path: str):
  291.         super().load(path)
  292.         self.scaler = jnp.load(os.path.join(os.getcwd(), path, "scaler.npy"))
Advertisement
Add Comment
Please, Sign In to add comment