ward-clustering / benchmarks /benchmark.py
tomaarsen's picture
tomaarsen HF Staff
Add Ward clustering source, benchmark media, and licenses
5650845 verified
Raw History Blame Contribute Delete
2.65 kB
"""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