Add YuE2 host, worker, and setup scripts.
Co-authored-by: Cursor <cursoragent@cursor.com>
This commit is contained in:
@@ -0,0 +1,104 @@
|
||||
"""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):
|
||||
# Same one-section rule as YuEGP field validation.
|
||||
import re
|
||||
sections = re.findall(r'\[([^\]]+)\]\s*([^\[]*)', text.strip(), re.S)
|
||||
sections = [(name, words.strip()) for name, words in sections if words.strip()]
|
||||
if len(sections) != 1:
|
||||
raise ValueError('YuE2 requires one non-empty lyric section. Combine the lyrics under one heading.')
|
||||
name, words = sections[0]
|
||||
return '[' + re.sub(r'\W+', '', name).lower() + ']\n' + words + '\n\n'
|
||||
|
||||
|
||||
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))
|
||||
emit(stage='loading', progress=1, message='Loading YuE2', model=model, vae=vae, duration=duration)
|
||||
import torch
|
||||
import soundfile as sf
|
||||
from yue2 import YuE2Pipeline
|
||||
if not torch.cuda.is_available():
|
||||
raise RuntimeError('YuE2 requires a CUDA GPU; CPU fallback is disabled.')
|
||||
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') 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
|
||||
Reference in New Issue
Block a user