/
vuron
/
adept
Обзор
Документация
Войти
/
vuron
/
adept
Код
Запросы
0
Задачи
Вики
Пакеты
0
Релизы
0
Аналитика
Безопасность
master
src/data/dataloader.cpp
91 строка
2 KB
kolkir
Refactor loops
08 фев 2025, 13:56
08 фев 2025, 13:56
b73a638
Код
Авторство
О чём код?
#include <adept/data/dataloader.hpp> #include <adept/irange.hpp> namespace adept { namespace detail { RandomSampler::RandomSampler(index_t size) : indicies_(size) { std::iota(indicies_.begin(), indicies_.end(), 0); reset(); } void RandomSampler::reset() { std::shuffle(indicies_.begin(), indicies_.end(), rng_); current_index_ = 0; } std::optional<std::vector<index_t>> RandomSampler::next(index_t batch_size) { if (current_index_ >= indicies_.size()) return std::nullopt; std::vector<index_t> batch_indicies; batch_indicies.reserve(batch_size); for (index_t i = 0; current_index_ < indicies_.size() && i < batch_size; ++current_index_, ++i) { batch_indicies.push_back(current_index_); } return batch_indicies; } } // namespace detail DataLoader::DataLoader(dataset_ptr_t dataset, index_t batch_size) : dataset_(dataset), batch_size_(batch_size) { if (!dataset) { THROW_ERROR("Dataloader can't be initialized with nullptr dataset"); } sampler_ = std::make_unique<detail::RandomSampler>(dataset->size()); } DataLoader::DataLoader(Dataset* dataset, index_t batch_size) : dataset_(std::shared_ptr<Dataset>(dataset, [](Dataset*) {})), batch_size_(batch_size) { if (!dataset) { THROW_ERROR("Dataloader can't be initialized with nullptr dataset"); } sampler_ = std::make_unique<detail::RandomSampler>(dataset->size()); } DataLoader::~DataLoader() {} Iterator<batch_t> DataLoader::begin() { reset(); return Iterator<batch_t>( std::make_unique<detail::ValidIterator<batch_t>>([this] { return this->next(); })); } Iterator<batch_t> DataLoader::end() const { return Iterator<batch_t>(std::make_unique<detail::EndIterator<batch_t>>()); } void DataLoader::reset() { sampler_->reset(); } std::optional<batch_t> DataLoader::next() { auto batch_indicies = sampler_->next(batch_size_); if (!batch_indicies) return std::nullopt; std::vector<batch_t> batch_items_(dataset_->item_size()); for (auto& b : batch_items_) { b.reserve(batch_indicies->size()); } for (auto i : irange(batch_indicies->size())) { auto item = dataset_->item((*batch_indicies)[i]); for (auto j : irange(dataset_->item_size())) { batch_items_[j].push_back(item[j]); } } // stack batch items batch_t batch; batch.reserve(dataset_->item_size()); for (auto j : irange(dataset_->item_size())) { batch.push_back(Tensor::stack(batch_items_[j])); } return batch; } size_t DataLoader::size() const { return dataset_->size() / batch_size_; } } // namespace adept