(
self,
numbers: list[int],
*,
batch_size=2,
repeat=1,
out_signature="number",
)
| 39 | |
| 40 | class BatchTest(TestCase): |
| 41 | def _input_config( |
| 42 | self, |
| 43 | numbers: list[int], |
| 44 | *, |
| 45 | batch_size=2, |
| 46 | repeat=1, |
| 47 | out_signature="number", |
| 48 | ) -> input_tf_data.Input.Config: |
| 49 | return input_tf_data.Input.default_config().set( |
| 50 | source=config_for_function(make_ds_fn).set( |
| 51 | numbers=numbers, out_signature=out_signature |
| 52 | ), |
| 53 | processor=config_for_function(input_tf_data.identity), |
| 54 | batcher=config_for_function(input_tf_data.batch).set( |
| 55 | global_batch_size=batch_size, |
| 56 | pad_example_fn=input_tf_data.default_pad_example_fn, |
| 57 | repeat=repeat, |
| 58 | ), |
| 59 | ) |
| 60 | |
| 61 | @parameterized.parameters(False, True) |
| 62 | def test_concatenation(self, is_training): |
no test coverage detected