Source code for ddtw.ddtw_variants

##################################################################################
#                                dDTW Toolbox                                    #
##################################################################################
#                                                                                #
# Authors: Johannes Zeitler and Meinard Müller, 2026                             #
#                                                                                #
# If you use this toolbox, please cite the accompanying paper:                   #
# Johannes Zeitler and Meinard Müller. dDTW: A Unified and Efficient Toolbox for #
#  Differentiable Sequence Alignment. Submitted 2026.                            #
##################################################################################


##################################################################################
# MIT License                                                                    #
#                                                                                #
# Copyright 2026 Johannes Zeitler and Meinard Müller                             #
#                                                                                #
# Permission is hereby granted, free of charge, to any person obtaining a copy   #
# of this software and associated documentation files (the "Software"), to deal  #
# in the Software without restriction, including without limitation the rights   #
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell      #
# copies of the Software, and to permit persons to whom the Software is          #
# furnished to do so, subject to the following conditions:                       #
#                                                                                #
# The above copyright notice and this permission notice shall be included in all #
# copies or substantial portions of the Software.                                #
#                                                                                #
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR     #
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,       #
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE    #
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER         #
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,  #
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE  #
# SOFTWARE.                                                                      #
##################################################################################

import torch
from .ddtw import dDTW
from .backend.backend_cuda_cpp import compute_ctc_initialization, compute_subseq_initialization
from itertools import product

