diff --git a/src/main/java/com/thealgorithms/streaming/DynamicTimeWarping.java b/src/main/java/com/thealgorithms/streaming/DynamicTimeWarping.java new file mode 100644 index 000000000000..52e188cee2c5 --- /dev/null +++ b/src/main/java/com/thealgorithms/streaming/DynamicTimeWarping.java @@ -0,0 +1,190 @@ +package com.thealgorithms.streaming; + +import java.util.Arrays; + +/** + * Dynamic time warping: how far apart two series are once one of them is allowed to be + * stretched and squeezed in time. + * + *

The Euclidean distance compares sample {@code i} with sample {@code i} and nothing else, so two + * recordings of the same gesture, one performed slightly faster, come out as far apart as two + * unrelated ones. Dynamic time warping instead looks for the cheapest way to line the two series up: + * every point of the first has to be matched to at least one point of the second and the other way + * round, the matching may never go backwards, and the cost is the sum over the matched pairs. That + * alignment is a shortest path through a grid, and the dynamic program is the obvious one: + * + *

+ * D[i][j] = |a[i] - b[j]| + min( D[i-1][j], D[i][j-1], D[i-1][j-1] )
+ * 
+ * + *

The three predecessors are exactly the three legal moves: consume a point of the first series, + * of the second, or of both. The answer is the bottom right corner. + * + *

Left alone, the alignment may match one point of a series against an arbitrarily long stretch of + * the other, which is rarely meaningful and costs {@code O(n * m)} regardless. The Sakoe-Chiba band + * forbids matches further apart in time than a given width, which both rules out those degenerate + * alignments and narrows the grid that has to be filled. The band has to be at least the difference + * in length, or no alignment exists at all. + * + *

Note that the result is not a metric: it does not satisfy the triangle inequality, so it can be + * used to rank candidates but not to index them without further care. + * + *

Usage

+ * + *
{@code
+ * double distance = DynamicTimeWarping.distance(query, candidate);
+ * double banded = DynamicTimeWarping.distance(query, candidate, 10);
+ * int[][] alignment = DynamicTimeWarping.path(query, candidate);
+ * }
+ * + *

The distance costs O(n * m) time and O(min(n, m)) memory; the path costs O(n * m) of both, + * because it has to remember the grid. + * + * @see Dynamic time warping + */ +public final class DynamicTimeWarping { + + private DynamicTimeWarping() { + } + + /** + * Returns the warping distance between two series. + * + * @param first the first series, left untouched + * @param second the second series, left untouched + * @return the cost of the cheapest alignment + * @throws IllegalArgumentException if a series is empty or holds a non-finite value + * @throws NullPointerException if a series is {@code null} + */ + public static double distance(double[] first, double[] second) { + return distance(first, second, Math.max(first.length, second.length)); + } + + /** + * Returns the warping distance between two series, with the alignment confined to a Sakoe-Chiba + * band. + * + * @param first the first series, left untouched + * @param second the second series, left untouched + * @param band how far apart in time two matched points may be, at least the difference in length + * @return the cost of the cheapest alignment inside the band + * @throws IllegalArgumentException if a series is empty or holds a non-finite value, or if the + * band is too narrow for any alignment to exist + * @throws NullPointerException if a series is {@code null} + */ + public static double distance(double[] first, double[] second, int band) { + requireSeries(first, "first"); + requireSeries(second, "second"); + requireBand(band, first.length, second.length); + + double[] previous = new double[second.length + 1]; + double[] current = new double[second.length + 1]; + Arrays.fill(previous, Double.POSITIVE_INFINITY); + previous[0] = 0.0; + + for (int i = 1; i <= first.length; i++) { + Arrays.fill(current, Double.POSITIVE_INFINITY); + int from = Math.max(1, i - band); + int to = Math.min(second.length, i + band); + for (int j = from; j <= to; j++) { + double cost = Math.abs(first[i - 1] - second[j - 1]); + double best = Math.min(previous[j], Math.min(current[j - 1], previous[j - 1])); + current[j] = cost + best; + } + double[] swap = previous; + previous = current; + current = swap; + } + return previous[second.length]; + } + + /** + * Returns the cheapest alignment itself. + * + * @param first the first series, left untouched + * @param second the second series, left untouched + * @return the matched pairs of indices, from {@code (0, 0)} to the two last indices + * @throws IllegalArgumentException if a series is empty or holds a non-finite value + * @throws NullPointerException if a series is {@code null} + */ + public static int[][] path(double[] first, double[] second) { + return path(first, second, Math.max(first.length, second.length)); + } + + /** + * Returns the cheapest alignment inside a Sakoe-Chiba band. + * + * @param first the first series, left untouched + * @param second the second series, left untouched + * @param band how far apart in time two matched points may be, at least the difference in length + * @return the matched pairs of indices, from {@code (0, 0)} to the two last indices + * @throws IllegalArgumentException if a series is empty or holds a non-finite value, or if the + * band is too narrow for any alignment to exist + * @throws NullPointerException if a series is {@code null} + */ + public static int[][] path(double[] first, double[] second, int band) { + requireSeries(first, "first"); + requireSeries(second, "second"); + requireBand(band, first.length, second.length); + + double[][] grid = new double[first.length + 1][second.length + 1]; + for (double[] row : grid) { + Arrays.fill(row, Double.POSITIVE_INFINITY); + } + grid[0][0] = 0.0; + + for (int i = 1; i <= first.length; i++) { + int from = Math.max(1, i - band); + int to = Math.min(second.length, i + band); + for (int j = from; j <= to; j++) { + double cost = Math.abs(first[i - 1] - second[j - 1]); + grid[i][j] = cost + Math.min(grid[i - 1][j], Math.min(grid[i][j - 1], grid[i - 1][j - 1])); + } + } + + int steps = 0; + int row = first.length; + int column = second.length; + int[][] reversed = new int[first.length + second.length][2]; + while (row > 0 && column > 0) { + reversed[steps][0] = row - 1; + reversed[steps][1] = column - 1; + steps++; + double diagonal = grid[row - 1][column - 1]; + double above = grid[row - 1][column]; + double left = grid[row][column - 1]; + if (diagonal <= above && diagonal <= left) { + row--; + column--; + } else if (above <= left) { + row--; + } else { + column--; + } + } + + int[][] alignment = new int[steps][2]; + for (int i = 0; i < steps; i++) { + alignment[i] = reversed[steps - 1 - i]; + } + return alignment; + } + + private static void requireSeries(double[] series, String name) { + if (series.length == 0) { + throw new IllegalArgumentException("The " + name + " series must not be empty"); + } + for (double value : series) { + if (!Double.isFinite(value)) { + throw new IllegalArgumentException("Samples must be finite, but the " + name + " series held " + value); + } + } + } + + private static void requireBand(int band, int firstLength, int secondLength) { + int minimum = Math.abs(firstLength - secondLength); + if (band < minimum) { + throw new IllegalArgumentException("The band must be at least the difference in length, " + minimum + ", but was " + band); + } + } +} diff --git a/src/test/java/com/thealgorithms/streaming/DynamicTimeWarpingTest.java b/src/test/java/com/thealgorithms/streaming/DynamicTimeWarpingTest.java new file mode 100644 index 000000000000..5be3b21eb101 --- /dev/null +++ b/src/test/java/com/thealgorithms/streaming/DynamicTimeWarpingTest.java @@ -0,0 +1,221 @@ +package com.thealgorithms.streaming; + +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertThrows; +import static org.junit.jupiter.api.Assertions.assertTrue; + +import java.util.Random; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; +import org.junit.jupiter.params.ParameterizedTest; +import org.junit.jupiter.params.provider.ValueSource; + +class DynamicTimeWarpingTest { + + private static double manhattan(double[] first, double[] second) { + double sum = 0.0; + for (int i = 0; i < first.length; i++) { + sum += Math.abs(first[i] - second[i]); + } + return sum; + } + + private static double[] sine(int length, double period, double phase) { + double[] series = new double[length]; + for (int i = 0; i < length; i++) { + series[i] = Math.sin(2 * Math.PI * i / period + phase); + } + return series; + } + + @Test + void rejectsEmptySeries() { + assertThrows(IllegalArgumentException.class, () -> DynamicTimeWarping.distance(new double[0], new double[] {1.0})); + assertThrows(IllegalArgumentException.class, () -> DynamicTimeWarping.distance(new double[] {1.0}, new double[0])); + assertThrows(IllegalArgumentException.class, () -> DynamicTimeWarping.path(new double[0], new double[] {1.0})); + } + + @ParameterizedTest + @ValueSource(doubles = {Double.NaN, Double.POSITIVE_INFINITY, Double.NEGATIVE_INFINITY}) + void rejectsNonFiniteSamples(double value) { + double[] good = {1.0, 2.0, 3.0}; + double[] bad = {1.0, value, 3.0}; + + assertThrows(IllegalArgumentException.class, () -> DynamicTimeWarping.distance(bad, good)); + assertThrows(IllegalArgumentException.class, () -> DynamicTimeWarping.distance(good, bad)); + } + + @Test + @DisplayName("a band narrower than the difference in length admits no alignment") + void rejectsABandThatIsTooNarrow() { + double[] shorter = {1.0, 2.0, 3.0}; + double[] longer = {1.0, 2.0, 3.0, 4.0, 5.0, 6.0}; + + assertThrows(IllegalArgumentException.class, () -> DynamicTimeWarping.distance(shorter, longer, 2)); + assertEquals(6.0, DynamicTimeWarping.distance(shorter, longer, 3), 1e-12, "the last point of the short series has to carry 3, 4, 5 and 6"); + } + + @Test + void aSeriesIsAtNoDistanceFromItself() { + double[] series = {1.0, 4.0, 2.0, 8.0, 3.0}; + + assertEquals(0.0, DynamicTimeWarping.distance(series, series)); + } + + @Test + @DisplayName("two flat series a fixed distance apart cost that distance once per sample") + void measuresAConstantOffset() { + double[] first = {1.0, 1.0, 1.0, 1.0}; + double[] second = {3.0, 3.0, 3.0, 3.0}; + + assertEquals(8.0, DynamicTimeWarping.distance(first, second), 1e-12, "the diagonal is the shortest path and every cell costs 2"); + } + + @Test + @DisplayName("a series shifted in value is cheaper than sample by sample, because warping reuses points") + void warpingBeatsTheStraightComparisonOnARamp() { + double[] first = {1.0, 2.0, 3.0, 4.0}; + double[] second = {3.0, 4.0, 5.0, 6.0}; + + assertEquals(6.0, DynamicTimeWarping.distance(first, second), 1e-12); + assertEquals(8.0, manhattan(first, second), 1e-12); + } + + @Test + void isSymmetric() { + Random random = new Random(3L); + for (int trial = 0; trial < 20; trial++) { + double[] first = new double[10 + random.nextInt(20)]; + double[] second = new double[10 + random.nextInt(20)]; + for (int i = 0; i < first.length; i++) { + first[i] = random.nextGaussian(); + } + for (int i = 0; i < second.length; i++) { + second[i] = random.nextGaussian(); + } + + assertEquals(DynamicTimeWarping.distance(first, second), DynamicTimeWarping.distance(second, first), 1e-9); + } + } + + @Test + @DisplayName("warping never costs more than matching sample by sample") + void neverExceedsTheStraightComparison() { + Random random = new Random(11L); + for (int trial = 0; trial < 50; trial++) { + double[] first = new double[30]; + double[] second = new double[30]; + for (int i = 0; i < first.length; i++) { + first[i] = random.nextGaussian(); + second[i] = random.nextGaussian(); + } + + assertTrue(DynamicTimeWarping.distance(first, second) <= manhattan(first, second) + 1e-9); + } + } + + @Test + @DisplayName("a band of zero forces the straight comparison") + void aBandOfZeroIsTheStraightComparison() { + double[] first = {1.0, 5.0, 2.0, 8.0}; + double[] second = {2.0, 4.0, 4.0, 7.0}; + + assertEquals(manhattan(first, second), DynamicTimeWarping.distance(first, second, 0), 1e-12); + } + + @Test + @DisplayName("a narrower band can only cost more") + void aNarrowerBandCostsAtLeastAsMuch() { + double[] first = sine(60, 12.0, 0.0); + double[] second = sine(60, 12.0, 0.9); + + double free = DynamicTimeWarping.distance(first, second); + double banded = DynamicTimeWarping.distance(first, second, 3); + double tight = DynamicTimeWarping.distance(first, second, 1); + + assertTrue(banded >= free - 1e-9, "banded " + banded + " should not be below free " + free); + assertTrue(tight >= banded - 1e-9, "tight " + tight + " should not be below banded " + banded); + } + + @Test + @DisplayName("a shift in time costs almost nothing, where a straight comparison is fooled") + void absorbsAShiftInTime() { + double[] first = sine(60, 12.0, 0.0); + double[] second = sine(60, 12.0, Math.PI / 3); + + double warping = DynamicTimeWarping.distance(first, second); + double straight = manhattan(first, second); + + assertTrue(warping < 0.25 * straight, "warping " + warping + " against straight " + straight); + } + + @Test + @DisplayName("a series stretched in time still matches the original") + void absorbsAStretch() { + double[] original = sine(40, 10.0, 0.0); + double[] stretched = new double[80]; + for (int i = 0; i < stretched.length; i++) { + stretched[i] = original[i / 2]; + } + + double warping = DynamicTimeWarping.distance(original, stretched); + + assertTrue(warping < 1.0, "a stretched copy should be close, but was " + warping); + } + + @Test + @DisplayName("the path runs from corner to corner without ever going backwards") + void thePathIsMonotoneAndComplete() { + double[] first = sine(30, 8.0, 0.0); + double[] second = sine(45, 12.0, 0.4); + + int[][] alignment = DynamicTimeWarping.path(first, second); + + assertEquals(0, alignment[0][0]); + assertEquals(0, alignment[0][1]); + assertEquals(first.length - 1, alignment[alignment.length - 1][0]); + assertEquals(second.length - 1, alignment[alignment.length - 1][1]); + for (int step = 1; step < alignment.length; step++) { + int rowStep = alignment[step][0] - alignment[step - 1][0]; + int columnStep = alignment[step][1] - alignment[step - 1][1]; + assertTrue(rowStep >= 0 && rowStep <= 1, "the path stepped " + rowStep + " rows"); + assertTrue(columnStep >= 0 && columnStep <= 1, "the path stepped " + columnStep + " columns"); + assertTrue(rowStep + columnStep > 0, "the path stood still"); + } + } + + @Test + @DisplayName("the cost of the path is the distance") + void thePathCostsWhatTheDistanceSays() { + double[] first = sine(25, 7.0, 0.0); + double[] second = sine(33, 9.0, 0.2); + + int[][] alignment = DynamicTimeWarping.path(first, second); + double cost = 0.0; + for (int[] pair : alignment) { + cost += Math.abs(first[pair[0]] - second[pair[1]]); + } + + assertEquals(DynamicTimeWarping.distance(first, second), cost, 1e-9); + } + + @Test + void handlesSeriesOfOneSample() { + assertEquals(3.0, DynamicTimeWarping.distance(new double[] {1.0}, new double[] {4.0}), 1e-12); + assertEquals(1, DynamicTimeWarping.path(new double[] {1.0}, new double[] {4.0}).length); + } + + @Test + void leavesTheSeriesUntouched() { + double[] first = {1.0, 2.0, 3.0}; + double[] second = {4.0, 5.0}; + double[] firstCopy = first.clone(); + double[] secondCopy = second.clone(); + + DynamicTimeWarping.distance(first, second); + DynamicTimeWarping.path(first, second); + + org.junit.jupiter.api.Assertions.assertArrayEquals(firstCopy, first); + org.junit.jupiter.api.Assertions.assertArrayEquals(secondCopy, second); + } +}