| 36 | } |
| 37 | |
| 38 | virtual int run(int argc, const char* argv[]) override { |
| 39 | if (argc != 2) { |
| 40 | cout << "usage: ./runTrainDemo.out DataLoaderTest /path/to/unzipped/mnist/data/" << endl; |
| 41 | return 0; |
| 42 | } |
| 43 | |
| 44 | const int testCount = 6; |
| 45 | int passedTestCount = 0; |
| 46 | |
| 47 | std::string root = argv[1]; |
| 48 | |
| 49 | // train data loader |
| 50 | const size_t trainDatasetSize = 60000; |
| 51 | auto trainDataset = MnistDataset::create(root, MnistDataset::Mode::TRAIN); |
| 52 | |
| 53 | auto trainSampler = std::make_shared<RandomSampler>(trainDataset.get<MnistDataset>()->size()); |
| 54 | |
| 55 | const size_t trainBatchSize = 7; |
| 56 | const size_t trainNumWorkers = 4; |
| 57 | auto trainConfig = std::make_shared<DataLoaderConfig>(trainBatchSize, trainNumWorkers); |
| 58 | |
| 59 | DataLoader trainDataLoader(trainDataset.mDataset, trainSampler, trainConfig); |
| 60 | |
| 61 | auto images = trainDataset.get<MnistDataset>()->images(); |
| 62 | auto labels = trainDataset.get<MnistDataset>()->labels(); |
| 63 | const int32_t kImageRows = 28; |
| 64 | const int32_t kImageColumns = 28; |
| 65 | |
| 66 | const size_t iterations = trainDatasetSize / trainBatchSize; |
| 67 | |
| 68 | auto samplerIndices = trainSampler->indices(); |
| 69 | sort(samplerIndices.begin(), samplerIndices.end()); |
| 70 | for (int i = 0; i < samplerIndices.size(); i++) { |
| 71 | MNN_ASSERT(samplerIndices[i] == i); |
| 72 | } |
| 73 | |
| 74 | for (int i = 0; i < iterations; i++) { |
| 75 | auto trainData = trainDataLoader.next(); |
| 76 | |
| 77 | for (int j = 0; j < trainData.size(); j++) { |
| 78 | auto index = int(trainData[j].first[1]->readMap<float>()[0]); |
| 79 | |
| 80 | auto data = trainData[j].first[0]->readMap<uint8_t>(); |
| 81 | auto label = trainData[j].second[0]->readMap<uint8_t>(); |
| 82 | |
| 83 | auto trueData = images->readMap<uint8_t>() + kImageRows * kImageColumns * index; |
| 84 | auto trueLabel = labels->readMap<uint8_t>() + index; |
| 85 | |
| 86 | for (int k = 0; k < kImageRows * kImageColumns; k++) { |
| 87 | MNN_ASSERT(data[k] == trueData[k]); |
| 88 | } |
| 89 | MNN_ASSERT(label[0] == trueLabel[0]); |
| 90 | } |
| 91 | } |
| 92 | trainDataLoader.clean(); |
| 93 | |
| 94 | passedTestCount++; |
| 95 | cout << "[" << passedTestCount << " / " << testCount << "] passed." << endl; |