YuE2-3B-Diffusers / transformer.py
OzzyGT's picture
OzzyGT HF Staff
Upload 12 files
7b4c0bf verified
Raw History Blame Contribute Delete
14.4 kB
# 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,)