ec1117

LD

Jan 16th, 2026
78
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
Python 19.84 KB | None | 0 0
  1. from functools import partial
  2. from typing import Callable, Optional, Sequence, Tuple
  3.  
  4. import flax.linen as nn
  5. import jax
  6. import jax.numpy as jnp
  7. import optax
  8. from flax.struct import PyTreeNode, dataclass, field
  9. from flax.training.train_state import TrainState
  10.  
  11. from flowrl.flow.continuous_ddpm import cosine_noise_schedule, linear_noise_schedule
  12. from flowrl.module.model import Model
  13. from flowrl.types import *
  14.  
  15. # ======= Langevin Dynamics Sampling =======
  16.  
  17. @dataclass
  18. class LangevinDynamics(Model):
  19.     state: TrainState
  20.     dropout_rng: PRNGKey = field(pytree_node=True)
  21.     x_dim: int = field(pytree_node=False, default=None)
  22.     grad_prediction: bool = field(pytree_node=False, default=True)
  23.     steps: int = field(pytree_node=False, default=None)
  24.     step_size: float = field(pytree_node=False, default=None)
  25.     noise_scale: float = field(pytree_node=False, default=None)
  26.     clip_sampler: bool = field(pytree_node=False, default=None)
  27.     x_min: float = field(pytree_node=False, default=None)
  28.     x_max: float = field(pytree_node=False, default=None)
  29.  
  30.     @classmethod
  31.     def create(
  32.         cls,
  33.         network: nn.Module,
  34.         rng: PRNGKey,
  35.         inputs: Sequence[jnp.ndarray],
  36.         x_dim: int,
  37.         grad_prediction: bool = True,
  38.         steps: int = 100,
  39.         step_size: float = 0.01,
  40.         noise_scale: float = 1.0,
  41.         clip_sampler: bool = False,
  42.         x_min: Optional[float] = None,
  43.         x_max: Optional[float] = None,
  44.         optimizer: Optional[optax.GradientTransformation] = None,
  45.         clip_grad_norm: float = None
  46.     ) -> 'LangevinDynamics':
  47.         ret = super().create(network, rng, inputs, optimizer, clip_grad_norm)
  48.  
  49.         return ret.replace(
  50.             x_dim=x_dim,
  51.             grad_prediction=grad_prediction,
  52.             steps=steps,
  53.             step_size=step_size,
  54.             noise_scale=noise_scale,
  55.             clip_sampler=clip_sampler,
  56.             x_min=x_min,
  57.             x_max=x_max,
  58.         )
  59.  
  60.     @partial(jax.jit, static_argnames=("training"))
  61.     def compute_grad(
  62.         self,
  63.         x: jnp.ndarray,
  64.         i: int,
  65.         condition: Optional[jnp.ndarray] = None,
  66.         training: bool = False,
  67.         params: Optional[Param] = None,
  68.         dropout_rng: Optional[PRNGKey] = None
  69.     ) -> jnp.ndarray:
  70.         original_shape = x.shape[:-1]
  71.         t = i * jnp.ones((*x.shape[:-1], 1), dtype=jnp.int32)
  72.  
  73.         x = x.reshape(-1, x.shape[-1])
  74.         t = t.reshape(-1, 1)
  75.         condition = condition.reshape(-1, condition.shape[-1])
  76.         if self.grad_prediction:
  77.             if training:
  78.                 grad = self.apply(
  79.                     {"params": params}, x, t, condition=condition, training=training, rngs={"dropout": dropout_rng}
  80.                 )
  81.             else:
  82.                 grad = self(x, t, condition=condition, training=training)
  83.             energy = jnp.zeros_like((*x.shape[:-1], 1), dtype=jnp.float32)
  84.         else:
  85.             if training:
  86.                 energy_and_grad_fn = jax.vmap(jax.value_and_grad(lambda x, t, condition: self.apply(
  87.                     {"params": params}, x, t, condition=condition, training=training, rngs={"dropout": dropout_rng}
  88.                 ).mean()))
  89.             else:
  90.                 energy_and_grad_fn = jax.vmap(jax.value_and_grad(lambda x, t, condition: self(x, t, condition=condition, training=training).mean()))
  91.             energy, grad = energy_and_grad_fn(x, t, condition)
  92.         return grad.reshape(*original_shape, self.x_dim), energy.reshape(*original_shape, 1)
  93.  
  94.     @partial(jax.jit, static_argnames=("training", "steps","step_size","noise_scale"))
  95.     def sample(
  96.         self,
  97.         rng: PRNGKey,
  98.         x_init: jnp.ndarray,
  99.         condition: Optional[jnp.ndarray] = None,
  100.         training: bool = False,
  101.         steps: Optional[int] = None,
  102.         step_size: Optional[float] = None,
  103.         noise_scale: Optional[float] = None,
  104.         params: Optional[Param] = None,
  105.     ) -> Tuple[PRNGKey, jnp.ndarray, Optional[jnp.ndarray]]:
  106.         steps = steps or self.steps
  107.         step_size = step_size or self.step_size
  108.         noise_scale = noise_scale or self.noise_scale
  109.  
  110.         def fn(input_tuple, i):
  111.             rng_, xt = input_tuple
  112.             rng_, noise_rng, dropout_rng_ = jax.random.split(rng_, 3)
  113.  
  114.             grad, energy = self.compute_grad(xt, i, condition=condition, training=training, params=params, dropout_rng=dropout_rng_)
  115.  
  116.             xt_1 = xt + step_size * grad
  117.             if self.clip_sampler:
  118.                 xt_1 = jnp.clip(xt_1, self.x_min, self.x_max)
  119.             noise = jax.random.normal(noise_rng, xt_1.shape, dtype=jnp.float32)
  120.             xt_1 += (i>1) * jnp.sqrt(2 * step_size * noise_scale) * noise
  121.  
  122.             return (rng_, xt_1), (xt, grad, energy)
  123.  
  124.         output, history = jax.lax.scan(fn, (rng, x_init), jnp.arange(steps, 0, -1), unroll=True)
  125.         rng, action = output
  126.         return rng, action, history
  127.  
  128.  
  129. @dataclass
  130. class AnnealedLangevinDynamics(LangevinDynamics):
  131.     state: TrainState
  132.     dropout_rng: PRNGKey = field(pytree_node=True)
  133.     x_dim: int = field(pytree_node=False, default=None)
  134.     grad_prediction: bool = field(pytree_node=False, default=True)
  135.     steps: int = field(pytree_node=False, default=None)
  136.     step_size: float = field(pytree_node=False, default=None)
  137.     noise_scale: float = field(pytree_node=False, default=None)
  138.     clip_sampler: bool = field(pytree_node=False, default=None)
  139.     x_min: float = field(pytree_node=False, default=None)
  140.     x_max: float = field(pytree_node=False, default=None)
  141.     t_schedule_n: float = field(pytree_node=False, default=None)
  142.     t_diffusion: Tuple[float, float] = field(pytree_node=False, default=None)
  143.     noise_schedule_func: Callable = field(pytree_node=False, default=None)
  144.  
  145.     @classmethod
  146.     def create(
  147.         cls,
  148.         network: nn.Module,
  149.         rng: PRNGKey,
  150.         inputs: Sequence[jnp.ndarray],
  151.         x_dim: int,
  152.         grad_prediction: bool,
  153.         steps: int,
  154.         step_size: float,
  155.         noise_scale: float,
  156.         noise_schedule: str,
  157.         noise_schedule_params: Optional[Dict]=None,
  158.         clip_sampler: bool = False,
  159.         x_min: Optional[float] = None,
  160.         x_max: Optional[float] = None,
  161.         t_schedule_n: float=1.0,
  162.         epsilon: float=0.001,
  163.         optimizer: Optional[optax.GradientTransformation]=None,
  164.         clip_grad_norm: float=None
  165.     ) -> 'AnnealedLangevinDynamics':
  166.         ret = super().create(
  167.             network,
  168.             rng,
  169.             inputs,
  170.             x_dim,
  171.             grad_prediction,
  172.             steps,
  173.             step_size,
  174.             noise_scale,
  175.             clip_sampler,
  176.             x_min,
  177.             x_max,
  178.             optimizer,
  179.             clip_grad_norm,
  180.         )
  181.  
  182.         if noise_schedule_params is None:
  183.             noise_schedule_params = {}
  184.         if noise_schedule == "cosine":
  185.             t_diffusion = [epsilon, 0.9946]
  186.         else:
  187.             t_diffusion = [epsilon, 1.0]
  188.         if noise_schedule == "linear":
  189.             noise_schedule_func = partial(linear_noise_schedule, **noise_schedule_params)
  190.         elif noise_schedule == "cosine":
  191.             noise_schedule_func = partial(cosine_noise_schedule, **noise_schedule_params)
  192.         elif noise_schedule == "none":
  193.             noise_schedule_func = lambda t: (jnp.ones_like(t), jnp.zeros_like(t))
  194.         else:
  195.             raise NotImplementedError(f"Unsupported noise schedule: {noise_schedule}")
  196.  
  197.         return ret.replace(
  198.             t_schedule_n=t_schedule_n,
  199.             t_diffusion=t_diffusion,
  200.             noise_schedule_func=noise_schedule_func,
  201.         )
  202.  
  203.     @partial(jax.jit, static_argnames=("training"))
  204.     def compute_grad(
  205.         self,
  206.         x: jnp.ndarray,
  207.         t: jnp.ndarray,
  208.         condition: Optional[jnp.ndarray] = None,
  209.         training: bool = False,
  210.         params: Optional[Param] = None,
  211.         dropout_rng: Optional[PRNGKey] = None
  212.     ) -> jnp.ndarray:
  213.         original_shape = x.shape[:-1]
  214.         t = t * jnp.ones((*x.shape[:-1], 1), dtype=jnp.int32)
  215.  
  216.         x = x.reshape(-1, x.shape[-1])
  217.         t = t.reshape(-1, 1)
  218.         condition = condition.reshape(-1, condition.shape[-1])
  219.         if self.grad_prediction:
  220.             if training:
  221.                 grad = self.apply(
  222.                     {"params": params}, x, t, condition=condition, training=training, rngs={"dropout": dropout_rng}
  223.                 )
  224.             else:
  225.                 grad = self(x, t, condition=condition, training=training)
  226.             energy = jnp.zeros((*x.shape[:-1], 1), dtype=jnp.float32)
  227.         else:
  228.             if training:
  229.                 energy_and_grad_fn = jax.vmap(jax.value_and_grad(lambda x, t, condition: self.apply(
  230.                     {"params": params}, x, t, condition=condition, training=training, rngs={"dropout": dropout_rng}
  231.                 ).mean()))
  232.             else:
  233.                 energy_and_grad_fn = jax.vmap(jax.value_and_grad(lambda x, t, condition: self(x, t, condition=condition, training=training).mean()))
  234.             energy, grad = energy_and_grad_fn(x, t, condition)
  235.         # alpha, sigma = self.noise_schedule_func(t)
  236.         # grad = alpha * grad - sigma * x
  237.         return grad.reshape(*original_shape, self.x_dim), energy.reshape(*original_shape, 1)
  238.  
  239.     def add_noise(self, rng: PRNGKey, x: jnp.ndarray) -> Tuple[PRNGKey, jnp.ndarray, jnp.ndarray, jnp.ndarray]:
  240.         rng, t_rng, noise_rng = jax.random.split(rng, 3)
  241.         t = jax.random.uniform(t_rng, (*x.shape[:-1], 1), dtype=jnp.float32, minval=self.t_diffusion[0], maxval=self.t_diffusion[1])
  242.         alpha, sigma = self.noise_schedule_func(t)
  243.         eps = jax.random.normal(noise_rng, x.shape, dtype=jnp.float32)
  244.         xt = alpha * x + sigma * eps
  245.         return rng, xt, t, eps
  246.  
  247.     @partial(jax.jit, static_argnames=("training", "steps","step_size","noise_scale"))
  248.     def sample(
  249.         self,
  250.         rng: PRNGKey,
  251.         x_init: jnp.ndarray,
  252.         condition: Optional[jnp.ndarray] = None,
  253.         training: bool = False,
  254.         steps: Optional[int] = None,
  255.         step_size: Optional[float] = None,
  256.         noise_scale: Optional[float] = None,
  257.         params: Optional[Param] = None,
  258.     ) -> Tuple[PRNGKey, jnp.ndarray, Optional[jnp.ndarray]]:
  259.         steps = steps or self.steps
  260.         # step_size = step_size or self.step_size
  261.         # noise_scale = noise_scale or self.noise_scale
  262.         t_schedule_n = 1.0
  263.         from flowrl.flow.continuous_ddpm import quad_t_schedule
  264.         ts = quad_t_schedule(steps, n=t_schedule_n, tmin=self.t_diffusion[0], tmax=self.t_diffusion[1])
  265.         alpha_hats = self.noise_schedule_func(ts)[0] ** 2
  266.         alphas = alpha_hats[1:] / alpha_hats[:-1]
  267.         alphas = jnp.concat([jnp.ones((1, )), alphas], axis=0)
  268.         betas = 1 - alphas
  269.         alpha1, alpha2 = self.noise_schedule_func(ts)
  270.  
  271.         t_proto = jnp.ones((*x_init.shape[:-1], 1), dtype=jnp.int32)
  272.  
  273.         def fn(input_tuple, i):
  274.             rng_, xt = input_tuple
  275.             rng_, dropout_rng_, key_ = jax.random.split(rng_, 3)
  276.             input_t = t_proto * ts[i]
  277.  
  278.             q_grad, energy = self.compute_grad(xt, ts[i], condition=condition, training=training, params=params, dropout_rng=dropout_rng_)
  279.             eps_theta = q_grad
  280.  
  281.             x0_hat = (xt - jnp.sqrt(1 - alpha_hats[i]) * eps_theta) / jnp.sqrt(alpha_hats[i])
  282.             x0_hat = jnp.clip(x0_hat, self.x_min, self.x_max) if self.clip_sampler else x0_hat
  283.  
  284.             mean_coef1 = jnp.sqrt(alpha_hats[i-1]) * betas[i] / (1 - alpha_hats[i])
  285.             mean_coef2 = jnp.sqrt(alphas[i]) * (1 - alpha_hats[i-1]) / (1 - alpha_hats[i])
  286.             xt_1 = mean_coef1 * x0_hat + mean_coef2 * xt
  287.             xt_1 += (i>1) * jnp.sqrt(betas[i]) * jax.random.normal(key_, xt_1.shape)
  288.  
  289.             return (rng_, xt_1), (xt, eps_theta, energy)
  290.  
  291.         output, history = jax.lax.scan(fn, (rng, x_init), jnp.arange(steps, 0, -1), unroll=True)
  292.         rng, action = output
  293.         return rng, action, history
  294.  
  295.  
  296. from flowrl.flow.continuous_ddpm import ContinuousDDPM, quad_t_schedule
  297.  
  298.  
  299. @dataclass
  300. class ContinuousDDPMLD(ContinuousDDPM):
  301.     state: TrainState
  302.     dropout_rng: PRNGKey = field(pytree_node=True)
  303.     x_dim: int = field(pytree_node=False, default=None)
  304.     steps: int = field(pytree_node=False, default=None)
  305.     clip_sampler: bool = field(pytree_node=False, default=None)
  306.     x_min: float = field(pytree_node=False, default=None)
  307.     x_max: float = field(pytree_node=False, default=None)
  308.     t_schedule_n: float = field(pytree_node=False, default=None)
  309.     t_diffusion: Tuple[float, float] = field(pytree_node=False, default=None)
  310.     noise_schedule_func: Callable = field(pytree_node=False, default=None)
  311.  
  312.     @partial(jax.jit, static_argnames=("model_fn", "training", "solver", "steps", "t_schedule_n"))
  313.     def sample(
  314.         self,
  315.         rng: PRNGKey,
  316.         model_fn: Callable,
  317.         xT: jnp.ndarray,
  318.         condition: Optional[jnp.ndarray] = None,
  319.         training: bool = False,
  320.         solver: str = "ddpm",
  321.         steps: Optional[int] = None,
  322.         t_schedule_n: Optional[float] = None,
  323.         params: Optional[Param] = None,
  324.     ) -> Tuple[PRNGKey, jnp.ndarray, Tuple[jnp.ndarray, jnp.ndarray]]:
  325.         steps = steps or self.steps
  326.         t_schedule_n = t_schedule_n or self.t_schedule_n
  327.  
  328.         ts = quad_t_schedule(steps, n=t_schedule_n, tmin=self.t_diffusion[0], tmax=self.t_diffusion[1])
  329.         alpha_hats = self.noise_schedule_func(ts)[0] ** 2
  330.         alphas = alpha_hats[1:] / alpha_hats[:-1]
  331.         alphas = jnp.concat([jnp.ones((1, )), alphas], axis=0)
  332.         betas = 1 - alphas
  333.         alpha1, alpha2 = self.noise_schedule_func(ts)
  334.  
  335.         t_proto = jnp.ones((*xT.shape[:-1], 1), dtype=jnp.int32)
  336.  
  337.         def fn(input_tuple, i):
  338.             rng_, xt = input_tuple
  339.             rng_, dropout_rng_, key_ = jax.random.split(rng_, 3)
  340.             input_t = t_proto * ts[i]
  341.  
  342.             energy, q_grad = model_fn(xt, input_t, condition=condition)
  343.             # q_grad = alpha1[i] * q_grad - alpha2[i] * xt
  344.  
  345.             if solver == "ddpm":
  346.                 eps_theta = - alpha2[i] * q_grad
  347.                 x0_hat = (xt - jnp.sqrt(1 - alpha_hats[i]) * eps_theta) / jnp.sqrt(alpha_hats[i])
  348.                 x0_hat = jnp.clip(x0_hat, self.x_min, self.x_max) if self.clip_sampler else x0_hat
  349.  
  350.                 mean_coef1 = jnp.sqrt(alpha_hats[i-1]) * betas[i] / (1 - alpha_hats[i])
  351.                 mean_coef2 = jnp.sqrt(alphas[i]) * (1 - alpha_hats[i-1]) / (1 - alpha_hats[i])
  352.                 xt_1 = mean_coef1 * x0_hat + mean_coef2 * xt
  353.                 xt_1 += (i>1) * jnp.sqrt(betas[i]) * jax.random.normal(key_, xt_1.shape)
  354.             else:
  355.                 raise NotImplementedError(f"Unsupported solver: {solver}")
  356.  
  357.             return (rng_, xt_1), (eps_theta, q_grad)
  358.  
  359.         output, history = jax.lax.scan(fn, (rng, xT), jnp.arange(steps, 0, -1), unroll=True)
  360.         rng, action = output
  361.         return rng, action, history
  362.  
  363.  
  364. @dataclass
  365. class IBCLangevinDynamics(Model):
  366.     state: TrainState
  367.     dropout_rng: PRNGKey = field(pytree_node=True)
  368.     x_dim: int = field(pytree_node=False, default=None)
  369.     steps: int = field(pytree_node=False, default=None)
  370.     schedule: str = field(pytree_node=False, default=None)
  371.     stepsize_init: float = field(pytree_node=False, default=None)
  372.     stepsize_final: float = field(pytree_node=False, default=None)
  373.     stepsize_decay: float = field(pytree_node=False, default=None)
  374.     stepsize_power: float = field(pytree_node=False, default=None)
  375.     noise_scale: float = field(pytree_node=False, default=None)
  376.     grad_clip: float | None = field(pytree_node=False, default=None)
  377.     drift_clip: float | None = field(pytree_node=False, default=None)
  378.     margin_clip: float | None = field(pytree_node=False, default=None)
  379.     x_min: float = field(pytree_node=False, default=None)
  380.     x_max: float = field(pytree_node=False, default=None)
  381.  
  382.     @classmethod
  383.     def create(
  384.         cls,
  385.         network: nn.Module,
  386.         rng: PRNGKey,
  387.         inputs: Sequence[jnp.ndarray],
  388.         x_dim: int,
  389.         steps: int = 100,
  390.         schedule: str = "polynomial",
  391.         stepsize_init: float = 1e-1,
  392.         stepsize_final: float = 1e-5,
  393.         stepsize_decay: float = 0.8,
  394.         stepsize_power: float = 2.0,
  395.         noise_scale: float = 1.0,
  396.         grad_clip: float | None = None,
  397.         drift_clip: float | None = None,
  398.         margin_clip: float | None = None,
  399.         optimizer: Optional[optax.GradientTransformation] = None,
  400.         clip_grad_norm: float = None
  401.     ) -> 'LangevinDynamics':
  402.         ret = super().create(network, rng, inputs, optimizer, clip_grad_norm)
  403.  
  404.         return ret.replace(
  405.             x_dim=x_dim,
  406.             steps=steps,
  407.             schedule=schedule,
  408.             stepsize_init=stepsize_init,
  409.             stepsize_final=stepsize_final,
  410.             stepsize_decay=stepsize_decay,
  411.             stepsize_power=stepsize_power,
  412.             noise_scale=noise_scale,
  413.             grad_clip=grad_clip,
  414.             drift_clip=drift_clip,
  415.             margin_clip=margin_clip,
  416.         )
  417.  
  418.     @partial(jax.jit, static_argnames=("model_fn", "training", "solver"))
  419.     def sample(
  420.         self,
  421.         rng: PRNGKey,
  422.         model_fn: Callable,
  423.         xT: jnp.ndarray,
  424.         condition: Optional[jnp.ndarray] = None,
  425.         training: bool = False,
  426.         solver: str = "ddpm",
  427.     ) -> Tuple[PRNGKey, jnp.ndarray, Tuple[jnp.ndarray, jnp.ndarray]]:
  428.         if self.schedule == "polynomial":
  429.             stepsizes = polynomial_schedule(
  430.                 self.stepsize_init,
  431.                 self.stepsize_final,
  432.                 self.stepsize_power,
  433.                 self.steps,
  434.             )
  435.         else:
  436.             stepsizes = exponential_schedule(
  437.                 self.stepsize_init,
  438.                 self.stepsize_decay,
  439.                 self.steps,
  440.             )
  441.         t_proto = jnp.ones((*xT.shape[:-1], 1), dtype=jnp.int32)
  442.  
  443.         def fn(input_tuple, i):
  444.             rng_, xt = input_tuple
  445.             rng_, dropout_rng_, key_ = jax.random.split(rng_, 3)
  446.  
  447.             energy, q_grad = model_fn(xt, t_proto * i, condition=condition)
  448.             if self.grad_clip is not None:
  449.                 q_grad = jnp.clip(q_grad, -self.grad_clip, self.grad_clip)
  450.  
  451.             drift = stepsizes[i] * (
  452.                 0.5 * q_grad +\
  453.                 jax.random.normal(key_, xt.shape) * self.noise_scale
  454.             )
  455.             if self.drift_clip is not None:
  456.                 drift = jnp.clip(drift, -self.drift_clip, self.drift_clip)
  457.             xt_1 = xt + drift
  458.             if self.margin_clip is not None:
  459.                 xt_1 = jnp.clip(xt_1, -self.margin_clip, self.margin_clip)
  460.             return (rng_, xt_1), (xt, q_grad, energy)
  461.  
  462.         output, history = jax.lax.scan(fn, (rng, xT), jnp.arange(self.steps), unroll=True)
  463.         rng, action = output
  464.         return rng, action, history
  465.  
  466. def exponential_schedule(init, decay, steps):
  467.     return init * (decay ** jnp.arange(steps))
  468.  
  469. def polynomial_schedule(init, final, power, steps):
  470.     return (init - final) * (1 - jnp.arange(steps) / (steps - 1)) ** power + final
  471.  
  472.  
  473. if __name__ == "__main__":
  474.     import flax
  475.     import flax.linen as nn
  476.  
  477.     rng = jax.random.PRNGKey(0)
  478.     ld = IBCLangevinDynamics.create(
  479.         network=flax.linen.Dense(10),
  480.         rng=rng,
  481.         inputs=(jnp.ones((1, 10)),),
  482.         x_dim=10,
  483.         steps=20,
  484.         schedule="polynomial",
  485.         stepsize_init=1e-1,
  486.         stepsize_final=1e-5,
  487.         stepsize_decay=0.8,
  488.         stepsize_power=2.0,
  489.         grad_clip=1.0,
  490.         drift_clip=1.0,
  491.         margin_clip=1.0,
  492.     )
  493.     model_fn = lambda x, i, condition: (jnp.ones((*x.shape[:-1], 1)), jnp.ones(x.shape))
  494.     xT = jnp.ones((128, 10))
  495.     condition = jnp.ones((128, 5))
  496.     r1, r2 = ld.sample(
  497.         rng,
  498.         model_fn,
  499.         xT,
  500.         condition,
  501.         training=False,
  502.     )
Advertisement
Add Comment
Please, Sign In to add comment