fix(sampling): reject non-finite temperature in SamplingParams.verify (#28153)
Signed-off-by: Ting Sun <suntcrick@gmail.com>
This commit is contained in:
@@ -14,6 +14,7 @@
|
|||||||
"""Sampling parameters for text generation."""
|
"""Sampling parameters for text generation."""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import math
|
||||||
from typing import Any, Dict, List, Optional, Union
|
from typing import Any, Dict, List, Optional, Union
|
||||||
|
|
||||||
# sre_parse is deprecated in Python 3.11+, use re._parser instead
|
# sre_parse is deprecated in Python 3.11+, use re._parser instead
|
||||||
@@ -118,9 +119,9 @@ class SamplingParams:
|
|||||||
self.top_k = TOP_K_ALL # whole vocabulary
|
self.top_k = TOP_K_ALL # whole vocabulary
|
||||||
|
|
||||||
def verify(self, vocab_size):
|
def verify(self, vocab_size):
|
||||||
if self.temperature < 0.0:
|
if not math.isfinite(self.temperature) or self.temperature < 0.0:
|
||||||
raise ValueError(
|
raise ValueError(
|
||||||
f"temperature must be non-negative, got {self.temperature}."
|
f"temperature must be a non-negative finite number, got {self.temperature}."
|
||||||
)
|
)
|
||||||
if not 0.0 < self.top_p <= 1.0:
|
if not 0.0 < self.top_p <= 1.0:
|
||||||
raise ValueError(f"top_p must be in (0, 1], got {self.top_p}.")
|
raise ValueError(f"top_p must be in (0, 1], got {self.top_p}.")
|
||||||
|
|||||||
@@ -92,6 +92,18 @@ class TestSamplingParamsVerify(CustomTestCase):
|
|||||||
with self.assertRaises(ValueError):
|
with self.assertRaises(ValueError):
|
||||||
sp.verify(self.VOCAB_SIZE)
|
sp.verify(self.VOCAB_SIZE)
|
||||||
|
|
||||||
|
def test_nan_temperature_raises(self):
|
||||||
|
"""verify() must reject NaN temperature; the bare < 0.0 check alone lets it through."""
|
||||||
|
sp = self._make(temperature=float("nan"))
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
sp.verify(self.VOCAB_SIZE)
|
||||||
|
|
||||||
|
def test_inf_temperature_raises(self):
|
||||||
|
"""verify() must reject non-finite (inf) temperature."""
|
||||||
|
sp = self._make(temperature=float("inf"))
|
||||||
|
with self.assertRaises(ValueError):
|
||||||
|
sp.verify(self.VOCAB_SIZE)
|
||||||
|
|
||||||
# --- top_p ---
|
# --- top_p ---
|
||||||
def test_top_p_negative_raises(self):
|
def test_top_p_negative_raises(self):
|
||||||
"""Test that verify() rejects negative top_p (valid range is (0, 1])."""
|
"""Test that verify() rejects negative top_p (valid range is (0, 1])."""
|
||||||
|
|||||||
Reference in New Issue
Block a user