File indexing completed on 2026-08-16 09:21:05
0001
0002
0003
0004
0005
0006
0007
0008
0009
0010
0011
0012
0013
0014
0015 #ifndef ROOT_INTERNAL_ML_RCLUSTERLOADER
0016 #define ROOT_INTERNAL_ML_RCLUSTERLOADER
0017
0018 #include <algorithm>
0019 #include <numeric>
0020 #include <random>
0021 #include <string>
0022 #include <utility>
0023 #include <vector>
0024
0025 #include "ROOT/ML/RFlat2DMatrix.hxx"
0026 #include "ROOT/ML/RFlat2DMatrixOperators.hxx"
0027 #include "ROOT/RDataFrame.hxx"
0028 #include "ROOT/RDFHelpers.hxx"
0029 #include "ROOT/RDF/Utils.hxx"
0030
0031 namespace ROOT::Experimental::Internal::ML {
0032
0033
0034
0035
0036
0037
0038
0039
0040
0041
0042 struct RClusterRange {
0043 std::size_t rdfIdx;
0044 std::uint64_t start;
0045 std::uint64_t end;
0046 std::size_t numEntries{
0047 static_cast<std::size_t>(end - start)};
0048
0049 std::size_t GetNumEntries() const { return numEntries; }
0050 void SetNumEntries(std::size_t num) { numEntries = num; }
0051 };
0052
0053
0054
0055
0056
0057
0058
0059 template <typename... ColTypes>
0060 class RClusterLoaderFunctor {
0061 std::size_t fOffset{};
0062 std::size_t fVecSizeIdx{};
0063 float fVecPadding{};
0064 std::vector<std::size_t> fMaxVecSizes{};
0065 RFlat2DMatrix &fChunkTensor;
0066
0067 std::size_t fNumChunkCols;
0068
0069 int fI;
0070 int fNumColumns;
0071
0072
0073
0074 template <typename T, std::enable_if_t<ROOT::Internal::RDF::IsDataContainer<T>::value, int> = 0>
0075 void AssignToTensor(const T &vec, int i, int numColumns)
0076 {
0077 std::size_t max_vec_size = fMaxVecSizes[fVecSizeIdx++];
0078 std::size_t vec_size = vec.size();
0079
0080 float *dst = fChunkTensor.GetData() + fOffset + numColumns * i;
0081 if (vec_size < max_vec_size)
0082 {
0083 std::copy(vec.begin(), vec.end(), dst);
0084 std::fill(dst + vec_size, dst + max_vec_size, fVecPadding);
0085 } else
0086 {
0087 std::copy(vec.begin(), vec.begin() + max_vec_size, dst);
0088 }
0089 fOffset += max_vec_size;
0090 }
0091
0092
0093
0094 template <typename T, std::enable_if_t<!ROOT::Internal::RDF::IsDataContainer<T>::value, int> = 0>
0095 void AssignToTensor(const T &val, int i, int numColumns)
0096 {
0097 fChunkTensor.GetData()[fOffset + numColumns * i] = val;
0098 fOffset++;
0099 }
0100
0101 public:
0102 RClusterLoaderFunctor(RFlat2DMatrix &chunkTensor, std::size_t numColumns,
0103 const std::vector<std::size_t> &maxVecSizes, float vecPadding, int i,
0104 std::size_t rowOffset = 0)
0105 : fChunkTensor(chunkTensor),
0106 fMaxVecSizes(maxVecSizes),
0107 fVecPadding(vecPadding),
0108 fI(i),
0109 fNumColumns(numColumns),
0110 fOffset(rowOffset * numColumns)
0111 {
0112 }
0113
0114 void operator()(const ColTypes &...cols)
0115 {
0116 fVecSizeIdx = 0;
0117 (AssignToTensor(cols, fI, fNumColumns), ...);
0118 }
0119 };
0120
0121
0122
0123
0124
0125
0126
0127
0128
0129
0130
0131
0132
0133
0134
0135
0136
0137
0138
0139
0140
0141
0142
0143
0144
0145
0146
0147
0148 template <typename... Args>
0149 class RClusterLoader {
0150 private:
0151 std::vector<ROOT::RDF::RNode> &fRdfs;
0152 std::vector<std::size_t> fRdfSizes;
0153 std::vector<std::string> fCols;
0154 std::vector<std::size_t> fVecSizes;
0155 float fVecPadding;
0156 float fValidationSplit;
0157 bool fShuffle;
0158 std::size_t fSetSeed;
0159
0160 std::size_t fNumCols;
0161 std::size_t fSumVecSizes;
0162 std::size_t fNumChunkCols;
0163
0164 std::vector<RClusterRange> fAllClusters;
0165 std::vector<RClusterRange> fTrainingClusters;
0166 std::vector<RClusterRange> fValidationClusters;
0167
0168 std::size_t fTotalEntries{0};
0169 std::size_t fNumTrainingEntries{0};
0170 std::size_t fNumValidationEntries{0};
0171
0172 bool fIsFiltered{false};
0173 bool fSplitDiscovered{false};
0174 std::size_t fAccumulatedFilteredForTrain{0};
0175
0176 public:
0177 RClusterLoader(std::vector<ROOT::RDF::RNode> &rdfs, const std::vector<std::string> &cols,
0178 const std::vector<std::size_t> &vecSizes, float vecPadding, float validationSplit, bool shuffle,
0179 std::size_t setSeed)
0180 : fRdfs(rdfs),
0181 fCols(cols),
0182 fVecSizes(vecSizes),
0183 fVecPadding(vecPadding),
0184 fValidationSplit(validationSplit),
0185 fShuffle(shuffle),
0186 fSetSeed(setSeed)
0187 {
0188 fNumCols = fCols.size();
0189 fSumVecSizes = std::accumulate(fVecSizes.begin(), fVecSizes.end(), 0UL);
0190 fNumChunkCols = fNumCols + fSumVecSizes - fVecSizes.size();
0191
0192 for (auto &rdf : fRdfs) {
0193
0194 if (!rdf.GetFilterNames().empty()) {
0195 fIsFiltered = true;
0196 break;
0197 }
0198 }
0199
0200 fRdfSizes.resize(fRdfs.size(), 0);
0201
0202
0203
0204 for (std::size_t rdfIdx = 0; rdfIdx < fRdfs.size(); ++rdfIdx) {
0205 for (const auto &r : ROOT::Internal::RDF::GetDatasetGlobalClusterBoundaries(fRdfs[rdfIdx])) {
0206 fAllClusters.push_back({rdfIdx, r.first, r.second});
0207 auto numEntries = r.second - r.first;
0208 fRdfSizes[rdfIdx] += numEntries;
0209 fTotalEntries += numEntries;
0210 }
0211 }
0212 }
0213
0214
0215
0216
0217 void SplitDataset()
0218 {
0219 if (fAllClusters.empty())
0220 throw std::runtime_error("RClusterLoader::SplitDataset: no clusters found.");
0221
0222 if (fIsFiltered) {
0223 return;
0224 }
0225
0226 if (fShuffle) {
0227
0228
0229
0230
0231
0232 std::mt19937 g(fSetSeed);
0233 std::uniform_int_distribution<int> coin(0, 1);
0234
0235 for (const RClusterRange &c : fAllClusters) {
0236 const std::size_t sz = c.GetNumEntries();
0237 const std::size_t trainSz = static_cast<std::size_t>((1.0f - fValidationSplit) * sz);
0238 const std::size_t valSz = sz - trainSz;
0239
0240
0241 bool trainIsPrefix = coin(g);
0242 const uint64_t trainStart = trainIsPrefix ? c.start : c.start + static_cast<std::uint64_t>(valSz);
0243 const uint64_t valStart = trainIsPrefix ? c.start + static_cast<std::uint64_t>(trainSz) : c.start;
0244
0245 if (trainSz > 0) {
0246 fTrainingClusters.push_back({c.rdfIdx, trainStart, trainStart + static_cast<std::uint64_t>(trainSz)});
0247 fNumTrainingEntries += trainSz;
0248 }
0249 if (valSz > 0) {
0250 fValidationClusters.push_back({c.rdfIdx, valStart, valStart + static_cast<std::uint64_t>(valSz)});
0251 fNumValidationEntries += valSz;
0252 }
0253 }
0254 } else {
0255
0256
0257
0258
0259 const std::size_t targetTraining = fTotalEntries - static_cast<std::size_t>(fValidationSplit * fTotalEntries);
0260
0261 std::size_t accumulated = 0;
0262 std::size_t splitIdx = 0;
0263 for (; splitIdx < fAllClusters.size(); ++splitIdx) {
0264 const std::size_t sz = fAllClusters[splitIdx].GetNumEntries();
0265 if (accumulated + sz > targetTraining) {
0266 break;
0267 }
0268 accumulated += sz;
0269 }
0270
0271
0272 fTrainingClusters.assign(fAllClusters.begin(), fAllClusters.begin() + splitIdx);
0273 fNumTrainingEntries = accumulated;
0274
0275 if (splitIdx < fAllClusters.size() && accumulated < targetTraining) {
0276
0277 const RClusterRange &boundary = fAllClusters[splitIdx];
0278 const std::uint64_t splitPoint = boundary.start + static_cast<std::uint64_t>(targetTraining - accumulated);
0279
0280 fTrainingClusters.push_back({boundary.rdfIdx, boundary.start, splitPoint});
0281 fValidationClusters.push_back({boundary.rdfIdx, splitPoint, boundary.end});
0282 fValidationClusters.insert(fValidationClusters.end(), fAllClusters.begin() + splitIdx + 1,
0283 fAllClusters.end());
0284
0285 fNumTrainingEntries += splitPoint - boundary.start;
0286 } else {
0287 fValidationClusters.assign(fAllClusters.begin() + splitIdx, fAllClusters.end());
0288 }
0289
0290 fNumValidationEntries = fTotalEntries - fNumTrainingEntries;
0291 }
0292
0293 if (fTrainingClusters.empty())
0294 throw std::runtime_error("RClusterLoader::SplitDataset: no entries for training after split. "
0295 "Reduce validation_split.");
0296
0297 if (fValidationSplit > 0.0f && fValidationClusters.empty())
0298 throw std::runtime_error("RClusterLoader::SplitDataset: no entries for validation after split. "
0299 "Increase validation_split.");
0300 }
0301
0302
0303
0304 void ShuffleTrainingClusters(std::size_t epochIdx)
0305 {
0306 if (!fShuffle) {
0307 return;
0308 }
0309
0310 std::mt19937 g(fSetSeed == 0 ? std::random_device{}() : fSetSeed ^ epochIdx);
0311 std::shuffle(fTrainingClusters.begin(), fTrainingClusters.end(), g);
0312 }
0313
0314
0315
0316 void ShuffleValidationClusters(std::size_t epochIdx)
0317 {
0318 if (!fShuffle) {
0319 return;
0320 }
0321 std::mt19937 g(fSetSeed == 0 ? std::random_device{}() : fSetSeed ^ epochIdx);
0322 std::shuffle(fValidationClusters.begin(), fValidationClusters.end(), g);
0323 }
0324
0325 void LoadClusterInto(RFlat2DMatrix &dest, std::size_t rdfIdx, std::uint64_t startRow, std::uint64_t endRow,
0326 std::size_t rowOffset = 0)
0327 {
0328 ROOT::RDF::RNode &rdf = fRdfs[rdfIdx];
0329 ROOT::Internal::RDF::ChangeBeginAndEndEntries(rdf, startRow, endRow);
0330 RClusterLoaderFunctor<Args...> func(dest, fNumChunkCols, fVecSizes, fVecPadding, 0, rowOffset);
0331 rdf.Foreach(func, fCols);
0332 ROOT::Internal::RDF::ChangeBeginAndEndEntries(rdf, 0, fRdfSizes[rdfIdx]);
0333 }
0334
0335
0336
0337
0338
0339
0340
0341
0342
0343
0344
0345
0346
0347
0348
0349 std::size_t LoadTrainingClusterInto(RFlat2DMatrix &dest, std::size_t rdfIdx, std::uint64_t startRow,
0350 std::uint64_t endRow, std::size_t rowOffset = 0)
0351 {
0352 if (fIsFiltered && !fSplitDiscovered) {
0353
0354 if (fAccumulatedFilteredForTrain == 0 && fNumTrainingEntries == 0) {
0355 std::vector<ROOT::RDF::RResultPtr<ULong64_t>> counts;
0356 counts.reserve(fRdfs.size());
0357 for (auto &rdf : fRdfs) {
0358 counts.push_back(rdf.Count());
0359 }
0360 ROOT::RDF::RunGraphs({counts.begin(), counts.end()});
0361
0362 std::size_t totalFiltered = 0;
0363 for (auto &c : counts) {
0364 totalFiltered += c.GetValue();
0365 }
0366 fNumTrainingEntries = static_cast<std::size_t>(totalFiltered * (1.0f - fValidationSplit));
0367 fNumValidationEntries = totalFiltered - fNumTrainingEntries;
0368 }
0369
0370 ROOT::RDF::RNode &rdf = fRdfs[rdfIdx];
0371
0372
0373 std::vector<ULong64_t> rdfEntries;
0374 rdfEntries.reserve(endRow - startRow);
0375
0376 RClusterLoaderFunctor<Args...> loader(dest, fNumChunkCols, fVecSizes, fVecPadding, 0, rowOffset);
0377 ROOT::Internal::RDF::ChangeBeginAndEndEntries(rdf, startRow, endRow);
0378
0379 std::vector<std::string> colsWithEntry;
0380 colsWithEntry.reserve(fCols.size() + 1);
0381 colsWithEntry.push_back("rdfentry_");
0382 colsWithEntry.insert(colsWithEntry.end(), fCols.begin(), fCols.end());
0383
0384 rdf.Foreach(
0385 [&](ULong64_t entry, const Args &...cols) {
0386 rdfEntries.push_back(entry);
0387 loader(cols...);
0388 },
0389 colsWithEntry);
0390
0391 ROOT::Internal::RDF::ChangeBeginAndEndEntries(rdf, 0, fRdfSizes[rdfIdx]);
0392
0393 const std::size_t totalFiltered = rdfEntries.size();
0394 if (totalFiltered == 0) {
0395 return 0;
0396 }
0397 std::sort(rdfEntries.begin(), rdfEntries.end());
0398
0399 const std::size_t trainRemaining = fNumTrainingEntries - fAccumulatedFilteredForTrain;
0400 const std::size_t trainCount =
0401 std::min(static_cast<std::size_t>(totalFiltered * (1.0f - fValidationSplit)), trainRemaining);
0402 const std::size_t valCount = totalFiltered - trainCount;
0403
0404 bool trainIsPrefix = true;
0405 if (fShuffle) {
0406
0407
0408 std::mt19937 g(fSetSeed + fAccumulatedFilteredForTrain);
0409 std::uniform_int_distribution<int> coin(0, 1);
0410 trainIsPrefix = coin(g);
0411 }
0412
0413
0414
0415
0416
0417
0418
0419 std::uint64_t boundary;
0420 if (trainIsPrefix) {
0421
0422 boundary = (trainCount < totalFiltered) ? rdfEntries[trainCount] : endRow;
0423 } else {
0424
0425 boundary = (valCount < totalFiltered) ? rdfEntries[valCount] : endRow;
0426 }
0427
0428 const std::uint64_t trainStart = trainIsPrefix ? startRow : boundary;
0429 const std::uint64_t trainEnd = trainIsPrefix ? boundary : endRow;
0430 const std::uint64_t valStart = trainIsPrefix ? boundary : startRow;
0431 const std::uint64_t valEnd = trainIsPrefix ? endRow : boundary;
0432
0433 if (trainCount > 0)
0434 fTrainingClusters.push_back({rdfIdx, trainStart, trainEnd, trainCount});
0435 if (valCount > 0)
0436 fValidationClusters.push_back({rdfIdx, valStart, valEnd, valCount});
0437
0438 fAccumulatedFilteredForTrain += trainCount;
0439 return trainCount;
0440 }
0441
0442 LoadClusterInto(dest, rdfIdx, startRow, endRow, rowOffset);
0443 return endRow - startRow;
0444 }
0445
0446
0447
0448 void LoadValidationClusterInto(RFlat2DMatrix &dest, std::size_t rdfIdx, std::uint64_t startRow, std::uint64_t endRow,
0449 std::size_t rowOffset = 0)
0450 {
0451 LoadClusterInto(dest, rdfIdx, startRow, endRow, rowOffset);
0452 }
0453
0454
0455
0456 void FinaliseSplitDiscovery()
0457 {
0458 if (fIsFiltered)
0459 fSplitDiscovered = true;
0460 }
0461
0462 bool IsSplitDiscovered() const { return !fIsFiltered || fSplitDiscovered; }
0463
0464
0465
0466 std::size_t GetNumTrainingEntries() const { return fNumTrainingEntries; }
0467 std::size_t GetNumValidationEntries() const { return fNumValidationEntries; }
0468 std::size_t GetNumChunkCols() const { return fNumChunkCols; }
0469
0470 const std::vector<RClusterRange> &GetTrainingClusters() const
0471 {
0472 return (fIsFiltered && !fSplitDiscovered) ? fAllClusters : fTrainingClusters;
0473 }
0474 const std::vector<RClusterRange> &GetValidationClusters() const { return fValidationClusters; }
0475
0476 std::size_t GetNumTrainingClusters() const
0477 {
0478 return (fIsFiltered && !fSplitDiscovered) ? fAllClusters.size() : fTrainingClusters.size();
0479 }
0480 std::size_t GetNumValidationClusters() const { return fValidationClusters.size(); }
0481 std::size_t GetNmTotalClusters() const { return fAllClusters.size(); }
0482 };
0483
0484 }
0485 #endif