27void EcalPnetVetoProcessor::configure(
29 auto disc_cut = parameters.
get<
double>(
"disc_cut");
30 disc_cut_logit_ = disc_cut <= 0 ? -std::numeric_limits<double>::infinity()
31 : std::log(disc_cut / (1. - disc_cut));
32 rt_ = std::make_unique<ldmx::ort::ONNXRuntime>(
33 parameters.
get<std::string>(
"model_path"));
36 collection_name_ = parameters.
get<std::string>(
"collection_name");
38 rec_coll_name_ = parameters.
get<std::string>(
"rec_coll_name");
39 ecal_rec_hits_passname_ =
40 parameters.
get<std::string>(
"ecal_rec_hits_passname");
41 ecal_sp_coll_name_ = parameters.
get<std::string>(
"ecal_sp_coll_name");
42 ecal_sp_hits_passname_ = parameters.
get<std::string>(
"ecal_sp_hits_passname");
43 track_pass_name_ = parameters.
get<std::string>(
"track_pass_name",
"");
44 track_collection_ = parameters.
get<std::string>(
"track_collection");
45 recoil_from_tracking_ = parameters.
get<
bool>(
"recoil_from_tracking");
51 const auto& ecal_geometry = getCondition<ldmx::EcalGeometry>(
52 ldmx::EcalGeometry::CONDITIONS_OBJECT_NAME);
55 const auto ecal_rec_hits =
event.getCollection<
ldmx::EcalHit>(
56 rec_coll_name_, ecal_rec_hits_passname_);
57 auto nhits = std::count_if(
58 ecal_rec_hits.begin(), ecal_rec_hits.end(),
59 [](
const ldmx::EcalHit& hit) { return hit.getEnergy() > 0; });
62 ldmx_log(debug) <<
"nhits = " << nhits <<
" MAX_NUM_HITS = " << MAX_NUM_HITS;
63 if (nhits < MAX_NUM_HITS) {
64 std::array<double, 3> etraj = {-999., -999., -999.};
65 std::array<double, 3> enorm = {-999., -999., -999.};
69 if (!recoil_from_tracking_ &&
70 event.
exists(ecal_sp_coll_name_, ecal_sp_hits_passname_)) {
72 ecal_sp_coll_name_, ecal_sp_hits_passname_);
73 double electron_pz_max = -1.0;
74 for (
auto const& hit : ecal_sp_hits) {
76 if (hit.getPdgID() != 11)
continue;
77 double electron_z = hit.getPosition()[2];
79 if (electron_z <= 239.0 || electron_z >= 240.0)
continue;
80 double electron_pz = hit.getMomentum()[2];
82 if (electron_pz > electron_pz_max) {
83 electron_pz_max = electron_pz;
90 ldmx_log(debug) <<
"Electron Found in the Ecal SP!";
93 ldmx_log(debug) <<
"ECAL SP pos_=(" << pos[0] <<
"," << pos[1] <<
","
95 ldmx_log(debug) <<
"ECAL SP mom=(" << mom[0] <<
"," << mom[1] <<
","
97 etraj = {pos[0], pos[1], pos[2]};
101 enorm = {mom[0] / pz, mom[1] / pz, 1.0};
104 }
else if (recoil_from_tracking_) {
106 auto recoil_tracks{
event.getCollection<
ldmx::Track>(track_collection_,
108 ldmx::TrackStateType ts_at_ecal = ldmx::AtECAL;
109 auto recoil_track_states_ecal =
111 if (!recoil_track_states_ecal.empty()) {
112 std::array<double, 3> pos = {recoil_track_states_ecal[0],
113 recoil_track_states_ecal[1],
114 recoil_track_states_ecal[2]};
115 std::array<double, 3> mom = {(recoil_track_states_ecal[3]),
116 (recoil_track_states_ecal[4]),
117 (recoil_track_states_ecal[5])};
118 ldmx_log(debug) <<
"Electron track pos_=(" << pos[0] <<
"," << pos[1]
119 <<
"," << pos[2] <<
")";
120 ldmx_log(debug) <<
"Electron track mom=(" << mom[0] <<
"," << mom[1]
121 <<
"," << mom[2] <<
")";
125 enorm = {mom[0] / pz, mom[1] / pz, 1.0};
128 ldmx_log(info) <<
" No recoil track at ECAL";
131 ldmx_log(fatal) <<
" No electron hit at scoring plane or no tracking";
134 if (etraj[0] == -999.) {
136 result.setDiscValue(-99);
137 result.setVetoResult(
false);
140 makeInputs(ecal_geometry, ecal_rec_hits, etraj, enorm);
142 auto logits = rt_->run(INPUT_NAMES, data_)[0];
145 auto prob = std::exp((logSoftmax(logits)[1]));
146 result.setDiscValue(prob);
148 double logit_diff = logits[1] - logits[0];
149 ldmx_log(debug) <<
"ParticleNet logit difference = " << logit_diff;
150 result.setVetoResult(logit_diff > disc_cut_logit_);
153 result.setDiscValue(-99);
154 result.setVetoResult(
false);
157 ldmx_log(debug) <<
"ParticleNet disc value = " << result.getDisc();
166 event.add(collection_name_, result);
169void EcalPnetVetoProcessor::makeInputs(
171 const std::vector<ldmx::EcalHit>& ecal_rec_hits,
172 std::array<double, 3> etraj, std::array<double, 3> enorm) {
174 for (
auto& v : data_) {
175 std::fill(v.begin(), v.end(), 0);
180 for (
const auto& hit : ecal_rec_hits) {
181 if (hit.getEnergy() <= 0) {
183 <<
"Hit with zero energy found, should not happen, skipping it.";
189 double delta_z = hit_z - etraj[2];
190 double etraj_x = etraj[0] + enorm[0] * delta_z;
191 double etraj_y = etraj[1] + enorm[1] * delta_z;
192 data_[0].at(COORDINATE_X_OFFSET + idx) = hit_x - etraj_x;
193 data_[0].at(COORDINATE_Y_OFFSET + idx) = hit_y - etraj_y;
194 data_[0].at(COORDINATE_Z_OFFSET + idx) = hit_z;
196 data_[1].at(FEATURE_X_OFFSET + idx) = hit_x - etraj_x;
197 data_[1].at(FEATURE_Y_OFFSET + idx) = hit_y - etraj_y;
198 data_[1].at(FEATURE_Z_OFFSET + idx) = hit_z;
199 data_[1].at(FEATURE_LAYER_ID_OFFSET + idx) =
id.layer();
200 data_[1].at(FEATURE_ENERGY_OFFSET + idx) = std::log(hit.getEnergy());
205 std::stringstream ss;
206 for (
unsigned iname = 0; iname < INPUT_NAMES.size(); ++iname) {
207 ss <<
"=== " << INPUT_NAMES[iname] <<
" ===";
208 for (
unsigned i = 0; i < INPUT_SIZES[iname]; ++i) {
209 ss << data_[iname].at(i) <<
", ";
210 if ((i + 1) % MAX_NUM_HITS == 0) {
215 ldmx_log(trace) << ss.str();
218std::vector<float> EcalPnetVetoProcessor::logSoftmax(
219 const std::vector<float>& logits) {
221 auto max_val = *std::max_element(logits.begin(), logits.end());
224 std::vector<float> exp_vals(logits.size());
225 for (
size_t i = 0; i < logits.size(); ++i) {
226 exp_vals[i] = std::exp(logits[i] - max_val);
229 float sum_exp = std::accumulate(exp_vals.begin(), exp_vals.end(), 0.0);
230 float log_sum_exp = max_val + std::log(sum_exp);
233 std::vector<float> result(logits.size());
234 for (
size_t i = 0; i < logits.size(); ++i) {
235 result[i] = logits[i] - log_sum_exp;