(self, dtype, remat_spec, drop_output, num_layers_total, unroll)
| 162 | unroll=(True, False, 1, 2), |
| 163 | ) |
| 164 | def test_repeat(self, dtype, remat_spec, drop_output, num_layers_total, unroll): |
| 165 | batch_size, num_layers = 14, 4 |
| 166 | cfg = _Ensemble.default_config().set(name="test", num_layers=num_layers_total, dtype=dtype) |
| 167 | cfg.repeat_layer.set( |
| 168 | remat_spec=remat_spec, |
| 169 | drop_output=drop_output, |
| 170 | unroll=unroll, |
| 171 | ) |
| 172 | layer: _Ensemble = cfg.instantiate(parent=None) |
| 173 | self.assertEqual( |
| 174 | PartitionSpec(None), |
| 175 | layer.create_parameter_specs_recursively()["repeat_layer"]["layer"]["inc"].mesh_axes, |
| 176 | ) |
| 177 | layer_params = layer.initialize_parameters_recursively(prng_key=jax.random.PRNGKey(1)) |
| 178 | logging.info("layer params=%s", layer_params) |
| 179 | |
| 180 | input_forward_state = layer.init_forward_state(batch_size) |
| 181 | if num_layers_total == num_layers: |
| 182 | method = "forward" |
| 183 | inputs = dict( |
| 184 | carry=jnp.arange(batch_size, dtype=dtype), |
| 185 | forward_state=input_forward_state, |
| 186 | ) |
| 187 | else: |
| 188 | method = "forward_first_n" |
| 189 | input_forward_state = _get_first_n(num_layers, input_forward_state) |
| 190 | inputs = dict( |
| 191 | carry=jnp.arange(batch_size, dtype=dtype), |
| 192 | forward_state=input_forward_state, |
| 193 | n=num_layers, |
| 194 | ) |
| 195 | |
| 196 | (carry, output_forward_state), output_collection = F( |
| 197 | layer, |
| 198 | prng_key=jax.random.PRNGKey(2), |
| 199 | state=layer_params, |
| 200 | inputs=inputs, |
| 201 | method=method, |
| 202 | is_training=True, |
| 203 | drop_output_collections=(), |
| 204 | ) |
| 205 | logging.info("forward_state=%s", output_forward_state) |
| 206 | logging.info("output_collection=%s", output_collection) |
| 207 | assert_allclose(carry, jnp.arange(num_layers, num_layers + batch_size, dtype=dtype)) |
| 208 | self.assertEqual(shapes(input_forward_state), shapes(output_forward_state)) |
| 209 | assert_allclose( |
| 210 | output_forward_state["repeat_layer"]["layer"], |
| 211 | jnp.reshape( |
| 212 | jnp.arange(batch_size)[None, :] + jnp.arange(num_layers, dtype=dtype)[:, None], |
| 213 | (num_layers, batch_size), |
| 214 | ), |
| 215 | ) |
| 216 | # Check output collection. |
| 217 | self.assertEqual( |
| 218 | OutputCollection( |
| 219 | state_updates={ |
| 220 | "repeat_layer": { |
| 221 | # State update values are stacked across layers. |
nothing calls this directly
no test coverage detected