news 2026/9/14 22:05:03

MNN PyMNN 之 data.Dataset:自定义训练数据源的完整实现与源码级解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
MNN PyMNN 之 data.Dataset:自定义训练数据源的完整实现与源码级解析

MNN PyMNN 之 data.Dataset:自定义训练数据源的完整实现与源码级解析

【免费下载链接】MNNMNN: A blazing-fast, lightweight inference engine battle-tested by Alibaba, powering high-performance on-device LLMs and Edge AI.项目地址: https://gitcode.com/GitHub_Trending/mn/MNN

本文围绕 MNN 的 Python 训练接口MNN.data.Dataset展开:它是一个虚基类,用户通过继承并重写__getitem____len__即可把任意数据(MNIST 图像、自建数据集等)接入 MNN 的训练数据管线。读完本文,你既能直接写出一个可运行的 Dataset 实现并配合DataLoader批量取数,也能看懂这条数据是如何通过 PyBind 桥接层(DatasetWrapper)被转换为 C++ 侧的Examplestd::pair<std::vector<VARP>, std::vector<VARP>>)并进入多线程 DataLoader 的。

一、Dataset 的核心契约:两个必须重写的方法

docs/pymnn/Dataset.mdMNN.data.Dataset的定义非常简洁:

class Dataset

Dataset 是一个虚基类,用户实现自己的 Dataset 需要继承基类并重写__getitem____len__方法。

其中有两个关键约定,直接决定了整个训练数据管线的数据形状:

  1. __getitem__(index)必须返回两个变量:一个是输入(inputs),另一个是目标(targets);
  2. 输入和目标各自都是一个 list:因为一个样本可能有多路输入(例如多模态),也可能有多路目标(例如多任务损失),用列表承载可以保持结构统一。

这两个约定不是随便定下的——它就是 C++ 侧Example类型在 Python 侧的镜像。从桥接代码 pymnn/src/data.h 中可以看到:

// typedef std::pair<std::vector<VARP>, std::vector<VARP>> Example; // Example ==> ([Var], [Var]) PyObject *ret = PyList_New(2); PyList_SetItem(ret, 0, toPyObj<VARP, toPyObj>(example.first)); PyList_SetItem(ret, 1, toPyObj<VARP, toPyObj>(example.second));

Example是一个 pair:first为输入 Var 列表,second为目标 Var 列表;Python 侧的(inputs, targets)元组与之严格对应。因此__getitem__返回的每个元素应当是MNN.exprVar(通常由expr.const构造的常量张量),而不是裸的 numpy 数组。

二、完整实战:MNIST 数据集的 Dataset 实现

官方文档给出了一个完整可运行的 MNIST 示例,这里在保留原结构的基础上补充参数注释:

try: import mnist except ImportError: print("please 'pip install mnist' before run this demo") import MNN import MNN.expr as expr class MnistDataset(MNN.data.Dataset): def __init__(self, training_dataset=True): super(MnistDataset, self).__init__() self.is_training_dataset = training_dataset if self.is_training_dataset: self.data = mnist.train_images() / 255.0 # 归一化到 0~1 self.labels = mnist.train_labels() else: self.data = mnist.test_images() / 255.0 self.labels = mnist.test_labels() def __getitem__(self, index): # 输入:float 张量,形状 [1, 28, 28],NCHW 布局(灰度图 C=1) dv = expr.const(self.data[index].flatten().tolist(), [1, 28, 28], expr.NCHW) # 目标:无形状的 uint8 标量张量,存放该样本的类别标签 dl = expr.const([self.labels[index]], [], expr.NCHW, expr.uint8) # first for inputs, and may have many inputs, so it's a list # second for targets, also, there may be more than one targets return [dv], [dl] def __len__(self): # size of the dataset if self.is_training_dataset: return 60000 else: return 10000

