Files
aigen/scripts/yue2-worker.py
T

141 lines
5.5 KiB
Python
Raw 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.
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
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 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():
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)
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)
pipe_kwargs = dict(style=style, lyrics=lyrics, cot='full', seed=seed)
# Duration is kept for library metadata and validation. Upstream one-shot
# requests do not take a seconds field; song length follows the plan.
with YuE2Pipeline.from_pretrained(model, vae=vae, device='cuda', backend=pipeline_backend) as pipe:
emit(stage='plan', message='Planning melody and chords', progress=5)
plan = pipe.plan(**pipe_kwargs)
emit(stage='semantic', message='Generating semantic tokens', progress=25)
semantic = pipe.generate_semantic(plan)
emit(stage='synthesize', message='Synthesizing acoustic latents', progress=55)
latents = pipe.synthesize(semantic)
emit(stage='decode', message='Decoding audio', progress=80)
audio = pipe.decode(latents)
# Context exit unloads the pipeline. Clear any residual CUDA cache.
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
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)
if __name__ == '__main__':
try:
main()
except Exception as error:
emit(stage='error', message=str(error), error=str(error))
raise