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,17 @@ struct StackParam : public o2::conf::ConfigurableParamHelper<StackParam> {
std::string transportPrimaryFileName = "";
std::string transportPrimaryFuncName = "";
bool transportPrimaryInvert = false;
// simnet.birth.v1: 34 raw birth features. The ONNX graph selects its inputs,
// embeds preprocessing/domain guards and returns [batch,1] class-1 probability.
// Class 1 means no recorded hits in the entire subtree. Invalid scores keep tracks.
std::string transportPrimaryOnnxCCDBPath = "";
// Double preserves validation cuts just above tied float32 scores.
double transportPrimaryOnnxThreshold = -1.; // require explicit configuration
int transportPrimaryOnnxOutputIndex = 0;
bool transportPrimaryOnnxApplySigmoid = false; // probability already in graph
// Saved roots-v3 models only saw simulation roots, including injected tracks.
// Enable only for a model trained/validated on transport secondaries too.
bool transportPrimaryOnnxSecondaries = false;

// 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
57 changes: 57 additions & 0 deletions Detectors/Base/README-track-pruning.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,57 @@
# ONNX transport pruning: simnet.birth.v1

O2 supplies one float32 `[batch,34]` input named `birth_features_v1`. The order is
`TrackTransportFeatureNames` in `TrackTransportUtils.h`, mirrored by
`TRANSPORT_FEATURES` in simnet's `src/training/qa_tools.py`:

```
charge_sign,mass,ekin,pt,eta,phi,theta,vx,vy,vz,r_xy,t_ns,
pdg_class,lifetime_ns,bending_radius_cm,z_at_r40,z_at_r250,medium_code,
has_gen_mother,pdg,abs_pdg,energy,px,py,pz,p,rapidity,dx_from_event,
dy_from_event,dz_from_event,r_from_event_xy,r3_from_event,mother_pdg,mother_abs_pdg
```

These are birth-time quantities, never hit labels or post-transport bookkeeping.
`has_gen_mother` uses the first mother, including transport mothers. Missing
ancestry, event vertex or medium is represented by NaN in the affected columns.
Birth-medium lookup uses a private per-thread navigator.

The model uses ONNX Gather to select/reorder its configured `labels_x` columns.
Only selected inputs are checked for finite values and the model's training
ranges. Unused NaN/Inf columns have no effect; multiplying them by zero would
not be sufficient. Scaling is embedded after selection, including any external
NN standard scaler. Invalid selected inputs return a NaN score, which O2 always
interprets as KEEP, including when inversion is enabled.

There is one float32 `[batch,1]` output, `probability_hit_free_subtree`. The NN's
sigmoid and the BDT's selection of probability column 1 are inside the graph.
Class 1 means no recorded detector hits in the track's entire descendant tree.
It does not mean electrically neutral or prove zero deposited energy.

Configuration:

```
Stack.transportPrimary=onnx
Stack.transportPrimaryOnnxCCDBPath=<raw ONNX object path>
Stack.transportPrimaryOnnxThreshold=<saved validation-selected threshold>
Stack.transportPrimaryOnnxOutputIndex=0
Stack.transportPrimaryOnnxApplySigmoid=false
Stack.transportPrimaryInvert=false
Stack.transportPrimaryOnnxSecondaries=false
```

The threshold has double precision to preserve cuts immediately above tied
float32 scores. The default -1 requires explicit configuration. The threshold
`nextafter(1,+inf)` is also accepted to represent a keep-all operating point.
CCDB validity and creation-time constraints must cover the simulation timestamp.

Pruning defaults to simulation roots, including injected tracks, matching the
current simnet training configs. Enable secondary pruning only after training
and validating that cohort. Feature selection does not change the training
population or make existing weights suitable for unseen populations.

With the updated training scripts, deploy `network/net.onnx` or
`network/bdt.onnx`. `*.core.onnx` files are intermediate models with selected
inputs, not O2-compatible exports. Earlier 19/21/25-input exports are rejected.
The full input/selection contract and operating point are saved beside the
model as `net.json` or `bdt.json` and embedded in ONNX metadata.
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, double)> mTransportTrack; //! ONNX decision for primaries and secondaries

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

Expand Down
185 changes: 185 additions & 0 deletions Detectors/Base/include/DetectorsBase/TrackTransportUtils.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,185 @@
// 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.

/// Helpers for the versioned simnet.birth.v1 raw-input transport contract.
#ifndef O2_TRACK_TRANSPORT_UTILS_H
#define O2_TRACK_TRANSPORT_UTILS_H

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

