跳过正文
  1. Posts/

ska::flat_hash_map 源码分析

·2416 字·5 分钟

背景
#

最近在调研各种 hashmap,发现 ska::flat_hash_map 性能优秀,于是来看看代码。 最大的特点是,它使用了带 probe count 上限的 robin hood hashing。

相关概念
#

Distance_from_desired
#

对于采用了 open addressing 的 hash 实现,当插入发生冲突时,会以一定方式(如线性探测、平方探测等)探测下一个可以插入的 slot。 因而实际插入的 slot 位置与理想的 slot 位置通常不相同,这段距离定义为 distance_from_desired。 在没有冲突的理想情况下,所有 distance_from_desired 的值应该都为 0。 distance_from_desired 的一种更常见的说法叫做 probe sequence lengths(PSL)。

robin hood hashing
#

robin hood hashing 的核心思想是“劫富济贫”: distance_from_desired 小的 slot 被认为更“富有”,distance_from_desired 大的 slot 被认为更“贫穷”。 具体来说,当插入一个新的元素时,如果当前位置元素的 distance_from_desired 小于待插入元素的 distance_from_desired,那么就将待插入元素放入当前位置,把当前位置的元素取出,继续寻找新的位置。

这样做使得所有元素的 distance_from_desired 分布更为平均,variance 更小。 这样的分布对 cache 更友好(几乎全部元素的 distance_from_desired 都小于一个 cache line 的长度,因此在 find 的时候只需要 fetch 一次 cache line),从而拥有更好的性能。

一般的 robin hood hashing 在 find 时,会用一个全局最大的 distance_from_desired 作为没有找到该元素的终止条件。 一种常见的改进是,不维护全局最大 distance_from_desired,而是在看到当前位置元素的 distance_from_desired 比要插入元素的 distance_from_desired 小时终止。

以插入 C 为例,探测与交换的具体过程如下:

robin hood hashing 的探测与劫富济贫过程
 1
 2  iterator find(const FindKey& key) {
 3    size_t index =
 4        hash_policy.index_for_hash(hash_object(key), num_slots_minus_one);
 5    EntryPointer it = entries + ptrdiff_t(index);
 6    for (int8_t distance = 0; it->distance_from_desired >= distance;
 7         ++distance, ++it) {
 8      if (compares_equal(key, it->value)) return {it};
 9    }
10    return end();
11  }

带上限的 robin hood hashing
#

一般的 robin hood hashing 在 insert 时,会不断寻找(包括可能的 swap 过程),直到找到一个空的 slot 为止。该过程在 hash table 较满时可能接近线性的时间复杂度。 ska::flat_hash_map 对这一点的改进是,限制了 insert 时尝试的上限次数,作者给出的经验值为 log(N),其中 N 为 slots 的个数。 这样保证每个 slot 的最大 distance_from_desired 不会超过 log(N)。

关键实现
#

emplace
#

插入一个元素,分析见注释。 其中 emplace 函数主要负责查找是否已经存在该元素,以及调整到合适的插入位置;emplace_new_key 函数执行真正的 emplace 操作。

