Download ddpm.py from jcandane/DIMAX: direct link, hf CLI and curl.
- Browser
- Download file 19.6 kB
-
https://huggingface.co/jcandane/DIMAX/resolve/main/ddpm.py
- Command line
-
hf download hf://jcandane/DIMAX/ddpm.py
-
curl -L -o ddpm.py https://huggingface.co/jcandane/DIMAX/resolve/main/ddpm.py
19.6 kB
| # src/dima/ddpm.py | |
| from __future__ import annotations | |
| from typing import Any, Optional | |
| import jax | |
| import jax.numpy as jnp | |
| from jax import random | |
| from flax import linen as nn | |
| from flax.training import train_state | |
| from flax import struct | |
| import optax | |
| # ------------------------------ | |
| # Schedules & embeddings | |
| # ------------------------------ | |
| def cosine_schedule(T: int, s: float = 0.008): | |
| """ | |
| Nichol & Dhariwal cosine schedule. | |
| Returns: | |
| alpha: (T,) | |
| beta: (T,) | |
| alpha_bar: (T,) | |
| """ | |
| steps = jnp.arange(T + 1, dtype=jnp.float32) | |
| f = jnp.cos(((steps / T + s) / (1.0 + s)) * jnp.pi / 2.0) ** 2 | |
| alpha_bar_all = f / f[0] | |
| alpha_bar = alpha_bar_all[1:] # (T,) | |
| alpha = alpha_bar / jnp.concatenate([jnp.array([1.0], dtype=jnp.float32), alpha_bar[:-1]]) | |
| beta = 1.0 - alpha | |
| return alpha, beta, alpha_bar | |
| def sinusoidal_embedding(t_idx: jnp.ndarray, dim: int) -> jnp.ndarray: | |
| """ | |
| t_idx: (B,1) int32 | |
| returns: (B,dim) | |
| """ | |
| if t_idx.ndim != 2 or t_idx.shape[1] != 1: | |
| raise ValueError("t_idx must have shape (B,1)") | |
| t = t_idx.astype(jnp.float32) | |
| half = dim // 2 | |
| denom = float(max(half - 1, 1)) | |
| freqs = jnp.exp(-jnp.log(10_000.0) * jnp.arange(half, dtype=jnp.float32) / denom) | |
| args = t * freqs | |
| emb = jnp.concatenate([jnp.sin(args), jnp.cos(args)], axis=-1) | |
| if dim % 2 == 1: | |
| emb = jnp.pad(emb, ((0, 0), (0, 1))) | |
| return emb | |
| # ------------------------------ | |
| # Model: epsilon predictor | |
| # ------------------------------ | |
| class EpsMLP(nn.Module): | |
| """Simple MLP epsilon-predictor for DDPM in R^D.""" | |
| hidden: int | |
| t_dim: int | |
| data_dim: int | |
| def __call__(self, x: jnp.ndarray, t_idx: jnp.ndarray) -> jnp.ndarray: | |
| # x: (B,D), t_idx: (B,1) | |
| t_emb = sinusoidal_embedding(t_idx, self.t_dim) | |
| t_h = nn.Dense(self.hidden)(t_emb) | |
| t_h = nn.gelu(t_h) | |
| h = nn.Dense(self.hidden)(x) | |
| h = nn.gelu(h + t_h) | |
| t_h2 = nn.Dense(self.hidden)(t_h) | |
| h = nn.Dense(self.hidden)(h) | |
| h = nn.gelu(h + t_h2) | |
| out = nn.Dense(self.data_dim)(h) | |
| return out | |
| # ------------------------------ | |
| # TrainState with EMA | |
| # ------------------------------ | |
| class TrainStateEMA(train_state.TrainState): | |
| """Flax TrainState extended with EMA params.""" | |
| ema_params: Any = struct.field(pytree_node=True) | |
| def apply_gradients(self, *, grads, ema_decay: float): | |
| updates, new_opt_state = self.tx.update(grads, self.opt_state, self.params) | |
| new_params = optax.apply_updates(self.params, updates) | |
| new_ema = optax.incremental_update(new_params, self.ema_params, step_size=1.0 - ema_decay) | |
| return self.replace( | |
| step=self.step + 1, | |
| params=new_params, | |
| opt_state=new_opt_state, | |
| ema_params=new_ema, | |
| ) | |
| # ------------------------------ | |
| # DDPM | |
| # ------------------------------ | |
| class DDPM: | |
| """ | |
| DDPM for D-dimensional latents. | |
| API expected by your DIMA wrapper: | |
| - DDPM(Z_train, ...): trains in __init__ (n_iter can be 0 to skip) | |
| - refine_latents(z0, t_start, key, add_noise) -> z_refined | |
| - __call__(...) delegates to refine_latents | |
| - sample(N) -> latent samples | |
| Added: | |
| - state_dict / save_local / load_local / from_state | |
| - upload_to_huggingface / download_from_huggingface | |
| """ | |
| def __init__( | |
| self, | |
| Z_iX: jnp.ndarray, | |
| *, | |
| T: int = 100, | |
| hidden_dim: int = 128, | |
| t_embed_dim: int = 64, | |
| learning_rate: float = 1e-3, | |
| n_iter: int = 20_000, | |
| ema_decay: float = 0.999, | |
| beta_max: float = 0.02, | |
| batch_size: Optional[int] = None, | |
| key: jax.Array = random.PRNGKey(0), | |
| verbose_every: int = 0, | |
| eps: float = 1e-5, | |
| ): | |
| Z_iX = jnp.asarray(Z_iX, dtype=jnp.float32) | |
| if Z_iX.ndim != 2: | |
| raise ValueError("Z_iX must be 2D (N,D).") | |
| self.D = int(Z_iX.shape[1]) | |
| self.T = int(T) | |
| self.key = key | |
| self.ema_decay = float(ema_decay) | |
| self.batch_size = batch_size | |
| self.verbose_every = int(verbose_every) | |
| self.eps = float(eps) | |
| self.beta_max = float(beta_max) | |
| # store scalars for serialization/rebuild | |
| self.learning_rate = float(learning_rate) | |
| self.hidden_dim = int(hidden_dim) | |
| self.t_embed_dim = int(t_embed_dim) | |
| # schedule (cosine, clipped by beta_max) | |
| alpha, beta, alpha_bar = cosine_schedule(self.T) | |
| beta = jnp.minimum(beta, self.beta_max) | |
| alpha = 1.0 - beta | |
| alpha_bar = jnp.cumprod(alpha) | |
| self.alpha_s = alpha.astype(jnp.float32) | |
| self.beta_s = beta.astype(jnp.float32) | |
| self.alpha_bar_s = alpha_bar.astype(jnp.float32) | |
| # model + optimizer + EMA | |
| self.model = EpsMLP(hidden=self.hidden_dim, t_dim=self.t_embed_dim, data_dim=self.D) | |
| params = self.model.init( | |
| self.key, | |
| jnp.zeros((1, self.D), dtype=jnp.float32), | |
| jnp.zeros((1, 1), dtype=jnp.int32), | |
| )["params"] | |
| tx = optax.adam(self.learning_rate) | |
| self.state = TrainStateEMA.create(apply_fn=self.model.apply, params=params, tx=tx, ema_params=params) | |
| # train | |
| if int(n_iter) > 0: | |
| self._train(Z_iX, int(n_iter)) | |
| # ---------- training ---------- | |
| def _loss(params, apply_fn, x_t, t_idx, eps_true): | |
| eps_pred = apply_fn({"params": params}, x_t, t_idx) | |
| return jnp.mean((eps_pred - eps_true) ** 2) | |
| def _train_step( | |
| state: TrainStateEMA, | |
| x0_batch: jnp.ndarray, | |
| key: jax.Array, | |
| alpha_bar_s: jnp.ndarray, | |
| ema_decay: float, | |
| eps: float, | |
| ): | |
| B = x0_batch.shape[0] | |
| key, k_eps, k_t = random.split(key, 3); | |
| eps_noise = random.normal(k_eps, shape=x0_batch.shape) # (B,D) | |
| t_idx = random.randint(k_t, shape=(B, 1), minval=0, maxval=alpha_bar_s.shape[0]) | |
| a_bar_t = jnp.take(alpha_bar_s, t_idx.squeeze(-1))[:, None] | |
| a_bar_t = jnp.clip(a_bar_t, eps, 1.0) | |
| x_t = jnp.sqrt(a_bar_t) * x0_batch + jnp.sqrt(1.0 - a_bar_t) * eps_noise | |
| def loss_fn(p): | |
| return DDPM._loss(p, state.apply_fn, x_t, t_idx, eps_noise) | |
| loss, grads = jax.value_and_grad(loss_fn)(state.params) | |
| new_state = state.apply_gradients(grads=grads, ema_decay=ema_decay) | |
| return new_state, loss, key | |
| def _train(self, Z_iX: jnp.ndarray, n_iter: int): | |
| N = int(Z_iX.shape[0]) | |
| bs = N if (self.batch_size is None) else min(int(self.batch_size), N) | |
| for it in range(n_iter): | |
| if bs >= N: | |
| batch = Z_iX | |
| else: | |
| self.key, k_perm = random.split(self.key) | |
| idx = random.permutation(k_perm, N)[:bs] | |
| batch = Z_iX[idx] | |
| self.state, loss, self.key = self._train_step( | |
| self.state, | |
| batch, | |
| self.key, | |
| self.alpha_bar_s, | |
| self.ema_decay, | |
| self.eps, | |
| ) | |
| if self.verbose_every and (it % self.verbose_every == 0 or it == n_iter - 1): | |
| print(f"iter {it:6d} loss {float(loss):.6f}", end="\r") | |
| if self.verbose_every: | |
| print("\ntraining complete.") | |
| # ---------- diffusion utilities ---------- | |
| def _posterior_variance(alpha_s, beta_s, alpha_bar_s, t): | |
| a_bar_t = alpha_bar_s[t] | |
| a_bar_prev = jnp.where(t > 0, alpha_bar_s[t - 1], jnp.array(1.0, dtype=alpha_bar_s.dtype)) | |
| return ((1.0 - a_bar_prev) / (1.0 - a_bar_t)) * beta_s[t] | |
| def _make_sampler_step(params_ema, apply_fn, alpha_s, beta_s, alpha_bar_s, eps: float): | |
| def step(carry, _): | |
| key, t, x = carry # x: (B,D) | |
| key, k = random.split(key) | |
| alpha_t = jnp.clip(alpha_s[t], eps, 1.0) | |
| a_bar_t = jnp.clip(alpha_bar_s[t], eps, 1.0) | |
| sqrt_alpha = jnp.sqrt(alpha_t) | |
| sqrt_one_minus_a_bar = jnp.sqrt(jnp.clip(1.0 - a_bar_t, eps, 1.0)) | |
| B = x.shape[0] | |
| t_batch = jnp.full((B, 1), t, dtype=jnp.int32) | |
| eps_pred = apply_fn({"params": params_ema}, x, t_batch) # (B,D) | |
| # predict x0 | |
| x0_hat = (x - sqrt_one_minus_a_bar * eps_pred) / jnp.sqrt(a_bar_t) | |
| a_bar_prev = jnp.where(t > 0, alpha_bar_s[t - 1], jnp.array(1.0, dtype=alpha_bar_s.dtype)) | |
| denom = jnp.clip(1.0 - a_bar_t, eps, 1.0) | |
| coef1 = jnp.sqrt(jnp.clip(a_bar_prev, eps, 1.0)) * beta_s[t] / denom | |
| coef2 = sqrt_alpha * (1.0 - a_bar_prev) / denom | |
| mean = coef1 * x0_hat + coef2 * x | |
| beta_tilde = DDPM._posterior_variance(alpha_s, beta_s, alpha_bar_s, t) | |
| sigma = jnp.sqrt(jnp.clip(beta_tilde, 0.0, 1.0)) | |
| z = random.normal(k, x.shape) | |
| z = jnp.where(t == 0, 0.0, z); | |
| x_prev = mean + sigma * z | |
| return (key, t - 1, x_prev), x_prev | |
| return step | |
| # ---------- API ---------- | |
| def refine_latents( | |
| self, | |
| z0: jnp.ndarray, | |
| t_start: int = 10, | |
| key: Optional[jax.Array] = None, | |
| add_noise: bool = True, | |
| ) -> jnp.ndarray: | |
| """ | |
| Refine latents by: | |
| (optional) forward-noise z0 to step t_start | |
| reverse-diffuse from t_start -> 0 using EMA params. | |
| """ | |
| z0 = jnp.asarray(z0, dtype=jnp.float32) | |
| if z0.ndim != 2 or z0.shape[1] != self.D: | |
| raise ValueError(f"z0 must have shape (B,{self.D}).") | |
| if not (0 <= int(t_start) < self.T): | |
| raise ValueError(f"t_start must be in [0, {self.T-1}]") | |
| t_start = int(t_start) | |
| if key is None: | |
| self.key, key = random.split(self.key) | |
| else: | |
| # advance internal RNG too | |
| self.key, _ = random.split(key) | |
| # forward diffuse to t_start | |
| key, k_eps = random.split(key) | |
| eps_noise = random.normal(k_eps, z0.shape) | |
| a_bar_t = jnp.clip(self.alpha_bar_s[t_start], self.eps, 1.0) | |
| if add_noise: | |
| z_t = jnp.sqrt(a_bar_t) * z0 + jnp.sqrt(1.0 - a_bar_t) * eps_noise | |
| else: | |
| z_t = z0 | |
| step = self._make_sampler_step( | |
| self.state.ema_params, | |
| self.state.apply_fn, | |
| self.alpha_s, | |
| self.beta_s, | |
| self.alpha_bar_s, | |
| self.eps, | |
| ) | |
| (final_key, _, _), trace = jax.lax.scan( | |
| step, | |
| (key, t_start, z_t), | |
| xs=None, | |
| length=t_start + 1, | |
| ) | |
| self.key = final_key | |
| return trace[-1] | |
| def __call__( | |
| self, | |
| z0: jnp.ndarray, | |
| t_start: int = 10, | |
| key: Optional[jax.Array] = None, | |
| add_noise: bool = True, | |
| ) -> jnp.ndarray: | |
| return self.refine_latents(z0, t_start=t_start, key=key, add_noise=add_noise) | |
| def reverse_from_T(self, x_T: jnp.ndarray) -> jnp.ndarray: | |
| x_T = jnp.asarray(x_T, dtype=jnp.float32) | |
| if x_T.ndim != 2 or x_T.shape[1] != self.D: | |
| raise ValueError(f"x_T must have shape (B,{self.D}).") | |
| step = self._make_sampler_step( | |
| self.state.ema_params, | |
| self.state.apply_fn, | |
| self.alpha_s, | |
| self.beta_s, | |
| self.alpha_bar_s, | |
| self.eps, | |
| ) | |
| self.key, k0 = random.split(self.key) | |
| (_, _, _), trace = jax.lax.scan( | |
| step, | |
| (k0, self.T - 1, x_T), | |
| xs=None, | |
| length=self.T, | |
| ) | |
| return trace[-1] | |
| def sample(self, N: int = 10_000) -> jnp.ndarray: | |
| self.key, k = random.split(self.key) | |
| noise = random.normal(k, (int(N), self.D)) | |
| return self.reverse_from_T(noise) | |
| # ------------------------- | |
| # Serialization helpers | |
| # ------------------------- | |
| def state_dict(self): | |
| """ | |
| Msgpack-safe checkpoint dict. | |
| Uses flax.serialization.to_state_dict to avoid tuple/namedtuple issues in opt_state. | |
| """ | |
| from flax import serialization as flax_ser # type: ignore | |
| import numpy as np | |
| return dict( | |
| # hyperparameters needed to rebuild the object | |
| T=int(self.T), | |
| D=int(self.D), | |
| hidden_dim=int(self.hidden_dim), | |
| t_embed_dim=int(self.t_embed_dim), | |
| learning_rate=float(self.learning_rate), | |
| ema_decay=float(self.ema_decay), | |
| beta_max=float(self.beta_max), | |
| batch_size=None if (self.batch_size is None) else int(self.batch_size), | |
| eps=float(self.eps), | |
| # RNG key (store as numpy for portability) | |
| key=np.asarray(self.key), | |
| # TrainStateEMA (params, ema_params, opt_state, step) in a serializable form | |
| train_state=flax_ser.to_state_dict(self.state), | |
| ) | |
| def save_local(self, weights_file: str = "ddpm.msgpack", config_file: Optional[str] = "ddpm_config.json") -> None: | |
| from flax import serialization as flax_ser # type: ignore | |
| import json as _json | |
| ckpt = self.state_dict() | |
| blob = flax_ser.msgpack_serialize(ckpt) | |
| with open(weights_file, "wb") as f: | |
| f.write(blob) | |
| if config_file is not None: | |
| cfg = dict( | |
| T=int(ckpt["T"]), | |
| D=int(ckpt["D"]), | |
| hidden_dim=int(ckpt["hidden_dim"]), | |
| t_embed_dim=int(ckpt["t_embed_dim"]), | |
| learning_rate=float(ckpt["learning_rate"]), | |
| ema_decay=float(ckpt["ema_decay"]), | |
| beta_max=float(ckpt["beta_max"]), | |
| batch_size=ckpt["batch_size"], | |
| eps=float(ckpt["eps"]), | |
| ) | |
| with open(config_file, "w") as f: | |
| _json.dump(cfg, f, indent=2) | |
| def from_state(cls, ckpt, *, key: Optional[jax.Array] = None) -> "DDPM": | |
| """ | |
| Rehydrate a DDPM instance without retraining, restoring full TrainStateEMA | |
| (including opt_state) via flax.serialization.from_state_dict. | |
| """ | |
| from flax import serialization as flax_ser # type: ignore | |
| import numpy as np | |
| T = int(ckpt["T"]) | |
| D = int(ckpt["D"]) | |
| hidden_dim = int(ckpt["hidden_dim"]) | |
| t_embed_dim = int(ckpt["t_embed_dim"]) | |
| learning_rate = float(ckpt["learning_rate"]) | |
| ema_decay = float(ckpt["ema_decay"]) | |
| beta_max = float(ckpt["beta_max"]) | |
| batch_size = ckpt.get("batch_size", None) | |
| eps = float(ckpt.get("eps", 1e-5)) | |
| if key is None: | |
| # prefer explicit passed key; else use checkpoint key; else default | |
| if "key" in ckpt: | |
| key = jnp.asarray(np.asarray(ckpt["key"])) | |
| else: | |
| key = random.PRNGKey(0) | |
| # Build a skeleton instance (no training) | |
| dummy = jnp.zeros((1, D), dtype=jnp.float32) | |
| obj = cls( | |
| dummy, | |
| T=T, | |
| hidden_dim=hidden_dim, | |
| t_embed_dim=t_embed_dim, | |
| learning_rate=learning_rate, | |
| n_iter=0, | |
| ema_decay=ema_decay, | |
| beta_max=beta_max, | |
| batch_size=batch_size, | |
| key=key, | |
| verbose_every=0, | |
| eps=eps, | |
| ) | |
| # Restore TrainStateEMA (including opt_state tuples) safely | |
| obj.state = flax_ser.from_state_dict(obj.state, ckpt["train_state"]) | |
| # Restore RNG key if present (keeps reproducibility) | |
| if "key" in ckpt: | |
| obj.key = jnp.asarray(np.asarray(ckpt["key"])) | |
| return obj | |
| def load_local(cls, weights_file: str = "ddpm.msgpack", *, key: Optional[jax.Array] = None) -> "DDPM": | |
| from flax import serialization as flax_ser # type: ignore | |
| with open(weights_file, "rb") as f: | |
| ckpt = flax_ser.msgpack_restore(f.read()) | |
| return cls.from_state(ckpt, key=key) | |
| def upload_to_huggingface( | |
| self, | |
| repo_id: str, | |
| *, | |
| weights_file: str = "ddpm.msgpack", | |
| config_file: str = "ddpm_config.json", | |
| token: Optional[str] = None, | |
| repo_type: str = "model", | |
| revision: Optional[str] = None, | |
| ) -> None: | |
| """ | |
| Upload DDPM weights to the Hugging Face Hub. | |
| Saves locally first (msgpack + JSON), then creates repo (if needed) and uploads files. | |
| """ | |
| try: | |
| from huggingface_hub import HfApi, HfFolder, upload_file # type: ignore | |
| except Exception as e: | |
| raise RuntimeError("huggingface_hub not installed. Install it (or `pip install dima[hf]`).") from e | |
| self.save_local(weights_file=weights_file, config_file=config_file) | |
| if token is None: | |
| token = HfFolder.get_token() | |
| if token is None: | |
| raise RuntimeError("No HF token found. Run `huggingface-cli login`, or pass `token=...`.") | |
| import os as _os | |
| api = HfApi() | |
| api.create_repo(repo_id=repo_id, repo_type=repo_type, exist_ok=True, token=token) | |
| wf = _os.path.basename(weights_file) | |
| cf = _os.path.basename(config_file) | |
| upload_file( | |
| path_or_fileobj=weights_file, | |
| path_in_repo=wf, | |
| repo_id=repo_id, | |
| token=token, | |
| repo_type=repo_type, | |
| revision=revision, | |
| ) | |
| upload_file( | |
| path_or_fileobj=config_file, | |
| path_in_repo=cf, | |
| repo_id=repo_id, | |
| token=token, | |
| repo_type=repo_type, | |
| revision=revision, | |
| ) | |
| def download_from_huggingface( | |
| cls, | |
| repo_id: str, | |
| *, | |
| weights_file: str = "ddpm.msgpack", | |
| token: Optional[str] = None, | |
| repo_type: str = "model", | |
| revision: Optional[str] = None, | |
| key: Optional[jax.Array] = None, | |
| ) -> "DDPM": | |
| """ | |
| Download DDPM weights from the Hugging Face Hub and return a rehydrated DDPM instance. | |
| """ | |
| try: | |
| from huggingface_hub import hf_hub_download # type: ignore | |
| except Exception as e: | |
| raise RuntimeError("huggingface_hub not installed. Install it (or `pip install dima[hf]`).") from e | |
| path = hf_hub_download( | |
| repo_id=repo_id, | |
| filename=weights_file, | |
| token=token, | |
| repo_type=repo_type, | |
| revision=revision, | |
| ) | |
| return cls.load_local(path, key=key) | |
| # Optional backward-compatible alias (if you ever used the other name) | |
| download_to_huggingface = download_from_huggingface | |
| __all__ = ["DDPM", "EpsMLP", "cosine_schedule", "sinusoidal_embedding"] |