(self)
| 221 | self.assertEqual(output_collection.module_outputs, {}) |
| 222 | |
| 223 | def test_parent_forward(self): |
| 224 | test_module: _ParentLayer = ( |
| 225 | _ParentLayer.default_config().set(name="test").instantiate(parent=None) |
| 226 | ) |
| 227 | state = test_module.initialize_parameters_recursively(prng_key=jax.random.PRNGKey(456)) |
| 228 | self.assertEqual({"child": {"moving_mean": 1.0}}, state) |
| 229 | y, output_collection = jax.jit(partial(F, test_module, is_training=True))( |
| 230 | prng_key=jax.random.PRNGKey(123), state=state, inputs=(jnp.asarray(5.0),) |
| 231 | ) |
| 232 | self.assertEqual(4, y) |
| 233 | self.assertEqual( |
| 234 | OutputCollection( |
| 235 | summaries={"child": {"x": 5}}, |
| 236 | state_updates={"child": {"moving_mean": 1.4}}, |
| 237 | module_outputs={}, |
| 238 | ), |
| 239 | output_collection, |
| 240 | ) |
| 241 | |
| 242 | @parameterized.parameters(("true", 10), ("false", 7)) |
| 243 | def test_not_too_many_stack_frames(self, enable_traceback, expected_frames): |
nothing calls this directly
no test coverage detected