几个细节值得注意:

  • 归一化前置mnist.train_images() / 255.0在数据加载阶段就把像素值从 0~255 缩放到 0~1,避免在网络内部再做一次 cast;C++ 侧的等价做法见 tools/train/source/demo/dataLoaderDemo.cpp 中的LambdaTransform_Cast到 float 再乘以1/255)。
  • expr.const的三个重载形态:输入张量传的是展平后的 list + 形状 [1, 28, 28] + expr.NCHW 布局;目标张量传的是形状为空的 list(标量)+ expr.NCHW + 显式 dtype expr.uint8。也就是说,expr.const的最后一个参数可用来指定元素类型,标签用uint8正是为了与 C++ 侧readMap<uint8_t>()的读取类型对齐。
  • 形状选择 [1, 28, 28] 而非 [28, 28]:单样本以 NCHW 的四维形状表达,这样 DataLoader 做 stack 批量化时可以直接沿 channel 维拼接成[n, 1, 28, 28],与 C++ 侧StackTransform的行为一致。

三、桥接层解析:Python 对象如何变成 C++ 的 Dataset

MNN.data.Dataset并非普通 Python 类。它由 PyMNN 的 C 扩展在初始化时注册:桥接代码末尾通过 pymnn/src/data.h 中的toDatasetdef_class_register(Dataset)完成绑定,Python 包侧则由 pymnn/pip_package/MNN/data/init.py 一行from _mnncengine._data import *导出。

真正的工作由DatasetWrapper完成(pymnn/src/data.h):

