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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 19:22:42