Creates a context in which `layer` has state initialized using `initialize_parameters_recursively`. The only difference between this and `bind_module()` is this calls `initialize_parameters_recursively`. Example: ``` cfg = Linear.default_config().set(input_dim=5, ou
(
layer: ConfigOr[L],
*,
is_training: bool = True,
prng_key: Optional[jax.random.PRNGKey] = None,
state: Optional[Nested[Tensor]] = None,
)
| 829 | |
| 830 | @contextlib.contextmanager |
| 831 | def bind_layer( |
| 832 | layer: ConfigOr[L], |
| 833 | *, |
| 834 | is_training: bool = True, |
| 835 | prng_key: Optional[jax.random.PRNGKey] = None, |
| 836 | state: Optional[Nested[Tensor]] = None, |
| 837 | ) -> Iterator[L]: |
| 838 | """Creates a context in which `layer` has state initialized using |
| 839 | `initialize_parameters_recursively`. |
| 840 | |
| 841 | The only difference between this and `bind_module()` is this calls |
| 842 | `initialize_parameters_recursively`. |
| 843 | |
| 844 | Example: |
| 845 | ``` |
| 846 | cfg = Linear.default_config().set(input_dim=5, output_dim=7) |
| 847 | with test_utils.bind_layer(cfg) as layer: |
| 848 | result = layer(jnp.ones(5)) |
| 849 | assert result.shape == (7,) |
| 850 | ``` |
| 851 | |
| 852 | Args: |
| 853 | layer: The layer to initialize. |
| 854 | is_training: Tell the layer it is in training or not. |
| 855 | prng_key: The PRNG key to use. If None, `jax.random.PRNGKey(0)`. |
| 856 | state: The state to use. If None, call `initialize_parameters_recursively()` to initialize |
| 857 | the state. |
| 858 | |
| 859 | Returns: |
| 860 | The Initialized module. |
| 861 | """ |
| 862 | if prng_key is None: |
| 863 | prng_key = jax.random.PRNGKey(0) |
| 864 | |
| 865 | init_key, ctx_key = jax.random.split(prng_key) |
| 866 | |
| 867 | with bind_module(layer, is_training=is_training, prng_key=ctx_key, state={}) as instance: |
| 868 | if state is None: |
| 869 | state = instance.initialize_parameters_recursively(prng_key=init_key) |
| 870 | current_context().state = state |
| 871 | yield instance |
| 872 | |
| 873 | |
| 874 | def initialize_parameters_with_prebuilt( |
nothing calls this directly
no test coverage detected