# 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 asdict, dataclass EOD = 151643 ABC_START, ABC_END = (151847, 151848) MUSIC_START, MUSIC_END = (151851, 151852) CODEC_OFFSET, CODEC_SIZE = (151853, 32768) CONTEXT = 24576 INSTRUCTIONS = { "off": "Generate music with codec tokens from the given conditions.", "melody": "Generate a melody-only ABC transcription without chord symbols, then generate music with codec tokens from the given conditions.", "full": "Generate a chord-annotated ABC transcription, then generate music with codec tokens from the given conditions.", } @dataclass(frozen=True) class Sampling: temperature: float = 1.0 top_p: float = 0.95 top_k: int = 100 repetition_penalty: float = 1.2 penalty_window: int = 50 min_tokens: int = 200 max_tokens: int = 9000 def __post_init__(self): if any(type(x) is not int for x in (self.top_k, self.penalty_window, self.min_tokens, self.max_tokens)): raise ValueError("Sampling counts must be integers") if not all(math.isfinite(x) for x in (self.temperature, self.top_p, self.repetition_penalty)): raise ValueError("Sampling numbers must be finite") if not 0 <= self.temperature <= 5 or not 0 < self.top_p <= 1 or self.top_k < 1: raise ValueError("Invalid sampling temperature/top_p/top_k") if self.repetition_penalty <= 0 or not 1 <= self.penalty_window <= 100: raise ValueError("Invalid repetition penalty/window") if not 0 <= self.min_tokens <= self.max_tokens or self.max_tokens < 1: raise ValueError("Require 0 <= min_tokens <= max_tokens") ABC_SAMPLING = Sampling(0.7, 0.9, 30, 1.005, 100, 32, 4096) SEMANTIC_SAMPLING = Sampling() def resolve_sampling(overrides, default): if overrides is None: return default if not isinstance(overrides, dict): raise TypeError("Sampling overrides must be a dictionary") return Sampling(**{**asdict(default), **overrides}) @dataclass(frozen=True) class SongRequest: style: str lyrics: str cot: str = "full" seed: int = 831001 abc: str | None = None def __post_init__(self): if self.cot not in INSTRUCTIONS: raise ValueError("cot must be off, melody or full") if not isinstance(self.style, str) or not isinstance(self.lyrics, str): raise TypeError("style and lyrics must be strings") if type(self.seed) is not int or not 0 <= self.seed < 2**63: raise ValueError("seed must be an integer in [0, 2**63)") if self.abc is not None and (self.cot == "off" or not isinstance(self.abc, str) or not self.abc.strip()): raise ValueError("External ABC requires nonempty text and cot=melody/full") def text(self): return f"{INSTRUCTIONS[self.cot]}\n[Tags]\n{self.style}\n[Lyrics]\n{self.lyrics}\n" def _check_abc_ids(abc_ids): abc_ids = list(abc_ids) if any(type(token) is not int or not 0 <= token < EOD for token in abc_ids): raise ValueError("ABC IDs must remain inside the ordinary text vocabulary") return abc_ids def token_prefixes(request, tokenizer, abc_ids=None): base = [EOD] + tokenizer.encode(request.text()) if request.cot == "off": return base + [ABC_START, ABC_END, MUSIC_START] if abc_ids is None: if request.abc is None: return base + [ABC_START] abc_ids = tokenizer.encode(request.abc) return base + [ABC_START] + _check_abc_ids(abc_ids) + [ABC_END, MUSIC_START] def negative_prefix(request, tokenizer, abc_ids=None): base = [EOD] + tokenizer.encode(INSTRUCTIONS[request.cot]) if request.cot == "off": return base + [MUSIC_START] if abc_ids is None: raise ValueError("Symbolic CFG must retain the exact positive-branch ABC IDs") return base + [ABC_START] + _check_abc_ids(abc_ids) + [ABC_END, MUSIC_START] def chunk_ranges(frames, prefix_tokens, context=CONTEXT): size = min((context - prefix_tokens - 3) // 2, CONTEXT) if frames < 1 or size < 1: raise ValueError("Empty codec or prefix leaves no acoustic context") return [(a, min(a + size, frames)) for a in range(0, frames, size)]