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])
| 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 | |
| 650 | def _map_label_sequences( |