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

Method test_forward

axlearn/audio/decoder_asr_test.py:439–468  ·  view source on GitHub ↗
(self, blank_id)

Source from the content-addressed store, hash-verified

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]

Callers

nothing calls this directly

Calls 4

load_goldenFunction · 0.90
_is_valid_ctc_seqFunction · 0.90
assert_allcloseFunction · 0.90
astypeMethod · 0.80

Tested by

no test coverage detected