详细代码
 1  template <typename Key, typename... Args>
 2  std::pair<iterator, bool> emplace(Key&& key, Args&&... args) {
 3    size_t index =
 4        hash_policy.index_for_hash(hash_object(key), num_slots_minus_one);
 5    EntryPointer current_entry = entries + ptrdiff_t(index);
 6    int8_t distance_from_desired = 0;
 7    // 插入前先查找是否存在
 8    // 只需要查找有限的距离
 9
10    // trick在于,初始 current_entry->distance_from_desired为-1
11    // 此时不会进入for loop,直接进行emplace_new_key。
12    // 该for loop有两层意义: 1.在index位置不为空时找到合适的位置: 空的slot或者更富有的slot(也就是current_entry->distance_from_desired < distance_from_desired的slot)
13    // 2.在该过程中找一下是否已经插入了该值
14    for (; current_entry->distance_from_desired >= distance_from_desired;
15         ++current_entry, ++distance_from_desired) {
16      if (compares_equal(key, current_entry->value))
17        return {{current_entry}, false};
18    }
19    return emplace_new_key(distance_from_desired, current_entry,
20                           std::forward<Key>(key), std::forward<Args>(args)...);
21  }
22
23
24
25
26  template <typename Key, typename... Args>
27  SKA_NOINLINE(std::pair<iterator, bool>)
28  emplace_new_key(int8_t distance_from_desired, EntryPointer current_entry,
29                  Key&& key, Args&&... args) {
30    using std::swap;
31    // num_slots_minus_one初始值为0,表示第一次进行插入,需要先进行grow,很合理。
32    // 如果得到max_load_factor,或者查找次数达到max_lookups,就进行rehash
33    if (num_slots_minus_one == 0 || distance_from_desired == max_lookups ||
34        num_elements + 1 >
35            (num_slots_minus_one + 1) * static_cast<double>(_max_load_factor)) {
36      grow();
37      return emplace(std::forward<Key>(key), std::forward<Args>(args)...);
38    } else if (current_entry->is_empty()) {
39      current_entry->emplace(distance_from_desired, std::forward<Key>(key),
40                             std::forward<Args>(args)...);
41      ++num_elements;
42      return {{current_entry}, true};
43    }
44
45    // 执行到这里,说明有更富有的slot。于是进行swap,转而为被换出的pair<key,value>找一个新的slot
46    // to_insert是当前要插入的,由于swap的发生,可能并不是最初要插入的那一对值
47    value_type to_insert(std::forward<Key>(key), std::forward<Args>(args)...);
48    swap(distance_from_desired, current_entry->distance_from_desired);
49    swap(to_insert, current_entry->value);
50    iterator result = {current_entry};
51    for (++distance_from_desired, ++current_entry;; ++current_entry) {
52      if (current_entry->is_empty()) {
53        // 如果被换过的slot后面某个slot是空的,就直接放置了
54        current_entry->emplace(distance_from_desired, std::move(to_insert));
55        ++num_elements;
56        return {result, true};
57      } else if (current_entry->distance_from_desired < distance_from_desired) {
58        // 在找新slot的过程中,仍然进行劫富济贫的操作, 转而为被换出的pair<key,value>找一个新的slot
59        swap(distance_from_desired, current_entry->distance_from_desired);
60        swap(to_insert, current_entry->value);
61        ++distance_from_desired;
62      } else {
63        // 如果没有空的slot,也没有更富有的slot,那就只能继续往前寻找了,直到达到上限
64        ++distance_from_desired;
65        if (distance_from_desired == max_lookups) {
66          // 如果找了max_lookups个位置还没找到,就进行rehash
67          swap(to_insert, result.current->value);
68          grow();
69          return emplace(std::move(to_insert));
70        }
71      }
72    }
73  }

