在推理引擎中,KV Cache 是一个很重要的组件,当前需要对 KV Cache 中的热点 key 做统计。key 的数量是海量的,如果为每个 key 都分配一个 int 类型来统计,会占用 X G 的内存,这种方式是不可接受的。因此,当前使用概率型数据结构 Count-Min Sketch(简称 CM Sketch)来实现。
Count-Min Sketch(简称 CM Sketch)是一种在大数据流处理中广泛使用的概率型数据结构。它利用极其有限的内存空间,能够近似地统计海量数据流中每个元素的出现频率。它的核心优势在于空间和时间复杂度都是常数级别 O(1),但代价是会存在一定的估算偏差(只会多算,不会少算)。
一、Count-Min Sketch 的算法原理
CM Sketch 的核心思想可以理解为多维度的计数型布隆过滤器(Bloom Filter)。它主要由一个二维计数矩阵和一组相互独立的哈希函数组成。
1. 核心结构
二维数组(矩阵):定义一个宽为 w、高为 d 的二维整数数组 C。初始时,矩阵中所有值均为 0。
哈希函数:挑选 d 个相互独立的哈希函数,记为 h_1, h_2, ..., h_d。每个哈希函数将输入的元素映射到 [0, w-1] 的范围内。
2. 核心操作
① 更新操作(Update)
当数据流中到来一个元素 $x$ 时(或者要给元素 $x$ 增加计数 $c$):
分别用 d 个哈希函数计算出 x 在每一行对应的列索引:col_i = h_i(x)(其中 i 范围是 1 到 d)。
将矩阵中对应位置的计数器加上 c:C[i][col_i] = C[i][col_i] + c。
② 查询操作(Query)
当要查询元素 x 出现的总次数时:
同样用 d 个哈希函数计算出 x 在每一行对应的列:col_i = h_i(x)。
找出这些对应位置中的最小值作为最终估算结果:hat{a}_x = min_{C[i][col_i]}。
3. 为什么取最小值?(误差来源)
因为多个不同的元素可能会经过同一个哈希函数映射到矩阵的同一个位置(发生哈希碰撞),导致该位置的计数器被重复叠加。
任何一个位置的值,都必定大于或等于元素的真实频率。
每一行由于哈希函数不同,碰撞的元素也不同。取所有行之中的最小值,可以最大程度地剔除其他元素碰撞带来的干扰,使其最接近真实值。
4. 滑动窗口衰减
- CMS 的格子值是只增不减的。
- 在长时间运行的系统中,随着总流量越来越大,所有格子最终都会被填满、冲突越来越严重,误差会无限放大。
- 架构解法:必须引入“滑动窗口衰减(Decay)”。例如在 Mooncake 或高性能网关中,每隔一段时间,后台线程会原子的把整个 CMS 矩阵的所有格子值整体右移 1 位(相当于数值除以 2),以此来“忘掉”历史,突出近期的热度。
二、Count-Min Sketch 的典型使用场景
由于CM Sketch 极其节省内存且速度飞快,它非常适合处理海量、高并发、允许一定误差的流式数据场景:
Top-K / 频繁项查找(Heavy Hitters)
网络流量监控:在骨干路由器中,实时统计哪些 IP 地址发送了大量的巨型数据包(发现 DDoS 攻击源)。
热门搜索词统计:实时计算过去一小时内搜索量最高的关键词。
数据流频次估算
防刷与限流:网络爬虫或恶意防刷系统,实时统计某个 API 被某个 IP 访问的频率。
网页/视频去重与推荐:估算某个用户对某一类内容的点击频次。
缓存淘汰策略(如 TinyLFU)
现代高性能缓存(如 Java 的 Caffeine 缓存库)使用 CM Sketch 的变体来记录数据的访问频率。通过极小的内存消耗,判断新来的数据是否比当前缓存中的老数据更“热”,从而决定是否淘汰老数据。
三、优缺点对比
优点 🌟 | 缺点 ⚠️ |
|---|---|
内存极小:可以处理数以亿计的去重数据,而内存只需几兆字节。 | 不支持准确的单元素查询:结果是一个近似值,对于低频元素,误差可能相对较大。 |
时间常数级:吞吐量极高,适合硬件芯片(如 FPGA/ASIC)或高并发网络设备。 | 只能高估,不能低估:估算值 hat{a} 大于或等于真实值。 |
支持合并:两个相同维度的 CM Sketch 矩阵可以直接按位置相加,实现分布式合并。 | 不支持删除操作:由于哈希碰撞,如果直接减去计数,会影响到其他共享该位置的元素(可用 Counting Bloom Filter 思想演进的变体解决)。 |
四、代码实现
参考mooncake中实现:
#pragma once #include <cstdint> #include <functional> #include <mutex> #include <string> #include <vector> namespace mooncake { // A simple Count-Min Sketch for tracking key access frequency. // Used by the frequency admission policy to decide whether a key // should be promoted into the local hot cache. class CountMinSketch { public: explicit CountMinSketch(size_t width = 4096, size_t depth = 4) : width_(width > 0 ? width : kDefaultWidth), depth_(depth > 0 ? depth : kDefaultDepth), table_(depth_, std::vector<uint8_t>(width_, 0)), total_increments_(0) {} // Increment the count for |key| and return the estimated min-count. // Automatically triggers decay when total_increments exceeds the // threshold (width * depth) to prevent counters from saturating. uint8_t increment(const std::string &key) { std::lock_guard<std::mutex> lock(mu_); uint8_t min_val = UINT8_MAX; for (size_t i = 0; i < depth_; ++i) { size_t idx = hash(key, i) % width_; if (table_[i][idx] < UINT8_MAX) { ++table_[i][idx]; } min_val = std::min(min_val, table_[i][idx]); } if (++total_increments_ >= width_ * depth_) { decayLocked(); } return min_val; } // Return the estimated count for |key| (read-only). uint8_t count(const std::string &key) const { std::lock_guard<std::mutex> lock(mu_); uint8_t min_val = UINT8_MAX; for (size_t i = 0; i < depth_; ++i) { size_t idx = hash(key, i) % width_; min_val = std::min(min_val, table_[i][idx]); } return min_val; } // Halve all counters (right-shift by 1). Useful for periodic aging. void decay() { std::lock_guard<std::mutex> lock(mu_); decayLocked(); } private: static constexpr size_t kDefaultWidth = 4096; static constexpr size_t kDefaultDepth = 4; size_t hash(const std::string &key, size_t seed) const { // Combine std::hash with a per-row seed to get independent hashes. size_t h = std::hash<std::string>{}(key); h ^= seed * 0x9e3779b97f4a7c15ULL + 0x517cc1b727220a95ULL; h ^= (h >> 33); h *= 0xff51afd7ed558ccdULL; h ^= (h >> 33); return h; } void decayLocked() { for (size_t i = 0; i < depth_; ++i) { for (size_t j = 0; j < width_; ++j) { table_[i][j] >>= 1; } } total_increments_ = 0; } const size_t width_; const size_t depth_; std::vector<std::vector<uint8_t>> table_; size_t total_increments_; mutable std::mutex mu_; }; } // namespace mooncake