Tighten YuE2 VRAM use on the 5080 and empty GPU before launch.

Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
Towsty
2026-09-14 20:01:06 -05:00
co-authored by Cursor
parent 62671d05da
commit ee84a99035
3 changed files with 89 additions and 37 deletions
+3
View File
@@ -856,6 +856,9 @@ const yue2 = createYue2Host({
if (healthy) { if (healthy) {
const queue = await fetchLocalQueue(healthy) const queue = await fetchLocalQueue(healthy)
if (!queue.ok || queue.running || queue.pending) throw new Error('Comfy is busy; YuE2 cannot start.') if (!queue.ok || queue.running || queue.pending) throw new Error('Comfy is busy; YuE2 cannot start.')
}
// Same gate as YuEGP: stop Comfy and refuse to launch while its Python still owns VRAM.
if (healthy || await processUp() || await pythonMainUp().catch(() => false)) {
await stopComfyProcesses() await stopComfyProcesses()
markAsleep() markAsleep()
} }
+79 -32
View File
@@ -1,6 +1,6 @@
"""Headless YuE2 adapter: plan → generate_semantic → synthesize → decode. """Headless YuE2 adapter: plan → generate_semantic → synthesize → decode.
No Comfy imports. No score editor, covers, or auto-fallback to YuEGP. No Comfy imports. No score editor, covers, or auto-fallback to YuEGP / ACE.
Stdout AIGEN_EVENT lines are consumed by the device-local host agent. Stdout AIGEN_EVENT lines are consumed by the device-local host agent.
""" """
import argparse import argparse
@@ -12,6 +12,9 @@ import sys
import time import time
import threading import threading
# Must be set before the first torch import on this 16GB Windows host.
os.environ.setdefault('PYTORCH_CUDA_ALLOC_CONF', 'expandable_segments:True')
def emit(**event): def emit(**event):
print('AIGEN_EVENT ' + json.dumps(event), flush=True) print('AIGEN_EVENT ' + json.dumps(event), flush=True)
@@ -28,22 +31,15 @@ def normalize_lyrics(text):
return '[song]\n' + text + '\n\n' return '[song]\n' + text + '\n\n'
def is_cuda_oom(error):
message = str(error).lower()
return 'out of memory' in message or ('cuda' in message and 'alloc' in message) or 'cudnn_status_alloc_failed' in message
def resolve_attention_backend(torch_mod): def resolve_attention_backend(torch_mod):
"""Pick a non-flash attention path for the Windows CUDA wheel on the 5080.""" """Never flash. Prefer torch-eager (no CUDA graphs) on the 16GB 5080."""
flash = False # Windows wheels expose flash ops without USE_FLASH_ATTENTION; never select flash.
try: # torch-eager disables GraphAR / CUDA graphs via YuE2Pipeline.backend.
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' return 'sdpa', 'torch-eager'
@@ -61,6 +57,28 @@ def patch_graph_attention(attention_backend):
GraphAR.__init__ = init GraphAR.__init__ = init
def free_cuda(torch_mod):
gc.collect()
if torch_mod.cuda.is_available():
torch_mod.cuda.empty_cache()
torch_mod.cuda.synchronize()
def memory_snapshot(torch_mod, stage):
if not torch_mod.cuda.is_available():
emit(stage=stage, message=f'YuE2 {stage}', cuda=False)
return
free, total = torch_mod.cuda.mem_get_info()
emit(
stage=stage,
message=f'YuE2 {stage}',
gpuAllocatedMiB=round(torch_mod.cuda.memory_allocated() / 1024 ** 2),
gpuReservedMiB=round(torch_mod.cuda.memory_reserved() / 1024 ** 2),
gpuFreeMiB=round(free / 1024 ** 2),
gpuTotalMiB=round(total / 1024 ** 2),
)
def main(): def main():
import psutil import psutil
parent = psutil.Process(os.getppid()) parent = psutil.Process(os.getppid())
@@ -93,27 +111,56 @@ def main():
import soundfile as sf import soundfile as sf
attention_backend, pipeline_backend = resolve_attention_backend(torch) attention_backend, pipeline_backend = resolve_attention_backend(torch)
emit(stage='loading', progress=1, message='Loading YuE2', model=model, vae=vae, emit(stage='loading', progress=1, message='Loading YuE2', model=model, vae=vae,
duration=duration, attention_backend=attention_backend, pipeline_backend=pipeline_backend) duration=duration, attention_backend=attention_backend, pipeline_backend=pipeline_backend,
cudaAllocConf=os.environ.get('PYTORCH_CUDA_ALLOC_CONF'))
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) patch_graph_attention(attention_backend)
pipe_kwargs = dict(style=style, lyrics=lyrics, cot='full', seed=seed) free_cuda(torch)
# Duration is kept for library metadata and validation. Upstream one-shot memory_snapshot(torch, 'memory-before-load')
# requests do not take a seconds field; song length follows the plan. # 16GB card: cap budget so YuE2 leaves headroom; tiled VAE at 512 frames.
with YuE2Pipeline.from_pretrained(model, vae=vae, device='cuda', backend=pipeline_backend) as pipe: pipe_load = dict(
emit(stage='plan', message='Planning melody and chords', progress=5) device='cuda',
plan = pipe.plan(**pipe_kwargs) backend=pipeline_backend,
emit(stage='semantic', message='Generating semantic tokens', progress=25) memory_budget_gib=12,
vae_core_frames=512,
offload_ar=True,
progress=False,
)
cot = 'full'
pipe_kwargs = dict(style=style, lyrics=lyrics, cot=cot, seed=seed)
with YuE2Pipeline.from_pretrained(model, vae=vae, **pipe_load) as pipe:
memory_snapshot(torch, 'memory-loaded')
emit(stage='plan', message='Planning melody and chords', progress=5, cot=cot)
try:
plan = pipe.plan(**pipe_kwargs)
except Exception as error:
if cot != 'full' or not is_cuda_oom(error):
raise
free_cuda(torch)
cot = 'melody'
pipe_kwargs['cot'] = cot
emit(stage='plan', message='Full CoT OOM; retrying melody CoT', progress=5, cot=cot,
error=str(error))
memory_snapshot(torch, 'memory-before-melody-plan')
plan = pipe.plan(**pipe_kwargs)
free_cuda(torch)
memory_snapshot(torch, 'memory-after-plan')
emit(stage='semantic', message='Generating semantic tokens', progress=25, cot=cot)
semantic = pipe.generate_semantic(plan) semantic = pipe.generate_semantic(plan)
emit(stage='synthesize', message='Synthesizing acoustic latents', progress=55) free_cuda(torch)
memory_snapshot(torch, 'memory-after-semantic')
emit(stage='synthesize', message='Synthesizing acoustic latents', progress=55, cot=cot)
latents = pipe.synthesize(semantic) latents = pipe.synthesize(semantic)
emit(stage='decode', message='Decoding audio', progress=80) free_cuda(torch)
audio = pipe.decode(latents) memory_snapshot(torch, 'memory-after-synthesize')
# Context exit unloads the pipeline. Clear any residual CUDA cache. emit(stage='decode', message='Decoding audio (tiled)', progress=80, cot=cot)
gc.collect() # full=False uses decode_tiled with vae_core_frames / halo; model offloads to CPU after.
if torch.cuda.is_available(): audio = pipe.decode(latents, full=False)
torch.cuda.empty_cache() free_cuda(torch)
memory_snapshot(torch, 'memory-after-decode')
free_cuda(torch)
wave = audio wave = audio
if hasattr(audio, 'detach'): if hasattr(audio, 'detach'):
wave = audio.detach().cpu().numpy() wave = audio.detach().cpu().numpy()
@@ -129,7 +176,7 @@ def main():
info = sf.info(str(target)) info = sf.info(str(target))
if info.frames <= 0: if info.frames <= 0:
raise RuntimeError('YuE2 produced empty audio.') raise RuntimeError('YuE2 produced empty audio.')
emit(stage='complete', message='Audio ready', progress=100, duration=info.duration) emit(stage='complete', message='Audio ready', progress=100, duration=info.duration, cot=cot)
if __name__ == '__main__': if __name__ == '__main__':
+7 -5
View File
@@ -1,5 +1,6 @@
"""CPU-only YuE2 worker regression tests; never import torch or run inference.""" """CPU-only YuE2 worker regression tests; never import torch or run inference."""
import importlib.util import importlib.util
import os
from pathlib import Path from pathlib import Path
import unittest import unittest
from types import SimpleNamespace from types import SimpleNamespace
@@ -26,13 +27,14 @@ class WorkerTests(unittest.TestCase):
backends=SimpleNamespace(cudnn=SimpleNamespace(is_available=lambda: True)), backends=SimpleNamespace(cudnn=SimpleNamespace(is_available=lambda: True)),
) )
attention, pipeline = worker.resolve_attention_backend(torch_mod) 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(attention, 'sdpa')
self.assertEqual(pipeline, 'torch-eager') self.assertEqual(pipeline, 'torch-eager')
self.assertNotEqual(attention, 'flash')
self.assertTrue(os.environ.get('PYTORCH_CUDA_ALLOC_CONF', '').startswith('expandable_segments'))
def test_cuda_oom_detection(self):
self.assertTrue(worker.is_cuda_oom(RuntimeError('CUDA out of memory. Tried to allocate 2.49 GiB')))
self.assertFalse(worker.is_cuda_oom(RuntimeError('bad lyrics')))
if __name__ == '__main__': if __name__ == '__main__':