141 lines
5.5 KiB
Python
141 lines
5.5 KiB
Python
"""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
|