namespace o2::data::detail
{
// Stable, append-only ABI. Models Gather their selected columns in the ONNX
// graph; unused NaN/Inf values must not inhibit an otherwise valid prediction.
constexpr std::array<std::string_view, 34> TrackTransportFeatureNames{
"charge_sign",
"mass",
"ekin",
"pt",
"eta",
"phi",
"theta",
"vx",
"vy",
"vz",
"r_xy",
"t_ns",
"pdg_class",
"lifetime_ns",
"bending_radius_cm",
"z_at_r40",
"z_at_r250",
"medium_code",
"has_gen_mother",
"pdg",
"abs_pdg",
"energy",
"px",
"py",
"pz",
"p",
"rapidity",
"dx_from_event",
"dy_from_event",
"dz_from_event",
"r_from_event_xy",
"r3_from_event",
"mother_pdg",
"mother_abs_pdg"};
constexpr size_t TrackTransportFeatureCount = TrackTransportFeatureNames.size();

inline bool validTrackTransportFeatures(const std::vector<float>& features)
{
// Numerical/domain checks belong in the graph AFTER selecting model inputs.
return features.size() == TrackTransportFeatureCount;
}

inline float trackTransportMediumCode(std::string_view medium)
{
return medium == "PIPE_VACUUM" ? 1.f : medium == "TPC_DriftGas2" ? 2.f : medium == "TPC_Air" ? 3.f : 0.f;
}

inline double trackTransportZAtRadius(const TParticle& p, double radius)
{
const double pt = std::hypot(p.Px(), p.Py());
if (pt <= 0.) {
return std::numeric_limits<double>::quiet_NaN();
}
const double r0 = std::hypot(p.Vx(), p.Vy());
if (r0 >= radius) {
return p.Vz();
}
const double b = p.Vx() * p.Px() + p.Vy() * p.Py();
const double c = r0 * r0 - radius * radius;
return p.Vz() + (-b + std::sqrt(std::max(0., b * b - pt * pt * c))) / (pt * pt) * p.Pz();
}

// Birth-time quantities only, in TrackTransportFeatureNames order. Missing
// context is NaN; a graph that does not select that context can still classify.
// motherPdg is the FIRST mother's PDG, including transport mothers.
inline std::vector<float> makeTrackTransportFeatures(
const TParticle& particle, double motherPdg, float mediumCode,
double eventX = std::numeric_limits<double>::quiet_NaN(),
double eventY = std::numeric_limits<double>::quiet_NaN(),
double eventZ = std::numeric_limits<double>::quiet_NaN())
{
const double px = particle.Px(), py = particle.Py(), pz = particle.Pz();
const double momentum = std::sqrt(px * px + py * py + pz * pz);
const double pt = std::hypot(px, py);
const int code = particle.GetPdgCode();
const auto absCode = std::abs(static_cast<int64_t>(code));
bool massKnown = false;
double mass = o2::O2DatabasePDG::Mass(code, massKnown);
const auto* pdg = TDatabasePDG::Instance()->GetParticle(code);
if (!massKnown) {
mass = pdg ? pdg->Mass() : 0.;
}
double charge = !pdg || pdg->Charge() == 0. ? 0. : std::copysign(1., pdg->Charge());
if (absCode >= 1000000000) {
charge = (absCode / 10000) % 1000 == 0 ? 0. : (code > 0 ? 1. : -1.);
}
int pdgClass = 9;
if (absCode == 22) {
pdgClass = 0;
} else if (absCode == 11) {
pdgClass = 1;
} else if (absCode == 13) {
pdgClass = 2;
} else if (absCode == 12 || absCode == 14 || absCode == 16) {
pdgClass = 3;
} else if (absCode >= 1000000000) {
pdgClass = 8;
} else if (absCode >= 1000 && absCode < 10000) {
pdgClass = charge != 0. ? 6 : 7;
} else if (absCode >= 100 && absCode < 1000) {
pdgClass = charge != 0. ? 4 : 5;
}
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;
const double dx = particle.Vx() - eventX, dy = particle.Vy() - eventY, dz = particle.Vz() - eventZ;
const double hasMother = std::isfinite(motherPdg) ? (motherPdg != 0. ? 1. : 0.) : missing;
return {static_cast<float>(charge), static_cast<float>(mass), static_cast<float>(energy - mass),
static_cast<float>(pt), static_cast<float>(eta), static_cast<float>(std::atan2(py, px)),
static_cast<float>(theta), static_cast<float>(particle.Vx()), static_cast<float>(particle.Vy()),
static_cast<float>(particle.Vz()), static_cast<float>(std::hypot(particle.Vx(), particle.Vy())),
static_cast<float>(particle.T() * 1.e9), static_cast<float>(pdgClass),
static_cast<float>(pdg ? pdg->Lifetime() * 1.e9 : 0.),
static_cast<float>(charge != 0. ? pt / (0.3 * 0.5) * 100. : 0.),
static_cast<float>(trackTransportZAtRadius(particle, 40.)),
static_cast<float>(trackTransportZAtRadius(particle, 250.)), mediumCode, static_cast<float>(hasMother),
static_cast<float>(code), static_cast<float>(absCode), static_cast<float>(energy),
static_cast<float>(px), static_cast<float>(py), static_cast<float>(pz), static_cast<float>(momentum),
static_cast<float>(rapidity), static_cast<float>(dx), static_cast<float>(dy), static_cast<float>(dz),
static_cast<float>(std::hypot(dx, dy)), static_cast<float>(std::sqrt(dx * dx + dy * dy + dz * dz)),
static_cast<float>(motherPdg), static_cast<float>(std::abs(motherPdg))};
}

inline bool validTrackTransportOutput(const std::vector<std::vector<int64_t>>& shapes, int index)
{
if (shapes.size() != 1 || shapes[0].empty() || shapes[0].size() > 2 ||
(shapes[0][0] != 1 && shapes[0][0] != -1) || index < 0) {
return false;
}
// The existing NN exports squeeze(-1), so [batch] is a single score.
return shapes[0].size() == 1 ? index == 0 : index < shapes[0][1];
}

inline bool transportFromOnnxScore(float score, double threshold, bool applySigmoid, bool invert = false)
{
if (!std::isfinite(threshold) || threshold < 0. || threshold > std::nextafter(1., std::numeric_limits<double>::infinity())) {
throw std::runtime_error("ONNX pruning requires an explicit probability threshold (or nextafter(1,+inf) to keep all)");
}
// Invalid predictions always keep the track, including with inversion enabled.
if (!std::isfinite(score) || (!applySigmoid && (score < 0.f || score > 1.f))) {
return true;
}
if (applySigmoid) {
score = score >= 0.f ? 1.f / (1.f + std::exp(-score)) : std::exp(score) / (1.f + std::exp(score));
}
const bool transport = static_cast<double>(score) < threshold; // class 1 is a hit-free subtree
return invert ? !transport : transport;
}
} // namespace o2::data::detail
#endif
Loading
Loading