213 lines
10 KiB
Python
213 lines
10 KiB
Python
"""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 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
|
||
offloader = ns['offload'].profile({'transformer': model, 'stage2': model2}, profile_no=cli.profile,
|
||
quantizeTransformer=cli.profile == 3, compile=compile_enabled, verboseLevel=1)
|
||
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)
|
||
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
|