@@ -97,21 +97,35 @@ __global__ void dtw_wavefront_kernel(
9797 // outside the loop, to avoid issues with __syncthreads convergence.
9898 __shared__ int s_pid;
9999
100- // Set up anti-diagonal buffer pointers (constant across pairs)
100+ // Two modes: 3-buffer (classic) for L<=1024, 2-buffer (double-buffer) for L>1024.
101+ // The 2-buffer mode saves max_L*sizeof(T) shared memory, improving occupancy
102+ // for long series at the cost of an extra sync + register pressure per anti-diag.
103+ // For medium series the 3-buffer mode is faster (no extra sync overhead).
104+ constexpr int DOUBLE_BUF_THRESHOLD = 1024 ;
105+ const bool use_double_buf = (max_L > DOUBLE_BUF_THRESHOLD ) && !preload;
106+
101107 T *s_row_buf = nullptr ;
102108 T *s_col_buf = nullptr ;
103- T *diag[3 ];
109+ T *diag_buf[3 ]; // [0],[1] always used; [2] only in 3-buffer mode
110+ int n_diag_bufs;
104111
105112 if (preload) {
106- s_row_buf = smem; // [max_L]
107- s_col_buf = smem + max_L; // [max_L]
108- diag[0 ] = smem + 2 * max_L;
109- diag[1 ] = smem + 3 * max_L;
110- diag[2 ] = smem + 4 * max_L;
113+ s_row_buf = smem;
114+ s_col_buf = smem + max_L;
115+ diag_buf[0 ] = smem + 2 * max_L;
116+ diag_buf[1 ] = smem + 3 * max_L;
117+ diag_buf[2 ] = smem + 4 * max_L; // 3-buffer always for preload
118+ n_diag_bufs = 3 ;
119+ } else if (use_double_buf) {
120+ diag_buf[0 ] = smem;
121+ diag_buf[1 ] = smem + max_L;
122+ diag_buf[2 ] = nullptr ; // not used
123+ n_diag_bufs = 2 ;
111124 } else {
112- diag[0 ] = smem;
113- diag[1 ] = smem + max_L;
114- diag[2 ] = smem + 2 * max_L;
125+ diag_buf[0 ] = smem;
126+ diag_buf[1 ] = smem + max_L;
127+ diag_buf[2 ] = smem + 2 * max_L;
128+ n_diag_bufs = 3 ;
115129 }
116130
117131 // ---------------------------------------------------------------------------
@@ -185,60 +199,92 @@ __global__ void dtw_wavefront_kernel(
185199
186200 const int total_diags = M + N_len - 1 ;
187201
188- for (int k = 0 ; k < total_diags; ++k) {
189- // Anti-diagonal k: cells (i, j) where i + j = k
190- const int i_min = max (0 , k - N_len + 1 );
191- const int i_max = min (k, M - 1 );
192- const int len_k = i_max - i_min + 1 ;
193-
194- // i_min for the two previous anti-diagonals
195- const int i_min_k1 = max (0 , (k - 1 ) - N_len + 1 );
196- const int i_min_k2 = max (0 , (k - 2 ) - N_len + 1 );
197-
198- // Which buffer slot for k, k-1, k-2
199- T *cur = diag[k % 3 ];
200- T *prev = diag[(k - 1 + 3 ) % 3 ]; // k-1
201- T *prev2 = diag[(k - 2 + 3 ) % 3 ]; // k-2
202-
203- for (int p = tid; p < len_k; p += nthreads) {
204- const int i = i_min + p;
205- const int j = k - i;
206-
207- // Banded check
208- if (use_band) {
209- const double center = slope * i;
210- const int j_low = (int )ceil (round (100.0 * (center - window)) / 100.0 );
211- const int j_high = (int )floor (round (100.0 * (center + window)) / 100.0 );
212- if (j < j_low || j > j_high) {
213- cur[p] = INF ;
214- continue ;
202+ if (use_double_buf) {
203+ // ── Double-buffer mode (L > 1024): 2 ping-pong buffers ───────────
204+ // Pre-fetch cost_diag from k-2 buffer into registers before overwriting.
205+ // Saves max_L*sizeof(T) shared memory → better occupancy for long series.
206+ for (int k = 0 ; k < total_diags; ++k) {
207+ const int i_min = max (0 , k - N_len + 1 );
208+ const int i_max = min (k, M - 1 );
209+ const int len_k = i_max - i_min + 1 ;
210+ const int i_min_k1 = max (0 , (k - 1 ) - N_len + 1 );
211+ const int i_min_k2 = max (0 , (k - 2 ) - N_len + 1 );
212+
213+ T *cur = diag_buf[k & 1 ]; // output for k (also holds k-2)
214+ T *prev = diag_buf[(k & 1 ) ^ 1 ]; // k-1
215+
216+ // Phase 1: cache cost_diag from k-2 (= cur) before overwriting
217+ constexpr int MAX_SI = 8 ;
218+ T cd[MAX_SI ];
219+ for (int p = tid, s = 0 ; p < len_k && s < MAX_SI ; p += nthreads, ++s) {
220+ int i = i_min + p;
221+ cd[s] = (k >= 2 && i > 0 && (k - i) > 0 )
222+ ? cur[(i - 1 ) - i_min_k2] : INF ;
223+ }
224+ __syncthreads ();
225+
226+ // Phase 2: compute anti-diag k
227+ for (int p = tid, s = 0 ; p < len_k && s < MAX_SI ; p += nthreads, ++s) {
228+ int i = i_min + p, j = k - i;
229+ if (use_band) {
230+ double center = slope * i;
231+ int jl = (int )ceil (round (100.0 * (center - window)) / 100.0 );
232+ int jh = (int )floor (round (100.0 * (center + window)) / 100.0 );
233+ if (j < jl || j > jh) { cur[p] = INF ; continue ; }
234+ }
235+ T diff = __ldg (&s_row[i]) - __ldg (&s_col[j]);
236+ T d = use_squared_l2 ? (diff * diff) : fabs (diff);
237+ if (i == 0 && j == 0 ) { cur[p] = d; }
238+ else {
239+ T ca = (i > 0 ) ? prev[(i-1 ) - i_min_k1] : INF ;
240+ T cl = (j > 0 ) ? prev[i - i_min_k1] : INF ;
241+ cur[p] = fmin (cd[s], fmin (ca, cl)) + d;
215242 }
216243 }
217-
218- // Series data: from shared memory if preloaded, else via __ldg from global
219- T diff = preload ? (s_row[i] - s_col[j])
220- : (__ldg (&s_row[i]) - __ldg (&s_col[j]));
221- T d = use_squared_l2 ? (diff * diff) : fabs (diff);
222-
223- if (i == 0 && j == 0 ) {
224- cur[p] = d;
225- } else {
226- // cost[i-1][j] is on anti-diag k-1 at position (i-1) - i_min(k-1)
227- T cost_above = (i > 0 ) ? prev[(i - 1 ) - i_min_k1] : INF ;
228- // cost[i][j-1] is on anti-diag k-1 at position i - i_min(k-1)
229- T cost_left = (j > 0 ) ? prev[i - i_min_k1] : INF ;
230- // cost[i-1][j-1] is on anti-diag k-2 at position (i-1) - i_min(k-2)
231- T cost_diag = (i > 0 && j > 0 ) ? prev2[(i - 1 ) - i_min_k2] : INF ;
232-
233- cur[p] = fmin (cost_diag, fmin (cost_above, cost_left)) + d;
244+ __syncthreads ();
245+ }
246+ } else {
247+ // ── 3-buffer mode (L <= 1024): classic rotating buffers ──────────
248+ // Faster for medium series (no extra sync or register pressure).
249+ for (int k = 0 ; k < total_diags; ++k) {
250+ const int i_min = max (0 , k - N_len + 1 );
251+ const int i_max = min (k, M - 1 );
252+ const int len_k = i_max - i_min + 1 ;
253+ const int i_min_k1 = max (0 , (k - 1 ) - N_len + 1 );
254+ const int i_min_k2 = max (0 , (k - 2 ) - N_len + 1 );
255+
256+ T *cur = diag_buf[k % 3 ];
257+ T *prev = diag_buf[(k - 1 + 3 ) % 3 ];
258+ T *prev2 = diag_buf[(k - 2 + 3 ) % 3 ];
259+
260+ for (int p = tid; p < len_k; p += nthreads) {
261+ int i = i_min + p, j = k - i;
262+ if (use_band) {
263+ double center = slope * i;
264+ int jl = (int )ceil (round (100.0 * (center - window)) / 100.0 );
265+ int jh = (int )floor (round (100.0 * (center + window)) / 100.0 );
266+ if (j < jl || j > jh) { cur[p] = INF ; continue ; }
267+ }
268+ T diff = preload ? (s_row[i] - s_col[j])
269+ : (__ldg (&s_row[i]) - __ldg (&s_col[j]));
270+ T d = use_squared_l2 ? (diff * diff) : fabs (diff);
271+ if (i == 0 && j == 0 ) { cur[p] = d; }
272+ else {
273+ T ca = (i > 0 ) ? prev[(i-1 ) - i_min_k1] : INF ;
274+ T cl = (j > 0 ) ? prev[i - i_min_k1] : INF ;
275+ T cd = (i > 0 && j > 0 ) ? prev2[(i-1 ) - i_min_k2] : INF ;
276+ cur[p] = fmin (cd, fmin (ca, cl)) + d;
277+ }
234278 }
279+ __syncthreads ();
235280 }
236- __syncthreads ();
237281 }
238282
239283 // Result is the last anti-diagonal (single cell: (M-1, N_len-1))
240284 if (tid == 0 ) {
241- distances[pid] = diag[(total_diags - 1 ) % 3 ][0 ];
285+ int last_buf = use_double_buf ? ((total_diags - 1 ) & 1 )
286+ : ((total_diags - 1 ) % 3 );
287+ distances[pid] = diag_buf[last_buf][0 ];
242288 }
243289
244290 // Non-persistent mode: exit after one pair
@@ -496,7 +542,11 @@ std::vector<double> launch_dtw_kernel(
496542 } else {
497543 // Wavefront kernel: shared memory and block size configuration
498544 const bool preload = (max_L <= 256 );
499- const size_t shared_mem = (preload ? 5 : 3 ) * max_L * sizeof (T);
545+ // L<=256: preload mode (2 series + 3 anti-diag buffers = 5)
546+ // 256<L<=1024: 3-buffer mode (3 anti-diag buffers)
547+ // L>1024: double-buffer mode (2 anti-diag buffers, saves occupancy)
548+ const size_t n_bufs = preload ? 5 : (max_L > 1024 ? 2 : 3 );
549+ const size_t shared_mem = n_bufs * max_L * sizeof (T);
500550
501551 // Block size heuristic tuned for the anti-diagonal wavefront pattern.
502552 int block_size;
0 commit comments