Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
RDataLoaderEngine.hxx
Go to the documentation of this file.
1// Author: Dante Niewenhuis, VU Amsterdam 07/2023
2// Author: Kristupas Pranckietis, Vilnius University 05/2024
3// Author: Nopphakorn Subsa-Ard, King Mongkut's University of Technology Thonburi (KMUTT) (TH) 08/2024
4// Author: Vincenzo Eduardo Padulano, CERN 10/2024
5// Author: Martin Føll, University of Oslo (UiO) & CERN 01/2026
6// Author: Silia Taider, CERN 02/2026
7
8/*************************************************************************
9 * Copyright (C) 1995-2026, Rene Brun and Fons Rademakers. *
10 * All rights reserved. *
11 * *
12 * For the licensing terms see $ROOTSYS/LICENSE. *
13 * For the list of contributors see $ROOTSYS/README/CREDITS. *
14 *************************************************************************/
15
16#ifndef ROOT_INTERNAL_ML_RDATALOADERENGINE
17#define ROOT_INTERNAL_ML_RDATALOADERENGINE
18
19#include <algorithm>
20#include <condition_variable>
21#include <memory>
22#include <mutex>
23#include <string>
24#include <string_view>
25#include <thread>
26#include <vector>
27
34#include "ROOT/ML/RSampler.hxx"
36
37// Empty namespace to create a hook for the Pythonization
39}
40
42/**
43 \class ROOT::Experimental::Internal::ML::RDataLoaderEngine
44\brief
45
46In this class, the processes of loading clusters (see RClusterLoader) and creating batches from those clusters (see
47RBatchLoader) are combined, allowing batches from the training and validation sets to be loaded directly from a dataset
48in an RDataFrame.
49*/
50
51template <typename... Args>
53private:
54 std::vector<std::string> fCols;
55 std::vector<std::size_t> fVecSizes;
56 std::size_t fBatchSize;
57 std::size_t fSetSeed;
58
59 // buffer quantities
60 std::size_t fBatchesInMemory;
61 std::size_t fBufferCapacity;
62 std::size_t fLowWatermark;
63 std::size_t fHighWatermark;
64
65 std::size_t fTrainingClusterIdx{0};
66 std::size_t fValidationClusterIdx{0};
67
68 float fTestSize;
69
70 std::unique_ptr<RDatasetLoader<Args...>> fDatasetLoader;
71 std::unique_ptr<RClusterLoader<Args...>> fClusterLoader;
72 std::unique_ptr<RBatchLoader> fTrainingBatchLoader;
73 std::unique_ptr<RBatchLoader> fValidationBatchLoader;
74 std::unique_ptr<RSampler> fTrainingSampler;
75 std::unique_ptr<RSampler> fValidationSampler;
76
77 std::unique_ptr<RFlat2DMatrixOperators> fTensorOperators;
78
79 std::vector<ROOT::RDF::RNode> fRdfs;
80
81 std::unique_ptr<std::thread> fLoadingThread;
82 std::condition_variable fLoadingCondition;
83 std::mutex fLoadingMutex;
84
88 std::string fSampleType;
91
92 bool fIsActive{false}; // Whether the loading thread is active
93
94 bool fEpochActive{false};
97
100
101 // flattened buffers for chunks and temporary tensors (rows * cols)
102 std::vector<RFlat2DMatrix> fTrainingDatasets;
103 std::vector<RFlat2DMatrix> fValidationDatasets;
104
107
110
111 std::size_t fTrainingEpochCount{0};
112 std::size_t fValidationEpochCount{0};
113
114 /// \brief Describe how the loader's columns map onto a batch-tensor row.
115 std::vector<RColumnLayout> MakeColumnLayout() const
116 {
117 std::vector<RColumnLayout> layout;
118 layout.reserve(sizeof...(Args));
119
120 std::size_t colIdx = 0;
121 std::size_t vecIdx = 0;
122 std::size_t offset = 0;
123 (
124 [&] {
126 const std::size_t width = isVector ? fVecSizes[vecIdx++] : 1;
127 layout.push_back({fCols[colIdx++], offset, width, isVector});
128 offset += width;
129 }(),
130 ...);
131
132 return layout;
133 }
134
135 /// \brief Opens a training or validation epoch and closes it again when done
136 struct REpochGuard {
139
140 REpochGuard(RDataLoaderEngine &engine, bool isTraining) : fEngine(engine), fIsTraining(isTraining)
141 {
142 // Same order as the pythonization's epoch context managers
144 if (fIsTraining) {
147 } else {
150 }
151 }
152
160 };
161
162public:
163 RDataLoaderEngine(const std::vector<ROOT::RDF::RNode> &rdfs, const std::size_t batchSize,
164 const std::size_t batchesInMemory, const std::vector<std::string> &cols,
165 const std::vector<std::size_t> &vecSizes = {}, const float vecPadding = 0.0,
166 const float testSize = 0.0, bool shuffle = true, bool dropRemainder = true,
167 const std::size_t setSeed = 0, bool loadEager = false, std::string sampleType = "",
168 float sampleRatio = 1.0, bool replacement = false)
169 : fRdfs(rdfs),
170 fCols(cols),
171 fVecSizes(vecSizes),
172 fBatchSize(batchSize),
173 fBatchesInMemory(batchesInMemory),
174 fTestSize(testSize),
175 fDropRemainder(dropRemainder),
176 fSetSeed(setSeed),
177 fShuffle(shuffle),
178 fLoadEager(loadEager),
179 fSampleType(sampleType),
180 fSampleRatio(sampleRatio),
181 fReplacement(replacement)
182 {
183 fTensorOperators = std::make_unique<RFlat2DMatrixOperators>(fShuffle, fSetSeed);
184
185 if (fLoadEager) {
186 fDatasetLoader = std::make_unique<RDatasetLoader<Args...>>(fRdfs, fTestSize, fCols, fVecSizes, vecPadding,
188 fDatasetLoader->SplitDatasets();
189
190 if (fSampleType == "") {
191 fDatasetLoader->ConcatenateDatasets();
192
193 fTrainingDataset = fDatasetLoader->ReleaseTrainingDataset();
194 fValidationDataset = fDatasetLoader->ReleaseValidationDataset();
195
196 fNumTrainingEntries = fDatasetLoader->GetNumTrainingEntries();
197 fNumValidationEntries = fDatasetLoader->GetNumValidationEntries();
198 }
199
200 else {
201 fTrainingDatasets = fDatasetLoader->ReleaseTrainingDatasets();
202 fValidationDatasets = fDatasetLoader->ReleaseValidationDatasets();
203
206 fValidationSampler = std::make_unique<RSampler>(fValidationDatasets, fSampleType, fSampleRatio,
208
209 fNumTrainingEntries = fTrainingSampler->GetNumEntries();
210 fNumValidationEntries = fValidationSampler->GetNumEntries();
211 }
212
213 // the dataset is in memory now, release the objects used to create it
214 fDatasetLoader.reset();
215 fRdfs.clear();
216 }
217
218 else {
219 // scan cluster boundaries
220 fClusterLoader = std::make_unique<RClusterLoader<Args...>>(fRdfs, fCols, fVecSizes, vecPadding, fTestSize,
222
223 // derive buffer quantities
225 // at least one batch, otherwise the refill threshold rounds down to 0 and nothing is ever loaded
228
229 // split cluster list into training and validation
230 fClusterLoader->SplitDataset();
231 fNumTrainingEntries = fClusterLoader->GetNumTrainingEntries();
232 fNumValidationEntries = fClusterLoader->GetNumValidationEntries();
233 }
234
235 fTrainingBatchLoader = std::make_unique<RBatchLoader>(fBatchSize, fCols, fLoadingMutex, fLoadingCondition,
237 fValidationBatchLoader = std::make_unique<RBatchLoader>(fBatchSize, fCols, fLoadingMutex, fLoadingCondition,
239 }
240
242
244 {
245 {
246 std::lock_guard<std::mutex> lock(fLoadingMutex);
247 if (!fIsActive)
248 return;
249 fIsActive = false;
250 }
251
252 fLoadingCondition.notify_all();
253
254 if (fLoadingThread) {
255 if (fLoadingThread->joinable()) {
256 fLoadingThread->join();
257 }
258 }
259
260 fLoadingThread.reset();
261 }
262
263 /// \brief Activate the loading process by spawning the loading thread.
264 void Activate()
265 {
266 {
267 std::lock_guard<std::mutex> lock(fLoadingMutex);
268 if (fIsActive)
269 return;
270
271 fIsActive = true;
272 }
273
274 if (fLoadEager) {
275 return;
276 }
277
278 fLoadingThread = std::make_unique<std::thread>(&RDataLoaderEngine::LoadData, this);
279 }
280
281 /// \brief Materialize one train/test split to disk by draining a full epoch through the normal batch
282 /// pipeline and Fill() each batch into \p filename instead of yielding it.
283 ///
284 /// Filters, shuffling, the train/validation split and the batch_size/drop_remainder settings
285 /// are all inherited from the loader's configuration.
286 /// \param outputFormat Either "ttree" or "rntuple".
287 void Save(std::string_view dataset_name, std::string_view filename, bool isTraining, std::string_view outputFormat)
288 {
289 // Cannot invoke mid-epoch
290 if (isTraining ? IsTrainingActive() : IsValidationActive())
291 throw std::runtime_error("RDataLoaderEngine::Save: this dataset is already being iterated elsewhere "
292 "(e.g. inside a training loop). Finish or stop that iteration before saving.");
293
294 REpochGuard epoch(*this, isTraining);
295 auto sink = CreateBatchSink(dataset_name, filename, MakeColumnLayout(), outputFormat);
296
297 while (true) {
298 RFlat2DMatrix batch = isTraining ? GetTrainBatch() : GetValidationBatch();
299 if (batch.GetSize() == 0)
300 break;
301 sink->FillBatch(batch);
302 }
303
304 sink->Commit();
305 }
306
307 /// \brief Activate the training epoch by starting the batchloader.
309 {
310 {
311 std::lock_guard<std::mutex> lock(fLoadingMutex);
314 if (!fLoadEager) {
315 // Shuffle the cluster indices at the beginning of each epoch
316 fClusterLoader->ShuffleTrainingClusters(fTrainingEpochCount++);
317 }
318 }
319
320 fTrainingBatchLoader->Activate();
321 fLoadingCondition.notify_all();
322 }
323
325 {
326 {
327 std::lock_guard<std::mutex> lock(fLoadingMutex);
328 fTrainingEpochActive = false;
329 }
330
331 fTrainingBatchLoader->Reset();
332 fTrainingBatchLoader->DeActivate();
333 fLoadingCondition.notify_all();
334 }
335
337 {
338 {
339 std::lock_guard<std::mutex> lock(fLoadingMutex);
342 if (!fLoadEager) {
343 fClusterLoader->ShuffleValidationClusters(fValidationEpochCount++);
344 }
345 }
346
347 fValidationBatchLoader->Activate();
348 fLoadingCondition.notify_all();
349 }
350
352 {
353 {
354 std::lock_guard<std::mutex> lock(fLoadingMutex);
356 }
357
358 fValidationBatchLoader->Reset();
359 fValidationBatchLoader->DeActivate();
360 fLoadingCondition.notify_all();
361 }
362
363 /// \brief Main loop for loading clusters and creating batches.
364 /// The producer (loading thread) will keep loading clusters and creating batches until the end of the epoch is
365 /// reached, or the generator is deactivated.
366 void LoadData()
367 {
368 std::unique_lock<std::mutex> lock(fLoadingMutex);
369
370 while (true) {
371 // Wait until we have work or shutdown
372 fLoadingCondition.wait(lock, [&] {
373 return !fIsActive ||
374 (fTrainingEpochActive && fTrainingClusterIdx < fClusterLoader->GetNumTrainingClusters()) ||
375 (fValidationEpochActive && fValidationClusterIdx < fClusterLoader->GetNumValidationClusters());
376 });
377
378 if (!fIsActive) {
379 break;
380 }
381
382 // Helper: check if validation queue below watermark and needs the producer
383 auto validationEmpty = [&] {
384 if (!fValidationEpochActive || fValidationClusterIdx >= fClusterLoader->GetNumValidationClusters())
385 return false;
386 if (fValidationBatchLoader->isProducerDone())
387 return false;
388 return fValidationBatchLoader->GetNumBatchQueue() < fLowWatermark / fBatchSize;
389 };
390
391 // -- TRAINING --
393 const std::size_t numTrainingClusters = fClusterLoader->GetNumTrainingClusters();
394
395 while (true) {
396 // Stop conditions (shutdown or epoch end)
398 break;
399
400 // No more chunks to load: signal consumers
401 if (fTrainingClusterIdx >= numTrainingClusters) {
402 fTrainingBatchLoader->MarkProducerDone();
403 break;
404 }
405
406 // In the case of training prefetching, we could start requesting data for the next training loop while
407 // validation is active and might need data. To avoid getting stuck in the training loop, we check if the
408 // validation queue is below watermark and if so, we break out of the training loop.
409 if (validationEmpty()) {
410 break;
411 }
412
413 // If queue is not empty, wait until it drains below watermark, or validation needs data, or we are
414 // deactivated.
415 if (fTrainingBatchLoader->GetNumBatchQueue() >= fLowWatermark / fBatchSize) {
416 fLoadingCondition.wait(lock, [&] {
417 return !fIsActive || !fTrainingEpochActive ||
418 fTrainingBatchLoader->GetNumBatchQueue() < (fLowWatermark / fBatchSize) ||
419 validationEmpty();
420 });
421 continue;
422 }
423
424 // Accumulate clusters to load, enough to fill the buffer, or until we run out of clusters
425 std::vector<RClusterRange> trainClustersToLoad;
426 auto accumulatedEntries = 0;
427 const bool discovering = !fClusterLoader->IsSplitDiscovered();
428 while (fTrainingClusterIdx < numTrainingClusters && accumulatedEntries < fBufferCapacity &&
429 (!discovering || trainClustersToLoad.empty())) {
430 const auto &cluster = fClusterLoader->GetTrainingClusters()[fTrainingClusterIdx++];
431 trainClustersToLoad.push_back(cluster);
432 accumulatedEntries += cluster.GetNumEntries();
433 }
434
435 const bool isLastBuffer = (fTrainingClusterIdx >= numTrainingClusters);
436
437 // Release lock while reading and loading data to allow the consumer to access the queue freely in
438 // parallel. The loading thread re-acquires the lock in CreateBatches when it needs to push batches to
439 // the queue.
440 lock.unlock();
441 RFlat2DMatrix stagingBuffer(accumulatedEntries, fClusterLoader->GetNumChunkCols());
442 std::size_t rowOffset = 0;
443
444 for (auto &cluster : trainClustersToLoad) {
445 auto loadedEntries = fClusterLoader->LoadTrainingClusterInto(stagingBuffer, cluster.rdfIdx,
446 cluster.start, cluster.end, rowOffset);
447 if (discovering) {
448 // For the first epoch, we might discover that the cluster has fewer entries than expected because
449 // of filters
450 cluster.SetNumEntries(loadedEntries);
451 }
452 rowOffset += cluster.GetNumEntries();
453 }
454
455 if (discovering && fNumTrainingEntries == 0 && fClusterLoader->GetNumTrainingEntries() > 0) {
456 fNumTrainingEntries = fClusterLoader->GetNumTrainingEntries();
457 fNumValidationEntries = fClusterLoader->GetNumValidationEntries();
458 fTrainingBatchLoader->RecalculateBatchCounts(fNumTrainingEntries);
459 fValidationBatchLoader->RecalculateBatchCounts(fNumValidationEntries);
460 }
461
462 if (rowOffset < static_cast<std::size_t>(accumulatedEntries)) {
463 stagingBuffer.Resize(rowOffset, stagingBuffer.GetCols());
464 }
465
466 RFlat2DMatrix shuffledStagingBuffer;
467 fTrainingBatchLoader->CreateBatches(
468 fTensorOperators->ShuffleTensor(shuffledStagingBuffer, stagingBuffer), isLastBuffer);
469
470 // Re-acquire the lock before the next iteration to check conditions and update indices
471 lock.lock();
472
473 if (isLastBuffer && discovering) {
474 fClusterLoader->FinaliseSplitDiscovery();
475 }
476 }
477 }
478
479 // -- VALIDATION --
481 const std::size_t numValidationClusters = fClusterLoader->GetNumValidationClusters();
482
483 while (true) {
484 // Stop conditions (shutdown or epoch end)
486 break;
487
488 // No more chunks to load: signal consumers
489 if (fValidationClusterIdx >= numValidationClusters) {
490 fValidationBatchLoader->MarkProducerDone();
491 break;
492 }
493
494 // If queue is not hungry, wait until it drains below watermark, or we are deactivated
495 if (fValidationBatchLoader->GetNumBatchQueue() >= (fLowWatermark / fBatchSize)) {
496 fLoadingCondition.wait(lock, [&] {
497 return !fIsActive || !fValidationEpochActive ||
498 fValidationBatchLoader->GetNumBatchQueue() < (fLowWatermark / fBatchSize);
499 });
500 continue;
501 }
502
503 // Accumulate clusters to load, enough to fill the buffer, or until we run out of clusters
504 std::vector<RClusterRange> valClustersToLoad;
505 auto accumulatedEntries = 0;
506 while (fValidationClusterIdx < numValidationClusters && accumulatedEntries < fBufferCapacity) {
507 const auto &cluster = fClusterLoader->GetValidationClusters()[fValidationClusterIdx++];
508 valClustersToLoad.push_back(cluster);
509 accumulatedEntries += cluster.GetNumEntries();
510 }
511
512 const bool isLastBuffer = (fValidationClusterIdx >= numValidationClusters);
513
514 lock.unlock();
515
516 RFlat2DMatrix stagingBuffer(accumulatedEntries, fClusterLoader->GetNumChunkCols());
517 std::size_t rowOffset = 0;
518
519 for (const auto &cluster : valClustersToLoad) {
520 fClusterLoader->LoadValidationClusterInto(stagingBuffer, cluster.rdfIdx, cluster.start, cluster.end,
521 rowOffset);
522 rowOffset += cluster.GetNumEntries();
523 }
524
525 RFlat2DMatrix shuffledStagingBuffer;
526 fValidationBatchLoader->CreateBatches(
527 fTensorOperators->ShuffleTensor(shuffledStagingBuffer, stagingBuffer), isLastBuffer);
528
529 lock.lock();
530 }
531 }
532 }
533 }
534
535 /// \brief Create training batches by first loading a chunk (see RClusterLoader) and split it into batches (see
536 /// RBatchLoader)
538 {
539 fTrainingBatchLoader->Activate();
540
541 if (fLoadEager) {
543 if (fSampleType == "") {
545 }
546
547 else {
549 }
550
551 fTrainingBatchLoader->CreateBatches(*source, true);
552 fTrainingBatchLoader->MarkProducerDone();
553 }
554 }
555
556 /// \brief Creates validation batches by first loading a chunk (see RClusterLoader), and then split it into batches
557 /// (see RBatchLoader)
559 {
560 fValidationBatchLoader->Activate();
561
562 if (fLoadEager) {
564 if (fSampleType == "") {
566 }
567
568 else {
570 }
571
572 fValidationBatchLoader->CreateBatches(*source, true);
573 fValidationBatchLoader->MarkProducerDone();
574 }
575 }
576
577 /// \brief Loads a training batch from the queue
579 {
580 // Get next batch if available
581 return fTrainingBatchLoader->GetBatch();
582 }
583
584 /// \brief Loads a validation batch from the queue
586 {
587 // Get next batch if available
588 return fValidationBatchLoader->GetBatch();
589 }
590
591 std::size_t NumberOfTrainingBatches() { return fTrainingBatchLoader->GetNumBatches(); }
592 std::size_t NumberOfValidationBatches() { return fValidationBatchLoader->GetNumBatches(); }
593
594 std::size_t TrainRemainderRows() { return fTrainingBatchLoader->GetNumRemainderRows(); }
595 std::size_t ValidationRemainderRows() { return fValidationBatchLoader->GetNumRemainderRows(); }
596
597 bool IsActive()
598 {
599 std::lock_guard<std::mutex> lock(fLoadingMutex);
600 return fIsActive;
601 }
602
604 {
605 std::lock_guard<std::mutex> lock(fLoadingMutex);
607 }
608
610 {
611 std::lock_guard<std::mutex> lock(fLoadingMutex);
613 }
614};
615
616} // namespace ROOT::Experimental::Internal::ML
617
618#endif // ROOT_INTERNAL_ML_RDATALOADERENGINE
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t WindowAttributes_t Float_t Float_t Float_t Int_t Int_t UInt_t UInt_t Rectangle_t Int_t Int_t Window_t TString Int_t GCValues_t GetPrimarySelectionOwner GetDisplay GetScreen GetColormap GetNativeEvent const char const char dpyName wid window const char font_name cursor keysym reg const char only_if_exist regb h Point_t winding char text const char depth char const char Int_t count const char ColorStruct_t color const char filename
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t WindowAttributes_t Float_t Float_t Float_t Int_t Int_t UInt_t UInt_t Rectangle_t Int_t Int_t Window_t TString Int_t GCValues_t GetPrimarySelectionOwner GetDisplay GetScreen GetColormap GetNativeEvent const char const char dpyName wid window const char font_name cursor keysym reg const char only_if_exist regb h Point_t winding char text const char depth char const char Int_t count const char ColorStruct_t color const char Pixmap_t Pixmap_t PictureAttributes_t attr const char char ret_data h unsigned char height h offset
Option_t Option_t width
Loads TTree/RNTuple clusters from one or more RDataFrames into RFlat2DMatrix buffers for ML training ...
In this class, the processes of loading clusters (see RClusterLoader) and creating batches from those...
void ActivateTrainingEpoch()
Activate the training epoch by starting the batchloader.
void Save(std::string_view dataset_name, std::string_view filename, bool isTraining, std::string_view outputFormat)
Materialize one train/test split to disk by draining a full epoch through the normal batch pipeline a...
RFlat2DMatrix GetTrainBatch()
Loads a training batch from the queue.
std::unique_ptr< RFlat2DMatrixOperators > fTensorOperators
void CreateValidationBatches()
Creates validation batches by first loading a chunk (see RClusterLoader), and then split it into batc...
void LoadData()
Main loop for loading clusters and creating batches.
void CreateTrainBatches()
Create training batches by first loading a chunk (see RClusterLoader) and split it into batches (see ...
RFlat2DMatrix GetValidationBatch()
Loads a validation batch from the queue.
void Activate()
Activate the loading process by spawning the loading thread.
std::unique_ptr< RDatasetLoader< Args... > > fDatasetLoader
std::vector< RColumnLayout > MakeColumnLayout() const
Describe how the loader's columns map onto a batch-tensor row.
RDataLoaderEngine(const std::vector< ROOT::RDF::RNode > &rdfs, const std::size_t batchSize, const std::size_t batchesInMemory, const std::vector< std::string > &cols, const std::vector< std::size_t > &vecSizes={}, const float vecPadding=0.0, const float testSize=0.0, bool shuffle=true, bool dropRemainder=true, const std::size_t setSeed=0, bool loadEager=false, std::string sampleType="", float sampleRatio=1.0, bool replacement=false)
std::unique_ptr< RClusterLoader< Args... > > fClusterLoader
std::unique_ptr< RBatchSink > CreateBatchSink(std::string_view dataset_name, std::string_view filename, std::vector< RColumnLayout > layout, std::string_view format)
Create the sink matching format.
Opens a training or validation epoch and closes it again when done.
Wrapper around ROOT::RVec<float> representing a 2D matrix.
void Resize(std::size_t rows, std::size_t cols)