MCPcopy Create free account
hub / github.com/apple/axlearn / test_parent_forward

Method test_parent_forward

axlearn/common/base_layer_test.py:223–240  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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):

Callers

nothing calls this directly

Calls 5

OutputCollectionClass · 0.90
instantiateMethod · 0.45
setMethod · 0.45
default_configMethod · 0.45

Tested by

no test coverage detected