194 lines
7.4 KiB
Python
194 lines
7.4 KiB
Python
"""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
|