Skip to content

Commit 312f9b5

Browse files
authored
Merge pull request #3881 from aalkin/expressions-update
DPL Analysis: enable arithmetic operations between ints and floats in expressions (filters and expression columns)
2 parents 5d19a39 + bae6bea commit 312f9b5

5 files changed

Lines changed: 58 additions & 22 deletions

File tree

Analysis/Tutorials/src/histograms.cxx

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -70,9 +70,9 @@ struct CTask {
7070
void process(aod::Tracks const& tracks)
7171
{
7272
for (auto& track : tracks) {
73-
if (track.pt() < pTCut)
73+
if (track.pt2() < pTCut * pTCut)
7474
continue;
75-
ptH->Fill(track.pt());
75+
ptH->Fill(std::sqrt(track.pt2()));
7676
trZ->Fill(track.z());
7777
}
7878
}
@@ -84,16 +84,16 @@ struct DTask {
8484
void init(InitContext const&)
8585
{
8686
list.setObject(new TList);
87-
list->Add(new TH1F("ptHist", "", 100, 0, 10));
87+
list->Add(new TH1F("pHist", "", 100, 0, 10));
8888
list->Add(new TH1F("etaHist", "", 102, -2.01, 2.01));
8989
}
9090

9191
void process(aod::Track const& track)
9292
{
93-
auto ptHist = dynamic_cast<TH1F*>(list->At(0));
93+
auto pHist = dynamic_cast<TH1F*>(list->At(0));
9494
auto etaHist = dynamic_cast<TH1F*>(list->At(1));
9595

96-
ptHist->Fill(track.pt());
96+
pHist->Fill(std::sqrt(track.p2()));
9797
etaHist->Fill(track.eta());
9898
}
9999
};

Framework/Core/include/Framework/AnalysisDataModel.h

Lines changed: 3 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -116,10 +116,8 @@ DECLARE_SOA_DYNAMIC_COLUMN(Pz, pz, [](float signed1Pt, float tgl) -> float {
116116
DECLARE_SOA_DYNAMIC_COLUMN(P, p, [](float signed1Pt, float tgl) -> float {
117117
return std::sqrt(1.f + tgl * tgl) / std::abs(signed1Pt);
118118
});
119-
DECLARE_SOA_DYNAMIC_COLUMN(P2, p2, [](float signed1Pt, float tgl) -> float {
120-
return (1.f + tgl * tgl) / (signed1Pt * signed1Pt);
121-
});
122119

120+
DECLARE_SOA_EXPRESSION_COLUMN(P2, p2, float, (1.f + aod::track::tgl * aod::track::tgl) / (aod::track::signed1Pt * aod::track::signed1Pt));
123121
DECLARE_SOA_EXPRESSION_COLUMN(Pt2, pt2, float, (1.f / aod::track::signed1Pt) * (1.f / aod::track::signed1Pt));
124122

125123
// TRACKPARCOV TABLE definition
@@ -210,10 +208,10 @@ DECLARE_SOA_TABLE_FULL(StoredTracks, "Tracks", "AOD", "TRACKPAR",
210208
track::Py<track::Signed1Pt, track::Snp, track::Alpha>,
211209
track::Pz<track::Signed1Pt, track::Tgl>,
212210
track::P<track::Signed1Pt, track::Tgl>,
213-
track::P2<track::Signed1Pt, track::Tgl>,
214211
track::Charge<track::Signed1Pt>);
215212

216-
DECLARE_SOA_EXTENDED_TABLE(Tracks, StoredTracks, "TRACKPAR", aod::track::Pt2);
213+
DECLARE_SOA_EXTENDED_TABLE(Tracks, StoredTracks, "TRACKPAR", aod::track::Pt2,
214+
aod::track::P2);
217215

218216
DECLARE_SOA_TABLE_FULL(StoredTracksCov, "TracksCov", "AOD", "TRACKPARCOV",
219217
track::SigmaY, track::SigmaZ, track::SigmaSnp, track::SigmaTgl, track::Sigma1Pt,

Framework/Core/include/Framework/Expressions.h

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -77,6 +77,7 @@ constexpr auto selectArrowType()
7777
}
7878

7979
std::shared_ptr<arrow::DataType> concreteArrowType(atype::type type);
80+
std::string upcastTo(atype::type f);
8081

8182
/// An expression tree node corresponding to a literal value
8283
struct LiteralNode {
@@ -294,11 +295,23 @@ inline Node operator+(Node left, T right)
294295
return Node{OpNode{BasicOp::Addition}, std::move(left), LiteralNode{right}};
295296
}
296297

298+
template <typename T>
299+
inline Node operator+(T left, Node right)
300+
{
301+
return Node{OpNode{BasicOp::Addition}, LiteralNode{left}, std::move(right)};
302+
}
303+
297304
template <typename T>
298305
inline Node operator-(Node left, T right)
299306
{
300307
return Node{OpNode{BasicOp::Subtraction}, std::move(left), LiteralNode{right}};
301308
}
309+
310+
template <typename T>
311+
inline Node operator-(T left, Node right)
312+
{
313+
return Node{OpNode{BasicOp::Subtraction}, LiteralNode{left}, std::move(right)};
314+
}
302315
/// semi-binary
303316
template <typename T>
304317
inline Node npow(Node left, T right)

Framework/Core/src/ExpressionHelpers.h

Lines changed: 2 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -41,12 +41,12 @@ static std::array<std::string, BasicOp::Abs + 1> binaryOperationsMap = {
4141
struct DatumSpec {
4242
/// datum spec either contains an index, a value of a literal or a binding label
4343
using datum_t = std::variant<std::monostate, size_t, LiteralNode::var_t, std::string>;
44-
datum_t datum;
44+
datum_t datum = std::monostate{};
4545
atype::type type = atype::NA;
4646
explicit DatumSpec(size_t index, atype::type type_) : datum{index}, type{type_} {}
4747
explicit DatumSpec(LiteralNode::var_t literal, atype::type type_) : datum{literal}, type{type_} {}
4848
explicit DatumSpec(std::string binding, atype::type type_) : datum{binding}, type{type_} {}
49-
DatumSpec() : datum{std::monostate{}} {}
49+
DatumSpec() = default;
5050
DatumSpec(DatumSpec const&) = default;
5151
DatumSpec(DatumSpec&&) = default;
5252
DatumSpec& operator=(DatumSpec const&) = default;
@@ -93,15 +93,6 @@ struct ColumnOperationSpec {
9393
template <typename... C>
9494
std::shared_ptr<gandiva::Projector> createProjectors(framework::pack<C...>, gandiva::SchemaPtr schema)
9595
{
96-
auto findField = [&schema](const char* label) -> gandiva::FieldPtr const {
97-
for (auto i = 0; i < schema->num_fields(); ++i) {
98-
if (schema->field(i)->name() == label) {
99-
return gandiva::FieldPtr{schema->field(i)};
100-
}
101-
}
102-
throw std::runtime_error(fmt::format("Cannot find field \"{}\"", label));
103-
};
104-
10596
std::shared_ptr<gandiva::Projector> projector;
10697
auto s = gandiva::Projector::Make(
10798
schema,

Framework/Core/src/Expressions.cxx

Lines changed: 35 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -70,6 +70,22 @@ std::shared_ptr<arrow::DataType> concreteArrowType(atype::type type)
7070
}
7171
}
7272

73+
std::string upcastTo(atype::type f)
74+
{
75+
switch (f) {
76+
case atype::INT32:
77+
return "castINT";
78+
case atype::INT64:
79+
return "castBIGINT";
80+
case atype::FLOAT:
81+
return "castFLOAT4";
82+
case atype::DOUBLE:
83+
return "castFLOAT8";
84+
default:
85+
throw std::runtime_error(fmt::format("Do not know how to cast to {}", f));
86+
}
87+
}
88+
7389
bool operator==(DatumSpec const& lhs, DatumSpec const& rhs)
7490
{
7591
return (lhs.datum == rhs.datum) && (lhs.type == rhs.type);
@@ -385,7 +401,11 @@ gandiva::NodePtr createExpressionTree(Operations const& opSpecs,
385401
if (lookup != fieldNodes.end()) {
386402
return lookup->second;
387403
}
388-
auto node = gandiva::TreeExprBuilder::MakeField(Schema->GetFieldByName(name));
404+
auto field = Schema->GetFieldByName(name);
405+
if (field == nullptr) {
406+
throw std::runtime_error(fmt::format("Cannot find field \"{}\"", name));
407+
}
408+
auto node = gandiva::TreeExprBuilder::MakeField(field);
389409
fieldNodes.insert({name, node});
390410
return node;
391411
}
@@ -396,6 +416,15 @@ gandiva::NodePtr createExpressionTree(Operations const& opSpecs,
396416
for (auto it = opSpecs.rbegin(); it != opSpecs.rend(); ++it) {
397417
auto leftNode = datumNode(it->left);
398418
auto rightNode = datumNode(it->right);
419+
420+
auto insertUpcastNode = [&](gandiva::NodePtr node, atype::type t) {
421+
if (t != it->type) {
422+
auto upcast = gandiva::TreeExprBuilder::MakeFunction(upcastTo(it->type), {node}, concreteArrowType(it->type));
423+
node = upcast;
424+
}
425+
return node;
426+
};
427+
399428
switch (it->op) {
400429
case BasicOp::LogicalOr:
401430
tree = gandiva::TreeExprBuilder::MakeOr({leftNode, rightNode});
@@ -405,8 +434,13 @@ gandiva::NodePtr createExpressionTree(Operations const& opSpecs,
405434
break;
406435
default:
407436
if (it->op < BasicOp::Exp) {
437+
if (it->type != atype::BOOL) {
438+
leftNode = insertUpcastNode(leftNode, it->left.type);
439+
rightNode = insertUpcastNode(rightNode, it->right.type);
440+
}
408441
tree = gandiva::TreeExprBuilder::MakeFunction(binaryOperationsMap[it->op], {leftNode, rightNode}, concreteArrowType(it->type));
409442
} else {
443+
leftNode = insertUpcastNode(leftNode, it->left.type);
410444
tree = gandiva::TreeExprBuilder::MakeFunction(binaryOperationsMap[it->op], {leftNode}, concreteArrowType(it->type));
411445
}
412446
break;

0 commit comments

Comments
 (0)