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
9 changes: 5 additions & 4 deletions lmdeploy/pytorch/engine/inputs_maker.py
Original file line number Diff line number Diff line change
Expand Up @@ -1229,16 +1229,17 @@ def create_model_inputs_delta_valid_only(self):

def update_running_seqs(self, running: 'SeqList', inputs: 'ModelInputs|None'):
"""Update running seqs."""
if self.config.role == EngineRole.Prefill:
# p node will not update running seqs
return

is_decoding = inputs is None
if self.long_context_chunker.enabled() and not is_decoding and inputs.is_chunk:
# long context chunk does not need to update running seqs
self.long_context_chunker.update_step(inputs)
return

# Prefill workers must advance chunk history above, but must not add
# completed prefills to the decode-running list.
if self.config.role == EngineRole.Prefill:
return

if is_decoding:
self.running_seqs = running
else:
Expand Down
3 changes: 3 additions & 0 deletions lmdeploy/serve/openai/chat_completions/serving.py
Original file line number Diff line number Diff line change
Expand Up @@ -303,6 +303,7 @@ async def _completion_stream_generator() -> AsyncGenerator[str, None]:
routed_experts=chunk.routed_experts,
output_ids=stream_output_ids)
if chunk.cache_block_ids is not None and chunk.is_last_delta:
response_json['cache_session_id'] = chat_runner.session.session_id
response_json['cache_block_ids'] = chunk.cache_block_ids
response_json['remote_token_ids'] = chunk.token_ids
yield f'data: {json.dumps(response_json)}\n\n'
Expand Down Expand Up @@ -383,6 +384,8 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]:
).model_dump()

if with_cache:
# Keep the OpenAI response ID independent of the engine's cache owner.
response['cache_session_id'] = chat_runner.session.session_id
response['cache_block_ids'] = cache_block_ids
response['remote_token_ids'] = remote_token_ids

Expand Down
5 changes: 4 additions & 1 deletion lmdeploy/serve/openai/endpoints/completions.py
Original file line number Diff line number Diff line change
Expand Up @@ -223,7 +223,7 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]:
# First chunk with role
final_usage = UsageInfo() if include_usage else None
try:
for generator in generators:
for generator, session in zip(generators, sessions):
async for res in generator:
logprobs = None
if request.logprobs and res.logprobs:
Expand All @@ -243,6 +243,7 @@ async def completion_stream_generator() -> AsyncGenerator[str, None]:
finish_reason=res.finish_reason,
logprobs=logprobs)
if res.cache_block_ids is not None:
response_json['cache_session_id'] = session.session_id
response_json['cache_block_ids'] = res.cache_block_ids
response_json['remote_token_ids'] = res.token_ids
yield f'data: {json.dumps(response_json)}\n\n'
Expand Down Expand Up @@ -333,6 +334,8 @@ async def _inner_call(i, generator, session):
).model_dump()

if with_cache:
# Cache ownership uses the engine's ID, not the public response ID.
response['cache_session_id'] = sessions[0].session_id
response['cache_block_ids'] = cache_block_ids
response['remote_token_ids'] = remote_token_ids

Expand Down
23 changes: 16 additions & 7 deletions lmdeploy/serve/proxy/proxy.py
Original file line number Diff line number Diff line change
Expand Up @@ -710,7 +710,12 @@ async def chat_completions_v1(request: ChatCompletionRequest, raw_request: Reque
rdma_config=node_manager.rdma_config,
))

remote_session_id = int(prefill_info.get('id')) if prefill_info.get('id') else 0
# The public response ID is not the engine's cache owner.
remote_session_id = 0 if node_manager.dummy_prefill else prefill_info.get('cache_session_id')
if type(remote_session_id) is not int or remote_session_id < 0:
return create_error_response(HTTPStatus.BAD_GATEWAY,
'Prefill response is missing a valid cache_session_id; '
'upgrade the Prefill server together with the proxy.')
Comment on lines +715 to +718
remote_block_ids = prefill_info.get('cache_block_ids') or []
remote_token_id = prefill_info.get('remote_token_ids')[-1] if prefill_info.get('remote_token_ids') else 0

