Fast Inference via vLLM

This commit is contained in:
Daniel Han 2025-02-05 02:21:15 -08:00
commit e3051647ae
2 changed files with 111 additions and 16 deletions

View file

@ -1634,9 +1634,18 @@ class FastLlamaModel:
model_patcher = None,
tokenizer_name = None,
trust_remote_code = False,
fast_inference = False, # uses vLLM
gpu_memory_utilization = 0.5,
float8_kv_cache = True,
random_state = 3407,
max_lora_rank = 16,
disable_log_stats = False,
**kwargs,
):
if trust_remote_code:
if fast_inference:
raise NotImplementedError("Unsloth: Fast inference does not support `trust_remote_code` yet.")
print(
"Unsloth: WARNING `trust_remote_code` is True.\n"\
"Are you certain you want to do remote code execution?"
@ -1650,9 +1659,9 @@ class FastLlamaModel:
statistics = \
f"==((====))== Unsloth {__version__}: Fast {model_patcher.__name__[4:-5]} patching. Transformers: {transformers_version}.\n"\
f" \\\ /| GPU: {gpu_stats.name}. Max memory: {max_memory} GB. Platform: {platform_system}.\n"\
f"O^O/ \_/ \\ Torch: {torch.__version__}. CUDA: {gpu_stats.major}.{gpu_stats.minor}. CUDA Toolkit: {torch.version.cuda}. Triton: {triton_version}\n"\
f"\ / Bfloat16 = {str(SUPPORTS_BFLOAT16).upper()}. FA [Xformers = {xformers_version}. FA2 = {HAS_FLASH_ATTENTION}]\n"\
f" {chr(92)}{chr(92)} /| GPU: {gpu_stats.name}. Max memory: {max_memory} GB. Platform: {platform_system}.\n"\
f"O^O/ {chr(92)}_/ {chr(92)} Torch: {torch.__version__}. CUDA: {gpu_stats.major}.{gpu_stats.minor}. CUDA Toolkit: {torch.version.cuda}. Triton: {triton_version}\n"\
f"{chr(92)} / Bfloat16 = {str(SUPPORTS_BFLOAT16).upper()}. FA [Xformers = {xformers_version}. FA2 = {HAS_FLASH_ATTENTION}]\n"\
f' "-____-" Free Apache license: http://github.com/unslothai/unsloth'
print(statistics)
@ -1680,7 +1689,11 @@ class FastLlamaModel:
assert(dtype == torch.float16 or dtype == torch.bfloat16 or dtype == torch.float32)
# RoPE Scaling
model_config = AutoConfig.from_pretrained(model_name, token = token)
model_config = AutoConfig.from_pretrained(
model_name,
token = token,
attn_implementation = "sdpa",
)
model_max_seq_length = model_config.max_position_embeddings
# Check if RoPE Scaling is even allowed
@ -1701,6 +1714,9 @@ class FastLlamaModel:
rope_scaling = max_seq_length / model_max_seq_length
if fast_inference:
raise NotImplementedError("Unsloth: Fast inference does not yet work with RoPE Scaling.")
logger.warning_once(
f"Unsloth: {model_name} can only handle sequence lengths of at most "\
f"{model_max_seq_length}.\nBut with kaiokendev's RoPE scaling of "\
@ -1742,17 +1758,55 @@ class FastLlamaModel:
# Cannot be None, since HF now checks for the config
if load_in_4bit: kwargs["quantization_config"] = bnb_config
model = AutoModelForCausalLM.from_pretrained(
model_name,
device_map = device_map,
torch_dtype = dtype,
# quantization_config = bnb_config,
token = token,
max_position_embeddings = max_position_embeddings,
trust_remote_code = trust_remote_code,
attn_implementation = "eager",
**kwargs,
)
if not fast_inference:
model = AutoModelForCausalLM.from_pretrained(
model_name,
device_map = device_map,
torch_dtype = dtype,
# quantization_config = bnb_config,
token = token,
max_position_embeddings = max_position_embeddings,
trust_remote_code = trust_remote_code,
attn_implementation = "eager",
**kwargs,
)
else:
from unsloth_zoo.vllm_utils import (
load_vllm,
get_vllm_state_dict,
convert_vllm_to_huggingface,
generate_batches,
)
allowed_args = inspect.getfullargspec(load_vllm).args
load_vllm_kwargs = dict(
model_name = model_name,
config = model_config,
gpu_memory_utilization = gpu_memory_utilization,
max_seq_length = max_seq_length,
dtype = dtype,
disable_log_stats = disable_log_stats,
float8_kv_cache = float8_kv_cache,
enable_lora = True,
max_lora_rank = max_lora_rank,
disable_log_stats = disable_log_stats,
)
for allowed_arg in allowed_args:
if allowed_arg not in load_vllm_kwargs and allowed_arg in kwargs:
load_vllm_kwargs[allowed_arg] = kwargs[allowed_arg]
pass
# Load vLLM first
llm = load_vllm(**load_vllm_kwargs)
# Convert to HF format
_, quant_state_dict = get_vllm_state_dict(llm, config = model_config)
model = convert_vllm_to_huggingface(quant_state_dict, model_config, dtype)
model.vllm_engine = llm
model.fast_generate = model.vllm_engine.generate
from functools import partial
model.fast_generate_batches = partial(generate_batches, model.vllm_engine)
pass
# Return old flag
os.environ["HF_HUB_ENABLE_HF_TRANSFER"] = old_hf_transfer
# We currently only support NVIDIA GPUs - AMD / Intel is a work in progress!

View file

@ -30,11 +30,11 @@ except:
from huggingface_hub.utils._token import get_token
pass
from huggingface_hub import HfFileSystem
import importlib.util
# [TODO] Move USE_MODELSCOPE to utils
USE_MODELSCOPE = os.environ.get("UNSLOTH_USE_MODELSCOPE", "0") == "1"
if USE_MODELSCOPE:
import importlib
if importlib.util.find_spec("modelscope") is None:
raise ImportError(f'You are using the modelscope hub, please install modelscope by `pip install modelscope -U`')
pass
@ -73,9 +73,25 @@ class FastLanguageModel(FastLlamaModel):
resize_model_vocab = None,
revision = None,
use_exact_model_name = False,
fast_inference = False, # uses vLLM
gpu_memory_utilization = 0.5,
float8_kv_cache = True,
random_state = 3407,
max_lora_rank = 16,
disable_log_stats = False,
*args, **kwargs,
):
if token is None: token = get_token()
if fast_inference:
if importlib.util.find_spec("vllm") is None:
raise ImportError(
"Unsloth: Please install vLLM before enabling `fast_inference`!\n"\
"You can do this in a terminal via `pip install vllm`"
)
pass
pass
old_model_name = model_name
if not use_exact_model_name:
@ -255,6 +271,24 @@ class FastLanguageModel(FastLlamaModel):
tokenizer_name = None
pass
if fast_inference:
from unsloth_zoo.vllm_utils import (
patch_vllm,
vllm_dynamic_quant_supported,
)
patch_vllm()
if model_name.endswith("unsloth-bnb-4bit"):
if not vllm_dynamic_quant_supported(model_name, model_config):
# Instead use -bnb-4bit variant
print(
f"Unsloth: Switching from Unsloth dynamic quant to normal quant since\n"\
f"we do not yet support fast inference for {model_name}"
)
model_name = model_name[:-len("unsloth-bnb-4bit")] + "bnb-4bit"
pass
pass
pass
model, tokenizer = dispatch_model.from_pretrained(
model_name = model_name,
max_seq_length = max_seq_length,
@ -268,6 +302,13 @@ class FastLanguageModel(FastLlamaModel):
tokenizer_name = tokenizer_name,
trust_remote_code = trust_remote_code,
revision = revision if not is_peft else None,
fast_inference = fast_inference,
gpu_memory_utilization = gpu_memory_utilization,
float8_kv_cache = float8_kv_cache,
random_state = random_state,
max_lora_rank = max_lora_rank,
disable_log_stats = disable_log_stats,
*args, **kwargs,
)