Source code for ddtw.ddtw

##################################################################################
#                                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.                            #
#                                                                                #
# Code based on:                                                                 #
# Mehran Maghoumi et al. "DeepNAG: Deep Non-Adversarial Gesture Generation".     #
#  International Conference on Intelligent User Interfaces, 2021.                #
#  https://github.com/Maghoumi/pytorch-softdtw-cuda/                             #
##################################################################################


##################################################################################
# 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
import torch.cuda

from .cost_function import get_cost_function
from .backend.backend_torch import _backend_torch
from .backend.backend_cuda_cpp import _backend_CUDA_CPP

try:
    from .backend.backend_cpu_numba import _backend_CPU_Numba
except ImportError:
    _backend_CPU_Numba = None


SUPPORTED_BACKENDS = ("auto", "torch", "cpu_numba", "cuda_cpp")
MIN_FUNCTION_IDS = {
    "softmin": 1,
    "sparsemin": 2,
    "smoothmin": 3,
    "hardmin": 4,
}


def _cuda_device_available(cuda_device):
    if not torch.cuda.is_available():
        return False
    if cuda_device is None:
        return True

    try:
        device = torch.device(cuda_device)
    except (TypeError, RuntimeError):
        return False

    if device.type != "cuda":
        return False
    if device.index is not None and device.index >= torch.cuda.device_count():
        return False
    return True


def _resolve_backend(backend, cuda_device):
    if backend == "auto":
        if _cuda_device_available(cuda_device):
            return "cuda_cpp"
        if _backend_CPU_Numba is not None:
            return "cpu_numba"
        return "torch"

    if backend not in SUPPORTED_BACKENDS:
        raise ValueError(
            "Unsupported backend: %s. Choose among %s."
            % (backend, SUPPORTED_BACKENDS)
        )
    if backend == "cpu_numba" and _backend_CPU_Numba is None:
        raise ImportError("backend='cpu_numba' requires numba to be installed")
    if backend == "cuda_cpp" and not _cuda_device_available(cuda_device):
        raise RuntimeError("backend='cuda_cpp' requires an available CUDA device")
    if backend == "torch" and cuda_device is not None:
        device = torch.device(cuda_device)
        if device.type == "cuda" and not _cuda_device_available(cuda_device):
            raise RuntimeError("backend='torch' with a CUDA device requires CUDA")
    return backend


def _backend_device(backend, cuda_device):
    if backend == "cuda_cpp":
        return "cuda" if cuda_device is None else cuda_device
    if backend == "torch" and cuda_device is not None:
        return cuda_device
    return "cpu"


