DIMAX / ddpmx.py
jcandane's picture
Upload 9 files
204f4a4 verified
Raw
History Blame Contribute Delete
25.7 kB
# how to use
# ddpm = DDPM(Z_train, T=100, n_iter=20_000, key=random.PRNGKey(0))
#
# Fast unconditional samples (try 15–30 first)
# z = ddpm.sample_dpmpp(N=4096, num_steps=20)
#
# Fast “refine latents” starting from an intermediate noise level
#z_ref = ddpm.refine_latents_dpmpp(z0, t_start=30, num_steps=20, add_noise=True)
# src/dima/ddpmx.py
from __future__ import annotations
from typing import Any, Optional, Dict
import os
import json
import numpy as np
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, serialization as flax_ser
import optax
# ---------------------------------------------------------------------
# Helpers (as in ddpmx.py)
# ---------------------------------------------------------------------
def _sigma_to_alpha_sigma_t(sigma: jnp.ndarray) -> tuple[jnp.ndarray, jnp.ndarray]:
"""
EDM-style sigma parameterization:
alpha_t = 1 / sqrt(1 + sigma^2)
sigma_t = sigma * alpha_t
so that x = alpha_t * x0 + sigma_t * eps
"""
alpha_t = 1.0 / jnp.sqrt(1.0 + sigma**2)
sigma_t = sigma * alpha_t
return alpha_t, sigma_t
def _make_lu_sigma_schedule(sigma_start: float, sigma_end: float, num_steps: int) -> np.ndarray:
"""
"Lu" schedule uniform in lambda = -log(sigma).
"""
sigma_start = float(max(sigma_start, 1e-12))
sigma_end = float(max(sigma_end, 1e-12))
lam_start = -np.log(sigma_start)
lam_end = -np.log(sigma_end)
lambdas = np.linspace(lam_start, lam_end, int(num_steps), dtype=np.float32)
sigmas = np.exp(-lambdas).astype(np.float32)
return sigmas
def cosine_schedule(T: int, s: float = 0.008):
"""
Nichol & Dhariwal cosine schedule.
Returns alpha, beta, alpha_bar with shape (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 or float32
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
class EpsMLP(nn.Module):
"""Simple MLP epsilon-predictor for DDPM in R^D."""
hidden: int
t_dim: int
data_dim: int
@nn.compact
def __call__(self, x: jnp.ndarray, t_idx: jnp.ndarray) -> jnp.ndarray:
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
@struct.dataclass
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,
)
# ---------------------------------------------------------------------
# DDPMX with HF upload/download integrated
# ---------------------------------------------------------------------
class DDPM:
"""
DDPM (+ fast DPM-Solver++(2M) sampler utilities) for D-dimensional latents.
Added persistence utilities:
- save_local / load_local
- upload_to_huggingface / download_from_huggingface
Serialization is done via flax.serialization.to_state_dict / from_state_dict
to avoid msgpack failures with non-serializable Python objects (e.g., tuples).
"""
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)
# store config for checkpointing
self.hidden_dim = int(hidden_dim)
self.t_embed_dim = int(t_embed_dim)
self.learning_rate = float(learning_rate)
self.ema_decay = float(ema_decay)
self.beta_max = float(beta_max)
self.batch_size = batch_size
self.verbose_every = int(verbose_every)
self.eps = float(eps)
self.key = key
# VP schedule (cosine + clip)
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)
# Precompute EDM-style "sigma_in" for DPM++ schedule interpolation
# sigma_in = sqrt((1 - a_bar) / a_bar)
self.sigma_in_train = jnp.sqrt(
jnp.clip(
(1.0 - self.alpha_bar_s) / jnp.clip(self.alpha_bar_s, self.eps, 1.0),
self.eps,
1e12,
)
).astype(jnp.float32)
# model + optimizer + EMA state
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,
)
if int(n_iter) > 0:
self._train(Z_iX, int(n_iter))
# -------------------------
# Training
# -------------------------
@staticmethod
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)
@staticmethod
@jax.jit
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)
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.")
# -------------------------
# Standard DDPM refine/sample
# -------------------------
@staticmethod
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]
@staticmethod
def _make_sampler_step(params_ema, apply_fn, alpha_s, beta_s, alpha_bar_s, eps: float):
@jax.jit
def step(carry, _):
key, t, x = carry
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 = 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)
x0_hat = (x - sqrt_one_minus * 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
def refine_latents(
self,
z0: jnp.ndarray,
t_start: int = 10,
key: Optional[jax.Array] = None,
add_noise: bool = True,
) -> jnp.ndarray:
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:
self.key, _ = random.split(key)
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)).astype(jnp.float32)
return self.reverse_from_T(noise)
# -------------------------
# DPM-Solver++(2M) schedule + sampler
# -------------------------
def _make_dpmpp_schedule(self, *, num_steps: int, t_start: int) -> tuple[jnp.ndarray, jnp.ndarray]:
"""
Returns:
sigmas_in: (K+1,) float32 decreasing, last one is 0
t_cont: (K,) float32 continuous "time" indices for model calls
"""
t_start = int(t_start)
if not (0 <= t_start < self.T):
raise ValueError(f"t_start must be in [0, {self.T-1}]")
if int(num_steps) < 1:
raise ValueError("num_steps must be >= 1")
sigma_start = float(self.sigma_in_train[t_start])
sigma_end = float(self.sigma_in_train[0])
sigmas_k = _make_lu_sigma_schedule(sigma_start, sigma_end, int(num_steps))
sigmas = np.concatenate([sigmas_k, np.array([0.0], np.float32)], axis=0)
sigma_train = np.array(self.sigma_in_train).astype(np.float32) # (T,)
log_sig_train = np.log(np.maximum(sigma_train, 1e-12))
t_train = np.arange(self.T, dtype=np.float32)
log_sig = np.log(np.maximum(sigmas[:-1], 1e-12))
t_cont = np.interp(log_sig, log_sig_train, t_train).astype(np.float32)
return jnp.array(sigmas, dtype=jnp.float32), jnp.array(t_cont, dtype=jnp.float32)
@staticmethod
@jax.jit
def _dpmpp_2m_midpoint_sample(
params_ema: Any,
apply_fn: Any,
x_start: jnp.ndarray, # (B,D)
sigmas_in: jnp.ndarray, # (K+1,)
t_cont: jnp.ndarray, # (K,)
eps: float,
) -> jnp.ndarray:
"""
DPM-Solver++ (2M, midpoint) sampler.
"""
sigma_s = sigmas_in[:-1] # (K,)
sigma_t = sigmas_in[1:] # (K,)
alpha_s, sigma_s_t = _sigma_to_alpha_sigma_t(sigma_s)
alpha_t, sigma_t_t = _sigma_to_alpha_sigma_t(sigma_t)
lambda_s = jnp.log(alpha_s) - jnp.log(sigma_s_t)
lambda_t = jnp.log(alpha_t) - jnp.log(sigma_t_t)
K = t_cont.shape[0]
is_first = jnp.arange(K) == 0
is_last = jnp.arange(K) == (K - 1)
def step(carry, inp):
x, m_prev, lam_prev = carry
(a_s, s_s, a_t, s_t, lam_s_i, lam_t_i, t_i, first_i, last_i) = inp
B = x.shape[0]
t_batch = jnp.full((B, 1), t_i, dtype=jnp.float32)
eps_pred = apply_fn({"params": params_ema}, x, t_batch)
a_s_b = jnp.clip(a_s, eps, 1.0)
x0 = (x - s_s * eps_pred) / a_s_b
h = lam_t_i - lam_s_i
exp_neg_h = jnp.exp(-h)
# 1st-order
x_first = (s_t / s_s) * x - (a_t * (exp_neg_h - 1.0)) * x0
def do_second(_):
h0 = lam_s_i - lam_prev
r0 = h0 / jnp.clip(h, 1e-12)
D1 = (x0 - m_prev) / jnp.clip(r0, 1e-12)
x_second = (s_t / s_s) * x - (a_t * (exp_neg_h - 1.0)) * (x0 + 0.5 * D1)
return x_second
x_next = jax.lax.cond(first_i | last_i, lambda _: x_first, do_second, operand=None)
return (x_next, x0, lam_s_i), x_next
xs = (
alpha_s, sigma_s_t,
alpha_t, sigma_t_t,
lambda_s, lambda_t,
t_cont, is_first, is_last
)
x0_init = jnp.zeros_like(x_start)
lam_init = jnp.array(0.0, dtype=jnp.float32)
(x_final, _, _), _ = jax.lax.scan(step, (x_start, x0_init, lam_init), xs)
return x_final
def refine_latents_dpmpp(
self,
z0: jnp.ndarray,
*,
t_start: int = 10,
num_steps: int = 20,
key: Optional[jax.Array] = None,
add_noise: bool = True,
) -> jnp.ndarray:
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}]")
if key is None:
self.key, key = random.split(self.key)
else:
self.key, _ = random.split(key)
# forward-noise to t_start
key, k_eps = random.split(key)
eps_noise = random.normal(k_eps, z0.shape)
a_bar = jnp.clip(self.alpha_bar_s[int(t_start)], self.eps, 1.0)
if add_noise:
x_start = jnp.sqrt(a_bar) * z0 + jnp.sqrt(1.0 - a_bar) * eps_noise
else:
x_start = z0
sigmas_in, t_cont = self._make_dpmpp_schedule(num_steps=int(num_steps), t_start=int(t_start))
x_final = self._dpmpp_2m_midpoint_sample(
self.state.ema_params,
self.state.apply_fn,
x_start,
sigmas_in,
t_cont,
self.eps,
)
return x_final
def reverse_from_T_dpmpp(self, x_T: jnp.ndarray, *, num_steps: int = 20) -> 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}).")
sigmas_in, t_cont = self._make_dpmpp_schedule(num_steps=int(num_steps), t_start=self.T - 1)
return self._dpmpp_2m_midpoint_sample(
self.state.ema_params,
self.state.apply_fn,
x_T,
sigmas_in,
t_cont,
self.eps,
)
def sample_dpmpp(self, N: int = 10_000, *, num_steps: int = 20) -> jnp.ndarray:
self.key, k = random.split(self.key)
x_T = random.normal(k, (int(N), self.D)).astype(jnp.float32)
return self.reverse_from_T_dpmpp(x_T, num_steps=int(num_steps))
# -----------------------------------------------------------------
# Persistence: state_dict / save_local / load_local
# -----------------------------------------------------------------
def _config_dict(self) -> Dict[str, Any]:
return {
"class_name": "DDPMX",
"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),
"key": np.array(self.key).tolist(),
}
def save_local(self, weights_file: str = "ddpmx_weights.msgpack", config_file: str = "ddpmx_config.json") -> None:
"""
Saves:
- config_file: JSON with hyperparams + PRNG key
- weights_file: msgpack with flax state_dict of TrainStateEMA
"""
cfg = self._config_dict()
with open(config_file, "w", encoding="utf-8") as f:
json.dump(cfg, f, indent=2, ensure_ascii=False)
# Robust serialization (avoid msgpack tuple errors)
state_sd = flax_ser.to_state_dict(self.state)
blob = flax_ser.msgpack_serialize(state_sd)
with open(weights_file, "wb") as f:
f.write(blob)
@classmethod
def load_local(
cls,
weights_file: str,
config_file: str,
*,
Z_iX: Optional[jnp.ndarray] = None,
) -> "DDPM":
"""
Reconstructs a DDPMX instance from local files.
Z_iX is only used to provide shape (N,D) for initialization; training is skipped.
If Z_iX is None, a dummy array of shape (1,D) is created.
"""
with open(config_file, "r", encoding="utf-8") as f:
cfg = json.load(f)
D = int(cfg["D"])
if Z_iX is None:
Z_iX = jnp.zeros((1, D), dtype=jnp.float32)
# Build a fresh instance with the same architecture, skip training
obj = cls(
Z_iX,
T=int(cfg["T"]),
hidden_dim=int(cfg["hidden_dim"]),
t_embed_dim=int(cfg["t_embed_dim"]),
learning_rate=float(cfg["learning_rate"]),
n_iter=0,
ema_decay=float(cfg["ema_decay"]),
beta_max=float(cfg["beta_max"]),
batch_size=cfg["batch_size"],
key=random.PRNGKey(0),
verbose_every=0,
eps=float(cfg["eps"]),
)
with open(weights_file, "rb") as f:
state_sd = flax_ser.msgpack_restore(f.read())
obj.state = flax_ser.from_state_dict(obj.state, state_sd)
key_list = cfg.get("key", None)
if key_list is not None:
obj.key = jnp.array(key_list, dtype=jnp.uint32)
return obj
# -----------------------------------------------------------------
# Hugging Face Hub: upload / download
# -----------------------------------------------------------------
def upload_to_huggingface(
self,
repo_id: str,
*,
token: Optional[str] = None,
weights_file: str = "ddpmx_weights.msgpack",
config_file: str = "ddpmx_config.json",
repo_type: str = "model",
revision: Optional[str] = None,
) -> Dict[str, str]:
"""
Saves locally and uploads (weights_file, config_file) to Hugging Face Hub.
"""
try:
from huggingface_hub import create_repo, upload_file
except Exception as e:
raise RuntimeError(
"huggingface_hub not installed. Install it (e.g., `pip install huggingface_hub`)."
) from e
self.save_local(weights_file=weights_file, config_file=config_file)
create_repo(repo_id, token=token, repo_type=repo_type, exist_ok=True)
w_name = os.path.basename(weights_file)
c_name = os.path.basename(config_file)
upload_file(
path_or_fileobj=weights_file,
path_in_repo=w_name,
repo_id=repo_id,
repo_type=repo_type,
token=token,
revision=revision,
)
upload_file(
path_or_fileobj=config_file,
path_in_repo=c_name,
repo_id=repo_id,
repo_type=repo_type,
token=token,
revision=revision,
)
return {"repo_id": repo_id, "weights": w_name, "config": c_name}
@classmethod
def download_from_huggingface(
cls,
repo_id: str,
*,
token: Optional[str] = None,
weights_file: str = "ddpmx_weights.msgpack",
config_file: str = "ddpmx_config.json",
repo_type: str = "model",
revision: Optional[str] = None,
cache_dir: Optional[str] = None,
Z_iX: Optional[jnp.ndarray] = None,
) -> "DDPM":
"""
Downloads (weights_file, config_file) from Hugging Face Hub and reconstructs the class.
"""
try:
from huggingface_hub import hf_hub_download
except Exception as e:
raise RuntimeError(
"huggingface_hub not installed. Install it (e.g., `pip install huggingface_hub`)."
) from e
w_name = os.path.basename(weights_file)
c_name = os.path.basename(config_file)
w_path = hf_hub_download(
repo_id=repo_id,
filename=w_name,
repo_type=repo_type,
token=token,
revision=revision,
cache_dir=cache_dir,
)
c_path = hf_hub_download(
repo_id=repo_id,
filename=c_name,
repo_type=repo_type,
token=token,
revision=revision,
cache_dir=cache_dir,
)
return cls.load_local(w_path, c_path, Z_iX=Z_iX)
__all__ = ["DDPM", "EpsMLP", "cosine_schedule", "sinusoidal_embedding"]