Ace-Step-Munk / acestep /core /generation /handler /conditioning_embed.py
OnyxMunk's picture
Upload folder using huggingface_hub
1e69a1f verified
Raw
History Blame Contribute Delete
6.66 kB
"""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,
)