(self, blank_id)
| 437 | |
| 438 | @parameterized.parameters([0, 1]) |
| 439 | def test_forward(self, blank_id): |
| 440 | vocab_size = 20 |
| 441 | golden = load_golden("axlearn.audio.decoder_asr_test", f"test_forward_blank{blank_id}") |
| 442 | # Load golden inputs. |
| 443 | logits = jnp.asarray(golden["inputs"]["logits"]) |
| 444 | input_lengths = jnp.asarray(golden["inputs"]["input_lengths"]) |
| 445 | target_labels = jnp.asarray(golden["inputs"]["target_labels"]) |
| 446 | max_seq_len = logits.shape[1] |
| 447 | |
| 448 | # Compute paddings and target_paddings from lengths. |
| 449 | paddings = (jnp.arange(max_seq_len) >= input_lengths[:, None]).astype(jnp.float32) |
| 450 | target_paddings = jnp.logical_or(vocab_size <= target_labels, target_labels < 0) |
| 451 | |
| 452 | # Compute CTC loss using optax (same as CTCDecoderModel.forward). |
| 453 | per_example_loss = optax.ctc_loss( |
| 454 | logits=logits, |
| 455 | logit_paddings=paddings, |
| 456 | labels=target_labels, |
| 457 | label_paddings=target_paddings, |
| 458 | blank_id=blank_id, |
| 459 | ) |
| 460 | |
| 461 | # Mask invalid sequences (optax returns large values; torch zero_infinity returns 0). |
| 462 | per_example_weight = _is_valid_ctc_seq( |
| 463 | paddings=paddings, target_labels=target_labels, target_paddings=target_paddings |
| 464 | ).astype(jnp.float32) |
| 465 | |
| 466 | # Compare against golden reference (torch CTC loss, weighted by validity). |
| 467 | ref_per_example_loss = jnp.asarray(golden["outputs"]["per_example_loss"]) |
| 468 | assert_allclose(ref_per_example_loss, per_example_loss * per_example_weight) |
| 469 | |
| 470 | def _check_summary( |
| 471 | self, summary_collection: dict[str, Any], name: str, value: Union[Tensor, WeightedSummary] |
nothing calls this directly
no test coverage detected