Instructions to use OzzyGT/YuE2-3B-Diffusers with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use OzzyGT/YuE2-3B-Diffusers with Diffusers:
pip install -U diffusers transformers accelerate
import torch from diffusers import DiffusionPipeline # switch to "mps" for apple devices pipe = DiffusionPipeline.from_pretrained("OzzyGT/YuE2-3B-Diffusers", dtype=torch.bfloat16, device_map="cuda") prompt = "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k" image = pipe(prompt).images[0] - Notebooks
- Google Colab
- Kaggle
Download transformer.py from OzzyGT/YuE2-3B-Diffusers: direct link, hf CLI and curl.
- Browser
- Download file 14.4 kB
-
https://huggingface.co/OzzyGT/YuE2-3B-Diffusers/resolve/main/transformer.py
- Command line
-
hf download hf://OzzyGT/YuE2-3B-Diffusers/transformer.py
-
curl -L -o transformer.py https://huggingface.co/OzzyGT/YuE2-3B-Diffusers/resolve/main/transformer.py
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 | |
| 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"] | |
| 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,) | |