from functools import partial import jax import jax.numpy as jnp from jax import Array # ============================================ # Noise Schedulers (how to compute z from x0 and noise) # ============================================ def add_noise(x0, noise, t, config, cond_seq_mask=None): """Flow-matching interpolation z = t*x0 + (1-t)*noise*scale, preserving cond tokens.""" t_expanded = t.reshape(-1, 1, 1) z = t_expanded * x0 + (1 - t_expanded) * noise * config.denoiser_noise_scale if cond_seq_mask is not None: z = cond_seq_mask * x0 + (1 - cond_seq_mask) * z return z # ============================================ # Time Schedulers (how to sample t) # ============================================ def sample_timesteps( rng, batch_size, P_mean=-0.8, P_std=0.8, time_schedule='logit_normal', ): """Sample timesteps using various time schedules. Args: rng: JAX random key batch_size: Number of samples P_mean: Mean for logit-normal distribution P_std: Std for logit-normal distribution time_schedule: 'logit_normal' or 'uniform' Returns: Sampled timesteps in [0, 1] """ if time_schedule == 'logit_normal': # Biased toward middle timesteps via sigmoid(N(P_mean, P_std)). z = jax.random.normal(rng, (batch_size,)) * P_std + P_mean return jax.nn.sigmoid(z) if time_schedule == 'uniform': return jax.random.uniform(rng, (batch_size,)) raise ValueError(f"Unknown time_schedule: {time_schedule}") def get_sampling_steps( rng, n_steps: int, time_schedule: str = "logit_normal", P_mean: float = -0.8, P_std: float = 0.8, ) -> Array: """Return a length-(n_steps+1) array of t values in [0, 1] for a sampling run. - "uniform": evenly-spaced linspace from 0 to 1 (deterministic). - "logit_normal": sorted logit-normal samples with 0 / 1 endpoints (random). """ if time_schedule == "uniform": return jnp.linspace(0.0, 1.0, n_steps + 1) if time_schedule == "logit_normal": steps = sample_timesteps( rng, batch_size=n_steps - 1, P_mean=P_mean, P_std=P_std, time_schedule=time_schedule, ) return jnp.concatenate([jnp.array([0.0]), jnp.sort(steps), jnp.array([1.0])]) raise ValueError(f"Unknown time_schedule: {time_schedule}") # ============================================ # CFG Scale Sampling (how to sample cfg scale) # ============================================ def sample_cfg_scale(rng, batch_size, cfg_min=0.0, cfg_max=3.0): """Sample CFG scale from log-uniform distribution in [cfg_min, cfg_max].""" u = jax.random.uniform(rng, (batch_size,)) a = jnp.float32(1.0 + cfg_min) b = jnp.float32(1.0 + cfg_max) return a * jnp.exp(u * jnp.log(b / a)) - 1.0 # ============================================ # Conditioning helpers (preserve clean tokens during sampling) # ============================================ def restore_cond(z_updated, cond_seq, cond_seq_mask): """Restore clean conditioning tokens in z after a denoising step.""" mask = cond_seq_mask target_ndim = max(z_updated.ndim, cond_seq.ndim) while mask.ndim < target_ndim: mask = mask[..., None] return jnp.where(mask > 0, cond_seq, z_updated) def restore_vx(v, x, cond_seq, cond_seq_mask): """Restore cond positions: x → clean cond_seq, v → 0 (cond tokens don't move).""" if cond_seq is not None: x = restore_cond(x, cond_seq, cond_seq_mask) v = restore_cond(v, jnp.zeros_like(cond_seq), cond_seq_mask) return v, x # ============================================ # Flow-matching forward passes (with optional self-cond / CFG) # ============================================ def net_out_to_v_x(net_out, z, t, t_eps=5e-2): """Convert x_pred network output to v and x. When the model returns a tuple (denoised_output, decoder_logits), decoder logits are discarded here (used separately in training). """ if isinstance(net_out, tuple): net_out = net_out[0] t_reshaped = t.reshape(-1, 1, 1) x = net_out v = (x - z) / jnp.maximum(1.0 - t_reshaped, t_eps) return v, x @partial(jax.jit, static_argnums=(0, 5, 6)) def _forward_sample_self_cond( model_apply_fn, model_params, z, t_batch, x_pred_prev, config, self_cond_cfg_scale, cond_seq, cond_seq_mask, ): """Forward pass with self-conditioning.""" t_eps = config.t_eps self_cond_prob = config.self_cond_prob _restore_vx = partial(restore_vx, cond_seq=cond_seq, cond_seq_mask=cond_seq_mask) if config.num_self_cond_cfg_tokens > 0: if x_pred_prev is None: x_pred_prev = restore_cond(jnp.zeros_like(z), cond_seq, cond_seq_mask) z_input_cond = jnp.concatenate([z, x_pred_prev], axis=-1) self_cond_scale_batch = jnp.full((z.shape[0],), self_cond_cfg_scale) net_out_cond = model_apply_fn( {"params": model_params}, z_input_cond, t_batch, deterministic=True, self_cond_cfg_scale=self_cond_scale_batch, ) v_cond, x_cond = net_out_to_v_x(net_out_cond, z, t_batch, t_eps) return _restore_vx(v_cond, x_cond) # No self-conditioning if self_cond_prob == 0: net_out = model_apply_fn( {"params": model_params}, z, t_batch, deterministic=True, ) v, x = net_out_to_v_x(net_out, z, t_batch, t_eps) return _restore_vx(v, x) # Combined unconditional and conditional forward pass if self_cond_cfg_scale != 1 or x_pred_prev is None: z_uncond = restore_cond(jnp.zeros_like(z), cond_seq, cond_seq_mask) z_input_uncond = jnp.concatenate([z, z_uncond], axis=-1) net_out_uncond = model_apply_fn( {"params": model_params}, z_input_uncond, t_batch, deterministic=True, ) v_uncond, x_uncond = net_out_to_v_x(net_out_uncond, z, t_batch, t_eps) v_uncond, x_uncond = _restore_vx(v_uncond, x_uncond) if self_cond_cfg_scale == 0.0 or x_pred_prev is None: return v_uncond, x_uncond z_input_cond = jnp.concatenate([z, x_pred_prev], axis=-1) net_out_cond = model_apply_fn( {"params": model_params}, z_input_cond, t_batch, deterministic=True, ) v_cond, x_cond = net_out_to_v_x(net_out_cond, z, t_batch, t_eps) v_cond, x_cond = _restore_vx(v_cond, x_cond) if self_cond_cfg_scale == 1: return v_cond, x_cond v_out = v_uncond + self_cond_cfg_scale * (v_cond - v_uncond) x_out = x_uncond + self_cond_cfg_scale * (x_cond - x_uncond) return _restore_vx(v_out, x_out) @partial(jax.jit, static_argnums=(0, 5, 6, 7)) def _forward_sample( model_apply_fn, model_params, z, t_batch, x_pred_prev, config, cfg_scale, self_cond_cfg_scale, cond_seq, cond_seq_mask, ): """Forward pass with optional self-conditioning and CFG.""" v_cond, x_cond = _forward_sample_self_cond( model_apply_fn, model_params, z, t_batch, x_pred_prev, config, self_cond_cfg_scale=self_cond_cfg_scale, cond_seq=cond_seq, cond_seq_mask=cond_seq_mask, ) if cfg_scale == 1.0: return v_cond, x_cond # Unconditional forward: zero out cond prefix, no self-cond state, no restore z_uncond = restore_cond(z, jnp.zeros_like(z), cond_seq_mask) x_pred_prev_uncond = ( None if x_pred_prev is None else restore_cond(x_pred_prev, jnp.zeros_like(x_pred_prev), cond_seq_mask) ) v_uncond, x_uncond = _forward_sample_self_cond( model_apply_fn, model_params, z_uncond, t_batch, x_pred_prev_uncond, config, self_cond_cfg_scale=self_cond_cfg_scale, cond_seq=jnp.zeros_like(cond_seq), cond_seq_mask=cond_seq_mask, ) v_out = v_uncond + cfg_scale * (v_cond - v_uncond) x_out = x_uncond + cfg_scale * (x_cond - x_uncond) return restore_vx(v_out, x_out, cond_seq, cond_seq_mask) @partial(jax.jit, static_argnums=(0, 6, 7, 8)) def _ode_step( model_apply_fn, model_params, z, t, t_next, x_pred_prev, config, cfg_scale, self_cond_cfg_scale, cond_seq, cond_seq_mask, ): """Single ODE (Euler) step for sampling.""" t_batch = jnp.full((z.shape[0],), t) v_pred, x_pred = _forward_sample( model_apply_fn=model_apply_fn, model_params=model_params, z=z, t_batch=t_batch, x_pred_prev=x_pred_prev, config=config, cfg_scale=cfg_scale, self_cond_cfg_scale=self_cond_cfg_scale, cond_seq=cond_seq, cond_seq_mask=cond_seq_mask, ) return z + (t_next - t) * v_pred, x_pred @partial(jax.jit, static_argnums=(0, 6, 7, 8)) def _sde_step( model_apply_fn, model_params, z, t, t_next, x_pred_prev, config, cfg_scale, self_cond_cfg_scale, cond_seq, cond_seq_mask, gamma, rng, ): """Per-step SDE-style sampler with hybrid (t-and-step) noise scaling. t_back = t * (1 - gamma * h), where h = t_next - t. alpha = 1 - gamma*h is the signal-preservation fraction, constant in t. gamma=0 degenerates to a plain ODE step. Uniform-N-step equivalence with old multiplicative gamma_old: gamma_hybrid = gamma_old * N. """ h = t_next - t alpha = jnp.clip(1.0 - gamma * h, 0.0, 1.0) t_back = alpha * t eps = jax.random.normal(rng, z.shape) * config.denoiser_noise_scale z_back = restore_cond(alpha * z + (1.0 - alpha) * eps, cond_seq, cond_seq_mask) t_batch = jnp.full((z.shape[0],), t_back) v_pred, x_pred = _forward_sample( model_apply_fn=model_apply_fn, model_params=model_params, z=z_back, t_batch=t_batch, x_pred_prev=x_pred_prev, config=config, cfg_scale=cfg_scale, self_cond_cfg_scale=self_cond_cfg_scale, cond_seq=cond_seq, cond_seq_mask=cond_seq_mask, ) return z_back + (t_next - t_back) * v_pred, x_pred