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
7 changes: 5 additions & 2 deletions src/engine/ov_genai/llm.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,7 @@
from src.engine.ov_genai.utils import extract_scheduler_config_from_loader
from src.engine.ov_genai.utils import (
apply_temperature,
extract_scheduler_config_from_loader,
)
import asyncio
import gc
import logging
Expand Down Expand Up @@ -322,7 +325,7 @@ def create_generation_config(self, config: OVGenAI_GenConfig) -> GenerationConfi
"""
generation_kwargs = self.model.get_generation_config() if self.model else GenerationConfig()
generation_kwargs.max_new_tokens = config.max_tokens
generation_kwargs.temperature = config.temperature
apply_temperature(generation_kwargs, config.temperature)
generation_kwargs.top_k = config.top_k
generation_kwargs.top_p = config.top_p
generation_kwargs.repetition_penalty = config.repetition_penalty
Expand Down
15 changes: 14 additions & 1 deletion src/engine/ov_genai/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,13 +2,26 @@
import logging
from typing import Literal

from openvino_genai import SchedulerConfig
from openvino_genai import GenerationConfig, SchedulerConfig

from src.server.schemas.modeling.contract_ovgenai_llm_and_vlm import SchedulerConfigSchema
from src.server.schemas.registration import ModelLoadConfig

logger = logging.getLogger(__name__)


def apply_temperature(generation_config: GenerationConfig, temperature: float) -> None:
"""Set the sampling temperature, falling back to greedy decoding at zero.

OpenAI treats temperature 0 as greedy, but OpenVINO GenAI rejects a
non-positive temperature while do_sample is true, and that failure unloads
the model.
"""
generation_config.temperature = temperature
if temperature <= 0:
generation_config.do_sample = False


def generate_ov_scheduler_config(scheduler_config: SchedulerConfigSchema) -> dict:
"""Generates a SchedulerConfig object from the scheduler config model.

Expand Down
7 changes: 5 additions & 2 deletions src/engine/ov_genai/vlm.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,7 @@
from src.engine.ov_genai.utils import extract_scheduler_config_from_loader
from src.engine.ov_genai.utils import (
apply_temperature,
extract_scheduler_config_from_loader,
)
import asyncio
import base64
import gc
Expand Down Expand Up @@ -415,7 +418,7 @@ def create_generation_config(self, config: OVGenAI_GenConfig) -> GenerationConfi
"""
generation_kwargs = self.model_path.get_generation_config() if self.model_path else GenerationConfig()
generation_kwargs.max_new_tokens = config.max_tokens
generation_kwargs.temperature = config.temperature
apply_temperature(generation_kwargs, config.temperature)
generation_kwargs.top_k = config.top_k
generation_kwargs.top_p = config.top_p
generation_kwargs.repetition_penalty = config.repetition_penalty
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -53,7 +53,7 @@ class OVGenAI_GenConfig(BaseModel):
description="Confidence threshold for accepting draft tokens (typically 0.3-0.5)"
)

stream: bool = Field(
stream: Optional[bool] = Field(
default=False,
description="Stream output in chunks of tokens."
)
Expand Down
8 changes: 8 additions & 0 deletions tests/unit/test_config_merge_unit.py
Original file line number Diff line number Diff line change
Expand Up @@ -244,3 +244,11 @@ def test_asr_yaml_default_yields_to_request() -> None:
assert config.language == "Chinese" # request wins
assert config.max_chunk_sec == 25.0 # yaml default
assert config.max_tokens == 1024 # engine default


def test_stream_null_is_accepted() -> None:
# OpenAI-compatible clients send `stream: null`, and the chat route forwards
# request.stream (None when the client omits it); a bare `bool` field 400s.
config = build_config(OVGenAI_GenConfig, request={}, defaults={}, messages=[], stream=None)

assert not config.stream
30 changes: 30 additions & 0 deletions tests/unit/test_ov_genai_llm_unit.py
Original file line number Diff line number Diff line change
Expand Up @@ -269,3 +269,33 @@ async def test_unload_model_resets_state(monkeypatch: pytest.MonkeyPatch, load_c
registry.register_unload.assert_called_once_with("model-name")
gc_mock.assert_called_once()



class DummyGenerationConfig:
pass


def test_create_generation_config_zero_temperature_disables_sampling(
monkeypatch: pytest.MonkeyPatch, load_config: ModelLoadConfig
) -> None:
monkeypatch.setattr(llm_module, "GenerationConfig", DummyGenerationConfig)
llm = OVGenAI_LLM(load_config)
llm.model = None

config = llm.create_generation_config(OVGenAI_GenConfig(temperature=0.0))

assert config.do_sample is False
assert config.temperature == 0.0


def test_create_generation_config_positive_temperature_keeps_sampling(
monkeypatch: pytest.MonkeyPatch, load_config: ModelLoadConfig
) -> None:
monkeypatch.setattr(llm_module, "GenerationConfig", DummyGenerationConfig)
llm = OVGenAI_LLM(load_config)
llm.model = None

config = llm.create_generation_config(OVGenAI_GenConfig(temperature=0.7))

assert getattr(config, "do_sample", None) is None
assert config.temperature == 0.7
17 changes: 17 additions & 0 deletions tests/unit/test_ov_genai_vlm_unit.py
Original file line number Diff line number Diff line change
Expand Up @@ -253,3 +253,20 @@ def test_resolve_prompt_strips_stray_token_for_text_only_chat(
prompt, images = vlm._resolve_prompt_and_images(config)
assert images == []
assert tok not in prompt


class DummyVlmGenerationConfig:
pass


def test_vlm_create_generation_config_zero_temperature_disables_sampling(
monkeypatch: pytest.MonkeyPatch, load_config: ModelLoadConfig
) -> None:
monkeypatch.setattr(vlm_module, "GenerationConfig", DummyVlmGenerationConfig)
vlm = OVGenAI_VLM(load_config)
vlm.model_path = None

config = vlm.create_generation_config(OVGenAI_GenConfig(temperature=0.0))

assert config.do_sample is False
assert config.temperature == 0.0