Tests the `Functional` class and `InvocationContext.functional()`.
(self)
| 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 |
nothing calls this directly
no test coverage detected