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

如何自动生成固定模板类型的全部实例化以适配Pybind11绑定?

问题

现有模板函数:

template<typename T1, typename T2, typename T3>
void func(std::string a, T1 arg1, T2 arg2, T3 arg3, bool b);

要求模板类型仅能取int或std::vector<int>,需要在编译时自动生成所有可能的实例化,用于Pybind11绑定Python扩展,避免手动编写所有组合(如func<int,int,std::vector<int>>等)。

原尝试代码如下(存在编译错误):

template <typename T1, typename T2, typename T3, typename T4>
void bind_func(py::module& m) {
  m.def("func", &func<T1, T2, T3, T4>);
}

//Recursive function to generate bindings for all combinations of int and
//vector<int> in three slots
template <unsigned int N, typename... Args> struct generateBindings{
  static void generate(py::module &m) {
    generateBindings<N-1, Args..., int>::generate(m);
    generateBindings<N-1, Args..., std::vector<int>>::generate(m);
  }
};

template <typename... Args>
struct generateBindings<0, Args...> {
  static void generate(py::module &m) {
    bind_func<Args...>(m);
  }
};

// 实例化方式
generateBindings<2,int>(m);
generateBindings<2,std::vector<int>>(m);
解决方案

原代码问题分析

  1. bind_func声明了4个模板参数,但目标func仅需要3个,参数数量不匹配,导致编译报错。
  2. 手动实例化generateBindings<2, int>和generateBindings<2, std::vector<int>>的方式错误,既没有覆盖所有3参数组合,递归逻辑的终止条件和参数传递也不符合需求。

方案一:针对固定3个模板参数的简化实现

直接生成int和std::vector<int>的所有8种组合,代码简洁易读:

#include <vector>
#include <string>
#include <pybind11/pybind11.h>

namespace py = pybind11;

template<typename T1, typename T2, typename T3>
void func(std::string a, T1 arg1, T2 arg2, T3 arg3, bool b) {
    // 函数实现
}

void bind_all_funcs(py::module& m) {
    // 生成所有8种组合
    m.def("func", &func<int, int, int>);
    m.def("func", &func<int, int, std::vector<int>>);
    m.def("func", &func<int, std::vector<int>, int>);
    m.def("func", &func<int, std::vector<int>, std::vector<int>>);
    m.def("func", &func<std::vector<int>, int, int>);
    m.def("func", &func<std::vector<int>, int, std::vector<int>>);
    m.def("func", &func<std::vector<int>, std::vector<int>, int>);
    m.def("func", &func<std::vector<int>, std::vector<int>, std::vector<int>>);
}

方案二:支持任意N个模板参数的通用实现

利用C++模板元编程递归生成所有类型组合,可轻松扩展到任意数量的模板参数:

步骤1:定义基础类型列表

首先定义允许使用的类型集合:

using allowed_types = std::tuple<int, std::vector<int>>;

步骤2:递归生成所有组合的模板结构体

// 递归终止:当参数数量达到N时,绑定对应的func实例
template<typename... Args>
struct generate_bindings_impl {
    static void bind(py::module& m) {
        m.def("func", &func<Args...>);
    }
};

// 递归生成组合:每次添加allowed_types中的一种类型,直到参数数量达标
template<size_t N, typename... CurrentArgs>
struct generate_bindings {
    static void generate(py::module& m) {
        // 遍历allowed_types中的每个类型,递归添加到参数列表
        std::apply([&m](auto... types) {
            (generate_bindings<N-1, CurrentArgs..., decltype(types)>::generate(m), ...);
        }, allowed_types{});
    }
};

// 特化:当N=0时,调用终止逻辑绑定函数
template<typename... CurrentArgs>
struct generate_bindings<0, CurrentArgs...> {
    static void generate(py::module& m) {
        generate_bindings_impl<CurrentArgs...>::bind(m);
    }
};

步骤3:对外暴露的调用接口

// 绑定N个模板参数的所有组合
template<size_t N>
void bind_func_combinations(py::module& m) {
    generate_bindings<N>::generate(m);
}

使用方式

在Pybind11模块初始化中调用:

PYBIND11_MODULE(my_module, m) {
    // 生成3个模板参数的所有组合(共2^3=8种)
    bind_func_combinations<3>(m);
}

方案优势

  • 通用实现支持任意数量的模板参数,只需修改bind_func_combinations的模板参数值即可扩展。
  • 编译时生成所有实例,无运行时开销。
  • 代码结构清晰,避免手动编写重复的绑定代码。

内容的提问来源于stack exchange,提问作者Arham

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 09:05:39