ddtw 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- ddtw/__init__.py +57 -0
- ddtw/backend/__init__.py +37 -0
- ddtw/backend/backend_cpu_numba.py +447 -0
- ddtw/backend/backend_cuda_cpp.py +215 -0
- ddtw/backend/backend_torch.py +378 -0
- ddtw/backend/cpp_extension.py +67 -0
- ddtw/backend/csrc/ddtw_cuda.cu +1105 -0
- ddtw/backend/csrc/ddtw_extension.cpp +64 -0
- ddtw/cost_function.py +122 -0
- ddtw/ddtw.py +436 -0
- ddtw/ddtw_variants.py +906 -0
- ddtw-0.1.0.dist-info/METADATA +290 -0
- ddtw-0.1.0.dist-info/RECORD +16 -0
- ddtw-0.1.0.dist-info/WHEEL +5 -0
- ddtw-0.1.0.dist-info/licenses/LICENSE +21 -0
- ddtw-0.1.0.dist-info/top_level.txt +1 -0
ddtw/__init__.py
ADDED
|
@@ -0,0 +1,57 @@
|
|
|
1
|
+
##################################################################################
|
|
2
|
+
# dDTW Toolbox #
|
|
3
|
+
##################################################################################
|
|
4
|
+
# #
|
|
5
|
+
# Authors: Johannes Zeitler and Meinard Müller, 2026 #
|
|
6
|
+
# #
|
|
7
|
+
# If you use this toolbox, please cite the accompanying paper: #
|
|
8
|
+
# Johannes Zeitler and Meinard Müller. dDTW: A Unified and Efficient Toolbox for #
|
|
9
|
+
# Differentiable Sequence Alignment. Submitted 2026. #
|
|
10
|
+
##################################################################################
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
##################################################################################
|
|
14
|
+
# MIT License #
|
|
15
|
+
# #
|
|
16
|
+
# Copyright 2026 Johannes Zeitler and Meinard Müller #
|
|
17
|
+
# #
|
|
18
|
+
# Permission is hereby granted, free of charge, to any person obtaining a copy #
|
|
19
|
+
# of this software and associated documentation files (the "Software"), to deal #
|
|
20
|
+
# in the Software without restriction, including without limitation the rights #
|
|
21
|
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell #
|
|
22
|
+
# copies of the Software, and to permit persons to whom the Software is #
|
|
23
|
+
# furnished to do so, subject to the following conditions: #
|
|
24
|
+
# #
|
|
25
|
+
# The above copyright notice and this permission notice shall be included in all #
|
|
26
|
+
# copies or substantial portions of the Software. #
|
|
27
|
+
# #
|
|
28
|
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR #
|
|
29
|
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, #
|
|
30
|
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE #
|
|
31
|
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER #
|
|
32
|
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, #
|
|
33
|
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE #
|
|
34
|
+
# SOFTWARE. #
|
|
35
|
+
##################################################################################
|
|
36
|
+
|
|
37
|
+
"""Public package interface for the dDTW toolbox."""
|
|
38
|
+
|
|
39
|
+
from .ddtw import dDTW
|
|
40
|
+
from .ddtw_variants import CTC
|
|
41
|
+
from .ddtw_variants import DTW
|
|
42
|
+
from .ddtw_variants import SDTW
|
|
43
|
+
from .ddtw_variants import partial_matching
|
|
44
|
+
from .ddtw_variants import smoothDTW
|
|
45
|
+
from .ddtw_variants import sparseDTW
|
|
46
|
+
from .ddtw_variants import subSDTW
|
|
47
|
+
|
|
48
|
+
__all__ = [
|
|
49
|
+
"dDTW",
|
|
50
|
+
"SDTW",
|
|
51
|
+
"DTW",
|
|
52
|
+
"smoothDTW",
|
|
53
|
+
"sparseDTW",
|
|
54
|
+
"subSDTW",
|
|
55
|
+
"CTC",
|
|
56
|
+
"partial_matching",
|
|
57
|
+
]
|
ddtw/backend/__init__.py
ADDED
|
@@ -0,0 +1,37 @@
|
|
|
1
|
+
##################################################################################
|
|
2
|
+
# dDTW Toolbox #
|
|
3
|
+
##################################################################################
|
|
4
|
+
# #
|
|
5
|
+
# Authors: Johannes Zeitler and Meinard Müller, 2026 #
|
|
6
|
+
# #
|
|
7
|
+
# If you use this toolbox, please cite the accompanying paper: #
|
|
8
|
+
# Johannes Zeitler and Meinard Müller. dDTW: A Unified and Efficient Toolbox for #
|
|
9
|
+
# Differentiable Sequence Alignment. Submitted 2026. #
|
|
10
|
+
##################################################################################
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
##################################################################################
|
|
14
|
+
# MIT License #
|
|
15
|
+
# #
|
|
16
|
+
# Copyright 2026 Johannes Zeitler and Meinard Müller #
|
|
17
|
+
# #
|
|
18
|
+
# Permission is hereby granted, free of charge, to any person obtaining a copy #
|
|
19
|
+
# of this software and associated documentation files (the "Software"), to deal #
|
|
20
|
+
# in the Software without restriction, including without limitation the rights #
|
|
21
|
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell #
|
|
22
|
+
# copies of the Software, and to permit persons to whom the Software is #
|
|
23
|
+
# furnished to do so, subject to the following conditions: #
|
|
24
|
+
# #
|
|
25
|
+
# The above copyright notice and this permission notice shall be included in all #
|
|
26
|
+
# copies or substantial portions of the Software. #
|
|
27
|
+
# #
|
|
28
|
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR #
|
|
29
|
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, #
|
|
30
|
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE #
|
|
31
|
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER #
|
|
32
|
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, #
|
|
33
|
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE #
|
|
34
|
+
# SOFTWARE. #
|
|
35
|
+
##################################################################################
|
|
36
|
+
|
|
37
|
+
"""Backend implementations for dDTW."""
|
|
@@ -0,0 +1,447 @@
|
|
|
1
|
+
##################################################################################
|
|
2
|
+
# dDTW Toolbox #
|
|
3
|
+
##################################################################################
|
|
4
|
+
# #
|
|
5
|
+
# Authors: Johannes Zeitler and Meinard Müller, 2026 #
|
|
6
|
+
# #
|
|
7
|
+
# If you use this toolbox, please cite the accompanying paper: #
|
|
8
|
+
# Johannes Zeitler and Meinard Müller. dDTW: A Unified and Efficient Toolbox for #
|
|
9
|
+
# Differentiable Sequence Alignment. Submitted 2026. #
|
|
10
|
+
# #
|
|
11
|
+
# Code based on: #
|
|
12
|
+
# Mehran Maghoumi et al. "DeepNAG: Deep Non-Adversarial Gesture Generation". #
|
|
13
|
+
# International Conference on Intelligent User Interfaces, 2021. #
|
|
14
|
+
# https://github.com/Maghoumi/pytorch-softdtw-cuda/ #
|
|
15
|
+
##################################################################################
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
##################################################################################
|
|
19
|
+
# MIT License #
|
|
20
|
+
# #
|
|
21
|
+
# Copyright 2026 Johannes Zeitler and Meinard Müller #
|
|
22
|
+
# #
|
|
23
|
+
# Permission is hereby granted, free of charge, to any person obtaining a copy #
|
|
24
|
+
# of this software and associated documentation files (the "Software"), to deal #
|
|
25
|
+
# in the Software without restriction, including without limitation the rights #
|
|
26
|
+
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell #
|
|
27
|
+
# copies of the Software, and to permit persons to whom the Software is #
|
|
28
|
+
# furnished to do so, subject to the following conditions: #
|
|
29
|
+
# #
|
|
30
|
+
# The above copyright notice and this permission notice shall be included in all #
|
|
31
|
+
# copies or substantial portions of the Software. #
|
|
32
|
+
# #
|
|
33
|
+
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR #
|
|
34
|
+
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, #
|
|
35
|
+
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE #
|
|
36
|
+
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER #
|
|
37
|
+
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, #
|
|
38
|
+
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE #
|
|
39
|
+
# SOFTWARE. #
|
|
40
|
+
##################################################################################
|
|
41
|
+
|
|
42
|
+
"""Numba-compiled CPU backend for differentiable DTW.
|
|
43
|
+
|
|
44
|
+
The dynamic-programming recurrences operate on zero-copy NumPy views of CPU
|
|
45
|
+
PyTorch tensors. Each batch item is independent and is processed by one Numba
|
|
46
|
+
worker; cells within an item retain their dependency-preserving order.
|
|
47
|
+
"""
|
|
48
|
+
|
|
49
|
+
import math
|
|
50
|
+
|
|
51
|
+
import torch
|
|
52
|
+
from numba import njit, prange
|
|
53
|
+
from torch.autograd import Function
|
|
54
|
+
|
|
55
|
+
|
|
56
|
+
MINFUNC_SOFTMIN = 1
|
|
57
|
+
MINFUNC_SPARSEMIN = 2
|
|
58
|
+
MINFUNC_SMOOTHMIN = 3
|
|
59
|
+
MINFUNC_HARDMIN = 4
|
|
60
|
+
|
|
61
|
+
|
|
62
|
+
@njit(inline="always")
|
|
63
|
+
def _all_inf(values):
|
|
64
|
+
for i in range(values.shape[0]):
|
|
65
|
+
if not math.isinf(values[i]):
|
|
66
|
+
return False
|
|
67
|
+
return True
|
|
68
|
+
|
|
69
|
+
|
|
70
|
+
@njit(inline="always")
|
|
71
|
+
def _softmin(values, gamma, grads):
|
|
72
|
+
size = values.shape[0]
|
|
73
|
+
if _all_inf(values):
|
|
74
|
+
value = math.inf
|
|
75
|
+
uniform = 1.0 / size
|
|
76
|
+
for i in range(size):
|
|
77
|
+
grads[i] = uniform
|
|
78
|
+
return value
|
|
79
|
+
|
|
80
|
+
minimum = values[0]
|
|
81
|
+
for i in range(1, size):
|
|
82
|
+
if values[i] < minimum:
|
|
83
|
+
minimum = values[i]
|
|
84
|
+
|
|
85
|
+
exp_sum = 0.0
|
|
86
|
+
for i in range(size):
|
|
87
|
+
weight = math.exp(-(values[i] - minimum) / gamma)
|
|
88
|
+
grads[i] = weight
|
|
89
|
+
exp_sum += weight
|
|
90
|
+
|
|
91
|
+
for i in range(size):
|
|
92
|
+
grads[i] /= exp_sum
|
|
93
|
+
return minimum - gamma * math.log(exp_sum)
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
@njit(inline="always")
|
|
97
|
+
def _smoothmin(values, gamma, grads):
|
|
98
|
+
size = values.shape[0]
|
|
99
|
+
for i in range(size):
|
|
100
|
+
if values[i] > 1.0e20:
|
|
101
|
+
values[i] = 1.0e20
|
|
102
|
+
|
|
103
|
+
_softmin(values, gamma, grads)
|
|
104
|
+
value = 0.0
|
|
105
|
+
for i in range(size):
|
|
106
|
+
value += values[i] * grads[i]
|
|
107
|
+
for i in range(size):
|
|
108
|
+
grads[i] *= 1.0 - (values[i] - value) / gamma
|
|
109
|
+
return value
|
|
110
|
+
|
|
111
|
+
|
|
112
|
+
@njit(inline="always")
|
|
113
|
+
def _sparsemin(values, gamma, grads, sparsemin_sort_buffer):
|
|
114
|
+
size = values.shape[0]
|
|
115
|
+
if _all_inf(values):
|
|
116
|
+
value = math.inf
|
|
117
|
+
uniform = 1.0 / size
|
|
118
|
+
for i in range(size):
|
|
119
|
+
grads[i] = uniform
|
|
120
|
+
return value
|
|
121
|
+
|
|
122
|
+
for i in range(size):
|
|
123
|
+
capped = min(values[i], 1.0e20)
|
|
124
|
+
projected = -capped / gamma
|
|
125
|
+
grads[i] = projected
|
|
126
|
+
sparsemin_sort_buffer[i] = projected
|
|
127
|
+
|
|
128
|
+
# Insertion sort in descending order. The number of steps is normally
|
|
129
|
+
# very small, so this avoids allocating or invoking a general sorter.
|
|
130
|
+
for i in range(1, size):
|
|
131
|
+
key = sparsemin_sort_buffer[i]
|
|
132
|
+
j = i - 1
|
|
133
|
+
while j >= 0 and sparsemin_sort_buffer[j] < key:
|
|
134
|
+
sparsemin_sort_buffer[j + 1] = sparsemin_sort_buffer[j]
|
|
135
|
+
j -= 1
|
|
136
|
+
sparsemin_sort_buffer[j + 1] = key
|
|
137
|
+
|
|
138
|
+
cumulative = sparsemin_sort_buffer[0] - 1.0
|
|
139
|
+
cmax = cumulative
|
|
140
|
+
rho = 1
|
|
141
|
+
for i in range(1, size):
|
|
142
|
+
cumulative += sparsemin_sort_buffer[i]
|
|
143
|
+
if sparsemin_sort_buffer[i] - cumulative / (i + 1) > 0.0:
|
|
144
|
+
rho = i + 1
|
|
145
|
+
cmax = cumulative
|
|
146
|
+
else:
|
|
147
|
+
break
|
|
148
|
+
theta = cmax / rho
|
|
149
|
+
|
|
150
|
+
value = -gamma / 2.0
|
|
151
|
+
for i in range(size):
|
|
152
|
+
projected = max(grads[i] - theta, 0.0)
|
|
153
|
+
grads[i] = projected
|
|
154
|
+
value -= projected * (projected + theta - 0.5 * projected) * gamma
|
|
155
|
+
return value
|
|
156
|
+
|
|
157
|
+
|
|
158
|
+
@njit(inline="always")
|
|
159
|
+
def _hardmin(values, grads):
|
|
160
|
+
size = values.shape[0]
|
|
161
|
+
if _all_inf(values):
|
|
162
|
+
value = math.inf
|
|
163
|
+
uniform = 1.0 / size
|
|
164
|
+
for i in range(size):
|
|
165
|
+
grads[i] = uniform
|
|
166
|
+
return value
|
|
167
|
+
|
|
168
|
+
minimum_index = 0
|
|
169
|
+
minimum = values[0]
|
|
170
|
+
for i in range(size):
|
|
171
|
+
grads[i] = 0.0
|
|
172
|
+
if values[i] < minimum:
|
|
173
|
+
minimum = values[i]
|
|
174
|
+
minimum_index = i
|
|
175
|
+
grads[minimum_index] = 1.0
|
|
176
|
+
return minimum
|
|
177
|
+
|
|
178
|
+
|
|
179
|
+
@njit(inline="always")
|
|
180
|
+
def _minimum(values, gamma, min_func_id, grads, sparsemin_sort_buffer):
|
|
181
|
+
if min_func_id == MINFUNC_SOFTMIN:
|
|
182
|
+
return _softmin(values, gamma, grads)
|
|
183
|
+
if min_func_id == MINFUNC_SPARSEMIN:
|
|
184
|
+
return _sparsemin(values, gamma, grads, sparsemin_sort_buffer)
|
|
185
|
+
if min_func_id == MINFUNC_SMOOTHMIN:
|
|
186
|
+
return _smoothmin(values, gamma, grads)
|
|
187
|
+
return _hardmin(values, grads)
|
|
188
|
+
|
|
189
|
+
|
|
190
|
+
@njit(parallel=True, nogil=True, cache=True)
|
|
191
|
+
def compute_dDTW_forward_numba(
|
|
192
|
+
C,
|
|
193
|
+
D,
|
|
194
|
+
G,
|
|
195
|
+
GE,
|
|
196
|
+
K,
|
|
197
|
+
W,
|
|
198
|
+
C_start,
|
|
199
|
+
B_end,
|
|
200
|
+
W_end,
|
|
201
|
+
cost_end,
|
|
202
|
+
grad_end,
|
|
203
|
+
cost_out,
|
|
204
|
+
gamma,
|
|
205
|
+
min_func_id,
|
|
206
|
+
list_N,
|
|
207
|
+
list_M,
|
|
208
|
+
step_sizes,
|
|
209
|
+
num_end_conditions,
|
|
210
|
+
candidate_costs_scratch,
|
|
211
|
+
sparsemin_scratch,
|
|
212
|
+
):
|
|
213
|
+
batch_size = C.shape[0]
|
|
214
|
+
num_steps = step_sizes.shape[0]
|
|
215
|
+
|
|
216
|
+
for b in prange(batch_size):
|
|
217
|
+
n_limit = int(list_N[b])
|
|
218
|
+
m_limit = int(list_M[b])
|
|
219
|
+
candidate_costs = candidate_costs_scratch[b]
|
|
220
|
+
sparsemin_sort_buffer = sparsemin_scratch[b]
|
|
221
|
+
|
|
222
|
+
for n in range(n_limit):
|
|
223
|
+
for m in range(m_limit):
|
|
224
|
+
for s in range(num_steps):
|
|
225
|
+
predecessor_n = n - int(step_sizes[s, 0])
|
|
226
|
+
predecessor_m = m - int(step_sizes[s, 1])
|
|
227
|
+
if predecessor_n >= 0 and predecessor_m >= 0:
|
|
228
|
+
candidate_costs[s] = (
|
|
229
|
+
W[b, n, m, s] * C[b, n, m]
|
|
230
|
+
+ D[b, predecessor_n, predecessor_m]
|
|
231
|
+
)
|
|
232
|
+
else:
|
|
233
|
+
candidate_costs[s] = math.inf
|
|
234
|
+
candidate_costs[num_steps] = C_start[b, n, m]
|
|
235
|
+
|
|
236
|
+
D[b, n, m] = _minimum(
|
|
237
|
+
candidate_costs,
|
|
238
|
+
gamma,
|
|
239
|
+
min_func_id,
|
|
240
|
+
K[b, n, m],
|
|
241
|
+
sparsemin_sort_buffer,
|
|
242
|
+
)
|
|
243
|
+
|
|
244
|
+
local_gradient = K[b, n, m, num_steps]
|
|
245
|
+
for s in range(num_steps):
|
|
246
|
+
local_gradient += K[b, n, m, s] * W[b, n, m, s]
|
|
247
|
+
G[b, n, m] = local_gradient
|
|
248
|
+
|
|
249
|
+
end_count = int(num_end_conditions[b])
|
|
250
|
+
for i in range(end_count):
|
|
251
|
+
n = int(B_end[b, i, 0])
|
|
252
|
+
m = int(B_end[b, i, 1])
|
|
253
|
+
cost_end[b, i] = D[b, n, m] + W_end[b, i] * C[b, n, m]
|
|
254
|
+
|
|
255
|
+
cost_out[b] = _minimum(
|
|
256
|
+
cost_end[b], gamma, min_func_id, grad_end[b], sparsemin_sort_buffer
|
|
257
|
+
)
|
|
258
|
+
for i in range(end_count):
|
|
259
|
+
n = int(B_end[b, i, 0])
|
|
260
|
+
m = int(B_end[b, i, 1])
|
|
261
|
+
GE[b, n, m] = grad_end[b, i]
|
|
262
|
+
|
|
263
|
+
|
|
264
|
+
@njit(parallel=True, nogil=True, cache=True)
|
|
265
|
+
def compute_dDTW_backward_numba(E, GE, K, list_N, list_M, step_sizes):
|
|
266
|
+
batch_size = E.shape[0]
|
|
267
|
+
num_steps = step_sizes.shape[0]
|
|
268
|
+
|
|
269
|
+
for b in prange(batch_size):
|
|
270
|
+
n_limit = int(list_N[b])
|
|
271
|
+
m_limit = int(list_M[b])
|
|
272
|
+
for n in range(n_limit - 1, -1, -1):
|
|
273
|
+
for m in range(m_limit - 1, -1, -1):
|
|
274
|
+
value = GE[b, n, m]
|
|
275
|
+
for s in range(num_steps):
|
|
276
|
+
successor_n = n + int(step_sizes[s, 0])
|
|
277
|
+
successor_m = m + int(step_sizes[s, 1])
|
|
278
|
+
if successor_n < n_limit and successor_m < m_limit:
|
|
279
|
+
value += (
|
|
280
|
+
E[b, successor_n, successor_m]
|
|
281
|
+
* K[b, successor_n, successor_m, s]
|
|
282
|
+
)
|
|
283
|
+
E[b, n, m] = value
|
|
284
|
+
|
|
285
|
+
|
|
286
|
+
def _numpy_view(tensor):
|
|
287
|
+
if tensor.device.type != "cpu":
|
|
288
|
+
raise ValueError("backend='cpu_numba' requires CPU tensors")
|
|
289
|
+
return tensor.detach().numpy()
|
|
290
|
+
|
|
291
|
+
|
|
292
|
+
def _clear_debug_matrices(cls):
|
|
293
|
+
cls.C_matrix = None
|
|
294
|
+
cls.D_matrix = None
|
|
295
|
+
cls.E_matrix = None
|
|
296
|
+
cls.G_matrix = None
|
|
297
|
+
cls.H_matrix = None
|
|
298
|
+
cls.GE_matrix = None
|
|
299
|
+
cls.C_start_matrix = None
|
|
300
|
+
|
|
301
|
+
|
|
302
|
+
class _backend_CPU_Numba(Function):
|
|
303
|
+
C_matrix = None
|
|
304
|
+
D_matrix = None
|
|
305
|
+
E_matrix = None
|
|
306
|
+
G_matrix = None
|
|
307
|
+
H_matrix = None
|
|
308
|
+
GE_matrix = None
|
|
309
|
+
C_start_matrix = None
|
|
310
|
+
|
|
311
|
+
@staticmethod
|
|
312
|
+
def forward(
|
|
313
|
+
ctx,
|
|
314
|
+
C,
|
|
315
|
+
min_function,
|
|
316
|
+
gamma,
|
|
317
|
+
step_sizes,
|
|
318
|
+
W,
|
|
319
|
+
list_N,
|
|
320
|
+
list_M,
|
|
321
|
+
B_start,
|
|
322
|
+
B_end,
|
|
323
|
+
num_start_conditions,
|
|
324
|
+
num_end_conditions,
|
|
325
|
+
W_start,
|
|
326
|
+
W_end,
|
|
327
|
+
store_debug,
|
|
328
|
+
):
|
|
329
|
+
dtype = C.dtype
|
|
330
|
+
if dtype not in (torch.float32, torch.float64):
|
|
331
|
+
raise TypeError(
|
|
332
|
+
"backend='cpu_numba' supports torch.float32 and torch.float64 tensors; "
|
|
333
|
+
f"got {dtype}"
|
|
334
|
+
)
|
|
335
|
+
batch_size, max_n, max_m = C.shape
|
|
336
|
+
num_directions = step_sizes.shape[0] + 1
|
|
337
|
+
|
|
338
|
+
C_start = torch.full_like(C, torch.inf)
|
|
339
|
+
for b in range(B_start.shape[0]):
|
|
340
|
+
for i in range(int(num_start_conditions[b])):
|
|
341
|
+
n = int(B_start[b, i, 0])
|
|
342
|
+
m = int(B_start[b, i, 1])
|
|
343
|
+
if n < list_N[b] and m < list_M[b]:
|
|
344
|
+
C_start[b, n, m] = C[b, n, m] * W_start[b, i]
|
|
345
|
+
|
|
346
|
+
GE = torch.zeros_like(C)
|
|
347
|
+
cost_end = torch.full(
|
|
348
|
+
(batch_size, B_end.shape[1]), torch.inf, device=C.device, dtype=dtype
|
|
349
|
+
)
|
|
350
|
+
grad_end = torch.zeros_like(cost_end)
|
|
351
|
+
D = torch.zeros_like(C)
|
|
352
|
+
G = torch.zeros_like(C)
|
|
353
|
+
K = torch.zeros(
|
|
354
|
+
(batch_size, max_n, max_m, num_directions),
|
|
355
|
+
device=C.device,
|
|
356
|
+
dtype=dtype,
|
|
357
|
+
)
|
|
358
|
+
cost_out = torch.zeros(batch_size, device=C.device, dtype=dtype)
|
|
359
|
+
candidate_costs_scratch = torch.empty(
|
|
360
|
+
(batch_size, num_directions), device=C.device, dtype=dtype
|
|
361
|
+
)
|
|
362
|
+
sparsemin_scratch = torch.empty(
|
|
363
|
+
(batch_size, max(num_directions, B_end.shape[1])),
|
|
364
|
+
device=C.device,
|
|
365
|
+
dtype=dtype,
|
|
366
|
+
)
|
|
367
|
+
|
|
368
|
+
with torch.no_grad():
|
|
369
|
+
compute_dDTW_forward_numba(
|
|
370
|
+
_numpy_view(C),
|
|
371
|
+
_numpy_view(D),
|
|
372
|
+
_numpy_view(G),
|
|
373
|
+
_numpy_view(GE),
|
|
374
|
+
_numpy_view(K),
|
|
375
|
+
_numpy_view(W),
|
|
376
|
+
_numpy_view(C_start),
|
|
377
|
+
_numpy_view(B_end),
|
|
378
|
+
_numpy_view(W_end),
|
|
379
|
+
_numpy_view(cost_end),
|
|
380
|
+
_numpy_view(grad_end),
|
|
381
|
+
_numpy_view(cost_out),
|
|
382
|
+
float(gamma.item()),
|
|
383
|
+
int(min_function),
|
|
384
|
+
_numpy_view(list_N),
|
|
385
|
+
_numpy_view(list_M),
|
|
386
|
+
_numpy_view(step_sizes),
|
|
387
|
+
_numpy_view(num_end_conditions),
|
|
388
|
+
_numpy_view(candidate_costs_scratch),
|
|
389
|
+
_numpy_view(sparsemin_scratch),
|
|
390
|
+
)
|
|
391
|
+
|
|
392
|
+
ctx.list_N = list_N
|
|
393
|
+
ctx.list_M = list_M
|
|
394
|
+
ctx.step_sizes = step_sizes
|
|
395
|
+
ctx.store_debug = store_debug
|
|
396
|
+
ctx.save_for_backward(G.detach(), K.detach(), GE.detach())
|
|
397
|
+
|
|
398
|
+
if store_debug:
|
|
399
|
+
_backend_CPU_Numba.C_matrix = C.detach()
|
|
400
|
+
_backend_CPU_Numba.D_matrix = D.detach()
|
|
401
|
+
_backend_CPU_Numba.G_matrix = G.detach()
|
|
402
|
+
_backend_CPU_Numba.GE_matrix = GE.detach()
|
|
403
|
+
_backend_CPU_Numba.C_start_matrix = C_start.detach()
|
|
404
|
+
_backend_CPU_Numba.E_matrix = None
|
|
405
|
+
_backend_CPU_Numba.H_matrix = None
|
|
406
|
+
else:
|
|
407
|
+
_clear_debug_matrices(_backend_CPU_Numba)
|
|
408
|
+
|
|
409
|
+
return cost_out
|
|
410
|
+
|
|
411
|
+
@staticmethod
|
|
412
|
+
def backward(ctx, grad_output):
|
|
413
|
+
G, K, GE = ctx.saved_tensors
|
|
414
|
+
E = torch.zeros_like(G)
|
|
415
|
+
|
|
416
|
+
with torch.no_grad():
|
|
417
|
+
compute_dDTW_backward_numba(
|
|
418
|
+
_numpy_view(E),
|
|
419
|
+
_numpy_view(GE),
|
|
420
|
+
_numpy_view(K),
|
|
421
|
+
_numpy_view(ctx.list_N),
|
|
422
|
+
_numpy_view(ctx.list_M),
|
|
423
|
+
_numpy_view(ctx.step_sizes),
|
|
424
|
+
)
|
|
425
|
+
|
|
426
|
+
grad_C = grad_output.view(-1, 1, 1) * E * G
|
|
427
|
+
|
|
428
|
+
if ctx.store_debug:
|
|
429
|
+
_backend_CPU_Numba.E_matrix = E.detach()
|
|
430
|
+
_backend_CPU_Numba.H_matrix = (E * G).detach()
|
|
431
|
+
|
|
432
|
+
return (
|
|
433
|
+
grad_C,
|
|
434
|
+
None,
|
|
435
|
+
None,
|
|
436
|
+
None,
|
|
437
|
+
None,
|
|
438
|
+
None,
|
|
439
|
+
None,
|
|
440
|
+
None,
|
|
441
|
+
None,
|
|
442
|
+
None,
|
|
443
|
+
None,
|
|
444
|
+
None,
|
|
445
|
+
None,
|
|
446
|
+
None,
|
|
447
|
+
)
|