import pytest import torch import kernels dpx = kernels.get_kernel("phanerozoic/dpx-decode", version=1, trust_remote_code=True) requires_cuda = pytest.mark.skipif(not torch.cuda.is_available(), reason="CUDA required") def ref_viterbi_int(emit, trans, prior): """Fixed-point reference replicating the kernel's packing, clamping, and renormalization exactly; bit-exact agreement is required.""" B, T, S = emit.shape paths = torch.empty(B, T, dtype=torch.int32) scores = torch.empty(B, dtype=torch.int64) for b in range(B): prev = (prior + emit[b, 0]).to(torch.int64) norm = 0 bps = torch.empty(T, S, dtype=torch.int64) for t in range(1, T): mx = prev.max() norm += int(mx) prev = prev - mx cand = torch.clamp(prev[:, None] + trans.to(torch.int64), min=-32768) packed = (cand << 16) | torch.arange(S)[:, None] best = packed.max(dim=0).values # torch max returns first argmax; the kernel's packed max selects # the largest index among score ties, replicated here: arg = (packed == best[None, :]).to(torch.int64).cumsum(0).argmax(0) bps[t] = arg prev = (best >> 16) + emit[b, t].to(torch.int64) j = int(prev.argmax()) # kernel picks the first maximal final state j = int((prev == prev.max()).nonzero()[0]) scores[b] = int(prev[j]) + norm paths[b, T - 1] = j for t in range(T - 1, 0, -1): j = int(bps[t, j]) paths[b, t - 1] = j return paths, scores @requires_cuda @pytest.mark.kernels_ci def test_viterbi_int_bit_exact(): torch.manual_seed(0) B, T, S = 3, 40, 33 emit = torch.randint(-5000, 0, (B, T, S), dtype=torch.int32) trans = torch.randint(-3000, 0, (S, S), dtype=torch.int32) prior = torch.randint(-1000, 0, (S,), dtype=torch.int32) path, score = dpx.viterbi(emit.cuda(), trans.cuda(), prior.cuda()) rp, rs = ref_viterbi_int(emit, trans, prior) assert torch.equal(path.cpu(), rp) assert torch.equal(score.cpu(), rs) @requires_cuda @pytest.mark.kernels_ci def test_viterbi_float_matches_reference(): torch.manual_seed(1) B, T, S = 2, 50, 24 emit = torch.randn(B, T, S, device="cuda") trans = torch.randn(S, S, device="cuda") path, score = dpx.viterbi(emit, trans) # float64 reference e, tr = emit.double().cpu(), trans.double().cpu() for b in range(B): prev = e[b, 0].clone() bps = torch.zeros(T, S, dtype=torch.long) for t in range(1, T): cand = prev[:, None] + tr best, arg = cand.max(dim=0) bps[t] = arg prev = best + e[b, t] j = int(prev.argmax()) ref_path = torch.empty(T, dtype=torch.int32) ref_path[T - 1] = j for t in range(T - 1, 0, -1): j = int(bps[t, j]) ref_path[t - 1] = j assert torch.equal(path[b].cpu(), ref_path) @requires_cuda @pytest.mark.kernels_ci def test_dtw_matches_reference(): torch.manual_seed(2) B, N, M = 2, 30, 45 cost = torch.rand(B, N, M, device="cuda") path, plen, D = dpx.dtw(cost) c = cost.double().cpu() for b in range(B): ref = torch.full((N, M), float("inf"), dtype=torch.float64) for i in range(N): for j in range(M): if i == 0 and j == 0: m = 0.0 else: up = ref[i - 1, j] if i > 0 else float("inf") left = ref[i, j - 1] if j > 0 else float("inf") ul = ref[i - 1, j - 1] if i > 0 and j > 0 else float("inf") m = min(up, left, ul) ref[i, j] = c[b, i, j] + m assert abs(float(D[b, N - 1, M - 1]) - float(ref[N - 1, M - 1])) < 1e-4 n = int(plen[b]) pts = path[b, :n].cpu() assert tuple(pts[0].tolist()) == (0, 0) and tuple(pts[-1].tolist()) == (N - 1, M - 1) steps = pts[1:] - pts[:-1] assert bool(((steps >= 0) & (steps <= 1)).all()) and bool((steps.sum(1) >= 1).all()) @requires_cuda @pytest.mark.kernels_ci def test_ctc_align_properties_and_score(): torch.manual_seed(3) B, T, C, L = 2, 60, 20, 8 log_probs = torch.log_softmax(torch.randn(B, T, C, device="cuda"), dim=-1) targets = torch.randint(1, C, (B, L), dtype=torch.int64, device="cuda") frames, score = dpx.ctc_forced_align(log_probs, targets, blank=0) for b in range(B): seq = frames[b].cpu().tolist() collapsed = [] prev = None for x in seq: if x != 0 and x != prev: collapsed.append(x) prev = x assert collapsed == targets[b].cpu().tolist(), "collapse(frames) must equal the transcript" # the reported score equals the sum of per-frame log-probs on the path s = sum(float(log_probs[b, t, seq[t]]) for t in range(T)) assert abs(s - float(score[b])) < 1e-3 @requires_cuda @pytest.mark.kernels_ci def test_deterministic(): torch.manual_seed(4) emit = torch.randn(2, 80, 64, device="cuda") trans = torch.randn(64, 64, device="cuda") p1, s1 = dpx.viterbi(emit, trans) p2, s2 = dpx.viterbi(emit, trans) assert torch.equal(p1, p2) and torch.equal(s1, s2)