(self)
| 1617 | self.assertGreater(loss, 0.0) |
| 1618 | |
| 1619 | def test_decode(self): |
| 1620 | encoder_dim, decoder_dim, num_heads, vocab_size = 5, 16, 4, 20 |
| 1621 | num_decodes = 5 |
| 1622 | cfg: LASDecoderModel.Config = self._set_up_las( |
| 1623 | encoder_dim=encoder_dim, |
| 1624 | decoder_dim=decoder_dim, |
| 1625 | num_heads=num_heads, |
| 1626 | vocab_size=vocab_size, |
| 1627 | ) |
| 1628 | pad_id = cfg.decoder.pad_token_id |
| 1629 | # Initialize layer parameters. |
| 1630 | layer: LASDecoderModel = cfg.set(name="test").instantiate(parent=None) |
| 1631 | prng_key = jax.random.PRNGKey(123) |
| 1632 | decode_key, init_key, input_key, prefix_key = jax.random.split(prng_key, num=4) |
| 1633 | layer_params = layer.initialize_parameters_recursively(init_key) |
| 1634 | |
| 1635 | batch_size, max_src_len, max_tgt_len = 3, 10, 8 |
| 1636 | src_len = jnp.array([10, 7, 5]) |
| 1637 | prefix_length = jnp.array([1, 3, 6]) |
| 1638 | |
| 1639 | # [batch_size, max_seq_len, dim]. |
| 1640 | inputs = jax.random.normal(input_key, [batch_size, max_src_len, cfg.input_dim]) * 1000 |
| 1641 | # [batch_size, max_seq_len]. |
| 1642 | paddings = jnp.arange(max_src_len) >= src_len[:, None] |
| 1643 | prefix = jax.random.randint( |
| 1644 | prefix_key, |
| 1645 | shape=[batch_size, max_tgt_len], |
| 1646 | # Prefix can consist of any tokens, including pad and eos. |
| 1647 | minval=0, |
| 1648 | maxval=vocab_size, |
| 1649 | ) |
| 1650 | # Explicitly fill positions >= prefix_length with pad_id (-1). |
| 1651 | # Note that each batch example may have a different prefix length. |
| 1652 | # [batch_size, max_tgt_len]. |
| 1653 | prefix_mask = sequence_mask(lengths=prefix_length, max_len=max_tgt_len) |
| 1654 | prefix = prefix * prefix_mask + pad_id * (1 - prefix_mask) |
| 1655 | |
| 1656 | @functools.partial(jax.jit, static_argnames=("method", "num_decodes", "logits_modifier")) |
| 1657 | def jit_method(inputs, prng_key, method, num_decodes, logits_modifier=None): |
| 1658 | if logits_modifier is not None: |
| 1659 | inputs["logits_modifier"] = logits_modifier |
| 1660 | outputs, _ = F( |
| 1661 | layer, |
| 1662 | inputs=dict(**inputs, max_decode_len=max_tgt_len, num_decodes=num_decodes), |
| 1663 | is_training=True, |
| 1664 | prng_key=prng_key, |
| 1665 | state=layer_params, |
| 1666 | method=method, |
| 1667 | ) |
| 1668 | return outputs |
| 1669 | |
| 1670 | decode_inputs = dict( |
| 1671 | input_batch=dict(inputs=inputs, paddings=paddings, prefix=prefix), |
| 1672 | ) |
| 1673 | |
| 1674 | # Beam search decode. |
| 1675 | beam_search_outputs: DecodeOutputs = jit_method( |
| 1676 | decode_inputs, |
nothing calls this directly
no test coverage detected