(self, prefill_states: bool, prefix_zero: bool)
| 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 |
nothing calls this directly
no test coverage detected