Skip to content

Commit 7baa88d

Browse files
committed
Fix sequence_length float coercion in ByteTokenizer
1 parent ee2b32f commit 7baa88d

2 files changed

Lines changed: 14 additions & 4 deletions

File tree

keras_hub/src/tokenizers/byte_tokenizer.py

Lines changed: 12 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -179,11 +179,21 @@ def __init__(
179179
"`sequence_length` must be an int, got bool: "
180180
f"{sequence_length}"
181181
)
182-
if not isinstance(sequence_length, numbers.Integral):
182+
183+
if isinstance(sequence_length, float):
184+
if not sequence_length.is_integer():
185+
raise ValueError(
186+
"`sequence_length` must be a whole number. "
187+
f"Received: {sequence_length}"
188+
)
189+
sequence_length = int(sequence_length)
190+
191+
elif not isinstance(sequence_length, numbers.Integral):
183192
raise ValueError(
184193
"`sequence_length` must be an int or None. "
185-
f"{sequence_length}"
194+
f"Received: {sequence_length}"
186195
)
196+
187197
if sequence_length <= 0:
188198
raise ValueError(
189199
"`sequence_length` must be > 0. "

keras_hub/src/tokenizers/byte_tokenizer_test.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -222,8 +222,8 @@ def test_sequence_length_valid_int(self):
222222
self.assertEqual(tokenizer.sequence_length, 5)
223223

224224
def test_sequence_length_valid_float(self):
225-
with self.assertRaises(ValueError):
226-
ByteTokenizer(sequence_length=5.0)
225+
tokenizer = ByteTokenizer(sequence_length=5.0)
226+
self.assertEqual(tokenizer.sequence_length, 5)
227227

228228
def test_sequence_length_invalid_float(self):
229229
with self.assertRaises(ValueError):

0 commit comments

Comments
 (0)