Skip to content

Commit 2a9128f

Browse files
authored
fix: add a mechanism for allowing more max tokens (#662)
<!-- SPDX-FileCopyrightText: Copyright (c) 2025-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. --> <!-- SPDX-License-Identifier: Apache-2.0 --> <!-- Thank you for contributing to Safe Synthesizer! --> # Summary <!-- Brief description of changes --> ## Pre-Review Checklist <!-- These checks should be completed before a PR is reviewed, --> <!-- but you can submit a draft early to indicate that the issue is being worked on. --> Ensure that the following pass: - [ ] `mise run format && mise run check` or via prek validation. - [ ] `mise run test` passes locally - [ ] `mise run test:e2e` passes locally - [ ] `mise run test:ci-container` passes locally (recommended) - [ ] GPU CI status check passes -- comment `/sync` on this PR to trigger a run (auto-triggers on ready-for-review) ## Pre-Merge Checklist <!-- These checks need to be completed before a PR is merged, --> <!-- but as PRs often change significantly during review, --> <!-- it's OK for them to be incomplete when review is first requested. --> - [ ] New or updated tests for any fix or new behavior - [ ] Updated documentation for new features and behaviors, including docstrings for API docs. ## Other Notes <!-- Please add the issue number that should be closed when this PR is merged. --> - Closes #<issue> <!-- This is an auto-generated comment: release notes by coderabbit.ai --> ## Summary by CodeRabbit * **New Features** * Added `max_tokens_multiplier` to control the per-request generation token budget, defaulting to **1.2**. * Enforced validation to allow only finite values strictly **greater than 0**. * **Bug Fixes** * Improved consistency of max-token budgeting across generation backends by applying the configured multiplier while keeping existing context-window clamping behavior. * **Tests** * Added unit tests for the new multiplier field (default, valid/invalid values including non-finite cases) and for multiplier behavior (including clamping). <!-- end of auto-generated comment: release notes by coderabbit.ai --> --------- Signed-off-by: Matt Kornfield <mkornfield@nvidia.com>
1 parent 4eeb170 commit 2a9128f

7 files changed

Lines changed: 91 additions & 9 deletions

File tree

‎src/nemo_safe_synthesizer/config/generate.py‎

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33

44
from __future__ import annotations
55

6+
import math
67
import warnings
78
from collections.abc import Mapping
89
from typing import Annotated, Any, ClassVar, Literal, Self
@@ -264,6 +265,18 @@ class GenerateParameters(Parameters, BaseModel):
264265
),
265266
] = 0.8
266267