Expand All @@ -724,7 +729,7 @@ async def chat_completions_v1(request: ChatCompletionRequest, raw_request: Reque

start = node_manager.pre_call(d_url)
if not node_manager.dummy_prefill:
node_manager.pd_connection_pool.shelf_prefill_session((p_url, d_url), prefill_info['id'])
node_manager.pd_connection_pool.shelf_prefill_session((p_url, d_url), remote_session_id)
if request.stream is True:
response = node_manager.stream_generate(request_dict, d_url, '/v1/chat/completions')
background_task = node_manager.create_background_tasks(d_url, start)
Expand All @@ -735,7 +740,7 @@ async def chat_completions_v1(request: ChatCompletionRequest, raw_request: Reque
resp = JSONResponse(json.loads(response))

if not node_manager.dummy_prefill:
node_manager.pd_connection_pool.unshelf_prefill_session((p_url, d_url), prefill_info['id'])
node_manager.pd_connection_pool.unshelf_prefill_session((p_url, d_url), remote_session_id)

return resp

Expand Down Expand Up @@ -855,9 +860,13 @@ async def completions_v1(request: CompletionRequest, raw_request: Request = None
except Exception as e:
logger.error(f'error Msg: {str(e)}')
return {'status': f'Connection error, cannot establish connection {(p_url, d_url)}'}
node_manager.pd_connection_pool.shelf_prefill_session((p_url, d_url), prefill_info['id'])

remote_session_id = int(prefill_info.get('id')) if prefill_info.get('id') else 0
# The public response ID is not the engine's cache owner.
remote_session_id = 0 if node_manager.dummy_prefill else prefill_info.get('cache_session_id')
if type(remote_session_id) is not int or remote_session_id < 0:
return create_error_response(HTTPStatus.BAD_GATEWAY,
'Prefill response is missing a valid cache_session_id; '
'upgrade the Prefill server together with the proxy.')
remote_block_ids = prefill_info.get('cache_block_ids') or []
remote_token_id = prefill_info.get('remote_token_ids')[-1] if prefill_info.get('remote_token_ids') else 0
request_dict['migration_request'] = MigrationRequest(
Expand All @@ -870,7 +879,7 @@ async def completions_v1(request: CompletionRequest, raw_request: Request = None

start = node_manager.pre_call(d_url)
if not node_manager.dummy_prefill:
node_manager.pd_connection_pool.shelf_prefill_session((p_url, d_url), prefill_info['id'])
node_manager.pd_connection_pool.shelf_prefill_session((p_url, d_url), remote_session_id)
if request.stream is True:
response = node_manager.stream_generate(request_dict, d_url, '/v1/completions')
background_task = node_manager.create_background_tasks(d_url, start)
Expand All @@ -880,7 +889,7 @@ async def completions_v1(request: CompletionRequest, raw_request: Request = None
node_manager.post_call(d_url, start)
resp = JSONResponse(json.loads(response))
if not node_manager.dummy_prefill:
node_manager.pd_connection_pool.unshelf_prefill_session((p_url, d_url), prefill_info.get('id'))
node_manager.pd_connection_pool.unshelf_prefill_session((p_url, d_url), remote_session_id)
return resp
else:
raise ValueError(f'No serving strategy named {node_manager.serving_strategy}')
Expand Down
32 changes: 32 additions & 0 deletions tests/pytorch/engine/test_inputs_maker.py
Original file line number Diff line number Diff line change
Expand Up @@ -1388,6 +1388,38 @@ def has_waiting_long_prefill(self):
assert calls == [(True, True), (False, False)]


@pytest.mark.parametrize('role', list(EngineRole))
@pytest.mark.parametrize('length', [511, 512, 513, 1106])
def test_chunk_history_advances_for_all_roles_without_decoding_on_prefill(role, length):
from lmdeploy.pytorch.config import CacheConfig, SchedulerConfig
from lmdeploy.pytorch.messages import SequenceMeta
from lmdeploy.pytorch.paging import Scheduler
from lmdeploy.pytorch.strategies.ar.sequence import ARSequenceStrategy

scheduler = Scheduler(
SchedulerConfig(max_batches=4, max_session_len=4096, max_request_output_len=2048),
CacheConfig(max_batches=4, block_size=16, num_cpu_blocks=0, num_gpu_blocks=256),
SequenceMeta(16, strategy=ARSequenceStrategy()))
seq = scheduler.add_session(0).add_sequence(list(range(length)), preserve_cache=True)
maker = InputsMakerAsync.__new__(InputsMakerAsync)
maker.config = SimpleNamespace(role=role)
maker.running_seqs = []
maker.long_context_chunker = LongContextChunker(512)
chunker = maker.long_context_chunker
chunker.set_seq(seq)
for expected in range(512, length, 512):
assert not chunker.is_last_chunk()
size, _ = chunker.next_chunk_size()
maker.update_running_seqs([seq], SimpleNamespace(is_chunk=True, max_q_seqlen=size))
assert seq.num_history_ids == expected
assert maker.running_seqs == []
assert chunker.is_last_chunk()
# The input builder clears the chunker before submitting the final chunk.
chunker.clear()
maker.update_running_seqs([seq], SimpleNamespace(is_chunk=True))
assert maker.running_seqs == ([] if role == EngineRole.Prefill else [seq])


def test_normal_prefill_can_update_running_while_long_chunker_is_active():
long_seq = _DummySeq(history_ids=0, token_ids=1024, all_multimodals={}, input_multimodals={})
short_seq = _DummySeq(history_ids=0, token_ids=16, all_multimodals={}, input_multimodals={})
Expand Down
224 changes: 224 additions & 0 deletions tests/test_lmdeploy/serve/openai/test_distserve.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,224 @@
# Copyright (c) OpenMMLab. All rights reserved.
"""Cache-owner IDs travel from the real API to migration ACK and scheduler GC.

Only generation/transport are simulated; no model, CUDA, or RDMA is needed.
"""
import asyncio
import importlib
import json
from types import SimpleNamespace
from unittest.mock import Mock

import httpx
import pytest
from fastapi import APIRouter, FastAPI

from lmdeploy.pytorch.config import CacheConfig, SchedulerConfig
from lmdeploy.pytorch.disagg.config import EngineRole, ServingStrategy
from lmdeploy.pytorch.disagg.conn.protocol import MigrationProtocol
from lmdeploy.pytorch.disagg.conn.proxy_conn import PDConnectionPool
from lmdeploy.pytorch.engine.engine import Engine
from lmdeploy.pytorch.messages import SequenceMeta
from lmdeploy.pytorch.paging.scheduler import Scheduler
from lmdeploy.pytorch.paging.seq_states.states import ToBeMigratedState
from lmdeploy.pytorch.strategies.ar.sequence import ARSequenceStrategy
from lmdeploy.serve.managers import SessionManager
from lmdeploy.serve.openai.api_server import ServerContext
from lmdeploy.serve.openai.chat_completions import register as register_chat
from lmdeploy.serve.openai.endpoints.completions import register as register_completion
from lmdeploy.serve.openai.protocol import ChatCompletionRequest, CompletionRequest, DeltaMessage


class _Parser:
tool_parser_cls = None

def __init__(self, request):
self.request = request
self.tool_parser = None
self.reasoning_tokens = 0

def stream_chunk(self, text, token_ids, **kwargs):
return [(DeltaMessage(content=text), False)]

def parse_complete(self, text, token_ids=None):
return text, None, None

def validate_complete(self, **kwargs):
return True


@pytest.fixture
def prefill():
scheduler = Scheduler(
SchedulerConfig(max_batches=4, max_session_len=128, max_request_output_len=64),
CacheConfig(max_batches=4, block_size=16, num_cpu_blocks=0, num_gpu_blocks=16),
SequenceMeta(16, strategy=ARSequenceStrategy()))
engine = object.__new__(Engine)
engine.scheduler = scheduler

class Generation:
model_name = 'fake-model'
epoch = 0
backend_config = SimpleNamespace(role=EngineRole.Prefill, logprobs_mode=None, adapters=[])
tokenizer = SimpleNamespace(model=SimpleNamespace(model=SimpleNamespace(model=None)))

def __init__(self):
self.session_mgr = SessionManager()
self.ids = []

async def preprocess(self, prompt, session, **kwargs):
return SimpleNamespace(session=session, gen_config=kwargs['gen_config'])

async def generate(self, prepared, **kwargs):
sid = prepared.session.session_id
self.ids.append(sid)
seq = scheduler.add_session(sid).add_sequence([1, 2], preserve_cache=prepared.gen_config.preserve_cache)
seq.state.to_state(ToBeMigratedState)
engine._on_end_session([SimpleNamespace(data={'session_id': sid, 'response': False})])
yield SimpleNamespace(response='ok', token_ids=[1], input_token_len=2, generate_token_len=1,
finish_reason='stop', logprobs=None, cache_block_ids=[0], cached_tokens=0,
routed_experts=None)

context = ServerContext()
context.async_engine = Generation()
context.response_parser_cls = _Parser
router = APIRouter()
register_chat(router, context)
register_completion(router, context)
app = FastAPI()
app.include_router(router)
return context, engine, app


def _request(chat, **kwargs):
if chat:
return ChatCompletionRequest(model='fake-model', messages=[dict(role='user', content='hi')], **kwargs)
return CompletionRequest(model='fake-model', prompt='hi', **kwargs)


@pytest.mark.parametrize('chat', [False, True])
@pytest.mark.parametrize('user_id', [-1, 42])
@pytest.mark.parametrize('stream', [False, True])
def test_prefill_cache_handle_is_not_the_public_response_id(prefill, chat, user_id, stream):
context, engine, app = prefill

async def run():
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url='http://prefill') as client:
for _ in range(3):
payload = _request(chat, session_id=user_id, stream=stream).model_dump()
payload.update(with_cache=True, preserve_cache=True)
response = await client.post('/v1/chat/completions' if chat else '/v1/completions', json=payload)
assert response.status_code == 200, response.text
if stream:
chunks = [json.loads(line[6:]) for line in response.text.splitlines()
if line.startswith('data: ') and line != 'data: [DONE]']
result = next(chunk for chunk in chunks if 'cache_session_id' in chunk)
else:
result = response.json()
owner = result['cache_session_id']
assert owner == context.async_engine.ids[-1]
assert result['id'].startswith('chatcmpl-') if chat else result['id'] == str(user_id)
assert len(engine.scheduler.sessions) == 1
assert engine.end_session(owner)
assert not engine.scheduler.sessions
assert not engine.scheduler.seq_manager._seq_map
assert not context.session_manager.sessions

asyncio.run(run())


@pytest.mark.parametrize('chat', [False, True])
@pytest.mark.parametrize('user_id', [-1, 42])
@pytest.mark.parametrize('stream', [False, True])
def test_proxy_passes_internal_cache_owner_to_decode(prefill, monkeypatch, chat, user_id, stream):
proxy = importlib.import_module('lmdeploy.serve.proxy.proxy')
context, engine, app = prefill
pool = PDConnectionPool()
pool.is_connected = lambda *args: True
pool.shelf_prefill_session = Mock(wraps=pool.shelf_prefill_session)
pool.unshelf_prefill_session = Mock(wraps=pool.unshelf_prefill_session)
owners = []

async def check_model(model):
return None

def acknowledge(payload):
owner = payload['migration_request']['remote_session_id']
assert owner == context.async_engine.ids[-1]
owners.append(owner)
assert engine.end_session(owner)

async def generate(payload, url, endpoint):
if url == 'd':
acknowledge(payload)
return json.dumps(dict(choices=[dict(finish_reason='stop')]))
async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url='http://prefill') as client:
response = await client.post(endpoint, json=payload)
assert response.status_code == 200
return response.text

async def stream_generate(payload, url, endpoint):
acknowledge(payload)
yield b'data: {"choices": [{"finish_reason": "stop"}]}\n\n'
yield b'data: [DONE]\n\n'

manager = SimpleNamespace(
serving_strategy=ServingStrategy.DistServe, check_request_model=check_model,
dummy_prefill=False, migration_protocol=MigrationProtocol.RDMA, rdma_config=None,
pd_connection_pool=pool, generate=generate, stream_generate=stream_generate,
get_node_url=lambda model, role: 'p' if role == EngineRole.Prefill else 'd',
pre_call=lambda url: 0, post_call=lambda *args: None,
create_background_tasks=lambda *args: None)
monkeypatch.setattr(proxy, 'node_manager', manager)

async def run():
request = _request(chat, session_id=user_id, stream=stream, max_tokens=8)
handler = proxy.chat_completions_v1 if chat else proxy.completions_v1
response = await handler(request)
assert response.status_code == 200
if stream:
async for _ in response.body_iterator:
pass
assert owners == [context.async_engine.ids[-1]]
pool.shelf_prefill_session.assert_called_once_with(('p', 'd'), owners[0])
pool.unshelf_prefill_session.assert_called_once_with(('p', 'd'), owners[0])
assert not engine.scheduler.sessions
assert not engine.scheduler.seq_manager._seq_map

asyncio.run(run())


@pytest.mark.parametrize('chat', [False, True])
@pytest.mark.parametrize('owner', [None, -1, '42', True])
def test_proxy_rejects_invalid_cache_owner_without_public_id_fallback(monkeypatch, chat, owner):
proxy = importlib.import_module('lmdeploy.serve.proxy.proxy')
pool = PDConnectionPool()
pool.is_connected = lambda *args: True
calls = []

async def check_model(model):
return None

async def generate(payload, url, endpoint):
calls.append(url)
result = dict(id='42', cache_block_ids=[0], remote_token_ids=[1])
if owner is not None:
result['cache_session_id'] = owner
return json.dumps(result)

manager = SimpleNamespace(
serving_strategy=ServingStrategy.DistServe, check_request_model=check_model,
dummy_prefill=False, pd_connection_pool=pool, generate=generate,
get_node_url=lambda model, role: 'p' if role == EngineRole.Prefill else 'd',
pre_call=lambda url: 0, post_call=lambda *args: None)
monkeypatch.setattr(proxy, 'node_manager', manager)

async def run():
handler = proxy.chat_completions_v1 if chat else proxy.completions_v1
response = await handler(_request(chat))
assert response.status_code == 502
assert 'cache_session_id' in response.body.decode()
assert calls == ['p']
assert not pool.migration_session_shelf

asyncio.run(run())
Loading