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
69 changes: 68 additions & 1 deletion neon_data_models/models/api/llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,6 +117,38 @@ def beam_search(self) -> bool:
def beam_search(self, value: bool):
self.extra_body["use_beam_search"] = value

@property
def thinking_token_budget(self) -> Optional[int]:
return self.extra_body.get('thinking_token_budget')

@thinking_token_budget.setter
def thinking_token_budget(self, value: Optional[int]):
if value is None:
self.extra_body.pop("thinking_token_budget", None)
self.extra_body.pop("skip_special_tokens", None)
if "chat_template_kwargs" in self.extra_body:
self.extra_body["chat_template_kwargs"].pop(
"add_thinking_start", None)
self.extra_body["chat_template_kwargs"].pop(
"enable_thinking", None)
elif isinstance(value, int):
assert value >= 0, "thinking_token_budget must be positive"
if value >= self.max_tokens:
raise ValueError(
"thinking_token_budget must be smaller than max_tokens")
# Reasoning parsers need the literal think tags in decoded text
self.extra_body["skip_special_tokens"] = False
template_kwargs = self.extra_body.setdefault(
"chat_template_kwargs", {})
self.extra_body["thinking_token_budget"] = value
if value == 0:
# Treat `0` as a request for no thinking at all
template_kwargs["add_thinking_start"] = False
template_kwargs["enable_thinking"] = False
else:
template_kwargs["add_thinking_start"] = True
template_kwargs.pop("enable_thinking", None)

@property
def best_of(self) -> int:
return self.extra_body['best_of']
Expand Down Expand Up @@ -187,6 +219,32 @@ def validate_request(self):
# If beam search is enabled, temperature must be set to 0.0
if self.beam_search:
assert self.temperature == 0.0, "Beam search requires temperature 0"

requested_budget = self.extra_body.get("thinking_token_budget")
add_thinking_start = self.extra_body.get(
"chat_template_kwargs", {}).get("add_thinking_start")
if requested_budget is not None and requested_budget < 0:
raise ValueError("thinking_token_budget must be positive")
if requested_budget:
if requested_budget >= self.max_tokens:
raise ValueError(
"thinking_token_budget must be smaller than max_tokens")
if add_thinking_start is not True:
raise ValueError("add_thinking_start must be True if "
"thinking_token_budget is set")
# Reasoning parsers need the literal think tags in decoded text
if self.extra_body.get("skip_special_tokens") is True:
raise ValueError("skip_special_tokens must be False if "
"add_thinking_start is True")
self.thinking_token_budget = requested_budget
elif requested_budget == 0 or add_thinking_start is False:
# Either parameter alone means no thinking is requested; set both
# so the request cannot ask the model to open a think block that
# it has no budget to complete
self.thinking_token_budget = 0
elif add_thinking_start is True:
raise ValueError("thinking_token_budget must be set if "
"add_thinking_start is True")
return self

@property
Expand All @@ -210,16 +268,25 @@ def to_completion_kwargs(self, mq2role: dict = None) -> dict:
history.insert(0, {"role": "system",
"content": self.persona.system_prompt})
history.append({"role": "user", "content": self.query})
extra_body = dict(self.extra_body)
if extra_body.get("thinking_token_budget") == 0:
# A zero budget requests no thinking, but an inference engine reads
# it as a budget to spend; it opens a think block and closes it
# after one generated token, leaving the model to complete its
# interrupted reasoning as response content
extra_body.pop("thinking_token_budget")
return {"model": self.model,
"messages": history,
"max_tokens": self.max_tokens,
"temperature": self.temperature,
"stream": self.stream,
"extra_body": self.extra_body}
"extra_body": extra_body}