# dDTW loss class
[docs] class dDTW(torch.nn.Module): 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 ): """Initialize the general dDTW loss function. See [1] for the graph formulation. [1] 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 or callable, optional Local cost function used when ``X`` and ``Y`` are passed to :meth:`forward`. Built-in strings are ``"MSE"``, ``"BCE"``, and ``"CTC"``. A callable must return a cost tensor with shape ``(B, N, M)``. Default: ``"MSE"``. min_function : str, optional Recursive minimum or differentiable approximation. Choose among ``"softmin"``, ``"sparsemin"``, ``"smoothmin"``, and ``"hardmin"``. Default: ``"softmin"``. gamma : float, optional Temperature parameter used by differentiable minimum functions. Default: ``1.0``. step_sizes : list of list of int, optional Alignment step sizes ``[dn, dm]``. Each step points from the current cell ``(n, m)`` to predecessor ``(n-dn, m-dm)``. Default: ``[[1, 0], [0, 1], [1, 1]]``. global_step_weights : list of float, optional Local cost weights associated with ``step_sizes``. Must contain one scalar per step. Default: ``[1.0, 1.0, 1.0]``. normalization : str, optional Normalization applied to each batch loss before averaging. Choose among ``"N"``, ``"M"``, ``"NM"``, ``"ctc"``, and ``"none"``. Default: ``"N"``. backend : str, optional Backend to use. Choose among ``"auto"``, ``"torch"``, ``"cpu_numba"``, and ``"cuda_cpp"``. ``"auto"`` tries CUDA first, then Numba CPU, then pure PyTorch. Default: ``"auto"``. dtype_float : torch.dtype, optional Floating-point dtype for internal tensors. The CUDA C++ backend currently requires ``torch.float32``. Default: ``torch.float32``. cuda_device : str or torch.device, optional CUDA device to use, for example ``"cuda:0"``. Default: ``None``. store_debug : bool, optional If ``True``, retain intermediate matrices on the backend class for inspection after forward/backward. Default: ``False``. """ super(dDTW, self).__init__() self.dtype_float=dtype_float self.store_debug = store_debug self.backend = _resolve_backend(backend, cuda_device) self.requested_backend = backend self.device = _backend_device(self.backend, cuda_device) if self.backend == "torch": self.core = _backend_torch self.compute_dDTW = _backend_torch.apply elif self.backend == "cpu_numba": self.core = _backend_CPU_Numba self.compute_dDTW = _backend_CPU_Numba.apply else: self.core = _backend_CUDA_CPP self.compute_dDTW = _backend_CUDA_CPP.apply if normalization not in ["M", "N", "NM", "ctc", "none"]: raise ValueError( "Unsupported normalization: %s. Choose among %s." % (normalization, ("M", "N", "NM", "ctc", "none")) ) self.normalization = normalization self.gamma = torch.tensor([gamma], device=self.device, dtype=self.dtype_float, requires_grad=False) self.step_sizes = torch.tensor(step_sizes, device=self.device, dtype=torch.int16, requires_grad=False) if self.step_sizes.dim() != 2 or self.step_sizes.shape[1] != 2: raise ValueError("step_sizes must have shape (num_steps, 2)") self.global_step_weights = torch.tensor(global_step_weights, device=self.device, dtype=self.dtype_float, requires_grad=False) if self.global_step_weights.dim() != 1 or self.global_step_weights.shape[0] != self.step_sizes.shape[0]: raise ValueError("global_step_weights must contain one scalar weight per step size") if min_function not in MIN_FUNCTION_IDS: raise ValueError( "Unsupported min_function: %s. Choose among %s." % (min_function, tuple(MIN_FUNCTION_IDS)) ) self.min_function = MIN_FUNCTION_IDS[min_function] if isinstance(cost_function, str): self.cost_function = get_cost_function(cost_function, self.backend == "cuda_cpp") else: self.cost_function = cost_function return None def _as_fixed_tensor(self, data, dtype): if isinstance(data, torch.Tensor): return data.detach().to(device=self.device, dtype=dtype) return torch.as_tensor(data, device=self.device, dtype=dtype) def _prepare_boundary_conditions(self, boundary_conditions, penalties, num_conditions, default_boundary, default_penalty, B): if num_conditions is not None: num_conditions = self._as_fixed_tensor(num_conditions, torch.int32) if num_conditions.shape != torch.Size([B]): raise ValueError("num_conditions must have shape (B,)") if boundary_conditions is None: boundary_tensor = default_boundary.to(device=self.device, dtype=torch.int16) if num_conditions is None: num_conditions = torch.ones(B, device=self.device, dtype=torch.int32, requires_grad=False) elif isinstance(boundary_conditions, torch.Tensor): if boundary_conditions.dim() != 3: raise ValueError("boundary_conditions tensor must have shape (B, num_conditions, 2)") if boundary_conditions.shape[0] != B: raise ValueError("boundary_conditions batch dimension must match the cost matrix batch size") if boundary_conditions.shape[2] != 2: raise ValueError("boundary_conditions must store two indices per condition") boundary_tensor = boundary_conditions.detach().to(device=self.device, dtype=torch.int16) if num_conditions is None: if penalties is not None and not isinstance(penalties, torch.Tensor): num_conditions = torch.as_tensor([len(batch_penalties) for batch_penalties in penalties], device=self.device, dtype=torch.int32) else: num_conditions = torch.full((B,), boundary_tensor.shape[1], device=self.device, dtype=torch.int32, requires_grad=False) else: if len(boundary_conditions) != B: raise ValueError("boundary_conditions must contain one condition list per batch item") num_conditions_list = [len(batch_conditions) for batch_conditions in boundary_conditions] if num_conditions is not None: expected_num_conditions = torch.as_tensor(num_conditions_list, dtype=torch.int32) if not torch.equal(num_conditions.cpu(), expected_num_conditions): raise ValueError("num_conditions does not match the provided boundary_conditions") max_conditions = max(num_conditions_list) if max_conditions <= 0: raise ValueError("Each batch must provide at least one boundary condition") padded_conditions = [] for batch_conditions in boundary_conditions: if isinstance(batch_conditions, torch.Tensor): batch_conditions = batch_conditions.detach().cpu().tolist() batch_conditions = [list(condition) for condition in batch_conditions] batch_conditions.extend([[0, 0] for _ in range(max_conditions - len(batch_conditions))]) padded_conditions.append(batch_conditions) boundary_tensor = torch.as_tensor(padded_conditions, device=self.device, dtype=torch.int16) num_conditions = torch.as_tensor(num_conditions_list, device=self.device, dtype=torch.int32) max_conditions = boundary_tensor.shape[1] if int(torch.max(num_conditions)) > max_conditions: raise ValueError("num_conditions cannot exceed the padded number of boundary conditions") if penalties is None: weights = torch.full((B, max_conditions), default_penalty, device=self.device, dtype=self.dtype_float, requires_grad=False) elif isinstance(penalties, torch.Tensor): if penalties.shape != torch.Size([B, max_conditions]): raise ValueError("penalties must have shape (B, max_conditions)") weights = penalties.detach().to(device=self.device, dtype=self.dtype_float) else: if len(penalties) != B: raise ValueError("penalties must contain one penalty list per batch item") padded_penalties = [] for batch_penalties, count in zip(penalties, num_conditions.detach().cpu().tolist()): if isinstance(batch_penalties, torch.Tensor): batch_penalties = batch_penalties.detach().cpu().tolist() if len(batch_penalties) != count: raise ValueError("Each penalty list must match the corresponding number of boundary conditions") padded_batch = list(batch_penalties) padded_batch.extend([default_penalty for _ in range(max_conditions - len(padded_batch))]) padded_penalties.append(padded_batch) weights = torch.as_tensor(padded_penalties, device=self.device, dtype=self.dtype_float) return boundary_tensor, num_conditions, weights
[docs] def forward(self, X=None, Y=None, C=None, B_start=None, B_end=None, list_N=None, list_M=None, local_step_weights=None, start_penalty=None, end_penalty=None, num_start_conditions=None, num_end_conditions=None): """Compute the dDTW loss. Pass either ``X`` and ``Y`` or a precomputed cost matrix ``C``. If ``X`` and ``Y`` are provided, ``self.cost_function`` computes ``C``. Parameters ---------- X : torch.Tensor, optional First input sequence with shape ``(B, N, D)``. Usually the model predictions. Y : torch.Tensor, optional Second input sequence with shape ``(B, M, D)``. Usually the target or reference sequence. C : torch.Tensor, optional Precomputed local cost matrix with shape ``(B, N, M)``. B_start : list or torch.Tensor, optional Start boundary conditions. For each batch item, stores one or more zero-based ``[n, m]`` cells. Tensor form must have shape ``(B, max_start_conditions, 2)``. If ``None``, defaults to ``[[0, 0]]`` for each batch item. B_end : list or torch.Tensor, optional End boundary conditions. For each batch item, stores one or more zero-based ``[n, m]`` cells. Tensor form must have shape ``(B, max_end_conditions, 2)``. If ``None``, defaults to ``[[list_N[b] - 1, list_M[b] - 1]]``. list_N : list or torch.Tensor, optional Active lengths along the ``X``/row axis with shape ``(B,)``. If ``None``, all batch items use the full padded length ``N``. list_M : list or torch.Tensor, optional Active lengths along the ``Y``/column axis with shape ``(B,)``. If ``None``, all batch items use the full padded length ``M``. local_step_weights : torch.Tensor, optional Cell-wise step weights with shape ``(B, N, M, S)``, where ``S`` is the number of configured steps. If ``None``, ``global_step_weights`` are broadcast to all cells. start_penalty : list or torch.Tensor, optional Multiplicative local-cost weights for start boundary conditions. Tensor form must have shape ``(B, max_start_conditions)``. If ``None``, defaults to ``1`` for every start condition. end_penalty : list or torch.Tensor, optional Multiplicative local-cost weights for end boundary conditions. Tensor form must have shape ``(B, max_end_conditions)``. If ``None``, defaults to ``0`` for every end condition. num_start_conditions : list or torch.Tensor, optional Number of valid start conditions for each batch item when ``B_start`` is padded. Shape ``(B,)``. num_end_conditions : list or torch.Tensor, optional Number of valid end conditions for each batch item when ``B_end`` is padded. Shape ``(B,)``. Returns ------- torch.Tensor Scalar batch-mean dDTW loss after the configured normalization. """ # cost matrix ##################################################################################### if X is not None: X = X.to(device=self.device) if Y is not None: Y = Y.to(device=self.device) if C is not None: C = C.to(device=self.device) if (X is not None) and (Y is not None): if (len(X.shape) == 3) and (len(Y.shape) == 3): X_ = X[:,:,:,None] Y_ = Y[:,:,:,None] #print("expanded dim.") C = self.cost_function(X_, Y_) else: C = self.cost_function(X, Y) else: if C is None: raise ValueError("Pass either X and Y, or a precomputed cost matrix C.") B = C.shape[0] N = C.shape[1] M = C.shape[2] ################################################################################################## # sequence lengths ################################################################################ if list_N is None: self.list_N = torch.full((B,), N, device=self.device, dtype=torch.int16, requires_grad=False) else: self.list_N = self._as_fixed_tensor(list_N, torch.int16) if list_M is None: self.list_M = torch.full((B,), M, device=self.device, dtype=torch.int16, requires_grad=False) else: self.list_M = self._as_fixed_tensor(list_M, torch.int16) #################################################################################################### # boundary conditions ############################################################################## default_B_start = torch.zeros((B, 1, 2), device=self.device, dtype=torch.int16, requires_grad=False) default_B_end = torch.stack((self.list_N - 1, self.list_M - 1), dim=-1).unsqueeze(1) self.B_start, self.num_start_conditions, self.W_start = self._prepare_boundary_conditions( B_start, start_penalty, num_start_conditions, default_B_start, 1.0, B) self.B_end, self.num_end_conditions, self.W_end = self._prepare_boundary_conditions( B_end, end_penalty, num_end_conditions, default_B_end, 0.0, B) ######################################################################################################## # step weights ######################################################################################### if local_step_weights is not None: self.step_weights = self._as_fixed_tensor(local_step_weights, self.dtype_float) else: self.step_weights = self.global_step_weights.expand(B,N,M,-1) ######################################################################################################## dDTW_cost = self.compute_dDTW(C, self.min_function, self.gamma, self.step_sizes, self.step_weights, self.list_N, self.list_M, self.B_start, self.B_end, self.num_start_conditions, self.num_end_conditions, self.W_start, self.W_end, self.store_debug) if self.normalization == "N": dDTW_cost_mean = torch.mean(dDTW_cost/self.list_N) elif self.normalization == "M": dDTW_cost_mean = torch.mean(dDTW_cost/self.list_M) elif self.normalization == "NM": dDTW_cost_mean = torch.mean(dDTW_cost/self.list_N/self.list_M) elif self.normalization == "ctc": dDTW_cost_mean = torch.mean(dDTW_cost/( (self.list_M-1)/2)) else: dDTW_cost_mean = torch.mean(dDTW_cost) return dDTW_cost_mean