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

Function bind_layer

axlearn/common/test_utils.py:831–871  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

829
830@contextlib.contextmanager
831def 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
874def initialize_parameters_with_prebuilt(

Callers

nothing calls this directly

Calls 3

current_contextFunction · 0.90
bind_moduleFunction · 0.85

Tested by

no test coverage detected