class LLMResponse(BaseModel):
response: str = Field(description="LLM Response to the input query")
reasoning: Optional[str] = Field(
None, description="Thinking/reasoning trace returned by the LLM")
history: List[Tuple[LlmMessageRole, str]] = Field(
description="List of (role, content) tuples in chronological order "
"(`response` is in the last list element)")
Expand Down
148 changes: 148 additions & 0 deletions tests/models/api/test_llm.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,6 +89,7 @@ def test_llm_request(self):
self.assertIsInstance(valid_request.persona, LLMPersona)
self.assertTrue(valid_request.stream)
self.assertFalse(valid_request.beam_search)
self.assertIsNone(valid_request.thinking_token_budget)
self.assertEqual(len(valid_request.history), len(test_history))
self.assertEqual(len(valid_request.to_completion_kwargs()['messages']),
2 * valid_request.max_history + 2)
Expand Down Expand Up @@ -122,6 +123,87 @@ def test_llm_request(self):
self.assertFalse(valid_no_stream.beam_search)
self.assertFalse(valid_no_stream.stream)

# Valid thinking_token_budget via extra_body
thinking_request = LLMRequest(
query=test_query, history=test_history, persona=test_persona,
model=test_model,
extra_body={"thinking_token_budget": 128,
"chat_template_kwargs": {"add_thinking_start": True}})
self.assertEqual(thinking_request.thinking_token_budget, 128)
self.assertTrue(thinking_request.extra_body["chat_template_kwargs"][
"add_thinking_start"])
self.assertFalse(thinking_request.extra_body["skip_special_tokens"])
self.assertEqual(
thinking_request.to_completion_kwargs()["extra_body"][
"thinking_token_budget"], 128)

# Valid thinking with explicit skip_special_tokens=False
thinking_no_skip = LLMRequest(
query=test_query, history=test_history, persona=test_persona,
model=test_model,
extra_body={"thinking_token_budget": 128,
"skip_special_tokens": False,
"chat_template_kwargs": {"add_thinking_start": True}})
self.assertFalse(thinking_no_skip.extra_body["skip_special_tokens"])

# A zero thinking_token_budget requests no thinking at all
zero_thinking = LLMRequest(
query=test_query, history=test_history, persona=test_persona,
model=test_model,
extra_body={"thinking_token_budget": 0,
"chat_template_kwargs": {"add_thinking_start": True}})
self.assertEqual(zero_thinking.thinking_token_budget, 0)
self.assertFalse(zero_thinking.extra_body["skip_special_tokens"])
self.assertFalse(zero_thinking.extra_body["chat_template_kwargs"][
"add_thinking_start"])
self.assertFalse(zero_thinking.extra_body["chat_template_kwargs"][
"enable_thinking"])
self.assertNotIn("thinking_token_budget",
zero_thinking.to_completion_kwargs()["extra_body"])

# Thinking may be disabled with `add_thinking_start` alone
disabled_thinking = LLMRequest(
query=test_query, history=test_history, persona=test_persona,
model=test_model,
extra_body={"chat_template_kwargs": {"add_thinking_start": False}})
self.assertEqual(disabled_thinking.thinking_token_budget, 0)
self.assertFalse(disabled_thinking.extra_body["skip_special_tokens"])
self.assertFalse(disabled_thinking.extra_body["chat_template_kwargs"][
"enable_thinking"])
self.assertNotIn(
"thinking_token_budget",
disabled_thinking.to_completion_kwargs()["extra_body"])

# thinking_token_budget property setter enables chat_template_kwargs
setter_request = LLMRequest(query=test_query, history=test_history,
persona=test_persona, model=test_model)
setter_request.thinking_token_budget = 64
self.assertEqual(setter_request.thinking_token_budget, 64)
self.assertTrue(setter_request.extra_body["chat_template_kwargs"][
"add_thinking_start"])
self.assertFalse(setter_request.extra_body["skip_special_tokens"])
# A zero budget via the property setter disables thinking
setter_request.thinking_token_budget = 0
self.assertEqual(setter_request.thinking_token_budget, 0)
self.assertNotIn("thinking_token_budget",
setter_request.to_completion_kwargs()["extra_body"])
self.assertFalse(setter_request.extra_body["chat_template_kwargs"][
"add_thinking_start"])
self.assertFalse(setter_request.extra_body["chat_template_kwargs"][
"enable_thinking"])
self.assertFalse(setter_request.extra_body["skip_special_tokens"])
# Clearing thinking_token_budget also clears add_thinking_start
setter_request.thinking_token_budget = None
self.assertIsNone(setter_request.thinking_token_budget)
self.assertNotIn("thinking_token_budget", setter_request.extra_body)
self.assertNotIn("add_thinking_start",
setter_request.extra_body.get(
"chat_template_kwargs", {}))
self.assertNotIn("enable_thinking",
setter_request.extra_body.get(
"chat_template_kwargs", {}))
self.assertNotIn("skip_special_tokens", setter_request.extra_body)

