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

Method test_repeat

axlearn/common/repeat_test.py:164–281  ·  view source on GitHub ↗
(self, dtype, remat_spec, drop_output, num_layers_total, unroll)

Source from the content-addressed store, hash-verified

162 unroll=(True, False, 1, 2),
163 )
164 def test_repeat(self, dtype, remat_spec, drop_output, num_layers_total, unroll):
165 batch_size, num_layers = 14, 4
166 cfg = _Ensemble.default_config().set(name="test", num_layers=num_layers_total, dtype=dtype)
167 cfg.repeat_layer.set(
168 remat_spec=remat_spec,
169 drop_output=drop_output,
170 unroll=unroll,
171 )
172 layer: _Ensemble = cfg.instantiate(parent=None)
173 self.assertEqual(
174 PartitionSpec(None),
175 layer.create_parameter_specs_recursively()["repeat_layer"]["layer"]["inc"].mesh_axes,
176 )
177 layer_params = layer.initialize_parameters_recursively(prng_key=jax.random.PRNGKey(1))
178 logging.info("layer params=%s", layer_params)
179
180 input_forward_state = layer.init_forward_state(batch_size)
181 if num_layers_total == num_layers:
182 method = "forward"
183 inputs = dict(
184 carry=jnp.arange(batch_size, dtype=dtype),
185 forward_state=input_forward_state,
186 )
187 else:
188 method = "forward_first_n"
189 input_forward_state = _get_first_n(num_layers, input_forward_state)
190 inputs = dict(
191 carry=jnp.arange(batch_size, dtype=dtype),
192 forward_state=input_forward_state,
193 n=num_layers,
194 )
195
196 (carry, output_forward_state), output_collection = F(
197 layer,
198 prng_key=jax.random.PRNGKey(2),
199 state=layer_params,
200 inputs=inputs,
201 method=method,
202 is_training=True,
203 drop_output_collections=(),
204 )
205 logging.info("forward_state=%s", output_forward_state)
206 logging.info("output_collection=%s", output_collection)
207 assert_allclose(carry, jnp.arange(num_layers, num_layers + batch_size, dtype=dtype))
208 self.assertEqual(shapes(input_forward_state), shapes(output_forward_state))
209 assert_allclose(
210 output_forward_state["repeat_layer"]["layer"],
211 jnp.reshape(
212 jnp.arange(batch_size)[None, :] + jnp.arange(num_layers, dtype=dtype)[:, None],
213 (num_layers, batch_size),
214 ),
215 )
216 # Check output collection.
217 self.assertEqual(
218 OutputCollection(
219 state_updates={
220 "repeat_layer": {
221 # State update values are stacked across layers.

Callers

nothing calls this directly

Calls 11

assert_allcloseFunction · 0.90
shapesFunction · 0.90
OutputCollectionClass · 0.90
get_recursivelyFunction · 0.90
_get_first_nFunction · 0.85
setMethod · 0.45
default_configMethod · 0.45
instantiateMethod · 0.45
init_forward_stateMethod · 0.45

Tested by

no test coverage detected