Skip to content
Merged
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
37 changes: 24 additions & 13 deletions tests/integration/defs/llmapi/test_llm_api_connector.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
# limitations under the License.

import math
import time
from unittest.mock import MagicMock, patch

import pytest
Expand Down Expand Up @@ -63,6 +64,16 @@ def enforce_single_worker(monkeypatch):
yield


def generate_and_sleep(model, *args, **kwargs):
# Some KV connector API calls are made after a full response is returned. We want to be able to track these calls.
# However, we don't have any indication of when all the calls are complete.
# To compensate for this, we sleep between the generate call and the return of the outputs.
# TODO(jthomson04): Surely there's a better way to do this?
outputs = model.generate(*args, **kwargs)
time.sleep(1)
return outputs


@pytest.mark.threadleak(enabled=False)
@pytest.mark.parametrize("use_overlap_scheduler", [True, False])
def test_connector_simple(enforce_single_worker, model_with_connector,
Expand All @@ -81,7 +92,7 @@ def test_connector_simple(enforce_single_worker, model_with_connector,

sampling_params = SamplingParams(max_tokens=NUM_TOKENS, ignore_eos=True)

model.generate(["Hello, world"], sampling_params)
generate_and_sleep(model, ["Hello, world"], sampling_params)

assert scheduler.update_state_after_alloc.call_count == 1

Expand Down Expand Up @@ -154,7 +165,7 @@ def test_connector_async_onboard(enforce_single_worker, model_with_connector,
worker.get_finished.side_effect = lambda finished_gen, load_async: (
finished_gen, load_async)

model.generate([
generate_and_sleep(model, [
"Lorem ipsum dolor sit amet, consectetur adipiscing elit. Sed do eiusmod tempor incididunt ut labore et dolore magna aliqua."
], SamplingParams(max_tokens=NUM_TOKENS, ignore_eos=True))

Expand Down Expand Up @@ -187,7 +198,7 @@ def test_connector_async_save(enforce_single_worker, model_with_connector,

sampling_params = SamplingParams(max_tokens=NUM_TOKENS, ignore_eos=True)

model.generate(["Hello, world"], sampling_params)
generate_and_sleep(model, ["Hello, world"], sampling_params)

assert scheduler.request_finished.call_count == 1

Expand Down Expand Up @@ -227,7 +238,7 @@ def test_connector_scheduler_output(enforce_single_worker, model_with_connector,

sampling_params = SamplingParams(max_tokens=32, ignore_eos=True)

model.generate([0] * NUM_INPUT_TOKENS, sampling_params)
generate_and_sleep(model, [0] * NUM_INPUT_TOKENS, sampling_params)

assert scheduler.update_state_after_alloc.call_count == 1
assert len(
Expand Down Expand Up @@ -277,7 +288,7 @@ def test_connector_scheduler_output(enforce_single_worker, model_with_connector,
assert len(scheduler.request_finished.call_args.args[1]) == math.ceil(
(NUM_INPUT_TOKENS + NUM_TOKENS) / BLOCK_SIZE)

model.generate([1] * NUM_INPUT_TOKENS, sampling_params)
generate_and_sleep(model, [1] * NUM_INPUT_TOKENS, sampling_params)

# The initial computed position should be 0, since we haven't yet onboarded any blocks.
assert scheduler.build_connector_meta.call_args_list[0].args[
Expand Down Expand Up @@ -306,7 +317,7 @@ def test_connector_scheduler_output_chunked_context(enforce_single_worker,

sampling_params = SamplingParams(max_tokens=BLOCK_SIZE, ignore_eos=True)

model.generate([0] * (CHUNK_SIZE * 2), sampling_params)
generate_and_sleep(model, [0] * (CHUNK_SIZE * 2), sampling_params)

assert scheduler.update_state_after_alloc.call_count == 1

Expand Down Expand Up @@ -379,16 +390,16 @@ def test_connector_disagg_prefill(enforce_single_worker, model_with_connector,
scheduler.request_finished.return_value = False
worker.get_finished.return_value = [], []

result = prefill_worker.generate([0] * 48,
sampling_params=sampling_params,
disaggregated_params=disaggregated_params)
result = generate_and_sleep(prefill_worker, [0] * 48,
sampling_params=sampling_params,
disaggregated_params=disaggregated_params)

gen_disagg_params = result.disaggregated_params
gen_disagg_params.request_type = "generation_only"

result = decode_worker.generate([0] * 48,
sampling_params=sampling_params,
disaggregated_params=gen_disagg_params)
generate_and_sleep(decode_worker, [0] * 48,
sampling_params=sampling_params,
disaggregated_params=gen_disagg_params)

assert scheduler.build_connector_meta.call_count == 1

Expand All @@ -411,7 +422,7 @@ def test_connector_multi_request(enforce_single_worker, model_with_connector):
model_fn, scheduler, worker = model_with_connector

model = model_fn(disable_overlap_scheduler=True,
kv_cache_config=KvCacheConfig(max_tokens=120))
kv_cache_config=KvCacheConfig(max_tokens=144))

sampling_params = SamplingParams(ignore_eos=True, max_tokens=4)

Expand Down
25 changes: 13 additions & 12 deletions tests/integration/test_lists/test-db/l0_a10.yml
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,19 @@ l0_a10:
- unittest/llmapi/apps/test_chat_utils.py
- unittest/llmapi/apps/test_tool_parsers.py
- unittest/llmapi/apps/test_harmony_channel_validation.py
- llmapi/test_llm_api_connector.py::test_connector_simple[True]
- llmapi/test_llm_api_connector.py::test_connector_simple[False]
- llmapi/test_llm_api_connector.py::test_connector_async_onboard[True]
- llmapi/test_llm_api_connector.py::test_connector_async_onboard[False]
- llmapi/test_llm_api_connector.py::test_connector_async_save[True]
- llmapi/test_llm_api_connector.py::test_connector_async_save[False]
- llmapi/test_llm_api_connector.py::test_connector_scheduler_output[True]
- llmapi/test_llm_api_connector.py::test_connector_scheduler_output[False]
- llmapi/test_llm_api_connector.py::test_connector_scheduler_output_chunked_context[True]
- llmapi/test_llm_api_connector.py::test_connector_scheduler_output_chunked_context[False]
- llmapi/test_llm_api_connector.py::test_connector_disagg_prefill[False]
- llmapi/test_llm_api_connector.py::test_connector_disagg_prefill[True]
- llmapi/test_llm_api_connector.py::test_connector_multi_request
- condition:
ranges:
system_gpu_count:
Expand Down Expand Up @@ -119,18 +132,6 @@ l0_a10:
- examples/test_bert.py::test_llm_bert_general[compare_hf-enable_remove_input_padding-use_attention_plugin-enable_context_fmha-tp:1-pp:1-float16-BertModel-bert/bert-base-uncased]
- unittest/trt/model/test_mistral.py
- unittest/trt/model/test_llama.py
- llmapi/test_llm_api_connector.py::test_connector_simple[True]
- llmapi/test_llm_api_connector.py::test_connector_simple[False]
- llmapi/test_llm_api_connector.py::test_connector_async_onboard[True]
- llmapi/test_llm_api_connector.py::test_connector_async_onboard[False]
- llmapi/test_llm_api_connector.py::test_connector_async_save[True]
- llmapi/test_llm_api_connector.py::test_connector_async_save[False]
- llmapi/test_llm_api_connector.py::test_connector_scheduler_output[True]
- llmapi/test_llm_api_connector.py::test_connector_scheduler_output[False]
- llmapi/test_llm_api_connector.py::test_connector_scheduler_output_chunked_context[True]
- llmapi/test_llm_api_connector.py::test_connector_scheduler_output_chunked_context[False]
- llmapi/test_llm_api_connector.py::test_connector_disagg_prefill[False]
- llmapi/test_llm_api_connector.py::test_connector_disagg_prefill[True]
- llmapi/test_llm_e2e.py::test_llmapi_load_engine_from_build_command[llama-llama-models/llama-7b-hf] # 5min
- llmapi/test_llm_e2e.py::test_llmapi_build_command_parameters_align[llama-llama-models-v2/TinyLlama-1.1B-Chat-v1.0]
- llmapi/test_llm_e2e.py::test_llmapi_load_engine_from_build_command_with_lora[llama-llama-models-v2/llama-v2-7b-hf]
Expand Down