Skip to content

第16章 模板编程

pybind11 对 C++ 模板提供了良好的支持,可以通过隐式实例化或显式参数指定来导出模板。

函数模板通过隐式实例化自动导出——只要在 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 不会自动推断——必须为每种使用的类型显式调用模板函数。编译器根据实际调用生成代码。

类模板同样需要显式实例化导出。

#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.first
42
>>> pair.second
'hello'

使用类型别名简化复杂的模板类型。

#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}

奇异递归模板模式(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;
};
// 派生类:Circle
class 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_;
};
// 派生类:Rectangle
class 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.53981633974483
rect.print_info()
Rectangle area: 24.0

关键洞察:CRTP 模式通过模板继承实现编译时多态,避免了虚函数的开销。所有调用在编译时解析,没有运行时查找。pybind11 可以直接处理继承自 CRTP 基类的模板类。

使用模板特化为特定类型提供特殊实现。

#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 作为泛型容器。