From b4d42099dc8c6e67d8137ced6b8707725f94d654 Mon Sep 17 00:00:00 2001 From: Juan Miguel Carceller Date: Thu, 26 Mar 2026 13:29:17 +0100 Subject: [PATCH 1/4] Add lazy reading to ROOTReader and RNTupleReader --- include/podio/RNTupleReader.h | 25 +++++ include/podio/ROOTFrameData.h | 7 ++ include/podio/ROOTReader.h | 9 ++ include/podio/Reader.h | 5 + src/RNTupleReader.cc | 205 ++++++++++++++++++++++++++++++++++ src/ROOTFrameData.cc | 18 ++- src/ROOTReader.cc | 63 +++++++++++ src/Reader.cc | 22 ++++ 8 files changed, 350 insertions(+), 4 deletions(-) diff --git a/include/podio/RNTupleReader.h b/include/podio/RNTupleReader.h index 80b731b3e..65b39f035 100644 --- a/include/podio/RNTupleReader.h +++ b/include/podio/RNTupleReader.h @@ -6,8 +6,10 @@ #include "podio/utilities/DatamodelRegistryIOHelpers.h" #include "podio/utilities/RootHelpers.h" +#include #include #include +#include #include #include @@ -91,6 +93,15 @@ class RNTupleReader { std::unique_ptr readEntry(std::string_view name, const unsigned entry, const std::vector& collsToRead = {}); + /// Like readEntry, but collections not in collsToRead are loaded lazily on + /// first access rather than eagerly. Also skips reading GenericParameters. + /// + /// The reader must not advance to the next entry while the returned FrameData + /// (and any Frame built from it) is still being accessed. This is guaranteed + /// in the DataSource context where each slot has its own reader. + std::unique_ptr readEntryLazy(const std::string& name, unsigned entry, + const std::vector& collsToRead); + /// Get the names of all the available Frame categories in the current file(s). /// /// @returns The names of the available categores from the file @@ -158,6 +169,9 @@ class RNTupleReader { std::unordered_map>> m_readers{}; std::unordered_map> m_metadata_readers{}; std::vector m_filenames{}; + /// Per-category list of filenames (parallel to m_readers[category]); may differ from m_filenames + /// if some files don't contain a given category. + std::unordered_map> m_categoryFilenames{}; std::unordered_map m_entries{}; // Map category to a vector that contains at how many entries each reader starts @@ -173,6 +187,17 @@ class RNTupleReader { std::vector m_availableCategories{}; std::unordered_map> m_idTables{}; + + /// Cache of partial readers, keyed by (category + "|" + sorted-collsToRead-key) + readerIndex. + /// Each partial reader is opened with a minimal RNTupleModel containing only the fields + /// needed for the requested collections + std::map, std::unique_ptr> + m_partialReaders{}; + + /// Return or lazily create a partial ROOT::RNTupleReader for the given category, + /// readerIndex, and set of collections. + root_compat::RNTupleReader& getOrCreatePartialReader(const std::string& category, size_t readerIndex, + const std::vector& collsToRead); }; } // namespace podio diff --git a/include/podio/ROOTFrameData.h b/include/podio/ROOTFrameData.h index 8fde8ac9f..b1b95596f 100644 --- a/include/podio/ROOTFrameData.h +++ b/include/podio/ROOTFrameData.h @@ -5,6 +5,7 @@ #include "podio/CollectionIDTable.h" #include "podio/GenericParameters.h" +#include #include #include #include @@ -17,6 +18,8 @@ class ROOTFrameData { public: using BufferMap = std::unordered_map; + /// Callback type for lazily reading collection buffers on demand + using LazyReadFn = std::function(const std::string&)>; ROOTFrameData() = delete; ~ROOTFrameData() = default; @@ -27,6 +30,9 @@ class ROOTFrameData { ROOTFrameData(BufferMap&& buffers, CollIDPtr&& idTable, podio::GenericParameters&& params); + /// Constructor with an optional lazy callback for loading collections on-demand + ROOTFrameData(BufferMap&& buffers, CollIDPtr&& idTable, podio::GenericParameters&& params, LazyReadFn&& lazyRead); + std::optional getCollectionBuffers(const std::string& name); podio::CollectionIDTable getIDTable() const; @@ -42,6 +48,7 @@ class ROOTFrameData { // This is co-owned by each FrameData and the original reader. (for now at least) CollIDPtr m_idTable{nullptr}; podio::GenericParameters m_parameters{}; + LazyReadFn m_lazyRead{nullptr}; }; } // namespace podio diff --git a/include/podio/ROOTReader.h b/include/podio/ROOTReader.h index f59c2941d..5cd483e6b 100644 --- a/include/podio/ROOTReader.h +++ b/include/podio/ROOTReader.h @@ -105,6 +105,15 @@ class ROOTReader { std::unique_ptr readEntry(std::string_view name, const unsigned entry, const std::vector& collsToRead = {}); + /// Like readEntry, but collections not in collsToRead are loaded lazily on + /// first access rather than eagerly. Also skips reading GenericParameters. + /// + /// The reader must not advance to the next entry while the returned FrameData + /// (and any Frame built from it) is still being accessed. This is guaranteed + /// in the DataSource context where each slot has its own reader. + std::unique_ptr readEntryLazy(const std::string& name, unsigned entry, + const std::vector& collsToRead); + /// Get the number of entries for the given name /// /// @param name The name of the category diff --git a/include/podio/Reader.h b/include/podio/Reader.h index 7d1f45822..cd5cee378 100644 --- a/include/podio/Reader.h +++ b/include/podio/Reader.h @@ -145,6 +145,11 @@ class Reader { return m_self->readFrame(name, index, collsToRead); } + /// Like readFrame but loads collections not in collsToRead lazily on first access. + /// Falls back to readFrame for non-ROOT backends. + /// See ROOTReader::readEntryLazy + podio::Frame readFrameLazy(const std::string& name, size_t index, const std::vector& collsToRead); + /// Read a specific frame of the "events" category /// /// @param index The event number to read diff --git a/src/RNTupleReader.cc b/src/RNTupleReader.cc index 4466ecddc..111cee55c 100644 --- a/src/RNTupleReader.cc +++ b/src/RNTupleReader.cc @@ -6,7 +6,10 @@ #include "podio/utilities/RootHelpers.h" #include "rootUtils.h" +#include #include +#include +#include #include #include @@ -110,6 +113,7 @@ void RNTupleReader::openFiles(const std::vector& filenames) { #else m_readers[category].emplace_back(root_compat::RNTupleReader::Open(category, filename)); #endif + m_categoryFilenames[category].emplace_back(filename); m_readerEntries[category].push_back(m_readerEntries[category].back() + m_readers[category].back()->GetNEntries()); } catch (const RException&) { @@ -138,6 +142,72 @@ std::vector RNTupleReader::getAvailableCategories() const { return cats; } +root_compat::RNTupleReader& RNTupleReader::getOrCreatePartialReader(const std::string& category, size_t readerIndex, + const std::vector& collsToRead) { + // Build a stable cache key: category + "|" + sorted collection names + std::vector sortedColls = collsToRead; + std::ranges::sort(sortedColls); + std::string collKey; + for (const auto& c : sortedColls) { + collKey += c; + collKey += ','; + } + + const auto cacheKey = std::make_tuple(category, collKey, readerIndex); + auto it = m_partialReaders.find(cacheKey); + if (it != m_partialReaders.end()) { + return *it->second; + } + + const auto& collInfo = m_collectionInfo[category]; + std::vector neededFieldNames; + for (const auto& coll : collInfo) { + if (std::ranges::find(collsToRead, coll.name) == collsToRead.end()) { + continue; + } + if (coll.isSubset) { + neededFieldNames.emplace_back(root_utils::subsetBranch(coll.name)); + } else { + neededFieldNames.emplace_back(coll.name); + const auto relVecNames = podio::DatamodelRegistry::instance().getRelationNames(coll.dataType); + for (const auto& relName : relVecNames.relations) { + neededFieldNames.emplace_back(root_utils::refBranch(coll.name, relName)); + } + for (const auto& vecName : relVecNames.vectorMembers) { + neededFieldNames.emplace_back(root_utils::vecBranch(coll.name, vecName)); + } + } + } + + // Build a minimal RNTupleModel from the full reader's descriptor + auto& fullReader = *m_readers[category][readerIndex]; + const auto& desc = fullReader.GetDescriptor(); + + ROOT::RCreateFieldOptions fieldOpts; + fieldOpts.SetEmulateUnknownTypes(true); + fieldOpts.SetReturnInvalidOnError(true); + + auto smallModel = ROOT::RNTupleModel::CreateBare(); + const auto& topFieldDesc = desc.GetFieldDescriptor(desc.GetFieldZeroId()); + for (const auto& fieldDesc : desc.GetFieldIterable(topFieldDesc)) { + const auto& fn = fieldDesc.GetFieldName(); + if (std::ranges::find(neededFieldNames, fn) != neededFieldNames.end()) { + auto field = fieldDesc.CreateField(desc, fieldOpts); + if (field) { + smallModel->AddField(std::move(field)); + } + } + } + smallModel->Freeze(); + + // Open a new partial reader with this model, using the same file as the full reader + const auto& filename = m_categoryFilenames[category][readerIndex]; + auto partialReader = root_compat::RNTupleReader::Open(std::move(smallModel), category, filename); + + auto [insertIt, _] = m_partialReaders.emplace(cacheKey, std::move(partialReader)); + return *insertIt->second; +} + std::unique_ptr RNTupleReader::readNextEntry(std::string_view category, const std::vector& collsToRead) { return readEntry(category, m_entries[category], collsToRead); @@ -241,4 +311,139 @@ std::unique_ptr RNTupleReader::readEntry(std::string_view categor return std::make_unique(std::move(buffers), m_idTables[category], std::move(parameters)); } +std::unique_ptr RNTupleReader::readEntryLazy(const std::string& category, const unsigned entNum, + const std::vector& collsToRead) { + if (m_totalEntries.find(category) == m_totalEntries.end()) { + getEntries(category); + } + if (entNum >= m_totalEntries[category]) { + return nullptr; + } + + if (m_collectionInfo.find(category) == m_collectionInfo.end()) { + if (!initCategory(category)) { + return nullptr; + } + } + + const auto& collInfo = m_collectionInfo[category]; + + m_entries[category] = entNum + 1; + + const auto upper = std::ranges::upper_bound(m_readerEntries[category], entNum); + const auto localEntry = entNum - *(upper - 1); + const auto readerIndex = static_cast(upper - 1 - m_readerEntries[category].begin()); + + // Create one REntry for this readEntryLazy call. + auto& activeReader = collsToRead.empty() ? *m_readers[category][readerIndex] + : getOrCreatePartialReader(category, readerIndex, collsToRead); + const auto dentry = activeReader.CreateEntry(); + + ROOTFrameData::BufferMap buffers; + for (const auto& coll : collInfo) { + if (!collsToRead.empty() && std::ranges::find(collsToRead, coll.name) == collsToRead.end()) { + continue; + } + const auto& collType = coll.dataType; + const auto& bufferFactory = podio::CollectionBufferFactory::instance(); + auto maybeBuffers = bufferFactory.createBuffers(collType, coll.schemaVersion, coll.isSubset); + if (!maybeBuffers) { + std::cerr << "WARNING: Buffers couldn't be created for collection " << coll.name << " of type " << coll.dataType + << " and schema version " << coll.schemaVersion << std::endl; + continue; + } + auto& collBuffers = maybeBuffers.value(); + + try { + if (coll.isSubset) { + const auto brName = root_utils::subsetBranch(coll.name); + const auto vec = new std::vector; + dentry->BindRawPtr(brName, vec); + collBuffers.references->at(0) = std::unique_ptr>(vec); + } else { + dentry->BindRawPtr(coll.name, collBuffers.data); + + const auto relVecNames = podio::DatamodelRegistry::instance().getRelationNames(collType); + for (size_t j = 0; j < relVecNames.relations.size(); ++j) { + const auto relName = relVecNames.relations[j]; + const auto vec = new std::vector; + const auto brName = root_utils::refBranch(coll.name, relName); + dentry->BindRawPtr(brName, vec); + collBuffers.references->at(j) = std::unique_ptr>(vec); + } + + for (size_t j = 0; j < relVecNames.vectorMembers.size(); ++j) { + const auto vecName = relVecNames.vectorMembers[j]; + const auto brName = root_utils::vecBranch(coll.name, vecName); + dentry->BindRawPtr(brName, collBuffers.vectorMembers->at(j).second); + } + } + } catch (const RException&) { + collBuffers.deleteBuffers = {}; + continue; + } + + buffers.emplace(coll.name, std::move(collBuffers)); + } + + activeReader.LoadEntry(localEntry, *dentry); + + // Lazy callback: called only for collections not in collsToRead + auto lazyRead = [this, category, localEntry, + readerIndex](const std::string& name) -> std::optional { + const auto& lazyCollInfo = m_collectionInfo[category]; + auto it = std::ranges::find(lazyCollInfo, name, &root_utils::CollectionWriteInfo::name); + if (it == lazyCollInfo.end()) { + return std::nullopt; + } + const root_utils::CollectionWriteInfo& coll = *it; + const auto& collType = coll.dataType; + const auto& bufferFactory = podio::CollectionBufferFactory::instance(); + auto maybeBuffers = bufferFactory.createBuffers(collType, coll.schemaVersion, coll.isSubset); + if (!maybeBuffers) { + return std::nullopt; + } + auto& collBuffers = maybeBuffers.value(); + + // Use a per-collection partial reader (cached) for cheap CreateEntry(). + auto& lazyPartialReader = getOrCreatePartialReader(category, readerIndex, {name}); + const auto lazyDentry = lazyPartialReader.CreateEntry(); + try { + if (coll.isSubset) { + const auto brName = root_utils::subsetBranch(coll.name); + const auto vec = new std::vector; + lazyDentry->BindRawPtr(brName, vec); + collBuffers.references->at(0) = std::unique_ptr>(vec); + } else { + lazyDentry->BindRawPtr(coll.name, collBuffers.data); + + const auto relVecNames = podio::DatamodelRegistry::instance().getRelationNames(collType); + for (size_t j = 0; j < relVecNames.relations.size(); ++j) { + const auto relName = relVecNames.relations[j]; + const auto vec = new std::vector; + const auto brName = root_utils::refBranch(coll.name, relName); + lazyDentry->BindRawPtr(brName, vec); + collBuffers.references->at(j) = std::unique_ptr>(vec); + } + + for (size_t j = 0; j < relVecNames.vectorMembers.size(); ++j) { + const auto vecName = relVecNames.vectorMembers[j]; + const auto brName = root_utils::vecBranch(coll.name, vecName); + lazyDentry->BindRawPtr(brName, collBuffers.vectorMembers->at(j).second); + } + } + } catch (const RException&) { + collBuffers.deleteBuffers = {}; + return std::nullopt; + } + + lazyPartialReader.LoadEntry(localEntry, *lazyDentry); + return {std::move(collBuffers)}; + }; + + // Skip GenericParameters. DataSource users never need event metadata + return std::make_unique(std::move(buffers), m_idTables[category], podio::GenericParameters{}, + std::move(lazyRead)); +} + } // namespace podio diff --git a/src/ROOTFrameData.cc b/src/ROOTFrameData.cc index cbd690cc8..a5fbf2e30 100644 --- a/src/ROOTFrameData.cc +++ b/src/ROOTFrameData.cc @@ -6,13 +6,23 @@ ROOTFrameData::ROOTFrameData(BufferMap&& buffers, CollIDPtr&& idTable, podio::Ge m_buffers(std::move(buffers)), m_idTable(std::move(idTable)), m_parameters(std::move(params)) { } +ROOTFrameData::ROOTFrameData(BufferMap&& buffers, CollIDPtr&& idTable, podio::GenericParameters&& params, + LazyReadFn&& lazyRead) : + m_buffers(std::move(buffers)), + m_idTable(std::move(idTable)), + m_parameters(std::move(params)), + m_lazyRead(std::move(lazyRead)) { +} + std::optional ROOTFrameData::getCollectionBuffers(const std::string& name) { auto bufferHandle = m_buffers.extract(name); - if (bufferHandle.empty()) { - return std::nullopt; + if (!bufferHandle.empty()) { + return {std::move(bufferHandle.mapped())}; } - - return {std::move(bufferHandle.mapped())}; + if (m_lazyRead) { + return m_lazyRead(name); + } + return std::nullopt; } podio::CollectionIDTable ROOTFrameData::getIDTable() const { diff --git a/src/ROOTReader.cc b/src/ROOTReader.cc index 65aaa7429..1ec767428 100644 --- a/src/ROOTReader.cc +++ b/src/ROOTReader.cc @@ -147,6 +147,69 @@ std::unique_ptr ROOTReader::readEntry(ROOTReader::CategoryInfo& c return std::make_unique(std::move(buffers), catInfo.table, std::move(parameters)); } +std::unique_ptr ROOTReader::readEntryLazy(const std::string& name, unsigned entNum, + const std::vector& collsToRead) { + auto& catInfo = getCategoryInfo(name); + catInfo.entry = entNum; + + if (!catInfo.chain) { + return nullptr; + } + if (catInfo.entry >= static_cast(catInfo.chain->GetEntries())) { + return nullptr; + } + + // Make sure to not silently ignore non-existant but requested collections + if (!collsToRead.empty()) { + for (const auto& collName : collsToRead) { + if (std::ranges::find(catInfo.storedClasses, collName, &detail::NamedCollInfo::name) == catInfo.storedClasses.end()) { + throw std::invalid_argument(collName + " is not available from Frame"); + } + } + } + + // After switching trees in the chain, branch pointers get invalidated so + // they need to be reassigned. + // NOTE: root 6.22/06 requires that we get completely new branches here, + // with 6.20/04 we could just re-set them + const auto preTreeNo = catInfo.chain->GetTreeNumber(); + const auto localEntry = catInfo.chain->LoadTree(catInfo.entry); + const auto treeChange = catInfo.chain->GetTreeNumber() != preTreeNo; + // Also need to make sure to handle the first event + const auto reloadBranches = treeChange || localEntry == 0; + + // Eagerly read only the requested collections + ROOTFrameData::BufferMap buffers; + for (size_t i = 0; i < catInfo.storedClasses.size(); ++i) { + if (!collsToRead.empty() && std::ranges::find(collsToRead, catInfo.storedClasses[i].name) == collsToRead.end()) { + continue; + } + auto collBuffers = getCollectionBuffers(catInfo, i, reloadBranches, localEntry); + if (collBuffers) { + buffers.emplace(catInfo.storedClasses[i].name, std::move(collBuffers.value())); + } + } + + catInfo.entry++; + + // Create a lazy-read callback for collections not loaded eagerly. + // Only safe if this reader won't advance to the next entry until all + // the reads for this event are complete + auto lazyRead = [this, &catInfo, localEntry, + reloadBranches](const std::string& collName) -> std::optional { + auto it = std::ranges::find(catInfo.storedClasses, collName, &detail::NamedCollInfo::name); + if (it == catInfo.storedClasses.end()) { + return std::nullopt; + } + const auto idx = static_cast(std::distance(catInfo.storedClasses.begin(), it)); + return getCollectionBuffers(catInfo, idx, reloadBranches, localEntry); + }; + + // Skip GenericParameters. DataSource users never need event metadata + return std::make_unique(std::move(buffers), catInfo.table, podio::GenericParameters{}, + std::move(lazyRead)); +} + std::optional ROOTReader::getCollectionBuffers(ROOTReader::CategoryInfo& catInfo, size_t iColl, bool reloadBranches, unsigned int localEntry) { diff --git a/src/Reader.cc b/src/Reader.cc index 922d62eea..b30df269b 100644 --- a/src/Reader.cc +++ b/src/Reader.cc @@ -89,6 +89,28 @@ Reader makeReader(const std::vector& filenames) { throw std::runtime_error("Unknown file extension: " + suffix); } +podio::Frame Reader::readFrameLazy(const std::string& name, size_t index, const std::vector& collsToRead) { + if (auto* rootReader = dynamic_cast*>(m_self.get())) { + auto maybeFrame = rootReader->m_reader->readEntryLazy(name, index, collsToRead); + if (maybeFrame) { + return podio::Frame(std::move(maybeFrame)); + } + throw std::runtime_error("Failed reading category " + name + " at frame " + std::to_string(index) + + " (reading beyond bounds?)"); + } +#if PODIO_ENABLE_RNTUPLE + if (auto* rntReader = dynamic_cast*>(m_self.get())) { + auto maybeFrame = rntReader->m_reader->readEntryLazy(name, index, collsToRead); + if (maybeFrame) { + return podio::Frame(std::move(maybeFrame)); + } + throw std::runtime_error("Failed reading category " + name + " at frame " + std::to_string(index) + + " (reading beyond bounds?)"); + } +#endif + return readFrame(name, index, collsToRead); +} + std::optional> Reader::getSizeStats(std::string_view category) { if (const auto* rootReader = dynamic_cast*>(m_self.get())) { return rootReader->m_reader->getSizeStats(category); From 5153bb4a68d5e99f99432edf4c5604f85fae524f Mon Sep 17 00:00:00 2001 From: Juan Miguel Carceller Date: Thu, 26 Mar 2026 13:29:27 +0100 Subject: [PATCH 2/4] Refactor DataSource --- include/podio/DataSource.h | 38 ++++++++---- src/DataSource.cc | 123 +++++++++++++++++++------------------ 2 files changed, 90 insertions(+), 71 deletions(-) diff --git a/include/podio/DataSource.h b/include/podio/DataSource.h index 33f6ab921..6ed137765 100644 --- a/include/podio/DataSource.h +++ b/include/podio/DataSource.h @@ -7,6 +7,7 @@ #include // ROOT +#include #include #include @@ -14,6 +15,7 @@ #include #include #include +#include #include #include @@ -100,47 +102,59 @@ class DataSource : public ROOT::RDF::RDataSource { std::string GetLabel() override { return "PODIO Datasource"; - }; + } + + // Legacy API + std::vector GetColumnReadersImpl(std::string_view, const std::type_info&) override { + return {}; + } + + std::size_t GetNFiles() const override { + return m_filePathList.size(); + } -protected: /// - /// @brief Type-erased vector of pointers to pointers to column - /// values --- one per slot. + /// @brief Returns a column reader for the given slot and column. /// - std::vector GetColumnReadersImpl(std::string_view name, const std::type_info& typeInfo) override; + std::unique_ptr GetColumnReaders(unsigned int slot, std::string_view name, + const std::type_info& tid) override; +protected: std::string AsString() override { return "Podio data source"; } private: - /// Number of slots/threads - unsigned int m_nSlots = 1; - /// Input filename std::vector m_filePathList = {}; /// Total number of events ULong64_t m_nEvents = 0; - /// Ranges of events available to be processed - std::vector> m_rangesAvailable = {}; - - /// Ranges of events available ever created + /// All entry ranges, fixed after SetNSlots std::vector> m_rangesAll = {}; + /// Cursor into m_rangesAll for GetEntryRanges, reset each Initialize() + size_t m_rangesCursor = 0; + /// Column names std::vector m_columnNames{}; /// Column types std::vector m_columnTypes = {}; + /// Fast column name -> index lookup + std::unordered_map m_columnIndex{}; + /// Collections, m_Collections[columnIndex][slotIndex] std::vector> m_Collections = {}; /// Active collections std::vector m_activeCollections = {}; + /// Names of active collections, kept in sync with m_activeCollections + std::vector m_activeCollectionNames{}; + /// Root podio readers std::vector> m_podioReaders = {}; diff --git a/src/DataSource.cc b/src/DataSource.cc index 5e4f0779f..9e223967f 100644 --- a/src/DataSource.cc +++ b/src/DataSource.cc @@ -6,21 +6,35 @@ #include // ROOT -#include +#include // STL #include -#include #include namespace podio { + +// Column reader that wraps a pointer to the per-slot CollectionBase* pointer +class PodioColumnReader : public ROOT::Detail::RDF::RColumnReaderBase { + const podio::CollectionBase** fPtr; + +public: + explicit PodioColumnReader(const podio::CollectionBase** ptr) : fPtr(ptr) { + } + void* GetImpl(Long64_t) override { + // Return the actual collection pointer (T*), not the address of the storage (T**) + // RColumnReaderBase::Get does *static_cast(GetImpl()), so we return the T* itself + return const_cast(static_cast(*fPtr)); + } +}; + DataSource::DataSource(const std::string& filePath, int nEvents, const std::vector& collNames) : DataSource(utils::expand_glob(filePath), nEvents, collNames) { } DataSource::DataSource(const std::vector& filePathList, int nEvents, const std::vector& collNames) : - m_nSlots{1}, m_filePathList{filePathList} { + m_filePathList{filePathList} { SetupInput(nEvents, collNames); } @@ -55,11 +69,12 @@ void DataSource::SetupInput(int nEvents, const std::vector& collsTo m_nEvents = nEventsInFiles; } - // Get collections stored in the files + // Get collections stored in the files and build fast lookup map std::vector collNames = frame.getAvailableCollections(); for (auto&& collName : collNames) { const podio::CollectionBase* coll = frame.get(collName); if (coll) { + m_columnIndex[collName] = m_columnNames.size(); m_columnNames.emplace_back(std::move(collName)); m_columnTypes.emplace_back(coll->getTypeName()); } @@ -67,62 +82,54 @@ void DataSource::SetupInput(int nEvents, const std::vector& collsTo } void DataSource::SetNSlots(unsigned int nSlots) { - m_nSlots = nSlots; - - if (m_nSlots > m_nEvents) { - throw std::runtime_error("podio::DataSource: Number of events too small!"); - } + RDataSource::SetNSlots(nSlots); - int eventsPerSlot = m_nEvents / m_nSlots; - for (size_t i = 0; i < (m_nSlots - 1); ++i) { + // Build one range per slot; if there are fewer events than slots, cap at m_nEvents ranges + const unsigned int effectiveSlots = std::min(fNSlots, static_cast(m_nEvents)); + const ULong64_t eventsPerSlot = m_nEvents / effectiveSlots; + for (size_t i = 0; i < (effectiveSlots - 1); ++i) { m_rangesAll.emplace_back(eventsPerSlot * i, eventsPerSlot * (i + 1)); } - m_rangesAll.emplace_back(eventsPerSlot * (m_nSlots - 1), m_nEvents); - m_rangesAvailable = m_rangesAll; + m_rangesAll.emplace_back(eventsPerSlot * (effectiveSlots - 1), m_nEvents); + m_rangesCursor = 0; - // Initialize set of addresses needed - m_Collections.resize(m_columnNames.size(), std::vector(m_nSlots, nullptr)); + // Collections indexed [column][slot] + m_Collections.resize(m_columnNames.size(), std::vector(fNSlots, nullptr)); // Initialize podio readers - for (size_t i = 0; i < m_nSlots; ++i) { + for (size_t i = 0; i < fNSlots; ++i) { m_podioReaders.emplace_back(std::make_unique(podio::makeReader(m_filePathList))); } - for (size_t i = 0; i < m_nSlots; ++i) { + for (size_t i = 0; i < fNSlots; ++i) { m_frames.emplace_back(std::make_unique()); } } void DataSource::Initialize() { + m_rangesCursor = 0; } std::vector> DataSource::GetEntryRanges() { - std::vector> rangesToBeProcessed; - for (auto& range : m_rangesAvailable) { - rangesToBeProcessed.emplace_back(range.first, range.second); - if (rangesToBeProcessed.size() >= m_nSlots) { - break; - } + if (m_rangesCursor >= m_rangesAll.size()) { + return {}; } - - if (m_rangesAvailable.size() > m_nSlots) { - m_rangesAvailable.erase(m_rangesAvailable.begin(), m_rangesAvailable.begin() + m_nSlots); - } else { - m_rangesAvailable.erase(m_rangesAvailable.begin(), m_rangesAvailable.end()); - } - - return rangesToBeProcessed; + const size_t end = std::min(m_rangesCursor + fNSlots, m_rangesAll.size()); + std::vector> result(m_rangesAll.cbegin() + m_rangesCursor, + m_rangesAll.cbegin() + end); + m_rangesCursor = end; + return result; } void DataSource::InitSlot(unsigned int, ULong64_t) { } bool DataSource::SetEntry(unsigned int slot, ULong64_t entry) { - m_frames[slot] = - std::make_unique(m_podioReaders[slot]->readFrame(podio::Category::Event, entry, m_columnNames)); + m_frames[slot] = std::make_unique( + m_podioReaders[slot]->readFrameLazy(podio::Category::Event, entry, m_activeCollectionNames)); - for (auto& collectionIndex : m_activeCollections) { - m_Collections[collectionIndex][slot] = m_frames[slot]->get(m_columnNames.at(collectionIndex)); + for (auto collectionIndex : m_activeCollections) { + m_Collections[collectionIndex][slot] = m_frames[slot]->get(m_columnNames[collectionIndex]); } return true; @@ -134,45 +141,43 @@ void DataSource::FinalizeSlot(unsigned int) { void DataSource::Finalize() { } -std::vector DataSource::GetColumnReadersImpl(std::string_view columnName, const std::type_info&) { - auto itr = std::find(m_columnNames.begin(), m_columnNames.end(), columnName); - if (itr == m_columnNames.end()) { - std::string errMsg = "podio::DataSource: Can't find requested column \""; - errMsg += columnName; - errMsg += "\"!"; - throw std::runtime_error(errMsg); - } - auto columnIndex = std::distance(m_columnNames.begin(), itr); - m_activeCollections.emplace_back(columnIndex); - - std::vector columnReaders(m_nSlots); - for (size_t slotIndex = 0; slotIndex < m_nSlots; ++slotIndex) { - columnReaders[slotIndex] = static_cast(&m_Collections[columnIndex][slotIndex]); - } - - return columnReaders; -} - const std::vector& DataSource::GetColumnNames() const { return m_columnNames; } bool DataSource::HasColumn(std::string_view columnName) const { - return std::find(m_columnNames.begin(), m_columnNames.end(), columnName) != m_columnNames.end(); + return m_columnIndex.count(std::string(columnName)) > 0; } std::string DataSource::GetTypeName(std::string_view columnName) const { - auto itr = std::find(m_columnNames.begin(), m_columnNames.end(), columnName); - if (itr == m_columnNames.end()) { + auto itr = m_columnIndex.find(std::string(columnName)); + if (itr == m_columnIndex.end()) { std::string errMsg = "podio::DataSource: Type name for \""; errMsg += columnName; errMsg += "\" not found!"; throw std::runtime_error(errMsg); } - auto typeIndex = std::distance(m_columnNames.begin(), itr); + return m_columnTypes.at(itr->second); +} + +std::unique_ptr +DataSource::GetColumnReaders(unsigned int slot, std::string_view columnName, const std::type_info&) { + auto itr = m_columnIndex.find(std::string(columnName)); + if (itr == m_columnIndex.end()) { + std::string errMsg = "podio::DataSource: Can't find requested column \""; + errMsg += columnName; + errMsg += "\"!"; + throw std::runtime_error(errMsg); + } + const auto columnIndex = itr->second; + + if (std::find(m_activeCollections.begin(), m_activeCollections.end(), columnIndex) == m_activeCollections.end()) { + m_activeCollections.emplace_back(columnIndex); + m_activeCollectionNames.emplace_back(m_columnNames[columnIndex]); + } - return m_columnTypes.at(typeIndex); + return std::make_unique(&m_Collections[columnIndex][slot]); } ROOT::RDataFrame CreateDataFrame(const std::vector& filePathList, From 7a4993b3bb9ab2e0e8eb6280efe49a3a7672ed7e Mon Sep 17 00:00:00 2001 From: Juan Miguel Carceller Date: Thu, 26 Mar 2026 13:49:02 +0100 Subject: [PATCH 3/4] Add new and remove old tests for DataSource --- tests/root_io/check_datasource_output.cpp | 119 ++++++++++++++ tests/root_io/read_datasource.py | 16 -- tests/root_io/use_datasource.py | 179 ++++++++++++++++++++++ 3 files changed, 298 insertions(+), 16 deletions(-) create mode 100644 tests/root_io/check_datasource_output.cpp delete mode 100644 tests/root_io/read_datasource.py create mode 100644 tests/root_io/use_datasource.py diff --git a/tests/root_io/check_datasource_output.cpp b/tests/root_io/check_datasource_output.cpp new file mode 100644 index 000000000..3b3307fb0 --- /dev/null +++ b/tests/root_io/check_datasource_output.cpp @@ -0,0 +1,119 @@ +/** + * Checker for the datasource snapshot output produced by use_datasource.py. + */ + +#include +#include + +#include +#include +#include +#include +#include + +#define CHECK(condition, msg) \ + if (!(condition)) { \ + throw std::runtime_error(std::string("check_datasource_output: ") + (msg)); \ + } + +#define CHECK_EQUAL(actual, expected, label) \ + if (std::abs((actual) - (expected)) >= std::numeric_limits::min()) { \ + throw std::runtime_error(std::string("check_datasource_output: ") + (label) + " value mismatch: got " + \ + std::to_string(actual) + ", expected " + std::to_string(expected)); \ + } + +int main(int argc, const char* argv[]) { + std::string inputFile = "datasource_snapshot.root"; + if (argc == 2) { + inputFile = argv[1]; + } else if (argc > 2) { + std::cout << "Usage: " << argv[0] << " [FILE]" << std::endl; + return 1; + } + + TFile f(inputFile.c_str(), "READ"); + if (f.IsZombie()) { + std::cerr << "Could not open " << inputFile << std::endl; + return EXIT_FAILURE; + } + + auto* tree = f.Get("events"); + CHECK(tree != nullptr, "TTree 'events' not found in " + inputFile) + CHECK(tree->GetEntries() == 10, "Expected 10 entries in snapshot, got " + std::to_string(tree->GetEntries())) + + int event_number = 0; + double hit_energy_0 = 0., hit_energy_1 = 0., hit_energy_sum = 0.; + double cluster_energy_0 = 0., cluster_energy_1 = 0., cluster_energy_2 = 0.; + + double composite_hit_energy_0 = 0., composite_hit_energy_1 = 0.; + + double sub_cluster_energy_0 = 0., sub_cluster_energy_1 = 0.; + + double link0_weight = 0., link0_from_energy = 0., link0_to_energy = 0.; + double link1_weight = 0., link1_from_energy = 0., link1_to_energy = 0.; + + tree->SetBranchAddress("event_number", &event_number); + tree->SetBranchAddress("hit_energy_0", &hit_energy_0); + tree->SetBranchAddress("hit_energy_1", &hit_energy_1); + tree->SetBranchAddress("hit_energy_sum", &hit_energy_sum); + tree->SetBranchAddress("cluster_energy_0", &cluster_energy_0); + tree->SetBranchAddress("cluster_energy_1", &cluster_energy_1); + tree->SetBranchAddress("cluster_energy_2", &cluster_energy_2); + tree->SetBranchAddress("composite_hit_energy_0", &composite_hit_energy_0); + tree->SetBranchAddress("composite_hit_energy_1", &composite_hit_energy_1); + tree->SetBranchAddress("sub_cluster_energy_0", &sub_cluster_energy_0); + tree->SetBranchAddress("sub_cluster_energy_1", &sub_cluster_energy_1); + tree->SetBranchAddress("link0_weight", &link0_weight); + tree->SetBranchAddress("link0_from_energy", &link0_from_energy); + tree->SetBranchAddress("link0_to_energy", &link0_to_energy); + tree->SetBranchAddress("link1_weight", &link1_weight); + tree->SetBranchAddress("link1_from_energy", &link1_from_energy); + tree->SetBranchAddress("link1_to_energy", &link1_to_energy); + + for (Long64_t entry = 0; entry < tree->GetEntries(); ++entry) { + tree->GetEntry(entry); + + const int i = event_number; // event_number == i by construction + CHECK(i >= 0 && i < 10, "event_number out of expected range [0,9]: " + std::to_string(i)) + const std::string ev = " (event " + std::to_string(i) + ")"; + + CHECK_EQUAL(hit_energy_0, 23. + i, "hit_energy_0" + ev) + CHECK_EQUAL(hit_energy_1, 12. + i, "hit_energy_1" + ev) + CHECK_EQUAL(hit_energy_sum, 35. + 2. * i, "hit_energy_sum" + ev) + + CHECK_EQUAL(cluster_energy_0, 23. + i, "cluster_energy_0" + ev) + CHECK_EQUAL(cluster_energy_1, 12. + i, "cluster_energy_1" + ev) + CHECK_EQUAL(cluster_energy_2, 35. + 2. * i, "cluster_energy_2" + ev) + + CHECK_EQUAL(hit_energy_sum, cluster_energy_2, "hit_energy_sum vs cluster_energy_2" + ev) + CHECK_EQUAL(hit_energy_0 + hit_energy_1, cluster_energy_2, "hit0+hit1 vs cluster_energy_2" + ev) + CHECK_EQUAL(cluster_energy_0 + cluster_energy_1, cluster_energy_2, "sub-cluster sum vs cluster_energy_2" + ev) + + CHECK_EQUAL(composite_hit_energy_0, 23. + i, "composite_hit_energy_0 (cluster->hit relation)" + ev) + CHECK_EQUAL(composite_hit_energy_1, 12. + i, "composite_hit_energy_1 (cluster->hit relation)" + ev) + + CHECK_EQUAL(composite_hit_energy_0, hit_energy_0, "composite_hit_energy_0 vs hit_energy_0" + ev) + CHECK_EQUAL(composite_hit_energy_1, hit_energy_1, "composite_hit_energy_1 vs hit_energy_1" + ev) + + CHECK_EQUAL(sub_cluster_energy_0, 23. + i, "sub_cluster_energy_0 (cluster->cluster relation)" + ev) + CHECK_EQUAL(sub_cluster_energy_1, 12. + i, "sub_cluster_energy_1 (cluster->cluster relation)" + ev) + + CHECK_EQUAL(sub_cluster_energy_0, cluster_energy_0, "sub_cluster_energy_0 vs cluster_energy_0" + ev) + CHECK_EQUAL(sub_cluster_energy_1, cluster_energy_1, "sub_cluster_energy_1 vs cluster_energy_1" + ev) + + CHECK_EQUAL(link0_weight, 0.0, "link0_weight" + ev) + CHECK_EQUAL(link0_from_energy, 23. + i, "link0_from_energy (link->hit relation)" + ev) + CHECK_EQUAL(link0_to_energy, 12. + i, "link0_to_energy (link->cluster relation)" + ev) + CHECK_EQUAL(link0_from_energy, hit_energy_0, "link0_from_energy vs hit_energy_0" + ev) + CHECK_EQUAL(link0_to_energy, cluster_energy_1, "link0_to_energy vs cluster_energy_1" + ev) + + CHECK_EQUAL(link1_weight, 0.5, "link1_weight" + ev) + CHECK_EQUAL(link1_from_energy, 12. + i, "link1_from_energy (link->hit relation)" + ev) + CHECK_EQUAL(link1_to_energy, 23. + i, "link1_to_energy (link->cluster relation)" + ev) + CHECK_EQUAL(link1_from_energy, hit_energy_1, "link1_from_energy vs hit_energy_1" + ev) + CHECK_EQUAL(link1_to_energy, cluster_energy_0, "link1_to_energy vs cluster_energy_0" + ev) + } + + std::cout << "check_datasource_output: all checks passed" << std::endl; + return EXIT_SUCCESS; +} diff --git a/tests/root_io/read_datasource.py b/tests/root_io/read_datasource.py deleted file mode 100644 index b8528568d..000000000 --- a/tests/root_io/read_datasource.py +++ /dev/null @@ -1,16 +0,0 @@ -#!/usr/bin/env python3 -"""Small test case for checking DataSource based creating RDataFrames is accessible from python""" - -import ROOT -from podio.data_source import CreateDataFrame # pylint: disable=import-error, no-name-in-module - -if ROOT.gSystem.Load("libTestDataModelDict") < 0: - raise RuntimeError("Could not load TestDataModel dictionary") - -rdf = CreateDataFrame("example_frame.root") - -assert rdf.Count().GetValue() == 10 - -rdf = CreateDataFrame("example_frame_?.root") - -assert rdf.Count().GetValue() == 20 diff --git a/tests/root_io/use_datasource.py b/tests/root_io/use_datasource.py new file mode 100644 index 000000000..698834196 --- /dev/null +++ b/tests/root_io/use_datasource.py @@ -0,0 +1,179 @@ +#!/usr/bin/env python3 +"""Test to exercise DataSource thoroughly""" + +import sys +import ROOT +from podio.data_source import CreateDataFrame # pylint: disable=import-error, no-name-in-module + +input_file = sys.argv[1] if len(sys.argv) > 1 else "example_frame.root" +snapshot_file = sys.argv[2] if len(sys.argv) > 2 else "datasource_snapshot.root" + +if ROOT.gSystem.Load("libTestDataModelDict") < 0: + raise RuntimeError("Could not load TestDataModel dictionary") + +ROOT.gInterpreter.ProcessLine("using namespace ROOT::VecOps;") + +ROOT.gInterpreter.Declare( + """ +#include +#include +#include +#include +""" +) + +# Declare helpers that extract quantities from the test collections +ROOT.gInterpreter.Declare( + """ +RVec getHitEnergies(const ExampleHitCollection& hits) { + RVec v; + v.reserve(hits.size()); + for (const auto& h : hits) { + v.push_back(h.energy()); + } + return v; +} + +RVec getClusterEnergies(const ExampleClusterCollection& clusters) { + RVec v; + v.reserve(clusters.size()); + for (const auto& c : clusters) { + v.push_back(c.energy()); + } + return v; +} + +int getEventNumber(const EventInfoCollection& info) { + return info[0].Number(); +} + +RVec getCompositeClusterHitEnergies(const ExampleClusterCollection& clusters) { + RVec v; + for (const auto& hit : clusters[2].Hits()) { + v.push_back(hit.energy()); + } + return v; +} + +RVec getSubClusterEnergies(const ExampleClusterCollection& clusters) { + RVec v; + for (const auto& sub : clusters[2].Clusters()) { + v.push_back(sub.energy()); + } + return v; +} + +// Follow each link From (hit) and To (cluster) relations and return: +// [weight0, hit_energy0, cluster_energy0, weight1, hit_energy1, cluster_energy1] +RVec getLinkInfo(const TestLinkCollection& links) { + RVec v; + for (const auto& link : links) { + v.push_back(link.getWeight()); + v.push_back(link.getFrom().energy()); + v.push_back(link.getTo().energy()); + } + return v; +} +""" +) + +rdf = CreateDataFrame(input_file) + +assert rdf.Count().GetValue() == 10, f"Expected 10 events in {input_file}" + +rdf = ( + rdf.Define("hit_energies", "getHitEnergies(hits)") + .Define("hit_energy_0", "hit_energies[0]") + .Define("hit_energy_1", "hit_energies[1]") + .Define("hit_energy_sum", "Sum(hit_energies)") + .Define("cluster_energies", "getClusterEnergies(clusters)") + .Define("cluster_energy_0", "cluster_energies[0]") + .Define("cluster_energy_1", "cluster_energies[1]") + .Define("cluster_energy_2", "cluster_energies[2]") + .Define("event_number", "getEventNumber(info)") + .Define("composite_hit_energies", "getCompositeClusterHitEnergies(clusters)") + .Define("composite_hit_energy_0", "composite_hit_energies[0]") + .Define("composite_hit_energy_1", "composite_hit_energies[1]") + .Define("sub_cluster_energies", "getSubClusterEnergies(clusters)") + .Define("sub_cluster_energy_0", "sub_cluster_energies[0]") + .Define("sub_cluster_energy_1", "sub_cluster_energies[1]") + .Define("link_info", "getLinkInfo(links)") + .Define("link0_weight", "link_info[0]") + .Define("link0_from_energy", "link_info[1]") + .Define("link0_to_energy", "link_info[2]") + .Define("link1_weight", "link_info[3]") + .Define("link1_from_energy", "link_info[4]") + .Define("link1_to_energy", "link_info[5]") +) + +# print(rdf.Describe()) +# print(rdf.Display(["event_number", "hit_energy_0", "cluster_energy_2", "composite_hit_energy_0", "sub_cluster_energy_0"]).AsString()) + +# 10 events, i = 0..9 +# hit_energy_0 = 23+i +mean_hit0 = rdf.Mean("hit_energy_0").GetValue() +assert abs(mean_hit0 - 27.5) < 1e-9, f"Mean of hit_energy_0 should be 27.5, got {mean_hit0}" + +# cluster_energy_2 = 35+2*i +mean_clu2 = rdf.Mean("cluster_energy_2").GetValue() +assert abs(mean_clu2 - 44.0) < 1e-9, f"Mean of cluster_energy_2 should be 44.0, got {mean_clu2}" + +# event_number = i +mean_evtnum = rdf.Mean("event_number").GetValue() +assert abs(mean_evtnum - 4.5) < 1e-9, f"Mean of event_number should be 4.5, got {mean_evtnum}" + +# composite_hit_energy_0 follows clusters[2].Hits()[0] == hits[0], energy = 23+i -> Mean = 27.5 +mean_chit0 = rdf.Mean("composite_hit_energy_0").GetValue() +assert ( + abs(mean_chit0 - 27.5) < 1e-9 +), f"Mean of composite_hit_energy_0 (cluster->hit relation) should be 27.5, got {mean_chit0}" + +# sub_cluster_energy_0 follows clusters[2].Clusters()[0] == clusters[0], energy = 23+i -> Mean = 27.5 +mean_sub0 = rdf.Mean("sub_cluster_energy_0").GetValue() +assert ( + abs(mean_sub0 - 27.5) < 1e-9 +), f"Mean of sub_cluster_energy_0 (cluster->cluster relation) should be 27.5, got {mean_sub0}" + +# link0_weight is always 0.0, link1_weight is always 0.5 +mean_w0 = rdf.Mean("link0_weight").GetValue() +assert abs(mean_w0 - 0.0) < 1e-9, f"Mean of link0_weight should be 0.0, got {mean_w0}" +mean_w1 = rdf.Mean("link1_weight").GetValue() +assert abs(mean_w1 - 0.5) < 1e-9, f"Mean of link1_weight should be 0.5, got {mean_w1}" + +# link0 goes hits[0]->clusters[1]: from-energy = 23+i -> Mean = 27.5, +# to-energy = 12+i -> Mean = 16.5 +mean_l0from = rdf.Mean("link0_from_energy").GetValue() +assert ( + abs(mean_l0from - 27.5) < 1e-9 +), f"Mean of link0_from_energy (link->hit relation) should be 27.5, got {mean_l0from}" +mean_l0to = rdf.Mean("link0_to_energy").GetValue() +assert ( + abs(mean_l0to - 16.5) < 1e-9 +), f"Mean of link0_to_energy (link->cluster relation) should be 16.5, got {mean_l0to}" + +rdf_even = rdf.Filter("event_number % 2 == 0") +assert rdf_even.Count().GetValue() == 5, "Expected 5 even-numbered events after filter" + +columns = [ + "event_number", + "hit_energy_0", + "hit_energy_1", + "hit_energy_sum", + "cluster_energy_0", + "cluster_energy_1", + "cluster_energy_2", + "composite_hit_energy_0", + "composite_hit_energy_1", + "sub_cluster_energy_0", + "sub_cluster_energy_1", + "link0_weight", + "link0_from_energy", + "link0_to_energy", + "link1_weight", + "link1_from_energy", + "link1_to_energy", +] + +rdf.Snapshot("events", snapshot_file, columns) + +print(f"All assertions passed, snapshot written to {snapshot_file}") From 3b332798cc58b11b78e66470e1711efd448b8b16 Mon Sep 17 00:00:00 2001 From: Juan Miguel Carceller Date: Thu, 26 Mar 2026 15:41:28 +0100 Subject: [PATCH 4/4] Fix compiler warning with GCC --- src/DataSource.cc | 2 ++ 1 file changed, 2 insertions(+) diff --git a/src/DataSource.cc b/src/DataSource.cc index 9e223967f..6103b9adc 100644 --- a/src/DataSource.cc +++ b/src/DataSource.cc @@ -21,6 +21,8 @@ class PodioColumnReader : public ROOT::Detail::RDF::RColumnReaderBase { public: explicit PodioColumnReader(const podio::CollectionBase** ptr) : fPtr(ptr) { } + PodioColumnReader(const PodioColumnReader&) = delete; + PodioColumnReader& operator=(const PodioColumnReader&) = delete; void* GetImpl(Long64_t) override { // Return the actual collection pointer (T*), not the address of the storage (T**) // RColumnReaderBase::Get does *static_cast(GetImpl()), so we return the T* itself