第16章 模板编程
pybind11 对 C++ 模板提供了良好的支持,可以通过隐式实例化或显式参数指定来导出模板。
16.1 函数模板导出
Section titled “16.1 函数模板导出”函数模板通过隐式实例化自动导出——只要在 Python 端调用,编译器就会生成相应的代码。
#include <pybind11/pybind11.h>#include <vector>#include <string>
namespace py = pybind11;
// 通用函数模板:交换两个值template <typename T>void swap(T& a, T& b) { T temp = a; a = b; b = temp;}
// 函数模板:获取最大值template <typename T>T max_value(const T& a, const T& b) { return a > b ? a : b;}
// 函数模板:容器求和template <typename T>T sum_values(const std::vector<T>& values) { T total = T(); for (const auto& v : values) { total += v; } return total;}
// 函数模板:查找元素template <typename T>int find_index(const std::vector<T>& vec, const T& value) { for (size_t i = 0; i < vec.size(); ++i) { if (vec[i] == value) { return static_cast<int>(i); } } return -1;}
PYBIND11_MODULE(template_function_module, m) { // 隐式实例化:pybind11 会为 int、double、std::string 生成绑定 m.def("swap", &swap<int>); m.def("swap", &swap<double>);
m.def("max_value", &max_value<int>); m.def("max_value", &max_value<double>);
m.def("sum_values", &sum_values<int>); m.def("sum_values", &sum_values<double>);
m.def("find_index", &find_index<int>); m.def("find_index", &find_index<double>);}>>> import template_function_module as m
>>> # swap 测试>>> a, b = 5, 10>>> m.swap(a, b)>>> a, b(10, 5)
>>> # max_value 测试>>> m.max_value(3, 7)7>>> m.max_value(3.14, 2.71)3.14
>>> # find_index 测试>>> m.find_index([1, 2, 3, 4, 5], 3)2>>> m.find_index([1, 2, 3, 4, 5], 10)-1关键洞察:函数模板需要显式指定要实例化的类型。pybind11 不会自动推断——必须为每种使用的类型显式调用模板函数。编译器根据实际调用生成代码。
16.2 类模板导出
Section titled “16.2 类模板导出”类模板同样需要显式实例化导出。
#include <pybind11/pybind11.h>#include <vector>#include <string>#include <stdexcept>
namespace py = pybind11;
// 通用栈类模板template <typename T>class Stack {public: Stack() = default;
void push(const T& value) { data_.push_back(value); }
T pop() { if (data_.empty()) { throw std::out_of_range("stack is empty"); } T value = data_.back(); data_.pop_back(); return value; }
const T& top() const { if (data_.empty()) { throw std::out_of_range("stack is empty"); } return data_.back(); }
bool empty() const { return data_.empty(); }
size_t size() const { return data_.size(); }
private: std::vector<T> data_;};
// 通用 Pair 类模板template <typename T1, typename T2>class Pair {public: Pair(const T1& first, const T2& second) : first_(first), second_(second) {}
T1 first; T2 second;};
// 智能包装器:自动处理返回值的生命周期template <typename T>class SmartWrapper {public: SmartWrapper(T value) : value_(std::move(value)) {}
T& get() { return value_; } const T& get() const { return value_; }
private: T value_;};
PYBIND11_MODULE(template_class_module, m) { // 导出 int 类型的 Stack py::class_<Stack<int>>(m, "IntStack") .def(py::init<>()) .def("push", &Stack<int>::push) .def("pop", &Stack<int>::pop) .def("top", &Stack<int>::top, py::return_value_policy::reference_internal) .def("empty", &Stack<int>::empty) .def("size", &Stack<int>::size);
// 导出 double 类型的 Stack py::class_<Stack<double>>(m, "DoubleStack") .def(py::init<>()) .def("push", &Stack<double>::push) .def("pop", &Stack<double>::pop) .def("top", &Stack<double>::top, py::return_value_policy::reference_internal) .def("empty", &Stack<double>::empty) .def("size", &Stack<double>::size);
// 导出 Pair 模板(int, std::string)实例 py::class_<Pair<int, std::string>>(m, "IntStringPair") .def(py::init<const std::string&, const std::string&>()) .def_readwrite("first", &Pair<int, std::string>::first) .def_readwrite("second", &Pair<int, std::string>::second);
// 导出 SmartWrapper<int> py::class_<SmartWrapper<int>>(m, "SmartIntWrapper") .def(py::init<int>()) .def("get", &SmartWrapper<int>::get, py::return_value_policy::reference_internal);}>>> import template_class_module as m
>>> # IntStack 测试>>> stack = m.IntStack()>>> stack.push(1)>>> stack.push(2)>>> stack.push(3)>>> stack.size()3>>> stack.top()3>>> stack.pop()3>>> stack.size()2
>>> # DoubleStack 测试>>> dstack = m.DoubleStack()>>> dstack.push(3.14)>>> dstack.push(2.71)>>> dstack.pop()2.71
>>> # Pair 测试>>> pair = m.IntStringPair(42, "hello")>>> pair.first42>>> pair.second'hello'16.3 模板别名
Section titled “16.3 模板别名”使用类型别名简化复杂的模板类型。
#include <pybind11/pybind11.h>#include <vector>#include <string>#include <memory>
namespace py = pybind11;
// 定义常用模板别名template <typename T>using Ptr = std::shared_ptr<T>;
template <typename T>using Vector = std::vector<T>;
template <typename K, typename V>using Map = std::map<K, V>;
// 模板类template <typename T>class Resource {public: Resource(T value) : value_(std::move(value)) {} T get() const { return value_; }private: T value_;};
PYBIND11_MODULE(template_alias_module, m) { // 使用别名导出模板类 py::class_<Resource<int>>(m, "IntResource") .def(py::init<int>()) .def("get", &Resource<int>::get);
py::class_<Resource<std::string>>(m, "StringResource") .def(py::init<const std::string&>()) .def("get", &Resource<std::string>::get);
// 导出返回 Vector 的函数 m.def("create_int_vector", []() { Vector<int> vec = {1, 2, 3, 4, 5}; return vec; // pybind11 会自动处理 std::vector 转换 });
// 导出返回 Map 的函数 m.def("create_string_map", []() { Map<std::string, int> m = { {"one", 1}, {"two", 2}, {"three", 3} }; return m; });}>>> import template_alias_module as m
>>> resource = m.IntResource(100)>>> resource.get()100
>>> string_res = m.StringResource("template")>>> string_res.get()'template'
>>> m.create_int_vector()[1, 2, 3, 4, 5]
>>> m.create_string_map(){'one': 1, 'two': 2, 'three': 3}16.4 编译时多态(CRTP 模式)
Section titled “16.4 编译时多态(CRTP 模式)”奇异递归模板模式(Curiously Recurring Template Pattern,CRTP)实现编译时多态。
#include <pybind11/pybind11.h>#include <string>#include <memory>
namespace py = pybind11;
// CRTP 基类template <typename Derived>class Shape {public: // 编译时多态:调用子类的实现 double area() { return static_cast<Derived*>(this)->area_impl(); }
std::string name() { return static_cast<Derived*>(this)->name_impl(); }
// 通用方法 void print_info() { py::gil_scoped_acquire gil; py::print(name(), "area:", area()); }
protected: Shape() = default; Shape(const Shape&) = default;};
// 派生类:Circleclass Circle : public Shape<Circle> {public: Circle(double radius) : radius_(radius) {}
double area_impl() const { return 3.14159265358979 * radius_ * radius_; }
std::string name_impl() const { return "Circle"; }
private: double radius_;};
// 派生类:Rectangleclass Rectangle : public Shape<Rectangle> {public: Rectangle(double width, double height) : width_(width), height_(height) {}
double area_impl() const { return width_ * height_; }
std::string name_impl() const { return "Rectangle"; }
private: double width_; double height_;};
// 工厂函数模板template <typename ShapeType, typename... Args>std::shared_ptr<ShapeType> create_shape(Args&&... args) { return std::make_shared<ShapeType>(std::forward<Args>(args)...);}
PYBIND11_MODULE(crtp_module, m) { // 导出 Circle py::class_<Circle, std::shared_ptr<Circle>>(m, "Circle") .def(py::init<double>(), py::arg("radius")) .def("area", &Circle::area) .def("name", &Circle::name) .def("print_info", &Circle::print_info);
// 导出 Rectangle py::class_<Rectangle, std::shared_ptr<Rectangle>>(m, "Rectangle") .def(py::init<double, double>(), py::arg("width"), py::arg("height")) .def("area", &Rectangle::area) .def("name", &Rectangle::name) .def("print_info", &Rectangle::print_info);
// 工厂函数 m.def("create_circle", &create_shape<Circle, double>); m.def("create_rectangle", &create_shape<Rectangle, double, double>);}>>> import crtp_module as m
>>> circle = m.create_circle(5.0)>>> circle.name()'Circle'>>> circle.area()78.53981633974483
>>> rect = m.create_rectangle(4.0, 6.0)>>> rect.name()'Rectangle'>>> rect.area()24.0
>>> circle.print_info()Circle area: 78.53981633974483rect.print_info()Rectangle area: 24.0关键洞察:CRTP 模式通过模板继承实现编译时多态,避免了虚函数的开销。所有调用在编译时解析,没有运行时查找。pybind11 可以直接处理继承自 CRTP 基类的模板类。
16.5 模板特化
Section titled “16.5 模板特化”使用模板特化为特定类型提供特殊实现。
#include <pybind11/pybind11.h>#include <vector>#include <string>#include <stdexcept>
namespace py = pybind11;
// 通用模板实现template <typename T>class Container {public: void add(const T& value) { data_.push_back(value); }
T get(size_t index) { if (index >= data_.size()) { throw std::out_of_range("index out of range"); } return data_[index]; }
size_t size() const { return data_.size(); }
private: std::vector<T> data_;};
// 特化模板:字符串处理template <>class Container<std::string> {public: void add(const std::string& value) { // 字符串特化:转换为大写 std::string upper = value; for (auto& c : upper) { c = std::toupper(c); } data_.push_back(upper); }
std::string get(size_t index) { if (index >= data_.size()) { throw std::out_of_range("index out of range"); } return data_[index]; }
size_t size() const { return data_.size(); }
// 特化版本特有的方法 void add_with_prefix(const std::string& prefix, const std::string& value) { data_.push_back(prefix + value); }
private: std::vector<std::string> data_;};
// 函数模板特化template <typename T>T process_value(T value) { return value * 2;}
// 函数模板特化:字符串template <>std::string process_value(std::string value) { return "processed: " + value;}
PYBIND11_MODULE(template_specialization_module, m) { // 导出 int Container py::class_<Container<int>>(m, "IntContainer") .def(py::init<>()) .def("add", &Container<int>::add) .def("get", &Container<int>::get) .def("size", &Container<int>::size);
// 导出 string Container(特化版本) py::class_<Container<std::string>>(m, "StringContainer") .def(py::init<>()) .def("add", &Container<std::string>::add) .def("get", &Container<std::string>::get) .def("size", &Container<std::string>::size) .def("add_with_prefix", &Container<std::string>::add_with_prefix);
// 导出函数模板特化 m.def("process_value", &process_value<int>); m.def("process_value", &process_value<std::string>);}>>> import template_specialization_module as m
>>> # int container(通用实现)>>> ic = m.IntContainer()>>> ic.add(42)>>> ic.add(100)>>> ic.get(0)42
>>> # string container(特化实现,自动转大写)>>> sc = m.StringContainer()>>> sc.add("hello")>>> sc.get(0)'HELLO'
>>> # 特化方法>>> sc.add_with_prefix("PRE:", "special")>>> sc.get(1)'PRE:SPECIAL'
>>> # 函数模板特化>>> m.process_value(5)10>>> m.process_value("test")'processed: test'关键洞察:模板特化允许为特定类型提供定制化实现。pybind11 可以导出特化版本和通用版本。需要显式实例化每种使用的类型。字符串特化常用于添加特定的大小写转换或其他字符串操作。
模板编程总结:
| 技术 | 说明 | 注意事项 |
|---|---|---|
| 函数模板 | 隐式实例化,需显式导出每种类型 | 显式调用 func<Type> |
| 类模板 | 需要为每个类型实例化类 | 使用 py::class_<TemplateType> |
| 模板别名 | 简化复杂模板类型的表达 | 纯语法层面,不影响绑定 |
| CRTP | 编译时多态,避免虚函数开销 | 导出派生类即可 |
| 特化 | 为特定类型提供定制实现 | 导出特化版本,显式实例化 |
实战建议:优先使用模板减少代码重复,但要注意 pybind11 需要显式实例化。如果模板参数不确定,考虑使用 py::object 作为泛型容器。