diff --git a/.github/workflows/sensevoice-container.yml b/.github/workflows/sensevoice-container.yml index a0b9c68..f4daac0 100644 --- a/.github/workflows/sensevoice-container.yml +++ b/.github/workflows/sensevoice-container.yml @@ -19,6 +19,7 @@ on: - tests/test_docker_context.sh - tests/test_device_env.py - tests/test_model_timestamps.py + - tests/test_ctc_alignment.py - utils/ctc_alignment.py push: branches: @@ -42,6 +43,7 @@ on: - tests/test_docker_context.sh - tests/test_device_env.py - tests/test_model_timestamps.py + - tests/test_ctc_alignment.py - utils/ctc_alignment.py workflow_dispatch: @@ -68,7 +70,7 @@ jobs: - name: Install CPU PyTorch for model contracts run: python -m pip install torch==2.12.1 --index-url https://download.pytorch.org/whl/cpu - name: Validate module and timestamp contracts - run: python -m unittest tests.test_model_timestamps -v + run: python -m unittest tests.test_model_timestamps tests.test_ctc_alignment -v build: needs: contract diff --git a/tests/test_ctc_alignment.py b/tests/test_ctc_alignment.py new file mode 100644 index 0000000..76564c9 --- /dev/null +++ b/tests/test_ctc_alignment.py @@ -0,0 +1,75 @@ +import unittest + +import torch + +from utils.ctc_alignment import ctc_forced_align + + +# Five frames of a two-label utterance. A sixth blank frame sits past the +# length and must not move the alignment. +_FRAMES = [ + [-1.99, -1.09, -2.18, -1.83, -1.36], + [-0.29, -4.06, -4.18, -3.26, -1.72], + [-0.93, -1.8, -2.21, -2.37, -1.44], + [-4.04, -4.09, -2.9, -2.12, -0.23], + [-3.59, -3.29, -0.24, -2.61, -2.63], +] + + +class CtcAlignmentTests(unittest.TestCase): + def test_frames_past_the_length_do_not_move_the_alignment(self): + frames = torch.tensor(_FRAMES) + extra = torch.tensor([[0.0, -8.0, -8.0, -8.0, -8.0]]) + emissions = torch.cat([frames, extra], dim=0).unsqueeze(0) + targets = torch.tensor([[1, 1]]) + aligned = ctc_forced_align(emissions, targets, torch.tensor([5]), torch.tensor([2])) + self.assertEqual(aligned[0, :5].tolist(), [1, 0, 0, 0, 1]) + + def test_full_length_alignment_stays(self): + emissions = torch.tensor(_FRAMES).unsqueeze(0) + aligned = ctc_forced_align(emissions, torch.tensor([[1, 1]]), torch.tensor([5]), torch.tensor([2])) + self.assertEqual(aligned[0].tolist(), [1, 0, 0, 0, 1]) + + def test_ignore_id_does_not_change_the_targets(self): + targets = torch.tensor([[1, -1]]) + emissions = torch.zeros(1, 3, 4) + ctc_forced_align(emissions, targets, torch.tensor([3]), torch.tensor([2])) + self.assertEqual(targets.tolist(), [[1, -1]]) + + def test_an_unequal_length_batch_matches_each_item_cropped(self): + frames = torch.tensor(_FRAMES) + extra = torch.tensor([[0.0, -8.0, -8.0, -8.0, -8.0]]) + padded = torch.cat([frames, extra]) + emissions = torch.stack([padded, padded]) + targets = torch.tensor([[1, 1], [1, 1]]) + input_lengths = torch.tensor([5, 6]) + target_lengths = torch.tensor([2, 2]) + batched = ctc_forced_align(emissions, targets, input_lengths, target_lengths) + for i, length in enumerate(input_lengths.tolist()): + single = ctc_forced_align( + emissions[i : i + 1, :length], targets[i : i + 1], torch.tensor([length]), target_lengths[i : i + 1] + ) + self.assertEqual(batched[i, :length].tolist(), single[0].tolist()) + + def test_lengths_on_the_cpu_work_with_emissions_on_an_accelerator(self): + if torch.cuda.is_available(): + device = torch.device("cuda") + elif torch.backends.mps.is_available(): + device = torch.device("mps") + else: + self.skipTest("needs an accelerator") + torch.manual_seed(0) + emissions = torch.randn(2, 7, 5).log_softmax(-1) + targets = torch.tensor([[1, 2], [2, 1]]) + input_lengths = torch.tensor([7, 5]) + target_lengths = torch.tensor([2, 2]) + expected = ctc_forced_align(emissions, targets, input_lengths, target_lengths) + aligned = ctc_forced_align( + emissions.to(device), targets.to(device), input_lengths, target_lengths.to(device) + ) + self.assertEqual(aligned.cpu()[0].tolist(), expected[0].tolist()) + self.assertEqual(aligned.cpu()[1, :5].tolist(), expected[1, :5].tolist()) + + +if __name__ == "__main__": + unittest.main() diff --git a/utils/ctc_alignment.py b/utils/ctc_alignment.py index f10fcf2..51e86d8 100644 --- a/utils/ctc_alignment.py +++ b/utils/ctc_alignment.py @@ -23,10 +23,13 @@ def ctc_forced_align( blank_id (int, optional): The index of blank symbol in CTC emission. (Default: 0) ignore_id (int, optional): The index of ignore symbol in CTC emission. (Default: -1) """ + targets = targets.clone() targets[targets == ignore_id] = blank batch_size, input_time_size, _ = log_probs.size() bsz_indices = torch.arange(batch_size, device=input_lengths.device) + # The frame masks are combined with score tensors, so keep them on that device. + mask_lengths = input_lengths.to(log_probs.device) _t_a_r_g_e_t_s_ = torch.cat( ( @@ -53,12 +56,19 @@ def ctc_forced_align( backpointers = torch.zeros((batch_size, input_time_size, padded_t), device=log_probs.device, dtype=targets.dtype) for t in range(1, input_time_size): + # Frames past input_lengths must not move that item's score. + active = t < mask_lengths + if active.ndim == 0: + active = active.view(1) prev = torch.stack( (best_score[:, 2:], best_score[:, 1:-1], torch.where(diff_labels, best_score[:, :-2], neg_inf)) ) prev_max_value, prev_max_idx = prev.max(dim=0) - best_score[:, padding_num:] = log_probs[:, t].gather(-1, _t_a_r_g_e_t_s_) + prev_max_value - backpointers[:, t, padding_num:] = prev_max_idx + updated = log_probs[:, t].gather(-1, _t_a_r_g_e_t_s_) + prev_max_value + best_score[:, padding_num:] = torch.where(active.unsqueeze(-1), updated, best_score[:, padding_num:]) + backpointers[:, t, padding_num:] = torch.where( + active.unsqueeze(-1), prev_max_idx, backpointers[:, t, padding_num:] + ) l1l2 = best_score.gather( -1, torch.stack((padding_num + target_lengths * 2 - 1, padding_num + target_lengths * 2), dim=-1) @@ -68,9 +78,13 @@ def ctc_forced_align( path[bsz_indices, input_lengths - 1] = padding_num + target_lengths * 2 - 1 + l1l2.argmax(dim=-1) for t in range(input_time_size - 1, 0, -1): + active = t < mask_lengths + if active.ndim == 0: + active = active.view(1) target_indices = path[:, t] prev_max_idx = backpointers[bsz_indices, t, target_indices] - path[:, t - 1] += target_indices - prev_max_idx + step = target_indices - prev_max_idx + path[:, t - 1] = torch.where(active, step, path[:, t - 1]) alignments = _t_a_r_g_e_t_s_.gather(dim=-1, index=(path - padding_num).clamp(min=0)) return alignments