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

Method test_extend_step

axlearn/common/encoder_test.py:171–281  ·  view source on GitHub ↗
(self, prefill_states: bool, prefix_zero: bool)

Source from the content-addressed store, hash-verified

169 prefix_zero=[True, False],
170 )
171 def test_extend_step(self, prefill_states: bool, prefix_zero: bool):
172 hidden_dim = 12
173 vocab_size = 24
174 num_heads = 4
175 num_layers = 2
176 source_length = 11
177
178 encoder = CausalEncoder.default_config().set(
179 dim=hidden_dim,
180 vocab_size=vocab_size,
181 dropout_rate=0,
182 attention_mask=CausalAttentionLogitBiasLayer.default_config(),
183 emb=bert_embedding_config(type_vocab_size=1, max_position_embeddings=source_length),
184 transformer=bert_transformer_config(num_layers=num_layers, num_heads=num_heads),
185 param_init=DefaultInitializer.default_config().set(
186 init_by_param_name={
187 PARAM_REGEXP_WEIGHT: WeightInitializer.default_config().set(
188 fan=None, scale=0.02, distribution="normal"
189 )
190 }
191 ),
192 pad_token_id=0,
193 )
194 set_layer_norm_eps_recursively(encoder, 1e-5)
195
196 layer = encoder.set(name="layer_test").instantiate(parent=None)
197 batch_size = 3
198
199 # We ignore padding ids (0) for now to simplify the mask generation process.
200 if prefix_zero:
201 prefix = jnp.zeros([batch_size, 1], dtype=jnp.int32)
202 else:
203 prefix = jax.random.randint(
204 jax.random.PRNGKey(123), [batch_size, 1], minval=1, maxval=vocab_size - 1
205 )
206 input_ids = jax.random.randint(
207 jax.random.PRNGKey(123),
208 [batch_size, source_length - 1],
209 minval=1,
210 maxval=vocab_size - 1,
211 )
212 input_ids = jnp.hstack([prefix, input_ids])
213
214 params = layer.initialize_parameters_recursively(prng_key=jax.random.PRNGKey(123))
215
216 ref_hidden_states, _ = F(
217 layer,
218 is_training=False,
219 prng_key=jax.random.PRNGKey(123),
220 state=params,
221 inputs=dict(
222 input_ids=input_ids,
223 input_segment_ids=input_ids != 0,
224 positions=jnp.arange(input_ids.shape[-1])[None, :],
225 ),
226 )
227 ref_hidden_states = ref_hidden_states["hidden_states"]
228

Callers

nothing calls this directly

Calls 9

bert_embedding_configFunction · 0.90
bert_transformer_configFunction · 0.90
assert_allcloseFunction · 0.90
setMethod · 0.45
default_configMethod · 0.45
instantiateMethod · 0.45
init_statesMethod · 0.45

Tested by

no test coverage detected