Spaces:
Running
Running
File size: 6,659 Bytes
1e69a1f | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 | """Reference/text embedding preprocessing helpers for conditioned generation."""
from typing import Dict, List, Tuple
import torch
from loguru import logger
class ConditioningEmbedMixin:
"""Mixin containing reference/text embedding preprocessing steps.
Depends on host members:
- Attributes: ``device``, ``dtype``, ``silence_latent``, ``text_encoder``.
- Methods: ``_ensure_silence_latent_on_device``, ``_load_model_context``,
``tiled_encode``.
"""
def infer_refer_latent(self, refer_audioss: List[List[torch.Tensor]]) -> Tuple[torch.Tensor, torch.Tensor]:
"""Infer packed reference-audio latents and order mask."""
refer_audio_order_mask = []
refer_audio_latents = []
self._ensure_silence_latent_on_device()
def _normalize_audio_2d(a: torch.Tensor) -> torch.Tensor:
if not isinstance(a, torch.Tensor):
raise TypeError(f"refer_audio must be a torch.Tensor, got {type(a)!r}")
if a.dim() == 3 and a.shape[0] == 1:
a = a.squeeze(0)
if a.dim() == 1:
a = a.unsqueeze(0)
if a.dim() != 2:
raise ValueError(f"refer_audio must be 1D/2D/3D(1,2,T); got shape={tuple(a.shape)}")
if a.shape[0] == 1:
a = torch.cat([a, a], dim=0)
return a[:2]
def _ensure_latent_3d(z: torch.Tensor) -> torch.Tensor:
if z.dim() == 4 and z.shape[0] == 1:
z = z.squeeze(0)
if z.dim() == 2:
z = z.unsqueeze(0)
return z
refer_encode_cache: Dict[int, torch.Tensor] = {}
for batch_idx, refer_audios in enumerate(refer_audioss):
if len(refer_audios) == 1 and torch.all(refer_audios[0] == 0.0):
refer_audio_latent = _ensure_latent_3d(self.silence_latent[:, :750, :])
refer_audio_latents.append(refer_audio_latent)
refer_audio_order_mask.append(batch_idx)
else:
for refer_audio in refer_audios:
cache_key = refer_audio.data_ptr()
if cache_key in refer_encode_cache:
refer_audio_latent = refer_encode_cache[cache_key].clone()
else:
refer_audio = _normalize_audio_2d(refer_audio)
with torch.inference_mode():
refer_audio_latent = self.tiled_encode(refer_audio, offload_latent_to_cpu=True)
refer_audio_latent = refer_audio_latent.to(self.device).to(self.dtype)
if refer_audio_latent.dim() == 2:
refer_audio_latent = refer_audio_latent.unsqueeze(0)
refer_audio_latent = _ensure_latent_3d(refer_audio_latent.transpose(1, 2))
refer_encode_cache[cache_key] = refer_audio_latent
refer_audio_latents.append(refer_audio_latent)
refer_audio_order_mask.append(batch_idx)
refer_audio_latents = torch.cat(refer_audio_latents, dim=0)
refer_audio_order_mask = torch.tensor(refer_audio_order_mask, device=self.device, dtype=torch.long)
return refer_audio_latents, refer_audio_order_mask
def infer_text_embeddings(self, text_token_idss):
"""Infer text-token embeddings via text encoder."""
with torch.inference_mode():
return self.text_encoder(input_ids=text_token_idss, lyric_attention_mask=None).last_hidden_state
def infer_lyric_embeddings(self, lyric_token_ids):
"""Infer lyric-token embeddings via text encoder embedding table."""
with torch.inference_mode():
return self.text_encoder.embed_tokens(lyric_token_ids)
def preprocess_batch(self, batch) -> Tuple:
"""Preprocess an already prepared batch for DiT model input."""
target_latents = batch["target_latents"]
src_latents = batch["src_latents"]
attention_mask = batch["latent_masks"]
audio_codes = batch.get("audio_codes", None)
audio_attention_mask = attention_mask
dtype = target_latents.dtype
device = target_latents.device
keys = batch["keys"]
with self._load_model_context("vae"):
refer_audio_acoustic_hidden_states_packed, refer_audio_order_mask = self.infer_refer_latent(
batch["refer_audioss"]
)
if refer_audio_acoustic_hidden_states_packed.dtype != dtype:
refer_audio_acoustic_hidden_states_packed = refer_audio_acoustic_hidden_states_packed.to(dtype)
chunk_mask = batch["chunk_masks"]
chunk_mask = chunk_mask.to(device).unsqueeze(-1).repeat(1, 1, target_latents.shape[2])
spans = batch["spans"]
text_token_idss = batch["text_token_idss"]
text_attention_mask = batch["text_attention_masks"]
lyric_token_idss = batch["lyric_token_idss"]
lyric_attention_mask = batch["lyric_attention_masks"]
text_inputs = batch["text_inputs"]
logger.info("[preprocess_batch] Inferring prompt embeddings...")
with self._load_model_context("text_encoder"):
text_hidden_states = self.infer_text_embeddings(text_token_idss)
logger.info("[preprocess_batch] Inferring lyric embeddings...")
lyric_hidden_states = self.infer_lyric_embeddings(lyric_token_idss)
is_covers = batch["is_covers"]
precomputed_lm_hints_25hz = batch.get("precomputed_lm_hints_25Hz", None)
non_cover_text_input_ids = batch.get("non_cover_text_input_ids", None)
non_cover_text_attention_masks = batch.get("non_cover_text_attention_masks", None)
non_cover_text_hidden_states = None
if non_cover_text_input_ids is not None:
logger.info("[preprocess_batch] Inferring non-cover text embeddings...")
non_cover_text_hidden_states = self.infer_text_embeddings(non_cover_text_input_ids)
return (
keys,
text_inputs,
src_latents,
target_latents,
text_hidden_states,
text_attention_mask,
lyric_hidden_states,
lyric_attention_mask,
audio_attention_mask,
refer_audio_acoustic_hidden_states_packed,
refer_audio_order_mask,
chunk_mask,
spans,
is_covers,
audio_codes,
lyric_token_idss,
precomputed_lm_hints_25hz,
non_cover_text_hidden_states,
non_cover_text_attention_masks,
)
|