Skip to content
Open
Show file tree
Hide file tree
Changes from 3 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
17 changes: 16 additions & 1 deletion tests/test_vllm_client_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,11 @@
from transformers import AutoModelForCausalLM, AutoProcessor, AutoTokenizer
from transformers.testing_utils import torch_device

from trl.generation.vllm_client import VLLMClient
from trl.generation.vllm_client import (
VLLMClient,
_format_http_host,
_normalize_communicator_host,
)
from trl.generation.vllm_generation import extract_logprobs
from trl.import_utils import is_vllm_available
from trl.scripts.vllm_serve import chunk_list
Expand All @@ -39,6 +43,17 @@
from vllm import LLM, SamplingParams


class TestVLLMClientAddressing(TrlTestCase):
def test_communicator_host_strips_ipv6_brackets(self):
assert _normalize_communicator_host("[2001:db8::1]") == "2001:db8::1"
assert _normalize_communicator_host("2001:db8::1") == "2001:db8::1"

def test_http_host_brackets_only_ipv6_literals(self):
assert _format_http_host("[2001:db8::1]") == "[2001:db8::1]"
assert _format_http_host("2001:db8::1") == "[2001:db8::1]"
assert _format_http_host("127.0.0.1") == "127.0.0.1"
assert _format_http_host("localhost") == "localhost"

class TestChunkList(TrlTestCase):
def test_even_split(self):
assert chunk_list([1, 2, 3, 4, 5, 6], 2) == [[1, 2, 3], [4, 5, 6]]
Expand Down
24 changes: 19 additions & 5 deletions trl/generation/vllm_client.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,6 @@
import base64
import copy
import logging
import socket
import time
from io import BytesIO
from urllib.parse import urlparse
Expand Down Expand Up @@ -48,6 +47,19 @@
logger = logging.getLogger(__name__)


def _normalize_communicator_host(host: str) -> str:
"""Return the bare host literal required by TCPStore/NCCL communicators."""
if host.startswith("[") and host.endswith("]"):
return host[1:-1]
return host


def _format_http_host(host: str) -> str:
"""Bracket an IPv6 literal when embedding it in an HTTP URL."""
host = _normalize_communicator_host(host)
return f"[{host}]" if ":" in host else host


def pil_to_base64(image):
buffer = BytesIO()
image.save(buffer, format="PNG")
Expand Down Expand Up @@ -155,13 +167,13 @@ def __init__(
if base_url is not None:
# Parse the base_url to extract host and port
parsed_url = urlparse(base_url)
self.host = socket.gethostbyname(parsed_url.hostname)
self.host = _normalize_communicator_host(parsed_url.hostname)
Comment thread
cursor[bot] marked this conversation as resolved.
Outdated
scheme = parsed_url.scheme or "http"
self.base_url = f"{scheme}://{parsed_url.netloc}{parsed_url.path}"
else:
self.host = host
self.host = _normalize_communicator_host(host)
self.server_port = server_port
self.base_url = f"http://{self.host}:{self.server_port}"
self.base_url = f"http://{_format_http_host(self.host)}:{self.server_port}"
self.group_port = group_port
self.check_server(connection_timeout) # check server and fail after timeout

Expand Down Expand Up @@ -193,7 +205,9 @@ def check_server(self, total_timeout: float = 0.0, retry_interval: float = 2.0):
else:
if response.status_code == 200:
if "X-Forwarded-For" in response.headers:
self.host = response.headers["X-Forwarded-For"]
self.host = _normalize_communicator_host(
response.headers["X-Forwarded-For"]
)
logger.info("Server is up!")
return None

Expand Down