You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何编写通用宏调度带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);
}

额外优化建议

  1. 封装策略分支:如果策略数量较多,可以把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)
  1. 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.27 18:51:12