# Adapted for diffusers from multimodal-art-projection/YuE at commit ef1936f2ee39fe8de486a0f47a481c95f8d4da87. # Licensed under Apache-2.0; see LICENSE. from __future__ import annotations import math from dataclasses import dataclass import torch import torch.nn.functional as F from diffusers import ConfigMixin, ModelMixin from diffusers.configuration_utils import register_to_config from diffusers.models.attention import AttentionModuleMixin from diffusers.models.attention_dispatch import dispatch_attention_fn from diffusers.models.embeddings import TimestepEmbedding, Timesteps from diffusers.utils import BaseOutput from torch import nn @dataclass class YuE2TransformerOutput(BaseOutput): """ Args: logits (`torch.Tensor` of shape `(batch_size, num_logits, vocab_size)`, *optional*): Next-token logits, returned for token inputs. sample (`torch.Tensor` of shape `(batch_size, num_frames, latent_channels)`, *optional*): Flow-matching velocity, returned for acoustic latent inputs. """ logits: torch.Tensor | None = None sample: torch.Tensor | None = None class YuE2RMSNorm(nn.Module): def __init__(self, dim: int, eps: float = 1e-6): super().__init__() self.weight = nn.Parameter(torch.ones(dim)) self.eps = eps def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: variance = hidden_states.float().pow(2).mean(-1, keepdim=True) return hidden_states * torch.rsqrt(variance + self.eps).to(hidden_states.dtype) * self.weight class YuE2RotaryEmbedding(nn.Module): def __init__(self, head_dim: int, theta: float): super().__init__() self.head_dim = head_dim self.theta = theta # Kept in float32 on the positions' device, so it is not a buffer that `.to(dtype)` would cast. self._inv_freq = None def forward(self, position_ids: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]: if self._inv_freq is None or self._inv_freq.device != position_ids.device: exponent = torch.arange(0, self.head_dim, 2, dtype=torch.float32, device=position_ids.device) self._inv_freq = 1.0 / self.theta ** (exponent / self.head_dim) angles = position_ids.float().unsqueeze(-1) * self._inv_freq return angles.cos(), angles.sin() def apply_rotary_emb(hidden_states: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: half = hidden_states.shape[-1] // 2 x1, x2 = hidden_states[..., :half], hidden_states[..., half:] cos, sin = cos.to(hidden_states.dtype), sin.to(hidden_states.dtype) return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1) class YuE2LatentPositionEmbedding(nn.Module): def __init__(self, max_frames: int, dim: int): super().__init__() position = torch.arange(max_frames, dtype=torch.float32).unsqueeze(1) div_term = torch.exp(torch.arange(0, dim, 2, dtype=torch.float32) * (-math.log(10000.0) / dim)) pe = torch.zeros(max_frames, dim) pe[:, 0::2] = torch.sin(position * div_term) pe[:, 1::2] = torch.cos(position * div_term) self.register_buffer("pe", pe) def forward(self, position_ids: torch.Tensor) -> torch.Tensor: return self.pe[position_ids] class YuE2AttnProcessor: _attention_backend = None _parallel_config = None def __call__( self, attn: "YuE2Attention", hidden_states: torch.Tensor, rotary_emb: tuple[torch.Tensor, torch.Tensor], kv_cache=None, layer_idx: int = 0, ) -> torch.Tensor: batch_size, seq_len, _ = hidden_states.shape query = attn.to_q(hidden_states).view(batch_size, seq_len, attn.heads, attn.head_dim) key = attn.to_k(hidden_states).view(batch_size, seq_len, attn.kv_heads, attn.head_dim) value = attn.to_v(hidden_states).view(batch_size, seq_len, attn.kv_heads, attn.head_dim) query = attn.norm_q(query) key = attn.norm_k(key) cos, sin = rotary_emb[0].unsqueeze(2), rotary_emb[1].unsqueeze(2) query = apply_rotary_emb(query, cos, sin) key = apply_rotary_emb(key, cos, sin) if kv_cache is not None: key, value = kv_cache.update(key, value, layer_idx) # Tokens attend causally to themselves; acoustic frames attend to the whole cached prefix and to each other. is_causal = seq_len > 1 and key.shape[1] == seq_len enable_gqa = attn.heads != attn.kv_heads if enable_gqa and query.device.type == "mps": key = key.repeat_interleave(attn.heads // attn.kv_heads, dim=2) value = value.repeat_interleave(attn.heads // attn.kv_heads, dim=2) enable_gqa = False hidden_states = dispatch_attention_fn( query, key, value, is_causal=is_causal, enable_gqa=enable_gqa, backend=self._attention_backend, parallel_config=self._parallel_config, ) hidden_states = hidden_states.reshape(batch_size, seq_len, -1) return attn.to_out[0](hidden_states) class YuE2Attention(nn.Module, AttentionModuleMixin): _default_processor_cls = YuE2AttnProcessor _available_processors = [YuE2AttnProcessor] def __init__(self, dim: int, heads: int, kv_heads: int, head_dim: int, eps: float): super().__init__() self.heads = heads self.kv_heads = kv_heads self.head_dim = head_dim self.to_q = nn.Linear(dim, heads * head_dim, bias=False) self.to_k = nn.Linear(dim, kv_heads * head_dim, bias=False) self.to_v = nn.Linear(dim, kv_heads * head_dim, bias=False) self.to_out = nn.ModuleList([nn.Linear(heads * head_dim, dim, bias=False)]) self.norm_q = YuE2RMSNorm(head_dim, eps) self.norm_k = YuE2RMSNorm(head_dim, eps) self.set_processor(YuE2AttnProcessor()) def forward( self, hidden_states: torch.Tensor, rotary_emb: tuple[torch.Tensor, torch.Tensor], kv_cache=None, layer_idx: int = 0, ) -> torch.Tensor: return self.processor(self, hidden_states, rotary_emb, kv_cache, layer_idx) class YuE2FeedForward(nn.Module): def __init__(self, dim: int, inner_dim: int): super().__init__() self.gate_proj = nn.Linear(dim, inner_dim, bias=False) self.up_proj = nn.Linear(dim, inner_dim, bias=False) self.down_proj = nn.Linear(inner_dim, dim, bias=False) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: return self.down_proj(F.silu(self.gate_proj(hidden_states)) * self.up_proj(hidden_states)) class YuE2TransformerBlock(nn.Module): """Mixture-of-transformers layer: token and acoustic streams have separate attention and feed-forward weights.""" def __init__(self, dim: int, heads: int, kv_heads: int, head_dim: int, inner_dim: int, eps: float): super().__init__() self.norm1 = YuE2RMSNorm(dim, eps) self.attn = YuE2Attention(dim, heads, kv_heads, head_dim, eps) self.norm2 = YuE2RMSNorm(dim, eps) self.ff = YuE2FeedForward(dim, inner_dim) self.acoustic_norm1 = YuE2RMSNorm(dim, eps) self.acoustic_attn = YuE2Attention(dim, heads, kv_heads, head_dim, eps) self.acoustic_norm2 = YuE2RMSNorm(dim, eps) self.acoustic_ff = YuE2FeedForward(dim, inner_dim) def forward( self, hidden_states: torch.Tensor, rotary_emb: tuple[torch.Tensor, torch.Tensor], kv_cache=None, layer_idx: int = 0, acoustic: bool = False, ) -> torch.Tensor: if acoustic: attn_output = self.acoustic_attn(self.acoustic_norm1(hidden_states), rotary_emb, kv_cache, layer_idx) hidden_states = hidden_states + attn_output hidden_states = hidden_states + self.acoustic_ff(self.acoustic_norm2(hidden_states)) else: attn_output = self.attn(self.norm1(hidden_states), rotary_emb, kv_cache, layer_idx) hidden_states = hidden_states + attn_output hidden_states = hidden_states + self.ff(self.norm2(hidden_states)) return hidden_states class YuE2TransformerModel(ModelMixin, ConfigMixin): r""" The YuE2 mixture-of-transformers. One set of layers carries two weight streams: - Token forward (`input_ids`): a causal language model over text, ABC score and codec tokens. It returns next-token logits and, given a `kv_cache`, appends the tokens' keys and values to it. - Acoustic forward (`latents`, `timestep`, `kv_cache`): the flow-matching velocity for one chunk of latent frames. The frames attend to the token prefix already stored in `kv_cache` and to each other. `kv_cache` is any object with `get_seq_length()` and `update(key, value, layer_idx)`, taking and returning `(batch_size, seq_len, heads, head_dim)` tensors. """ _no_split_modules = ["YuE2TransformerBlock"] _repeated_blocks = ["YuE2TransformerBlock"] _skip_keys = ["kv_cache"] @register_to_config def __init__( self, hidden_size: int = 2048, num_layers: int = 28, num_attention_heads: int = 16, num_key_value_heads: int = 8, attention_head_dim: int = 128, intermediate_size: int = 6144, vocab_size: int = 184704, norm_eps: float = 1e-6, rope_theta: float = 1000000.0, max_position_embeddings: int = 24576, latent_channels: int = 64, max_latent_frames: int = 24576, timestep_shift: float = 1.0, ): super().__init__() self.embed_tokens = nn.Embedding(vocab_size, hidden_size) self.rotary_emb = YuE2RotaryEmbedding(attention_head_dim, rope_theta) self.proj_in = nn.Linear(latent_channels, hidden_size) self.time_proj = Timesteps(256, flip_sin_to_cos=True, downscale_freq_shift=0) self.time_embedding = TimestepEmbedding(256, hidden_size) self.latent_pos_embed = YuE2LatentPositionEmbedding(max_latent_frames, hidden_size) self.transformer_blocks = nn.ModuleList( [ YuE2TransformerBlock( hidden_size, num_attention_heads, num_key_value_heads, attention_head_dim, intermediate_size, norm_eps, ) for _ in range(num_layers) ] ) self.norm_out = YuE2RMSNorm(hidden_size, norm_eps) self.lm_head = nn.Linear(hidden_size, vocab_size, bias=False) self.proj_out = nn.Linear(hidden_size, latent_channels) def forward( self, input_ids: torch.LongTensor | None = None, latents: torch.Tensor | None = None, timestep: float | None = None, kv_cache=None, logits_to_keep: int = 0, return_dict: bool = True, ) -> YuE2TransformerOutput | tuple[torch.Tensor]: """ Args: input_ids (`torch.LongTensor` of shape `(batch_size, seq_len)`, *optional*): Token IDs for the token forward. Pass either `input_ids` or `latents`. latents (`torch.Tensor` of shape `(batch_size, num_frames, latent_channels)`, *optional*): Noisy acoustic latents for the acoustic forward. timestep (`float`, *optional*): Flow-matching time in logit space, required with `latents`. kv_cache (*optional*): Key/value cache. The token forward appends to it; the acoustic forward reads the token prefix from it. logits_to_keep (`int`, defaults to 0): Number of trailing positions to compute logits for; 0 keeps all positions. return_dict (`bool`, defaults to `True`): Whether to return a [`YuE2TransformerOutput`] instead of a plain tuple. Returns: [`YuE2TransformerOutput`] or `tuple`: logits for token inputs, velocity (`sample`) for latent inputs. """ if (input_ids is None) == (latents is None): raise ValueError("Pass exactly one of `input_ids` or `latents`.") past_length = kv_cache.get_seq_length() if kv_cache is not None else 0 if latents is not None: if timestep is None or past_length == 0: raise ValueError("The acoustic forward needs `timestep` and a `kv_cache` holding the token prefix.") # Each chunk is framed by a start and an end frame, both zero. hidden_states = self.proj_in(F.pad(latents, (0, 0, 1, 1))) num_frames = hidden_states.shape[1] timestep = torch.sigmoid(torch.tensor(timestep, dtype=hidden_states.dtype, device=hidden_states.device)) shift = self.config.timestep_shift timestep = shift * timestep / (1 + (shift - 1) * timestep) timesteps_proj = self.time_proj(timestep.expand(num_frames)).to(hidden_states.dtype) hidden_states = hidden_states + self.time_embedding(timesteps_proj)[None] frame_ids = torch.arange(num_frames, device=hidden_states.device).clamp( max=self.config.max_latent_frames - 1 ) hidden_states = hidden_states + self.latent_pos_embed(frame_ids)[None] seq_len = num_frames else: seq_len = input_ids.shape[1] if past_length and seq_len > 1: raise ValueError("Prefill the prompt in one call, then pass one new token per call.") hidden_states = self.embed_tokens(input_ids) position_ids = torch.arange(past_length, past_length + seq_len, device=hidden_states.device)[None] rotary_emb = self.rotary_emb(position_ids) for layer_idx, block in enumerate(self.transformer_blocks): hidden_states = block(hidden_states, rotary_emb, kv_cache, layer_idx, acoustic=latents is not None) hidden_states = self.norm_out(hidden_states) if latents is not None: sample = self.proj_out(hidden_states)[:, 1:-1] return YuE2TransformerOutput(sample=sample) if return_dict else (sample,) if logits_to_keep: hidden_states = hidden_states[:, -logits_to_keep:] logits = self.lm_head(hidden_states) return YuE2TransformerOutput(logits=logits) if return_dict else (logits,)