diff --git a/Analysis/Tutorials/src/histograms.cxx b/Analysis/Tutorials/src/histograms.cxx index 15679ff94c30a..17a4efa2e9151 100644 --- a/Analysis/Tutorials/src/histograms.cxx +++ b/Analysis/Tutorials/src/histograms.cxx @@ -70,9 +70,9 @@ struct CTask { void process(aod::Tracks const& tracks) { for (auto& track : tracks) { - if (track.pt() < pTCut) + if (track.pt2() < pTCut * pTCut) continue; - ptH->Fill(track.pt()); + ptH->Fill(std::sqrt(track.pt2())); trZ->Fill(track.z()); } } @@ -84,16 +84,16 @@ struct DTask { void init(InitContext const&) { list.setObject(new TList); - list->Add(new TH1F("ptHist", "", 100, 0, 10)); + list->Add(new TH1F("pHist", "", 100, 0, 10)); list->Add(new TH1F("etaHist", "", 102, -2.01, 2.01)); } void process(aod::Track const& track) { - auto ptHist = dynamic_cast(list->At(0)); + auto pHist = dynamic_cast(list->At(0)); auto etaHist = dynamic_cast(list->At(1)); - ptHist->Fill(track.pt()); + pHist->Fill(std::sqrt(track.p2())); etaHist->Fill(track.eta()); } }; diff --git a/Framework/Core/include/Framework/AnalysisDataModel.h b/Framework/Core/include/Framework/AnalysisDataModel.h index edecf64b4ec13..6c1e38aff0afc 100644 --- a/Framework/Core/include/Framework/AnalysisDataModel.h +++ b/Framework/Core/include/Framework/AnalysisDataModel.h @@ -116,10 +116,8 @@ DECLARE_SOA_DYNAMIC_COLUMN(Pz, pz, [](float signed1Pt, float tgl) -> float { DECLARE_SOA_DYNAMIC_COLUMN(P, p, [](float signed1Pt, float tgl) -> float { return std::sqrt(1.f + tgl * tgl) / std::abs(signed1Pt); }); -DECLARE_SOA_DYNAMIC_COLUMN(P2, p2, [](float signed1Pt, float tgl) -> float { - return (1.f + tgl * tgl) / (signed1Pt * signed1Pt); -}); +DECLARE_SOA_EXPRESSION_COLUMN(P2, p2, float, (1.f + aod::track::tgl * aod::track::tgl) / (aod::track::signed1Pt * aod::track::signed1Pt)); DECLARE_SOA_EXPRESSION_COLUMN(Pt2, pt2, float, (1.f / aod::track::signed1Pt) * (1.f / aod::track::signed1Pt)); // TRACKPARCOV TABLE definition @@ -194,10 +192,10 @@ DECLARE_SOA_TABLE_FULL(StoredTracks, "Tracks", "AOD", "TRACKPAR", track::Py, track::Pz, track::P, - track::P2, track::Charge); -DECLARE_SOA_EXTENDED_TABLE(Tracks, StoredTracks, "TRACKPAR", aod::track::Pt2); +DECLARE_SOA_EXTENDED_TABLE(Tracks, StoredTracks, "TRACKPAR", aod::track::Pt2, + aod::track::P2); DECLARE_SOA_TABLE(TracksCov, "AOD", "TRACKPARCOV", track::CYY, track::CZY, track::CZZ, track::CSnpY, diff --git a/Framework/Core/include/Framework/Expressions.h b/Framework/Core/include/Framework/Expressions.h index 150ab8d8c2aa5..fe8681c283da5 100644 --- a/Framework/Core/include/Framework/Expressions.h +++ b/Framework/Core/include/Framework/Expressions.h @@ -77,6 +77,7 @@ constexpr auto selectArrowType() } std::shared_ptr concreteArrowType(atype::type type); +std::string upcastTo(atype::type f); /// An expression tree node corresponding to a literal value struct LiteralNode { @@ -294,11 +295,23 @@ inline Node operator+(Node left, T right) return Node{OpNode{BasicOp::Addition}, std::move(left), LiteralNode{right}}; } +template +inline Node operator+(T left, Node right) +{ + return Node{OpNode{BasicOp::Addition}, LiteralNode{left}, std::move(right)}; +} + template inline Node operator-(Node left, T right) { return Node{OpNode{BasicOp::Subtraction}, std::move(left), LiteralNode{right}}; } + +template +inline Node operator-(T left, Node right) +{ + return Node{OpNode{BasicOp::Subtraction}, LiteralNode{left}, std::move(right)}; +} /// semi-binary template inline Node npow(Node left, T right) diff --git a/Framework/Core/src/ExpressionHelpers.h b/Framework/Core/src/ExpressionHelpers.h index df2aa56a19ed4..44bf582808323 100644 --- a/Framework/Core/src/ExpressionHelpers.h +++ b/Framework/Core/src/ExpressionHelpers.h @@ -41,12 +41,12 @@ static std::array binaryOperationsMap = { struct DatumSpec { /// datum spec either contains an index, a value of a literal or a binding label using datum_t = std::variant; - datum_t datum; + datum_t datum = std::monostate{}; atype::type type = atype::NA; explicit DatumSpec(size_t index, atype::type type_) : datum{index}, type{type_} {} explicit DatumSpec(LiteralNode::var_t literal, atype::type type_) : datum{literal}, type{type_} {} explicit DatumSpec(std::string binding, atype::type type_) : datum{binding}, type{type_} {} - DatumSpec() : datum{std::monostate{}} {} + DatumSpec() = default; DatumSpec(DatumSpec const&) = default; DatumSpec(DatumSpec&&) = default; DatumSpec& operator=(DatumSpec const&) = default; @@ -93,15 +93,6 @@ struct ColumnOperationSpec { template std::shared_ptr createProjectors(framework::pack, gandiva::SchemaPtr schema) { - auto findField = [&schema](const char* label) -> gandiva::FieldPtr const { - for (auto i = 0; i < schema->num_fields(); ++i) { - if (schema->field(i)->name() == label) { - return gandiva::FieldPtr{schema->field(i)}; - } - } - throw std::runtime_error(fmt::format("Cannot find field \"{}\"", label)); - }; - std::shared_ptr projector; auto s = gandiva::Projector::Make( schema, diff --git a/Framework/Core/src/Expressions.cxx b/Framework/Core/src/Expressions.cxx index f401333a2f2bf..5204caab15a52 100644 --- a/Framework/Core/src/Expressions.cxx +++ b/Framework/Core/src/Expressions.cxx @@ -70,6 +70,22 @@ std::shared_ptr concreteArrowType(atype::type type) } } +std::string upcastTo(atype::type f) +{ + switch (f) { + case atype::INT32: + return "castINT"; + case atype::INT64: + return "castBIGINT"; + case atype::FLOAT: + return "castFLOAT4"; + case atype::DOUBLE: + return "castFLOAT8"; + default: + throw std::runtime_error(fmt::format("Do not know how to cast to {}", f)); + } +} + bool operator==(DatumSpec const& lhs, DatumSpec const& rhs) { return (lhs.datum == rhs.datum) && (lhs.type == rhs.type); @@ -385,7 +401,11 @@ gandiva::NodePtr createExpressionTree(Operations const& opSpecs, if (lookup != fieldNodes.end()) { return lookup->second; } - auto node = gandiva::TreeExprBuilder::MakeField(Schema->GetFieldByName(name)); + auto field = Schema->GetFieldByName(name); + if (field == nullptr) { + throw std::runtime_error(fmt::format("Cannot find field \"{}\"", name)); + } + auto node = gandiva::TreeExprBuilder::MakeField(field); fieldNodes.insert({name, node}); return node; } @@ -396,6 +416,15 @@ gandiva::NodePtr createExpressionTree(Operations const& opSpecs, for (auto it = opSpecs.rbegin(); it != opSpecs.rend(); ++it) { auto leftNode = datumNode(it->left); auto rightNode = datumNode(it->right); + + auto insertUpcastNode = [&](gandiva::NodePtr node, atype::type t) { + if (t != it->type) { + auto upcast = gandiva::TreeExprBuilder::MakeFunction(upcastTo(it->type), {node}, concreteArrowType(it->type)); + node = upcast; + } + return node; + }; + switch (it->op) { case BasicOp::LogicalOr: tree = gandiva::TreeExprBuilder::MakeOr({leftNode, rightNode}); @@ -405,8 +434,13 @@ gandiva::NodePtr createExpressionTree(Operations const& opSpecs, break; default: if (it->op < BasicOp::Exp) { + if (it->type != atype::BOOL) { + leftNode = insertUpcastNode(leftNode, it->left.type); + rightNode = insertUpcastNode(rightNode, it->right.type); + } tree = gandiva::TreeExprBuilder::MakeFunction(binaryOperationsMap[it->op], {leftNode, rightNode}, concreteArrowType(it->type)); } else { + leftNode = insertUpcastNode(leftNode, it->left.type); tree = gandiva::TreeExprBuilder::MakeFunction(binaryOperationsMap[it->op], {leftNode}, concreteArrowType(it->type)); } break;