Files
aigen/scripts/yuegp-worker.py
T

232 lines
11 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 adapter for pinned deepbeepmeep/YuEGP inference functions.
No Comfy imports, no Gradio import/server, and no automatic profile fallback.
Stdout AIGEN_EVENT lines are consumed by the device-local host agent.
"""
import argparse
import ast
import hashlib
import gc
import json
import os
from pathlib import Path
import re
import sys
import time
import threading
from types import SimpleNamespace, FunctionType
YUEGP_REVISION = '2d72ff734b7a127324353c0dcd0f95ca4cc0b797'
def emit(**event):
print('AIGEN_EVENT ' + json.dumps(event), flush=True)
def normalize_lyrics(text):
# Accept our existing chips (Verse 1, Pre-Chorus); upstream only accepts \w+.
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('YuEGP 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 load_functions(root):
source = root / 'inference' / 'gradio_server.py'
text = source.read_text(encoding='utf-8')
if hashlib.sha256(text.encode()).hexdigest() != '567b91b3100c4fb2c95496e83b135e9d55fd03714e83912f9586037805fdecc9':
raise RuntimeError('YuEGP inference source differs from the validated revision.')
tree = ast.parse(text)
names = {'BlockTokenRangeProcessor', 'load_audio_mono', 'encode_audio',
'stage1_inference', 'stage2_generate', 'stage2_inference'}
# Load the library functions only. Never execute upstream CLI/UI startup.
nodes = [n for n in tree.body if isinstance(n, (ast.FunctionDef, ast.ClassDef)) and n.name in names]
if {n.name for n in nodes} != names:
raise RuntimeError('Unsupported YuEGP source: reinstall the pinned revision.')
ns = {'__file__': str(source)}
imports = [n for n in tree.body if isinstance(n, (ast.Import, ast.ImportFrom))
and not (isinstance(n, ast.Import) and any(a.name == 'gradio' for a in n.names))]
exec(compile(ast.Module(body=imports + nodes, type_ignores=[]), str(source), 'exec'), ns)
return ns
def load_codec_on_cpu(ns, torch):
with torch.device('cpu'):
config = ns['OmegaConf'].load('xcodec_mini_infer/final_ckpt/config.yaml')
codec = ns['SoundStream'](**config.generator.config)
codec.load_state_dict(torch.load('xcodec_mini_infer/final_ckpt/ckpt_00360000.pth', map_location='cpu', weights_only=False)['codec_model'])
codec.eval()
return codec
def memory_profile_options(profile, compile_enabled, platform, total_ram):
options = dict(profile_no=profile, quantizeTransformer=profile == 3,
compile=compile_enabled, verboseLevel=1)
# Windows CUDA pinned allocations can fail before inference on 32GB hosts.
# This affects transfer staging only, not weight precision or VRAM budgets.
if platform == 'win32' and total_ram < 48 * 1024 ** 3:
options['pinnedMemory'] = False
return options
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)
parser.add_argument('--profile', type=int, choices=[1, 3], default=1)
parser.add_argument('--compile', action='store_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'])
tags = ' '.join(str(request['tags']).split())
seed = int(request['seed'])
max_tokens = duration * 100
if max_tokens >= 16000:
raise ValueError('YuEGP supports at most 150 seconds per section.')
compile_enabled = False
if cli.compile:
import triton # Explicit opt-in AND an actual successful import required.
compile_enabled = True
os.chdir(root / 'inference')
sys.path[:0] = [str(root / 'inference'), str(root / 'inference/xcodec_mini_infer'),
str(root / 'inference/xcodec_mini_infer/descriptaudiocodec')]
emit(stage='loading', progress=1, message=f'Loading YuEGP profile {cli.profile}', profile=cli.profile,
compile=compile_enabled, sections=1, maxNewTokens=max_tokens)
ns = load_functions(root)
torch, np, sf = ns['torch'], ns['np'], ns['sf']
if not torch.cuda.is_available():
raise RuntimeError('YuEGP requires a CUDA GPU; CPU fallback is disabled.')
# The upstream transformer patch is part of YuEGP, in its own environment.
import transformers.generation.utils as generation
if 'callback' not in Path(generation.__file__).read_text(encoding='utf-8'):
raise RuntimeError('YuEGP transformer functions are missing. Run setup-yuegp.ps1.')
attention = 'sdpa'
try:
import flash_attn
attention = 'flash_attention_2'
except ImportError:
pass
ns['random'].seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
device = torch.device('cuda:0')
stage1 = request.get('stage1Model') or 'm-a-p/YuE-s1-7B-anneal-en-cot'
stage2 = request.get('stage2Model') or 'm-a-p/YuE-s2-1B-general'
model = ns['AutoModelForCausalLM'].from_pretrained(stage1, torch_dtype=torch.bfloat16, attn_implementation=attention).eval()
model2 = ns['AutoModelForCausalLM'].from_pretrained(stage2, torch_dtype=torch.float16, attn_implementation=attention).eval()
if not compile_enabled:
# Transformers 4.48 selects DynamicCache with None; the literal
# string "dynamic" is not an accepted configuration value.
model.generation_config.cache_implementation = None
model2.generation_config.cache_implementation = None
model._validate_model_kwargs = lambda _: None
model2._validate_model_kwargs = lambda _: None
profile_options = memory_profile_options(cli.profile, compile_enabled, sys.platform, psutil.virtual_memory().total)
emit(stage='loading', message='Preparing YuEGP memory', profileOptions=profile_options)
offloader = ns['offload'].profile({'transformer': model, 'stage2': model2}, **profile_options)
def memory_checkpoint(stage):
torch.cuda.synchronize()
free, total = torch.cuda.mem_get_info()
emit(stage=stage, message=f'YuEGP {stage}', gpuFreeMiB=round(free / 1024 ** 2),
gpuTotalMiB=round(total / 1024 ** 2), gpuAllocatedMiB=round(torch.cuda.memory_allocated() / 1024 ** 2),
ramAvailableMiB=round(psutil.virtual_memory().available / 1024 ** 2))
memory_checkpoint('memory-ready')
args = SimpleNamespace(use_audio_prompt=False, use_dual_tracks_prompt=False, rescale=True,
output_dir=str(output), cuda_idx=0)
ns.update(model=model, model_stage2=model2, device=device, codec_model=None, args=args,
mmtokenizer=ns['_MMSentencePieceTokenizer']('./mm_tokenizer_v0.2_hf/tokenizer.model'),
codectool=ns['CodecManipulator']('xcodec', 0, 1), codectool_stage2=ns['CodecManipulator']('xcodec', 0, 8),
stage1_output_dir=str(output / 'stage1'),
split_lyrics=lambda _: [lyrics], get_song_id=lambda *a: 'song')
(output / 'stage1').mkdir(exist_ok=True)
(output / 'stage2').mkdir(exist_ok=True)
state = {}
last = [0.0]
def callback(done, total):
now = time.monotonic()
if now - last[0] < 1 and done < total:
return
last[0] = now
stage = state.get('stage', 'Generating')
# Percent is explicitly local to the reported stage, never a song ETA.
emit(stage=stage, message=stage, step=int(done), maxStep=int(total),
progress=round(100 * done / max(1, total), 1))
emit(stage='stage1', message='Generating song tokens', progress=0)
stems = ns['stage1_inference'](tags, lyrics, 1, max_tokens, seed, state, callback)
memory_checkpoint('stage1-finished')
emit(stage='stage2', message='Generating audio detail', progress=0)
results = ns['stage2_inference'](model2, stems, str(output / 'stage2'),
batch_size=20 if cli.profile == 1 else 4, state=state, callback=callback)
# Release transformer allocations before codec/vocoder decode.
offloader.unload_all()
gc.collect()
torch.cuda.empty_cache()
# MMGP sets the default device to CUDA. Explicit CPU construction prevents
# decoder initialization from competing with the language models for VRAM.
emit(stage='decoding', message='Loading audio decoder', progress=0)
codec = load_codec_on_cpu(ns, torch)
codec.to(device)
emit(stage='decoding', message='Decoding and mixing audio', progress=0)
low_tracks = []
for path in results:
codes = np.load(path)
with torch.no_grad():
wave = codec.decode(torch.as_tensor(codes.astype(np.int16), dtype=torch.long).unsqueeze(0).permute(1, 0, 2).to(device))
low_tracks.append(wave.cpu().squeeze().numpy())
sf.write(str(output / 'mix16.wav'), low_tracks[0] + low_tracks[1], 16000)
with torch.device('cpu'):
vocal_decoder, inst_decoder = ns['build_codec_model']('xcodec_mini_infer/decoders/config.yaml',
'xcodec_mini_infer/decoders/decoder_131000.pth', 'xcodec_mini_infer/decoders/decoder_151000.pth')
# Keep upstream neural decoding, but write WAV directly: no MP3/FFmpeg
# backend dependency in the isolated Windows environment.
process_audio = ns['process_audio']
audio_globals = dict(process_audio.__globals__)
def save_wave(wave, path, sample_rate, rescale=False):
data = wave.detach().cpu().numpy()
peak = float(np.max(np.abs(data)))
if rescale and peak > 0.99:
data = data * (0.99 / peak)
sf.write(str(path), data.T, sample_rate, subtype='PCM_16')
audio_globals['save_audio'] = save_wave
decode_audio = FunctionType(process_audio.__code__, audio_globals)
tracks = []
for path in results:
instrumental = '_itrack' in path
decoder = inst_decoder if instrumental else vocal_decoder
with torch.no_grad():
tracks.append(decode_audio(path, str(output / ('instrumental.wav' if instrumental else 'vocal.wav')),
True, args, decoder, codec))
decoder.to('cpu')
torch.cuda.empty_cache()
mixed = (tracks[0] + tracks[1]).detach().cpu().squeeze().numpy()
sf.write(str(output / 'mix44.wav'), mixed, 44100)
ns['replace_low_freq_with_energy_matched'](a_file=str(output / 'mix16.wav'), b_file=str(output / 'mix44.wav'),
c_file=str(output / 'audio.wav'), cutoff_freq=5500.0)
info = sf.info(str(output / 'audio.wav'))
if info.frames <= 0:
raise RuntimeError('YuEGP 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