"""Headless YuE2 adapter: plan → generate_semantic → synthesize → decode. 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. """ import argparse import gc import json import os from pathlib import Path import sys import time 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): print('AIGEN_EVENT ' + json.dumps(event), flush=True) def normalize_lyrics(text): # YuE2 accepts multiple [Verse]/[Chorus] headings. Only wrap when none exist. import re 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 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): """Never flash. Prefer torch-eager (no CUDA graphs) on the 16GB 5080.""" # Windows wheels expose flash ops without USE_FLASH_ATTENTION; never select flash. # torch-eager disables GraphAR / CUDA graphs via YuE2Pipeline.backend. 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 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(): import psutil parent = psutil.Process(os.getppid()) def parent_watchdog(): while parent.is_running(): time.sleep(2) os._exit(2) # The host died: never leave an orphan consuming the GPU. threading.Thread(target=parent_watchdog, daemon=True).start() parser = argparse.ArgumentParser() parser.add_argument('--root', required=True) parser.add_argument('--request', required=True) cli = parser.parse_args() request = json.loads(Path(cli.request).read_text(encoding='utf-8-sig')) root = Path(cli.root).resolve() output = Path(cli.request).resolve().parent duration = int(request.get('duration', 60)) if not 30 <= duration <= 150: raise ValueError('Duration must be 30–150 seconds.') lyrics = normalize_lyrics(request['lyrics']) style = ' '.join(str(request['tags']).split()) seed = int(request['seed']) model = request.get('model') or os.environ.get('YUE2_MODEL') or 'm-a-p/YuE2-3B' vae = request.get('vae') or os.environ.get('YUE2_VAE') or 'm-a-p/YuE2-Vae' os.chdir(root) if str(root) not in sys.path: sys.path.insert(0, str(root)) 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, cudaAllocConf=os.environ.get('PYTORCH_CUDA_ALLOC_CONF')) 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) free_cuda(torch) memory_snapshot(torch, 'memory-before-load') # Prefer tiled VAE; do not cap the process with set_per_process_memory_fraction. pipe_load = dict( device='cuda', backend=pipeline_backend, vae_core_frames=512, offload_ar=True, progress=False, ) cot = 'full' pipe_kwargs = dict(style=style, lyrics=lyrics, cot=cot, seed=seed) # YuE2Pipeline.__init__ always calls set_per_process_memory_fraction; skip it on this 16GB host. _set_fraction = torch.cuda.set_per_process_memory_fraction torch.cuda.set_per_process_memory_fraction = lambda *args, **kwargs: None try: pipe_cm = YuE2Pipeline.from_pretrained(model, vae=vae, **pipe_load) finally: torch.cuda.set_per_process_memory_fraction = _set_fraction with pipe_cm 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) free_cuda(torch) memory_snapshot(torch, 'memory-after-semantic') emit(stage='synthesize', message='Synthesizing acoustic latents', progress=55, cot=cot) latents = pipe.synthesize(semantic) free_cuda(torch) memory_snapshot(torch, 'memory-after-synthesize') emit(stage='decode', message='Decoding audio (tiled)', progress=80, cot=cot) # full=False uses decode_tiled with vae_core_frames / halo; model offloads to CPU after. audio = pipe.decode(latents, full=False) free_cuda(torch) memory_snapshot(torch, 'memory-after-decode') free_cuda(torch) wave = audio if hasattr(audio, 'detach'): wave = audio.detach().cpu().numpy() import numpy as np wave = np.asarray(wave) if wave.ndim == 1: pass elif wave.shape[0] <= 8 and wave.shape[0] < wave.shape[-1]: wave = wave.T sample_rate = 48000 target = output / 'audio.wav' sf.write(str(target), wave, sample_rate, subtype='PCM_16') info = sf.info(str(target)) if info.frames <= 0: raise RuntimeError('YuE2 produced empty audio.') emit(stage='complete', message='Audio ready', progress=100, duration=info.duration, cot=cot) if __name__ == '__main__': try: main() except Exception as error: emit(stage='error', message=str(error), error=str(error)) raise