(self)
| 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( |
nothing calls this directly
no test coverage detected