Test beam search and sample decoding from a randomly initialized model.
(
self,
stack_cfg: BaseStackedTransformerLayer.Config,
num_decodes: int,
prefix_length: utils.Tensor,
method: Literal["sample_decode", "beam_search_decode"],
pad_token_id: int,
)
| 172 | ) |
| 173 | # pylint: disable-next=too-many-statements |
| 174 | def test_decode( |
| 175 | self, |
| 176 | stack_cfg: BaseStackedTransformerLayer.Config, |
| 177 | num_decodes: int, |
| 178 | prefix_length: utils.Tensor, |
| 179 | method: Literal["sample_decode", "beam_search_decode"], |
| 180 | pad_token_id: int, |
| 181 | ): |
| 182 | """Test beam search and sample decoding from a randomly initialized model.""" |
| 183 | with jax.checking_leaks(): |
| 184 | batch_size, src_len, tgt_len, vocab_size = prefix_length.shape[0], 11, 10, 6 |
| 185 | bos_id = 1 |
| 186 | init_key, prefix_key, source_key, method_key = jax.random.split( |
| 187 | jax.random.PRNGKey(0), num=4 |
| 188 | ) |
| 189 | |
| 190 | if isinstance(stack_cfg, RepeatedTransformerLayer.Config): |
| 191 | remat_spec = RematSpec(prevent_cse=False) |
| 192 | else: |
| 193 | remat_spec = None |
| 194 | |
| 195 | cfg = _model_config( |
| 196 | vocab_size=vocab_size, source_len=src_len, target_len=tgt_len, remat_spec=remat_spec |
| 197 | ) |
| 198 | cfg.encoder.pad_token_id = pad_token_id |
| 199 | cfg.decoder.pad_token_id = pad_token_id |
| 200 | model = cfg.set(name="test").instantiate(parent=None) |
| 201 | params = model.initialize_parameters_recursively(init_key) |
| 202 | |
| 203 | prefix = jax.random.randint( |
| 204 | prefix_key, |
| 205 | shape=[batch_size, tgt_len], |
| 206 | # Prefix can consist of any tokens, including pad and eos. |
| 207 | minval=0, |
| 208 | maxval=vocab_size, |
| 209 | ) |
| 210 | # Explicitly fill positions >= prefix_length with pad_token_id. |
| 211 | # Note that each batch example may have a different prefix length. |
| 212 | # [batch_size, tgt_len]. |
| 213 | prefix_mask = utils.sequence_mask(lengths=prefix_length, max_len=tgt_len) |
| 214 | prefix = prefix * prefix_mask + pad_token_id * (1 - prefix_mask) |
| 215 | # Set last token to a non-pad token, to fix the prefix length. |
| 216 | oh_indices = jax.nn.one_hot(prefix_length - 1, tgt_len, dtype=prefix.dtype) |
| 217 | prefix = prefix * (1 - oh_indices) + bos_id * oh_indices |
| 218 | |
| 219 | source_ids = jax.random.randint( |
| 220 | source_key, minval=1, maxval=vocab_size, shape=(batch_size, src_len) |
| 221 | ) |
| 222 | source_mask = dummy_padding_mask(batch_size=batch_size, max_seq_len=src_len) |
| 223 | source_ids = source_ids * source_mask + pad_token_id * (1 - source_mask) |
| 224 | inputs = dict( |
| 225 | input_batch=dict(prefix=prefix, source=dict(input_ids=source_ids)), |
| 226 | max_sequence_length=tgt_len, |
| 227 | num_decodes=num_decodes, |
| 228 | ) |
| 229 | if method == "sample_decode": |
| 230 | # Modify logits so that we will always sample the last token ID. |
| 231 | inputs["logits_modifier"] = lambda logits: ( |
nothing calls this directly
no test coverage detected