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
| # Adapted for diffusers from multimodal-art-projection/YuE at commit ef1936f2ee39fe8de486a0f47a481c95f8d4da87. | |
| # Licensed under Apache-2.0; see LICENSE. | |
| from dataclasses import replace | |
| import torch | |
| from diffusers import ModelMixin | |
| from diffusers.configuration_utils import FrozenDict | |
| from diffusers.modular_pipelines import ( | |
| ComponentSpec, | |
| InputParam, | |
| LoopSequentialPipelineBlocks, | |
| ModularPipelineBlocks, | |
| OutputParam, | |
| SequentialPipelineBlocks, | |
| ) | |
| from diffusers.utils.dynamic_modules_utils import get_class_from_dynamic_module | |
| from transformers import PreTrainedTokenizer | |
| from .cuda_graph import GraphAR | |
| from .guidance import YuE2SemanticGuider | |
| from .nar import YuE2PrefixKVCache, solve_midpoint, song_chunks | |
| from .protocol import ( | |
| ABC_SAMPLING, | |
| CODEC_OFFSET, | |
| CONTEXT, | |
| SEMANTIC_SAMPLING, | |
| SongRequest, | |
| negative_prefix, | |
| resolve_sampling, | |
| token_prefixes, | |
| ) | |
| from .sampling import generate_tokens | |
| def callback_inputs(): | |
| return [ | |
| InputParam( | |
| "cancelled", type_hint=object, default=None, description="Optional callable returning whether to cancel." | |
| ), | |
| InputParam( | |
| "on_token", type_hint=object, default=None, description="Optional callback receiving stage and token ID." | |
| ), | |
| ] | |
| def check_cancelled(callback): | |
| if callback is not None and callback(): | |
| raise InterruptedError("Generation cancelled") | |
| class YuE2PrepareInputsStep(ModularPipelineBlocks): | |
| model_name = "yue2" | |
| def description(self): | |
| return "Validates the song request." | |
| def inputs(self): | |
| return [ | |
| InputParam("style", type_hint=str, required=True, description="Music style description."), | |
| InputParam("lyrics", type_hint=str, required=True, description="Lyrics with song section tags."), | |
| InputParam("cot", type_hint=str, default="full", description="Score mode: full, melody, or off."), | |
| InputParam( | |
| "seed", | |
| type_hint=int, | |
| default=831001, | |
| description="Request seed. Each stage reseeds from it, as the original release does.", | |
| ), | |
| InputParam("abc", type_hint=str, default=None, description="Optional supplied ABC score."), | |
| ] | |
| def intermediate_outputs(self): | |
| return [OutputParam("request", type_hint=SongRequest, description="Validated song request.")] | |
| def __call__(self, components, state): | |
| block_state = self.get_block_state(state) | |
| block_state.request = SongRequest( | |
| style=block_state.style, | |
| lyrics=block_state.lyrics, | |
| cot=block_state.cot, | |
| seed=block_state.seed, | |
| abc=block_state.abc, | |
| ) | |
| self.set_block_state(state, block_state) | |
| return components, state | |
| class YuE2PlanStep(ModularPipelineBlocks): | |
| model_name = "yue2" | |
| _requirements = {"tiktoken": ">=0.12.0"} | |
| def description(self): | |
| return "Plans the song as an ABC score (or uses the supplied one) and builds the semantic-stage prompt." | |
| def expected_components(self): | |
| return [ComponentSpec("transformer", ModelMixin), ComponentSpec("tokenizer", PreTrainedTokenizer)] | |
| def inputs(self): | |
| return [ | |
| InputParam("request", type_hint=SongRequest, required=True, description="Validated song request."), | |
| InputParam( | |
| "abc_sampling", type_hint=dict, default=None, description="Overrides for the ABC sampling defaults." | |
| ), | |
| InputParam("use_cuda_graph", type_hint=bool, default=False, description="Decode tokens with CUDA graphs."), | |
| ] + callback_inputs() | |
| def intermediate_outputs(self): | |
| return [ | |
| OutputParam("score", type_hint=str, description="Generated or supplied ABC score; None with cot='off'."), | |
| OutputParam("abc_ids", type_hint=list, description="Exact score token IDs."), | |
| OutputParam("prefix", type_hint=list, description="Semantic-stage prompt token IDs."), | |
| OutputParam("abc_truncated", type_hint=bool, description="Whether planning hit its token limit."), | |
| OutputParam("abc_timing", type_hint=dict, description="Planning timings."), | |
| ] | |
| def __call__(self, components, state): | |
| block_state = self.get_block_state(state) | |
| check_cancelled(block_state.cancelled) | |
| request = block_state.request | |
| score, abc_ids, timing, truncated = None, [], {}, False | |
| if request.cot != "off": | |
| if request.abc is not None: | |
| score = request.abc | |
| abc_ids = components.tokenizer.encode(score) | |
| else: | |
| abc_ids, timing, truncated = generate_tokens( | |
| components.transformer, | |
| token_prefixes(request, components.tokenizer), | |
| resolve_sampling(block_state.abc_sampling, ABC_SAMPLING), | |
| request.seed, | |
| "abc", | |
| components._execution_device, | |
| cancelled=block_state.cancelled, | |
| on_token=block_state.on_token, | |
| graph_decoder=GraphAR if block_state.use_cuda_graph else None, | |
| ) | |
| score = components.tokenizer.decode(abc_ids) | |
| block_state.score = score | |
| block_state.abc_ids = abc_ids | |
| block_state.prefix = token_prefixes(request, components.tokenizer, abc_ids) | |
| block_state.abc_truncated = truncated | |
| block_state.abc_timing = timing | |
| self.set_block_state(state, block_state) | |
| return components, state | |
| class YuE2SemanticStep(ModularPipelineBlocks): | |
| model_name = "yue2" | |
| _requirements = {"tiktoken": ">=0.12.0"} | |
| def description(self): | |
| return "Generates the song's codec (semantic) tokens, with guidance from the `guider` component." | |
| def expected_components(self): | |
| return [ | |
| ComponentSpec("transformer", ModelMixin), | |
| ComponentSpec("tokenizer", PreTrainedTokenizer), | |
| ComponentSpec( | |
| "guider", YuE2SemanticGuider, config=FrozenDict({"scale": None}), default_creation_method="from_config" | |
| ), | |
| ] | |
| def inputs(self): | |
| return [ | |
| InputParam("request", type_hint=SongRequest, required=True, description="Validated song request."), | |
| InputParam("abc_ids", type_hint=list, required=True, description="Exact score token IDs."), | |
| InputParam("prefix", type_hint=list, required=True, description="Semantic-stage prompt token IDs."), | |
| InputParam( | |
| "semantic_sampling", | |
| type_hint=dict, | |
| default=None, | |
| description="Overrides for the semantic sampling defaults.", | |
| ), | |
| InputParam("use_cuda_graph", type_hint=bool, default=False, description="Decode tokens with CUDA graphs."), | |
| ] + callback_inputs() | |
| def intermediate_outputs(self): | |
| return [ | |
| OutputParam("semantic_tokens", type_hint=list, description="Codec token IDs, 25 per second of audio."), | |
| OutputParam("semantic_truncated", type_hint=bool, description="Whether generation hit its token limit."), | |
| OutputParam("semantic_timing", type_hint=dict, description="Semantic-stage timings."), | |
| ] | |
| def __call__(self, components, state): | |
| block_state = self.get_block_state(state) | |
| check_cancelled(block_state.cancelled) | |
| request = block_state.request | |
| if token_prefixes(request, components.tokenizer, block_state.abc_ids) != block_state.prefix: | |
| raise ValueError("The semantic prompt disagrees with the request and score IDs") | |
| guider = components.guider | |
| negative = None | |
| if guider.scale_for(request.cot) != 1: | |
| negative = negative_prefix(request, components.tokenizer, block_state.abc_ids) | |
| tokens, timing, truncated = generate_tokens( | |
| components.transformer, | |
| block_state.prefix, | |
| resolve_sampling(block_state.semantic_sampling, SEMANTIC_SAMPLING), | |
| request.seed, | |
| "semantic", | |
| components._execution_device, | |
| negative=negative, | |
| combine_logits=lambda conditional, unconditional: guider.combine(conditional, unconditional, request.cot), | |
| legacy_off=request.cot == "off", | |
| cancelled=block_state.cancelled, | |
| on_token=block_state.on_token, | |
| graph_decoder=GraphAR if block_state.use_cuda_graph else None, | |
| ) | |
| block_state.semantic_tokens = [int(t) - CODEC_OFFSET for t in tokens] | |
| block_state.semantic_truncated = truncated | |
| block_state.semantic_timing = timing | |
| self.set_block_state(state, block_state) | |
| return components, state | |
| class YuE2PrepareChunksStep(ModularPipelineBlocks): | |
| model_name = "yue2" | |
| def description(self): | |
| return "Splits the song into acoustic chunks and draws their seeded noise." | |
| def inputs(self): | |
| return [ | |
| InputParam("request", type_hint=SongRequest, required=True, description="Validated song request."), | |
| InputParam("prefix", type_hint=list, required=True, description="Semantic-stage prompt token IDs."), | |
| InputParam("semantic_tokens", type_hint=list, required=True, description="Codec token IDs."), | |
| InputParam( | |
| "acoustic_context", | |
| type_hint=int, | |
| default=CONTEXT, | |
| description="Context limit that sets the acoustic chunk size.", | |
| ), | |
| ] | |
| def intermediate_outputs(self): | |
| return [OutputParam("chunks", type_hint=list, description="Acoustic chunk tokens and CPU FP32 noise.")] | |
| def __call__(self, components, state): | |
| block_state = self.get_block_state(state) | |
| block_state.chunks = song_chunks( | |
| block_state.prefix, | |
| block_state.semantic_tokens, | |
| block_state.request.seed, | |
| block_state.acoustic_context, | |
| ) | |
| self.set_block_state(state, block_state) | |
| return components, state | |
| class YuE2ChunkConditionStep(ModularPipelineBlocks): | |
| model_name = "yue2" | |
| def description(self): | |
| return "Runs the chunk's tokens through the token stream once and keeps their keys and values." | |
| def expected_components(self): | |
| return [ComponentSpec("transformer", ModelMixin)] | |
| def inputs(self): | |
| return [ | |
| InputParam("chunks", type_hint=list, required=True, description="Prepared acoustic chunks."), | |
| InputParam("cancelled", type_hint=object, default=None, description="Cancellation callable."), | |
| ] | |
| def __call__(self, components, block_state, k): | |
| check_cancelled(block_state.cancelled) | |
| block_state.chunk = block_state.chunks[k] | |
| block_state.kv_cache = YuE2PrefixKVCache() | |
| ids = torch.tensor([block_state.chunk.ar_tokens], device=components._execution_device) | |
| components.transformer(ids, kv_cache=block_state.kv_cache, logits_to_keep=1) | |
| return components, block_state | |
| class YuE2ChunkSynthesizeStep(ModularPipelineBlocks): | |
| model_name = "yue2" | |
| def description(self): | |
| return "Solves the chunk's flow-matching ODE with the midpoint method." | |
| def expected_components(self): | |
| return [ComponentSpec("transformer", ModelMixin)] | |
| def inputs(self): | |
| return [ | |
| InputParam("chunks", type_hint=list, required=True, description="Prepared acoustic chunks."), | |
| InputParam("ode_steps", type_hint=int, default=32, description="Midpoint steps per chunk."), | |
| InputParam("cancelled", type_hint=object, default=None, description="Cancellation callable."), | |
| InputParam( | |
| "on_progress", type_hint=object, default=None, description="Callback(stage, completed, total)." | |
| ), | |
| ] | |
| def __call__(self, components, block_state, k): | |
| report = None | |
| if block_state.on_progress is not None: | |
| def report(done, total): | |
| block_state.on_progress("synthesis", k * total + done, len(block_state.chunks) * total) | |
| block_state.chunk_latents = solve_midpoint( | |
| components.transformer, | |
| block_state.kv_cache, | |
| block_state.chunk.noise, | |
| components._execution_device, | |
| block_state.ode_steps, | |
| block_state.cancelled, | |
| report, | |
| ) | |
| return components, block_state | |
| class YuE2CollectChunkStep(ModularPipelineBlocks): | |
| model_name = "yue2" | |
| def description(self): | |
| return "Appends the chunk's latents." | |
| def __call__(self, components, block_state, k): | |
| block_state.latent_chunks.append(block_state.chunk_latents) | |
| return components, block_state | |
| class YuE2AcousticChunkLoop(LoopSequentialPipelineBlocks): | |
| model_name = "yue2" | |
| block_classes = [YuE2ChunkConditionStep, YuE2ChunkSynthesizeStep, YuE2CollectChunkStep] | |
| block_names = ["condition", "synthesize", "collect"] | |
| def description(self): | |
| return "Synthesizes acoustic latents chunk by chunk." | |
| def loop_inputs(self): | |
| return [InputParam("chunks", type_hint=list, required=True, description="Prepared acoustic chunks.")] | |
| def loop_intermediate_outputs(self): | |
| return [OutputParam("latents", type_hint=torch.Tensor, description="Acoustic latents [frames, 64], CPU FP32.")] | |
| def __call__(self, components, state): | |
| block_state = self.get_block_state(state) | |
| if not block_state.chunks: | |
| raise ValueError("At least one acoustic chunk is required") | |
| block_state.latent_chunks = [] | |
| for k in range(len(block_state.chunks)): | |
| try: | |
| components, block_state = self.loop_step(components, block_state, k=k) | |
| finally: | |
| block_state.kv_cache = None | |
| block_state.latents = torch.cat(block_state.latent_chunks, dim=0) | |
| self.set_block_state(state, block_state) | |
| return components, state | |
| class YuE2SynthesizeStep(SequentialPipelineBlocks): | |
| model_name = "yue2" | |
| block_classes = [YuE2PrepareChunksStep, YuE2AcousticChunkLoop] | |
| block_names = ["prepare_chunks", "chunk_loop"] | |
| def description(self): | |
| return "Turns codec tokens into acoustic latents." | |
| class YuE2DecodeStep(ModularPipelineBlocks): | |
| model_name = "yue2" | |
| def description(self): | |
| return ( | |
| "Decodes latents to 48 kHz stereo audio in overlapping tiles. The VAE runs in FP32 with TF32 and cuDNN " | |
| "autotuning disabled, as in the original release." | |
| ) | |
| def expected_components(self): | |
| return [ComponentSpec("vae", ModelMixin)] | |
| def inputs(self): | |
| return [ | |
| InputParam( | |
| "latents", | |
| type_hint=torch.Tensor, | |
| required=True, | |
| description="Acoustic latents, [frames, 64] or [1, 64, frames].", | |
| ), | |
| InputParam("vae_tile_frames", type_hint=int, default=1024, description="Latent frames decoded per tile."), | |
| InputParam( | |
| "vae_tile_overlap_frames", | |
| type_hint=int, | |
| default=16, | |
| description="Context frames decoded on each side of a tile and cropped away.", | |
| ), | |
| InputParam("cancelled", type_hint=object, default=None, description="Cancellation callable."), | |
| InputParam( | |
| "on_progress", type_hint=object, default=None, description="Callback(stage, completed, total)." | |
| ), | |
| ] | |
| def intermediate_outputs(self): | |
| return [ | |
| OutputParam("audios", type_hint=torch.Tensor, description="Stereo audio [1, 2, samples], CPU FP32."), | |
| OutputParam("sample_rate", type_hint=int, description="Sample rate in Hz."), | |
| ] | |
| def __call__(self, components, state): | |
| block_state = self.get_block_state(state) | |
| check_cancelled(block_state.cancelled) | |
| vae = components.vae | |
| device = components._execution_device | |
| latents = torch.as_tensor(block_state.latents, dtype=torch.float32) | |
| if latents.ndim == 2: | |
| latents = latents.T.unsqueeze(0) | |
| frames = latents.shape[-1] | |
| tile, overlap = block_state.vae_tile_frames, block_state.vae_tile_overlap_frames | |
| if tile < 1 or overlap < 0: | |
| raise ValueError("vae_tile_frames must be positive and vae_tile_overlap_frames nonnegative") | |
| pieces = [] | |
| starts = range(0, frames, tile) | |
| with torch.backends.cudnn.flags(enabled=True, benchmark=False, deterministic=True, allow_tf32=False): | |
| for index, start in enumerate(starts): | |
| check_cancelled(block_state.cancelled) | |
| end = min(frames, start + tile) | |
| left, right = max(0, start - overlap), min(frames, end + overlap) | |
| decoded = vae.decode(latents[..., left:right].to(device=device, dtype=torch.float32)).sample | |
| crop = (start - left) * vae.hop_length | |
| if end < frames: | |
| decoded = decoded[..., crop : crop + (end - start) * vae.hop_length] | |
| else: | |
| decoded = decoded[..., crop:] | |
| pieces.append(decoded.cpu()) | |
| if block_state.on_progress is not None: | |
| block_state.on_progress("decode", index + 1, len(starts)) | |
| audio = torch.cat(pieces, dim=-1) | |
| if not torch.isfinite(audio).all(): | |
| raise FloatingPointError("Decoded audio is non-finite") | |
| block_state.audios = audio.clamp(-1, 1) | |
| block_state.sample_rate = vae.config.sampling_rate | |
| self.set_block_state(state, block_state) | |
| return components, state | |
| class YuE2Blocks(SequentialPipelineBlocks): | |
| model_name = "yue2" | |
| block_names = ["prepare", "plan", "semantic", "synthesize", "decode"] | |
| block_classes = [YuE2PrepareInputsStep, YuE2PlanStep, YuE2SemanticStep, YuE2SynthesizeStep, YuE2DecodeStep] | |
| def __init__(self, components_repo=None, components_revision=None, trust_components_code=False): | |
| super().__init__() | |
| self.components_repo = components_repo | |
| self.components_revision = components_revision | |
| self.component_types = {} | |
| if components_repo is not None: | |
| if not trust_components_code: | |
| raise ValueError("Set trust_components_code=True to load code from the model repository") | |
| for name, class_name in ( | |
| ("transformer", "YuE2TransformerModel"), | |
| ("vae", "YuE2VAE"), | |
| ("tokenizer", "YuE2Tokenizer"), | |
| ): | |
| self.component_types[name] = get_class_from_dynamic_module( | |
| components_repo, | |
| module_file=f"{name}.py", | |
| class_name=class_name, | |
| revision=components_revision, | |
| trust_remote_code=True, | |
| ) | |
| def expected_components(self): | |
| specs = super().expected_components | |
| if self.components_repo is None: | |
| return specs | |
| return [ | |
| replace( | |
| spec, | |
| type_hint=self.component_types[spec.name], | |
| pretrained_model_name_or_path=str(self.components_repo), | |
| subfolder=spec.name, | |
| revision=self.components_revision, | |
| ) | |
| if spec.default_creation_method == "from_pretrained" | |
| else spec | |
| for spec in specs | |
| ] | |
| def description(self): | |
| return "YuE2 score planning, semantic generation, midpoint synthesis, and stereo decoding." | |