diff --git a/DataFormats/simulation/include/SimulationDataFormat/StackParam.h b/DataFormats/simulation/include/SimulationDataFormat/StackParam.h index b76112b41b541..3e5e766829f7c 100644 --- a/DataFormats/simulation/include/SimulationDataFormat/StackParam.h +++ b/DataFormats/simulation/include/SimulationDataFormat/StackParam.h @@ -28,6 +28,15 @@ struct StackParam : public o2::conf::ConfigurableParamHelper { 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"); diff --git a/Detectors/Base/CMakeLists.txt b/Detectors/Base/CMakeLists.txt index 76e4ed9f741fd..effadfb3b5c12 100644 --- a/Detectors/Base/CMakeLists.txt +++ b/Detectors/Base/CMakeLists.txt @@ -51,6 +51,7 @@ o2_add_library(DetectorsBase O2::SimulationDataFormat O2::SimConfig O2::CCDB + O2::ML O2::GPUDataTypes MC::VMC TBB::tbb diff --git a/Detectors/Base/include/DetectorsBase/Stack.h b/Detectors/Base/include/DetectorsBase/Stack.h index 479981a65477a..37e3547317b60 100644 --- a/Detectors/Base/include/DetectorsBase/Stack.h +++ b/Detectors/Base/include/DetectorsBase/Stack.h @@ -177,6 +177,10 @@ class Stack : public FairGenericStack std::vector 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(mTransportTrack); } + /// Clone for worker (used in MT mode only) FairGenericStack* CloneStack() const override; @@ -301,6 +305,8 @@ class Stack : public FairGenericStack TransportFcn mTransportPrimary = [](const TParticle& p, const std::vector& particles) { return false; }; //! a function to inhibit the tracking of a particle + std::function mTransportTrack; //! ONNX decision for primaries and secondaries + // storage for track references std::vector* mTrackRefs = nullptr; //! diff --git a/Detectors/Base/include/DetectorsBase/TrackTransportUtils.h b/Detectors/Base/include/DetectorsBase/TrackTransportUtils.h new file mode 100644 index 0000000000000..41f4bac2d6708 --- /dev/null +++ b/Detectors/Base/include/DetectorsBase/TrackTransportUtils.h @@ -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 +#include +#include +#include +#include +#include +#include + +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 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::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(pdg), static_cast(std::abs(pdg)), static_cast(chargeSign), + static_cast(mass), static_cast(energy), static_cast(energy - mass), + static_cast(px), static_cast(py), static_cast(pz), static_cast(momentum), + static_cast(pt), static_cast(eta), static_cast(std::atan2(py, px)), static_cast(theta), + static_cast(rapidity), static_cast(particle.Vx()), static_cast(particle.Vy()), + static_cast(particle.Vz()), static_cast(particle.T() * 1.e9), static_cast(dx), + static_cast(dy), static_cast(dz), static_cast(std::hypot(particle.Vx(), particle.Vy())), + static_cast(std::hypot(dx, dy)), static_cast(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 diff --git a/Detectors/Base/src/Stack.cxx b/Detectors/Base/src/Stack.cxx index a00c21c0589b9..5729408803fef 100644 --- a/Detectors/Base/src/Stack.cxx +++ b/Detectors/Base/src/Stack.cxx @@ -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 #include #include // for NULL #include +#include +#include +#include +#include +#include +#include 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 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 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 lock(mMutex); + std::vector> inputs{o2::data::detail::makeTrackTransportFeatures(particle, eventX, eventY, eventZ)}; + auto output = mModel.inference(inputs); + if (static_cast(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 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 void insertInVector(std::vector& v, I index, T e) @@ -100,16 +194,31 @@ Stack::Stack(Int_t size) transportPrimary = o2::conf::GetFromMacro(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(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&) { 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& particles) { return !transportPrimary; }; + if (param.transportPrimaryInvert && param.transportPrimary != "onnx") { + mTransportPrimary = [transportPrimary](const TParticle& p, const std::vector& particles) { return !transportPrimary(p, particles); }; } else { mTransportPrimary = transportPrimary; } @@ -133,7 +242,9 @@ Stack::Stack(const Stack& rhs) mMinHits(rhs.mMinHits), mEnergyCut(rhs.mEnergyCut), mTrackRefs(new std::vector), - mIsG4Like(rhs.mIsG4Like) + mIsG4Like(rhs.mIsG4Like), + mTransportPrimary(rhs.mTransportPrimary), + mTransportTrack(rhs.mTransportTrack) { LOG(debug) << "copy constructor called"; mTracks = new std::vector(); @@ -169,6 +280,8 @@ Stack& Stack::operator=(const Stack& rhs) mMinHits = rhs.mMinHits; mEnergyCut = rhs.mEnergyCut; mIsG4Like = rhs.mIsG4Like; + mTransportPrimary = rhs.mTransportPrimary; + mTransportTrack = rhs.mTransportTrack; return *this; } @@ -269,6 +382,34 @@ void Stack::handleTransportPrimary(TParticle& p) } } +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(mPrimaryParticles.size())) { + inhibit(mPrimaryParticles[id]); + (*mTracks)[id].setToBeDone(false); + (*mTracks)[id].setInhibited(true); + } else if (id >= 0 && id < static_cast(mTrackIDtoParticlesEntry.size())) { + const int entry = mTrackIDtoParticlesEntry[id]; + if (entry >= 0 && entry < static_cast(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); diff --git a/Detectors/Base/test/testStack.cxx b/Detectors/Base/test/testStack.cxx index f6d32d3bf7157..042d64fb3302f 100644 --- a/Detectors/Base/test/testStack.cxx +++ b/Detectors/Base/test/testStack.cxx @@ -15,6 +15,7 @@ #include #include "DetectorsBase/Detector.h" #include "DetectorsBase/Stack.h" +#include "DetectorsBase/TrackTransportUtils.h" #include "SimulationDataFormat/BaseHits.h" #include "TFile.h" #include "TMCProcess.h" @@ -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::quiet_NaN(), 0.5f, false), std::runtime_error); + BOOST_CHECK_THROW(transportFromOnnxScore(std::numeric_limits::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); +} diff --git a/Steer/src/O2MCApplication.cxx b/Steer/src/O2MCApplication.cxx index dba61328c2d9c..33cdb03793064 100644 --- a/Steer/src/O2MCApplication.cxx +++ b/Steer/src/O2MCApplication.cxx @@ -9,6 +9,7 @@ // granted to it by virtue of its status as an Intergovernmental Organization // or submit itself to any jurisdiction. +#include #include #include @@ -200,6 +201,25 @@ void O2MCApplicationBase::PreTrack() // dispatch now to function in FairRoot FairMCApplication::PreTrack(); + + auto* stack = static_cast(GetStack()); + if (stack->hasTrackTransportModel() && fMC->TrackLength() == 0.) { + // Geant4 owns its secondary queue: clearing a stack bit alone cannot stop + // those tracks. Classify the engine's birth state, then explicitly stop it. + // Never classify a resumed track after it has already travelled or hit. + TLorentzVector position, momentum; + fMC->TrackPosition(position); + fMC->TrackMomentum(momentum); + TParticle particle(fMC->TrackPid(), 0, -1, -1, -1, -1, + momentum.Px(), momentum.Py(), momentum.Pz(), momentum.E(), + position.X(), position.Y(), position.Z(), fMC->TrackTime()); + if (!fMCEventHeader) { + throw std::runtime_error("ONNX track pruning requires the MC event vertex"); + } + if (!stack->transportTrack(particle, fMCEventHeader->GetX(), fMCEventHeader->GetY(), fMCEventHeader->GetZ())) { + fMC->StopTrack(); + } + } } void O2MCApplicationBase::ConstructGeometry()