Allow multi-section YuE2 lyrics and force non-flash attention on Windows.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Towsty
2026-09-14 19:28:43 -05:00
co-authored by Cursor
parent fae3b23e37
commit 62671d05da
6 changed files with 85 additions and 25 deletions
+45 -9
View File
@@ -18,14 +18,47 @@ def emit(**event):
def normalize_lyrics(text):
# Same one-section rule as YuEGP field validation.
# YuE2 accepts multiple [Verse]/[Chorus] headings. Only wrap when none exist.
import re
sections = re.findall(r'\[([^\]]+)\]\s*([^\[]*)', text.strip(), re.S)
sections = [(name, words.strip()) for name, words in sections if words.strip()]
if len(sections) != 1:
raise ValueError('YuE2 requires one non-empty lyric section. Combine the lyrics under one heading.')
name, words = sections[0]
return '[' + re.sub(r'\W+', '', name).lower() + ']\n' + words + '\n\n'
text = str(text or '').strip()
if not text:
raise ValueError('YuE2 requires non-empty lyrics.')
if re.search(r'\[[^\]]+\]', text):
return text if text.endswith('\n') else text + '\n'
return '[song]\n' + text + '\n\n'
def resolve_attention_backend(torch_mod):
"""Pick a non-flash attention path for the Windows CUDA wheel on the 5080."""
flash = False
try:
check = getattr(torch_mod.backends.cuda, 'is_flash_attention_available', None)
flash = bool(check()) if callable(check) else False
except Exception:
flash = False
# Cognito/YuE2-Windows: when flash is unavailable, fall back to cudnn.
# This host never enables flash: the torch wheel exposes the op without USE_FLASH_ATTENTION.
if sys.platform == 'win32' or not flash:
if torch_mod.cuda.is_available() and torch_mod.backends.cudnn.is_available():
return 'cudnn', 'torch'
return 'sdpa', 'torch-eager'
if torch_mod.cuda.is_available() and torch_mod.backends.cudnn.is_available():
return 'cudnn', 'torch'
return 'sdpa', 'torch-eager'
def patch_graph_attention(attention_backend):
from yue2.cuda_graph import GraphAR
original = GraphAR.__init__
def init(self, model, prefixes, max_tokens, *, capture=True, attention_backend='auto', fuse_projections=False):
if attention_backend in ('auto', 'flash'):
attention_backend = patch_graph_attention.forced
return original(self, model, prefixes, max_tokens, capture=capture,
attention_backend=attention_backend, fuse_projections=fuse_projections)
patch_graph_attention.forced = attention_backend
GraphAR.__init__ = init
def main():
@@ -56,16 +89,19 @@ def main():
os.chdir(root)
if str(root) not in sys.path:
sys.path.insert(0, str(root))
emit(stage='loading', progress=1, message='Loading YuE2', model=model, vae=vae, duration=duration)
import torch
import soundfile as sf
attention_backend, pipeline_backend = resolve_attention_backend(torch)
emit(stage='loading', progress=1, message='Loading YuE2', model=model, vae=vae,
duration=duration, attention_backend=attention_backend, pipeline_backend=pipeline_backend)
from yue2 import YuE2Pipeline
if not torch.cuda.is_available():
raise RuntimeError('YuE2 requires a CUDA GPU; CPU fallback is disabled.')
patch_graph_attention(attention_backend)
pipe_kwargs = dict(style=style, lyrics=lyrics, cot='full', seed=seed)
# Duration is kept for library metadata and validation. Upstream one-shot
# requests do not take a seconds field; song length follows the plan.
with YuE2Pipeline.from_pretrained(model, vae=vae, device='cuda') as pipe:
with YuE2Pipeline.from_pretrained(model, vae=vae, device='cuda', backend=pipeline_backend) as pipe:
emit(stage='plan', message='Planning melody and chords', progress=5)
plan = pipe.plan(**pipe_kwargs)
emit(stage='semantic', message='Generating semantic tokens', progress=25)