####################### classical softDTW (SDTW) #################################
[docs] class SDTW(dDTW): def __init__(self, cost_function = "MSE", gamma = 1.0, step_sizes = [[1,0], [0,1], [1,1]], global_step_weights = [1.0, 1.0, 1.0], normalization = "N", backend="auto", dtype_float = torch.float32, cuda_device=None, store_debug=False ): """Initialize SDTW loss function, see [1, 2]. [1] Marco Cuturi and Mathieu Blondel. Soft-DTW: A Differentiable Loss Function for Time-Series. In Proceedings of the International Conference on Neural Information Processing Systems (NIPS), vol. 2, pages 2292-2300, 2013. [2] Johannes Zeitler and Meinard Müller. A Unified Perspective on CTC and Soft-DTW Using Differentiable DTW. IEEE Transactions on Audio, Speech and Language Processing, vol. 34, pages 936-951, 2026. Parameters ---------- cost_function : str Local cost function for pair-wise comparison of sequence elements. Choose among ("MSE", "BCE", "CTC"). Default: "MSE". gamma : float Softmin temperature hyperparameter. Default: 1.0 step_sizes : list Alignment step sizes in [n,m] direction, given as list of tuples. Default: [[1,0], [0,1], [1,1]] global_step_weights: list Step weights associated to the step sizes. Default: [1.0, 1.0, 1.0] normalization : str Normalization of SDTW cost. Choose among ("M", "N", "NM", "ctc", "none"). "N": divide by N. "M": divide by M. "NM": divide by (N*M). "ctc": divide by (M-1)/2. "none": no normalization. Default: "N" backend : str Backend to use. Choose among ("auto", "torch", "cpu_numba", "cuda_cpp"). Default: "auto" dtype_float: torch.dtype Number format for internal computations. Default: torch.float32 cuda_device : str or torch.device, optional CUDA device to use, for example ``"cuda:0"``. Default: ``None``. store_debug : bool Whether to retain intermediate backend matrices for inspection. Default: False. """ super().__init__(cost_function=cost_function, min_function="softmin", gamma=gamma, step_sizes=step_sizes, global_step_weights=global_step_weights, normalization=normalization, backend=backend, dtype_float=dtype_float, cuda_device=cuda_device, store_debug=store_debug)
[docs] def forward(self, X=None, Y=None, C=None, list_N=None, list_M=None): """ Compute the SDTW loss. Parameters ---------- X : torch.tensor [shape=(B, N, D)] Input sequence, usually the DNN predictions. Y : torch.tensor [shape=(B, M, D)] Input sequence, usually the weak targets. C : torch.tensor [shape=(B, N, M)] Pre-computed local cost matrix C. list_N : torch.tensor [shape=(B)] Sequence lengths of X (<= N) of the individual batch elements. If None, defaults to [N, N, ..., N] list_M : torch.tensor [shape=(B)] Sequence lengths of Y (<= M) of the individual batch elements. If None, defaults to [M, M, ..., M] Returns ------- torch.Tensor Scalar batch-mean SDTW loss. """ return super().forward(X=X, Y=Y, C=C, list_N=list_N, list_M=list_M)
################################################################################### ############################ hard DTW #############################################
[docs] class DTW(dDTW): def __init__(self, cost_function = "MSE", step_sizes = [[1,0], [0,1], [1,1]], global_step_weights = [1.0, 1.0, 1.0], normalization = "N", backend="auto", dtype_float = torch.float32, cuda_device=None, store_debug=False ): """ Initialize DTW loss function, see [3]. [3]. Meinard Müller. Fundamentals of Music Processing - Using Python and Jupyter Notebooks. Springer Verlag, 2nd edition, 2021. Parameters ---------- cost_function : str Local cost function for pair-wise comparison of sequence elements. Choose among ("MSE", "BCE", "CTC"). Default: "MSE". step_sizes : list Alignment step sizes in [n,m] direction, given as list of tuples. Default: [[1,0], [0,1], [1,1]] global_step_weights: list Step weights associated to the step sizes. Default: [1.0, 1.0, 1.0] normalization : str Normalization of SDTW cost. Choose among ("M", "N", "NM", "ctc", "none"). "N": divide by N. "M": divide by M. "NM": divide by (N*M). "ctc": divide by (M-1)/2. "none": no normalization. Default: "N" backend : str Backend to use. Choose among ("auto", "torch", "cpu_numba", "cuda_cpp"). Default: "auto" dtype_float: torch.dtype Number format for internal computations. Default: torch.float32 cuda_device : str or torch.device, optional CUDA device to use, for example ``"cuda:0"``. Default: ``None``. store_debug : bool Whether to retain intermediate backend matrices for inspection. Default: False. """ super().__init__(cost_function=cost_function, min_function="hardmin", step_sizes=step_sizes, global_step_weights=global_step_weights, normalization=normalization, backend=backend, dtype_float=dtype_float, cuda_device=cuda_device, store_debug=store_debug)
[docs] def forward(self, X=None, Y=None, C=None, list_N=None, list_M=None): """ Compute the DTW loss. Parameters ---------- X : torch.tensor [shape=(B, N, D)] Input sequence, usually the DNN predictions. Y : torch.tensor [shape=(B, M, D)] Input sequence, usually the weak targets. C : torch.tensor [shape=(B, N, M)] Pre-computed local cost matrix C. list_N : torch.tensor [shape=(B)] Sequence lengths of X (<= N) of the individual batch elements. If None, defaults to [N, N, ..., N] list_M : torch.tensor [shape=(B)] Sequence lengths of Y (<= M) of the individual batch elements. If None, defaults to [M, M, ..., M] Returns ------- torch.Tensor Scalar batch-mean DTW loss. With a precomputed ``C``, gradients with respect to ``C`` mark the selected hard warping path. """ return super().forward(X=X, Y=Y, C=C, list_N=list_N, list_M=list_M)
################################################################################### ################################# smooth DTW ######################################
[docs] class smoothDTW(dDTW): def __init__(self, cost_function = "MSE", gamma = 1.0, step_sizes = [[1,0], [0,1], [1,1]], global_step_weights = [1.0, 1.0, 1.0], normalization = "N", backend="auto", dtype_float = torch.float32, cuda_device=None, store_debug=False ): """ Initialize smoothDTW loss function, see [4]. [4]. Isma Hadji, K. Derpanis, and A. Jepson. Representation learning via global temporal alignment and cycle-consistency. In IEE/CVF Conference on Computer Vision and Pattern Recognition (CVPR), pages 11068-11077, 2021. Parameters ---------- cost_function : str Local cost function for pair-wise comparison of sequence elements. Choose among ("MSE", "BCE", "CTC"). Default: "MSE". gamma : float Softmin temperature hyperparameter. Default: 1.0 step_sizes : list Alignment step sizes in [n,m] direction, given as list of tuples. Default: [[1,0], [0,1], [1,1]] global_step_weights: list Step weights associated to the step sizes. Default: [1.0, 1.0, 1.0] normalization : str Normalization of SDTW cost. Choose among ("M", "N", "NM", "ctc", "none"). "N": divide by N. "M": divide by M. "NM": divide by (N*M). "ctc": divide by (M-1)/2. "none": no normalization. Default: "N" backend : str Backend to use. Choose among ("auto", "torch", "cpu_numba", "cuda_cpp"). Default: "auto" dtype_float: torch.dtype Number format for internal computations. Default: torch.float32 cuda_device : str or torch.device, optional CUDA device to use, for example ``"cuda:0"``. Default: ``None``. store_debug : bool Whether to retain intermediate backend matrices for inspection. Default: False. """ super().__init__(cost_function=cost_function, min_function="smoothmin", gamma=gamma, step_sizes=step_sizes, global_step_weights=global_step_weights, normalization=normalization, backend=backend, dtype_float=dtype_float, cuda_device=cuda_device, store_debug=store_debug)
[docs] def forward(self, X=None, Y=None, C=None, list_N=None, list_M=None): """ Compute the smoothDTW loss. Parameters ---------- X : torch.tensor [shape=(B, N, D)] Input sequence, usually the DNN predictions. Y : torch.tensor [shape=(B, M, D)] Input sequence, usually the weak targets. C : torch.tensor [shape=(B, N, M)] Pre-computed local cost matrix C. list_N : torch.tensor [shape=(B)] Sequence lengths of X (<= N) of the individual batch elements. If None, defaults to [N, N, ..., N] list_M : torch.tensor [shape=(B)] Sequence lengths of Y (<= M) of the individual batch elements. If None, defaults to [M, M, ..., M] Returns ------- torch.Tensor Scalar batch-mean smoothDTW loss. """ return super().forward(X=X, Y=Y, C=C, list_N=list_N, list_M=list_M)
################################################################################### ########################## sparse DTW (sparseDTW) #################################
[docs] class sparseDTW(dDTW): def __init__(self, cost_function = "MSE", gamma = 1.0, step_sizes = [[1,0], [0,1], [1,1]], global_step_weights = [1.0, 1.0, 1.0], normalization = "N", backend="auto", dtype_float = torch.float32, cuda_device=None, store_debug=False ): """ Initialize sparseDTW loss function, see [5]. [5] Arthur Mensch and Mathieu Blondel. Differentiable Dynamic Programming for Structured Prediction and Attention. In Proceedings of the International Converence on Machine Learning (ICML), pages 3459-3468, Stockholm, Sweden, 2018. Parameters ---------- cost_function : str Local cost function for pair-wise comparison of sequence elements. Choose among ("MSE", "BCE", "CTC"). Default: "MSE". gamma : float Sparsemin temperature hyperparameter. Default: 1.0 step_sizes : list Alignment step sizes in [n,m] direction, given as list of tuples. Default: [[1,0], [0,1], [1,1]] global_step_weights: list Step weights associated to the step sizes. Default: [1.0, 1.0, 1.0] normalization : str Normalization of SDTW cost. Choose among ("M", "N", "NM", "ctc", "none"). "N": divide by N. "M": divide by M. "NM": divide by (N*M). "ctc": divide by (M-1)/2. "none": no normalization. Default: "N" backend : str Backend to use. Choose among ("auto", "torch", "cpu_numba", "cuda_cpp"). Default: "auto" dtype_float: torch.dtype Number format for internal computations. Default: torch.float32 cuda_device : str or torch.device, optional CUDA device to use, for example ``"cuda:0"``. Default: ``None``. store_debug : bool Whether to retain intermediate backend matrices for inspection. Default: False. """ super().__init__(cost_function=cost_function, min_function="sparsemin", gamma=gamma, step_sizes=step_sizes, global_step_weights=global_step_weights, normalization=normalization, backend=backend, dtype_float=dtype_float, cuda_device=cuda_device, store_debug=store_debug)
[docs] def forward(self, X=None, Y=None, C=None, list_N=None, list_M=None): """ Compute the sparseDTW loss. Parameters ---------- X : torch.tensor [shape=(B, N, D)] Input sequence, usually the DNN predictions. Y : torch.tensor [shape=(B, M, D)] Input sequence, usually the weak targets. C : torch.tensor [shape=(B, N, M)] Pre-computed local cost matrix C. list_N : torch.tensor [shape=(B)] Sequence lengths of X (<= N) of the individual batch elements. If None, defaults to [N, N, ..., N] list_M : torch.tensor [shape=(B)] Sequence lengths of Y (<= M) of the individual batch elements. If None, defaults to [M, M, ..., M] Returns ------- torch.Tensor Scalar batch-mean sparseDTW loss. """ return super().forward(X=X, Y=Y, C=C, list_N=list_N, list_M=list_M)
################################################################################### ########################### subsequence SDTW ######################################
[docs] class subSDTW(dDTW): def __init__(self, cost_function = "MSE", min_function="softmin", gamma = 1.0, step_sizes = [[1,0], [0,1], [1,1]], global_step_weights = [1.0, 1.0, 1.0], normalization = "N", backend="auto", dtype_float = torch.float32, cuda_device=None, store_debug=False, sub_X=True, # whether to do sub-sequence along the X-direction (crop predictions) sub_Y=True, # whether to do sub-sequence along the Y-direction (crop targets) compensate_subseq=True # whether to compensate for shorter sequences ): """ Initialize subsequence SDTW loss function, see [7]. [7] Johannes Zeitler and Meinard Müller. Subsequence Soft Dynamic Time Warping. In Proceedings of the IEEE International Conference on Acoustics, Speech, and Signal Processing (CASSP), Barcelona, Spain, 2026. Parameters ---------- cost_function : str Local cost function for pair-wise comparison of sequence elements. Choose among ("MSE", "BCE", "CTC"). Default: "MSE". min_function : str Minimum function or approximation thereof. Choose among ("softmin", "sparsemin", "smoothmin", "hardmin"). Default: "softmin". gamma : float Min. function temperature hyperparameter. Default: 1.0 step_sizes : list Alignment step sizes in [n,m] direction, given as list of tuples. Default: [[1,0], [0,1], [1,1]] global_step_weights: list Step weights associated to the step sizes. Default: [1.0, 1.0, 1.0] normalization : str Normalization of SDTW cost. Choose among ("M", "N", "NM", "ctc", "none"). "N": divide by N. "M": divide by M. "NM": divide by (N*M). "ctc": divide by (M-1)/2. "none": no normalization. Default: "N" backend : str Backend to use. Choose among ("auto", "torch", "cpu_numba", "cuda_cpp"). Default: "auto" dtype_float: torch.dtype Number format for internal computations. Default: torch.float32 cuda_device : str or torch.device, optional CUDA device to use, for example ``"cuda:0"``. Default: ``None``. store_debug : bool Whether to retain intermediate backend matrices for inspection. Default: False. sub_X : bool Whether to allow subsequence starts and ends along the X-direction (row axis). Default: True. sub_Y : bool Whether to allow subsequence starts and ends along the Y-direction (column axis). Default: True. compensate_subseq : bool Whether to add start/end penalties for skipped prefixes or suffixes. Default: True. """ super().__init__(cost_function=cost_function, min_function=min_function, gamma=gamma, step_sizes=step_sizes, global_step_weights=global_step_weights, normalization=normalization, backend=backend, dtype_float=dtype_float, cuda_device=cuda_device, store_debug=store_debug) self.sub_X = sub_X self.sub_Y = sub_Y self.compensate_subseq = compensate_subseq
[docs] def forward(self, X=None, Y=None, C=None, list_N=None, list_M=None): """ Compute the subsequence SDTW loss. Parameters ---------- X : torch.tensor [shape=(B, N, D)] Input sequence, usually the DNN predictions. Y : torch.tensor [shape=(B, M, D)] Input sequence, usually the weak targets. C : torch.tensor [shape=(B, N, M)] Pre-computed local cost matrix C. list_N : torch.tensor [shape=(B)] Sequence lengths of X (<= N) of the individual batch elements. If None, defaults to [N, N, ..., N] list_M : torch.tensor [shape=(B)] Sequence lengths of Y (<= M) of the individual batch elements. If None, defaults to [M, M, ..., M] Returns ------- torch.Tensor Scalar batch-mean subsequence SDTW loss. """ if (X is not None) and (Y is not None): B = X.shape[0] N_max = X.shape[1] M_max = Y.shape[1] else: B = C.shape[0] N_max = C.shape[1] M_max = C.shape[2] if self.backend == "cuda_cpp": list_N = (torch.full((B,), N_max, device=self.device, dtype=torch.int64) if list_N is None else torch.as_tensor(list_N, device=self.device, dtype=torch.int64)) list_M = (torch.full((B,), M_max, device=self.device, dtype=torch.int64) if list_M is None else torch.as_tensor(list_M, device=self.device, dtype=torch.int64)) B_start, B_end, start_penalty, end_penalty, num_conditions = compute_subseq_initialization( list_N, list_M, self.global_step_weights, N_max, M_max, self.sub_X, self.sub_Y, self.compensate_subseq, ) else: B_start = [[] for _ in range(B)] B_end = [[] for _ in range(B)] start_penalty = [[] for _ in range(B)] end_penalty = [[] for _ in range(B)] if list_N is None: list_N = [N_max for _ in range(B)] if list_M is None: list_M = [M_max for _ in range(B)] for b in range(B): N = list_N[b] M = list_M[b] B_start[b].append([0,0]) start_penalty[b].append(1) B_end[b].append([N-1,M-1]) end_penalty[b].append(0) if self.sub_X: for n in range(1,N-1): B_start[b].append([n, 0]) B_end[b].append([n, M-1]) if self.compensate_subseq: start_penalty[b].append(1 + n*self.global_step_weights[0]) end_penalty[b].append( (N-1 - n)*self.global_step_weights[0]) else: start_penalty[b].append(1) end_penalty[b].append(0) if self.sub_Y: for m in range(1,M-1): B_start[b].append([0, m]) B_end[b].append([N-1,m]) if self.compensate_subseq: start_penalty[b].append(1 + m*self.global_step_weights[1]) end_penalty[b].append( (M-1-m)*self.global_step_weights[1]) else: start_penalty[b].append(1) end_penalty[b].append(0) num_conditions = None return super().forward(X=X, Y=Y, C=C, list_N=list_N, list_M=list_M, B_start=B_start, B_end=B_end, start_penalty=start_penalty, end_penalty=end_penalty, num_start_conditions=num_conditions, num_end_conditions=num_conditions)
################################################################################### ################################## CTC ############################################
[docs] class CTC(dDTW): # the algorithm assumes step sizes [(1,0), (1,1), (1,2)] def __init__(self, blank_penalty_weight = 1.0, blank_index=0, gamma = 1.0, global_step_weights = [1.0, 1.0, 1.0], backend="auto", dtype_float = torch.float32, cuda_device=None, store_debug=False ): """ Initialize CTC loss function, parameterized within the dDTW framework, see [8, 2]. [2] Johannes Zeitler and Meinard Müller. A Unified Perspective on CTC and Soft-DTW Using Differentiable DTW. IEEE Transactions on Audio, Speech and Language Processing, vol. 34, pages 936-951, 2026. [8] Alex Graves et al. Connectionist Temporal Classification: Labelling Unsegmented Sequence Data with Recurrent Neural Networks. In Proceedings of the International Conference on Machine Learning (ICML), pages 369-376, Pittsburgh, Pennsylvania, USA, 2006. Parameters ---------- blank_penalty_weight : float Penalty for alignment of the blank symbol. Default: 1.0 (no penalty). blank_index : int Class index of the blank symbol. Default: 0. gamma : float Softmin temperature hyperparameter. Default: 1.0 global_step_weights: list Step weights associated to the CTC step sizes ``[[1, 0], [1, 1], [1, 2]]``. Default: [1.0, 1.0, 1.0]. backend : str Backend to use. Choose among ("auto", "torch", "cpu_numba", "cuda_cpp"). Default: "auto" dtype_float: torch.dtype Number format for internal computations. Default: torch.float32 cuda_device : str or torch.device, optional CUDA device to use, for example ``"cuda:0"``. Default: ``None``. store_debug : bool Whether to retain intermediate backend matrices for inspection. Default: False. Notes ----- The CTC variant fixes ``cost_function="CTC"``, ``min_function="softmin"``, ``step_sizes=[[1, 0], [1, 1], [1, 2]]``, and ``normalization="ctc"``. The normalization divides each batch item by its unexpanded target length before averaging. """ super().__init__(cost_function="CTC", min_function="softmin", gamma=gamma, step_sizes=[[1,0], [1,1], [1,2]], global_step_weights=global_step_weights, normalization="ctc", backend=backend, dtype_float=dtype_float, cuda_device=cuda_device, store_debug=store_debug) self.blankP = blank_penalty_weight self.blank_index = blank_index
[docs] def forward(self, X=None, Y=None, list_N=None, list_M=None): """ Compute the CTC loss in the dDTW graph formulation. Parameters ---------- X : torch.tensor [shape=(B, N, D)] Input log-probabilities. The CTC local cost is the negative log-probability of the active target or blank state. Y : torch.tensor [shape=(B, M)] Integer target-label indices before blank expansion. Labels should use ``blank_index`` only for padding beyond ``list_M``. list_N : torch.tensor [shape=(B)] Sequence lengths of X (<= N) of the individual batch elements. If None, defaults to [N, N, ..., N] list_M : torch.tensor [shape=(B)] Target-label lengths before CTC blank expansion. If None, defaults to [M, M, ..., M]. Returns ------- torch.Tensor Scalar batch-mean CTC loss. """ X = X.to(device=self.device) Y = Y.to(device=self.device) if Y.dim() == 2: CTC_targets = True else: CTC_targets = False B = X.shape[0] N_max = X.shape[1] M_max = Y.shape[1] M_e = 2*M_max + 1 device = torch.device(self.device) dtype = X.dtype D = X.shape[2] if self.backend == "cuda_cpp": list_N_tensor = (torch.full((B,), N_max, device=device, dtype=torch.int64) if list_N is None else torch.as_tensor(list_N, device=device, dtype=torch.int64)) list_M_tensor = (torch.full((B,), M_max, device=device, dtype=torch.int64) if list_M is None else torch.as_tensor(list_M, device=device, dtype=torch.int64)) W, Y_e, list_M_e, B_start, B_end = compute_ctc_initialization( X, Y.to(device=device, dtype=torch.int64), list_N_tensor, list_M_tensor, self.global_step_weights, self.blankP, self.blank_index, ) list_N = list_N_tensor else: if list_N is None: list_N = [N_max for _ in range(B)] if list_M is None: list_M = [M_max for _ in range(B)] Y_e = torch.zeros((B, M_e, D), device=device, dtype=dtype) W = torch.zeros((B, N_max, M_e, 3), device=device, dtype=dtype) for i_w, w in enumerate(self.global_step_weights): W[:,:,1::2,i_w] = w # transition into an acutal target W[:,:,0::2,i_w] = self.blankP # transition into blank B_start = [] B_end = [] list_M_e = [] for b in range(B): list_M_e.append(2*list_M[b]+1) B_start.append([[0,0], [0,1]]) B_end.append([[list_N[b]-1, 2*list_M[b]+1-1], [list_N[b]-1, 2*list_M[b]+1-2]]) last_tgt = None for m_c in range(list_M[b]): Y_e[b,2*m_c,self.blank_index] = 1 Y_e[b,2*m_c+1, Y[b,m_c]] = 1 # skipping blank is never allowed W[b,:,2*m_c,-1] = 1e20 tgt = Y[b,m_c] if tgt == last_tgt: # skipping identical targets is not allowed W[b,:,2*m_c+1,-1] = 1e20 # but we must allow to go through a blank symbol with a (1,1) step W[b,:,2*m_c, 1] = self.global_step_weights[1] last_tgt = tgt # skipping the last blank is also not allowed W[b,:,2*list_M[b],-1] = 1e20 Y_e[b,2*list_M[b],self.blank_index] = 1 return super().forward(X=X, Y=Y_e, list_N=list_N, list_M=list_M_e, B_start=B_start, B_end=B_end, local_step_weights=W)
################################################################################### ############################ partial matching #####################################
[docs] class partial_matching(dDTW): def __init__(self, cost_function = "CTC", min_function="hardmin", gamma = 1.0, normalization = "none", backend="auto", dtype_float = torch.float32, cuda_device=None, store_debug=False ): """ Initialize partial matching loss function, parameterized within the dDTW framework, see [9, 2]. [2] Johannes Zeitler and Meinard Müller. A Unified Perspective on CTC and Soft-DTW Using Differentiable DTW. IEEE Transactions on Audio, Speech and Language Processing, vol. 34, pages 936-951, 2026. [9] Pavel A. Pevzner. Computational Molecular Biology: An Algorithmic Approach. MIT Press, 2000. Parameters ---------- cost_function : str Local cost function used when ``X`` and ``Y`` are supplied. Default: "CTC". min_function : str Minimum function or differentiable approximation thereof. Default: "hardmin". gamma : float Temperature parameter for differentiable minimum functions. Default: 1.0. normalization : str Normalization of PM cost. Choose among ("M", "N", "NM", "ctc", "none"). "N": divide by N. "M": divide by M. "NM": divide by (N*M). "ctc": divide by (M-1)/2. "none": no normalization. Default: "none". backend : str Backend to use. Choose among ("auto", "torch", "cpu_numba", "cuda_cpp"). Default: "auto" dtype_float: torch.dtype Number format for internal computations. Default: torch.float32 cuda_device : str or torch.device, optional CUDA device to use, for example ``"cuda:0"``. Default: ``None``. store_debug : bool Whether to retain intermediate backend matrices for inspection. Default: False. Notes ----- This variant fixes ``step_sizes=[[1, 0], [0, 1], [1, 1]]`` and ``global_step_weights=[0.0, 0.0, 1.0]``. Horizontal and vertical moves therefore do not accumulate local cost; only diagonal matches do. """ super().__init__(cost_function=cost_function, min_function=min_function, gamma=gamma, step_sizes=[[1,0], [0,1], [1,1]], global_step_weights=[0., 0., 1.], normalization=normalization, backend=backend, dtype_float=dtype_float, cuda_device=cuda_device, store_debug=store_debug)
[docs] def forward(self, X=None, Y=None, C=None, list_N=None, list_M=None): """ Compute the partial matching loss. Parameters ---------- X : torch.tensor [shape=(B, N, D)] Input sequence, usually the DNN predictions. Y : torch.tensor [shape=(B, M, D)] Input sequence, usually the weak targets. C : torch.tensor [shape=(B, N, M)] Pre-computed local cost matrix C. To compare against a score-maximizing partial matching reference, pass ``C=-S`` for score matrix ``S``. list_N : torch.tensor [shape=(B)] Sequence lengths of X (<= N) of the individual batch elements. If None, defaults to [N, N, ..., N] list_M : torch.tensor [shape=(B)] Sequence lengths of Y (<= M) of the individual batch elements. If None, defaults to [M, M, ..., M] Returns ------- torch.Tensor Scalar batch-mean partial matching loss. """ if (X is not None) and (Y is not None): B = X.shape[0] N_max = X.shape[1] M_max = Y.shape[1] else: B = C.shape[0] N_max = C.shape[1] M_max = C.shape[2] if list_N is None: list_N = [N_max for _ in range(B)] if list_M is None: list_M = [M_max for _ in range(B)] # all cells of the cost matrix I = [] for b in range(B): I.append([[n,m] for n,m in product(range(list_N[b]), range(list_M[b]))]) return super().forward(X=X, Y=Y, C=C, list_N=list_N, list_M=list_M, B_start=I, B_end=I)
###################################################################################