Instructions to use OzzyGT/YuE2-Modular with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- Diffusers
How to use OzzyGT/YuE2-Modular 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-Modular", 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
File size: 9,829 Bytes
2577656 | 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 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 | # Adapted for diffusers from multimodal-art-projection/YuE at commit ef1936f2ee39fe8de486a0f47a481c95f8d4da87.
# Licensed under Apache-2.0; see LICENSE.
"""CUDA-graph token decoding for one request or its two guidance branches.
The graph replays the transformer's token stream for one new token per branch. The caller combines the branch logits,
samples once and passes that token to `step`. Each branch keeps its own RoPE positions and cache slots.
"""
from __future__ import annotations
import torch
def apply_rotary_emb(hidden_states, cos, sin):
# Same arithmetic as `apply_rotary_emb` in the model repository's transformer.py.
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 _BranchPrefillCache:
"""Writes one branch's prompt keys and values into the graph's sequence-major buffers."""
def __init__(self, keys, values, branch):
self.keys = [buffer[branch : branch + 1] for buffer in keys]
self.values = [buffer[branch : branch + 1] for buffer in values]
self.seen_tokens = 0
def get_seq_length(self):
return self.seen_tokens
def update(self, key, value, layer_idx):
end = self.seen_tokens + key.shape[1]
if end > self.keys[layer_idx].shape[1]:
raise ValueError("Prefix exceeds preallocated KV capacity")
self.keys[layer_idx][:, self.seen_tokens : end].copy_(key)
self.values[layer_idx][:, self.seen_tokens : end].copy_(value)
if layer_idx == len(self.keys) - 1:
self.seen_tokens = end
return self.keys[layer_idx][:, :end], self.values[layer_idx][:, :end]
class GraphAR:
"""Fixed-capacity decoding. `prefill()` returns the first logits; up to `max_tokens - 1` `step(token)` calls follow.
Returned logits share output storage and stay valid until the next step.
"""
def __init__(self, transformer, prefixes, max_tokens, device):
if transformer.training:
raise ValueError("GraphAR requires transformer.eval()")
if not 1 <= len(prefixes) <= 2:
raise ValueError("GraphAR supports one request or exactly two guidance branches")
if device.type != "cuda":
raise ValueError("CUDA graphs require a CUDA device")
try:
from diffusers.hooks.group_offloading import _get_group_onload_device
_get_group_onload_device(transformer)
except ValueError:
pass
else:
raise ValueError(
"CUDA graphs cannot capture a group-offloaded transformer: streamed weights change address between "
"steps. Pass use_cuda_graph=False, or keep the transformer resident on the GPU."
)
config = transformer.config
prefixes = [[int(token) for token in prefix] for prefix in prefixes]
for prefix in prefixes:
if not prefix or not all(0 <= token < config.vocab_size for token in prefix):
raise ValueError("Prefixes require valid token IDs")
if len(prefix) + max_tokens > config.max_position_embeddings:
raise ValueError("Prefix plus generation budget exceeds model context; no length was shortened")
self.transformer, self.device, self.dtype = transformer, device, transformer.dtype
self.prefixes, self.max_tokens, self.branches = prefixes, int(max_tokens), len(prefixes)
self.capacity = max(map(len, prefixes)) + self.max_tokens
self.graph = self.output = None
# PyTorch's variable-length FlashAttention takes per-branch key lengths on the GPU, so no mask is built.
flash = (
self.dtype in {torch.bfloat16, torch.float16}
and config.attention_head_dim % 8 == 0
and config.attention_head_dim <= 256
and "seqused_k" in str(torch.ops.aten._flash_attention_forward.default._schema)
)
if not flash:
raise ValueError(
"CUDA graphs need PyTorch's variable-length FlashAttention with a BF16/FP16 transformer; "
"pass use_cuda_graph=False."
)
self.ready, self.closed, self.steps = False, False, 0
# Sequence-major buffers make the packed FlashAttention view contiguous without copying the cache each step.
shape = (self.branches, self.capacity, config.num_key_value_heads, config.attention_head_dim)
self.keys = [torch.zeros(shape, device=device, dtype=self.dtype) for _ in range(config.num_layers)]
self.values = [torch.zeros(shape, device=device, dtype=self.dtype) for _ in range(config.num_layers)]
self.positions = torch.tensor([len(prefix) for prefix in prefixes], dtype=torch.long, device=device)
self.initial_positions = self.positions.clone()
self.cu_q = torch.arange(self.branches + 1, dtype=torch.int32, device=device)
self.cu_k = self.cu_q * self.capacity
self.tokens = torch.tensor([[prefix[-1]] for prefix in prefixes], dtype=torch.long, device=device)
@torch.inference_mode()
def _decode(self):
transformer = self.transformer
config = transformer.config
heads, kv_heads, head_dim = config.num_attention_heads, config.num_key_value_heads, config.attention_head_dim
cos, sin = transformer.rotary_emb(self.positions[:, None])
cos, sin = cos.unsqueeze(2), sin.unsqueeze(2)
hidden_states = transformer.embed_tokens(self.tokens)
used_lengths = (self.positions + 1).to(torch.int32)
slots = self.positions[:, None, None, None].expand(self.branches, 1, kv_heads, head_dim)
for block, keys, values in zip(transformer.transformer_blocks, self.keys, self.values):
attn = block.attn
normalized = block.norm1(hidden_states)
query = attn.norm_q(attn.to_q(normalized).view(self.branches, 1, heads, head_dim))
key = attn.norm_k(attn.to_k(normalized).view(self.branches, 1, kv_heads, head_dim))
value = attn.to_v(normalized).view(self.branches, 1, kv_heads, head_dim)
query = apply_rotary_emb(query, cos, sin)
key = apply_rotary_emb(key, cos, sin)
keys.scatter_(1, slots, key)
values.scatter_(1, slots, value)
# Every slot is allocated, but only each branch's filled prefix and current token are visible.
# The packed variable-length entry point respects seqused_k; the 4D fixed-batch one ignores it.
attn_output = torch.ops.aten._flash_attention_forward(
query[:, 0],
keys.view(-1, kv_heads, head_dim),
values.view(-1, kv_heads, head_dim),
self.cu_q,
self.cu_k,
1,
self.capacity,
0.0,
False,
False,
seqused_k=used_lengths,
)[0][:, None]
hidden_states = hidden_states + attn.to_out[0](attn_output.reshape(self.branches, 1, -1))
hidden_states = hidden_states + block.ff(block.norm2(hidden_states))
output = transformer.lm_head(transformer.norm_out(hidden_states))[:, 0]
self.positions.add_(1)
return output
@torch.inference_mode()
def _capture(self):
with torch.cuda.device(self.device):
current = torch.cuda.current_stream(self.device)
warmup = torch.cuda.Stream(device=self.device)
warmup.wait_stream(current)
with torch.cuda.stream(warmup):
for _ in range(3):
self.positions.copy_(self.initial_positions)
self._decode()
self.positions.copy_(self.initial_positions)
current.wait_stream(warmup)
torch.cuda.synchronize(self.device)
self.graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(self.graph):
self.output = self._decode()
# Warmup and capture wrote one future slot; the first real step overwrites it before it becomes visible.
self.positions.copy_(self.initial_positions)
@torch.inference_mode()
def prefill(self):
if self.closed or self.ready:
raise RuntimeError("prefill must be called exactly once on an open GraphAR")
logits = []
for branch, prefix in enumerate(self.prefixes):
cache = _BranchPrefillCache(self.keys, self.values, branch)
ids = torch.tensor([prefix], dtype=torch.long, device=self.device)
logits.append(self.transformer(ids, kv_cache=cache, logits_to_keep=1).logits[:, -1])
if any(param.device != self.device for param in self.transformer.parameters()):
raise ValueError("CUDA graphs need the whole transformer resident on the GPU; pass use_cuda_graph=False.")
if self.max_tokens > 1:
self._capture()
self.ready = True
return torch.cat(logits, dim=0)
@torch.inference_mode()
def step(self, token):
if self.closed or not self.ready:
raise RuntimeError("Call prefill before step and do not use a closed GraphAR")
if self.steps >= self.max_tokens - 1:
raise ValueError("Requested generation budget is exhausted")
self.tokens.copy_(token.reshape(1, 1).expand(self.branches, 1))
self.graph.replay()
self.steps += 1
return self.output
def close(self):
self.graph = self.output = None
self.keys.clear()
self.values.clear()
self.tokens = self.positions = self.initial_positions = None
self.cu_q = self.cu_k = None
self.closed = True
|