Skip to content

Commit 154fd48

Browse files
Fix CLIPTokenizer inconsistency between Python and TF paths (#2732)
1 parent c5938a2 commit 154fd48

2 files changed

Lines changed: 24 additions & 36 deletions

File tree

keras_hub/src/models/clip/clip_tokenizer.py

Lines changed: 3 additions & 36 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,8 @@
1-
import inspect
2-
31
import tokenizers
42
from tokenizers import decoders
53
from tokenizers import models
64
from tokenizers import normalizers
75
from tokenizers import pre_tokenizers
8-
from tokenizers import processors
96

107
from keras_hub.src.api_export import keras_hub_export
118
from keras_hub.src.models.clip.clip_backbone import CLIPBackbone
@@ -153,24 +150,6 @@ def set_vocabulary_and_merges(self, vocabulary, merges):
153150
super().set_vocabulary_and_merges(vocabulary, merges)
154151
if self.pad_with_end_token:
155152
self.pad_token_id = self.end_token_id
156-
if getattr(self, "_tokenizer") is not None:
157-
# tokenizers <=0.22 use `cls`, >= 0.23 use `cls_token` because `cls`
158-
# collides with the first argument of the Python `__new__`.
159-
cls_token_arg = (
160-
"cls"
161-
if "cls"
162-
in inspect.signature(processors.RobertaProcessing).parameters
163-
else "cls_token"
164-
)
165-
preprocessing_args = {
166-
"sep": (str(self.end_token), self.end_token_id),
167-
cls_token_arg: (str(self.start_token), self.start_token_id),
168-
"add_prefix_space": False,
169-
"trim_offsets": False,
170-
}
171-
self._tokenizer.post_processor = processors.RobertaProcessing(
172-
**preprocessing_args
173-
)
174153

175154
def _bpe_merge_and_update_cache_tf(self, tokens):
176155
"""Process unseen tokens and add to cache."""
@@ -254,28 +233,16 @@ def process_unseen_tokens():
254233
if self.sequence_length:
255234
output_shape = tokens.shape.as_list()
256235
output_shape[-1] = self.sequence_length
257-
tokens = tokens.to_tensor(shape=output_shape)
236+
tokens = tokens.to_tensor(
237+
shape=output_shape, default_value=self.pad_token_id
238+
)
258239

259240
# Convert to a dense output if input in scalar
260241
if unbatched:
261242
tokens = tf.squeeze(tokens, 0)
262243
tf.ensure_shape(tokens, shape=[self.sequence_length])
263244
return tokens
264245

265-
def _tokenize_tokenizers(self, inputs):
266-
outputs = super()._tokenize_tokenizers(inputs)
267-
is_batched = True
268-
if isinstance(outputs, str):
269-
is_batched = False
270-
outputs = [outputs]
271-
elif isinstance(outputs, list) and isinstance(outputs[0], int):
272-
is_batched = False
273-
outputs = [outputs]
274-
outputs = [output[1:-1] for output in outputs]
275-
if not is_batched:
276-
outputs = outputs[0]
277-
return outputs
278-
279246
@preprocessing_function
280247
def _detokenize_tf(self, inputs):
281248
self._check_vocabulary()

keras_hub/src/models/clip/clip_tokenizer_test.py

Lines changed: 21 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
import pytest
2+
import tensorflow as tf
23

34
from keras_hub.src.models.clip.clip_tokenizer import CLIPTokenizer
45
from keras_hub.src.tests.test_case import TestCase
@@ -35,6 +36,26 @@ def test_pad_with_end_token(self):
3536
tokenizer = CLIPTokenizer(**init_kwargs)
3637
self.assertEqual(tokenizer.pad_token_id, tokenizer.end_token_id)
3738

39+
def test_python_tf_consistency(self):
40+
init_kwargs = self.init_kwargs.copy()
41+
init_kwargs["sequence_length"] = 10
42+
init_kwargs["pad_with_end_token"] = True
43+
tokenizer = CLIPTokenizer(**init_kwargs)
44+
input_data = ["airplane", "airplane airport"]
45+
46+
# Python workflow
47+
python_output = tokenizer(input_data)
48+
49+
# TF workflow
50+
ds = tf.data.Dataset.from_tensor_slices(input_data)
51+
ds = ds.map(tokenizer)
52+
tf_outputs = list(ds.as_numpy_iterator())
53+
54+
self.assertAllEqual(python_output, tf_outputs)
55+
self.assertAllEqual(python_output[0][-1], tokenizer.pad_token_id)
56+
self.assertAllEqual(tf_outputs[0][-1], tokenizer.pad_token_id)
57+
self.assertAllEqual(tokenizer.pad_token_id, tokenizer.end_token_id)
58+
3859
def test_errors_missing_special_tokens(self):
3960
with self.assertRaises(ValueError):
4061
CLIPTokenizer(vocabulary={"foo": 0, "bar": 1}, merges=["fo o"])

0 commit comments

Comments
 (0)