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

Method test_forward

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

Source from the content-addressed store, hash-verified

1554 return cfg
1555
1556 def test_forward(self):
1557 encoder_dim, decoder_dim, num_heads, vocab_size = 8, 6, 3, 20
1558 cfg: LASDecoderModel.Config = self._set_up_las(
1559 encoder_dim=encoder_dim,
1560 decoder_dim=decoder_dim,
1561 num_heads=num_heads,
1562 vocab_size=vocab_size,
1563 )
1564 bos_id = eos_id = cfg.decoder.eos_token_id
1565 pad_id = cfg.decoder.pad_token_id
1566 # Initialize layer parameters.
1567 layer: LASDecoderModel = cfg.set(name="test").instantiate(parent=None)
1568 prng_key = jax.random.PRNGKey(123)
1569 prng_key, init_key, input_key = jax.random.split(prng_key, num=3)
1570 layer_params = layer.initialize_parameters_recursively(init_key)
1571
1572 batch_size, max_src_len = 3, 10
1573 target_labels = jnp.array(
1574 [
1575 [14, 8, 17, 19, 17, eos_id], # length 5.
1576 [17, 4, 18, eos_id, pad_id, pad_id], # length 3.
1577 [eos_id, pad_id, pad_id, pad_id, pad_id, pad_id], # length 0.
1578 ]
1579 )
1580 input_ids = jnp.concatenate(
1581 [jnp.full([batch_size, 1], bos_id), target_labels[:, :-1]], axis=1
1582 )
1583
1584 src_len = np.array([10, 0, 7])
1585 target_len = np.array([6, 4, 1])
1586 # [batch_size, src_len, am_dim].
1587 inputs = jax.random.normal(input_key, [batch_size, max_src_len, encoder_dim]) * 1000
1588 paddings = jnp.arange(max_src_len)[None, :] >= src_len[:, None]
1589
1590 @jax.jit
1591 def jit_forward(input_batch):
1592 (loss, aux_outputs), _ = F(
1593 layer,
1594 inputs=dict(input_batch=input_batch),
1595 is_training=True,
1596 prng_key=prng_key,
1597 state=layer_params,
1598 )
1599 return loss, aux_outputs
1600
1601 # Compute test loss.
1602 loss, aux_outputs = jit_forward(
1603 dict(
1604 inputs=inputs,
1605 paddings=paddings,
1606 target_labels=target_labels,
1607 target=dict(input_ids=input_ids),
1608 )
1609 )
1610 # Empty source example has weight 0.
1611 expected_weight = target_len * (src_len > 0)
1612 self.assertNestedAllClose(aux_outputs["per_example_weight"], expected_weight)
1613 assert_allclose(

Callers

nothing calls this directly

Calls 6

_set_up_lasMethod · 0.95
assert_allcloseFunction · 0.90
assertNestedAllCloseMethod · 0.80
instantiateMethod · 0.45
setMethod · 0.45

Tested by

no test coverage detected