Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 3 additions & 1 deletion .github/workflows/sensevoice-container.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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:

Expand All @@ -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
Expand Down
75 changes: 75 additions & 0 deletions tests/test_ctc_alignment.py
Original file line number Diff line number Diff line change
@@ -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()
20 changes: 17 additions & 3 deletions utils/ctc_alignment.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
(
Expand All @@ -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:])
Comment thread
LauraGPT marked this conversation as resolved.
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)
Expand All @@ -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