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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.25 18:16:05