| from __future__ import annotations |
|
|
| import torch |
|
|
| from .model import TextToLatentRFDiT |
|
|
|
|
| def _make_rng(seed: int, device: torch.device) -> tuple[torch.Generator, torch.device]: |
| |
| try: |
| return torch.Generator(device=device).manual_seed(seed), device |
| except RuntimeError: |
| return torch.Generator(device="cpu").manual_seed(seed), torch.device("cpu") |
|
|
|
|
| def sample_logit_normal_t( |
| batch_size: int, |
| device: torch.device, |
| mean: float = 0.0, |
| std: float = 1.0, |
| t_min: float = 1e-3, |
| t_max: float = 0.999, |
| ) -> torch.Tensor: |
| z = torch.randn(batch_size, device=device) * std + mean |
| t = torch.sigmoid(z) |
| return t.clamp(min=t_min, max=t_max) |
|
|
|
|
| def sample_stratified_logit_normal_t( |
| batch_size: int, |
| device: torch.device, |
| mean: float = 0.0, |
| std: float = 1.0, |
| t_min: float = 1e-3, |
| t_max: float = 0.999, |
| ) -> torch.Tensor: |
| """ |
| Stratified sampling for logit-normal timesteps. |
| |
| u ~ stratified U(0, 1), z = mean + std * Phi^{-1}(u), t = sigmoid(z) |
| """ |
| if batch_size <= 0: |
| return torch.empty((0,), device=device) |
| u = ( |
| torch.arange(batch_size, device=device, dtype=torch.float32) |
| + torch.rand(batch_size, device=device) |
| ) / float(batch_size) |
| u = u.clamp(1e-6, 1.0 - 1e-6) |
| |
| z = torch.erfinv(2.0 * u - 1.0) * (2.0**0.5) |
| z = z * std + mean |
| t = torch.sigmoid(z) |
| |
| t = t[torch.randperm(batch_size, device=device)] |
| return t.clamp(min=t_min, max=t_max) |
|
|
|
|
| def rf_interpolate(x0: torch.Tensor, noise: torch.Tensor, t: torch.Tensor) -> torch.Tensor: |
| |
| return (1.0 - t[:, None, None]) * x0 + t[:, None, None] * noise |
|
|
|
|
| def rf_velocity_target(x0: torch.Tensor, noise: torch.Tensor) -> torch.Tensor: |
| |
| return noise - x0 |
|
|
|
|
| def rf_predict_x0(x_t: torch.Tensor, v_pred: torch.Tensor, t: torch.Tensor) -> torch.Tensor: |
| |
| return x_t - t[:, None, None] * v_pred |
|
|
|
|
| def temporal_score_rescale( |
| v_pred: torch.Tensor, |
| x_t: torch.Tensor, |
| t: float | torch.Tensor, |
| rescale_k: float, |
| rescale_sigma: float, |
| ) -> torch.Tensor: |
| """ |
| Temporal score rescaling from https://arxiv.org/pdf/2510.01184. |
| """ |
| t_value = float(t.item()) if isinstance(t, torch.Tensor) else float(t) |
| if t_value >= 1.0: |
| return v_pred |
| one_minus_t = 1.0 - t_value |
| snr = (one_minus_t * one_minus_t) / (t_value * t_value) |
| sigma_sq = float(rescale_sigma) * float(rescale_sigma) |
| ratio = (snr * sigma_sq + 1.0) / (snr * sigma_sq / float(rescale_k) + 1.0) |
| return (ratio * (one_minus_t * v_pred + x_t) - x_t) / one_minus_t |
|
|
|
|
| def scale_speaker_kv_cache( |
| context_kv_cache: list[tuple[torch.Tensor, ...]], |
| scale: float, |
| max_layers: int | None = None, |
| ) -> None: |
| """ |
| In-place scaling of speaker K/V tensors in precomputed context cache. |
| """ |
| if max_layers is None: |
| n_layers = len(context_kv_cache) |
| else: |
| n_layers = max(0, min(int(max_layers), len(context_kv_cache))) |
| for i in range(n_layers): |
| layer_kv = context_kv_cache[i] |
| if len(layer_kv) < 4: |
| raise ValueError( |
| f"Expected at least 4 tensors in context KV cache entry, got {len(layer_kv)}" |
| ) |
| k_speaker = layer_kv[2] |
| v_speaker = layer_kv[3] |
| k_speaker.mul_(scale) |
| v_speaker.mul_(scale) |
|
|
|
|
| @torch.inference_mode() |
| def sample_euler_rf_cfg( |
| model: TextToLatentRFDiT, |
| text_input_ids: torch.Tensor, |
| text_mask: torch.Tensor, |
| ref_latent: torch.Tensor | None, |
| ref_mask: torch.Tensor | None, |
| sequence_length: int, |
| caption_input_ids: torch.Tensor | None = None, |
| caption_mask: torch.Tensor | None = None, |
| num_steps: int = 40, |
| cfg_scale_text: float = 3.0, |
| cfg_scale_caption: float = 3.0, |
| cfg_scale_speaker: float = 5.0, |
| cfg_guidance_mode: str = "independent", |
| cfg_min_t: float = 0.5, |
| cfg_max_t: float = 1.0, |
| seed: int = 0, |
| cfg_scale: float | None = None, |
| truncation_factor: float | None = None, |
| rescale_k: float | None = None, |
| rescale_sigma: float | None = None, |
| use_context_kv_cache: bool = True, |
| speaker_kv_scale: float | None = None, |
| speaker_kv_max_layers: int | None = None, |
| speaker_kv_min_t: float | None = None, |
| ) -> torch.Tensor: |
| """ |
| Euler sampling over RF ODE with text/reference/caption conditioning CFG. |
| |
| Returns: |
| latent sequence in patched space, shape (B, sequence_length, patched_latent_dim) |
| """ |
| device = model.device |
| dtype = model.dtype |
| batch_size = text_input_ids.shape[0] |
| latent_dim = model.cfg.patched_latent_dim |
|
|
| rng, rng_device = _make_rng(seed=seed, device=device) |
| x_t = torch.randn( |
| (batch_size, sequence_length, latent_dim), device=rng_device, dtype=dtype, generator=rng |
| ) |
| if rng_device != device: |
| x_t = x_t.to(device=device) |
| if truncation_factor is not None: |
| x_t = x_t * float(truncation_factor) |
|
|
| if cfg_scale is not None: |
| |
| cfg_scale_text = float(cfg_scale) |
| cfg_scale_caption = float(cfg_scale) |
| cfg_scale_speaker = float(cfg_scale) |
| if not model.cfg.use_speaker_condition: |
| cfg_scale_speaker = 0.0 |
| speaker_kv_scale = None |
|
|
| cfg_guidance_mode = str(cfg_guidance_mode).strip().lower() |
| if cfg_guidance_mode not in {"independent", "joint", "alternating"}: |
| raise ValueError( |
| f"Unsupported cfg_guidance_mode={cfg_guidance_mode!r}. " |
| "Expected one of: independent, joint, alternating." |
| ) |
|
|
| init_scale = 0.999 |
| t_schedule = torch.linspace(1.0, 0.0, num_steps + 1, device=device) * init_scale |
| use_independent_cfg = cfg_guidance_mode == "independent" |
| use_joint_cfg = cfg_guidance_mode == "joint" |
| use_alternating_cfg = cfg_guidance_mode == "alternating" |
|
|
| ( |
| text_state_cond, |
| text_mask_cond, |
| speaker_state_cond, |
| speaker_mask_cond, |
| caption_state_cond, |
| caption_mask_cond, |
| ) = model.encode_conditions( |
| text_input_ids=text_input_ids, |
| text_mask=text_mask, |
| ref_latent=ref_latent, |
| ref_mask=ref_mask, |
| caption_input_ids=caption_input_ids, |
| caption_mask=caption_mask, |
| ) |
| text_state_uncond = torch.zeros_like(text_state_cond) |
| text_mask_uncond = torch.zeros_like(text_mask_cond) |
| speaker_state_uncond = None |
| speaker_mask_uncond = None |
| if model.cfg.use_speaker_condition: |
| if speaker_state_cond is None or speaker_mask_cond is None: |
| raise RuntimeError( |
| "Speaker conditioning is enabled but encoded speaker state is missing." |
| ) |
| speaker_state_uncond = torch.zeros_like(speaker_state_cond) |
| speaker_mask_uncond = torch.zeros_like(speaker_mask_cond) |
| caption_state_uncond = None |
| caption_mask_uncond = None |
| if model.cfg.use_caption_condition: |
| if caption_state_cond is None or caption_mask_cond is None: |
| raise RuntimeError( |
| "Caption conditioning is enabled but encoded caption state is missing." |
| ) |
| caption_state_uncond = torch.zeros_like(caption_state_cond) |
| caption_mask_uncond = torch.zeros_like(caption_mask_cond) |
|
|
| has_text_cfg = cfg_scale_text > 0 |
| has_caption_cfg = ( |
| model.cfg.use_caption_condition |
| and cfg_scale_caption > 0 |
| and caption_mask_cond is not None |
| and bool(caption_mask_cond.any().item()) |
| ) |
| has_speaker_cfg = cfg_scale_speaker > 0 |
|
|
| def _bundle( |
| *, |
| text_state: torch.Tensor, |
| text_mask_val: torch.Tensor, |
| speaker_state: torch.Tensor | None, |
| speaker_mask_val: torch.Tensor | None, |
| caption_state: torch.Tensor | None, |
| caption_mask_val: torch.Tensor | None, |
| ) -> tuple[ |
| torch.Tensor, |
| torch.Tensor, |
| torch.Tensor | None, |
| torch.Tensor | None, |
| torch.Tensor | None, |
| torch.Tensor | None, |
| ]: |
| return ( |
| text_state, |
| text_mask_val, |
| speaker_state, |
| speaker_mask_val, |
| caption_state, |
| caption_mask_val, |
| ) |
|
|
| cond_bundle = _bundle( |
| text_state=text_state_cond, |
| text_mask_val=text_mask_cond, |
| speaker_state=speaker_state_cond, |
| speaker_mask_val=speaker_mask_cond, |
| caption_state=caption_state_cond, |
| caption_mask_val=caption_mask_cond, |
| ) |
| enabled_cfg_names: list[str] = [] |
| cfg_scales: dict[str, float] = {} |
| if has_text_cfg: |
| enabled_cfg_names.append("text") |
| cfg_scales["text"] = float(cfg_scale_text) |
| if has_speaker_cfg: |
| enabled_cfg_names.append("speaker") |
| cfg_scales["speaker"] = float(cfg_scale_speaker) |
| if has_caption_cfg: |
| enabled_cfg_names.append("caption") |
| cfg_scales["caption"] = float(cfg_scale_caption) |
|
|
| independent_bundles = [cond_bundle] |
| independent_names = ["cond"] |
| if use_independent_cfg: |
| for name in enabled_cfg_names: |
| independent_names.append(name) |
| independent_bundles.append( |
| _bundle( |
| text_state=text_state_uncond if name == "text" else text_state_cond, |
| text_mask_val=text_mask_uncond if name == "text" else text_mask_cond, |
| speaker_state=( |
| speaker_state_uncond if name == "speaker" else speaker_state_cond |
| ), |
| speaker_mask_val=( |
| speaker_mask_uncond if name == "speaker" else speaker_mask_cond |
| ), |
| caption_state=( |
| caption_state_uncond if name == "caption" else caption_state_cond |
| ), |
| caption_mask_val=( |
| caption_mask_uncond if name == "caption" else caption_mask_cond |
| ), |
| ) |
| ) |
| cfg_batch_mult = len(independent_bundles) |
|
|
| def _cat_optional_tensors(values: list[torch.Tensor | None]) -> torch.Tensor | None: |
| present = [value for value in values if value is not None] |
| if not present: |
| return None |
| if len(present) != len(values): |
| raise ValueError("Cannot concatenate optional condition tensors with mixed presence.") |
| return torch.cat(present, dim=0) |
|
|
| independent_text_state = torch.cat([bundle[0] for bundle in independent_bundles], dim=0) |
| independent_text_mask = torch.cat([bundle[1] for bundle in independent_bundles], dim=0) |
| independent_speaker_state = _cat_optional_tensors([bundle[2] for bundle in independent_bundles]) |
| independent_speaker_mask = _cat_optional_tensors([bundle[3] for bundle in independent_bundles]) |
| independent_caption_state = _cat_optional_tensors([bundle[4] for bundle in independent_bundles]) |
| independent_caption_mask = _cat_optional_tensors([bundle[5] for bundle in independent_bundles]) |
|
|
| joint_uncond_bundle = _bundle( |
| text_state=text_state_uncond, |
| text_mask_val=text_mask_uncond, |
| speaker_state=speaker_state_uncond, |
| speaker_mask_val=speaker_mask_uncond, |
| caption_state=caption_state_uncond, |
| caption_mask_val=caption_mask_uncond, |
| ) |
|
|
| alternating_bundles: dict[ |
| str, |
| tuple[ |
| torch.Tensor, |
| torch.Tensor, |
| torch.Tensor | None, |
| torch.Tensor | None, |
| torch.Tensor | None, |
| torch.Tensor | None, |
| ], |
| ] = { |
| "text": _bundle( |
| text_state=text_state_uncond, |
| text_mask_val=text_mask_uncond, |
| speaker_state=speaker_state_cond, |
| speaker_mask_val=speaker_mask_cond, |
| caption_state=caption_state_cond, |
| caption_mask_val=caption_mask_cond, |
| ), |
| "caption": _bundle( |
| text_state=text_state_cond, |
| text_mask_val=text_mask_cond, |
| speaker_state=speaker_state_cond, |
| speaker_mask_val=speaker_mask_cond, |
| caption_state=caption_state_uncond, |
| caption_mask_val=caption_mask_uncond, |
| ), |
| } |
| if has_speaker_cfg: |
| alternating_bundles["speaker"] = _bundle( |
| text_state=text_state_cond, |
| text_mask_val=text_mask_cond, |
| speaker_state=speaker_state_uncond, |
| speaker_mask_val=speaker_mask_uncond, |
| caption_state=caption_state_cond, |
| caption_mask_val=caption_mask_cond, |
| ) |
|
|
| |
| effective_use_context_kv_cache = bool(use_context_kv_cache or (speaker_kv_scale is not None)) |
|
|
| context_kv_cond = None |
| context_kv_cfg = None |
| context_kv_joint_uncond = None |
| context_kv_alternating: dict[str, list[tuple[torch.Tensor, ...]]] = {} |
| if effective_use_context_kv_cache: |
| context_kv_cond = model.build_context_kv_cache( |
| text_state=text_state_cond, |
| speaker_state=speaker_state_cond, |
| caption_state=caption_state_cond, |
| ) |
| if use_independent_cfg and cfg_batch_mult > 1: |
| context_kv_cfg = model.build_context_kv_cache( |
| text_state=independent_text_state, |
| speaker_state=independent_speaker_state, |
| caption_state=independent_caption_state, |
| ) |
| elif use_joint_cfg: |
| if enabled_cfg_names: |
| context_kv_joint_uncond = model.build_context_kv_cache( |
| text_state=joint_uncond_bundle[0], |
| speaker_state=joint_uncond_bundle[2], |
| caption_state=joint_uncond_bundle[4], |
| ) |
| elif use_alternating_cfg: |
| for name in enabled_cfg_names: |
| bundle = alternating_bundles[name] |
| context_kv_alternating[name] = model.build_context_kv_cache( |
| text_state=bundle[0], |
| speaker_state=bundle[2], |
| caption_state=bundle[4], |
| ) |
| if speaker_kv_scale is not None: |
| scale_speaker_kv_cache( |
| context_kv_cache=context_kv_cond, |
| scale=float(speaker_kv_scale), |
| max_layers=speaker_kv_max_layers, |
| ) |
| if context_kv_cfg is not None: |
| scale_speaker_kv_cache( |
| context_kv_cache=context_kv_cfg, |
| scale=float(speaker_kv_scale), |
| max_layers=speaker_kv_max_layers, |
| ) |
| for cache in context_kv_alternating.values(): |
| scale_speaker_kv_cache( |
| context_kv_cache=cache, |
| scale=float(speaker_kv_scale), |
| max_layers=speaker_kv_max_layers, |
| ) |
| speaker_kv_active = speaker_kv_scale is not None |
|
|
| for i in range(num_steps): |
| t = t_schedule[i] |
| t_next = t_schedule[i + 1] |
| tt = torch.full((batch_size,), t, device=device, dtype=dtype) |
|
|
| use_cfg = bool(enabled_cfg_names) and (cfg_min_t <= t.item() <= cfg_max_t) |
| if use_cfg: |
| if use_independent_cfg: |
| x_t_cfg = torch.cat([x_t] * cfg_batch_mult, dim=0).to(dtype) |
| tt_cfg = tt.repeat(cfg_batch_mult) |
| v_out = model.forward_with_encoded_conditions( |
| x_t=x_t_cfg, |
| t=tt_cfg, |
| text_state=independent_text_state, |
| text_mask=independent_text_mask, |
| speaker_state=independent_speaker_state, |
| speaker_mask=independent_speaker_mask, |
| caption_state=independent_caption_state, |
| caption_mask=independent_caption_mask, |
| context_kv_cache=context_kv_cfg, |
| ) |
| chunks = v_out.chunk(cfg_batch_mult, dim=0) |
| v = chunks[0] |
| for name, chunk in zip(independent_names[1:], chunks[1:], strict=True): |
| v = v + cfg_scales[name] * (chunks[0] - chunk) |
| else: |
| v_cond = model.forward_with_encoded_conditions( |
| x_t=x_t.to(dtype), |
| t=tt, |
| text_state=text_state_cond, |
| text_mask=text_mask_cond, |
| speaker_state=speaker_state_cond, |
| speaker_mask=speaker_mask_cond, |
| caption_state=caption_state_cond, |
| caption_mask=caption_mask_cond, |
| context_kv_cache=context_kv_cond, |
| ) |
| if use_joint_cfg: |
| if len(enabled_cfg_names) > 1: |
| joint_scales = [cfg_scales[name] for name in enabled_cfg_names] |
| if max(joint_scales) - min(joint_scales) > 1e-6: |
| raise ValueError( |
| "cfg_guidance_mode='joint' expects equal enabled guidance scales; " |
| "set matching text/speaker/caption scales or use --cfg-scale." |
| ) |
| joint_scale = cfg_scales[enabled_cfg_names[0]] |
| v_uncond_joint = model.forward_with_encoded_conditions( |
| x_t=x_t.to(dtype), |
| t=tt, |
| text_state=joint_uncond_bundle[0], |
| text_mask=joint_uncond_bundle[1], |
| speaker_state=joint_uncond_bundle[2], |
| speaker_mask=joint_uncond_bundle[3], |
| caption_state=joint_uncond_bundle[4], |
| caption_mask=joint_uncond_bundle[5], |
| context_kv_cache=context_kv_joint_uncond, |
| ) |
| v = v_cond + joint_scale * (v_cond - v_uncond_joint) |
| elif use_alternating_cfg: |
| alt_name = enabled_cfg_names[i % len(enabled_cfg_names)] |
| alt_bundle = alternating_bundles[alt_name] |
| v_uncond_alt = model.forward_with_encoded_conditions( |
| x_t=x_t.to(dtype), |
| t=tt, |
| text_state=alt_bundle[0], |
| text_mask=alt_bundle[1], |
| speaker_state=alt_bundle[2], |
| speaker_mask=alt_bundle[3], |
| caption_state=alt_bundle[4], |
| caption_mask=alt_bundle[5], |
| context_kv_cache=context_kv_alternating.get(alt_name), |
| ) |
| v = v_cond + cfg_scales[alt_name] * (v_cond - v_uncond_alt) |
| else: |
| raise RuntimeError(f"Unexpected cfg_guidance_mode: {cfg_guidance_mode}") |
| else: |
| v = model.forward_with_encoded_conditions( |
| x_t=x_t.to(dtype), |
| t=tt, |
| text_state=text_state_cond, |
| text_mask=text_mask_cond, |
| speaker_state=speaker_state_cond, |
| speaker_mask=speaker_mask_cond, |
| caption_state=caption_state_cond, |
| caption_mask=caption_mask_cond, |
| context_kv_cache=context_kv_cond, |
| ) |
|
|
| if rescale_k is not None and rescale_sigma is not None: |
| v = temporal_score_rescale( |
| v_pred=v, |
| x_t=x_t, |
| t=t, |
| rescale_k=float(rescale_k), |
| rescale_sigma=float(rescale_sigma), |
| ) |
|
|
| if ( |
| speaker_kv_active |
| and speaker_kv_min_t is not None |
| and (t_next < speaker_kv_min_t) |
| and (t >= speaker_kv_min_t) |
| ): |
| inv_scale = 1.0 / float(speaker_kv_scale) |
| scale_speaker_kv_cache( |
| context_kv_cache=context_kv_cond, |
| scale=inv_scale, |
| max_layers=speaker_kv_max_layers, |
| ) |
| if context_kv_cfg is not None: |
| scale_speaker_kv_cache( |
| context_kv_cache=context_kv_cfg, |
| scale=inv_scale, |
| max_layers=speaker_kv_max_layers, |
| ) |
| for cache in context_kv_alternating.values(): |
| scale_speaker_kv_cache( |
| context_kv_cache=cache, |
| scale=inv_scale, |
| max_layers=speaker_kv_max_layers, |
| ) |
| speaker_kv_active = False |
|
|
| x_t = x_t + v * (t_next - t) |
|
|
| return x_t |
|
|