| 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 |
| |
| |
| 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()) |
| |
| 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) |
| |
| 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" |
| |
| 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) |
|
|