File size: 2,653 Bytes
5650845 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 | """Batched CUDA Ward linkage and maxclust on precomputed float32 distances."""
import os
os.environ["CUDA_LAUNCH_BLOCKING"] = "0"
import numpy as np
import torch
from kernels.benchmark import Benchmark
from scipy.cluster.hierarchy import fcluster, is_valid_linkage, linkage
from scipy.spatial.distance import squareform
class WardBenchmark(Benchmark):
seed = 42
def _setup(self, batch, tokens):
torch.set_num_threads(1)
torch.backends.cuda.matmul.allow_tf32 = False
embeddings = torch.nn.functional.normalize(torch.randn(batch, tokens, 128, device=self.device), dim=-1)
self.distances = (1 - embeddings @ embeddings.transpose(1, 2)).clamp(0, 2)
self.condensed = [squareform(matrix.cpu().numpy(), checks=False) for matrix in self.distances]
self.clusters = tokens // 2
trees, labels, counts = self.kernel.ward(self.distances, self.clusters)
for tree, actual_labels, count in zip(trees.cpu().numpy(), labels.cpu().numpy(), counts.cpu().numpy()):
assert is_valid_linkage(tree)
expected_labels = fcluster(tree, self.clusters, criterion="maxclust") - 1
np.testing.assert_array_equal(actual_labels, expected_labels)
assert count == len(np.unique(expected_labels))
def _run(self):
trees, self.labels, self.counts = self.kernel.ward(self.distances, self.clusters)
# Float32 ties can change merge IDs. Compare merge heights with SciPy.
self.out = trees[:, :, 2]
def _reference(self):
heights = []
for distances in self.condensed:
tree = linkage(distances, method="ward")
fcluster(tree, self.clusters, criterion="maxclust")
heights.append(tree[:, 2])
return torch.as_tensor(np.stack(heights), device=self.device)
def setup_b32_n128(self):
self._setup(32, 128)
def setup_b32_n512(self):
self._setup(32, 512)
def setup_b32_n1024(self):
self._setup(32, 1024)
def setup_b128_n128(self):
self._setup(128, 128)
def setup_b128_n512(self):
self._setup(128, 512)
def setup_b512_n128(self):
self._setup(512, 128)
benchmark_b32_n128 = _run
benchmark_b32_n512 = _run
benchmark_b32_n1024 = _run
benchmark_b128_n128 = _run
benchmark_b128_n512 = _run
benchmark_b512_n128 = _run
verify_b32_n128 = _reference
verify_b32_n512 = _reference
verify_b32_n1024 = _reference
verify_b128_n128 = _reference
verify_b128_n512 = _reference
verify_b512_n128 = _reference
|