PyTorch C++前端使用MNIST数据集子集的实现问题
PyTorch C++前端实现MNIST数据集子集的解决方法
问题场景
想在PyTorch C前端使用MNIST的精简数据集,但C端没有Python版的torch.utils.data.Subset或RandomizedSubsetSampler这类工具。按照官方MNIST示例加载数据集后,自定义了SubsetSampler却出现编译错误。
加载数据集的代码:
auto train_dataset = torch::data::datasets::MNIST(kDataRoot) .map(torch::data::transforms::Normalize<>(0.1307, 0.3081)) .map(torch::data::transforms::Stack<>()); const size_t train_dataset_size = train_dataset.size().value();
自定义的SubsetSampler代码:
class SubsetSampler : public torch::data::samplers::Sampler<> { public: explicit SubsetSampler(std::vector<size_t> indices) : indices_(std::move(indices)) {} // Return the next batch of indices. c10::optional<std::vector<size_t>> next(size_t batch_size) override { std::vector<size_t> batch; while (batch.size() < batch_size && current_ < indices_.size()) { batch.push_back(indices_[current_++]); } if (batch.empty()) { return c10::nullopt; // No more data. } return batch; } // Reset the sampler's state. void reset() { current_ = 0; } // Return the total number of samples. c10::optional<size_t> size() { return indices_.size(); } private: std::vector<size_t> indices_; size_t current_ = 0; };
编译时出现的错误:
error: no type named ‘BatchRequestType’ in ‘class std::shared_ptr<SubsetSampler>’ 23 | class StatelessDataLoader : public DataLoaderBase< | ^~~~~~~~~~~~~~~~~~~
解决方法
方法一:修复自定义SubsetSampler的实现
编译报错的核心原因是自定义Sampler没有满足torch::data::samplers::Sampler基类的要求:
- 必须公开
BatchRequestType类型别名,指定采样返回的批次索引类型 - 虚函数
reset()和size()需要加上override关键字,确保正确重写基类方法
修正后的SubsetSampler代码:
class SubsetSampler : public torch::data::samplers::Sampler<> { public: // 必须添加的类型别名,指定批次请求的类型 using BatchRequestType = std::vector<size_t>; explicit SubsetSampler(std::vector<size_t> indices) : indices_(std::move(indices)) {} c10::optional<std::vector<size_t>> next(size_t batch_size) override { std::vector<size_t> batch; while (batch.size() < batch_size && current_ < indices_.size()) { batch.push_back(indices_[current_++]); } return batch.empty() ? c10::nullopt : c10::make_optional(batch); } // 加上override关键字重写基类方法 void reset() override { current_ = 0; } // 加上override关键字重写基类方法 c10::optional<size_t> size() override { return indices_.size(); } private: std::vector<size_t> indices_; size_t current_ = 0; };
使用时,创建Sampler实例并传入DataLoader:
// 生成想要的子集索引,比如取前1000个样本 std::vector<size_t> subset_indices; subset_indices.reserve(1000); for (size_t i = 0; i < 1000; ++i) { subset_indices.push_back(i); } auto sampler = std::make_shared<SubsetSampler>(subset_indices); auto train_loader = torch::data::make_data_loader( std::move(train_dataset), torch::data::DataLoaderOptions().batch_size(64).sampler(sampler) );
方法二:直接实现Subset数据集(更贴近Python用法)
如果不想修改Sampler,可以直接封装一个SubsetDataset类,继承自PyTorch的Dataset接口,直接对原数据集做子集封装,逻辑更直观:
template <typename Dataset> class SubsetDataset : public torch::data::datasets::Dataset<SubsetDataset<Dataset>> { public: using SampleType = typename Dataset::SampleType; // 构造函数:传入原数据集和子集索引 SubsetDataset(Dataset dataset, std::vector<size_t> indices) : dataset_(std::move(dataset)), indices_(std::move(indices)) {} // 获取指定索引的样本 SampleType get(size_t index) override { return dataset_.get(indices_[index]); } // 返回子集的大小 c10::optional<size_t> size() const override { return indices_.size(); } private: Dataset dataset_; std::vector<size_t> indices_; };
使用示例:
// 生成子集索引,比如随机选2000个样本 std::vector<size_t> subset_indices; subset_indices.reserve(2000); std::mt19937 rng(std::random_device{}()); std::uniform_int_distribution<size_t> dist(0, train_dataset_size - 1); for (size_t i = 0; i < 2000; ++i) { subset_indices.push_back(dist(rng)); } // 创建子集数据集 auto subset_dataset = SubsetDataset(std::move(train_dataset), subset_indices); // 正常创建DataLoader auto train_loader = torch::data::make_data_loader( std::move(subset_dataset), torch::data::DataLoaderOptions().batch_size(64) );
内容的提问来源于stack exchange,提问作者IlBowsta
相关产品推荐
相关产品推荐

