Skip to content
Draft
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
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,15 @@ struct StackParam : public o2::conf::ConfigurableParamHelper<StackParam> {
std::string transportPrimaryFileName = "";
std::string transportPrimaryFuncName = "";
bool transportPrimaryInvert = false;
// Used when transportPrimary="onnx". The model is fetched as raw ONNX bytes
// and class 1 means "this track and all descendants produce no hits".
// Despite the legacy parameter name, ONNX runs at PreTrack for primaries
// AND secondaries. Output is a single [batch, scores] float tensor; select
// the class-1 score below. Disable ApplySigmoid for probability outputs.
std::string transportPrimaryOnnxCCDBPath = "";
float transportPrimaryOnnxThreshold = 0.5f;
int transportPrimaryOnnxOutputIndex = 0;
bool transportPrimaryOnnxApplySigmoid = true;

// boilerplate stuff + make principal key "Stack"
O2ParamDef(StackParam, "Stack");
Expand Down
1 change: 1 addition & 0 deletions Detectors/Base/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,7 @@ o2_add_library(DetectorsBase
O2::SimulationDataFormat
O2::SimConfig
O2::CCDB
O2::ML
O2::GPUDataTypes
MC::VMC
TBB::tbb
Expand Down
6 changes: 6 additions & 0 deletions Detectors/Base/include/DetectorsBase/Stack.h
Original file line number Diff line number Diff line change
Expand Up @@ -177,6 +177,10 @@ class Stack : public FairGenericStack

std::vector<MCTrack> const* const getMCTracks() const { return mTracks; }

/// Classify a birth track at PreTrack; false means stop transport before any hit.
bool transportTrack(const TParticle& particle, double eventX, double eventY, double eventZ);
bool hasTrackTransportModel() const { return static_cast<bool>(mTransportTrack); }

/// Clone for worker (used in MT mode only)
FairGenericStack* CloneStack() const override;

Expand Down Expand Up @@ -301,6 +305,8 @@ class Stack : public FairGenericStack

TransportFcn mTransportPrimary = [](const TParticle& p, const std::vector<TParticle>& particles) { return false; }; //! a function to inhibit the tracking of a particle

std::function<bool(const TParticle&, double, double, double)> mTransportTrack; //! ONNX decision for primaries and secondaries

// storage for track references
std::vector<o2::TrackReference>* mTrackRefs = nullptr; //!

Expand Down
87 changes: 87 additions & 0 deletions Detectors/Base/include/DetectorsBase/TrackTransportUtils.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,87 @@
// Copyright 2019-2020 CERN and copyright holders of ALICE O2.
// See https://alice-o2.web.cern.ch/copyright for details of the copyright holders.
// All rights not expressly granted are reserved.
//
// This software is distributed under the terms of the GNU General Public
// License v3 (GPL Version 3), copied verbatim in the file "COPYING".
//
// In applying this license CERN does not waive the privileges and immunities
// granted to it by virtue of its status as an Intergovernmental Organization
// or submit itself to any jurisdiction.

/// \file Stack.h
/// \brief Definition of the Stack class
/// \author M. Al-Turany - June 2014

#ifndef O2_TRACK_TRANSPORT_UTILS_H
#define O2_TRACK_TRANSPORT_UTILS_H

#include "SimulationDataFormat/O2DatabasePDG.h"
#include <TParticle.h>
#include <TParticlePDG.h>
#include <algorithm>
#include <cmath>
#include <limits>
#include <stdexcept>
#include <vector>

namespace o2::data::detail
{
// Same feature order, units, PDG masses and missing values as
// extract_o2_kine_training.C. Normalisation/imputation belongs in the graph.
inline std::vector<float> makeTrackTransportFeatures(const TParticle& particle, double eventX, double eventY, double eventZ)
{
const double px = particle.Px();
const double py = particle.Py();
const double pz = particle.Pz();
const double momentum = std::sqrt(px * px + py * py + pz * pz);
const double pt = std::hypot(px, py);
bool massKnown = false;
double mass = o2::O2DatabasePDG::Mass(particle.GetPdgCode(), massKnown);
const auto* pdgInfo = particle.GetPDG();
if (!massKnown) mass = pdgInfo ? pdgInfo->Mass() : 0.;
const double missing = std::numeric_limits<double>::quiet_NaN();
const double energy = std::sqrt(std::max(0., mass * mass + momentum * momentum));
const double eta = momentum > std::abs(pz) ? 0.5 * std::log((momentum + pz) / (momentum - pz)) : missing;
const double theta = momentum > 0. ? std::acos(pz / momentum) : missing;
const double rapidity = energy > std::abs(pz) ? 0.5 * std::log((energy + pz) / (energy - pz)) : missing;
double chargeSign = pdgInfo == nullptr || pdgInfo->Charge() == 0. ? 0. : std::copysign(1., pdgInfo->Charge());
const int code = particle.GetPdgCode();
if (std::abs(code) >= 1000000000) {
chargeSign = (std::abs(code) / 10000) % 1000 == 0 ? 0. : (code > 0 ? 1. : -1.);
}
const double dx = particle.Vx() - eventX;
const double dy = particle.Vy() - eventY;
const double dz = particle.Vz() - eventZ;
const double pdg = particle.GetPdgCode();

return {static_cast<float>(pdg), static_cast<float>(std::abs(pdg)), static_cast<float>(chargeSign),
static_cast<float>(mass), static_cast<float>(energy), static_cast<float>(energy - mass),
static_cast<float>(px), static_cast<float>(py), static_cast<float>(pz), static_cast<float>(momentum),
static_cast<float>(pt), static_cast<float>(eta), static_cast<float>(std::atan2(py, px)), static_cast<float>(theta),
static_cast<float>(rapidity), static_cast<float>(particle.Vx()), static_cast<float>(particle.Vy()),
static_cast<float>(particle.Vz()), static_cast<float>(particle.T() * 1.e9), static_cast<float>(dx),
static_cast<float>(dy), static_cast<float>(dz), static_cast<float>(std::hypot(particle.Vx(), particle.Vy())),
static_cast<float>(std::hypot(dx, dy)), static_cast<float>(std::sqrt(dx * dx + dy * dy + dz * dz))};
}


inline bool transportFromOnnxScore(float score, float threshold, bool applySigmoid)
{
if (!std::isfinite(threshold) || threshold < 0.f || threshold > 1.f) {
throw std::runtime_error("ONNX pruning threshold must be finite and in [0,1]");
}
if (!std::isfinite(score)) {
throw std::runtime_error("ONNX pruning returned a non-finite score; refusing to reject a track");
}
if (applySigmoid) {
score = score >= 0.f ? 1.f / (1.f + std::exp(-score)) : std::exp(score) / (1.f + std::exp(score));
}
if (score < 0.f || score > 1.f) {
throw std::runtime_error("ONNX pruning probability is outside [0,1]");
}
// Class 1: this track and all descendants produce zero detector hits.
return score < threshold;
}
} // namespace o2::data::detail
#endif
149 changes: 145 additions & 4 deletions Detectors/Base/src/Stack.cxx
Original file line number Diff line number Diff line change
Expand Up @@ -25,23 +25,117 @@
#include "SimulationDataFormat/BaseHits.h"
#include "SimulationDataFormat/StackParam.h"
#include "CommonUtils/ConfigurationMacroHelper.h"
#include "CCDB/BasicCCDBManager.h"
#include "ML/OrtInterface.h"
#include "DetectorsBase/TrackTransportUtils.h"

#include "TLorentzVector.h" // for TLorentzVector
#include "TParticle.h" // for TParticle
#include "TRefArray.h" // for TRefArray
#include "TVirtualMC.h" // for VMC
#include "TMCProcess.h" // for VMC Particle Production Process
#include "TParticlePDG.h"

#include <algorithm>
#include <cassert>
#include <cstddef> // for NULL
#include <cmath>
#include <map>
#include <memory>
#include <limits>
#include <mutex>
#include <stdexcept>
#include <unordered_map>

using std::cout;
using std::endl;
using std::pair;
using namespace o2::data;

namespace
{
// Feature contract used by the sim-pruning models, in order:
// pdg, abs_pdg, charge_sign, mass, energy, ekin, px, py, pz, p, pt, eta,
// phi, theta, rapidity, vx, vy, vz, t_ns, dx/dy/dz_from_event, r_xy,
// r_from_event_xy, r3_from_event. Input normalisation can be embedded in the
// ONNX graph, keeping this code independent of model topology.
constexpr size_t OnnxFeatureCount = 25;

class OnnxTrackTransport
{
public:
explicit OnnxTrackTransport(const o2::sim::StackParam& param)
: mThreshold(param.transportPrimaryOnnxThreshold),
mOutputIndex(param.transportPrimaryOnnxOutputIndex),
mApplySigmoid(param.transportPrimaryOnnxApplySigmoid)
{
if (!std::isfinite(mThreshold) || mThreshold < 0.f || mThreshold > 1.f) {
throw std::runtime_error("ONNX pruning threshold must be finite and in [0,1]");
}
if (param.transportPrimaryOnnxCCDBPath.empty()) {
throw std::runtime_error("Stack.transportPrimaryOnnxCCDBPath must be configured");
}

auto& ccdbManager = o2::ccdb::BasicCCDBManager::instance();
auto& ccdb = ccdbManager.getCCDBAccessor();
std::map<std::string, std::string> headers;
const auto createdNotAfter = ccdbManager.getCreatedNotAfter();
const auto createdNotBefore = ccdbManager.getCreatedNotBefore();
ccdb.loadFileToMemory(mModelBytes, param.transportPrimaryOnnxCCDBPath, {},
ccdbManager.getTimestamp(), &headers, {},
createdNotAfter ? std::to_string(createdNotAfter) : "",
createdNotBefore ? std::to_string(createdNotBefore) : "");
if (mModelBytes.empty()) {
throw std::runtime_error("failed to retrieve ONNX model from CCDB path " + param.transportPrimaryOnnxCCDBPath);
}

std::unordered_map<std::string, std::string> options{{"model-path", param.transportPrimaryOnnxCCDBPath},
{"device-type", "CPU"},
{"intra-op-num-threads", "1"},
{"inter-op-num-threads", "1"},
{"enable-optimizations", "99"},
{"logging-level", "2"},
{"onnx-environment-name", "primary-transport-pruning"}};
mModel.init(options);
mModel.initSessionFromBuffer(mModelBytes.data(), mModelBytes.size());

const auto inputShapes = mModel.getNumInputNodes();
if (inputShapes.size() != 1 || inputShapes[0].size() != 2 ||
(inputShapes[0][0] != 1 && inputShapes[0][0] != -1) ||
inputShapes[0][1] != OnnxFeatureCount) {
throw std::runtime_error("primary transport ONNX model must have one float input with 25 features");
}
const auto outputShapes = mModel.getNumOutputNodes();
if (outputShapes.size() != 1 || outputShapes[0].size() != 2 ||
(outputShapes[0][0] != 1 && outputShapes[0][0] != -1) ||
outputShapes[0][1] <= 0 || mOutputIndex < 0 || mOutputIndex >= outputShapes[0][1]) {
throw std::runtime_error("track transport ONNX model must have one [batch, scores] output and a valid score index");
}
}

bool transport(const TParticle& particle, double eventX, double eventY, double eventZ)
{
// OrtModel mutates its shape buffers during inference. Stack clones share
// this classifier, so protect the session and those buffers together.
std::lock_guard<std::mutex> lock(mMutex);
std::vector<std::vector<float>> inputs{o2::data::detail::makeTrackTransportFeatures(particle, eventX, eventY, eventZ)};
auto output = mModel.inference<float, float>(inputs);
if (static_cast<size_t>(mOutputIndex) >= output.size()) {
throw std::runtime_error("Stack.transportPrimaryOnnxOutputIndex is outside the model output");
}
return o2::data::detail::transportFromOnnxScore(output[mOutputIndex], mThreshold, mApplySigmoid);
}

private:
std::vector<char> mModelBytes; // Must outlive the session (members are destroyed in reverse order).
o2::ml::OrtModel mModel;
std::mutex mMutex;
float mThreshold;
int mOutputIndex;
bool mApplySigmoid;
};
} // namespace

// small helper function to append to vector at arbitrary position
template <typename T, typename I>
void insertInVector(std::vector<T>& v, I index, T e)
Expand Down Expand Up @@ -100,16 +194,31 @@
transportPrimary = o2::conf::GetFromMacro<o2::data::Stack::TransportFcn>(param.transportPrimaryFileName,
param.transportPrimaryFuncName,
"o2::data::Stack::TransportFcn", "stack_transport_primary");
if (!mTransportPrimary) {
if (!transportPrimary) {
LOG(fatal) << "Failed to retrieve external \'transportPrimary\' function: problem with configuration ";
}
LOG(info) << "Successfully retrieve external \'transportPrimary\' frunction: " << param.transportPrimaryFileName;
} else if (param.transportPrimary.compare("onnx") == 0) {
try {
auto classifier = std::make_shared<OnnxTrackTransport>(param);
mTransportTrack = [classifier, invert = param.transportPrimaryInvert](const TParticle& p, double x, double y, double z) {
const bool transport = classifier->transport(p, x, y, z);
return invert ? !transport : transport;
};
// ONNX runs at PreTrack, where the true event vertex and both primary
// and secondary birth states are available, also in parallel simulation.
transportPrimary = [](const TParticle&, const std::vector<TParticle>&) { return true; };
LOG(info) << "Successfully configured ONNX track transport pruning from CCDB path "
<< param.transportPrimaryOnnxCCDBPath;
} catch (const std::exception& error) {
LOG(fatal) << "Failed to configure ONNX track transport pruning: " << error.what();
}
} else {
LOG(fatal) << "unsupported \'trasportPrimary\' mode: " << param.transportPrimary;
}

if (param.transportPrimaryInvert) {
mTransportPrimary = [transportPrimary](const TParticle& p, const std::vector<TParticle>& particles) { return !transportPrimary; };
if (param.transportPrimaryInvert && param.transportPrimary != "onnx") {
mTransportPrimary = [transportPrimary](const TParticle& p, const std::vector<TParticle>& particles) { return !transportPrimary(p, particles); };
} else {
mTransportPrimary = transportPrimary;
}
Expand All @@ -133,7 +242,9 @@
mMinHits(rhs.mMinHits),
mEnergyCut(rhs.mEnergyCut),
mTrackRefs(new std::vector<o2::TrackReference>),
mIsG4Like(rhs.mIsG4Like)
mIsG4Like(rhs.mIsG4Like),
mTransportPrimary(rhs.mTransportPrimary),
mTransportTrack(rhs.mTransportTrack)
{
LOG(debug) << "copy constructor called";
mTracks = new std::vector<MCTrack>();
Expand Down Expand Up @@ -169,6 +280,8 @@
mMinHits = rhs.mMinHits;
mEnergyCut = rhs.mEnergyCut;
mIsG4Like = rhs.mIsG4Like;
mTransportPrimary = rhs.mTransportPrimary;
mTransportTrack = rhs.mTransportTrack;

return *this;
}
Expand All @@ -193,8 +306,8 @@
TMCProcess proc2)
{
// printf("Pushing %s toBeDone %5d parentId %5d pdgCode %5d is %5d entries %5d \n",
// proc == kPPrimary ? "Primary: " : "Secondary: ",

Check failure on line 309 in Detectors/Base/src/Stack.cxx

View workflow job for this annotation

GitHub Actions / PR formatting / whitespace

Tab characters found

Indent code using spaces instead of tabs.
// toBeDone, parentId, pdgCode, is, mNumberOfEntriesInParticles);

Check failure on line 310 in Detectors/Base/src/Stack.cxx

View workflow job for this annotation

GitHub Actions / PR formatting / whitespace

Tab characters found

Indent code using spaces instead of tabs.

//
// This method is called
Expand Down Expand Up @@ -269,6 +382,34 @@
}
}

bool Stack::transportTrack(const TParticle& particle, double eventX, double eventY, double eventZ)
{
if (!mTransportTrack || mTransportTrack(particle, eventX, eventY, eventZ)) {
return true;
}
// Keep bookkeeping and ancestry, but mark the birth track as inhibited.
// Do not alter the primary completion count: its PreTrack/FinishPrimary
// lifecycle has already started and will still be completed by the engine.
auto inhibit = [](TParticle& p) {
p.SetBit(ParticleStatus::kToBeDone, 0);
p.SetBit(ParticleStatus::kInhibited, 1);
};
inhibit(mCurrentParticle);
const int id = mIndexOfCurrentTrack;
if (id >= 0 && id < static_cast<int>(mPrimaryParticles.size())) {
inhibit(mPrimaryParticles[id]);
(*mTracks)[id].setToBeDone(false);
(*mTracks)[id].setInhibited(true);
} else if (id >= 0 && id < static_cast<int>(mTrackIDtoParticlesEntry.size())) {
const int entry = mTrackIDtoParticlesEntry[id];
if (entry >= 0 && entry < static_cast<int>(mParticles.size())) {
mParticles[entry].setToBeDone(false);
mParticles[entry].setInhibited(true);
}
}
return false;
}

void Stack::PushTrack(int toBeDone, TParticle& p)
{
// printf("stack -> Pushing Primary toBeDone %5d %5d parentId %5d pdgCode %5d is %5d entries %5d \n", toBeDone, p.TestBit(ParticleStatus::kToBeDone), p.GetFirstMother(), p.GetPdgCode(), p.GetStatusCode(), mNumberOfEntriesInParticles);
Expand Down
35 changes: 35 additions & 0 deletions Detectors/Base/test/testStack.cxx
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
#include <boost/test/unit_test.hpp>
#include "DetectorsBase/Detector.h"
#include "DetectorsBase/Stack.h"
#include "DetectorsBase/TrackTransportUtils.h"
#include "SimulationDataFormat/BaseHits.h"
#include "TFile.h"
#include "TMCProcess.h"
Expand Down Expand Up @@ -180,3 +181,37 @@ BOOST_AUTO_TEST_CASE(Offsetting_keeps_an_invalid_index_invalid)
BOOST_CHECK_EQUAL(o2::base::Detector::offsetTrackIndex(7, nprimaries, primaryOffset, secondaryOffset), 107);
BOOST_CHECK_EQUAL(o2::base::Detector::offsetTrackIndex(-1, nprimaries, primaryOffset, secondaryOffset), -1);
}

BOOST_AUTO_TEST_CASE(Track_transport_features_match_training_units)
{
// A secondary displaced from the actual event vertex, with negative phi.
TParticle p(211, 0, 0, -1, -1, -1, 0., -2., 0., 2.1, 11., 22., 33., 7.e-9);
auto f = o2::data::detail::makeTrackTransportFeatures(p, 1., 2., 3.);
BOOST_REQUIRE_EQUAL(f.size(), 25);
BOOST_CHECK_CLOSE(f[18], 7.f, 1.e-4f); // nanoseconds, not seconds
BOOST_CHECK_CLOSE(f[12], -std::acos(-1.f) / 2.f, 1.e-4f); // atan2 range
BOOST_CHECK_EQUAL(f[19], 10.f);
BOOST_CHECK_EQUAL(f[20], 20.f);
BOOST_CHECK_EQUAL(f[21], 30.f);
BOOST_CHECK_CLOSE(f[23], std::sqrt(500.f), 1.e-4f);
BOOST_CHECK_EQUAL(f[2], 1.f);
// Undefined angular inputs retain CSV missing-value semantics.
p.SetMomentum(0., 0., 0., 0.);
f = o2::data::detail::makeTrackTransportFeatures(p, 1., 2., 3.);
BOOST_CHECK(std::isnan(f[11]));
BOOST_CHECK(std::isnan(f[13]));
}

BOOST_AUTO_TEST_CASE(Track_transport_class_one_rejects_and_invalid_scores_fail)
{
using o2::data::detail::transportFromOnnxScore;
BOOST_CHECK(transportFromOnnxScore(0.1f, 0.5f, false));
BOOST_CHECK(!transportFromOnnxScore(0.9f, 0.5f, false));
BOOST_CHECK(!transportFromOnnxScore(0.f, 0.5f, true));
BOOST_CHECK(transportFromOnnxScore(-1000.f, 0.5f, true));
BOOST_CHECK(!transportFromOnnxScore(1000.f, 0.5f, true));
BOOST_CHECK_THROW(transportFromOnnxScore(std::numeric_limits<float>::quiet_NaN(), 0.5f, false), std::runtime_error);
BOOST_CHECK_THROW(transportFromOnnxScore(std::numeric_limits<float>::infinity(), 0.5f, true), std::runtime_error);
BOOST_CHECK_THROW(transportFromOnnxScore(2.f, 0.5f, false), std::runtime_error);
BOOST_CHECK_THROW(transportFromOnnxScore(0.5f, -1.f, false), std::runtime_error);
}
Loading
Loading