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

Method test_decode

axlearn/audio/decoder_asr_test.py:1619–1695  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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,

Callers

nothing calls this directly

Calls 5

_set_up_lasMethod · 0.95
sequence_maskFunction · 0.90
instantiateMethod · 0.45
setMethod · 0.45

Tested by

no test coverage detected