Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
RBatchLoader.cxx
Go to the documentation of this file.
1
3
4#include <algorithm>
5#include <numeric>
6#include <utility>
7
9
10RBatchLoader::RBatchLoader(std::size_t batchSize, const std::vector<std::string> &cols, std::mutex &sharedMutex,
11 std::condition_variable &sharedCV, const std::vector<std::size_t> &vecSizes,
12 std::size_t numEntries, bool dropRemainder)
13 : fBatchSize(batchSize),
14 fCols(cols),
15 fLock(sharedMutex),
16 fCV(sharedCV),
17 fVecSizes(vecSizes),
18 fNumEntries(numEntries),
19 fDropRemainder(dropRemainder)
20{
21 fSumVecSizes = std::accumulate(fVecSizes.begin(), fVecSizes.end(), 0);
22 fNumColumns = fCols.size() + fSumVecSizes - fVecSizes.size();
23
24 RecalculateBatchCounts(numEntries);
25
26 fPrimaryLeftoverBatch = std::make_unique<RFlat2DMatrix>();
27 fSecondaryLeftoverBatch = std::make_unique<RFlat2DMatrix>();
28}
29
30/// \brief Activate the batchloader. This means that batches can be created and loaded.
32{
33 {
34 std::lock_guard<std::mutex> lock(fLock);
35 if (fIsActive)
36 return;
37 fIsActive = true;
38 fProducerDone = false;
39 }
40
41 fCV.notify_all();
42}
43
44/// \brief DeActivate the batchloader. This means that no more batches are created.
45/// Batches can still be returned if they are already loaded.
47{
48 {
49 std::lock_guard<std::mutex> lock(fLock);
50 if (!fIsActive)
51 return;
52 fIsActive = false;
53 }
54
55 fCV.notify_all();
56}
57
58/// \brief Reset the batchloader state.
60{
61 {
62 std::lock_guard<std::mutex> lock(fLock);
63
64 while (!fBatchQueue.empty()) {
65 fBatchQueue.pop();
66 }
67
68 fCurrentBatch.reset();
69 fPrimaryLeftoverBatch = std::make_unique<RFlat2DMatrix>();
70 fSecondaryLeftoverBatch = std::make_unique<RFlat2DMatrix>();
71 }
72
73 fCV.notify_all();
74}
75
76/// \brief Signal that the producer has finished pushing all batches for this epoch.
78{
79 fProducerDone = true;
80 fCV.notify_all();
81}
82
83/// \brief Return a batch of data as a unique pointer.
84/// After the batch has been processed, it should be destroyed.
85/// \param[in] chunkTensor Tensor with the data from the chunk
86/// \param[in] idxs Index of batch in the chunk
87/// \return Batch
88std::unique_ptr<RFlat2DMatrix> RBatchLoader::CreateBatch(RFlat2DMatrix &chunkTensor, std::size_t idxs)
89{
90 auto batch = std::make_unique<RFlat2DMatrix>(fBatchSize, fNumColumns);
91 std::copy(chunkTensor.GetData() + (idxs * fBatchSize * fNumColumns),
92 chunkTensor.GetData() + ((idxs + 1) * fBatchSize * fNumColumns), batch->GetData());
93
94 return batch;
95}
96
97/// \brief Loading the batch from the queue.
98/// \return Batch
100{
101 std::unique_lock<std::mutex> lock(fLock);
102
103 // Wait until:
104 // - there is data in the queue
105 // - or producer declares "done"
106 // - or we are deactivated
107 fCV.wait(lock, [&] { return !fBatchQueue.empty() || fProducerDone || !fIsActive; });
108
109 if (fBatchQueue.empty()) {
110 // producer done and no queued data -> end-of-epoch signal
111 fCurrentBatch = std::make_unique<RFlat2DMatrix>();
112 return *fCurrentBatch;
113 }
114
115 fCurrentBatch = std::move(fBatchQueue.front());
116 fBatchQueue.pop();
117 // Notify the loading thread that the queue has drained
118 fCV.notify_all();
119
120 // move the buffer over to the caller
121 return std::move(*fCurrentBatch);
122}
123
124/// \brief Creating the batches from a chunk and add them to the queue.
125/// \param[in] chunkTensor Tensor with the data from the chunk
126/// \param[in] isLastBatch Check if the batch in the chunk is the last one
127void RBatchLoader::CreateBatches(RFlat2DMatrix &chunkTensor, bool isLastBatch)
128{
129 std::size_t chunkSize = chunkTensor.GetRows();
130 std::size_t numCols = chunkTensor.GetCols();
131 std::size_t numBatches = chunkSize / fBatchSize;
132 std::size_t leftoverBatchSize = chunkSize % fBatchSize;
133
134 // create a vector of batches
135 std::vector<std::unique_ptr<RFlat2DMatrix>> batches;
136
137 // fill the full batches from the chunk into a vector
138 for (std::size_t i = 0; i < numBatches; i++) {
139 batches.emplace_back(CreateBatch(chunkTensor, i));
140 }
141
142 // copy the remaining entries from the chunk into a leftover batch
143 RFlat2DMatrix LeftoverBatch(leftoverBatchSize, numCols);
144 std::copy(chunkTensor.GetData() + (numBatches * fBatchSize * numCols),
145 chunkTensor.GetData() + (numBatches * fBatchSize * numCols + leftoverBatchSize * numCols),
146 LeftoverBatch.GetData());
147
148 // calculate how many empty slots are left in fPrimaryLeftoverBatch
149 std::size_t PrimaryLeftoverSize = fPrimaryLeftoverBatch->GetRows();
150 std::size_t emptySlots = fBatchSize - PrimaryLeftoverSize;
151
152 // copy LeftoverBatch to end of fPrimaryLeftoverBatch
153 if (emptySlots >= leftoverBatchSize) {
154 fPrimaryLeftoverBatch->Resize(PrimaryLeftoverSize + leftoverBatchSize, numCols);
155 std::copy(LeftoverBatch.GetData(), LeftoverBatch.GetData() + (leftoverBatchSize * fNumColumns),
156 fPrimaryLeftoverBatch->GetData() + (PrimaryLeftoverSize * numCols));
157
158 // copy LeftoverBatch to end of fPrimaryLeftoverBatch and add it to the batch
159 if (emptySlots == leftoverBatchSize) {
160 auto copy = std::make_unique<RFlat2DMatrix>(fBatchSize, fNumColumns);
161 std::copy(fPrimaryLeftoverBatch->GetData(), fPrimaryLeftoverBatch->GetData() + (fBatchSize * fNumColumns),
162 copy->GetData());
163 batches.emplace_back(std::move(copy));
164
165 // reset fPrimaryLeftoverBatch and fSecondaryLeftoverBatch
167 fSecondaryLeftoverBatch = std::make_unique<RFlat2DMatrix>();
168 }
169 }
170
171 // copy LeftoverBatch to both fPrimaryLeftoverBatch and fSecondaryLeftoverBatch
172 else if (emptySlots < leftoverBatchSize) {
173 // copy the first part of LeftoverBatch to end of fPrimaryLeftoverTrainingBatch
174 fPrimaryLeftoverBatch->Resize(fBatchSize, numCols);
175 std::copy(LeftoverBatch.GetData(), LeftoverBatch.GetData() + (emptySlots * numCols),
176 fPrimaryLeftoverBatch->GetData() + (PrimaryLeftoverSize * numCols));
177
178 // copy the last part of LeftoverBatch to the end of fSecondaryLeftoverBatch
179 fSecondaryLeftoverBatch->Resize(leftoverBatchSize - emptySlots, numCols);
180 std::copy(LeftoverBatch.GetData() + (emptySlots * numCols),
181 LeftoverBatch.GetData() + (leftoverBatchSize * numCols), fSecondaryLeftoverBatch->GetData());
182
183 // add fPrimaryLeftoverBatch to the batch vector
184 auto copy = std::make_unique<RFlat2DMatrix>(fBatchSize, fNumColumns);
185 std::copy(fPrimaryLeftoverBatch->GetData(), fPrimaryLeftoverBatch->GetData() + (fBatchSize * fNumColumns),
186 copy->GetData());
187 batches.emplace_back(std::move(copy));
188
189 // exchange fPrimaryLeftoverBatch and fSecondaryLeftoverBatch
191 // reset fSecondaryLeftoverTrainingBatch
192 fSecondaryLeftoverBatch = std::make_unique<RFlat2DMatrix>();
193 }
194
195 // copy the content of fPrimaryLeftoverBatch to the leftover batch from the chunk
196 if (isLastBatch) {
198 auto copy = std::make_unique<RFlat2DMatrix>(fLeftoverBatchSize, fNumColumns);
199 std::copy(fPrimaryLeftoverBatch->GetData(),
200 fPrimaryLeftoverBatch->GetData() + (fLeftoverBatchSize * fNumColumns), copy->GetData());
201 batches.emplace_back(std::move(copy));
202 }
203
204 fPrimaryLeftoverBatch = std::make_unique<RFlat2DMatrix>();
205 fSecondaryLeftoverBatch = std::make_unique<RFlat2DMatrix>();
206 }
207
208 {
209 std::lock_guard<std::mutex> lock(fLock);
210 for (auto &batch : batches) {
211 fBatchQueue.push(std::move(batch));
212 }
213 }
214
215 fCV.notify_all();
216}
217
218/// \brief Recalculate batch counts from the given number of entries.
219/// Used at construction or when the true entry count is discovered lazily (filtered case).
220void RBatchLoader::RecalculateBatchCounts(std::size_t numEntries)
221{
222 fNumEntries = numEntries;
223
224 if (fBatchSize == 0) {
226 }
227
230
231 const std::size_t numLeftoverBatches = fLeftoverBatchSize == 0 ? 0 : 1;
233}
234} // namespace ROOT::Experimental::Internal::ML
RFlat2DMatrix GetBatch()
Loading the batch from the queue.
std::unique_ptr< RFlat2DMatrix > fSecondaryLeftoverBatch
std::queue< std::unique_ptr< RFlat2DMatrix > > fBatchQueue
void RecalculateBatchCounts(std::size_t numEntries)
Recalculate batch counts from the given number of entries.
void Reset()
Reset the batchloader state.
void Activate()
Activate the batchloader. This means that batches can be created and loaded.
void MarkProducerDone()
Signal that the producer has finished pushing all batches for this epoch.
void CreateBatches(RFlat2DMatrix &chunkTensor, bool isLastBatch)
Creating the batches from a chunk and add them to the queue.
void DeActivate()
DeActivate the batchloader.
std::unique_ptr< RFlat2DMatrix > fPrimaryLeftoverBatch
std::unique_ptr< RFlat2DMatrix > fCurrentBatch
RBatchLoader(std::size_t batchSize, const std::vector< std::string > &cols, std::mutex &sharedMutex, std::condition_variable &sharedCV, const std::vector< std::size_t > &vecSizes={}, std::size_t numEntries=0, bool dropRemainder=false)
std::unique_ptr< RFlat2DMatrix > CreateBatch(RFlat2DMatrix &chunkTensor, std::size_t idxs)
Return a batch of data as a unique pointer.
Wrapper around ROOT::RVec<float> representing a 2D matrix.