如何编写通用宏调度带CacheEvictStrategy模板参数的模板函数?
问题描述
我有一个模板类Cache,包含Get、Put、Evict等方法,以及CacheEvictStrategy(FIFO、LRU、LFU)三种淘汰策略。
template<typename KeyType, typename ElemType> class Cache; template<typename KeyType, typename ElemType> void Cache::Get(uint32_t num_query, KeyType* queries, ElemType* result, bool* find_mask){ }
由于三种CacheEvictStrategy共享大量代码,我将通用逻辑写在Cache::Get中,它会调用另一模板函数GetInternal,该函数会根据不同的CacheEvictStrategy执行不同逻辑:
template<typename KeyType, typename ElemType, CacheEvictStrategy Strategy> __global__ void GetInternal(uint32_t num_query, KeyType* queries, ElemType* result, bool* find_mask); template<typename KeyType, typename ElemType> void Cache::Get(uint32_t num_query, KeyType* queries, ElemType* result, bool* find_mask){ // 通用逻辑 GetInternal<KeyType, ElemType, Strategy><<<grid,block>>>(num_query,queries,result,find_mask); }
由于Strategy是在运行时确定的,我希望编写一个宏来进行调度。但我还有Put、Evict等函数,希望编写一个可接收函数和CacheEvictStrategy的通用宏,而非为每个函数单独编写宏(如下所示):
#define DISPATCH_GET(strategy,grid,block,...){ \ switch(strategy){ \ case LRU: \ GetInternal<KeyType,ElemType,LRU><<<grid, block>>>(__VA_ARGS__); \ /* 其他case */ \ } \ } #define DISPATCH_PUT #define DISPATCH_EVICT
请问是否有可行的解决方案?希望能得到相关建议。
解决方案
可以编写一个通用调度宏,将内部模板函数、模板参数和运行时策略作为参数传入,避免为每个操作重复编写switch逻辑。核心思路是让宏自动生成针对不同策略的分支,填充模板参数并调用对应的内核函数。
通用调度宏实现
#define DISPATCH_CACHE_OP(INTERNAL_FUNC, KeyType, ElemType, strategy, grid, block, ...) \ do { \ switch(strategy) { \ case FIFO: \ INTERNAL_FUNC<KeyType, ElemType, FIFO><<<grid, block>>>(__VA_ARGS__); \ break; \ case LRU: \ INTERNAL_FUNC<KeyType, ElemType, LRU><<<grid, block>>>(__VA_ARGS__); \ break; \ case LFU: \ INTERNAL_FUNC<KeyType, ElemType, LFU><<<grid, block>>>(__VA_ARGS__); \ break; \ default: \ /* 处理无效策略的逻辑,比如断言或报错 */ \ assert(false && "Unsupported cache eviction strategy"); \ break; \ } \ } while(0)
使用示例
在Cache::Get、Cache::Put等方法中直接调用该宏即可:
template<typename KeyType, typename ElemType> void Cache::Get(uint32_t num_query, KeyType* queries, ElemType* result, bool* find_mask){ // 通用逻辑处理 dim3 grid(...); // 假设已计算好grid和block维度 dim3 block(...); DISPATCH_CACHE_OP(GetInternal, KeyType, ElemType, this->strategy, grid, block, num_query, queries, result, find_mask); } template<typename KeyType, typename ElemType> void Cache::Put(uint32_t num_put, KeyType* keys, ElemType* elems){ // 通用逻辑处理 dim3 grid(...); dim3 block(...); DISPATCH_CACHE_OP(PutInternal, KeyType, ElemType, this->strategy, grid, block, num_put, keys, elems); }
额外优化建议
- 封装策略分支:如果策略数量较多,可以把switch分支抽成单独的宏片段,让通用宏复用它,进一步减少冗余:
#define CACHE_STRATEGY_CASES(INTERNAL_FUNC, KeyType, ElemType, grid, block, ...) \ case FIFO: \ INTERNAL_FUNC<KeyType, ElemType, FIFO><<<grid, block>>>(__VA_ARGS__); \ break; \ case LRU: \ INTERNAL_FUNC<KeyType, ElemType, LRU><<<grid, block>>>(__VA_ARGS__); \ break; \ case LFU: \ INTERNAL_FUNC<KeyType, ElemType, LFU><<<grid, block>>>(__VA_ARGS__); \ break; #define DISPATCH_CACHE_OP(INTERNAL_FUNC, KeyType, ElemType, strategy, grid, block, ...) \ do { \ switch(strategy) { \ CACHE_STRATEGY_CASES(INTERNAL_FUNC, KeyType, ElemType, grid, block, __VA_ARGS__) \ default: \ assert(false && "Unsupported cache eviction strategy"); \ break; \ } \ } while(0)
- C++17类型安全替代方案:如果可以使用C++17及以上标准,建议用
std::visit配合变体类型替代宏,获得更好的类型安全性:
// 定义策略对应的类型标签 struct FIFOTag {}; struct LRUTag {}; struct LFUTag {}; // 用std::variant存储运行时策略 using CacheStrategyVariant = std::variant<FIFOTag, LRUTag, LFUTag>; // 定义调度函数 template<typename InternalFunc, typename KeyType, typename ElemType, typename... Args> void dispatch(CacheStrategyVariant variant, dim3 grid, dim3 block, Args&&... args) { std::visit([&](auto tag) { using StrategyTag = decltype(tag); if constexpr (std::is_same_v<StrategyTag, FIFOTag>) { InternalFunc<KeyType, ElemType, FIFO><<<grid, block>>>(std::forward<Args>(args)...); } else if constexpr (std::is_same_v<StrategyTag, LRUTag>) { InternalFunc<KeyType, ElemType, LRU><<<grid, block>>>(std::forward<Args>(args)...); } else if constexpr (std::is_same_v<StrategyTag, LFUTag>) { InternalFunc<KeyType, ElemType, LFU><<<grid, block>>>(std::forward<Args>(args)...); } }, variant); } // 使用方式 void Cache::Get(...) { CacheStrategyVariant variant = (this->strategy == FIFO) ? FIFOTag{} : (this->strategy == LRU) ? LRUTag{} : LFUTag{}; dispatch<GetInternal, KeyType, ElemType>(variant, grid, block, num_query, queries, result, find_mask); }
这种方式避免了宏的预处理魔法,类型检查更严格,适合现代C++项目。
内容的提问来源于stack exchange,提问作者shadow
相关产品推荐
相关产品推荐

