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

Method align

axlearn/audio/decoder_asr.py:623–647  ·  view source on GitHub ↗

Given an input_batch that contains both audio and labels, outputs text-audio alignment. Args: input_batch: See `CTCDecoderModel`'s forward interface. `input_batch` should contain: * inputs: A Tensor of shape [batch_size, num_frames, dim]. * padding

(self, input_batch: Nested[Tensor])

Source from the content-addressed store, hash-verified

621 )
622
623 def align(self, input_batch: Nested[Tensor]) -> Nested[Tensor]:
624 """Given an input_batch that contains both audio and labels, outputs text-audio alignment.
625 Args:
626 input_batch: See `CTCDecoderModel`'s forward interface. `input_batch` should contain:
627 * inputs: A Tensor of shape [batch_size, num_frames, dim].
628 * paddings: A 0/1 Tensor of shape [batch_size, num_frames].
629 * target_labels: A Tensor of shape [batch_size, label_length].
630 target_labels < 0 means this is a padding position.
631 Returns:
632 A NestedTensor, converted from `ctc_aligner.AlignmentOutput` object
633 """
634 logits = self.predict(input_batch)
635 log_posterior = jax.nn.log_softmax(logits, axis=-1)
636 log_pos_paddings = cast(Tensor, input_batch["paddings"])
637 labels = cast(Tensor, input_batch["target_labels"])
638 label_paddings = jnp.where(labels >= 0, 0, 1)
639
640 alignment_output = ctc_aligner.ctc_forced_alignment(
641 log_pos=log_posterior,
642 log_pos_paddings=log_pos_paddings,
643 labels=labels,
644 label_paddings=label_paddings,
645 blank_id=self.config.blank_id,
646 )
647 return alignment_output.asdict()
648
649
650def _map_label_sequences(

Callers

nothing calls this directly

Calls 3

predictMethod · 0.95
castFunction · 0.85
asdictMethod · 0.80

Tested by

no test coverage detected