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

Method test_decode

axlearn/common/encoder_decoder_test.py:174–273  ·  view source on GitHub ↗

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

Source from the content-addressed store, hash-verified

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: (

Callers

nothing calls this directly

Calls 6

RematSpecClass · 0.90
dummy_padding_maskFunction · 0.90
_model_configFunction · 0.70
instantiateMethod · 0.45
setMethod · 0.45

Tested by

no test coverage detected