# Validate `llm` history input
old_history = [("user", "hello"),
("llm", "Hi, how can I help you today?"),
Expand Down Expand Up @@ -151,6 +233,65 @@ def test_llm_request(self):
LLMRequest(query=test_query, history=test_history,
persona=test_persona, model=test_model, stream=False,
beam_search=True, best_of=1)
# Invalid thinking_token_budget without add_thinking_start
with self.assertRaises(ValidationError):
LLMRequest(query=test_query, history=test_history,
persona=test_persona, model=test_model,
extra_body={"thinking_token_budget": 128})
# Invalid add_thinking_start without thinking_token_budget
with self.assertRaises(ValidationError):
LLMRequest(
query=test_query, history=test_history, persona=test_persona,
model=test_model,
extra_body={"chat_template_kwargs": {
"add_thinking_start": True}})
# Invalid negative thinking_token_budget
with self.assertRaises(ValidationError):
LLMRequest(
query=test_query, history=test_history, persona=test_persona,
model=test_model,
extra_body={"thinking_token_budget": -1,
"chat_template_kwargs": {
"add_thinking_start": True}})
# Invalid add_thinking_start=False with thinking_token_budget
with self.assertRaises(ValidationError):
LLMRequest(
query=test_query, history=test_history, persona=test_persona,
model=test_model,
extra_body={"thinking_token_budget": 128,
"chat_template_kwargs": {
"add_thinking_start": False}})
# Invalid skip_special_tokens=True with thinking enabled
with self.assertRaises(ValidationError):
LLMRequest(
query=test_query, history=test_history, persona=test_persona,
model=test_model,
extra_body={"thinking_token_budget": 128,
"skip_special_tokens": True,
"chat_template_kwargs": {
"add_thinking_start": True}})
# Invalid thinking_token_budget equal to max_tokens
with self.assertRaises(ValidationError):
LLMRequest(
query=test_query, history=test_history, persona=test_persona,
model=test_model, max_tokens=128,
extra_body={"thinking_token_budget": 128,
"chat_template_kwargs": {
"add_thinking_start": True}})
# Invalid thinking_token_budget greater than max_tokens
with self.assertRaises(ValidationError):
LLMRequest(
query=test_query, history=test_history, persona=test_persona,
model=test_model, max_tokens=128,
extra_body={"thinking_token_budget": 129,
"chat_template_kwargs": {
"add_thinking_start": True}})
# Invalid negative thinking_token_budget via property setter
with self.assertRaises(AssertionError):
setter_request.thinking_token_budget = -1
# Invalid thinking_token_budget via property setter
with self.assertRaises(ValueError):
setter_request.thinking_token_budget = setter_request.max_tokens
# Invalid history
test_history.append(("invalid_key", "okay"))
with self.assertRaises(ValidationError):
Expand All @@ -167,8 +308,15 @@ def test_llm_response(self):
# Valid response with valid history
response = LLMResponse(response=valid_response, history=valid_history)
self.assertEqual(response.response, valid_response)
self.assertIsNone(response.reasoning)
self.assertEqual(response.history, valid_history)

# Valid response with a reasoning trace
reasoning = "We need answer the user's greeting."
response = LLMResponse(response=valid_response, history=valid_history,
reasoning=reasoning)
self.assertEqual(response.reasoning, reasoning)

# Valid response with legacy history
response = LLMResponse(response=valid_response, history=legacy_history)
self.assertEqual(response.response, valid_response)
Expand Down
Loading