MLIR中处理可选/不存在模板参数的模板特化替代方案
MLIR下转换Pass通用模板的类型适配问题
问题场景与代码示例
我正在开发MLIR下转换(Lowering)Pass,现有大量下转换逻辑相似的Op,因此尝试通过模板实现通用化处理。部分Op并非针对所有无符号/有符号16/32/64位类型都具备对应的Intrinsic:例如某Op没有32位版本的Intrinsic,另一Op仅支持无符号类型的Intrinsic。
当前遇到的问题是:编译器会判定createIntrinsic中所有模板类型参数均被使用,因此无法简单传入空类作为NONEXISTENT参数;但我可以保证,不会出现调用matchAndRewrite时传入会触发不存在Intrinsic分支的Op。
原始代码如下:
template <typename SomeOp, typename Intrinsic_U16, typename Intrinsic_U32, typename Intrinsic_U64, typename Intrinsic_S16, typename Intrinsic_S32, typename Intrinsic_S64> struct GeneralMLIROpLowering { template <typename T> Value createIntrinsic(...) { return rewriter.create<T>(...); } void matchAndRewrite(SomeOp Op) { switch (getOpType(Op)) { case OpType::U16: createIntrinsic<Intrinsic_U16>(...); case OpType::U32: createIntrinsic<Intrinsic_U32>(...); case OpType::U64: createIntrinsic<Intrinsic_U64>(...); ... default: error(); } } }; using SpecificMLIROpLowering = GeneralMLIROpLowering<SomeIntrinsic_U16, SomeIntrinsic_U32, SomeIntrinsic_U64, ...>; // 问题场景:该Op没有U32版本的Intrinsic using SpecificMLIROpLowering = GeneralMLIROpLowering<SomeIntrinsic_U16, NONEXISTENT, SomeIntrinsic_U64, ...>;
解决方案思路
1. 空标签类型+SFINAE禁用无效分支
定义一个空的标签类型,结合SFINAE让createIntrinsic在传入空类型时不会实例化无效的rewriter.create调用,同时通过断言保证不会走到对应分支:
struct NoneType {}; template <typename SomeOp, typename Intrinsic_U16, typename Intrinsic_U32, typename Intrinsic_U64, typename Intrinsic_S16, typename Intrinsic_S32, typename Intrinsic_S64> struct GeneralMLIROpLowering { // 仅针对有效Intrinsic类型启用的重载 template <typename T, std::enable_if_t<!std::is_same_v<T, NoneType>, int> = 0> Value createIntrinsic(...) { return rewriter.create<T>(...); } // 针对NoneType的重载,仅做断言永远不会被调用 template <typename T, std::enable_if_t<std::is_same_v<T, NoneType>, int> = 0> Value createIntrinsic(...) { llvm_unreachable("This intrinsic branch should never be triggered!"); } void matchAndRewrite(SomeOp Op) { switch (getOpType(Op)) { case OpType::U16: createIntrinsic<Intrinsic_U16>(...); break; // 注意补充break避免分支穿透 case OpType::U32: createIntrinsic<Intrinsic_U32>(...); break; case OpType::U64: createIntrinsic<Intrinsic_U64>(...); break; ... default: llvm_unreachable("Unknown Op type"); } } }; // 使用示例:传入NoneType替代不存在的U32 Intrinsic using SpecificMLIROpLowering = GeneralMLIROpLowering<SomeOp, SomeIntrinsic_U16, NoneType, SomeIntrinsic_U64, ...>;
2. 显式可选类型包装器+编译期检查
定义包装器类型标记是否存在有效Intrinsic,配合static_assert在编译期提前拦截误用:
// 有效Intrinsic的包装器 template <typename T> struct OptionalIntrinsic { using type = T; static constexpr bool has_intrinsic = true; }; // 无有效Intrinsic的包装器特化 template <> struct OptionalIntrinsic<void> { static constexpr bool has_intrinsic = false; }; template <typename SomeOp, typename Intrinsic_U16 = OptionalIntrinsic<void>, typename Intrinsic_U32 = OptionalIntrinsic<void>, typename Intrinsic_U64 = OptionalIntrinsic<void>, typename Intrinsic_S16 = OptionalIntrinsic<void>, typename Intrinsic_S32 = OptionalIntrinsic<void>, typename Intrinsic_S64 = OptionalIntrinsic<void>> struct GeneralMLIROpLowering { template <typename Wrapper> std::enable_if_t<Wrapper::has_intrinsic, Value> createIntrinsic(...) { return rewriter.create<typename Wrapper::type>(...); } template <typename Wrapper> std::enable_if_t<!Wrapper::has_intrinsic, Value> createIntrinsic(...) { llvm_unreachable("Unsupported intrinsic type"); } void matchAndRewrite(SomeOp Op) { switch (getOpType(Op)) { case OpType::U16: static_assert(Intrinsic_U16::has_intrinsic, "U16 intrinsic is required for this Op"); createIntrinsic<Intrinsic_U16>(...); break; case OpType::U32: static_assert(Intrinsic_U32::has_intrinsic, "U32 intrinsic is required for this Op"); createIntrinsic<Intrinsic_U32>(...); break; ... } } }; // 使用示例:仅传入有Intrinsic的类型包装器 using SpecificMLIROpLowering = GeneralMLIROpLowering<SomeOp, OptionalIntrinsic<SomeIntrinsic_U16>, OptionalIntrinsic<void>, // 标记无U32 Intrinsic OptionalIntrinsic<SomeIntrinsic_U64>, ...>;
3. 可变参数模板+类型标签映射
如果Op支持的类型差异较大,改用类型标签+可变参数模板,只传入该Op支持的类型组合,避免冗余参数:
// 定义类型标签 struct U16Tag {}; struct U32Tag {}; struct U64Tag {}; struct S16Tag {}; struct S32Tag {}; struct S64Tag {}; // 类型标签到Intrinsic的映射模板 template <typename Tag> struct IntrinsicMap; template <> struct IntrinsicMap<U16Tag> { using type = SomeIntrinsic_U16; }; template <> struct IntrinsicMap<U32Tag> { using type = SomeIntrinsic_U32; }; // ... 其他类型的映射特化 template <typename SomeOp, typename... SupportedTags> struct GeneralMLIROpLowering { template <typename Tag> Value createIntrinsic(...) { return rewriter.create<typename IntrinsicMap<Tag>::type>(...); } void matchAndRewrite(SomeOp Op) { auto opType = getOpType(Op); if constexpr ((std::is_same_v<SupportedTags, U16Tag> && ...) && opType == OpType::U16) { createIntrinsic<U16Tag>(...); } else if constexpr ((std::is_same_v<SupportedTags, U32Tag> && ...) && opType == OpType::U32) { createIntrinsic<U32Tag>(...); } // ... 其他支持类型的判断 else { llvm_unreachable("Unsupported type for this Op lowering"); } } }; // 使用示例:仅传入该Op支持的U16、U64标签 using SpecificMLIROpLowering = GeneralMLIROpLowering<SomeOp, U16Tag, U64Tag>;
内容的提问来源于stack exchange,提问作者Each One Chew
相关产品推荐
相关产品推荐

