import gradio as gr import subprocess import os import tempfile import spaces import torch import sys import uuid import re import numpy as np import json from omegaconf import OmegaConf import torchaudio from torchaudio.transforms import Resample import soundfile as sf from tqdm import tqdm from einops import rearrange from transformers import AutoTokenizer, AutoModelForCausalLM, LogitsProcessor, LogitsProcessorList from collections import Counter # Disable FlashAttention installation completely print("Skipping FlashAttention installation for CPU mode...") # Download model dependencies from huggingface_hub import snapshot_download folder_path = './xcodec_mini_infer' if not os.path.exists(folder_path): os.mkdir(folder_path) print(f"Folder created at: {folder_path}") else: print(f"Folder already exists at: {folder_path}") snapshot_download( repo_id="m-a-p/xcodec_mini_infer", local_dir="./xcodec_mini_infer", ) sys.path.append(os.path.join(os.path.dirname(os.path.abspath(__file__)), 'xcodec_mini_infer')) sys.path.append(os.path.join(os.path.dirname(os.path.abspath(__file__)), 'xcodec_mini_infer', 'descriptaudiocodec')) from codecmanipulator import CodecManipulator from mmtokenizer import _MMSentencePieceTokenizer from models.soundstream_hubert_new import SoundStream device = "cpu" print("Loading model on CPU...") model = AutoModelForCausalLM.from_pretrained( "m-a-p/YuE-s1-7B-anneal-en-cot", torch_dtype=torch.float32, attn_implementation="eager", # disable flash attention ) model.to(device) model.eval() print("YuE model loaded successfully on CPU.") # Codec setup basic_model_config = './xcodec_mini_infer/final_ckpt/config.yaml' resume_path = './xcodec_mini_infer/final_ckpt/ckpt_00360000.pth' mmtokenizer = _MMSentencePieceTokenizer("./mm_tokenizer_v0.2_hf/tokenizer.model") codectool = CodecManipulator("xcodec", 0, 1) model_config = OmegaConf.load(basic_model_config) codec_model = eval(model_config.generator.name)(**model_config.generator.config).to(device) parameter_dict = torch.load(resume_path, map_location='cpu') codec_model.load_state_dict(parameter_dict['codec_model']) codec_model.eval() print("Codec model loaded successfully.") # ---------------- Utility Classes ---------------- # class BlockTokenRangeProcessor(LogitsProcessor): def __init__(self, start_id, end_id): self.blocked_token_ids = list(range(start_id, end_id)) def __call__(self, input_ids, scores): scores[:, self.blocked_token_ids] = -float("inf") return scores def load_audio_mono(filepath, sampling_rate=16000): audio, sr = torchaudio.load(filepath) audio = torch.mean(audio, dim=0, keepdim=True) if sr != sampling_rate: resampler = Resample(orig_freq=sr, new_freq=sampling_rate) audio = resampler(audio) return audio def split_lyrics(lyrics: str): pattern = r"\[(\w+)\]\s*(.*?)(?=\s*\n\[|\Z)" segments = re.findall(pattern, lyrics, re.DOTALL) structured_lyrics = [f"[{seg[0]}]\n{seg[1].strip()}\n\n" for seg in segments] return structured_lyrics # ---------------- Generation Logic ---------------- # @spaces.GPU(duration=178) def generate_music( genre_txt=None, lyrics_txt=None, run_n_segments=2, max_new_tokens=45, use_audio_prompt=False, audio_prompt_path="", progress=gr.Progress() ): if use_audio_prompt and not audio_prompt_path: raise FileNotFoundError("Please provide an audio prompt file when 'Use Audio Prompt' is enabled!") with tempfile.TemporaryDirectory() as output_dir: stage1_output_dir = os.path.join(output_dir, f"stage1") os.makedirs(stage1_output_dir, exist_ok=True) genres = genre_txt.strip() lyrics = split_lyrics(lyrics_txt + "\n") full_lyrics = "\n".join(lyrics) prompt_texts = [f"Generate music from lyrics.\n[Genre] {genres}\n{full_lyrics}"] prompt_texts += lyrics random_id = uuid.uuid4() top_p = 0.93 temperature = 1.0 repetition_penalty = 1.2 start_of_segment = mmtokenizer.tokenize('[start_of_segment]') end_of_segment = mmtokenizer.tokenize('[end_of_segment]') run_n_segments = min(run_n_segments, len(lyrics)) + 1 raw_output = None def generator(): nonlocal raw_output for i, p in enumerate(tqdm(prompt_texts[:run_n_segments])): if i == 0: continue section_text = p.replace('[start_of_segment]', '').replace('[end_of_segment]', '') if i == 1: head_id = mmtokenizer.tokenize(prompt_texts[0]) prompt_ids = head_id + start_of_segment + mmtokenizer.tokenize(section_text) + [mmtokenizer.soa] + codectool.sep_ids else: prompt_ids = end_of_segment + start_of_segment + mmtokenizer.tokenize(section_text) + [mmtokenizer.soa] + codectool.sep_ids prompt_ids = torch.as_tensor(prompt_ids).unsqueeze(0).to(device) input_ids = torch.cat([raw_output, prompt_ids], dim=1) if i > 1 else prompt_ids max_context = 16384 - max_new_tokens - 1 if input_ids.shape[-1] > max_context: input_ids = input_ids[:, -(max_context):] with torch.inference_mode(): output_seq = model.generate( input_ids=input_ids, max_new_tokens=max_new_tokens, do_sample=True, top_p=top_p, temperature=temperature, repetition_penalty=repetition_penalty, eos_token_id=mmtokenizer.eoa, pad_token_id=mmtokenizer.eoa, logits_processor=LogitsProcessorList([BlockTokenRangeProcessor(0, 32002)]), num_beams=1 ) if i > 1: raw_output = torch.cat([raw_output, prompt_ids, output_seq[:, input_ids.shape[-1]:]], dim=1) else: raw_output = output_seq return raw_output raw_output = generator() ids = raw_output[0].cpu().numpy() soa_idx = np.where(ids == mmtokenizer.soa)[0].tolist() eoa_idx = np.where(ids == mmtokenizer.eoa)[0].tolist() vocals, instrumentals = [], [] for i in range(len(soa_idx)): codec_ids = ids[soa_idx[i] + 1:eoa_idx[i]] codec_ids = codec_ids[:2 * (len(codec_ids) // 2)] vocals_ids = codectool.ids2npy(rearrange(codec_ids, "(n b) -> b n", b=2)[0]) vocals.append(vocals_ids) instrumentals_ids = codectool.ids2npy(rearrange(codec_ids, "(n b) -> b n", b=2)[1]) instrumentals.append(instrumentals_ids) vocals = np.concatenate(vocals, axis=1) instrumentals = np.concatenate(instrumentals, axis=1) def save_audio(wav: torch.Tensor, path, sample_rate: int, rescale: bool = False): limit = 0.99 max_val = wav.abs().max() wav = wav * min(limit / max_val, 1) if rescale else wav.clamp(-limit, limit) torchaudio.save(str(path), wav, sample_rate=sample_rate, encoding='PCM_S', bits_per_sample=16) recons_output_dir = os.path.join(output_dir, "recons") os.makedirs(recons_output_dir, exist_ok=True) mix_stem = (vocals + instrumentals) / 2 mix_path = os.path.join(recons_output_dir, f"mix_{random_id}.wav") save_audio(torch.tensor(mix_stem).unsqueeze(0), mix_path, 16000) return mix_path # ---------------- Gradio Interface ---------------- # with gr.Blocks() as demo: with gr.Column(): gr.Markdown("# 🎵 YuE: CPU Music Generator (Simplified Demo)") with gr.Row(): with gr.Column(): genre_txt = gr.Textbox(label="Genre") lyrics_txt = gr.Textbox(label="Lyrics") submit_btn = gr.Button("Generate Song 🎶") with gr.Column(): music_out = gr.Audio(label="Generated Music") submit_btn.click( fn=generate_music, inputs=[genre_txt, lyrics_txt], outputs=[music_out] ) demo.queue().launch(show_error=True, share=True)