在不使用constexpr的情况下使用模板
在我实际的代码中,我有一个模板函数,其行为会根据输入的数据类型而变化。我必须找到一种方法,将输入数组的数据类型传递给该函数。变量tA、tB和tC是一个enum的一部分,这个enum提示输入数组的数据类型。若有助于理解,在实际代码中,输入数组是void*。问题在于,在我的实际代码中(我无法提供)变量tA、tB和tC不能是constexpr,这会在编译时导致错误:
error: non-type template argument is not a constant expression
56 | some_template_function<DataType<tA>,DataType<tB>,DataType<tD>>();
下面是在constexpr条件下可运行的示例代码:
#include <complex>
template<typename TA, typename TB, typename TC>
void some_template_function(){
TA a;
TB b;
TC c;
}
typedef int MY_Datatype;
enum
{
/* IEEE754 float32: 1 sign bit, 8 exponent bits, 23 explicit significand bits */
MY_F32 = 0,
/* IEEE754 float64: 1 sign bit, 11 exponent bits, 52 explicit significand bits */
MY_F64 = 1,
/* Complex IEEE754 float32, stored with consecutive real and imaginary parts packed into 8 bytes */
MY_C32 = 2,
/* Complex IEEE754 float64, stored with consecutive real and imaginary parts packed into 16 bytes */
MY_C64 = 3,
/* IEEE754 float16: 1 sign bit, 5 exponent bits, 10 explicit significand bits */
MY_F16 = 4,
/* bfloat16: 1 sign bit, 8 exponent bits, 7 explicit significand bits */
MY_BF16 = 5,
/* Aliases */
MY_FLOAT = MY_F32,
MY_DOUBLE = MY_F64,
MY_SCOMPLEX = MY_C32,
MY_DCOMPLEX = MY_C64,
};
template<MY_Datatype dtype>
constexpr auto select_cuda_datatype() {
if constexpr (dtype == MY_F32) return float{};
else if constexpr (dtype == MY_F64) return double{};
else if constexpr (dtype == MY_C32) return std::complex<float>{};
else if constexpr (dtype == MY_C64) return std::complex<double>{};
}
template<MY_Datatype dtype>
using DataType = decltype(select_cuda_datatype<dtype>());
int main(){
// /!\ Cannot use constexpr in my actual application/code
constexpr MY_Datatype tA = MY_F64;
constexpr MY_Datatype tB = MY_F64;
constexpr MY_Datatype tD = MY_F64;
some_template_function<DataType<tA>,DataType<tB>,DataType<tD>>();
printf("Finished main!\n");
return 0;
}
解决方案
std::variant 与 std::visit 可能实现分派:
using cuda_datatype_var = std::variant<
std::type_identity<float>,
std::type_identity<double>,
std::type_identity<std::complex<float>>,
std::type_identity<std::complex<double>>
>;
cuda_datatype_var select_cuda_datatype(MY_Datatype dtype) {
switch (dtype) {
case MY_F32: return std::type_identity<float>{};
case MY_F64: return std::type_identity<double>{};
case MY_C32: return std::type_identity<std::complex<float>>{};
case MY_C64: return std::type_identity<std::complex<double>>{};
}
throw std::runtime_error("Unsuported value");
}
void dispatch(cuda_datatype_var var1,
cuda_datatype_var var2,
cuda_datatype_var var3)
{
std::visit([]<typename T1, typename T2, typename T3>(std::type_identity<T1>,
std::type_identity<T2>,
std::type_identity<T3>) {
some_template_function<T1, T2, T3>();
}, var1, var2, var3);
}
void dispatch(MY_Datatype t1, MY_Datatype t2, MY_Datatype t3)
{
dispatch(select_cuda_datatype(t1),
select_cuda_datatype(t2),
select_cuda_datatype(t3));
}
站内所有文章版权归属LeftHeroAI导航站,无授权禁止任何主体转载、抄袭、复制内容,亦不得私自架设镜像站点。一经侵权,本站将通过法律途径追责。