diff --git a/tests/v1/e2e/general/test_correctness_sliding_window.py b/tests/v1/e2e/general/test_correctness_sliding_window.py index 01d60444170..a8a29203d9e 100644 --- a/tests/v1/e2e/general/test_correctness_sliding_window.py +++ b/tests/v1/e2e/general/test_correctness_sliding_window.py @@ -33,7 +33,7 @@ model_config = { @pytest.mark.parametrize("seed", [1]) @pytest.mark.parametrize("disable_hybrid_kv_cache_manager", [True, False]) def test_sliding_window_retrieval( - model, batch_size, seed, disable_hybrid_kv_cache_manager + model, batch_size, seed, disable_hybrid_kv_cache_manager, vllm_runner ): """ The test does a bunch of assignments "x1 = 10\nx2 = 33\n..." and then @@ -48,34 +48,39 @@ def test_sliding_window_retrieval( test_config = model_config[model] - llm = LLM( - model=model, + with vllm_runner( + model, + max_model_len=None, + enable_chunked_prefill=None, disable_hybrid_kv_cache_manager=disable_hybrid_kv_cache_manager, enforce_eager=enforce_eager, - ) - sampling_params = SamplingParams(temperature=0.0, max_tokens=100) + ) as runner: + llm = runner.get_llm() + sampling_params = SamplingParams(temperature=0.0, max_tokens=100) - prompts, answer, indices = prep_prompts(batch_size, ln_range=test_config.ln_range) + prompts, answer, indices = prep_prompts( + batch_size, ln_range=test_config.ln_range + ) - check_length(prompts, llm, test_config.sliding_window) + check_length(prompts, llm, test_config.sliding_window) - # Fresh generation - responses = llm.generate(prompts, sampling_params) - check_answers( - indices, - answer, - [response.outputs[0].text for response in responses], - accept_rate=1.0, - ) + # Fresh generation + responses = llm.generate(prompts, sampling_params) + check_answers( + indices, + answer, + [response.outputs[0].text for response in responses], + accept_rate=1.0, + ) - # Re-generate with the same prompts to test prefix caching - responses = llm.generate(prompts, sampling_params) - check_answers( - indices, - answer, - [response.outputs[0].text for response in responses], - accept_rate=1.0, - ) + # Re-generate with the same prompts to test prefix caching + responses = llm.generate(prompts, sampling_params) + check_answers( + indices, + answer, + [response.outputs[0].text for response in responses], + accept_rate=1.0, + ) def check_length(prompts: list[str], llm: LLM, sliding_window: int):