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
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),
172 fBatchSize(batchSize),
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,
189
190 if (fSampleType == "") {
191 fDatasetLoader->ConcatenateDatasets();
192
193 fTrainingDataset = fDatasetLoader->GetTrainingDataset();
194 fValidationDataset = fDatasetLoader->GetValidationDataset();
195
196 fNumTrainingEntries = fDatasetLoader->GetNumTrainingEntries();
197 fNumValidationEntries = fDatasetLoader->GetNumValidationEntries();
198 }
199
200 else {
201 fTrainingDatasets = fDatasetLoader->GetTrainingDatasets();
202 fValidationDatasets = fDatasetLoader->GetValidationDatasets();
203
206 fValidationSampler = std::make_unique<RSampler>(fValidationDatasets, fSampleType, fSampleRatio,
208
209 fNumTrainingEntries = fTrainingSampler->GetNumEntries();
210 fNumValidationEntries = fValidationSampler->GetNumEntries();
211 }
212 }
213
214 else {
215 // scan cluster boundaries
216 fClusterLoader = std::make_unique<RClusterLoader<Args...>>(fRdfs, fCols, fVecSizes, vecPadding, fTestSize,
218
219 // derive buffer quantities
221 // at least one batch, otherwise the refill threshold rounds down to 0 and nothing is ever loaded
224
225 // split cluster list into training and validation
226 fClusterLoader->SplitDataset();
227 fNumTrainingEntries = fClusterLoader->GetNumTrainingEntries();
228 fNumValidationEntries = fClusterLoader->GetNumValidationEntries();
229 }
230
231 fTrainingBatchLoader = std::make_unique<RBatchLoader>(fBatchSize, fCols, fLoadingMutex, fLoadingCondition,
233 fValidationBatchLoader = std::make_unique<RBatchLoader>(fBatchSize, fCols, fLoadingMutex, fLoadingCondition,
235 }
236
238
240 {
241 {
242 std::lock_guard<std::mutex> lock(fLoadingMutex);
243 if (!fIsActive)
244 return;
245 fIsActive = false;
246 }
247
248 fLoadingCondition.notify_all();
249
250 if (fLoadingThread) {
251 if (fLoadingThread->joinable()) {
252 fLoadingThread->join();
253 }
254 }
255
256 fLoadingThread.reset();
257 }
258
259 /// \brief Activate the loading process by spawning the loading thread.
260 void Activate()
261 {
262 {
263 std::lock_guard<std::mutex> lock(fLoadingMutex);
264 if (fIsActive)
265 return;
266
267 fIsActive = true;
268 }
269
270 if (fLoadEager) {
271 return;
272 }
273
274 fLoadingThread = std::make_unique<std::thread>(&RDataLoaderEngine::LoadData, this);
275 }
276
277 /// \brief Materialize one train/test split to disk by draining a full epoch through the normal batch
278 /// pipeline and Fill() each batch into \p filename instead of yielding it.
279 ///
280 /// Filters, shuffling, the train/validation split and the batch_size/drop_remainder settings
281 /// are all inherited from the loader's configuration.
282 /// \param outputFormat Either "ttree" or "rntuple".
283 void Save(std::string_view dataset_name, std::string_view filename, bool isTraining, std::string_view outputFormat)
284 {
285 // Cannot invoke mid-epoch
287 throw std::runtime_error("RDataLoaderEngine::Save: this dataset is already being iterated elsewhere "
288 "(e.g. inside a training loop). Finish or stop that iteration before saving.");
289
292
293 while (true) {
295 if (batch.GetSize() == 0)
296 break;
297 sink->FillBatch(batch);
298 }
299
300 sink->Commit();
301 }
302
303 /// \brief Activate the training epoch by starting the batchloader.
305 {
306 {
307 std::lock_guard<std::mutex> lock(fLoadingMutex);
310 if (!fLoadEager) {
311 // Shuffle the cluster indices at the beginning of each epoch
312 fClusterLoader->ShuffleTrainingClusters(fTrainingEpochCount++);
313 }
314 }
315
316 fTrainingBatchLoader->Activate();
317 fLoadingCondition.notify_all();
318 }
319
321 {
322 {
323 std::lock_guard<std::mutex> lock(fLoadingMutex);
324 fTrainingEpochActive = false;
325 }
326
327 fTrainingBatchLoader->Reset();
328 fTrainingBatchLoader->DeActivate();
329 fLoadingCondition.notify_all();
330 }
331
333 {
334 {
335 std::lock_guard<std::mutex> lock(fLoadingMutex);
338 if (!fLoadEager) {
339 fClusterLoader->ShuffleValidationClusters(fValidationEpochCount++);
340 }
341 }
342
343 fValidationBatchLoader->Activate();
344 fLoadingCondition.notify_all();
345 }
346
348 {
349 {
350 std::lock_guard<std::mutex> lock(fLoadingMutex);
352 }
353
354 fValidationBatchLoader->Reset();
355 fValidationBatchLoader->DeActivate();
356 fLoadingCondition.notify_all();
357 }
358
359 /// \brief Main loop for loading clusters and creating batches.
360 /// The producer (loading thread) will keep loading clusters and creating batches until the end of the epoch is
361 /// reached, or the generator is deactivated.
362 void LoadData()
363 {
364 std::unique_lock<std::mutex> lock(fLoadingMutex);
365
366 while (true) {
367 // Wait until we have work or shutdown
368 fLoadingCondition.wait(lock, [&] {
369 return !fIsActive ||
370 (fTrainingEpochActive && fTrainingClusterIdx < fClusterLoader->GetNumTrainingClusters()) ||
371 (fValidationEpochActive && fValidationClusterIdx < fClusterLoader->GetNumValidationClusters());
372 });
373
374 if (!fIsActive) {
375 break;
376 }
377
378 // Helper: check if validation queue below watermark and needs the producer
379 auto validationEmpty = [&] {
380 if (!fValidationEpochActive || fValidationClusterIdx >= fClusterLoader->GetNumValidationClusters())
381 return false;
382 if (fValidationBatchLoader->isProducerDone())
383 return false;
384 return fValidationBatchLoader->GetNumBatchQueue() < fLowWatermark / fBatchSize;
385 };
386
387 // -- TRAINING --
389 const std::size_t numTrainingClusters = fClusterLoader->GetNumTrainingClusters();
390
391 while (true) {
392 // Stop conditions (shutdown or epoch end)
394 break;
395
396 // No more chunks to load: signal consumers
398 fTrainingBatchLoader->MarkProducerDone();
399 break;
400 }
401
402 // In the case of training prefetching, we could start requesting data for the next training loop while
403 // validation is active and might need data. To avoid getting stuck in the training loop, we check if the
404 // validation queue is below watermark and if so, we break out of the training loop.
405 if (validationEmpty()) {
406 break;
407 }
408
409 // If queue is not empty, wait until it drains below watermark, or validation needs data, or we are
410 // deactivated.
411 if (fTrainingBatchLoader->GetNumBatchQueue() >= fLowWatermark / fBatchSize) {
412 fLoadingCondition.wait(lock, [&] {
413 return !fIsActive || !fTrainingEpochActive ||
414 fTrainingBatchLoader->GetNumBatchQueue() < (fLowWatermark / fBatchSize) ||
416 });
417 continue;
418 }
419
420 // Accumulate clusters to load, enough to fill the buffer, or until we run out of clusters
421 std::vector<RClusterRange> trainClustersToLoad;
422 auto accumulatedEntries = 0;
423 const bool discovering = !fClusterLoader->IsSplitDiscovered();
425 (!discovering || trainClustersToLoad.empty())) {
426 const auto &cluster = fClusterLoader->GetTrainingClusters()[fTrainingClusterIdx++];
427 trainClustersToLoad.push_back(cluster);
428 accumulatedEntries += cluster.GetNumEntries();
429 }
430
432
433 // Release lock while reading and loading data to allow the consumer to access the queue freely in
434 // parallel. The loading thread re-acquires the lock in CreateBatches when it needs to push batches to
435 // the queue.
436 lock.unlock();
438 std::size_t rowOffset = 0;
439
440 for (auto &cluster : trainClustersToLoad) {
441 auto loadedEntries = fClusterLoader->LoadTrainingClusterInto(stagingBuffer, cluster.rdfIdx,
442 cluster.start, cluster.end, rowOffset);
443 if (discovering) {
444 // For the first epoch, we might discover that the cluster has fewer entries than expected because
445 // of filters
446 cluster.SetNumEntries(loadedEntries);
447 }
448 rowOffset += cluster.GetNumEntries();
449 }
450
451 if (discovering && fNumTrainingEntries == 0 && fClusterLoader->GetNumTrainingEntries() > 0) {
452 fNumTrainingEntries = fClusterLoader->GetNumTrainingEntries();
453 fNumValidationEntries = fClusterLoader->GetNumValidationEntries();
454 fTrainingBatchLoader->RecalculateBatchCounts(fNumTrainingEntries);
455 fValidationBatchLoader->RecalculateBatchCounts(fNumValidationEntries);
456 }
457
458 if (rowOffset < static_cast<std::size_t>(accumulatedEntries)) {
459 stagingBuffer.Resize(rowOffset, stagingBuffer.GetCols());
460 }
461
465
466 // Re-acquire the lock before the next iteration to check conditions and update indices
467 lock.lock();
468
469 if (isLastBuffer && discovering) {
470 fClusterLoader->FinaliseSplitDiscovery();
471 }
472 }
473 }
474
475 // -- VALIDATION --
477 const std::size_t numValidationClusters = fClusterLoader->GetNumValidationClusters();
478
479 while (true) {
480 // Stop conditions (shutdown or epoch end)
482 break;
483
484 // No more chunks to load: signal consumers
486 fValidationBatchLoader->MarkProducerDone();
487 break;
488 }
489
490 // If queue is not hungry, wait until it drains below watermark, or we are deactivated
491 if (fValidationBatchLoader->GetNumBatchQueue() >= (fLowWatermark / fBatchSize)) {
492 fLoadingCondition.wait(lock, [&] {
493 return !fIsActive || !fValidationEpochActive ||
494 fValidationBatchLoader->GetNumBatchQueue() < (fLowWatermark / fBatchSize);
495 });
496 continue;
497 }
498
499 // Accumulate clusters to load, enough to fill the buffer, or until we run out of clusters
500 std::vector<RClusterRange> valClustersToLoad;
501 auto accumulatedEntries = 0;
503 const auto &cluster = fClusterLoader->GetValidationClusters()[fValidationClusterIdx++];
504 valClustersToLoad.push_back(cluster);
505 accumulatedEntries += cluster.GetNumEntries();
506 }
507
509
510 lock.unlock();
511
513 std::size_t rowOffset = 0;
514
515 for (const auto &cluster : valClustersToLoad) {
516 fClusterLoader->LoadValidationClusterInto(stagingBuffer, cluster.rdfIdx, cluster.start, cluster.end,
517 rowOffset);
518 rowOffset += cluster.GetNumEntries();
519 }
520
524
525 lock.lock();
526 }
527 }
528 }
529 }
530
531 /// \brief Create training batches by first loading a chunk (see RClusterLoader) and split it into batches (see
532 /// RBatchLoader)
534 {
535 fTrainingBatchLoader->Activate();
536
537 if (fLoadEager) {
538 if (fSampleType == "") {
540 }
541
542 else {
544 }
545
546 fTrainingBatchLoader->CreateBatches(fSampledTrainingDataset, true);
547 fTrainingBatchLoader->MarkProducerDone();
548 }
549 }
550
551 /// \brief Creates validation batches by first loading a chunk (see RClusterLoader), and then split it into batches
552 /// (see RBatchLoader)
554 {
555 fValidationBatchLoader->Activate();
556
557 if (fLoadEager) {
558 if (fSampleType == "") {
560 }
561
562 else {
564 }
565
567 fValidationBatchLoader->MarkProducerDone();
568 }
569 }
570
571 /// \brief Loads a training batch from the queue
573 {
574 // Get next batch if available
575 return fTrainingBatchLoader->GetBatch();
576 }
577
578 /// \brief Loads a validation batch from the queue
580 {
581 // Get next batch if available
582 return fValidationBatchLoader->GetBatch();
583 }
584
585 std::size_t NumberOfTrainingBatches() { return fTrainingBatchLoader->GetNumBatches(); }
586 std::size_t NumberOfValidationBatches() { return fValidationBatchLoader->GetNumBatches(); }
587
588 std::size_t TrainRemainderRows() { return fTrainingBatchLoader->GetNumRemainderRows(); }
589 std::size_t ValidationRemainderRows() { return fValidationBatchLoader->GetNumRemainderRows(); }
590
591 bool IsActive()
592 {
593 std::lock_guard<std::mutex> lock(fLoadingMutex);
594 return fIsActive;
595 }
596
598 {
599 std::lock_guard<std::mutex> lock(fLoadingMutex);
601 }
602
604 {
605 std::lock_guard<std::mutex> lock(fLoadingMutex);
607 }
608};
609
610} // namespace ROOT::Experimental::Internal::ML
611
612#endif // ROOT_INTERNAL_ML_RDATALOADERENGINE
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
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
void SplitDatasets()
Split the dataframes in a training and validation dataset.
const_iterator end() const
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.