rehash
#

 1  void rehash(size_t num_buckets) {
 2    num_buckets = std::max(
 3        num_buckets,
 4        static_cast<size_t>(
 5            std::ceil(num_elements / static_cast<double>(_max_load_factor))));
 6    if (num_buckets == 0) {
 7      reset_to_empty_state();
 8      return;
 9    }
10    auto new_prime_index = hash_policy.next_size_over(num_buckets);
11    if (num_buckets == bucket_count()) return;
12    int8_t new_max_lookups = compute_max_lookups(num_buckets);
13    // 额外分配了max_lookups个slots,避免了find时bound checking的开销
14    EntryPointer new_buckets(
15        AllocatorTraits::allocate(*this, num_buckets + new_max_lookups));
16    EntryPointer special_end_item =
17        new_buckets + static_cast<ptrdiff_t>(num_buckets + new_max_lookups - 1);
18    for (EntryPointer it = new_buckets; it != special_end_item; ++it)
19      it->distance_from_desired = -1;
20    special_end_item->distance_from_desired = Entry::special_end_value;
21    std::swap(entries, new_buckets);
22    std::swap(num_slots_minus_one, num_buckets);
23    --num_slots_minus_one;
24    hash_policy.commit(new_prime_index);
25    int8_t old_max_lookups = max_lookups;
26    max_lookups = new_max_lookups;
27    num_elements = 0;
28    // new_buckets其实是旧的entries
29    // num_buckets其实也是旧的值,因为已经被swap了
30    for (EntryPointer
31             it = new_buckets,
32             end = it + static_cast<ptrdiff_t>(num_buckets + old_max_lookups);
33         it != end; ++it) {
34      if (it->has_value()) {
35        emplace(std::move(it->value));
36        it->destroy_value();
37      }
38    }
39    deallocate_data(new_buckets, num_buckets, old_max_lookups);
40  }

一些其他 trick
#

通过多分配 log(N) 个 slot 消除 bound checking 的开销
#

代码见上面的 rehash 部分。 由于每个元素的最大 distance_from_desired 不会超过 log(N),因此可以保证查找时不需要做 bound checking,使得 find 部分的实现非常简洁。

使用素数个 slot 而不是 2 的整数次幂个 slot
#

2 的整数次幂个 slot 是一种很常见的实现。这种实现的主要好处是在将 hash 转换为 index 时,避免了代价高昂的取模操作,而是用代价很小的按位与(&)替代。

但是使用 2 的整数次幂个 slot 的缺点是,取模后得到的结果较少,比起使用素数个 slot 更容易发生冲突。 ska::flat_hash_map 的作者借鉴了 boost::multi_index 中的做法,将变量展开为 compile time const(对 compile time const 做取模运算要远远快于对变量做取模运算),从而减小了这部分开销的影响。

 1
 2struct prime_number_hash_policy {
 3  static size_t mod0(size_t) {
 4    return 0llu;
 5  }
 6  static size_t mod2(size_t hash) {
 7    return hash % 2llu;
 8  }
 9  static size_t mod3(size_t hash) {
10    return hash % 3llu;
11  }
12  static size_t mod5(size_t hash) {
13    return hash % 5llu;
14  }
15  static size_t mod7(size_t hash) {
16    return hash % 7llu;
17  }
18  static size_t mod11(size_t hash) {
19    return hash % 11llu;
20  }
21  static size_t mod13(size_t hash) {
22    return hash % 13llu;
23  }
24  static size_t mod17(size_t hash) {
25    return hash % 17llu;
26  }
27  static size_t mod23(size_t hash) {
28    return hash % 23llu;
29  }
30  static size_t mod29(size_t hash) {
31    return hash % 29llu;
32  }
33  static size_t mod37(size_t hash) {
34    return hash % 37llu;
35  }
36  static size_t mod47(size_t hash) {
37    return hash % 47llu;
38  }
39  ...

相关文章

tensorRT 模型兼容性说明

·525 字·2 分钟
名词说明 # CUDA. 一般来说指的是CUDA SDK. 目前经常使用的是CUDA 8.0和CUDA 10.1两个版本. 8.0和10.1都是SDK的版本号. CUDNN. The NVIDIA CUDA® Deep Neural Network library (cuDNN). 是一个可以为神经网络提供GPU加速的库 compute capability. 是GPU的固有参数,可以理解为GPU的版本.越新的显卡该数值往往越高. tensorRT.NVIDIA TensorRT™ is an SDK for high-performance deep learning inference. 是一个深度学习推理库,旨在提供高性能的推理速度. plan file,也称为 engine plan. 是生成的tensorRT 模型文件. 兼容性说明 # Engine plan 的兼容性依赖于GPU的compute capability 和 TensorRT 版本, 不依赖于CUDA和CUDNN版本.