class DatasetWrapper : public Dataset { public: using Dataset::Dataset; DatasetWrapper(PyObject* py_dataset) { Py_INCREF(py_dataset); this->py_dataset = py_dataset; } Example get(size_t index) override { auto getfunc = PyObject_GetAttrString(py_dataset, "__getitem__"); auto arg = PyTuple_New(1); PyTuple_SetItem(arg, 0, PyLong_FromLong(index)); auto res = PyObject_CallObject(getfunc, arg); ... auto py_example = PyTuple_GetItem(res, 0); auto py_example_second = PyTuple_GetItem(res, 1); auto example = std::make_pair( toVars(py_example), // inputs -> std::vector<VARP> toVars(py_example_second) // targets -> std::vector<VARP> ); return example; } size_t size() override { auto sizefunc = PyObject_GetAttrString(py_dataset, "__len__"); auto res = PyObject_CallObject(sizefunc, NULL); ... } private: PyObject *py_dataset = nullptr; };

从源码结构看,调用链是:

  1. 当你MnistDataset(...)实例化时,PyMNNDataset_init(pymnn/src/data.h)为该实例创建一个DatasetWrapper,其内部持有 Python 对象的引用(Py_INCREF),也就是说C++ 侧拿到的是一个"反向代理",每次取数都会回调回 Python
  2. C++ 侧任何需要样本的地方调用Dataset::get(index)DatasetWrapper就通过PyObject_GetAttrString找到你的__getitem__并调用,把返回的(inputs, targets)元组转换成Example
  3. size()同理回调你的__len__

这个"反向回调"设计带来两个实践结论:

  • 数据预处理可以完全留在 Python 生态里:numpy 归一化、解码、增广等都不需要写 C++;
  • __getitem__会被 DataLoader 的 worker 线程反复调用,因此其中不应有重复计算,MNIST 示例把mnist.train_images()一次性缓存到self.data就是出于这个考虑。

四、C++ 侧的 Dataset 基类:BatchDataset / Dataset / DatasetPtr

Python 的MNN.data.Dataset对应 C++ 的MNN::Train::Dataset,定义在 tools/train/source/data/Dataset.hpp:

class MNN_PUBLIC BatchDataset { public: virtual ~BatchDataset() = default; // get batch using given indices virtual std::vector<Example> getBatch(std::vector<size_t> indices) = 0; // size of the dataset virtual size_t size() = 0; }; class MNN_PUBLIC Dataset : public BatchDataset { public: // return a specific example with given index virtual Example get(size_t index) = 0; std::vector<Example> getBatch(std::vector<size_t> indices) { std::vector<Example> batch; batch.reserve(indices.size()); for (const auto i : indices) { batch.emplace_back(get(i)); } MNN_ASSERT(batch.size() != 0); return batch; } };

这里的层次设计很清晰:

  • BatchDataset是最底层抽象,只要求实现"按索引列表取一批"和"数据集大小";
  • Dataset继承它并只要求实现单样本的get(index)getBatch默认实现为循环调用get收集——DatasetWrapper覆写的正是这个get,所以 Python 实现者只需关心单样本;
  • 外层的DatasetPtr(Dataset.hpp)持有std::shared_ptr<BatchDataset>并提供便捷方法createLoader(batchSize, stack, shuffle, numWorkers),把数据集直接转成 DataLoader。

五、与 DataLoader 配合:批量、打乱与多 worker 取数

Dataset 只是"单样本来源",真正的批量流水线由DataLoader承担。Python 侧构造签名见 pymnn/src/data.h:

static PyObject* PyMNNDataLoader_new(PyTypeObject *type, PyObject *args, PyObject *kwargs) { PyObject* dataset = nullptr; int batch_size, num_workers = 0; int shuffle = 1; static char *kwlist[] = { "dataset", "batch_size", "shuffle", "num_workers", NULL }; if (!PyArg_ParseTupleAndKeywords(args, kwargs, "Oi|ii", kwlist, &dataset, &batch_size, &shuffle, &num_workers)) { ... } std::shared_ptr<Dataset> dataset_ = std::move(toDataset(dataset)); ... self->ptr = DataLoader::makeDataLoader(dataset_, batch_size, true, shuffle, num_workers); return (PyObject*)self; }

对照 C++ 侧的工厂方法签名(tools/train/source/data/DataLoader.hpp)可以读出完整参数语义:

参数含义默认值
datasetDataset 实例(Python)或shared_ptr<BatchDataset>(C++)必填
batch_size每个 batch 的样本数必填
stack是否把样本沿第 0 维 stack 成[n, ...]批量张量true
shuffle是否使用随机采样器打乱顺序true(Python 绑定中为 1)
num_workers后台预取 worker 线程数0(同步取数)

Python 绑定固定传stack=true,即[1, 28, 28]的样本会被堆叠成[n, 1, 28, 28]——这与 dataLoaderDemo.cpp 中注释的// the stack transform, stack [1, 28, 28] to [n, 1, 28, 28]是同一行为。

DataLoader 内部(DataLoader.hpp)使用BlockingQueue作为任务队列和数据队列,并维护一组std::threadworker:num_workers > 0时样本在后台线程中预取,主线程调next()只负责消费,从而把 Python__getitem__的回调开销与训练计算重叠起来。对外暴露的接口为next()reset()iterNumber()size();Python 绑定对应暴露了iter_numbersize两个只读属性和reset()next()两个方法(pymnn/src/data.h)。

一个典型的训练循环形态为:

dataset = MnistDataset(training_dataset=True) loader = MNN.data.DataLoader(dataset, batch_size=7, shuffle=1, num_workers=4) for i in range(loader.iter_number): example = loader.next() # (inputs: [Var...], targets: [Var...]) # 前向、计算 loss、反传 ... loader.reset() # 数据集耗尽后需 reset 采样器内部状态

reset()的作用在 C++ demo 中有明确注释(dataLoaderDemo.cpp):

// this will reset the sampler's internal state, not necessary here trainDataLoader->reset(); // this will reset the sampler's internal state, necessary here, because the test dataset is exhausted testDataLoader->reset();

即遍历完一个 epoch 后必须reset()重新采样,否则next()拿不到数据。

六、对照 C++ 训练 demo:同一套数据管线的原生写法

tools/train/source/demo/dataLoaderDemo.cpp 展示了不经过 Python 时相同管线的完整用法,可作为理解 Dataset/DataLoader 协作的参照:

// 归一化 transform:0~255 -> 0~1(单样本级) static Example func(Example example) { auto cast = _Cast(example.first[0], halide_type_of<float>()); example.first[0] = _Multiply(cast, _Const(1.0f / 255.0f)); return example; } ... auto trainDataLoader = std::shared_ptr<DataLoader>( DataLoader::makeDataLoader(trainDataset, {trainTransform}, trainBatchSize, true, trainNumWorkers)); ... for (int i = 0; i < iterations; i++) { auto trainData = trainDataLoader->next(); auto data = trainData[0].first[0]->readMap<float>(); auto label = trainData[0].second[0]->readMap<uint8_t>(); cout << "index: " << i << " train label: " << int(label[0]) << endl; }

这里trainData[0].first[0]取的就是 batch 中第 0 个样本的输入 Var,readMap<float>()直接读出指针级内存——Example[Var], [Var]结构在整个 C++/Python 两侧是完全一致的。更小的验证程序见 tools/train/source/demo/dataLoaderTest.cpp,纯 C++ 数据集实现与 MNIST 加载逻辑可参考tools/train/source/demo下的示例文件;C++ 侧数据集文档见 docs/train/data.md,DataLoader 的 Python API 说明见 docs/pymnn/DataLoader.md。

七、实现自定义 Dataset 的检查清单

综合文档约定与桥接源码,实现一个可用的MNN.data.Dataset时逐项确认:

  1. 继承MNN.data.Dataset并调用super().__init__()(Python 3 风格写作super()亦可);
  2. __len__返回数据集总样本数(int);
  3. __getitem__(index)返回(inputs, targets)元组,两侧都是 list,即使只有一路输入/目标也写成[dv], [dl]
  4. 元素必须是MNN.expr的 Var(expr.const构造),并指定正确的形状、布局(expr.NCHW)与 dtype(浮点输入、expr.uint8标签);
  5. 大对象(原始图像数组等)在__init__中一次性加载/缓存,__getitem__保持轻量,因为它会被 worker 线程高频回调;
  6. 训练用shuffle=1,验证/测试用shuffle=0;epoch 结束后调用reset()
  7. 通过MNN.data.DataLoader(dataset, batch_size, shuffle, num_workers)消费,batch_size决定 stack 后张量的第 0 维大小。

这样实现出的 Dataset 与 MNN 训练栈(优化器、损失、autograd)即可无缝衔接:DataLoader.next()产出的Example就是Example类型的[Var], [Var]结构,可直接作为网络输入参与前向与反传。

【免费下载链接】MNNMNN: A blazing-fast, lightweight inference engine battle-tested by Alibaba, powering high-performance on-device LLMs and Edge AI.项目地址: https://gitcode.com/GitHub_Trending/mn/MNN

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/14 22:04:44

地界范围图批量导出实战:模板化出图全流程解析

干了这么多年测绘和GIS内业&#xff0c;我最怕的不是外业跑杆&#xff0c;而是项目收尾时那一堆出图任务。尤其是在土地确权、土地整治、造林工程这类项目里&#xff0c;每个地块都要配一张地界范围图&#xff0c;少则几十张&#xff0c;多则上千张。以前在CAD或者ArcGIS里一张…

作者头像 李华
网站建设 2026/9/14 22:04:42

OpenGL体渲染实战:nii医学数据到Qt/PyQt5集成全指南

提到 OpenGL&#xff0c;很多老图形程序员会心一笑&#xff0c;很多新手则一头雾水。作为一门拥有跨平台影响力的图形 API&#xff0c;OpenGL 从 90 年代活到今天&#xff0c;依然是医学可视化、CAD、仿真、Qt 桌面应用里最常见的技术底座。我在做医学影像渲染和桌面工具时和它…

作者头像 李华
网站建设 2026/9/14 22:03:02

SpringBoot 3.x整合Swagger实现API文档自动化

1. SpringBoot 3.x整合Swagger的必要性在现代Web应用开发中&#xff0c;API文档的维护一直是个痛点。传统的手写文档方式存在更新不及时、格式不统一等问题。Swagger作为一套开源的API文档工具链&#xff0c;通过注解方式自动生成可视化文档&#xff0c;完美解决了这些问题。Sp…

作者头像 李华
网站建设 2026/9/14 22:02:59

混合动力汽车油耗计算的动态规划算法与MATLAB实现

1. 混合动力汽车油耗计算的核心挑战混合动力汽车&#xff08;HEV&#xff09;的油耗计算一直是汽车工程领域的难点问题。与传统燃油车不同&#xff0c;HEV同时具备发动机和电机两套动力系统&#xff0c;能量流动路径复杂多变。我在参与某插电混动车型开发时&#xff0c;发现传统…

作者头像 李华