Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -5,10 +5,10 @@ s3tokenizer
diffusers
torch==2.7.1
torchaudio==2.7.1
triton>=3.0.0
# triton>=3.0.0
transformers==4.57.1
accelerate==1.10.1
onnxruntime
onnxruntime-gpu
# onnxruntime-gpu
einops
gradio
32 changes: 21 additions & 11 deletions soulxpodcast/engine/llm_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@
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
Expand All @@ -26,10 +26,15 @@ 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"
if torch.cuda.is_available():
self.device = "cuda:0"
elif torch.backends.mps.is_available():
self.device = "mps"
else:
self.device = "cpu"
self.model = AutoModelForCausalLM.from_pretrained(model, torch_dtype=torch.bfloat16, device_map=self.device)
self.config = config
self.pad_token_id = self.tokenizer.pad_token_id
Expand All @@ -40,19 +45,19 @@ 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,
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():
with torch.no_grad():
input_len = len(prompt)
generated_ids = self.model.generate(
input_ids = torch.tensor([prompt], dtype=torch.int64).to(self.device),
Expand All @@ -78,14 +83,19 @@ def generate(
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"
if torch.cuda.is_available():
self.device = "cuda:0"
elif torch.backends.mps.is_available():
self.device = "mps"
else:
self.device = "cpu"
os.environ["VLLM_USE_V1"] = "0"
if SUPPORT_VLLM:
self.model = LLM(model=model, enforce_eager=True, dtype="bfloat16", max_model_len=8192, enable_prefix_caching=True,)
Expand All @@ -103,7 +113,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
Expand Down
51 changes: 34 additions & 17 deletions soulxpodcast/models/soulxpodcast.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,13 @@ 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.audio_tokenizer = s3tokenizer.load_model("speech_tokenizer_v2_25hz").to(self.device).eval()
if self.config.llm_engine == "hf":
self.llm = HFLLMEngine(**self.config.__dict__)
elif self.config.llm_engine == "vllm":
Expand All @@ -38,21 +44,21 @@ 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],
spk_ids: list[list[int]],
Expand All @@ -66,7 +72,7 @@ def forward_longform(

# 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
Expand All @@ -81,18 +87,18 @@ def forward_longform(
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_mel_len = torch.tensor([prompt_mel_len]).to(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]).to(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)

# 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 += [self.config.hf_config.eos_token_id]
Expand All @@ -108,7 +114,7 @@ def forward_longform(
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)
Expand All @@ -128,33 +134,44 @@ def forward_longform(
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|>

# 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
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 generation 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):
device_type = "cuda" if torch.cuda.is_available() else ("mps" if torch.backends.mps.is_available() else "cpu")
# MPS autocast support might be limited, fallback to cpu or disable if needed, but trying generic first or just safely using device type
# actually torch.amp.autocast needs 'cuda', 'cpu', 'xpu', 'hpu' etc. 'mps' is supported in recent pytorch.

# Since we are on Mac, we should dynamically set this.
autocast_device = "cuda" if torch.cuda.is_available() else "cpu" # Safe fallback for now as MPS autocast can be tricky or we can try "mps" if we are sure.
# Given the previous error "Torch not compiled with CUDA", keeping "cuda" will fail or warn.
# Let's use the device type we detected.
if torch.backends.mps.is_available():
autocast_device = "mps"

with torch.amp.autocast(device_type=autocast_device, dtype=torch.float16 if self.config.hf_config.fp16_flow else torch.float32):
generated_mels, generated_mels_lens = self.flow(
flow_input.cuda(), flow_inputs_len.cuda(),
prompt_mels, prompt_mels_lens, spk_emb.cuda(),
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
)

Expand Down