File indexing completed on 2026-08-16 09:21:06
0001
0002
0003
0004
0005
0006
0007
0008
0009
0010
0011 #ifndef ROOT_INTERNAL_ML_RDATASETLOADER
0012 #define ROOT_INTERNAL_ML_RDATASETLOADER
0013
0014 #include <algorithm>
0015 #include <memory>
0016 #include <numeric>
0017 #include <string>
0018 #include <type_traits>
0019 #include <vector>
0020
0021 #include "ROOT/ML/RFlat2DMatrix.hxx"
0022 #include "ROOT/ML/RFlat2DMatrixOperators.hxx"
0023 #include "ROOT/RDataFrame.hxx"
0024 #include "ROOT/RDF/Utils.hxx"
0025
0026 namespace ROOT::Experimental::Internal::ML {
0027
0028
0029
0030
0031
0032
0033
0034 template <typename... ColTypes>
0035 class RDatasetLoaderFunctor {
0036 std::size_t fOffset{};
0037 std::size_t fVecSizeIdx{};
0038 float fVecPadding{};
0039 std::vector<std::size_t> fMaxVecSizes{};
0040 RFlat2DMatrix &fDatasetTensor;
0041
0042 std::size_t fNumDatasetCols;
0043
0044 int fI;
0045 int fNumColumns;
0046
0047
0048
0049 template <typename T, std::enable_if_t<ROOT::Internal::RDF::IsDataContainer<T>::value, int> = 0>
0050 void AssignToTensor(const T &vec, int i, int numColumns)
0051 {
0052 std::size_t max_vec_size = fMaxVecSizes[fVecSizeIdx++];
0053 std::size_t vec_size = vec.size();
0054 if (vec_size < max_vec_size)
0055 {
0056 std::copy(vec.begin(), vec.end(), &fDatasetTensor.GetData()[fOffset + numColumns * i]);
0057 std::fill(&fDatasetTensor.GetData()[fOffset + numColumns * i + vec_size],
0058 &fDatasetTensor.GetData()[fOffset + numColumns * i + max_vec_size], fVecPadding);
0059 } else
0060 {
0061 std::copy(vec.begin(), vec.begin() + max_vec_size, &fDatasetTensor.GetData()[fOffset + numColumns * i]);
0062 }
0063 fOffset += max_vec_size;
0064 }
0065
0066
0067
0068 template <typename T, std::enable_if_t<!ROOT::Internal::RDF::IsDataContainer<T>::value, int> = 0>
0069 void AssignToTensor(const T &val, int i, int numColumns)
0070 {
0071 fDatasetTensor.GetData()[fOffset + numColumns * i] = val;
0072 fOffset++;
0073 }
0074
0075 public:
0076 RDatasetLoaderFunctor(RFlat2DMatrix &datasetTensor, std::size_t numColumns,
0077 const std::vector<std::size_t> &maxVecSizes, float vecPadding, int i)
0078 : fDatasetTensor(datasetTensor),
0079 fMaxVecSizes(maxVecSizes),
0080 fVecPadding(vecPadding),
0081 fI(i),
0082 fNumColumns(numColumns)
0083 {
0084 }
0085
0086 void operator()(const ColTypes &...cols)
0087 {
0088 fVecSizeIdx = 0;
0089 (AssignToTensor(cols, fI, fNumColumns), ...);
0090 }
0091 };
0092
0093
0094
0095
0096
0097
0098
0099
0100
0101
0102 template <typename... Args>
0103 class RDatasetLoader {
0104 private:
0105 std::size_t fNumEntries;
0106 float fValidationSplit;
0107
0108 std::vector<std::size_t> fVecSizes;
0109 std::size_t fSumVecSizes;
0110 std::size_t fVecPadding;
0111 std::size_t fNumDatasetCols;
0112
0113 std::vector<RFlat2DMatrix> fTrainingDatasets;
0114 std::vector<RFlat2DMatrix> fValidationDatasets;
0115
0116 RFlat2DMatrix fTrainingDataset;
0117 RFlat2DMatrix fValidationDataset;
0118
0119 std::size_t fNumTrainingEntries;
0120 std::size_t fNumValidationEntries;
0121 std::unique_ptr<RFlat2DMatrixOperators> fTensorOperators;
0122
0123 std::vector<ROOT::RDF::RNode> f_rdfs;
0124 std::vector<std::string> fCols;
0125 std::size_t fNumCols;
0126 std::size_t fSetSeed;
0127
0128 bool fNotFiltered;
0129 bool fShuffle;
0130
0131 ROOT::RDF::RResultPtr<std::vector<ULong64_t>> fEntries;
0132
0133 public:
0134 RDatasetLoader(const std::vector<ROOT::RDF::RNode> &rdfs, const float validationSplit,
0135 const std::vector<std::string> &cols, const std::vector<std::size_t> &vecSizes = {},
0136 const float vecPadding = 0.0, bool shuffle = true, const std::size_t setSeed = 0)
0137 : f_rdfs(rdfs),
0138 fCols(cols),
0139 fVecSizes(vecSizes),
0140 fVecPadding(vecPadding),
0141 fValidationSplit(validationSplit),
0142 fShuffle(shuffle),
0143 fSetSeed(setSeed)
0144 {
0145 fTensorOperators = std::make_unique<RFlat2DMatrixOperators>(fShuffle, fSetSeed);
0146 fNumCols = fCols.size();
0147 fSumVecSizes = std::accumulate(fVecSizes.begin(), fVecSizes.end(), 0);
0148
0149 fNumDatasetCols = fNumCols + fSumVecSizes - fVecSizes.size();
0150 }
0151
0152
0153
0154
0155
0156
0157 void SplitDataframe(ROOT::RDF::RNode &rdf, RFlat2DMatrix &TrainingDataset, RFlat2DMatrix &ValidationDataset)
0158 {
0159 ROOT::RDF::RResultPtr<std::vector<ULong64_t>> Entries = rdf.Take<ULong64_t>("rdfentry_");
0160 const std::size_t NumEntries = Entries->size();
0161
0162
0163 Entries->push_back((*Entries)[NumEntries - 1] + 1);
0164
0165
0166 std::size_t NumValidationEntries = static_cast<std::size_t>(fValidationSplit * NumEntries);
0167 std::size_t NumTrainingEntries = NumEntries - NumValidationEntries;
0168
0169 RFlat2DMatrix Dataset({NumEntries, fNumDatasetCols});
0170
0171 bool NotFiltered = rdf.GetFilterNames().empty();
0172 if (NotFiltered) {
0173 RDatasetLoaderFunctor<Args...> func(Dataset, fNumDatasetCols, fVecSizes, fVecPadding, 0);
0174 rdf.Foreach(func, fCols);
0175 }
0176
0177 else {
0178 std::size_t datasetEntry = 0;
0179 for (std::size_t j = 0; j < NumEntries; j++) {
0180 RDatasetLoaderFunctor<Args...> func(Dataset, fNumDatasetCols, fVecSizes, fVecPadding, datasetEntry);
0181 ROOT::Internal::RDF::ChangeBeginAndEndEntries(rdf, (*Entries)[j], (*Entries)[j + 1]);
0182 rdf.Foreach(func, fCols);
0183 datasetEntry++;
0184 }
0185 }
0186
0187
0188 ROOT::Internal::RDF::ChangeBeginAndEndEntries(rdf, (*Entries)[0], (*Entries)[NumEntries]);
0189
0190 RFlat2DMatrix ShuffledDataset({NumEntries, fNumDatasetCols});
0191 fTensorOperators->ShuffleTensor(ShuffledDataset, Dataset);
0192 fTensorOperators->SliceTensor(TrainingDataset, ShuffledDataset, {{0, NumTrainingEntries}, {0, fNumDatasetCols}});
0193 fTensorOperators->SliceTensor(ValidationDataset, ShuffledDataset,
0194 {{NumTrainingEntries, NumEntries}, {0, fNumDatasetCols}});
0195 }
0196
0197
0198
0199 void SplitDatasets()
0200 {
0201 fNumEntries = 0;
0202 fNumTrainingEntries = 0;
0203 fNumValidationEntries = 0;
0204
0205 for (auto &rdf : f_rdfs) {
0206 RFlat2DMatrix TrainingDataset;
0207 RFlat2DMatrix ValidationDataset;
0208
0209 SplitDataframe(rdf, TrainingDataset, ValidationDataset);
0210 fTrainingDatasets.push_back(TrainingDataset);
0211 fValidationDatasets.push_back(ValidationDataset);
0212
0213 fNumTrainingEntries += TrainingDataset.GetRows();
0214 fNumValidationEntries += ValidationDataset.GetRows();
0215 fNumEntries += TrainingDataset.GetRows() + ValidationDataset.GetRows();
0216 }
0217 }
0218
0219
0220
0221 void ConcatenateDatasets()
0222 {
0223 fTensorOperators->ConcatenateTensors(fTrainingDataset, fTrainingDatasets);
0224 fTensorOperators->ConcatenateTensors(fValidationDataset, fValidationDatasets);
0225 }
0226
0227 std::vector<RFlat2DMatrix> GetTrainingDatasets() { return fTrainingDatasets; }
0228 std::vector<RFlat2DMatrix> GetValidationDatasets() { return fValidationDatasets; }
0229
0230 RFlat2DMatrix GetTrainingDataset() { return fTrainingDataset; }
0231 RFlat2DMatrix GetValidationDataset() { return fValidationDataset; }
0232
0233 std::size_t GetNumTrainingEntries() { return fTrainingDataset.GetRows(); }
0234 std::size_t GetNumValidationEntries() { return fValidationDataset.GetRows(); }
0235 };
0236
0237 }
0238 #endif