(self, is_training)
| 60 | |
| 61 | @parameterized.parameters(False, True) |
| 62 | def test_concatenation(self, is_training): |
| 63 | cfg = ConcatenatedInput.default_config().set( |
| 64 | name="input", |
| 65 | is_training=is_training, |
| 66 | inputs=[ |
| 67 | self._input_config([1, 2, 3], batch_size=2), |
| 68 | self._input_config([11, 12, 13, 14, 15, 16], batch_size=3, repeat=None), |
| 69 | ], |
| 70 | ) |
| 71 | dataset = cfg.instantiate(parent=None) |
| 72 | batch_index = 0 |
| 73 | expected_train_batches = [ |
| 74 | {"index": jnp.asarray([0, 1]), "number": jnp.asarray([1, 2])}, |
| 75 | {"index": jnp.asarray([0, 1, 2]), "number": jnp.asarray([11, 12, 13])}, |
| 76 | {"index": jnp.asarray([3, 4, 5]), "number": jnp.asarray([14, 15, 16])}, |
| 77 | {"index": jnp.asarray([0, 1, 2]), "number": jnp.asarray([11, 12, 13])}, |
| 78 | {"index": jnp.asarray([3, 4, 5]), "number": jnp.asarray([14, 15, 16])}, |
| 79 | ] |
| 80 | expected_eval_batches = [ |
| 81 | {"index": jnp.asarray([0, 1]), "number": jnp.asarray([1, 2])}, |
| 82 | {"index": jnp.asarray([2, 0]), "number": jnp.asarray([3, 0])}, |
| 83 | {"index": jnp.asarray([0, 1, 2]), "number": jnp.asarray([11, 12, 13])}, |
| 84 | {"index": jnp.asarray([3, 4, 5]), "number": jnp.asarray([14, 15, 16])}, |
| 85 | ] |
| 86 | expected_batches = expected_train_batches if is_training else expected_eval_batches |
| 87 | for batch in dataset.dataset(): |
| 88 | print(batch) |
| 89 | if batch_index >= len(expected_batches): |
| 90 | break |
| 91 | self.assertNestedAllClose(as_tensor(expected_batches[batch_index]), batch) |
| 92 | batch_index += 1 |
| 93 | self.assertEqual(batch_index, len(expected_batches)) |
| 94 | |
| 95 | @parameterized.parameters(False, True) |
| 96 | def test_zipinput(self, is_training): |
nothing calls this directly
no test coverage detected