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
12 changes: 12 additions & 0 deletions requirements.macos.txt
Original file line number Diff line number Diff line change
@@ -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
93 changes: 68 additions & 25 deletions soulxpodcast/engine/llm_engine.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -20,42 +17,73 @@
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,
prompt: list[str],
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,
Expand All @@ -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

Expand All @@ -103,12 +146,12 @@ 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
output = {
"text": self.tokenizer.decode(generated_ids),
"token_ids": list(generated_ids),
}
return output
return output
Loading