diff --git a/requirements.macos.txt b/requirements.macos.txt new file mode 100644 index 0000000..9ac9661 --- /dev/null +++ b/requirements.macos.txt @@ -0,0 +1,12 @@ +librosa +numpy +scipy +s3tokenizer +diffusers +torch==2.7.1 +torchaudio==2.7.1 +transformers==4.57.1 +accelerate==1.10.1 +onnxruntime +einops +gradio diff --git a/soulxpodcast/engine/llm_engine.py b/soulxpodcast/engine/llm_engine.py index 6694974..4aa9aeb 100644 --- a/soulxpodcast/engine/llm_engine.py +++ b/soulxpodcast/engine/llm_engine.py @@ -1,15 +1,12 @@ import os -import types -import atexit -from time import perf_counter from functools import partial from dataclasses import fields, asdict import torch -import torch.multiprocessing as mp from transformers import AutoTokenizer, AutoModelForCausalLM, StoppingCriteriaList from transformers import EosTokenCriteria, RepetitionPenaltyLogitsProcessor -try: + +try: from vllm import LLM from vllm import SamplingParams as VllmSamplingParams from vllm.inputs import TokensPrompt as TokensPrompt @@ -20,19 +17,45 @@ from soulxpodcast.config import Config, SamplingParams from soulxpodcast.models.modules.sampler import _ras_sample_hf_engine + +def _resolve_device() -> tuple[torch.device, str]: + if torch.cuda.is_available(): + return torch.device("cuda:0"), "cuda" + if torch.backends.mps.is_available() and torch.backends.mps.is_built(): + return torch.device("mps"), "mps" + return torch.device("cpu"), "cpu" + + +def _resolve_dtype(device_type: str) -> torch.dtype: + if device_type in {"cuda", "mps"}: + return torch.float16 + return torch.float32 + + class HFLLMEngine: def __init__(self, model, **kwargs): config_fields = {field.name for field in fields(Config)} config_kwargs = {k: v for k, v in kwargs.items() if k in config_fields} config = Config(model, **config_kwargs) - + self.tokenizer = AutoTokenizer.from_pretrained(model, use_fast=True) - config.eos = config.hf_config.eos_token_id # speech eos token; - self.device = "cuda:0" if torch.cuda.is_available() else "cpu" - self.model = AutoModelForCausalLM.from_pretrained(model, torch_dtype=torch.bfloat16, device_map=self.device) + config.eos = config.hf_config.eos_token_id # speech eos token + + self.device, self.device_type = _resolve_device() + dtype = _resolve_dtype(self.device_type) + + self.model = AutoModelForCausalLM.from_pretrained( + model, + dtype=dtype, + low_cpu_mem_usage=True, + ).to(self.device) + self.model.eval() + self.config = config self.pad_token_id = self.tokenizer.pad_token_id + if self.pad_token_id is None: + self.pad_token_id = self.config.hf_config.eos_token_id def generate( self, @@ -40,22 +63,27 @@ def generate( sampling_param: SamplingParams, past_key_values=None, ) -> dict: - + stopping_criteria = StoppingCriteriaList([EosTokenCriteria(eos_token_id=self.config.hf_config.eos_token_id)]) if sampling_param.use_ras: - sample_hf_engine_handler = partial(_ras_sample_hf_engine, - use_ras=sampling_param.use_ras, - win_size=sampling_param.win_size, tau_r=sampling_param.tau_r) + sample_hf_engine_handler = partial( + _ras_sample_hf_engine, + use_ras=sampling_param.use_ras, + win_size=sampling_param.win_size, + tau_r=sampling_param.tau_r, + ) else: sample_hf_engine_handler = None + rep_pen_processor = RepetitionPenaltyLogitsProcessor( penalty=sampling_param.repetition_penalty, - prompt_ignore_length=len(prompt) - ) # exclude the input prompt, consistent with vLLM implementation; - with torch.no_grad(): + prompt_ignore_length=len(prompt), + ) # exclude the input prompt, consistent with vLLM implementation + + with torch.no_grad(): input_len = len(prompt) generated_ids = self.model.generate( - input_ids = torch.tensor([prompt], dtype=torch.int64).to(self.device), + input_ids=torch.tensor([prompt], dtype=torch.int64, device=self.device), do_sample=True, top_k=sampling_param.top_k, top_p=sampling_param.top_p, @@ -66,31 +94,46 @@ def generate( past_key_values=past_key_values, custom_generate=sample_hf_engine_handler, use_cache=True, - logits_processor=[rep_pen_processor] + logits_processor=[rep_pen_processor], + pad_token_id=self.pad_token_id, + eos_token_id=self.config.hf_config.eos_token_id, ) generated_ids = generated_ids[:, input_len:].cpu().numpy().tolist()[0] + output = { "text": self.tokenizer.decode(generated_ids), "token_ids": generated_ids, } return output + class VLLMEngine: def __init__(self, model, **kwargs): - config_fields = {field.name for field in fields(Config)} config_kwargs = {k: v for k, v in kwargs.items() if k in config_fields} config = Config(model, **config_kwargs) - + self.tokenizer = AutoTokenizer.from_pretrained(config.model, use_fast=True) - config.eos = config.hf_config.eos_token_id # speech eos token; - self.device = "cuda:0" if torch.cuda.is_available() else "cpu" + config.eos = config.hf_config.eos_token_id # speech eos token + + self.device, self.device_type = _resolve_device() os.environ["VLLM_USE_V1"] = "0" + + if self.device_type != "cuda": + raise RuntimeError("vLLM requires CUDA; use llm_engine='hf' on macOS/CPU") + if SUPPORT_VLLM: - self.model = LLM(model=model, enforce_eager=True, dtype="bfloat16", max_model_len=8192, enable_prefix_caching=True,) + self.model = LLM( + model=model, + enforce_eager=True, + dtype="bfloat16", + max_model_len=8192, + enable_prefix_caching=True, + ) else: raise ImportError("Not Support VLLM now!!!") + self.config = config self.pad_token_id = self.tokenizer.pad_token_id @@ -103,7 +146,7 @@ def generate( sampling_param.stop_token_ids = [self.config.hf_config.eos_token_id] with torch.no_grad(): generated_ids = self.model.generate( - TokensPrompt(prompt_token_ids=prompt), + TokensPrompt(prompt_token_ids=prompt), VllmSamplingParams(**asdict(sampling_param)), use_tqdm=False, )[0].outputs[0].token_ids @@ -111,4 +154,4 @@ def generate( "text": self.tokenizer.decode(generated_ids), "token_ids": list(generated_ids), } - return output \ No newline at end of file + return output diff --git a/soulxpodcast/models/soulxpodcast.py b/soulxpodcast/models/soulxpodcast.py index 05b2a01..9bf296f 100644 --- a/soulxpodcast/models/soulxpodcast.py +++ b/soulxpodcast/models/soulxpodcast.py @@ -1,28 +1,42 @@ import time from datetime import datetime +from contextlib import nullcontext from tqdm import tqdm from itertools import chain -from copy import deepcopy -import numpy as np import s3tokenizer import torch -from transformers import AutoTokenizer, AutoModelForCausalLM, DynamicCache -from soulxpodcast.config import Config, SamplingParams, AutoPretrainedConfig +from transformers import DynamicCache +from soulxpodcast.config import Config, AutoPretrainedConfig from soulxpodcast.engine.llm_engine import ( HFLLMEngine, VLLMEngine ) from soulxpodcast.models.modules.flow import CausalMaskedDiffWithXvec from soulxpodcast.models.modules.hifigan import HiFTGenerator + class SoulXPodcast(torch.nn.Module): def __init__(self, config: Config = None): super().__init__() self.config = Config() if config is None else config - self.audio_tokenizer = s3tokenizer.load_model("speech_tokenizer_v2_25hz").cuda().eval() + if torch.cuda.is_available(): + self.device = torch.device("cuda") + elif torch.backends.mps.is_available(): + self.device = torch.device("mps") + else: + self.device = torch.device("cpu") + self.device_type = self.device.type + + self.audio_tokenizer = s3tokenizer.load_model("speech_tokenizer_v2_25hz") + if self.device_type == "cuda": + self.audio_tokenizer = self.audio_tokenizer.cuda() + elif hasattr(self.audio_tokenizer, "to"): + self.audio_tokenizer = self.audio_tokenizer.to(self.device) + self.audio_tokenizer.eval() + if self.config.llm_engine == "hf": self.llm = HFLLMEngine(**self.config.__dict__) elif self.config.llm_engine == "vllm": @@ -38,39 +52,37 @@ def __init__(self, config: Config = None): tqdm.write(f"[{timestamp}] - [INFO] - Casting flow to fp16") self.flow.half() self.flow.load_state_dict(torch.load(f"{self.config.model}/flow.pt", map_location="cpu", weights_only=True), strict=True) - self.flow.cuda().eval() + self.flow.to(self.device).eval() self.hift = HiFTGenerator() hift_state_dict = {k.replace('generator.', ''): v for k, v in torch.load(f"{self.config.model}/hift.pt", map_location="cpu", weights_only=True).items()} self.hift.load_state_dict(hift_state_dict, strict=True) - self.hift.cuda().eval() + self.hift.to(self.device).eval() - @torch.inference_mode() def forward_longform( self, prompt_mels_for_llm, prompt_mels_lens_for_llm: torch.Tensor, prompt_text_tokens_for_llm: list[list[int]], text_tokens_for_llm: list[list[int]], - prompt_mels_for_flow_ori, + prompt_mels_for_flow_ori, spk_emb_for_flow: torch.Tensor, - sampling_params: SamplingParams | list[SamplingParams], + sampling_params, spk_ids: list[list[int]], use_dialect_prompt: bool = False, dialect_prompt_text_tokens_for_llm: list[list[int]] = None, dialect_prefix: list[list[int]] = None, - **kwargs, # for compatibility + **kwargs, ): prompt_size, turn_size = len(prompt_mels_for_llm), len(text_tokens_for_llm) # Audio tokenization prompt_speech_tokens_ori, prompt_speech_tokens_lens_ori = self.audio_tokenizer.quantize( - prompt_mels_for_llm.cuda(), prompt_mels_lens_for_llm.cuda() + prompt_mels_for_llm.to(self.device), prompt_mels_lens_for_llm.to(self.device) ) - # align speech token with speech feat as to reduce - # the noise ratio during the generation process. + # Align speech token with speech feat to reduce noise ratio during generation. prompt_speech_tokens = [] prompt_mels_for_flow, prompt_mels_lens_for_flow = [], [] @@ -80,11 +92,11 @@ def forward_longform( prompt_mel = prompt_mels_for_flow_ori[prompt_index] prompt_mel_len = prompt_mel.shape[0] if prompt_speech_token_len * 2 > prompt_mel_len: - prompt_speech_token = prompt_speech_token[:int(prompt_mel_len/2)] - prompt_mel_len = torch.tensor([prompt_mel_len]).cuda() + prompt_speech_token = prompt_speech_token[:int(prompt_mel_len / 2)] + prompt_mel_len = torch.tensor([prompt_mel_len], device=self.device) else: - prompt_mel = prompt_mel.detach().clone()[:prompt_speech_token_len * 2].cuda() - prompt_mel_len = torch.tensor([prompt_speech_token_len * 2]).cuda() + prompt_mel = prompt_mel.detach().clone()[:prompt_speech_token_len * 2].to(self.device) + prompt_mel_len = torch.tensor([prompt_speech_token_len * 2], device=self.device) prompt_speech_tokens.append(prompt_speech_token) prompt_mels_for_flow.append(prompt_mel) prompt_mels_lens_for_flow.append(prompt_mel_len) @@ -92,77 +104,81 @@ def forward_longform( # Prepare LLM inputs prompt_inputs = [] history_inputs = [] - + for i in range(prompt_size): - speech_tokens_i = [token+self.config.hf_config.speech_token_offset for token in prompt_speech_tokens[i].tolist()] + speech_tokens_i = [token + self.config.hf_config.speech_token_offset for token in prompt_speech_tokens[i].tolist()] speech_tokens_i += [self.config.hf_config.eos_token_id] - if use_dialect_prompt and len(dialect_prompt_text_tokens_for_llm[i])>0: + if use_dialect_prompt and len(dialect_prompt_text_tokens_for_llm[i]) > 0: dialect_prompt_input = prompt_text_tokens_for_llm[i] + speech_tokens_i + dialect_prompt_text_tokens_for_llm[i] - if i>0: + if i > 0: dialect_prompt_input = dialect_prefix[0] + dialect_prompt_input prompt_input = self.llm.generate(dialect_prompt_input, sampling_params, past_key_values=None)['token_ids'] - prompt_inputs.append(dialect_prefix[i+1]+dialect_prompt_text_tokens_for_llm[i] + prompt_input) - history_inputs.append(dialect_prefix[i+1]+dialect_prompt_text_tokens_for_llm[i] + prompt_input) + prompt_inputs.append(dialect_prefix[i + 1] + dialect_prompt_text_tokens_for_llm[i] + prompt_input) + history_inputs.append(dialect_prefix[i + 1] + dialect_prompt_text_tokens_for_llm[i] + prompt_input) else: - prompt_inputs.append(prompt_text_tokens_for_llm[i] + speech_tokens_i ) - history_inputs.append(prompt_text_tokens_for_llm[i] + speech_tokens_i ) + prompt_inputs.append(prompt_text_tokens_for_llm[i] + speech_tokens_i) + history_inputs.append(prompt_text_tokens_for_llm[i] + speech_tokens_i) generated_wavs, results_dict = [], {} - + # LLM generation inputs = list(chain.from_iterable(prompt_inputs)) cache_config = AutoPretrainedConfig().from_dataclass(self.llm.config.hf_config) past_key_values = DynamicCache(config=cache_config) valid_turn_size = prompt_size for i in range(turn_size): - - # # set ratio: reach the reset cache ratio; - if valid_turn_size > self.config.max_turn_size or len(inputs)>self.config.turn_tokens_threshold: - assert self.config.max_turn_size >= self.config.prompt_context + self.config.history_context, "Invalid Long history size setting, " - prompt_text_bound = max(self.config.prompt_context, len(history_inputs)-self.config.history_text_context-self.config.history_context) + # Reach reset cache ratio. + if valid_turn_size > self.config.max_turn_size or len(inputs) > self.config.turn_tokens_threshold: + assert self.config.max_turn_size >= self.config.prompt_context + self.config.history_context, "Invalid Long history size setting" + prompt_text_bound = max(self.config.prompt_context, len(history_inputs) - self.config.history_text_context - self.config.history_context) inputs = list(chain.from_iterable( - history_inputs[:self.config.prompt_context]+ \ - history_inputs[prompt_text_bound:-self.config.history_context]+ \ + history_inputs[:self.config.prompt_context] + + history_inputs[prompt_text_bound:-self.config.history_context] + prompt_inputs[-self.config.history_context:] )) valid_turn_size = self.config.prompt_context + len(history_inputs) - prompt_text_bound past_key_values = DynamicCache(config=cache_config) valid_turn_size += 1 - + inputs.extend(text_tokens_for_llm[i]) - start_time = time.time() llm_outputs = self.llm.generate(inputs, sampling_params, past_key_values=past_key_values) inputs.extend(llm_outputs['token_ids']) - prompt_inputs.append(text_tokens_for_llm[i]+llm_outputs['token_ids']) - history_inputs.append(text_tokens_for_llm[i][:-1]) # remove the <|audio_start|> - + prompt_inputs.append(text_tokens_for_llm[i] + llm_outputs['token_ids']) + history_inputs.append(text_tokens_for_llm[i][:-1]) # remove <|audio_start|> + # Prepare Flow inputs turn_spk = spk_ids[i] - generated_speech_tokens = [token - self.config.hf_config.speech_token_offset for token in llm_outputs['token_ids'][:-1]] # ignore last eos + generated_speech_tokens = [token - self.config.hf_config.speech_token_offset for token in llm_outputs['token_ids'][:-1]] prompt_speech_token = prompt_speech_tokens[turn_spk].tolist() flow_input = torch.tensor([prompt_speech_token + generated_speech_tokens]) flow_inputs_len = torch.tensor([len(prompt_speech_token) + len(generated_speech_tokens)]) - # Flow generation and HiFi-GAN generation + # Flow and HiFi-GAN generation start_idx = spk_ids[i] prompt_mels = prompt_mels_for_flow[start_idx][None] prompt_mels_lens = prompt_mels_lens_for_flow[start_idx][None] - spk_emb = spk_emb_for_flow[start_idx:start_idx+1] - - # Flow generation - with torch.amp.autocast("cuda", dtype=torch.float16 if self.config.hf_config.fp16_flow else torch.float32): + spk_emb = spk_emb_for_flow[start_idx:start_idx + 1] + + autocast_ctx = ( + torch.amp.autocast("cuda", dtype=torch.float16 if self.config.hf_config.fp16_flow else torch.float32) + if self.device_type == "cuda" + else nullcontext() + ) + with autocast_ctx: generated_mels, generated_mels_lens = self.flow( - flow_input.cuda(), flow_inputs_len.cuda(), - prompt_mels, prompt_mels_lens, spk_emb.cuda(), - streaming=False, finalize=True + flow_input.to(self.device), + flow_inputs_len.to(self.device), + prompt_mels, + prompt_mels_lens, + spk_emb.to(self.device), + streaming=False, + finalize=True, ) - # HiFi-GAN generation mel = generated_mels[:, :, prompt_mels_lens[0].item():generated_mels_lens[0].item()] wav, _ = self.hift(speech_feat=mel) generated_wavs.append(wav) - # Save the generated wav; results_dict['generated_wavs'] = generated_wavs - return results_dict \ No newline at end of file + return results_dict diff --git a/soulxpodcast/utils/commons.py b/soulxpodcast/utils/commons.py index 4ff87d0..2f57674 100644 --- a/soulxpodcast/utils/commons.py +++ b/soulxpodcast/utils/commons.py @@ -7,4 +7,5 @@ def set_all_random_seed(seed): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) - torch.cuda.manual_seed_all(seed) \ No newline at end of file + if torch.cuda.is_available(): + torch.cuda.manual_seed_all(seed) diff --git a/soulxpodcast/utils/infer_utils.py b/soulxpodcast/utils/infer_utils.py index 78b764d..00e7d8a 100644 --- a/soulxpodcast/utils/infer_utils.py +++ b/soulxpodcast/utils/infer_utils.py @@ -1,7 +1,5 @@ import re -import json import torch -import argparse from tqdm import tqdm from datetime import datetime @@ -15,57 +13,60 @@ def initiate_model(seed, model_path, llm_engine, fp16_flow): set_all_random_seed(seed) - + hf_config = SoulXPodcastLLMConfig.from_initial_and_json( initial_values={"fp16_flow": fp16_flow}, - json_file=f"{model_path}/soulxpodcast_config.json" + json_file=f"{model_path}/soulxpodcast_config.json", ) + if llm_engine == "vllm": import importlib.util - if not importlib.util.find_spec("vllm"): + + if (not importlib.util.find_spec("vllm")) or (not torch.cuda.is_available()): llm_engine = "hf" - timestamp = datetime.now().strftime('%Y-%m-%d %H:%M:%S,%f')[:-3] - tqdm.write(f"[{timestamp}] - [WARNING]: No install VLLM, switch to hf engine.") + timestamp = datetime.now().strftime("%Y-%m-%d %H:%M:%S,%f")[:-3] + tqdm.write(f"[{timestamp}] - [WARNING]: VLLM unavailable on current device, switch to hf engine.") config = Config(model=model_path, enforce_eager=True, llm_engine=llm_engine, hf_config=hf_config) model = SoulXPodcast(config) dataset = PodcastInferHandler(model.llm.tokenizer, None, config) - + return model, dataset def process_single_input(dataset, target_text_list, prompt_wav_list, prompt_text_list, use_dialect_prompt, dialect_prompt_text_list): spks, texts = [], [] for target_text in target_text_list: - pattern = r'(\[S[1-9]\])(.+)' + pattern = r"(\[S[1-9]\])(.+)" match = re.match(pattern, target_text) - text, spk = match.group(2), int(match.group(1)[2])-1 + text, spk = match.group(2), int(match.group(1)[2]) - 1 spks.append(spk) texts.append(text) - - dataitem = {"key": "001", "prompt_text": prompt_text_list, "prompt_wav": prompt_wav_list, - "text": texts, "spk": spks, } + + dataitem = { + "key": "001", + "prompt_text": prompt_text_list, + "prompt_wav": prompt_wav_list, + "text": texts, + "spk": spks, + } if use_dialect_prompt: dataitem.update({ - "dialect_prompt_text": dialect_prompt_text_list + "dialect_prompt_text": dialect_prompt_text_list, }) - dataset.update_datasource( - [ - dataitem - ] - ) + dataset.update_datasource([dataitem]) - # assert one data only; + # assert one data only data = dataset[0] prompt_mels_for_llm, prompt_mels_lens_for_llm = s3tokenizer.padding(data["log_mel"]) # [B, num_mels=128, T] spk_emb_for_flow = torch.tensor(data["spk_emb"]) prompt_mels_for_flow = torch.nn.utils.rnn.pad_sequence(data["mel"], batch_first=True, padding_value=0) # [B, T', num_mels=80] - prompt_mels_lens_for_flow = torch.tensor(data['mel_len']) + prompt_mels_lens_for_flow = torch.tensor(data["mel_len"]) text_tokens_for_llm = data["text_tokens"] prompt_text_tokens_for_llm = data["prompt_text_tokens"] spk_ids = data["spks_list"] - sampling_params = SamplingParams(use_ras=True,win_size=25,tau_r=0.2) + sampling_params = SamplingParams(use_ras=True, win_size=25, tau_r=0.2) infos = [data["info"]] processed_data = { "prompt_mels_for_llm": prompt_mels_for_llm, @@ -89,7 +90,7 @@ def process_single_input(dataset, target_text_list, prompt_wav_list, prompt_text def check_models(model_path, inputs): - if inputs['use_dialect_prompt']: - assert 'dialect' in model_path, "Dialect prompt is used, you should use a dialect model." - - return True \ No newline at end of file + if inputs["use_dialect_prompt"]: + assert "dialect" in model_path, "Dialect prompt is used, you should use a dialect model." + + return True