(self, *, vocab_cfg: InstantiableConfig, newlines_replaced_with: str)
| 267 | |
| 268 | class NumBytesTest(test_utils.TestCase): |
| 269 | def _test_num_bytes(self, *, vocab_cfg: InstantiableConfig, newlines_replaced_with: str): |
| 270 | vocab = vocab_cfg.instantiate() |
| 271 | |
| 272 | pad_id = vocab.pad_id |
| 273 | newline_id = vocab.encode("\n").pop() |
| 274 | newlines_replaced_with_id = vocab.encode(newlines_replaced_with).pop() |
| 275 | |
| 276 | # Test num_bytes computes expected value. |
| 277 | ids = tf.constant( |
| 278 | [vocab.eos_id, newlines_replaced_with_id, newline_id, pad_id, pad_id, pad_id], |
| 279 | dtype=tf.int32, |
| 280 | ) |
| 281 | self.assertEqual( |
| 282 | 3, |
| 283 | input_text.num_bytes( |
| 284 | ids, sp_vocab=vocab, newlines_replaced_with=newlines_replaced_with |
| 285 | ), |
| 286 | ) |
| 287 | |
| 288 | @parameterized.parameters( |
| 289 | dict( |
no test coverage detected