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