Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 5 additions & 5 deletions Analysis/Tutorials/src/histograms.cxx
Original file line number Diff line number Diff line change
Expand Up @@ -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());
}
}
Expand All @@ -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<TH1F*>(list->At(0));
auto pHist = dynamic_cast<TH1F*>(list->At(0));
auto etaHist = dynamic_cast<TH1F*>(list->At(1));

ptHist->Fill(track.pt());
pHist->Fill(std::sqrt(track.p2()));
etaHist->Fill(track.eta());
}
};
Expand Down
8 changes: 3 additions & 5 deletions Framework/Core/include/Framework/AnalysisDataModel.h
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -194,10 +192,10 @@ DECLARE_SOA_TABLE_FULL(StoredTracks, "Tracks", "AOD", "TRACKPAR",
track::Py<track::Signed1Pt, track::Snp, track::Alpha>,
track::Pz<track::Signed1Pt, track::Tgl>,
track::P<track::Signed1Pt, track::Tgl>,
track::P2<track::Signed1Pt, track::Tgl>,
track::Charge<track::Signed1Pt>);

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,
Expand Down
13 changes: 13 additions & 0 deletions Framework/Core/include/Framework/Expressions.h
Original file line number Diff line number Diff line change
Expand Up @@ -77,6 +77,7 @@ constexpr auto selectArrowType()
}

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

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

template <typename T>
inline Node operator+(T left, Node right)
{
return Node{OpNode{BasicOp::Addition}, LiteralNode{left}, std::move(right)};
}

template <typename T>
inline Node operator-(Node left, T right)
{
return Node{OpNode{BasicOp::Subtraction}, std::move(left), LiteralNode{right}};
}

template <typename T>
inline Node operator-(T left, Node right)
{
return Node{OpNode{BasicOp::Subtraction}, LiteralNode{left}, std::move(right)};
}
/// semi-binary
template <typename T>
inline Node npow(Node left, T right)
Expand Down
13 changes: 2 additions & 11 deletions Framework/Core/src/ExpressionHelpers.h
Original file line number Diff line number Diff line change
Expand Up @@ -41,12 +41,12 @@ static std::array<std::string, BasicOp::Abs + 1> binaryOperationsMap = {
struct DatumSpec {
/// datum spec either contains an index, a value of a literal or a binding label
using datum_t = std::variant<std::monostate, size_t, LiteralNode::var_t, std::string>;
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;
Expand Down Expand Up @@ -93,15 +93,6 @@ struct ColumnOperationSpec {
template <typename... C>
std::shared_ptr<gandiva::Projector> createProjectors(framework::pack<C...>, 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<gandiva::Projector> projector;
auto s = gandiva::Projector::Make(
schema,
Expand Down
36 changes: 35 additions & 1 deletion Framework/Core/src/Expressions.cxx
Original file line number Diff line number Diff line change
Expand Up @@ -70,6 +70,22 @@ std::shared_ptr<arrow::DataType> 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);
Expand Down Expand Up @@ -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;
}
Expand All @@ -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});
Expand All @@ -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;
Expand Down