Allow multi-section YuE2 lyrics and force non-flash attention on Windows.
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
+5
-4
@@ -34,7 +34,7 @@
|
|||||||
</select>
|
</select>
|
||||||
<span class="block text-xs text-zinc-400">One lyric section. Compile off. Profile 3 is never selected automatically.</span>
|
<span class="block text-xs text-zinc-400">One lyric section. Compile off. Profile 3 is never selected automatically.</span>
|
||||||
</label>
|
</label>
|
||||||
<p v-else-if="engineFamily === 'yue2'" class="text-xs text-zinc-400">YuE2-3B · one lyric section · plan → synthesize → decode · no score editor</p>
|
<p v-else-if="engineFamily === 'yue2'" class="text-xs text-zinc-400">YuE2-3B · lyrics with section tags · plan → synthesize → decode · no score editor</p>
|
||||||
<div v-if="!advanced" class="space-y-4">
|
<div v-if="!advanced" class="space-y-4">
|
||||||
<label v-if="engineFamily === 'ace'" class="block text-sm">Steps<input v-model.number="steps" type="number" :min="stepsMin" :max="stepsMax" class="mt-2 w-full rounded-xl border border-white/15 bg-zinc-900 p-3"></label>
|
<label v-if="engineFamily === 'ace'" class="block text-sm">Steps<input v-model.number="steps" type="number" :min="stepsMin" :max="stepsMax" class="mt-2 w-full rounded-xl border border-white/15 bg-zinc-900 p-3"></label>
|
||||||
<label class="block text-sm">Seed<input v-model="seed" placeholder="random" class="mt-2 w-full rounded-xl border border-white/15 bg-zinc-900 p-3"></label>
|
<label class="block text-sm">Seed<input v-model="seed" placeholder="random" class="mt-2 w-full rounded-xl border border-white/15 bg-zinc-900 p-3"></label>
|
||||||
@@ -109,7 +109,7 @@
|
|||||||
:max="durationMax"
|
:max="durationMax"
|
||||||
step="5"
|
step="5"
|
||||||
>
|
>
|
||||||
<span class="mt-1 block text-[11px] text-zinc-500">{{ engineFamily === 'yue' || engineFamily === 'yue2' ? 'Approximate length · one lyric section' : `${durationMin}–${durationMax} seconds` }}</span>
|
<span class="mt-1 block text-[11px] text-zinc-500">{{ engineFamily === 'yue' ? 'Approximate length · one lyric section' : engineFamily === 'yue2' ? 'Approximate length · multiple sections OK' : `${durationMin}–${durationMax} seconds` }}</span>
|
||||||
</label>
|
</label>
|
||||||
<label v-if="advanced && engineFamily === 'ace'" class="block text-sm">
|
<label v-if="advanced && engineFamily === 'ace'" class="block text-sm">
|
||||||
<span class="mb-1 block font-medium text-zinc-300">Steps · {{ steps }}</span>
|
<span class="mb-1 block font-medium text-zinc-300">Steps · {{ steps }}</span>
|
||||||
@@ -333,7 +333,8 @@ const blockReason = computed(() => {
|
|||||||
if ((engineFamily.value === 'yue' || engineFamily.value === 'yue2') && instrumental.value) return 'YuE requires lyrics. Use ACE for instrumental music.'
|
if ((engineFamily.value === 'yue' || engineFamily.value === 'yue2') && instrumental.value) return 'YuE requires lyrics. Use ACE for instrumental music.'
|
||||||
if ((engineFamily.value === 'yue' || engineFamily.value === 'yue2') && duration.value > 150) return 'YuE supports up to 150 seconds per section.'
|
if ((engineFamily.value === 'yue' || engineFamily.value === 'yue2') && duration.value > 150) return 'YuE supports up to 150 seconds per section.'
|
||||||
if (!tags.value.trim()) return 'Add genre and style tags.'
|
if (!tags.value.trim()) return 'Add genre and style tags.'
|
||||||
if (engineFamily.value === 'yue' || engineFamily.value === 'yue2') { const problem = yueLyricsProblem(lyrics.value); if (problem) return problem }
|
if (engineFamily.value === 'yue') { const problem = yueLyricsProblem(lyrics.value); if (problem) return problem }
|
||||||
|
if (engineFamily.value === 'yue2' && !lyrics.value.trim()) return 'Write lyrics for YuE2.'
|
||||||
if (!instrumental.value && !lyrics.value.trim()) return 'Write lyrics, or turn on Instrumental.'
|
if (!instrumental.value && !lyrics.value.trim()) return 'Write lyrics, or turn on Instrumental.'
|
||||||
return ''
|
return ''
|
||||||
})
|
})
|
||||||
@@ -341,7 +342,7 @@ const blockReason = computed(() => {
|
|||||||
function selectEngineFamily(family: 'ace' | 'yue' | 'yue2') {
|
function selectEngineFamily(family: 'ace' | 'yue' | 'yue2') {
|
||||||
engineFamily.value = family
|
engineFamily.value = family
|
||||||
if (family === 'yue' || family === 'yue2') { instrumental.value = false; duration.value = Math.min(150, duration.value) }
|
if (family === 'yue' || family === 'yue2') { instrumental.value = false; duration.value = Math.min(150, duration.value) }
|
||||||
if ((family === 'yue' || family === 'yue2') && lyrics.value === DEFAULT_MUSIC_LYRICS) lyrics.value = '[Verse 1]\n'
|
if (family === 'yue' && lyrics.value === DEFAULT_MUSIC_LYRICS) lyrics.value = '[Verse 1]\n'
|
||||||
if (family === 'ace' && ace15.value && steps.value === MUSIC_STEPS_DEFAULT) {
|
if (family === 'ace' && ace15.value && steps.value === MUSIC_STEPS_DEFAULT) {
|
||||||
steps.value = MUSIC_STEPS_DEFAULT_15
|
steps.value = MUSIC_STEPS_DEFAULT_15
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -10,9 +10,7 @@ export function validateYue2Request(body) {
|
|||||||
if (!/^[a-zA-Z0-9-]{12,80}$/.test(body.id || '')) throw new Error('Invalid job ID.')
|
if (!/^[a-zA-Z0-9-]{12,80}$/.test(body.id || '')) throw new Error('Invalid job ID.')
|
||||||
const tags = String(body.tags || '').trim()
|
const tags = String(body.tags || '').trim()
|
||||||
const lyrics = String(body.lyrics || '').trim()
|
const lyrics = String(body.lyrics || '').trim()
|
||||||
if (!tags || tags.length > 2000 || !lyrics || lyrics.length > 8000) throw new Error('Genre tags and one lyric section are required.')
|
if (!tags || tags.length > 2000 || !lyrics || lyrics.length > 8000) throw new Error('Genre tags and non-empty lyrics are required.')
|
||||||
const sections = [...lyrics.matchAll(/\[([^\]]+)\]\s*([^\[]*)/gs)].filter(m => m[2].trim())
|
|
||||||
if (sections.length !== 1) throw new Error('YuE2 requires one non-empty lyric section.')
|
|
||||||
return { id: body.id, duration, seed: body.seed, tags, lyrics }
|
return { id: body.id, duration, seed: body.seed, tags, lyrics }
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+45
-9
@@ -18,14 +18,47 @@ def emit(**event):
|
|||||||
|
|
||||||
|
|
||||||
def normalize_lyrics(text):
|
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
|
import re
|
||||||
sections = re.findall(r'\[([^\]]+)\]\s*([^\[]*)', text.strip(), re.S)
|
text = str(text or '').strip()
|
||||||
sections = [(name, words.strip()) for name, words in sections if words.strip()]
|
if not text:
|
||||||
if len(sections) != 1:
|
raise ValueError('YuE2 requires non-empty lyrics.')
|
||||||
raise ValueError('YuE2 requires one non-empty lyric section. Combine the lyrics under one heading.')
|
if re.search(r'\[[^\]]+\]', text):
|
||||||
name, words = sections[0]
|
return text if text.endswith('\n') else text + '\n'
|
||||||
return '[' + re.sub(r'\W+', '', name).lower() + ']\n' + words + '\n\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():
|
def main():
|
||||||
@@ -56,16 +89,19 @@ def main():
|
|||||||
os.chdir(root)
|
os.chdir(root)
|
||||||
if str(root) not in sys.path:
|
if str(root) not in sys.path:
|
||||||
sys.path.insert(0, str(root))
|
sys.path.insert(0, str(root))
|
||||||
emit(stage='loading', progress=1, message='Loading YuE2', model=model, vae=vae, duration=duration)
|
|
||||||
import torch
|
import torch
|
||||||
import soundfile as sf
|
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
|
from yue2 import YuE2Pipeline
|
||||||
if not torch.cuda.is_available():
|
if not torch.cuda.is_available():
|
||||||
raise RuntimeError('YuE2 requires a CUDA GPU; CPU fallback is disabled.')
|
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)
|
pipe_kwargs = dict(style=style, lyrics=lyrics, cot='full', seed=seed)
|
||||||
# Duration is kept for library metadata and validation. Upstream one-shot
|
# Duration is kept for library metadata and validation. Upstream one-shot
|
||||||
# requests do not take a seconds field; song length follows the plan.
|
# 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)
|
emit(stage='plan', message='Planning melody and chords', progress=5)
|
||||||
plan = pipe.plan(**pipe_kwargs)
|
plan = pipe.plan(**pipe_kwargs)
|
||||||
emit(stage='semantic', message='Generating semantic tokens', progress=25)
|
emit(stage='semantic', message='Generating semantic tokens', progress=25)
|
||||||
|
|||||||
@@ -71,10 +71,13 @@ export default defineEventHandler(async (event) => {
|
|||||||
if ((engine === 'yue' || engine === 'yue2') && instrumental) throw createError({ statusCode: 400, statusMessage: 'YuE currently requires lyrics. Use ACE for instrumental music.' })
|
if ((engine === 'yue' || engine === 'yue2') && instrumental) throw createError({ statusCode: 400, statusMessage: 'YuE currently requires lyrics. Use ACE for instrumental music.' })
|
||||||
const yueProfile = body.yueProfile ?? 1
|
const yueProfile = body.yueProfile ?? 1
|
||||||
if (engine === 'yue' && yueProfile !== 1 && yueProfile !== 3) throw createError({ statusCode: 400, statusMessage: 'Choose YuEGP profile 1 or manual fallback 3.' })
|
if (engine === 'yue' && yueProfile !== 1 && yueProfile !== 3) throw createError({ statusCode: 400, statusMessage: 'Choose YuEGP profile 1 or manual fallback 3.' })
|
||||||
if ((engine === 'yue' || engine === 'yue2') && !instrumental) {
|
if (engine === 'yue' && !instrumental) {
|
||||||
const problem = yueLyricsProblem(lyrics)
|
const problem = yueLyricsProblem(lyrics)
|
||||||
if (problem) throw createError({ statusCode: 400, statusMessage: problem })
|
if (problem) throw createError({ statusCode: 400, statusMessage: problem })
|
||||||
}
|
}
|
||||||
|
if (engine === 'yue2' && !instrumental && !lyrics) {
|
||||||
|
throw createError({ statusCode: 400, statusMessage: 'YuE2 needs lyrics.' })
|
||||||
|
}
|
||||||
const duration = clampMusicDuration(body.duration)
|
const duration = clampMusicDuration(body.duration)
|
||||||
if ((engine === 'yue' || engine === 'yue2') && duration > 150) throw createError({ statusCode: 400, statusMessage: 'YuE supports up to 150 seconds per section.' })
|
if ((engine === 'yue' || engine === 'yue2') && duration > 150) throw createError({ statusCode: 400, statusMessage: 'YuE supports up to 150 seconds per section.' })
|
||||||
const steps = clampMusicSteps(
|
const steps = clampMusicSteps(
|
||||||
|
|||||||
@@ -2,6 +2,7 @@
|
|||||||
import importlib.util
|
import importlib.util
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
import unittest
|
import unittest
|
||||||
|
from types import SimpleNamespace
|
||||||
|
|
||||||
spec = importlib.util.spec_from_file_location('yue2_worker', Path(__file__).parents[1] / 'scripts/yue2-worker.py')
|
spec = importlib.util.spec_from_file_location('yue2_worker', Path(__file__).parents[1] / 'scripts/yue2-worker.py')
|
||||||
worker = importlib.util.module_from_spec(spec)
|
worker = importlib.util.module_from_spec(spec)
|
||||||
@@ -9,12 +10,29 @@ spec.loader.exec_module(worker)
|
|||||||
|
|
||||||
|
|
||||||
class WorkerTests(unittest.TestCase):
|
class WorkerTests(unittest.TestCase):
|
||||||
def test_lyrics_preserve_words_and_normalize_ui_headings(self):
|
def test_lyrics_allow_multiple_sections_and_wrap_plain_text(self):
|
||||||
self.assertEqual(worker.normalize_lyrics('[Pre-Chorus]\nEvery word stays\n[Outro]\n'), '[prechorus]\nEvery word stays\n\n')
|
multi = '[Verse 1]\nHello\n\n[Chorus]\nSing it'
|
||||||
self.assertEqual(worker.normalize_lyrics('[Verse 1]\nHello'), '[verse1]\nHello\n\n')
|
self.assertEqual(worker.normalize_lyrics(multi), multi + '\n')
|
||||||
for lyrics in ['', 'No heading', '[Verse]\nA\n[Chorus]\nB']:
|
self.assertEqual(worker.normalize_lyrics('[Verse 1]\nHello'), '[Verse 1]\nHello\n')
|
||||||
|
self.assertEqual(worker.normalize_lyrics('No heading yet'), '[song]\nNo heading yet\n\n')
|
||||||
with self.assertRaises(ValueError):
|
with self.assertRaises(ValueError):
|
||||||
worker.normalize_lyrics(lyrics)
|
worker.normalize_lyrics('')
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
worker.normalize_lyrics(' ')
|
||||||
|
|
||||||
|
def test_attention_backend_never_selects_flash(self):
|
||||||
|
torch_mod = SimpleNamespace(
|
||||||
|
cuda=SimpleNamespace(is_available=lambda: True),
|
||||||
|
backends=SimpleNamespace(cudnn=SimpleNamespace(is_available=lambda: True)),
|
||||||
|
)
|
||||||
|
attention, pipeline = worker.resolve_attention_backend(torch_mod)
|
||||||
|
self.assertEqual(attention, 'cudnn')
|
||||||
|
self.assertEqual(pipeline, 'torch')
|
||||||
|
self.assertNotEqual(attention, 'flash')
|
||||||
|
torch_mod.backends.cudnn.is_available = lambda: False
|
||||||
|
attention, pipeline = worker.resolve_attention_backend(torch_mod)
|
||||||
|
self.assertEqual(attention, 'sdpa')
|
||||||
|
self.assertEqual(pipeline, 'torch-eager')
|
||||||
|
|
||||||
|
|
||||||
if __name__ == '__main__':
|
if __name__ == '__main__':
|
||||||
|
|||||||
+6
-2
@@ -11,9 +11,13 @@ import { createYueGpHost, validateYueGpRequest } from '../scripts/yuegp-host.mjs
|
|||||||
|
|
||||||
const request = { id: 'test-song-123456', tags: 'pop, warm vocals', lyrics: '[Verse 1]\nA quiet morning, a new day', seed: 0 }
|
const request = { id: 'test-song-123456', tags: 'pop, warm vocals', lyrics: '[Verse 1]\nA quiet morning, a new day', seed: 0 }
|
||||||
|
|
||||||
test('YuE2 defaults are 60 seconds and keep Yue lyric validation', () => {
|
test('YuE2 defaults are 60 seconds and allow multiple lyric sections', () => {
|
||||||
assert.deepEqual(validateYue2Request(request), { ...request, duration: 60 })
|
assert.deepEqual(validateYue2Request(request), { ...request, duration: 60 })
|
||||||
assert.throws(() => validateYue2Request({ ...request, lyrics: '[Verse]\nA\n[Chorus]\nB' }))
|
assert.deepEqual(
|
||||||
|
validateYue2Request({ ...request, lyrics: '[Verse]\nA\n[Chorus]\nB' }).lyrics,
|
||||||
|
'[Verse]\nA\n[Chorus]\nB'
|
||||||
|
)
|
||||||
|
assert.throws(() => validateYue2Request({ ...request, lyrics: '' }))
|
||||||
assert.throws(() => validateYue2Request({ ...request, duration: 20 }))
|
assert.throws(() => validateYue2Request({ ...request, duration: 20 }))
|
||||||
assert.throws(() => validateYue2Request({ ...request, id: '../escape' }))
|
assert.throws(() => validateYue2Request({ ...request, id: '../escape' }))
|
||||||
})
|
})
|
||||||
|
|||||||
Reference in New Issue
Block a user