Files
aigen/scripts/yue2-worker.py

201 lines
7.8 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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 15 <= duration <= 150:
raise ValueError('Target length must be 15–150 seconds.')
# VAE downsampling_ratio 1920 @ 48 kHz → 25 semantic / latent frames per second.
max_tokens = max(200, min(9000, duration * 25))
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, maxTokens=max_tokens, attention_backend=attention_backend,
pipeline_backend=pipeline_backend,
cudaAllocConf=os.environ.get('PYTORCH_CUDA_ALLOC_CONF'))
from yue2 import YuE2Pipeline
from yue2.protocol import Sampling
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)
# Cap semantic length from target seconds. Do not FFmpeg-trim after decode.
semantic_sampling = Sampling(max_tokens=max_tokens, min_tokens=min(200, max_tokens))
# 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,
maxTokens=max_tokens, targetSeconds=duration)
semantic = pipe.generate_semantic(plan, sampling=semantic_sampling)
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