Files
ray-project--ray/doc/source/llm/examples/batch/vllm-with-lora.ipynb
T
2026-07-13 13:17:40 +08:00

134 lines
4.9 KiB
Plaintext

{
"cells": [
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Batch Inference with LoRA Adapters\n",
"\n",
"In this example, we show how to perform batch inference using Ray Data LLM with LLM and a LoRA adapter. \n",
"\n",
"To run this example, we need to install the following dependencies:\n",
"\n",
"```bash\n",
"pip install -qU \"ray[llm]\"\n",
"```"
]
},
{
"cell_type": "code",
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
"import ray\n",
"from ray.data.llm import build_processor, vLLMEngineProcessorConfig\n",
"\n",
"# 1. Construct a vLLM processor config.\n",
"processor_config = vLLMEngineProcessorConfig(\n",
" # The base model.\n",
" model_source=\"unsloth/Llama-3.2-1B-Instruct\",\n",
" # vLLM engine config.\n",
" engine_kwargs=dict(\n",
" # Enable LoRA in the vLLM engine; otherwise you won't be able to\n",
" # process requests with LoRA adapters.\n",
" enable_lora=True,\n",
" # You need to set the LoRA rank for the adapter.\n",
" # The LoRA rank is the value of \"r\" in the LoRA config.\n",
" # If you want to use multiple LoRA adapters in this pipeline,\n",
" # please specify the maximum LoRA rank among all of them.\n",
" max_lora_rank=32,\n",
" # The maximum number of LoRA adapters vLLM cached. \"1\" means\n",
" # vLLM only caches one LoRA adapter at a time, so if your dataset\n",
" # needs more than one LoRA adapters, then there would be context\n",
" # switching. On the other hand, while increasing max_loras reduces\n",
" # the context switching, it increases the memory footprint.\n",
" max_loras=1,\n",
" # Older GPUs (e.g. T4) don't support bfloat16. You should remove\n",
" # this line if you're using later GPUs.\n",
" dtype=\"half\",\n",
" # Reduce the model length to fit small GPUs. You should remove\n",
" # this line if you're using large GPUs.\n",
" max_model_len=1024,\n",
" ),\n",
" # The batch size used in Ray Data.\n",
" batch_size=16,\n",
" # Use one GPU in this example.\n",
" concurrency=1,\n",
" # If you save the LoRA adapter in S3, you can set the following path.\n",
" # dynamic_lora_loading_path=\"s3://your-lora-bucket/\",\n",
")\n",
"\n",
"# 2. Construct a processor using the processor config.\n",
"processor = build_processor(\n",
" processor_config,\n",
" # Convert the input data to the OpenAI chat form.\n",
" preprocess=lambda row: dict(\n",
" # If you specify \"model\" in a request, and the model is different\n",
" # from the model you specify in the processor config, then this\n",
" # is the LoRA adapter. The \"model\" here can be a LoRA adapter\n",
" # available in the HuggingFace Hub or a local path.\n",
" #\n",
" # If you set dynamic_lora_loading_path, then only specify the LoRA\n",
" # path under dynamic_lora_loading_path.\n",
" model=\"EdBergJr/Llama32_Baha_3\",\n",
" messages=[\n",
" {\"role\": \"system\",\n",
" \"content\": \"You are a calculator. Please only output the answer \"\n",
" \"of the given equation.\"},\n",
" {\"role\": \"user\", \"content\": f\"{row['id']} ** 3 = ?\"},\n",
" ],\n",
" sampling_params=dict(\n",
" temperature=0.3,\n",
" max_tokens=20,\n",
" detokenize=False,\n",
" ),\n",
" ),\n",
" # Only keep the generated text in the output dataset.\n",
" postprocess=lambda row: {\n",
" \"resp\": row[\"generated_text\"],\n",
" },\n",
")\n",
"\n",
"# 3. Synthesize a dataset with 30 rows.\n",
"ds = ray.data.range(30)\n",
"# 4. Apply the processor to the dataset. Note that this line won't kick off\n",
"# anything because processor is execution lazily.\n",
"ds = processor(ds)\n",
"# Materialization kicks off the pipeline execution.\n",
"ds = ds.materialize()\n",
"\n",
"# 5. Print all outputs.\n",
"for out in ds.take_all():\n",
" print(out)\n",
" print(\"==========\")\n",
"\n",
"# 6. Shutdown Ray to release resources.\n",
"ray.shutdown()"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "base",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.11.11"
},
"orphan": true
},
"nbformat": 4,
"nbformat_minor": 2
}