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

Method test_functional

axlearn/common/module_test.py:286–306  ·  view source on GitHub ↗

Tests the `Functional` class and `InvocationContext.functional()`.

(self)

Source from the content-addressed store, hash-verified

284 ctx.add_summary("summary", value)
285
286 def test_functional(self):
287 """Tests the `Functional` class and `InvocationContext.functional()`."""
288 with test_utils.bind_layer(Linear.default_config().set(input_dim=5, output_dim=5)) as layer:
289
290 def fn(x: Tensor, y: Tensor) -> Tensor:
291 current_context().add_state_update("my_state", y)
292 return layer(x)
293
294 args = [jnp.ones(5)]
295 kwargs = dict(y=jnp.zeros(3))
296 new_fn = current_context().functional(fn)
297 old_ctx = current_context()
298 result, output_collection = new_fn(*args, **kwargs)
299 self.assertIs(current_context(), old_ctx)
300 self.assertNestedEqual(current_context().output_collection, new_output_collection())
301 # The below line would cause an output conflict error if we had called fn() instead of
302 # new_fn() on the line above. But it doesn't since new_fn() restores the context to its
303 # original state after the call.
304 result2 = fn(*args, **kwargs)
305 self.assertNestedEqual(result, result2)
306 self.assertNestedEqual(output_collection, current_context().output_collection)
307
308 def test_functional_with_method_call(self):
309 """Demonstrates usage of `InvocationContext.functional()` with a module method instead of an

Callers

nothing calls this directly

Calls 7

current_contextFunction · 0.90
new_output_collectionFunction · 0.90
functionalMethod · 0.80
assertNestedEqualMethod · 0.80
fnFunction · 0.70
setMethod · 0.45
default_configMethod · 0.45

Tested by

no test coverage detected