268+
max_tokens_multiplier: Annotated[
269+
float,
270+
ValueValidator(value_func=lambda v: math.isfinite(v) and v > 0),
271+
Field(
272+
title="max_tokens_multiplier",
273+
description=(
274+
"Multiplier on the longest training example when sizing per-sample "
275+
"max_tokens. Must be a finite value > 0. Default 1.2."
276+
),
277+
),
278+
] = 1.2 # mirrors llm.metadata.GENERATION_MAX_TOKENS_SAFETY_MULTIPLIER (kept a literal to avoid a config->llm import cycle)
279+
267280
structured_generation: StructuredGenerationParameters = Field(
268281
description="Structured generation parameters controlling schema-constrained output.",
269282
default_factory=StructuredGenerationParameters,

‎src/nemo_safe_synthesizer/generation/timeseries_backend.py‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -921,7 +921,10 @@ def generate(
921921
top_p=self.config.generation.top_p,
922922
top_k=FIXED_RUNTIME_GENERATE_ARGS["top_k"],
923923
min_p=FIXED_RUNTIME_GENERATE_ARGS["min_p"],
924-
max_tokens=self.model_metadata.generation_max_tokens_for(self._get_prompt_token_count()),
924+
max_tokens=self.model_metadata.generation_max_tokens_for(
925+
self._get_prompt_token_count(),
926+
multiplier=self.config.generation.max_tokens_multiplier,
927+
),
925928
skip_special_tokens=True,
926929
include_stop_str_in_output=False,
927930
ignore_eos=False,

‎src/nemo_safe_synthesizer/generation/vllm_backend.py‎

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -746,7 +746,10 @@ def _run_generation(self, data_actions_fn: utils.DataActionsFn | None) -> None:
746746
top_p=self.config.generation.top_p,
747747
top_k=FIXED_RUNTIME_GENERATE_ARGS["top_k"],
748748
min_p=FIXED_RUNTIME_GENERATE_ARGS["min_p"],
749-
max_tokens=self.model_metadata.generation_max_tokens_for(self._get_prompt_token_count()),
749+
max_tokens=self.model_metadata.generation_max_tokens_for(
750+
self._get_prompt_token_count(),
751+
multiplier=self.config.generation.max_tokens_multiplier,
752+
),
750753
skip_special_tokens=not need_special_token_outputs,
751754
include_stop_str_in_output=need_special_token_outputs,
752755
ignore_eos=False,

‎src/nemo_safe_synthesizer/llm/metadata.py‎

Lines changed: 18 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -470,13 +470,13 @@ def max_seq_length(self) -> int:
470470
rsf = self.rope_scaling.factor
471471
return int((self.base_max_seq_length or DEFAULT_MAX_SEQ_LENGTH) * rsf)
472472

473-
def generation_max_tokens_for(self, prompt_len: int) -> int:
473+
def generation_max_tokens_for(self, prompt_len: int, multiplier: float | None = None) -> int:
474474
"""Per-sample ``max_tokens`` ceiling, prompt-aware.
475475
476476
Returns the smaller of:
477477
478-
1. ``int(max_tokens_per_example * GENERATION_MAX_TOKENS_SAFETY_MULTIPLIER)``
479-
when the assembler stat is populated, else ``max_seq_length``.
478+
1. ``int(max_tokens_per_example * multiplier)`` when the assembler stat
479+
is populated, else ``max_seq_length``.
480480
2. ``max_seq_length - prompt_len`` -- vLLM raises when
481481
``len(prompt) + max_tokens > max_model_len``
482482
(`vllm#33418 <https://github.com/vllm-project/vllm/issues/33418>`_).
@@ -488,15 +488,29 @@ def generation_max_tokens_for(self, prompt_len: int) -> int:
488488
clamp is a defensive belt for legacy adapters where the assembler
489489
stat is missing and for prompts longer than those seen in training.
490490
491+
The default ``multiplier`` (``GENERATION_MAX_TOKENS_SAFETY_MULTIPLIER``)
492+
adds only a small jitter margin, which is enough for most tables but
493+
too tight for long, unbounded free-text columns: a model that
494+
over-generates slightly past the longest training example truncates
495+
mid-JSON and yields no parseable record. Callers wire the user-facing
496+
``generation.max_tokens_multiplier`` knob through here to widen the
497+
budget (bounded by the context window) for such datasets.
498+
491499
Args:
492500
prompt_len: Tokenized length of the prompt this sample will
493501
run against. Pass ``0`` to disable the prompt clamp.
502+
multiplier: Margin applied to ``max_tokens_per_example``. Defaults
503+
to ``GENERATION_MAX_TOKENS_SAFETY_MULTIPLIER`` when ``None`` so
504+
non-generation callers (e.g. the training eval callback) keep
505+
the legacy sizing.
494506
495507
Returns:
496508
Non-negative ``max_tokens`` value safe to feed to ``SamplingParams``.
497509
"""
510+
if multiplier is None:
511+
multiplier = GENERATION_MAX_TOKENS_SAFETY_MULTIPLIER
498512
if self.max_tokens_per_example and self.max_tokens_per_example > 0:
499-
sized = int(self.max_tokens_per_example * GENERATION_MAX_TOKENS_SAFETY_MULTIPLIER)
513+
sized = int(self.max_tokens_per_example * multiplier)
500514
else:
501515
sized = self.max_seq_length
502516
return max(0, min(sized, self.max_seq_length - prompt_len))

‎tests/config/test_generate.py‎

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -217,3 +217,22 @@ def test_legacy_keys_with_no_nested_dict_are_migrated(self) -> None:
217217
assert params.structured_generation.backend == "xgrammar"
218218
assert params.structured_generation.schema_method == "structural_tag"
219219
assert params.structured_generation.use_single_sequence is True
220+
221+
222+
@pytest.mark.unit
223+
class TestMaxTokensMultiplier:
224+
def test_default_matches_metadata_constant(self) -> None:
225+
"""The config default mirrors the metadata safety-margin constant."""
226+
from nemo_safe_synthesizer.llm.metadata import GENERATION_MAX_TOKENS_SAFETY_MULTIPLIER
227+
228+
assert GenerateParameters().max_tokens_multiplier == GENERATION_MAX_TOKENS_SAFETY_MULTIPLIER
229+
230+
def test_accepts_widened_value(self) -> None:
231+
"""Users can widen the budget for long free-text datasets."""
232+
assert GenerateParameters(max_tokens_multiplier=1.8).max_tokens_multiplier == 1.8
233+
234+
@pytest.mark.parametrize("value", [0, -0.5, float("inf"), float("-inf"), float("nan")])
235+
def test_rejects_non_positive(self, value: float) -> None:
236+
"""Non-positive and non-finite multipliers are rejected by the validator."""
237+
with pytest.raises(ValidationError):
238+
GenerateParameters(max_tokens_multiplier=value)

‎tests/generation/test_vllm_backend.py‎

Lines changed: 7 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1048,8 +1048,10 @@ def capture_and_stop(**kwargs):
10481048
assert captured["max_tokens"] == 4_200
10491049
# Engine is not initialized in this plumbing test, so the cached
10501050
# prompt-token count falls back to 0; the helper is still called
1051-
# exactly once with that value.
1052-
mock_model_metadata.generation_max_tokens_for.assert_called_once_with(0)
1051+
# exactly once with that value plus the configured budget multiplier.
1052+
mock_model_metadata.generation_max_tokens_for.assert_called_once_with(
1053+
0, multiplier=base_params.generation.max_tokens_multiplier
1054+
)
10531055

10541056
def test_passes_cached_prompt_token_count_when_engine_initialized(
10551057
self, base_params, mock_model_metadata, mock_schema, mock_workdir
@@ -1079,7 +1081,9 @@ def capture_and_stop(**kwargs):
10791081
backend.generate()
10801082

10811083
assert captured["max_tokens"] == 4_096
1082-
mock_model_metadata.generation_max_tokens_for.assert_called_once_with(37)
1084+
mock_model_metadata.generation_max_tokens_for.assert_called_once_with(
1085+
37, multiplier=base_params.generation.max_tokens_multiplier
1086+
)
10831087
# Cached: a second access does not retokenize.
10841088
assert backend._get_prompt_token_count() == 37
10851089
fake_tokenizer.encode.assert_called_once_with(backend.prompt)

‎tests/llm/test_metadata.py‎

Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -689,6 +689,32 @@ def test_metadata_generation_max_tokens_for_never_returns_negative(self, sample_
689689
# Beyond-context prompts still produce a safe value vLLM can accept.
690690
assert sample_model_metadata.generation_max_tokens_for(sample_model_metadata.max_seq_length + 64) == 0
691691

692+
def test_metadata_generation_max_tokens_for_none_multiplier_uses_default(self, sample_model_metadata):
693+
"""``multiplier=None`` reproduces the legacy safety-margin sizing."""
694+
sample_model_metadata.max_tokens_per_example = 1000
695+
expected = int(1000 * GENERATION_MAX_TOKENS_SAFETY_MULTIPLIER)
696+
assert sample_model_metadata.generation_max_tokens_for(10, multiplier=None) == expected
697+
# Explicitly passing the default constant matches the None path.
698+
assert (
699+
sample_model_metadata.generation_max_tokens_for(10, multiplier=GENERATION_MAX_TOKENS_SAFETY_MULTIPLIER)
700+
== expected
701+
)
702+
703+
def test_metadata_generation_max_tokens_for_custom_multiplier_widens_budget(self, sample_model_metadata):
704+
"""A larger multiplier widens the stat-derived budget (the long-text fix)."""
705+
sample_model_metadata.max_tokens_per_example = 1000
706+
# 1000 * 1.8 = 1800, still within the 2048 window given a small prompt.
707+
assert sample_model_metadata.generation_max_tokens_for(10, multiplier=1.8) == 1800
708+
709+
def test_metadata_generation_max_tokens_for_custom_multiplier_still_clamped_to_window(self, sample_model_metadata):
710+
"""The prompt/window clamp still binds even with an aggressive multiplier."""
711+
assert sample_model_metadata.max_seq_length == 2048
712+
sample_model_metadata.max_tokens_per_example = 1500 # * 4.0 = 6000, far past the window
713+
prompt_len = 48
714+
assert sample_model_metadata.generation_max_tokens_for(prompt_len, multiplier=4.0) == (
715+
sample_model_metadata.max_seq_length - prompt_len
716+
)
717+
692718
def test_metadata_max_tokens_per_example_round_trips_through_metadata_json(self, sample_model_metadata):
693719
"""``max_tokens_per_example`` persists through save → load."""
694720
sample_model_metadata.max_tokens_per_example = 1500

0 commit comments

Comments
 (0)