chore: import upstream snapshot with attribution
This commit is contained in:
@@ -0,0 +1,93 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Example Python client for `examples/applications/api_server/server.py`
|
||||
Start the demo server:
|
||||
python examples/applications/api_server/server.py --model <model_name>
|
||||
|
||||
NOTE: The API server is used only for demonstration and simple performance
|
||||
benchmarks. It is not intended for production use.
|
||||
For production use, we recommend `vllm serve` and the OpenAI client API.
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
from argparse import Namespace
|
||||
from collections.abc import Iterable
|
||||
|
||||
import requests
|
||||
|
||||
|
||||
def clear_line(n: int = 1) -> None:
|
||||
LINE_UP = "\033[1A"
|
||||
LINE_CLEAR = "\x1b[2K"
|
||||
for _ in range(n):
|
||||
print(LINE_UP, end=LINE_CLEAR, flush=True)
|
||||
|
||||
|
||||
def post_http_request(
|
||||
prompt: str, api_url: str, n: int = 1, stream: bool = False
|
||||
) -> requests.Response:
|
||||
headers = {"User-Agent": "Test Client"}
|
||||
pload = {
|
||||
"prompt": prompt,
|
||||
"n": n,
|
||||
"temperature": 0.0,
|
||||
"max_tokens": 16,
|
||||
"stream": stream,
|
||||
}
|
||||
response = requests.post(api_url, headers=headers, json=pload, stream=stream)
|
||||
return response
|
||||
|
||||
|
||||
def get_streaming_response(response: requests.Response) -> Iterable[list[str]]:
|
||||
for chunk in response.iter_lines(
|
||||
chunk_size=8192, decode_unicode=False, delimiter=b"\n"
|
||||
):
|
||||
if chunk:
|
||||
data = json.loads(chunk.decode("utf-8"))
|
||||
output = data["text"]
|
||||
yield output
|
||||
|
||||
|
||||
def get_response(response: requests.Response) -> list[str]:
|
||||
data = json.loads(response.content)
|
||||
output = data["text"]
|
||||
return output
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--host", type=str, default="localhost")
|
||||
parser.add_argument("--port", type=int, default=8000)
|
||||
parser.add_argument("--n", type=int, default=1)
|
||||
parser.add_argument("--prompt", type=str, default="San Francisco is a")
|
||||
parser.add_argument("--stream", action="store_true")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main(args: Namespace):
|
||||
prompt = args.prompt
|
||||
api_url = f"http://{args.host}:{args.port}/generate"
|
||||
n = args.n
|
||||
stream = args.stream
|
||||
|
||||
print(f"Prompt: {prompt!r}\n", flush=True)
|
||||
response = post_http_request(prompt, api_url, n, stream)
|
||||
|
||||
if stream:
|
||||
num_printed_lines = 0
|
||||
for h in get_streaming_response(response):
|
||||
clear_line(num_printed_lines)
|
||||
num_printed_lines = 0
|
||||
for i, line in enumerate(h):
|
||||
num_printed_lines += 1
|
||||
print(f"Beam candidate {i}: {line!r}", flush=True)
|
||||
else:
|
||||
output = get_response(response)
|
||||
for i, line in enumerate(output):
|
||||
print(f"Beam candidate {i}: {line!r}", flush=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,186 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
NOTE: This API server is used only for demonstrating usage of AsyncEngine
|
||||
and simple performance benchmarks. It is not intended for production use.
|
||||
For production use, we recommend using our OpenAI compatible server.
|
||||
We are also not going to accept PRs modifying this file, please
|
||||
change `vllm/entrypoints/openai/api_server.py` instead.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import ssl
|
||||
from argparse import Namespace
|
||||
from collections.abc import AsyncGenerator
|
||||
from typing import Any
|
||||
|
||||
from fastapi import FastAPI, Request
|
||||
from fastapi.responses import JSONResponse, Response, StreamingResponse
|
||||
|
||||
import vllm.envs as envs
|
||||
from vllm.engine.arg_utils import AsyncEngineArgs
|
||||
from vllm.engine.async_llm_engine import AsyncLLMEngine
|
||||
from vllm.entrypoints.launcher import serve_http
|
||||
from vllm.entrypoints.serve.utils.api_utils import with_cancellation
|
||||
from vllm.logger import init_logger
|
||||
from vllm.sampling_params import SamplingParams
|
||||
from vllm.usage.usage_lib import UsageContext
|
||||
from vllm.utils import random_uuid
|
||||
from vllm.utils.argparse_utils import FlexibleArgumentParser
|
||||
from vllm.utils.system_utils import set_ulimit
|
||||
from vllm.version import __version__ as VLLM_VERSION
|
||||
|
||||
logger = init_logger("api_server")
|
||||
|
||||
app = FastAPI()
|
||||
engine = None
|
||||
|
||||
|
||||
@app.get("/health")
|
||||
async def health() -> Response:
|
||||
"""Health check."""
|
||||
return Response(status_code=200)
|
||||
|
||||
|
||||
@app.post("/generate")
|
||||
async def generate(request: Request) -> Response:
|
||||
"""Generate completion for the request.
|
||||
|
||||
The request should be a JSON object with the following fields:
|
||||
- prompt: the prompt to use for the generation.
|
||||
- stream: whether to stream the results or not.
|
||||
- other fields: the sampling parameters (See `SamplingParams` for details).
|
||||
"""
|
||||
request_dict = await request.json()
|
||||
return await _generate(request_dict, raw_request=request)
|
||||
|
||||
|
||||
@with_cancellation
|
||||
async def _generate(request_dict: dict, raw_request: Request) -> Response:
|
||||
prompt = request_dict.pop("prompt")
|
||||
stream = request_dict.pop("stream", False)
|
||||
# Since SamplingParams is created fresh per request, safe to skip clone
|
||||
sampling_params = SamplingParams(**request_dict, skip_clone=True)
|
||||
request_id = random_uuid()
|
||||
|
||||
assert engine is not None
|
||||
results_generator = engine.generate(prompt, sampling_params, request_id)
|
||||
|
||||
# Streaming case
|
||||
async def stream_results() -> AsyncGenerator[bytes, None]:
|
||||
async for request_output in results_generator:
|
||||
prompt = request_output.prompt
|
||||
assert prompt is not None
|
||||
text_outputs = [prompt + output.text for output in request_output.outputs]
|
||||
ret = {"text": text_outputs}
|
||||
yield (json.dumps(ret) + "\n").encode("utf-8")
|
||||
|
||||
if stream:
|
||||
return StreamingResponse(stream_results())
|
||||
|
||||
# Non-streaming case
|
||||
final_output = None
|
||||
try:
|
||||
async for request_output in results_generator:
|
||||
final_output = request_output
|
||||
except asyncio.CancelledError:
|
||||
return Response(status_code=499)
|
||||
|
||||
assert final_output is not None
|
||||
prompt = final_output.prompt
|
||||
assert prompt is not None
|
||||
text_outputs = [prompt + output.text for output in final_output.outputs]
|
||||
ret = {"text": text_outputs}
|
||||
return JSONResponse(ret)
|
||||
|
||||
|
||||
def build_app(args: Namespace) -> FastAPI:
|
||||
global app
|
||||
|
||||
app.root_path = args.root_path
|
||||
return app
|
||||
|
||||
|
||||
async def init_app(
|
||||
args: Namespace,
|
||||
llm_engine: AsyncLLMEngine | None = None,
|
||||
) -> FastAPI:
|
||||
app = build_app(args)
|
||||
|
||||
global engine
|
||||
|
||||
engine_args = AsyncEngineArgs.from_cli_args(args)
|
||||
engine = (
|
||||
llm_engine
|
||||
if llm_engine is not None
|
||||
else AsyncLLMEngine.from_engine_args(
|
||||
engine_args, usage_context=UsageContext.API_SERVER
|
||||
)
|
||||
)
|
||||
app.state.engine_client = engine
|
||||
app.state.args = args
|
||||
return app
|
||||
|
||||
|
||||
async def run_server(
|
||||
args: Namespace, llm_engine: AsyncLLMEngine | None = None, **uvicorn_kwargs: Any
|
||||
) -> None:
|
||||
logger.info("vLLM API server version %s", VLLM_VERSION)
|
||||
logger.info("args: %s", args)
|
||||
|
||||
set_ulimit()
|
||||
|
||||
app = await init_app(args, llm_engine)
|
||||
assert engine is not None
|
||||
|
||||
shutdown_task = await serve_http(
|
||||
app,
|
||||
sock=None,
|
||||
enable_ssl_refresh=args.enable_ssl_refresh,
|
||||
host=args.host,
|
||||
port=args.port,
|
||||
log_level=args.log_level,
|
||||
timeout_keep_alive=envs.VLLM_HTTP_TIMEOUT_KEEP_ALIVE,
|
||||
ssl_keyfile=args.ssl_keyfile,
|
||||
ssl_certfile=args.ssl_certfile,
|
||||
ssl_ca_certs=args.ssl_ca_certs,
|
||||
ssl_cert_reqs=args.ssl_cert_reqs,
|
||||
**uvicorn_kwargs,
|
||||
)
|
||||
|
||||
await shutdown_task
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
parser = FlexibleArgumentParser()
|
||||
parser.add_argument("--host", type=str, default=None)
|
||||
parser.add_argument("--port", type=parser.check_port, default=8000)
|
||||
parser.add_argument("--ssl-keyfile", type=str, default=None)
|
||||
parser.add_argument("--ssl-certfile", type=str, default=None)
|
||||
parser.add_argument(
|
||||
"--ssl-ca-certs", type=str, default=None, help="The CA certificates file"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--enable-ssl-refresh",
|
||||
action="store_true",
|
||||
default=False,
|
||||
help="Refresh SSL Context when SSL certificate files change",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--ssl-cert-reqs",
|
||||
type=int,
|
||||
default=int(ssl.CERT_NONE),
|
||||
help="Whether client certificate is required (see stdlib ssl module's)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--root-path",
|
||||
type=str,
|
||||
default=None,
|
||||
help="FastAPI root_path when app is behind a path based routing proxy",
|
||||
)
|
||||
parser.add_argument("--log-level", type=str, default="debug")
|
||||
parser = AsyncEngineArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
|
||||
asyncio.run(run_server(args))
|
||||
@@ -0,0 +1,112 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Example for starting a Gradio OpenAI Chatbot Webserver
|
||||
Start vLLM API server:
|
||||
vllm serve meta-llama/Llama-2-7b-chat-hf
|
||||
|
||||
Start Gradio OpenAI Chatbot Webserver:
|
||||
python examples/applications/chatbot/gradio_openai_chatbot_webserver.py \
|
||||
-m meta-llama/Llama-2-7b-chat-hf
|
||||
|
||||
Note that `pip install --upgrade gradio` is needed to run this example.
|
||||
More details: https://github.com/gradio-app/gradio
|
||||
|
||||
If your antivirus software blocks the download of frpc for gradio,
|
||||
you can install it manually by following these steps:
|
||||
|
||||
1. Download this file: https://cdn-media.huggingface.co/frpc-gradio-0.3/frpc_linux_amd64
|
||||
2. Rename the downloaded file to: frpc_linux_amd64_v0.3
|
||||
3. Move the file to this location: /home/user/.cache/huggingface/gradio/frpc
|
||||
"""
|
||||
|
||||
import argparse
|
||||
|
||||
import gradio as gr
|
||||
from openai import OpenAI
|
||||
|
||||
|
||||
def predict(message, history, client, model_name, temp, stop_token_ids):
|
||||
messages = [
|
||||
{"role": "system", "content": "You are a great AI assistant."},
|
||||
*history,
|
||||
{"role": "user", "content": message},
|
||||
]
|
||||
|
||||
# Send request to OpenAI API (vLLM server)
|
||||
stream = client.chat.completions.create(
|
||||
model=model_name,
|
||||
messages=messages,
|
||||
temperature=temp,
|
||||
stream=True,
|
||||
extra_body={
|
||||
"repetition_penalty": 1,
|
||||
"stop_token_ids": [int(id.strip()) for id in stop_token_ids.split(",")]
|
||||
if stop_token_ids
|
||||
else [],
|
||||
},
|
||||
)
|
||||
|
||||
# Collect all chunks and concatenate them into a full message
|
||||
full_message = ""
|
||||
for chunk in stream:
|
||||
full_message += chunk.choices[0].delta.content or ""
|
||||
|
||||
# Return the full message as a single response
|
||||
return full_message
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Chatbot Interface with Customizable Parameters"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model-url", type=str, default="http://localhost:8000/v1", help="Model URL"
|
||||
)
|
||||
parser.add_argument(
|
||||
"-m", "--model", type=str, required=True, help="Model name for the chatbot"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--temp", type=float, default=0.8, help="Temperature for text generation"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--stop-token-ids", type=str, default="", help="Comma-separated stop token IDs"
|
||||
)
|
||||
parser.add_argument("--host", type=str, default=None)
|
||||
parser.add_argument("--port", type=int, default=8001)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def build_gradio_interface(client, model_name, temp, stop_token_ids):
|
||||
def chat_predict(message, history):
|
||||
return predict(message, history, client, model_name, temp, stop_token_ids)
|
||||
|
||||
return gr.ChatInterface(
|
||||
fn=chat_predict,
|
||||
title="Chatbot Interface",
|
||||
description="A simple chatbot powered by vLLM",
|
||||
)
|
||||
|
||||
|
||||
def main():
|
||||
# Parse the arguments
|
||||
args = parse_args()
|
||||
|
||||
# Set OpenAI's API key and API base to use vLLM's API server
|
||||
openai_api_key = "EMPTY"
|
||||
openai_api_base = args.model_url
|
||||
|
||||
# Create an OpenAI client
|
||||
client = OpenAI(api_key=openai_api_key, base_url=openai_api_base)
|
||||
|
||||
# Define the Gradio chatbot interface using the predict function
|
||||
gradio_interface = build_gradio_interface(
|
||||
client, args.model, args.temp, args.stop_token_ids
|
||||
)
|
||||
|
||||
gradio_interface.queue().launch(
|
||||
server_name=args.host, server_port=args.port, share=True
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,75 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""Example for starting a Gradio Webserver
|
||||
Start vLLM API server:
|
||||
python examples/applications/api_server/server.py \
|
||||
--model meta-llama/Llama-2-7b-chat-hf
|
||||
|
||||
Start Webserver:
|
||||
python examples/applications/chatbot/gradio_webserver.py
|
||||
|
||||
Note that `pip install --upgrade gradio` is needed to run this example.
|
||||
More details: https://github.com/gradio-app/gradio
|
||||
|
||||
If your antivirus software blocks the download of frpc for gradio,
|
||||
you can install it manually by following these steps:
|
||||
|
||||
1. Download this file: https://cdn-media.huggingface.co/frpc-gradio-0.3/frpc_linux_amd64
|
||||
2. Rename the downloaded file to: frpc_linux_amd64_v0.3
|
||||
3. Move the file to this location: /home/user/.cache/huggingface/gradio/frpc
|
||||
"""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
|
||||
import gradio as gr
|
||||
import requests
|
||||
|
||||
|
||||
def http_bot(prompt):
|
||||
headers = {"User-Agent": "vLLM Client"}
|
||||
pload = {
|
||||
"prompt": prompt,
|
||||
"stream": True,
|
||||
"max_tokens": 128,
|
||||
}
|
||||
response = requests.post(args.model_url, headers=headers, json=pload, stream=True)
|
||||
|
||||
for chunk in response.iter_lines(
|
||||
chunk_size=8192, decode_unicode=False, delimiter=b"\n"
|
||||
):
|
||||
if chunk:
|
||||
data = json.loads(chunk.decode("utf-8"))
|
||||
output = data["text"][0]
|
||||
yield output
|
||||
|
||||
|
||||
def build_demo():
|
||||
with gr.Blocks() as demo:
|
||||
gr.Markdown("# vLLM text completion demo\n")
|
||||
inputbox = gr.Textbox(label="Input", placeholder="Enter text and press ENTER")
|
||||
outputbox = gr.Textbox(
|
||||
label="Output", placeholder="Generated result from the model"
|
||||
)
|
||||
inputbox.submit(http_bot, [inputbox], [outputbox])
|
||||
return demo
|
||||
|
||||
|
||||
def parse_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--host", type=str, default=None)
|
||||
parser.add_argument("--port", type=int, default=8001)
|
||||
parser.add_argument(
|
||||
"--model-url", type=str, default="http://localhost:8000/generate"
|
||||
)
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def main(args):
|
||||
demo = build_demo()
|
||||
demo.queue().launch(server_name=args.host, server_port=args.port, share=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
args = parse_args()
|
||||
main(args)
|
||||
@@ -0,0 +1,311 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
vLLM Chat Assistant - A Streamlit Web Interface
|
||||
|
||||
A streamlined chat interface that quickly integrates
|
||||
with vLLM API server.
|
||||
|
||||
Features:
|
||||
- Multiple chat sessions management
|
||||
- Streaming response display
|
||||
- Configurable API endpoint
|
||||
- Real-time chat history
|
||||
- Reasoning Display: Optional thinking process visualization
|
||||
|
||||
Requirements:
|
||||
pip install streamlit openai
|
||||
|
||||
Usage:
|
||||
# Start the app with default settings
|
||||
streamlit run streamlit_openai_chatbot_webserver.py
|
||||
|
||||
# Start with custom vLLM API endpoint
|
||||
VLLM_API_BASE="http://your-server:8000/v1" \
|
||||
streamlit run streamlit_openai_chatbot_webserver.py
|
||||
|
||||
# Enable debug mode
|
||||
streamlit run streamlit_openai_chatbot_webserver.py \
|
||||
--logger.level=debug
|
||||
"""
|
||||
|
||||
import os
|
||||
from datetime import datetime
|
||||
|
||||
import streamlit as st
|
||||
from openai import OpenAI
|
||||
|
||||
# Get command line arguments from environment variables
|
||||
openai_api_key = os.getenv("VLLM_API_KEY", "EMPTY")
|
||||
openai_api_base = os.getenv("VLLM_API_BASE", "http://localhost:8000/v1")
|
||||
|
||||
# Initialize session states for managing chat sessions
|
||||
if "sessions" not in st.session_state:
|
||||
st.session_state.sessions = {}
|
||||
|
||||
if "current_session" not in st.session_state:
|
||||
st.session_state.current_session = None
|
||||
|
||||
if "messages" not in st.session_state:
|
||||
st.session_state.messages = []
|
||||
|
||||
if "active_session" not in st.session_state:
|
||||
st.session_state.active_session = None
|
||||
|
||||
# Add new session state for reasoning
|
||||
if "show_reasoning" not in st.session_state:
|
||||
st.session_state.show_reasoning = {}
|
||||
|
||||
# Initialize session state for API base URL
|
||||
if "api_base_url" not in st.session_state:
|
||||
st.session_state.api_base_url = openai_api_base
|
||||
|
||||
|
||||
def create_new_chat_session():
|
||||
"""Create a new chat session with timestamp as unique identifier.
|
||||
|
||||
This function initializes a new chat session by:
|
||||
1. Generating a timestamp-based session ID
|
||||
2. Creating an empty message list for the new session
|
||||
3. Setting the new session as both current and active session
|
||||
4. Resetting the messages list for the new session
|
||||
|
||||
Returns:
|
||||
None
|
||||
|
||||
Session State Updates:
|
||||
- sessions: Adds new empty message list with timestamp key
|
||||
- current_session: Sets to new session ID
|
||||
- active_session: Sets to new session ID
|
||||
- messages: Resets to empty list
|
||||
"""
|
||||
session_id = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
st.session_state.sessions[session_id] = []
|
||||
st.session_state.current_session = session_id
|
||||
st.session_state.active_session = session_id
|
||||
st.session_state.messages = []
|
||||
|
||||
|
||||
def switch_to_chat_session(session_id):
|
||||
"""Switch the active chat context to a different session.
|
||||
|
||||
Args:
|
||||
session_id (str): The timestamp ID of the session to switch to
|
||||
|
||||
This function handles chat session switching by:
|
||||
1. Setting the specified session as current
|
||||
2. Updating the active session marker
|
||||
3. Loading the messages history from the specified session
|
||||
|
||||
Session State Updates:
|
||||
- current_session: Updated to specified session_id
|
||||
- active_session: Updated to specified session_id
|
||||
- messages: Loaded from sessions[session_id]
|
||||
"""
|
||||
st.session_state.current_session = session_id
|
||||
st.session_state.active_session = session_id
|
||||
st.session_state.messages = st.session_state.sessions[session_id]
|
||||
|
||||
|
||||
def get_llm_response(messages, model, reason, content_ph=None, reasoning_ph=None):
|
||||
"""Generate and stream LLM response with optional reasoning process.
|
||||
|
||||
Args:
|
||||
messages (list): List of conversation message dicts with 'role' and 'content'
|
||||
model (str): The model identifier to use for generation
|
||||
reason (bool): Whether to enable and display reasoning process
|
||||
content_ph (streamlit.empty): Placeholder for streaming response content
|
||||
reasoning_ph (streamlit.empty): Placeholder for streaming reasoning process
|
||||
|
||||
Returns:
|
||||
tuple: (str, str)
|
||||
- First string contains the complete response text
|
||||
- Second string contains the complete reasoning text (if enabled)
|
||||
|
||||
Features:
|
||||
- Streams both reasoning and response text in real-time
|
||||
- Handles model API errors gracefully
|
||||
- Supports live updating of thinking process
|
||||
- Maintains separate content and reasoning displays
|
||||
|
||||
Raises:
|
||||
Exception: Wrapped in error message if API call fails
|
||||
|
||||
Note:
|
||||
The function uses streamlit placeholders for live updates.
|
||||
When reason=True, the reasoning process appears above the response.
|
||||
"""
|
||||
full_text = ""
|
||||
think_text = ""
|
||||
live_think = None
|
||||
# Build request parameters
|
||||
params = {"model": model, "messages": messages, "stream": True}
|
||||
if reason:
|
||||
params["extra_body"] = {"chat_template_kwargs": {"enable_thinking": True}}
|
||||
|
||||
try:
|
||||
response = client.chat.completions.create(**params)
|
||||
if isinstance(response, str):
|
||||
if content_ph:
|
||||
content_ph.markdown(response)
|
||||
return response, ""
|
||||
|
||||
# Prepare reasoning expander above content
|
||||
if reason and reasoning_ph:
|
||||
exp = reasoning_ph.expander("💭 Thinking Process (live)", expanded=True)
|
||||
live_think = exp.empty()
|
||||
|
||||
# Stream chunks
|
||||
for chunk in response:
|
||||
delta = chunk.choices[0].delta
|
||||
# Stream reasoning first
|
||||
if reason and hasattr(delta, "reasoning") and live_think:
|
||||
rc = delta.reasoning
|
||||
if rc:
|
||||
think_text += rc
|
||||
live_think.markdown(think_text + "▌")
|
||||
# Then stream content
|
||||
if hasattr(delta, "content") and delta.content and content_ph:
|
||||
full_text += delta.content
|
||||
content_ph.markdown(full_text + "▌")
|
||||
|
||||
# Finalize displays: reasoning remains above, content below
|
||||
if reason and live_think:
|
||||
live_think.markdown(think_text)
|
||||
if content_ph:
|
||||
content_ph.markdown(full_text)
|
||||
|
||||
return full_text, think_text
|
||||
except Exception as e:
|
||||
st.error(f"Error details: {str(e)}")
|
||||
return f"Error: {str(e)}", ""
|
||||
|
||||
|
||||
# Sidebar - API Settings first
|
||||
st.sidebar.title("API Settings")
|
||||
new_api_base = st.sidebar.text_input(
|
||||
"API Base URL:", value=st.session_state.api_base_url
|
||||
)
|
||||
if new_api_base != st.session_state.api_base_url:
|
||||
st.session_state.api_base_url = new_api_base
|
||||
st.rerun()
|
||||
|
||||
st.sidebar.divider()
|
||||
|
||||
# Sidebar - Session Management
|
||||
st.sidebar.title("Chat Sessions")
|
||||
if st.sidebar.button("New Session"):
|
||||
create_new_chat_session()
|
||||
|
||||
|
||||
# Display all sessions in reverse chronological order
|
||||
for session_id in sorted(st.session_state.sessions.keys(), reverse=True):
|
||||
# Mark the active session with a pinned button
|
||||
if session_id == st.session_state.active_session:
|
||||
st.sidebar.button(
|
||||
f"📍 {session_id}",
|
||||
key=session_id,
|
||||
type="primary",
|
||||
on_click=switch_to_chat_session,
|
||||
args=(session_id,),
|
||||
)
|
||||
else:
|
||||
st.sidebar.button(
|
||||
f"Session {session_id}",
|
||||
key=session_id,
|
||||
on_click=switch_to_chat_session,
|
||||
args=(session_id,),
|
||||
)
|
||||
|
||||
# Main interface
|
||||
st.title("vLLM Chat Assistant")
|
||||
|
||||
# Initialize OpenAI client with API settings
|
||||
client = OpenAI(api_key=openai_api_key, base_url=st.session_state.api_base_url)
|
||||
|
||||
# Get and display current model id
|
||||
models = client.models.list()
|
||||
model = models.data[0].id
|
||||
st.markdown(f"**Model**: {model}")
|
||||
|
||||
# Initialize first session if none exists
|
||||
if st.session_state.current_session is None:
|
||||
create_new_chat_session()
|
||||
st.session_state.active_session = st.session_state.current_session
|
||||
|
||||
# Update the chat history display section
|
||||
for idx, msg in enumerate(st.session_state.messages):
|
||||
# Render user messages normally
|
||||
if msg["role"] == "user":
|
||||
with st.chat_message("user"):
|
||||
st.write(msg["content"])
|
||||
# Render assistant messages with reasoning above
|
||||
else:
|
||||
# If reasoning exists for this assistant message, show it above the content
|
||||
if idx in st.session_state.show_reasoning:
|
||||
with st.expander("💭 Thinking Process", expanded=False):
|
||||
st.markdown(st.session_state.show_reasoning[idx])
|
||||
with st.chat_message("assistant"):
|
||||
st.write(msg["content"])
|
||||
|
||||
|
||||
# Setup & Cache reasoning support check
|
||||
@st.cache_data(show_spinner=False)
|
||||
def server_supports_reasoning():
|
||||
"""Check if the current model supports reasoning capability.
|
||||
|
||||
Returns:
|
||||
bool: True if the model supports reasoning, False otherwise
|
||||
"""
|
||||
resp = client.chat.completions.create(
|
||||
model=model,
|
||||
messages=[{"role": "user", "content": "Hi"}],
|
||||
stream=False,
|
||||
)
|
||||
return hasattr(resp.choices[0].message, "reasoning") and bool(
|
||||
resp.choices[0].message.reasoning
|
||||
)
|
||||
|
||||
|
||||
# Check support
|
||||
supports_reasoning = server_supports_reasoning()
|
||||
|
||||
# Add reasoning toggle in sidebar if supported
|
||||
reason = False # Default to False
|
||||
if supports_reasoning:
|
||||
reason = st.sidebar.checkbox("Enable Reasoning", value=False)
|
||||
else:
|
||||
st.sidebar.markdown(
|
||||
"<span style='color:gray;'>Reasoning unavailable for this model.</span>",
|
||||
unsafe_allow_html=True,
|
||||
)
|
||||
# reason remains False
|
||||
|
||||
# Update the input handling section
|
||||
if prompt := st.chat_input("Type your message here..."):
|
||||
# Save and display user message
|
||||
st.session_state.messages.append({"role": "user", "content": prompt})
|
||||
st.session_state.sessions[st.session_state.current_session] = (
|
||||
st.session_state.messages
|
||||
)
|
||||
with st.chat_message("user"):
|
||||
st.write(prompt)
|
||||
|
||||
# Prepare LLM messages
|
||||
msgs = [
|
||||
{"role": m["role"], "content": m["content"]} for m in st.session_state.messages
|
||||
]
|
||||
|
||||
# Stream assistant response
|
||||
with st.chat_message("assistant"):
|
||||
# Placeholders: reasoning above, content below
|
||||
reason_ph = st.empty()
|
||||
content_ph = st.empty()
|
||||
full, think = get_llm_response(msgs, model, reason, content_ph, reason_ph)
|
||||
# Determine index for this new assistant message
|
||||
message_index = len(st.session_state.messages)
|
||||
# Save assistant reply
|
||||
st.session_state.messages.append({"role": "assistant", "content": full})
|
||||
# Persist reasoning in session state if any
|
||||
if reason and think:
|
||||
st.session_state.show_reasoning[message_index] = think
|
||||
@@ -0,0 +1,257 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
Retrieval Augmented Generation (RAG) Implementation with Langchain
|
||||
==================================================================
|
||||
|
||||
This script demonstrates a RAG implementation using LangChain, Milvus
|
||||
and vLLM. RAG enhances LLM responses by retrieving relevant context
|
||||
from a document collection.
|
||||
|
||||
Features:
|
||||
- Web content loading and chunking
|
||||
- Vector storage with Milvus
|
||||
- Embedding generation with vLLM
|
||||
- Question answering with context
|
||||
|
||||
Prerequisites:
|
||||
1. Install dependencies:
|
||||
pip install -U vllm \
|
||||
langchain_milvus langchain_openai \
|
||||
langchain_community beautifulsoup4 \
|
||||
langchain-text-splitters
|
||||
|
||||
2. Start services:
|
||||
# Start embedding service (port 8000)
|
||||
vllm serve ssmits/Qwen2-7B-Instruct-embed-base
|
||||
|
||||
# Start chat service (port 8001)
|
||||
vllm serve qwen/Qwen1.5-0.5B-Chat --port 8001
|
||||
|
||||
Usage:
|
||||
python retrieval_augmented_generation_with_langchain.py
|
||||
|
||||
Notes:
|
||||
- Ensure both vLLM services are running before executing
|
||||
- Default ports: 8000 (embedding), 8001 (chat)
|
||||
- First run may take time to download models
|
||||
"""
|
||||
|
||||
import argparse
|
||||
from argparse import Namespace
|
||||
from typing import Any
|
||||
|
||||
from langchain_community.document_loaders import WebBaseLoader
|
||||
from langchain_core.documents import Document
|
||||
from langchain_core.output_parsers import StrOutputParser
|
||||
from langchain_core.prompts import PromptTemplate
|
||||
from langchain_core.runnables import RunnablePassthrough
|
||||
from langchain_milvus import Milvus
|
||||
from langchain_openai import ChatOpenAI, OpenAIEmbeddings
|
||||
from langchain_text_splitters import RecursiveCharacterTextSplitter
|
||||
|
||||
|
||||
def load_and_split_documents(config: dict[str, Any]):
|
||||
"""
|
||||
Load and split documents from web URL
|
||||
"""
|
||||
try:
|
||||
loader = WebBaseLoader(web_paths=(config["url"],))
|
||||
docs = loader.load()
|
||||
|
||||
text_splitter = RecursiveCharacterTextSplitter(
|
||||
chunk_size=config["chunk_size"],
|
||||
chunk_overlap=config["chunk_overlap"],
|
||||
)
|
||||
return text_splitter.split_documents(docs)
|
||||
except Exception as e:
|
||||
print(f"Error loading document from {config['url']}: {str(e)}")
|
||||
raise
|
||||
|
||||
|
||||
def init_vectorstore(config: dict[str, Any], documents: list[Document]):
|
||||
"""
|
||||
Initialize vector store with documents
|
||||
"""
|
||||
return Milvus.from_documents(
|
||||
documents=documents,
|
||||
embedding=OpenAIEmbeddings(
|
||||
model=config["embedding_model"],
|
||||
openai_api_key=config["vllm_api_key"],
|
||||
openai_api_base=config["vllm_embedding_endpoint"],
|
||||
),
|
||||
connection_args={"uri": config["uri"]},
|
||||
drop_old=True,
|
||||
)
|
||||
|
||||
|
||||
def init_llm(config: dict[str, Any]):
|
||||
"""
|
||||
Initialize llm
|
||||
"""
|
||||
return ChatOpenAI(
|
||||
model=config["chat_model"],
|
||||
openai_api_key=config["vllm_api_key"],
|
||||
openai_api_base=config["vllm_chat_endpoint"],
|
||||
)
|
||||
|
||||
|
||||
def get_qa_prompt():
|
||||
"""
|
||||
Get question answering prompt template
|
||||
"""
|
||||
template = """You are an assistant for question-answering tasks.
|
||||
Use the following pieces of retrieved context to answer the question.
|
||||
If you don't know the answer, just say that you don't know.
|
||||
Use three sentences maximum and keep the answer concise.
|
||||
Question: {question}
|
||||
Context: {context}
|
||||
Answer:
|
||||
"""
|
||||
return PromptTemplate.from_template(template)
|
||||
|
||||
|
||||
def format_docs(docs: list[Document]):
|
||||
"""
|
||||
Format documents for prompt
|
||||
"""
|
||||
return "\n\n".join(doc.page_content for doc in docs)
|
||||
|
||||
|
||||
def create_qa_chain(retriever: Any, llm: ChatOpenAI, prompt: PromptTemplate):
|
||||
"""
|
||||
Set up question answering chain
|
||||
"""
|
||||
return (
|
||||
{
|
||||
"context": retriever | format_docs,
|
||||
"question": RunnablePassthrough(),
|
||||
}
|
||||
| prompt
|
||||
| llm
|
||||
| StrOutputParser()
|
||||
)
|
||||
|
||||
|
||||
def get_parser() -> argparse.ArgumentParser:
|
||||
"""
|
||||
Parse command line arguments
|
||||
"""
|
||||
parser = argparse.ArgumentParser(description="RAG with vLLM and langchain")
|
||||
|
||||
# Add command line arguments
|
||||
parser.add_argument(
|
||||
"--vllm-api-key", default="EMPTY", help="API key for vLLM compatible services"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vllm-embedding-endpoint",
|
||||
default="http://localhost:8000/v1",
|
||||
help="Base URL for embedding service",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vllm-chat-endpoint",
|
||||
default="http://localhost:8001/v1",
|
||||
help="Base URL for chat service",
|
||||
)
|
||||
parser.add_argument("--uri", default="./milvus.db", help="URI for Milvus database")
|
||||
parser.add_argument(
|
||||
"--url",
|
||||
default=("https://docs.vllm.ai/en/latest/getting_started/quickstart.html"),
|
||||
help="URL of the document to process",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--embedding-model",
|
||||
default="ssmits/Qwen2-7B-Instruct-embed-base",
|
||||
help="Model name for embeddings",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--chat-model", default="qwen/Qwen1.5-0.5B-Chat", help="Model name for chat"
|
||||
)
|
||||
parser.add_argument(
|
||||
"-i", "--interactive", action="store_true", help="Enable interactive Q&A mode"
|
||||
)
|
||||
parser.add_argument(
|
||||
"-k", "--top-k", type=int, default=3, help="Number of top results to retrieve"
|
||||
)
|
||||
parser.add_argument(
|
||||
"-c",
|
||||
"--chunk-size",
|
||||
type=int,
|
||||
default=1000,
|
||||
help="Chunk size for document splitting",
|
||||
)
|
||||
parser.add_argument(
|
||||
"-o",
|
||||
"--chunk-overlap",
|
||||
type=int,
|
||||
default=200,
|
||||
help="Chunk overlap for document splitting",
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
def init_config(args: Namespace):
|
||||
"""
|
||||
Initialize configuration settings from command line arguments
|
||||
"""
|
||||
|
||||
return {
|
||||
"vllm_api_key": args.vllm_api_key,
|
||||
"vllm_embedding_endpoint": args.vllm_embedding_endpoint,
|
||||
"vllm_chat_endpoint": args.vllm_chat_endpoint,
|
||||
"uri": args.uri,
|
||||
"embedding_model": args.embedding_model,
|
||||
"chat_model": args.chat_model,
|
||||
"url": args.url,
|
||||
"chunk_size": args.chunk_size,
|
||||
"chunk_overlap": args.chunk_overlap,
|
||||
"top_k": args.top_k,
|
||||
}
|
||||
|
||||
|
||||
def main():
|
||||
# Parse command line arguments
|
||||
args = get_parser().parse_args()
|
||||
|
||||
# Initialize configuration
|
||||
config = init_config(args)
|
||||
|
||||
# Load and split documents
|
||||
documents = load_and_split_documents(config)
|
||||
|
||||
# Initialize vector store and retriever
|
||||
vectorstore = init_vectorstore(config, documents)
|
||||
retriever = vectorstore.as_retriever(search_kwargs={"k": config["top_k"]})
|
||||
|
||||
# Initialize llm and prompt
|
||||
llm = init_llm(config)
|
||||
prompt = get_qa_prompt()
|
||||
|
||||
# Set up QA chain
|
||||
qa_chain = create_qa_chain(retriever, llm, prompt)
|
||||
|
||||
# Interactive mode
|
||||
if args.interactive:
|
||||
print("\nWelcome to Interactive Q&A System!")
|
||||
print("Enter 'q' or 'quit' to exit.")
|
||||
|
||||
while True:
|
||||
question = input("\nPlease enter your question: ")
|
||||
if question.lower() in ["q", "quit"]:
|
||||
print("\nThank you for using! Goodbye!")
|
||||
break
|
||||
|
||||
output = qa_chain.invoke(question)
|
||||
print(output)
|
||||
else:
|
||||
# Default single question mode
|
||||
question = "How to install vLLM?"
|
||||
output = qa_chain.invoke(question)
|
||||
print("-" * 50)
|
||||
print(output)
|
||||
print("-" * 50)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -0,0 +1,225 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# SPDX-FileCopyrightText: Copyright contributors to the vLLM project
|
||||
"""
|
||||
RAG (Retrieval Augmented Generation) Implementation with LlamaIndex
|
||||
================================================================
|
||||
|
||||
This script demonstrates a RAG system using:
|
||||
- LlamaIndex: For document indexing and retrieval
|
||||
- Milvus: As vector store backend
|
||||
- vLLM: For embedding and text generation
|
||||
|
||||
Features:
|
||||
1. Document Loading & Processing
|
||||
2. Embedding & Storage
|
||||
3. Query Processing
|
||||
|
||||
Requirements:
|
||||
1. Install dependencies:
|
||||
pip install llama-index llama-index-readers-web \
|
||||
llama-index-llms-openai-like \
|
||||
llama-index-embeddings-openai-like \
|
||||
llama-index-vector-stores-milvus \
|
||||
|
||||
2. Start services:
|
||||
# Start embedding service (port 8000)
|
||||
vllm serve ssmits/Qwen2-7B-Instruct-embed-base
|
||||
|
||||
# Start chat service (port 8001)
|
||||
vllm serve qwen/Qwen1.5-0.5B-Chat --port 8001
|
||||
|
||||
Usage:
|
||||
python retrieval_augmented_generation_with_llamaindex.py
|
||||
|
||||
Notes:
|
||||
- Ensure both vLLM services are running before executing
|
||||
- Default ports: 8000 (embedding), 8001 (chat)
|
||||
- First run may take time to download models
|
||||
"""
|
||||
|
||||
import argparse
|
||||
from argparse import Namespace
|
||||
from typing import Any
|
||||
|
||||
from llama_index.core import Settings, StorageContext, VectorStoreIndex
|
||||
from llama_index.core.node_parser import SentenceSplitter
|
||||
from llama_index.embeddings.openai_like import OpenAILikeEmbedding
|
||||
from llama_index.llms.openai_like import OpenAILike
|
||||
from llama_index.readers.web import SimpleWebPageReader
|
||||
from llama_index.vector_stores.milvus import MilvusVectorStore
|
||||
|
||||
|
||||
def init_config(args: Namespace):
|
||||
"""Initialize configuration with command line arguments"""
|
||||
return {
|
||||
"url": args.url,
|
||||
"embedding_model": args.embedding_model,
|
||||
"chat_model": args.chat_model,
|
||||
"vllm_api_key": args.vllm_api_key,
|
||||
"embedding_endpoint": args.embedding_endpoint,
|
||||
"chat_endpoint": args.chat_endpoint,
|
||||
"db_path": args.db_path,
|
||||
"chunk_size": args.chunk_size,
|
||||
"chunk_overlap": args.chunk_overlap,
|
||||
"top_k": args.top_k,
|
||||
}
|
||||
|
||||
|
||||
def load_documents(url: str) -> list:
|
||||
"""Load and process web documents"""
|
||||
return SimpleWebPageReader(html_to_text=True).load_data([url])
|
||||
|
||||
|
||||
def setup_models(config: dict[str, Any]):
|
||||
"""Configure embedding and chat models"""
|
||||
Settings.embed_model = OpenAILikeEmbedding(
|
||||
api_base=config["embedding_endpoint"],
|
||||
api_key=config["vllm_api_key"],
|
||||
model_name=config["embedding_model"],
|
||||
)
|
||||
|
||||
Settings.llm = OpenAILike(
|
||||
model=config["chat_model"],
|
||||
api_key=config["vllm_api_key"],
|
||||
api_base=config["chat_endpoint"],
|
||||
context_window=128000,
|
||||
is_chat_model=True,
|
||||
is_function_calling_model=False,
|
||||
)
|
||||
|
||||
Settings.transformations = [
|
||||
SentenceSplitter(
|
||||
chunk_size=config["chunk_size"],
|
||||
chunk_overlap=config["chunk_overlap"],
|
||||
)
|
||||
]
|
||||
|
||||
|
||||
def setup_vector_store(db_path: str) -> MilvusVectorStore:
|
||||
"""Initialize vector store"""
|
||||
sample_emb = Settings.embed_model.get_text_embedding("test")
|
||||
print(f"Embedding dimension: {len(sample_emb)}")
|
||||
return MilvusVectorStore(uri=db_path, dim=len(sample_emb), overwrite=True)
|
||||
|
||||
|
||||
def create_index(documents: list, vector_store: MilvusVectorStore):
|
||||
"""Create document index"""
|
||||
storage_context = StorageContext.from_defaults(vector_store=vector_store)
|
||||
return VectorStoreIndex.from_documents(
|
||||
documents,
|
||||
storage_context=storage_context,
|
||||
)
|
||||
|
||||
|
||||
def query_document(index: VectorStoreIndex, question: str, top_k: int):
|
||||
"""Query document with given question"""
|
||||
query_engine = index.as_query_engine(similarity_top_k=top_k)
|
||||
return query_engine.query(question)
|
||||
|
||||
|
||||
def get_parser() -> argparse.ArgumentParser:
|
||||
"""Parse command line arguments"""
|
||||
parser = argparse.ArgumentParser(description="RAG with vLLM and LlamaIndex")
|
||||
|
||||
# Add command line arguments
|
||||
parser.add_argument(
|
||||
"--url",
|
||||
default=("https://docs.vllm.ai/en/latest/getting_started/quickstart.html"),
|
||||
help="URL of the document to process",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--embedding-model",
|
||||
default="ssmits/Qwen2-7B-Instruct-embed-base",
|
||||
help="Model name for embeddings",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--chat-model", default="qwen/Qwen1.5-0.5B-Chat", help="Model name for chat"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--vllm-api-key", default="EMPTY", help="API key for vLLM compatible services"
|
||||
)
|
||||
parser.add_argument(
|
||||
"--embedding-endpoint",
|
||||
default="http://localhost:8000/v1",
|
||||
help="Base URL for embedding service",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--chat-endpoint",
|
||||
default="http://localhost:8001/v1",
|
||||
help="Base URL for chat service",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--db-path", default="./milvus_demo.db", help="Path to Milvus database"
|
||||
)
|
||||
parser.add_argument(
|
||||
"-i", "--interactive", action="store_true", help="Enable interactive Q&A mode"
|
||||
)
|
||||
parser.add_argument(
|
||||
"-c",
|
||||
"--chunk-size",
|
||||
type=int,
|
||||
default=1000,
|
||||
help="Chunk size for document splitting",
|
||||
)
|
||||
parser.add_argument(
|
||||
"-o",
|
||||
"--chunk-overlap",
|
||||
type=int,
|
||||
default=200,
|
||||
help="Chunk overlap for document splitting",
|
||||
)
|
||||
parser.add_argument(
|
||||
"-k", "--top-k", type=int, default=3, help="Number of top results to retrieve"
|
||||
)
|
||||
|
||||
return parser
|
||||
|
||||
|
||||
def main():
|
||||
# Parse command line arguments
|
||||
args = get_parser().parse_args()
|
||||
|
||||
# Initialize configuration
|
||||
config = init_config(args)
|
||||
|
||||
# Load documents
|
||||
documents = load_documents(config["url"])
|
||||
|
||||
# Setup models
|
||||
setup_models(config)
|
||||
|
||||
# Setup vector store
|
||||
vector_store = setup_vector_store(config["db_path"])
|
||||
|
||||
# Create index
|
||||
index = create_index(documents, vector_store)
|
||||
|
||||
if args.interactive:
|
||||
print("\nEntering interactive mode. Type 'quit' to exit.")
|
||||
while True:
|
||||
# Get user question
|
||||
question = input("\nEnter your question: ")
|
||||
|
||||
# Check for exit command
|
||||
if question.lower() in ["quit", "exit", "q"]:
|
||||
print("Exiting interactive mode...")
|
||||
break
|
||||
|
||||
# Get and print response
|
||||
print("\n" + "-" * 50)
|
||||
print("Response:\n")
|
||||
response = query_document(index, question, config["top_k"])
|
||||
print(response)
|
||||
print("-" * 50)
|
||||
else:
|
||||
# Single query mode
|
||||
question = "How to install vLLM?"
|
||||
response = query_document(index, question, config["top_k"])
|
||||
print("-" * 50)
|
||||
print("Response:\n")
|
||||
print(response)
|
||||
print("-" * 50)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
Reference in New Issue
Block a user