Skip to content

Commit 6ca7c23

Browse files
committed
[ntuple] add support for inheritance in SoA fields
1 parent d282aaf commit 6ca7c23

5 files changed

Lines changed: 236 additions & 78 deletions

File tree

tree/ntuple/inc/ROOT/RField/RFieldSoA.hxx

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -21,6 +21,7 @@
2121
#include <ROOT/RNTupleTypes.hxx>
2222

2323
#include <cstddef>
24+
#include <functional>
2425
#include <memory>
2526
#include <mutex>
2627
#include <string_view>
@@ -88,8 +89,13 @@ class RSoAField : public RFieldBase {
8889
RSoAField(std::string_view fieldName, const RSoAField &source); ///< Used by CloneImpl
8990
RSoAField(std::string_view fieldName, TClass *clSoA);
9091

91-
/// Called during construction, picks up the (nested) member fields of the underlying record type(s)
92+
/// Called during construction, picks up the (nested) member fields of the underlying record type(s) and its
93+
/// base classes.
9294
void CollectRecordMemberFields();
95+
/// For a nested SoA struct (either as a member of as a base class), use their fRecordMemberFields in this class,
96+
/// i.e. "unroll" the vectors in the nested SoA struct into the SoA base class.
97+
void GraftNestedMemberFields(const RSoAField &nestedSoA, std::size_t offsetInParent,
98+
std::function<RFieldBase *(const std::string &)> fnRecordFieldFinder);
9399

94100
void ReconstructSplitFields() const;
95101

tree/ntuple/src/RFieldMeta.cxx

Lines changed: 104 additions & 35 deletions
Original file line numberDiff line numberDiff line change
@@ -686,42 +686,110 @@ ROOT::Experimental::RSoAField::RSoAField(std::string_view fieldName, std::string
686686
{
687687
}
688688

689+
void ROOT::Experimental::RSoAField::GraftNestedMemberFields(
690+
const RSoAField &nestedSoA, std::size_t offsetInParent,
691+
std::function<RFieldBase *(const std::string &)> fnRecordFieldFinder)
692+
{
693+
const std::size_t nNestedRecordMemberFields = nestedSoA.fRecordMemberFields.size();
694+
695+
// The qualified field name of fields in nestedSoA->fRecordMemberFields will have a "<field name>._0."
696+
// prefix because these fields are rooted in a collection named after the nested SoA field.
697+
//
698+
// E.g., in the following example:
699+
//
700+
// struct SoA_A { struct Record_A {
701+
// SoA_B fB; Record_B fB;
702+
// }; };
703+
//
704+
// struct SoA_B { struct Record_B {
705+
// ROOT::RVec<float> fX; float fX;
706+
// }; };
707+
//
708+
// The on-disk schema of SoA_A is "collection of Record_A", and the on-disk schema of SoA_B is
709+
// "collection of Record_B".
710+
// The qualified field name of fX in the subfield hiararchy of SoA_A is <field name>._0.fB.fX.
711+
// The qualified field name of fX in the subfield hiararchy of SoA_B is <field name>._0.fX.
712+
const auto lenPrefix = nestedSoA.GetFieldName().length() + strlen("._0.");
713+
714+
for (std::size_t i = 0; i < nNestedRecordMemberFields; ++i) {
715+
const auto fieldNameForMatching =
716+
nestedSoA.GetFieldName() + "." + nestedSoA.fRecordMemberFields[i]->GetQualifiedFieldName().substr(lenPrefix);
717+
718+
fRecordMemberFields.emplace_back(fnRecordFieldFinder(fieldNameForMatching));
719+
fRecordMemberDeleters.emplace_back(GetDeleterOf(*nestedSoA.fRecordMemberFields[i]));
720+
fSoAMemberOffsets.emplace_back(offsetInParent + nestedSoA.fSoAMemberOffsets[i]);
721+
}
722+
}
723+
689724
void ROOT::Experimental::RSoAField::CollectRecordMemberFields()
690725
{
691-
std::vector<RFieldBase *> realRecordMemberFields;
726+
// Build a map of all subfields (nested) of the underlying record type. Map the fully qualified name of the
727+
// subfields to their field pointer, so that we can later match the subfields of the SoA class to their corresponding
728+
// fields in the underlying record type. Note that the members of the SoA class and the underlying record type
729+
// can have different ordering. However, the base classes of the SoA class and the underlying record type must match
730+
// in order.
731+
732+
std::vector<RFieldBase *> realRecordMemberFields; // Contains all subfields of the underlying record type
733+
// Qualified field name --> index in realRecordMemberFields
734+
std::unordered_map<std::string, std::size_t> recordFieldNameToIdx;
735+
736+
// Count the top-level subfields of the underlying record type for cross-check with the SoA type
692737
unsigned int nDirectRecordSubfields = 0;
738+
unsigned int nDirectRecordBases = 0;
739+
693740
for (auto itr = fSubfields[0]->begin(), iEnd = fSubfields[0]->end(); itr != iEnd; ++itr) {
694-
if (itr->GetParent() == fSubfields[0].get())
695-
nDirectRecordSubfields++;
741+
if (itr->GetParent() == fSubfields[0].get()) {
742+
(itr->GetFieldName()[0] == ':') ? nDirectRecordBases++ : nDirectRecordSubfields++;
743+
}
744+
745+
// Build the qualified field name for matching. We root the qualified field name at the underlying record type.
746+
auto qualifiedName = itr->GetFieldName();
747+
auto parent = itr->GetParent();
748+
while (parent != fSubfields[0].get()) {
749+
qualifiedName = parent->GetFieldName() + "." + qualifiedName;
750+
parent = parent->GetParent();
751+
}
752+
recordFieldNameToIdx[qualifiedName] = realRecordMemberFields.size();
753+
696754
realRecordMemberFields.emplace_back(&(*itr));
697755
}
698756

699-
for (const auto f : realRecordMemberFields) {
700-
if (f->GetFieldName()[0] == ':') {
701-
throw RException(R__FAIL("SoA fields with inheritance are currently unsupported"));
702-
}
757+
// Base classes are treated as unrolled nested SoA classes
758+
const auto *soaBases = fSoAClass->GetListOfBases();
759+
if (soaBases->GetSize() != static_cast<Int_t>(nDirectRecordBases)) {
760+
throw RException(R__FAIL(std::string("number of base classes don't match between SoA class ") + GetFieldName() +
761+
" and its underlying record type"));
703762
}
763+
for (unsigned int i = 0; i < static_cast<unsigned int>(soaBases->GetSize()); ++i) {
764+
auto base = static_cast<TBaseClass *>(soaBases->At(i));
765+
if (base->GetDelta() < 0) {
766+
throw RException(R__FAIL(std::string("virtual inheritance is not supported: ") + GetTypeName() +
767+
" virtually inherits from " + base->GetName()));
768+
}
769+
TClass *cl = base->GetClassPointer();
704770

705-
// Build map name --> index in realRecordMemberFields vector. Cut the name prefix up to the collection subfields
706-
// so that it matches the data member name
707-
std::unordered_map<std::string, std::size_t> recordFieldNameToIdx;
708-
{
709-
recordFieldNameToIdx.reserve(realRecordMemberFields.size());
710-
const auto lenFieldNamePrefix = GetFieldName().length() + strlen("._0.");
711-
for (std::size_t i = 0; i < realRecordMemberFields.size(); ++i) {
712-
recordFieldNameToIdx[realRecordMemberFields[i]->GetQualifiedFieldName().substr(lenFieldNamePrefix)] = i;
771+
const auto baseFieldName = std::string(":_") + std::to_string(i);
772+
773+
// SoA class `A` inherits from a SoA class `B` whose underlying record type is `X` if and only if the underlying
774+
// record type of `A` inherits from a type `X`.
775+
const auto underlyingBaseTypeName = ROOT::Internal::GetRNTupleSoARecord(cl);
776+
auto recordBaseField = realRecordMemberFields[recordFieldNameToIdx[baseFieldName]];
777+
if (underlyingBaseTypeName != recordBaseField->GetTypeName()) {
778+
throw RException(R__FAIL(std::string("inheritance of SoA class ") + GetFieldName() +
779+
" does not match its underlying record type"));
713780
}
714-
}
715781

716-
const auto *bases = fSoAClass->GetListOfBases();
717-
assert(bases);
718-
for (auto baseClass : ROOT::Detail::TRangeStaticCast<TBaseClass>(*bases)) {
719-
if (baseClass->GetDelta() < 0) {
720-
throw RException(R__FAIL(std::string("virtual inheritance is not supported: ") + GetTypeName() +
721-
" virtually inherits from " + baseClass->GetName()));
782+
std::unique_ptr<RSoAField> soaBaseField;
783+
try {
784+
soaBaseField = std::make_unique<RSoAField>(baseFieldName, cl->GetName());
785+
} catch (const RException &e) {
786+
throw RException(R__FAIL(std::string("invalid field type in base class: ") + cl->GetName() + " of SoA field " +
787+
GetFieldName() + " (" + e.what() + ")"));
722788
}
723-
// At a later point, we will support inheritance
724-
throw RException(R__FAIL("SoA fields with inheritance are currently unsupported"));
789+
790+
GraftNestedMemberFields(*soaBaseField, base->GetDelta(), [&](const std::string &name) {
791+
return realRecordMemberFields[recordFieldNameToIdx[name]];
792+
});
725793
}
726794

727795
unsigned int nMembers = 0;
@@ -753,17 +821,9 @@ void ROOT::Experimental::RSoAField::CollectRecordMemberFields()
753821
underlyingField->GetTypeName()));
754822
}
755823

756-
const auto lenFieldNamePrefix = soaField->GetFieldName().length() + strlen("._0.");
757-
758-
for (std::size_t i = 0; i < soaField->fRecordMemberFields.size(); ++i) {
759-
const auto f = soaField->fRecordMemberFields[i];
760-
const auto fieldNameForMatching =
761-
dmField->GetFieldName() + "." + f->GetQualifiedFieldName().substr(lenFieldNamePrefix);
762-
763-
fRecordMemberFields.emplace_back(realRecordMemberFields[recordFieldNameToIdx[fieldNameForMatching]]);
764-
fRecordMemberDeleters.emplace_back(GetDeleterOf(*soaField->fRecordMemberFields[i]));
765-
fSoAMemberOffsets.emplace_back(dataMember->GetOffset() + soaField->fSoAMemberOffsets[i]);
766-
}
824+
GraftNestedMemberFields(*soaField, dataMember->GetOffset(), [&](const std::string &name) {
825+
return realRecordMemberFields[recordFieldNameToIdx[name]];
826+
});
767827
} else if (auto vecField = dynamic_cast<RRVecField *>(dmField.get())) {
768828
if (vecField->begin()->GetTypeName() != underlyingField->GetTypeName() ||
769829
vecField->begin()->GetTypeAlias() != underlyingField->GetTypeAlias()) {
@@ -945,6 +1005,15 @@ void ROOT::Experimental::RSoAField::ReconstructSplitFields() const
9451005
fSplitFields = std::make_unique<std::vector<std::unique_ptr<ROOT::RFieldBase>>>();
9461006
fSplitOffsets = std::make_unique<std::vector<std::size_t>>();
9471007

1008+
unsigned int iBase = 0;
1009+
for (auto base : ROOT::Detail::TRangeStaticCast<TBaseClass>(*fSoAClass->GetListOfBases())) {
1010+
TClass *cl = base->GetClassPointer();
1011+
auto baseField = RFieldBase::Create(std::string(":_" + std::to_string(iBase)), cl->GetName()).Unwrap();
1012+
fSplitFields->emplace_back(std::move(baseField));
1013+
fSplitOffsets->emplace_back(base->GetDelta());
1014+
iBase++;
1015+
}
1016+
9481017
for (auto dataMember : ROOT::Detail::TRangeStaticCast<TDataMember>(*fSoAClass->GetListOfDataMembers())) {
9491018
if ((dataMember->Property() & kIsStatic) || !dataMember->IsPersistent())
9501019
continue;

tree/ntuple/test/SoAField.hxx

Lines changed: 38 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -23,26 +23,6 @@ struct SoAVersionMismatch {
2323
ClassDefNV(SoAVersionMismatch, 4);
2424
};
2525

26-
struct RecordBase {
27-
ClassDefNV(RecordBase, 2);
28-
};
29-
30-
struct RecordDerived : public RecordBase {
31-
ClassDefNV(RecordDerived, 2);
32-
};
33-
34-
struct SoABase {
35-
ClassDefNV(SoABase, 2);
36-
};
37-
38-
struct SoAOnDerivedRecord {
39-
ClassDefNV(SoAOnDerivedRecord, 2);
40-
};
41-
42-
struct SoADerivedOnBaseRecord : public SoABase {
43-
ClassDefNV(SoADerivedOnBaseRecord, 2);
44-
};
45-
4626
struct RecordSimple {
4727
float fX;
4828
float fY;
@@ -152,4 +132,42 @@ struct SoADotBadNestedType {
152132
ClassDefNV(SoADotBadNestedType, 2);
153133
};
154134

135+
struct RecordBase {
136+
float fBase;
137+
ClassDefNV(RecordBase, 2);
138+
};
139+
140+
struct SoABase {
141+
ROOT::RVec<float> fBase;
142+
ClassDefNV(SoABase, 2);
143+
};
144+
145+
struct RecordDerived : public RecordBase {
146+
float fDerived;
147+
ClassDefNV(RecordDerived, 2);
148+
};
149+
150+
struct SoADerived : public SoABase {
151+
ROOT::RVec<float> fDerived;
152+
ClassDefNV(SoADerived, 2);
153+
};
154+
155+
struct RecordDerivedMulti : public RecordDerived, RecordDot {
156+
float fMulti;
157+
ClassDefNV(RecordDerivedMulti, 2);
158+
};
159+
160+
struct SoADerivedMulti : public SoADerived, SoADot {
161+
ROOT::RVec<float> fMulti;
162+
ClassDefNV(SoADerivedMulti, 2);
163+
};
164+
165+
struct SoADerivedFail1 : public RecordDerived {
166+
ClassDefNV(SoADerivedFail1, 2);
167+
};
168+
169+
struct SoADerivedFail2 : public SoABase {
170+
ClassDefNV(SoADerivedFail2, 2);
171+
};
172+
155173
#endif // ROOT_RNTuple_Test_SoAField

tree/ntuple/test/SoAFieldLinkDef.h

Lines changed: 9 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -6,12 +6,6 @@
66
#pragma link C++ options=rntupleSoARecord(xyz) class SoAUnknownRecord+;
77
#pragma link C++ options=rntupleSoARecord(Record) class SoAVersionMismatch+;
88

9-
#pragma link C++ class RecordBase+;
10-
#pragma link C++ class RecordDerived+;
11-
#pragma link C++ options=rntupleSoARecord(RecordBase) class SoABase+;
12-
#pragma link C++ options=rntupleSoARecord(RecordDerived) class SoAOnDerivedRecord+;
13-
#pragma link C++ options=rntupleSoARecord(RecordBase) class SoADerivedOnBaseRecord+;
14-
159
#pragma link C++ class RecordSimple+;
1610
#pragma link C++ options=rntupleSoARecord(RecordSimple) class SoASimple+;
1711
#pragma link C++ options=rntupleSoARecord(RecordSimple) class SoASimpleSwapped+;
@@ -31,4 +25,13 @@
3125
#pragma link C++ options=rntupleSoARecord(RecordDot) class SoADot+;
3226
#pragma link C++ options=rntupleSoARecord(RecordDot) class SoADotBadNestedType+;
3327

28+
#pragma link C++ class RecordBase+;
29+
#pragma link C++ class RecordDerived+;
30+
#pragma link C++ class RecordDerivedMulti+;
31+
#pragma link C++ options=rntupleSoARecord(RecordBase) class SoABase+;
32+
#pragma link C++ options=rntupleSoARecord(RecordDerived) class SoADerived+;
33+
#pragma link C++ options=rntupleSoARecord(RecordDerivedMulti) class SoADerivedMulti+;
34+
#pragma link C++ options=rntupleSoARecord(RecordDerived) class SoADerivedFail1+;
35+
#pragma link C++ options=rntupleSoARecord(RecordDerived) class SoADerivedFail2+;
36+
3437
#endif // __CLING__

tree/ntuple/test/ntuple_soa.cxx

Lines changed: 78 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -70,22 +70,6 @@ TEST(RNTuple, SoACheck)
7070
EXPECT_THAT(e.what(), ::testing::HasSubstr("version mismatch between SoA type and underlying record type"));
7171
}
7272

73-
try {
74-
auto f = std::make_unique<RSoAField>("f", "SoAOnDerivedRecord");
75-
FAIL() << "creating SoA field on derived record should fail";
76-
} catch (const ROOT::RException &e) {
77-
EXPECT_THAT(e.what(), ::testing::HasSubstr("SoA fields with inheritance are currently unsupported"));
78-
}
79-
try {
80-
auto f = std::make_unique<RSoAField>("f", "SoADerivedOnBaseRecord");
81-
FAIL() << "creating a derived SoA field should fail";
82-
} catch (const ROOT::RException &e) {
83-
EXPECT_THAT(e.what(), ::testing::HasSubstr("SoA fields with inheritance are currently unsupported"));
84-
}
85-
{
86-
EXPECT_NO_THROW(auto f = std::make_unique<RSoAField>("f", "SoABase"));
87-
}
88-
8973
try {
9074
auto f = std::make_unique<RSoAField>("f", "SoASimpleBadArray");
9175
FAIL() << "creating SoA field with arrays fail";
@@ -514,3 +498,81 @@ R"({
514498
// clang-format on
515499
EXPECT_EQ(expected, os.str());
516500
}
501+
502+
TEST(RNTuple, SoADerived)
503+
{
504+
ROOT::TestSupport::FileRaii fileGuard("test_rntuple_soa_derived.root");
505+
506+
{
507+
auto model = ROOT::RNTupleModel::Create();
508+
509+
try {
510+
model->AddField(std::make_unique<RSoAField>("x", "SoADerivedFail1"));
511+
FAIL() << "SoADerivedFail1 should fail";
512+
} catch (const ROOT::RException &e) {
513+
EXPECT_THAT(e.what(),
514+
testing::HasSubstr("inheritance of SoA class x does not match its underlying record type"));
515+
}
516+
try {
517+
model->AddField(std::make_unique<RSoAField>("x", "SoADerivedFail2"));
518+
FAIL() << "SoADerivedFail2 should fail";
519+
} catch (const ROOT::RException &e) {
520+
EXPECT_THAT(e.what(), testing::HasSubstr("missing SoA members"));
521+
}
522+
523+
model->AddField(std::make_unique<RSoAField>("derived", "SoADerived"));
524+
model->AddField(std::make_unique<RSoAField>("multi", "SoADerivedMulti"));
525+
auto writer = ROOT::RNTupleWriter::Recreate(std::move(model), "ntpl", fileGuard.GetPath());
526+
527+
auto derivedSoA = writer->GetModel().GetDefaultEntry().GetPtr<SoADerived>("derived");
528+
derivedSoA->fBase = {1.0, 2.0};
529+
derivedSoA->fDerived = {3.0, 4.0};
530+
531+
auto derivedMultiSoA = writer->GetModel().GetDefaultEntry().GetPtr<SoADerivedMulti>("multi");
532+
derivedMultiSoA->fBase = {5.0, 6.0};
533+
derivedMultiSoA->fDerived = {7.0, 8.0};
534+
derivedMultiSoA->fX = {9.0, 10.0};
535+
derivedMultiSoA->fY = {11.0, 12.0};
536+
derivedMultiSoA->fProperties.fColor = {13, 14};
537+
derivedMultiSoA->fProperties.fSize = {15.0, 16.0};
538+
derivedMultiSoA->fMulti = {17.0, 18.0};
539+
540+
writer->Fill();
541+
}
542+
543+
auto reader = ROOT::RNTupleReader::Open("ntpl", fileGuard.GetPath());
544+
545+
std::ostringstream os;
546+
reader->Show(0, os);
547+
548+
// clang-format off
549+
std::string expected{
550+
R"({
551+
"derived": {
552+
":_0": {
553+
"fBase": [1, 2]
554+
},
555+
"fDerived": [3, 4]
556+
},
557+
"multi": {
558+
":_0": {
559+
":_0": {
560+
"fBase": [5, 6]
561+
},
562+
"fDerived": [7, 8]
563+
},
564+
":_1": {
565+
"fX": [9, 10],
566+
"fY": [11, 12],
567+
"fProperties": {
568+
"fColor": [13, 14],
569+
"fSize": [15, 16]
570+
}
571+
},
572+
"fMulti": [17, 18]
573+
}
574+
}
575+
)" };
576+
// clang-format on
577+
EXPECT_EQ(expected, os.str());
578+
}

0 commit comments

Comments
 (0)