Belle II Software light-2609-luna
ModeSelectorModule.py
1#!/usr/bin/env python3
2
9
10"""
11ModeSelector Module for event-level B meson classification.
12
13This module evaluates a two-stage neural network on FEI B meson candidates
14to compute BplusScore, an event-level score based on all candidates.
15
16The module:
171. Collects features from all B candidates in the event
182. Runs a category network to classify as B0/B+/continuum
193. Runs a main network using category output to compute final scores
204. Stores results as ExtraInfo
21"""
22
23import basf2 as b2
24import numpy as np
25from modeSelector import config
26from ROOT import Belle2
27from variables import variables as vm
28
29
30class ModeSelectorModule(b2.Module):
31 """
32 Event-level B meson classifier using neural networks.
33
34 This module processes FEI B meson candidates and computes a score using
35 information from all candidates in the event.
36
37 Args:
38 particle_lists (list): List of B meson particle list names
39 cat_model_path (str): Path to the basf2 MVA weightfile for the category network,
40 as produced by convert_to_onnx.py. Pass None to load from the conditions
41 database via payload_cat_model.
42 main_model_path (str): Path to the basf2 MVA weightfile for the main network,
43 as produced by convert_to_onnx.py. Pass None to load from the conditions
44 database via payload_main_model.
45 payload_cat_model (str): Conditions DB payload name for the category model. None
46 (default) uses the name derived from the contract version this release
47 implements, which is the normal case.
48 payload_main_model (str): Same for the main model.
49 output_variable (str): Name of ExtraInfo variable for output score
50 training_mode (bool): Expose features and MC truth for training instead of running
51 the networks.
52 skip_nn_evaluation (bool): Fill placeholder outputs instead of running the networks
53 (debugging and timing only).
54 store_fei_calib_weight (bool): Store modeSelector_feiCalibWeight (MC only).
55 debug (bool): Print the feature vector and network outputs for the first events.
56 debug_max_events (int): Number of events printed when debug is True.
57
58 See modeSelector.modeSelector() for the full description of each parameter.
59 """
60
61 AUXILIARY_OUTPUT_PREFIX = 'modeSelector'
62
64 self,
65 particle_lists,
66 payload_cat_model=None,
67 payload_main_model=None,
68 output_variable='BplusScore',
69 cat_model_path=None,
70 main_model_path=None,
71 training_mode=False,
72 skip_nn_evaluation=False,
73 store_fei_calib_weight=False,
74 debug=False,
75 debug_max_events=10,
76 ):
77 """Initialise module parameters and training data buffers."""
78 super().__init__()
79
80 self.particle_lists = particle_lists if isinstance(particle_lists, list) else [particle_lists]
81
82 self.cat_model_path = cat_model_path
83
84 self.main_model_path = main_model_path
85
86 self.output_variable = output_variable
87 # None means: use the name derived from the contract version this release implements
88
89 self.payload_cat_model_given = payload_cat_model is not None
90
91 self.payload_main_model_given = payload_main_model is not None
92
93 self.payload_cat_model = payload_cat_model if payload_cat_model is not None else config.DEFAULT_CAT_PAYLOAD
94
95 self.payload_main_model = payload_main_model if payload_main_model is not None else config.DEFAULT_MAIN_PAYLOAD
96
100 self.contract_version = config.MODEL_CONTRACT_VERSION
101
102 self.training_mode = training_mode
103
104 self.skip_nn_evaluation = skip_nn_evaluation
105
106 self.store_fei_calib_weight = store_fei_calib_weight
107
109
110 self.debug = debug
111
112 self.debug_max_events = debug_max_events
113
114 self.event_count = 0
115
117
119
121
123
125
127
129
131
133
135
137
139
140 # Feature configuration (from modeSelector.config)
141
142 self.n_input_ids = config.N_INPUT_IDS
143
144 self.event_features = config.EVENT_FEATURES
145
146 def _load_weightfile(self, path, label):
147 """Load a basf2 MVA weightfile, raises error for unsupported extensions."""
148 if not path.endswith('.root'):
149 b2.B2FATAL(
150 "ModeSelector: " + label + " model path '" + path + "' is not a basf2 MVA "
151 "weightfile. Pass the file produced by convert_to_onnx.py "
152 "(must have a .root extension), not a raw ONNX file."
153 )
155
156 def _feature_selection_from_weightfile(self, variables, label):
157 """Recover the raw feature indices a model was trained on from its weightfile.
158
159 Returns None for legacy weightfiles that only carry placeholder variable names,
160 in which case the caller falls back to config.HAS_INPUTS.
161 """
162 prefix = config.FEATURE_VAR_PREFIX
163 indices = []
164 for name in variables:
165 name = str(name)
166 if not name.startswith(prefix):
167 continue
168 try:
169 indices.append(int(name[len(prefix):]))
170 except ValueError:
171 b2.B2FATAL(
172 "ModeSelector: " + label + " model weightfile has a malformed feature "
173 "name '" + name + "'. Re-export the model with convert_to_onnx.py."
174 )
175 return indices if indices else None
176
178 """Check that SUPPORTED_CONTRACT_VERSIONS is consistent with the current contract version.
179
180 A mismatch is a developer error in the release, not a problem with the payload.
181 The same rule is applied by convert_to_onnx.py before exporting.
182 """
183 error = config.contract_consistency_error()
184 if error:
185 b2.B2FATAL("ModeSelector: " + error)
186
187 def _check_contract(self, weightfile, identifier, label):
188 """Check a model's contract version and return it together with the training name.
189
190 The contract version is an extra element of the weightfile; the identifier holds the
191 training name. The contract version must be one this release supports: the current
192 one, or an older one whose code path the release still provides. A newer version, or
193 an older one that is no longer supported, would run without error and give wrong
194 results, so it is fatal. Weightfiles written before the contract version was recorded
195 are treated as contract version 1 and are subject to the same support rule.
196 """
197 training = str(identifier)
198 if not weightfile.hasContractVersion():
199 b2.B2WARNING(
200 "ModeSelector: " + label + " model weightfile carries no contract version, "
201 "so it is treated as contract version 1. Re-export it with convert_to_onnx.py."
202 )
203 # Still subject to the support rule: without it, a release that dropped contract 1
204 # would run such a model with the wrong code path.
205 self._require_supported_contract(1, label)
206 return 1, training
207
208 # getContractVersion() also returns the default if the stored value is not an
209 # integer, which hasContractVersion() above has already ruled out as 'not stored'.
210 version = weightfile.getContractVersion(-1)
211 if version < 0:
212 b2.B2FATAL(
213 "ModeSelector: " + label + " model weightfile has a malformed contract "
214 "version. Re-export it with convert_to_onnx.py."
215 )
216
217 self._require_supported_contract(version, label)
218 b2.B2INFO(
219 "ModeSelector: " + label + " model training '" + training
220 + "' (contract v" + str(version) + ")"
221 )
222 return version, training
223
224 def _require_supported_contract(self, version, label):
225 """Stop unless this release can run models built for the given contract version."""
226 if version in config.SUPPORTED_CONTRACT_VERSIONS:
227 return
228 current = config.MODEL_CONTRACT_VERSION
229 if version > current:
230 b2.B2FATAL(
231 "ModeSelector: " + label + " model was built for contract version " + str(version)
232 + ", newer than contract version " + str(current) + " implemented by this release. "
233 "Use a newer release for this payload."
234 )
235 b2.B2FATAL(
236 "ModeSelector: " + label + " model was built for contract version " + str(version)
237 + ", which this release no longer supports (supported: "
238 + ", ".join(str(v) for v in sorted(config.SUPPORTED_CONTRACT_VERSIONS))
239 + "). Use an older release for this payload."
240 )
241
242 def _select_contract_version(self, cat_version, main_version):
243 """Return the contract version to run with, and warn when it is older than the current one.
244
245 Both networks are evaluated by one code path, so they must share a contract version.
246 """
247 if cat_version != main_version:
248 b2.B2FATAL(
249 "ModeSelector: the category model was built for contract version " + str(cat_version)
250 + " but the main model for contract version " + str(main_version)
251 + ". Both must come from the same contract."
252 )
253 current = config.MODEL_CONTRACT_VERSION
254 if cat_version < current:
255 b2.B2WARNING(
256 "ModeSelector: the models were built for contract version " + str(cat_version)
257 + ", older than contract version " + str(current) + " implemented by this release. "
258 "The module falls back to the behaviour of contract version " + str(cat_version)
259 + ", which this release still supports. See the ModeSelector performance "
260 "recommendations for what changed in contract version " + str(current)
261 + " and whether using its models is worthwhile."
262 )
263 return cat_version
264
265 def _check_same_training(self, cat_training, main_training):
266 """Check that the category and main models come from the same training.
267
268 The main network takes the category network outputs as inputs, so a pair from
269 different trainings runs without error but gives wrong results. Such a pair shares
270 the contract version, so only the training name in the weightfile
271 identifier, which convert_to_onnx.py writes identically into both weightfiles, can
272 tell them apart.
273 """
274 if cat_training != main_training:
275 b2.B2FATAL(
276 "ModeSelector: the category model comes from training '" + cat_training
277 + "' but the main model from training '" + main_training + "'. The main network "
278 "takes the category outputs as inputs, so both must come from the same training. "
279 "Check that the two payloads were exported and uploaded together."
280 )
281
282 def _check_output_classes(self, options, expected, label):
283 """Check the model's output size against what the module's interpretation assumes.
284
285 The outputs are read positionally -- for the main network the class index is the
286 input_id and the last three entries are the background classes -- so a model with a
287 different number of outputs would be misread rather than rejected.
288 """
289 n_classes = int(options.m_nClasses)
290 if n_classes != expected:
291 b2.B2FATAL(
292 "ModeSelector: " + label + " model produces " + str(n_classes) + " output "
293 "classes but this software interprets " + str(expected) + ". The payload does "
294 "not match this release."
295 )
296
298 """Warn when the configured globaltags also serve payloads for a newer contract version.
299
300 Payload names are derived from the contract version, so a single performance
301 globaltag can hold models for several releases at once. Finding a newer one means a
302 newer release would use a different, possibly improved model. Looking up a payload
303 that does not exist only queries the metadata; the file of a newer payload is
304 fetched only if it is actually present.
305 """
306 # Contract versions are bumped rarely, and a globaltag may drop payloads for versions
307 # no supported release uses, so look a few versions ahead instead of stopping at the
308 # first gap.
309 probe_range = 10
310 current = config.MODEL_CONTRACT_VERSION
311 newest = None
312 for version in range(current + 1, current + 1 + probe_range):
313 name = config.payload_names(version)[0]
314 accessor = Belle2.DBAccessorBase(Belle2.DBStoreEntry.c_RawFile, name, False)
315 if accessor.getFilename():
316 newest = (version, name)
317 if newest is not None:
318 b2.B2WARNING(
319 "ModeSelector: the configured globaltags also provide payloads for contract version "
320 + str(newest[0]) + " ('" + newest[1] + "'), which newer releases use. This release "
321 "implements contract version " + str(current) + ". See the ModeSelector performance "
322 "recommendations for what changed and whether moving to a newer release is worthwhile."
323 )
324
325 def initialize(self):
326 """Called at the beginning of processing."""
327 Belle2.PyStoreObj('EventExtraInfo').registerInDataStore()
328
329 # Build feature index mapping (needed for both inference and training)
331
332 if self.training_mode:
333 b2.B2INFO("ModeSelector: Running in TRAINING mode (saving features, no NN inference)")
334
335 self.has_inputs = None
336
337 bp_calib_map = config.get_fei_calibration_map(521)
338
339 self._bp_calib_rest = config.get_fei_calibration_rest(521)
340
341 self._bp_calib_lookup = np.full(config.N_BP_MODES, self._bp_calib_rest, dtype=np.float32)
342 for _dm, _w in bp_calib_map.items():
343 if 0 <= _dm < config.N_BP_MODES:
344 self._bp_calib_lookup[_dm] = _w
345
346 b0_calib_map = config.get_fei_calibration_map(511)
347
348 self._b0_calib_rest = config.get_fei_calibration_rest(511)
349
350 self._b0_calib_lookup = np.full(config.N_B0_MODES, self._b0_calib_rest, dtype=np.float32)
351 for _dm, _w in b0_calib_map.items():
352 if 0 <= _dm < config.N_B0_MODES:
353 self._b0_calib_lookup[_dm] = _w
354 return
355
356 if self.skip_nn_evaluation:
357 self.has_inputs = list(config.HAS_INPUTS)
358
359 self.cat_input_size = len(self.has_inputs)
360
362 b2.B2INFO("ModeSelector: Running with NN evaluation disabled")
363 b2.B2INFO(
364 f"ModeSelector: Using placeholder outputs with {self.cat_input_size} selected features"
365 )
366 return
367
368 # A non-default payload name is the wrong way to pick a training: the names are
369 # deliberately version-free so that the globaltag decides which models are served.
370 # Only warn for names that are actually used, i.e. not overridden by a local file.
371 # An explicitly given payload name bypasses the name derived from the contract version.
372 # Local weightfiles take precedence, so names they override are not reported.
373 explicit = []
374 if not self.cat_model_path and self.payload_cat_model_given:
375 explicit.append(self.payload_cat_model + " (default " + config.DEFAULT_CAT_PAYLOAD + ")")
376 if not self.main_model_path and self.payload_main_model_given:
377 explicit.append(self.payload_main_model + " (default " + config.DEFAULT_MAIN_PAYLOAD + ")")
378 if explicit:
379 b2.B2WARNING(
380 "ModeSelector: not using the default payloads, explicitly requested: " + ", ".join(explicit)
381 + ". Normally the payload names are left unset, so they are derived from the contract"
382 " version this release implements, and the training is selected by prepending the"
383 " performance globaltag that serves it."
384 )
385
386 # Load models through the basf2 MVA Expert framework.
387 # Single-threaded ONNX execution is enforced by the framework in mva/methods/src/ONNX.cc.
389
391
392 # Local weightfiles are loaded once. Payloads are looked up in beginRun, so a
393 # globaltag can serve different models for different run ranges.
394
395 self._cat_local_wf = None
396
397 self._main_local_wf = None
398
399 self._cat_accessor = None
400
401 self._main_accessor = None
402 if self.cat_model_path:
403 self._cat_local_wf = self._load_weightfile(self.cat_model_path, 'category')
404 else:
406 Belle2.DBStoreEntry.c_RawFile, self.payload_cat_model, True
407 )
408 if self.main_model_path:
409 self._main_local_wf = self._load_weightfile(self.main_model_path, 'main')
410 else:
412 Belle2.DBStoreEntry.c_RawFile, self.payload_main_model, True
413 )
414
416
418 if self._cat_accessor is None and self._main_accessor is None:
420
421 def beginRun(self):
422 """Load the models from the conditions database if the payloads changed."""
423 if self.training_mode or self.skip_nn_evaluation:
424 return
425 if self._cat_accessor is None and self._main_accessor is None:
426 return
427
428 # Same pattern as MVAExpert, but comparing checksums instead of hasChanged(): for
429 # payloads accessed as plain files (c_RawFile) the changed flag can stay unset for
430 # the payload already loaded when the accessor was created, so the models would
431 # never be loaded.
432 checksums = tuple(
433 accessor.getChecksum() if accessor is not None else None
434 for accessor in (self._cat_accessor, self._main_accessor)
435 )
436 if checksums == self._loaded_checksums:
437 return
438
439 first_load = self._loaded_checksums is None
440 cat_wf = self._cat_local_wf
441 if cat_wf is None:
442 cat_wf = self._load_payload(self._cat_accessor, self.payload_cat_model, 'category')
443 main_wf = self._main_local_wf
444 if main_wf is None:
445 main_wf = self._load_payload(self._main_accessor, self.payload_main_model, 'main')
446 if not first_load:
447 b2.B2INFO("ModeSelector: model payloads changed, reloading the models")
448 self._setup_models(cat_wf, main_wf)
449 self._loaded_checksums = checksums
450 if first_load and not self.cat_model_path:
452
453 def _load_payload(self, accessor, payload_name, label):
454 """Load a basf2 MVA weightfile from the conditions database payload of the current run."""
455 filename = accessor.getFilename()
456 if not filename:
457 b2.B2FATAL(
458 "ModeSelector: " + label + " model payload '"
459 + payload_name
460 + "' not found in the conditions database. This release implements "
461 "contract version " + str(config.MODEL_CONTRACT_VERSION) + ", so the "
462 "configured performance globaltag must provide payloads for it."
463 )
465
466 def _setup_models(self, cat_wf, main_wf):
467 """Create the experts for a pair of weightfiles and validate them.
468
469 Called once for local weightfiles, and for each run in which a payload changes.
470 Everything derived from the weightfiles (contract version, feature selection, input
471 sizes) is set here, so a new pair of models can differ in all of them.
472 """
473 import basf2_mva
474
475
476 self.cat_expert = self._onnx_interface.getExpert()
477 self.cat_expert.load(cat_wf)
478
479 self.main_expert = self._onnx_interface.getExpert()
480 self.main_expert.load(main_wf)
481
482 cat_opts = basf2_mva.GeneralOptions()
483 cat_wf.getOptions(cat_opts)
484 main_opts = basf2_mva.GeneralOptions()
485 main_wf.getOptions(main_opts)
486
487 cat_version, cat_training = self._check_contract(cat_wf, cat_opts.m_identifier, 'category')
488 main_version, main_training = self._check_contract(main_wf, main_opts.m_identifier, 'main')
489 self._check_same_training(cat_training, main_training)
490 self.contract_version = self._select_contract_version(cat_version, main_version)
491 self._check_output_classes(cat_opts, config.NUM_CAT_LABELS, 'category')
492 self._check_output_classes(main_opts, config.N_INPUT_IDS + 3, 'main')
493
494 # The payload records which raw feature indices its model was trained on, so a
495 # retraining can change the selection without a software release.
496 cat_selection = self._feature_selection_from_weightfile(cat_opts.m_variables, 'category')
497 main_selection = self._feature_selection_from_weightfile(main_opts.m_variables, 'main')
498 if cat_selection is None or main_selection is None:
499 b2.B2WARNING(
500 "ModeSelector: weightfile does not record its feature selection, falling back "
501 "to config.HAS_INPUTS. This only works if the payload was trained against the "
502 "same selection. Re-export the model with convert_to_onnx.py."
503 )
504 self.has_inputs = list(config.HAS_INPUTS)
505 selection_source = 'config.HAS_INPUTS'
506 elif cat_selection != main_selection:
507 b2.B2FATAL(
508 "ModeSelector: the category and main models were trained on different feature "
509 "selections, so they come from different trainings and must not be combined."
510 )
511 else:
512 self.has_inputs = cat_selection
513 selection_source = 'the weightfile'
514
515 # Input sizes derived from the weightfile variable list
516 self.cat_input_size = len(cat_opts.m_variables)
517 self.main_input_size = len(main_opts.m_variables)
518
519 # Pre-allocate datasets; m_input elements are overwritten per event
520
521 self.cat_dataset = Belle2.MVA.SingleDataset(cat_opts, [0.0] * self.cat_input_size, 1.0)
522
523 self.main_dataset = Belle2.MVA.SingleDataset(main_opts, [0.0] * self.main_input_size, 1.0)
524
525 b2.B2INFO(f"ModeSelector: Loaded category model (input size: {self.cat_input_size})")
526 b2.B2INFO(f"ModeSelector: Loaded main model (input size: {self.main_input_size})")
527 if self.has_inputs:
528 b2.B2INFO(f"ModeSelector: Using {len(self.has_inputs)} selected features from {selection_source}")
529
531 """
532 Build the feature index mapping from config.
533
534 The training used sparse matrices with specific non-zero columns.
535 We need to match that structure.
536 """
537
538 self.feature_blocks = [(name, transform) for name, _, transform in config.FEATURE_BLOCKS]
539
540 # Register basf2 aliases where alias name differs from variable string
541 for name, var, _ in config.FEATURE_BLOCKS:
542 if name != var:
543 vm.addAlias(name, var)
544
545
546 self.deltaM_cut = config.DELTA_M_CUT
547
548 @staticmethod
549 def _assign_rank_extra_info(particle_score_input_id_triples, variable_name):
550 """Assign deterministic descending ranks to the given particles."""
551 ranked_pairs = sorted(
552 particle_score_input_id_triples,
553 key=lambda triple: (-triple[1], triple[2])
554 )
555 for rank, (particle, _, _) in enumerate(ranked_pairs, start=1):
556 particle.addExtraInfo(variable_name, rank)
557
558 def _get_input_id(self, particle):
559 """
560 Compute input_id from decay mode ID, B type, and particle/antiparticle sign.
561
562 Encoding:
563 B+ sector (input_ids 0 to 2*N_BP_MODES-1):
564 anti-B+ (PDG=-521): dmID * 2 + 0
565 B+ (PDG=+521): dmID * 2 + 1
566 B0 sector (input_ids 2*N_BP_MODES to N_INPUT_IDS-1):
567 anti-B0 (PDG=-511): N_BP_MODES*2 + dmID * 2 + 0
568 B0 (PDG=+511): N_BP_MODES*2 + dmID * 2 + 1
569 """
570 dm_id = int(vm.evaluate('extraInfo(decayModeID)', particle))
571 pdg = int(vm.evaluate('PDG', particle))
572 is_particle = 1 if pdg > 0 else 0
573 offset = 0 if abs(pdg) == 521 else config.N_BP_MODES * 2
574 return offset + dm_id * 2 + is_particle
575
576 def _extract_particle_features(self, particle):
577 """
578 Extract features from a single particle.
579
580 Returns:
581 tuple: (input_id, feature_dict)
582 """
583 input_id = self._get_input_id(particle)
584
585 features = {}
586 for name, transform in self.feature_blocks:
587 val = vm.evaluate(name, particle)
588 if np.isnan(val):
589 val = None
590 features[name] = val
591
592 # Apply D* deltaMassDiff cuts
593 for key in ['Dst0_deltaMassDiff', 'Dstp_deltaMassDiff']:
594 val = features[key]
595 if val is not None:
596 if val < self.deltaM_cut[0] or val > self.deltaM_cut[1]:
597 features[key] = None
598
599 return input_id, features
600
601 def _build_feature_array(self, candidates_data, event_features):
602 """
603 Build the feature array for neural network input.
604
605 Parameters:
606 candidates_data: List of (input_id, features) tuples
607 event_features: Dict of event-level features
608
609 Returns:
610 numpy.ndarray: Feature array ready for NN input
611 """
612 # Initialize feature matrix (n_feature_blocks * n_input_ids)
613 n_blocks = len(self.feature_blocks)
614 feature_matrix = np.zeros((n_blocks, self.n_input_ids), dtype=np.float32)
615
616 # Deduplicate by input_id: keep candidate with highest sigProb
617 best_by_input_id = {}
618 for input_id, features in candidates_data:
619 if input_id < 0 or input_id >= self.n_input_ids:
620 continue
621 sig_prob = features.get('sigProb')
622 if sig_prob is None:
623 sig_prob = -1
624 if input_id not in best_by_input_id or sig_prob > best_by_input_id[input_id][1]:
625 best_by_input_id[input_id] = (features, sig_prob)
626
627 # Fill in per-candidate features using deduplicated candidates
628 for input_id, (features, _) in best_by_input_id.items():
629 for block_idx, (feat_name, transform) in enumerate(self.feature_blocks):
630 val = features.get(feat_name)
631 if val is not None:
632 feature_matrix[block_idx, input_id] = transform(val)
633
634 # Flatten to 1D (matches training format)
635 flat_features = feature_matrix.flatten()
636
637 # Add event-level features
638 event_feat_values = [
639 event_features.get(name, 0.0) for name in self.event_features
640 ]
641
642 # Add ncandidates / 10 (all candidates before deduplication)
643 n_candidates = len(candidates_data)
644 event_feat_values.append(n_candidates / 10.0)
645
646 # Find best candidate (highest sigProb) using deduped data
647 if best_by_input_id:
648 # Best candidate overall (highest sigProb among unique input_ids)
649 max_input_id = max(best_by_input_id.keys(), key=lambda k: best_by_input_id[k][1])
650
651 # Best candidate in B+ sector (input_id < N_BP_MODES * 2) and B0 sector (>= N_BP_MODES * 2)
652 bp_sector_ids = {k: v for k, v in best_by_input_id.items() if k < config.N_BP_MODES * 2}
653 b0_sector_ids = {k: v for k, v in best_by_input_id.items() if k >= config.N_BP_MODES * 2}
654
655 best_is_bp = max_input_id < config.N_BP_MODES * 2
656 if best_is_bp and b0_sector_ids:
657 scnd_max_input_id = max(b0_sector_ids.keys(), key=lambda k: b0_sector_ids[k][1])
658 elif not best_is_bp and bp_sector_ids:
659 scnd_max_input_id = max(bp_sector_ids.keys(), key=lambda k: bp_sector_ids[k][1])
660 else:
661 scnd_max_input_id = -10
662 else:
663 max_input_id = 0
664 scnd_max_input_id = -10
665
666 event_feat_values.append(max_input_id / 50.0)
667 event_feat_values.append(scnd_max_input_id / 50.0)
668
669 # Add experiment number (from EventMetaData)
671 event_feat_values.append(experiment / 10.0)
672
673 # Concatenate all features
674 all_features = np.concatenate([flat_features, np.array(event_feat_values, dtype=np.float32)])
675
676 return all_features, max_input_id, n_candidates
677
679 """Return the raw experiment id."""
680 event_meta = Belle2.PyStoreObj('EventMetaData')
681 experiment = int(event_meta.getExperiment()) if event_meta.isValid() else 0
682
683 if self.training_mode or experiment in config.ALLOWED_EXPERIMENTS:
684 return experiment
685
687 self._unsupported_experiment_counts[experiment] = (
688 self._unsupported_experiment_counts.get(experiment, 0) + 1
689 )
690 return config.DEFAULT_EXPERIMENT
691
693 """Extract event-level features from EventShapeContainer."""
694 event_features = {}
695
696 # Event shape variables are stored in EventShapeContainer
697 event_shape = Belle2.PyStoreObj('EventShapeContainer')
698 if event_shape.isValid():
699 obj = event_shape.obj()
700 # Map feature names to EventShapeContainer methods
701 if 'sphericity' in self.event_features:
702 # Sphericity = 3/2 * (lambda2 + lambda3)
703 event_features['sphericity'] = 1.5 * (obj.getSphericityEigenvalue(1) + obj.getSphericityEigenvalue(2))
704 if 'thrust' in self.event_features:
705 event_features['thrust'] = obj.getThrust()
706 if 'thrustAxisCosTheta' in self.event_features:
707 thrust_axis = obj.getThrustAxis()
708 # Compute cosTheta = z / |r|
709 import math
710 r = math.sqrt(thrust_axis.X()**2 + thrust_axis.Y()**2 + thrust_axis.Z()**2)
711 event_features['thrustAxisCosTheta'] = thrust_axis.Z() / r if r > 0 else 0.0
712 if 'aplanarity' in self.event_features:
713 # Aplanarity = 1.5 * smallest sphericity eigenvalue
714 event_features['aplanarity'] = 1.5 * obj.getSphericityEigenvalue(2)
715 if 'foxWolframR2' in self.event_features:
716 # R2 = H2/H0
717 h0 = obj.getFWMoment(0)
718 h2 = obj.getFWMoment(2)
719 event_features['foxWolframR2'] = h2 / h0 if h0 != 0 else 0.0
720 if 'harmonicMomentThrust0' in self.event_features:
721 event_features['harmonicMomentThrust0'] = obj.getHarmonicMomentThrust(0)
722 if 'harmonicMomentThrust1' in self.event_features:
723 event_features['harmonicMomentThrust1'] = obj.getHarmonicMomentThrust(1)
724 if 'harmonicMomentThrust2' in self.event_features:
725 event_features['harmonicMomentThrust2'] = obj.getHarmonicMomentThrust(2)
726
727 return event_features
728
729 def _compute_fei_calib_weight(self, best_bp, best_b0):
730 """
731 Compute event-level FEI calibration weight from the best candidates.
732
733 Based on the reconstructed candidate, checking for truth-compatible tag PDG
734 and DeltaP < DELTA_P_THRESH. Returns FEI_CALIB_CONT for continuum events.
735 Returns NaN when reco conditions are not met (non-continuum events where
736 tag PDG does not match or DeltaP is above threshold).
737 """
738 bp_sig = -1.0
739 b0_sig = -1.0
740 if best_bp is not None:
741 v = vm.evaluate('extraInfo(SignalProbability)', best_bp)
742 if not np.isnan(v):
743 bp_sig = float(v)
744 if best_b0 is not None:
745 v = vm.evaluate('extraInfo(SignalProbability)', best_b0)
746 if not np.isnan(v):
747 b0_sig = float(v)
748
749 best = best_bp if bp_sig >= b0_sig else best_b0
750 if best is None:
751 return config.FEI_CALIB_CONT
752
753 is_cont_v = vm.evaluate('isContinuumEvent', best)
754 if np.isnan(is_cont_v) or int(is_cont_v) == 1:
755 return config.FEI_CALIB_CONT
756
757 pdg_v = vm.evaluate('PDG', best)
758 dm_v = vm.evaluate('extraInfo(decayModeID)', best)
759 tag_pdg_v = vm.evaluate('mostcommonBTagPDG', best)
760 dp_v = vm.evaluate('mostcommonBTagDeltaP', best)
761
762 if np.isnan(pdg_v) or np.isnan(dm_v):
763 return config.FEI_CALIB_CONT
764
765 abs_pdg = abs(int(pdg_v))
766 if abs_pdg not in (511, 521):
767 return config.FEI_CALIB_CONT
768
769 calib_map = config.get_fei_calibration_map(abs_pdg)
770 calib_rest = config.get_fei_calibration_rest(abs_pdg)
771
772 # Reco path: truth-compatible tag PDG and low DeltaP
773 tag_is_gen = (not np.isnan(tag_pdg_v)) and config.truth_tag_matches_pdg(pdg_v, tag_pdg_v, dm_v)
774 dp_ok = (not np.isnan(dp_v)) and float(dp_v) < config.DELTA_P_THRESH
775 if tag_is_gen and dp_ok:
776 iid = self._get_input_id(best)
777 if abs_pdg == 521:
778 dm = max(0, min(iid // 2, config.N_BP_MODES - 1))
779 else:
780 dm = max(0, min((iid - config.N_BP_MODES * 2) // 2, config.N_B0_MODES - 1))
781 return calib_map.get(dm, calib_rest)
782
783 return np.nan
784
785 def _get_placeholder_outputs(self, features):
786 """Return deterministic placeholder category and main-network outputs."""
787 cat_output = np.array([0.2, 0.7, 0.1], dtype=np.float32)
788 main_output = np.full(config.N_INPUT_IDS + 3, 0.01, dtype=np.float32)
789 main_output[config.N_INPUT_IDS] = 0.005
790 main_output[config.N_INPUT_IDS + 1] = 0.002
791 main_output[config.N_INPUT_IDS + 2] = 0.001
792 return cat_output, main_output
793
794 def _extract_mc_truth_scalars(self, particle):
795 """Extract compact MC truth scalars used by training output."""
796 if particle is None:
797 return {
798 'pdg': np.nan,
799 'dm': np.nan,
800 'sigprob': np.nan,
801 'is_cont': np.nan,
802 'tag_pdg': np.nan,
803 'gen_dm_id': np.nan,
804 'gen_calib_weight': np.nan,
805 'dp': np.nan,
806 }
807
808 raw = {
809 'pdg': vm.evaluate('PDG', particle),
810 'dm': vm.evaluate('extraInfo(decayModeID)', particle),
811 'sigprob': vm.evaluate('extraInfo(SignalProbability)', particle),
812 'is_cont': vm.evaluate('isContinuumEvent', particle),
813 'tag_pdg': vm.evaluate('mostcommonBTagPDG', particle),
814 'gen_dm_id': vm.evaluate('extraInfo(genDecayModeID)', particle),
815 'gen_calib_weight': vm.evaluate('extraInfo(genFEICalibWeight)', particle),
816 'dp': vm.evaluate('mostcommonBTagDeltaP', particle),
817 }
818 return {
819 name: np.nan if np.isnan(value) else value
820 for name, value in raw.items()
821 }
822
824 self, bp_truth, b0_truth, best_bp_iid, best_bp_dp, best_b0_iid, best_b0_dp
825 ):
826 """
827 Compute the per-event training scalar fields written to EventExtraInfo.
828
829 best_bp_iid/best_bp_dp/best_b0_iid/best_b0_dp come from the truth-tag-matched
830 particle_by_input_id scan already performed in event() and are only stored for
831 the labels; everything else is derived from the best-B+/best-B0 MC truth scalars
832 for this event. In particular, the FEI calibration weight is defined for the
833 highest-sigProb candidate, so its DeltaP requirement uses that candidate's own
834 DeltaP, as at inference in _compute_fei_calib_weight().
835 """
836 bp_pdg = bp_truth['pdg']
837 b0_pdg = b0_truth['pdg']
838 bp_dm = bp_truth['dm']
839 b0_dm = b0_truth['dm']
840 bp_sig = -1.0 if np.isnan(bp_truth['sigprob']) else float(bp_truth['sigprob'])
841 b0_sig = -1.0 if np.isnan(b0_truth['sigprob']) else float(b0_truth['sigprob'])
842 bp_is_cont = bp_truth['is_cont']
843 b0_is_cont = b0_truth['is_cont']
844 bp_gen_pdg = bp_truth['tag_pdg']
845 b0_gen_pdg = b0_truth['tag_pdg']
846 bp_gen_dm_id = -1 if np.isnan(bp_truth['gen_dm_id']) else int(bp_truth['gen_dm_id'])
847 b0_gen_dm_id = -1 if np.isnan(b0_truth['gen_dm_id']) else int(b0_truth['gen_dm_id'])
848 bp_gen_calib_weight = 1.0 if np.isnan(bp_truth['gen_calib_weight']) else float(bp_truth['gen_calib_weight'])
849 b0_gen_calib_weight = 1.0 if np.isnan(b0_truth['gen_calib_weight']) else float(b0_truth['gen_calib_weight'])
850
851 bp_is_best = 1 if bp_sig >= b0_sig else 0
852 best_sigprob = max(bp_sig, b0_sig)
853
854 is_cont_f = bp_is_cont if bp_is_best else b0_is_cont
855 is_cont = 1 if is_cont_f == 1.0 else 0
856
857 gen_pdg_f = bp_gen_pdg if bp_is_best else b0_gen_pdg
858 gen_pdg = -1 if is_cont == 1 else (-1 if np.isnan(gen_pdg_f) else int(gen_pdg_f))
859
860 bp_dm_i = -1 if np.isnan(bp_dm) else int(bp_dm)
861 b0_dm_i = -1 if np.isnan(b0_dm) else int(b0_dm)
862
863 best_bp_sigprob_iid = -1
864 if not np.isnan(bp_pdg) and not np.isnan(bp_dm):
865 offset = 0 if abs(int(bp_pdg)) == 521 else config.N_BP_MODES * 2
866 sign = 1 if bp_pdg > 0 else 0
867 best_bp_sigprob_iid = offset + bp_dm_i * 2 + sign
868
869 best_b0_sigprob_iid = -1
870 if not np.isnan(b0_pdg) and not np.isnan(b0_dm):
871 offset = 0 if abs(int(b0_pdg)) == 521 else config.N_BP_MODES * 2
872 sign = 1 if b0_pdg > 0 else 0
873 best_b0_sigprob_iid = offset + b0_dm_i * 2 + sign
874
875 bp_tag_is_gen = int(config.truth_tag_matches_pdg(bp_pdg, bp_gen_pdg, bp_dm))
876 b0_tag_is_gen = int(config.truth_tag_matches_pdg(b0_pdg, b0_gen_pdg, b0_dm))
877
878 use_bp = bp_is_best == 1
879 tag_is_gen_ev = bp_tag_is_gen if use_bp else b0_tag_is_gen
880 best_dp_ev = bp_truth['dp'] if use_bp else b0_truth['dp']
881 sigprob_iid_ev = best_bp_sigprob_iid if use_bp else best_b0_sigprob_iid
882 bb_mask_ev = is_cont != 1
883 use_reco_ev = (tag_is_gen_ev == 1) and (best_dp_ev < config.DELTA_P_THRESH) and (sigprob_iid_ev >= 0)
884
885 fei_calib_weight = config.FEI_CALIB_CONT
886 bp_threshold = config.N_BP_MODES * 2
887
888 if bb_mask_ev and use_reco_ev:
889 if use_bp:
890 dm = min(max(best_bp_sigprob_iid // 2, 0), config.N_BP_MODES - 1)
891 fei_calib_weight = float(self._bp_calib_lookup[dm])
892 else:
893 dm = min(max((best_b0_sigprob_iid - bp_threshold) // 2, 0), config.N_B0_MODES - 1)
894 fei_calib_weight = float(self._b0_calib_lookup[dm])
895 elif bb_mask_ev:
896 abs_pdg_ev = abs(gen_pdg)
897 gen_dm_ev = bp_gen_dm_id if use_bp else b0_gen_dm_id
898 if abs_pdg_ev == 521:
899 fei_calib_weight = (
900 float(self._bp_calib_lookup[gen_dm_ev])
901 if 0 <= gen_dm_ev < config.N_BP_MODES else self._bp_calib_rest
902 )
903 elif abs_pdg_ev == 511:
904 fei_calib_weight = (
905 float(self._b0_calib_lookup[gen_dm_ev])
906 if 0 <= gen_dm_ev < config.N_B0_MODES else self._b0_calib_rest
907 )
908
909 return {
910 'is_cont': is_cont,
911 'gen_pdg': gen_pdg,
912 'bp_gen_decay_mode_id': bp_gen_dm_id,
913 'b0_gen_decay_mode_id': b0_gen_dm_id,
914 'bp_gen_fei_calib_weight': bp_gen_calib_weight,
915 'b0_gen_fei_calib_weight': b0_gen_calib_weight,
916 'bp_is_best': bp_is_best,
917 'best_sigprob': best_sigprob,
918 'best_bp_sigprob_iid': best_bp_sigprob_iid,
919 'best_b0_sigprob_iid': best_b0_sigprob_iid,
920 'bp_tag_is_gen': bp_tag_is_gen,
921 'b0_tag_is_gen': b0_tag_is_gen,
922 'fei_calib_weight': fei_calib_weight,
923 'best_bp_iid': best_bp_iid,
924 'best_bp_dp': best_bp_dp,
925 'best_b0_iid': best_b0_iid,
926 'best_b0_dp': best_b0_dp,
927 }
928
929 def _print_debug_info(self, candidates_data, event_features, all_features, max_input_id):
930 """Print debug information for preprocessing."""
931 # Get event identification
932 event_meta = Belle2.PyStoreObj('EventMetaData')
933 if event_meta.isValid():
934 exp = event_meta.getExperiment()
935 run = event_meta.getRun()
936 evt = event_meta.getEvent()
937 event_id_str = f"exp={exp}, run={run}, evt={evt}"
938 else:
939 event_id_str = "unknown"
940
941 b2.B2DEBUG(10, f"{'='*60}")
942 b2.B2DEBUG(10, f"DEBUG Event {self.event_count} ({event_id_str})")
943 b2.B2DEBUG(10, f"{'='*60}")
944 b2.B2DEBUG(10, f"Number of candidates: {len(candidates_data)}")
945 b2.B2DEBUG(10, f"Max input_id (best candidate): {max_input_id}")
946
947 b2.B2DEBUG(10, "--- Candidates ---")
948 for i, (input_id, features) in enumerate(candidates_data):
949 b2.B2DEBUG(10, f" Candidate {i}: input_id={input_id}")
950 for key, val in features.items():
951 b2.B2DEBUG(10, f" {key}: {val}")
952
953 b2.B2DEBUG(10, "--- Event Features ---")
954 for key, val in event_features.items():
955 b2.B2DEBUG(10, f" {key}: {val}")
956
957 b2.B2DEBUG(10, "--- Feature Array (non-zero, first 5 blocks) ---")
958 n_blocks = len(self.feature_blocks)
959 for block_idx in range(min(5, n_blocks)):
960 block_name = self.feature_blocks[block_idx][0]
961 start = block_idx * self.n_input_ids
962 end = start + self.n_input_ids
963 block_data = all_features[start:end]
964 non_zero = [(j, v) for j, v in enumerate(block_data) if v != 0]
965 if non_zero:
966 b2.B2DEBUG(10, f" {block_name}: {non_zero[:10]}...")
967
968 # Log last few features (event-level)
969 n_candidate_features = n_blocks * self.n_input_ids
970 event_feat_start = n_candidate_features
971 b2.B2DEBUG(10, f"--- Event-level features (indices {event_feat_start}+) ---")
972 event_feat_names = self.event_features + ['ncandidates/10', 'max_input_id/50', 'scnd_max_input_id/50', '__experiment__/10']
973 for i, name in enumerate(event_feat_names):
974 idx = event_feat_start + i
975 if idx < len(all_features):
976 b2.B2DEBUG(10, f" {name}: {all_features[idx]}")
977
978 b2.B2DEBUG(10, f"--- Total feature array shape: {len(all_features)} ---")
979
980 def terminate(self):
981 """Called at the end of processing."""
982 if self._presel_violation_count > 0:
983 frac = self._presel_violation_count / max(self._presel_total_candidates, 1)
984 _sep = "=" * 70
985 b2.B2WARNING(
986 f"\n{_sep}\n"
987 f"ModeSelector: {self._presel_violation_count}/{self._presel_total_candidates} "
988 f"candidates ({100.0 * frac:.2f}%) did not pass the required preselections "
989 f"(Mbc > {config.PRESELECTION_MBC_MIN}, "
990 f"{config.PRESELECTION_DELTAE_MIN} < deltaE < {config.PRESELECTION_DELTAE_MAX}, "
991 f"cosTBTO < {config.PRESELECTION_COSTBTO_MAX}). "
992 f"Apply these cuts before running ModeSelector to match the training setup.\n{_sep}"
993 )
994
996 unsupported_summary = ", ".join(
997 f"{exp} ({count} event(s))"
998 for exp, count in sorted(self._unsupported_experiment_counts.items())
999 )
1000 b2.B2WARNING(
1001 "ModeSelector: unsupported EventMetaData experiment ids encountered in "
1002 f"{self._unsupported_experiment_event_count} event(s); replaced with default "
1003 f"experiment {config.DEFAULT_EXPERIMENT}. Observed unsupported values: "
1004 f"{unsupported_summary}."
1005 )
1006
1007 if self._inference_event_count > 0:
1008 missing_base_count = self._inference_event_count - self._empty_predicted_sector_count
1009 missing_frac = 0.0
1010 if missing_base_count > 0:
1011 missing_frac = self._missing_top_mode_count / missing_base_count
1012 high_conf_missing_base_count = self._high_conf_event_count - self._empty_predicted_sector_high_conf_count
1013 high_conf_missing_frac = 0.0
1014 if high_conf_missing_base_count > 0:
1015 high_conf_missing_frac = self._missing_top_mode_high_conf_count / high_conf_missing_base_count
1016 b2.B2INFO(
1017 "ModeSelector: missing top mode among event candidates "
1018 "(excluding events with no predicted-sector candidate): "
1019 f"{self._missing_top_mode_count}/{missing_base_count} event(s) "
1020 f"({100.0 * missing_frac:.3f}%), high-confidence (sigProb>{config.HIGH_CONF_SIGPROB_MIN}) "
1021 f"{self._missing_top_mode_high_conf_count}/{high_conf_missing_base_count} event(s) "
1022 f"({100.0 * high_conf_missing_frac:.3f}%)."
1023 )
1024 if high_conf_missing_frac > config.MONITOR_WARN_FRACTION:
1025 b2.B2WARNING(
1026 "ModeSelector: missing top mode high-confidence fraction exceeds threshold "
1027 f"({100.0 * high_conf_missing_frac:.3f}% > {100.0 * config.MONITOR_WARN_FRACTION:.3f}%)."
1028 )
1029
1030 if self._inference_event_count > 0:
1031 fallback_frac = self._empty_predicted_sector_count / self._inference_event_count
1032 high_conf_fallback_frac = 0.0
1033 if self._high_conf_event_count > 0:
1034 high_conf_fallback_frac = self._empty_predicted_sector_high_conf_count / self._high_conf_event_count
1035 b2.B2INFO(
1036 "ModeSelector: predicted sector had no candidate in "
1037 f"{self._empty_predicted_sector_count}/{self._inference_event_count} event(s) "
1038 f"({100.0 * fallback_frac:.3f}%), high-confidence (sigProb>{config.HIGH_CONF_SIGPROB_MIN}) "
1039 f"{self._empty_predicted_sector_high_conf_count}/{self._high_conf_event_count} event(s) "
1040 f"({100.0 * high_conf_fallback_frac:.3f}%); used score fallback."
1041 )
1042
1043 def event(self):
1044 """Called for each event."""
1045 # Check that EventShapeCalculator was run (once per job)
1046 if not self._event_shape_checked:
1047 self._event_shape_checked = True
1048 _es = Belle2.PyStoreObj('EventShapeContainer')
1049 if not _es.isValid():
1050 b2.B2FATAL(
1051 "ModeSelector: EventShapeContainer not found. "
1052 "Run the EventShapeCalculator module before ModeSelector."
1053 )
1054
1055 # Collect all candidates from all particle lists
1056 candidates_data = []
1057 best_bp = None
1058 best_bp_sig = -1
1059 best_b0 = None
1060 best_b0_sig = -1
1061 particle_by_input_id = {} # input_id -> Particle with highest sigProb
1062 _sigprob_by_input_id = {} # input_id -> sigProb for deduplication
1063
1064 for list_name in self.particle_lists:
1065 plist = Belle2.PyStoreObj(list_name)
1066 if not plist.isValid():
1067 continue
1068
1069 for i in range(plist.getListSize()):
1070 particle = plist.obj().getParticle(i)
1071
1072 # Check that DstarVeto was run (once per job, on first particle)
1073 if not self._dstar_veto_checked:
1074 self._dstar_veto_checked = True
1075 if not particle.hasExtraInfo('Dst0_deltaMassDiff'):
1076 b2.B2FATAL(
1077 "ModeSelector: D* veto ExtraInfo (Dst0_deltaMassDiff) not found on B candidates. "
1078 "Run the DstarVeto module before ModeSelector."
1079 )
1080
1081 input_id, features = self._extract_particle_features(particle)
1082 candidates_data.append((input_id, features))
1083
1084 # Check preselection compliance (negligible overhead; values already extracted)
1085 self._presel_total_candidates += 1
1086 _mbc = features.get('Mbc')
1087 _de = features.get('deltaE')
1088 _cos = features.get('cosTBTO')
1089 if (
1090 (_mbc is None or _mbc <= config.PRESELECTION_MBC_MIN)
1091 or (_de is None or not (config.PRESELECTION_DELTAE_MIN < _de < config.PRESELECTION_DELTAE_MAX))
1092 or (_cos is None or _cos >= config.PRESELECTION_COSTBTO_MAX)
1093 ):
1094 self._presel_violation_count += 1
1095
1096 # Track best B+ and B0 by sigProb
1097 sig_prob = features.get('sigProb', -1) or -1
1098 pdg = abs(int(vm.evaluate('PDG', particle)))
1099 if pdg == 521 and sig_prob > best_bp_sig:
1100 best_bp_sig = sig_prob
1101 best_bp = particle
1102 elif pdg == 511 and sig_prob > best_b0_sig:
1103 best_b0_sig = sig_prob
1104 best_b0 = particle
1105
1106 # Track best particle per input_id (same deduplication as _build_feature_array)
1107 if 0 <= input_id < self.n_input_ids:
1108 prev_sig = _sigprob_by_input_id.get(input_id, -2)
1109 if sig_prob > prev_sig:
1110 _sigprob_by_input_id[input_id] = sig_prob
1111 particle_by_input_id[input_id] = particle
1112
1113 if not candidates_data:
1114 return
1115
1116 # Get event-level features
1117 event_features = self._get_event_features()
1118
1119 # Build feature array
1120 all_features, max_input_id, n_candidates = self._build_feature_array(candidates_data, event_features)
1121
1122 # Training mode: expose features + MC truth via EventExtraInfo/ExtraInfo for
1123 # variablesToNtuple to dump in the steering script, skip NN inference.
1124 if self.training_mode:
1125 event_extra_info = Belle2.PyStoreObj('EventExtraInfo')
1126 if not event_extra_info.isValid():
1127 event_extra_info.create()
1128
1129 for i, value in enumerate(all_features):
1130 event_extra_info.addExtraInfo(f'modeSelector_feat_{i:04d}', float(value))
1131
1132 bp_truth = self._extract_mc_truth_scalars(best_bp)
1133 b0_truth = self._extract_mc_truth_scalars(best_b0)
1134
1135 bp_threshold = config.N_BP_MODES * 2
1136 best_bp_iid = -1
1137 best_bp_dp = np.inf
1138 best_b0_iid = -1
1139 best_b0_dp = np.inf
1140
1141 for iid, particle in particle_by_input_id.items():
1142 cand_truth = self._extract_mc_truth_scalars(particle)
1143 is_cont = cand_truth['is_cont']
1144 tag_pdg = cand_truth['tag_pdg']
1145 pdg = cand_truth['pdg']
1146 dp = cand_truth['dp']
1147 dm_id = cand_truth['dm']
1148 if np.isnan(is_cont) or np.isnan(tag_pdg) or np.isnan(pdg) or np.isnan(dp) or np.isnan(dm_id):
1149 continue
1150 if is_cont == 1.0 or not config.truth_tag_matches_pdg(pdg, tag_pdg, dm_id):
1151 continue
1152
1153 if int(iid) < bp_threshold:
1154 if (best_bp_iid < 0) or (dp < best_bp_dp):
1155 best_bp_iid = int(iid)
1156 best_bp_dp = dp
1157 else:
1158 if (best_b0_iid < 0) or (dp < best_b0_dp):
1159 best_b0_iid = int(iid)
1160 best_b0_dp = dp
1161
1162 # Mark the deduplicated best-per-input_id signal candidates so the
1163 # per-candidate ntuple dump in the steering script can pick them out.
1164 for iid, particle in particle_by_input_id.items():
1165 is_signal = vm.evaluate('isSignal', particle)
1166 if not np.isnan(is_signal) and int(is_signal) == 1:
1167 particle.addExtraInfo('modeSelector_trainSigInputId', int(iid))
1168
1169 training_scalars = self._compute_training_event_scalars(
1170 bp_truth, b0_truth, best_bp_iid, best_bp_dp, best_b0_iid, best_b0_dp
1171 )
1172 for name, value in training_scalars.items():
1173 event_extra_info.addExtraInfo(f'modeSelector_tr_{name}', float(value))
1174 return
1175
1176 # --- Inference mode ---
1177
1178 # Select features based on has_inputs (indices used during training)
1179 if self.has_inputs is not None:
1180 features = all_features[self.has_inputs]
1181 else:
1182 features = all_features
1183
1184 # Verify feature size matches model expectation
1185 if len(features) != self.cat_input_size:
1186 b2.B2FATAL(
1187 f"ModeSelector: category model input size mismatch: "
1188 f"got {len(features)} features, model expects {self.cat_input_size}. "
1189 f"The ONNX model and current config must match."
1190 )
1191
1192 if self.skip_nn_evaluation:
1193 cat_output, main_output = self._get_placeholder_outputs(features)
1194 charged_cat = 1.0 if cat_output[1] > cat_output[0] else 0.0
1195 else:
1196 # Run category network
1197 for i, v in enumerate(features.tolist()):
1198 self.cat_dataset.m_input[i] = v
1199 cat_output = list(self.cat_expert.applyMulticlass(self.cat_dataset)[0])
1200
1201 # Determine predicted category (0=B0, 1=B+, 2=continuum)
1202 charged_cat = 1.0 if cat_output[1] > cat_output[0] else 0.0
1203
1204 # Build main network input (features + cat_output + charged_cat)
1205 main_vals = features.tolist() + cat_output + [charged_cat]
1206
1207 if len(main_vals) != self.main_input_size:
1208 b2.B2FATAL(
1209 f"ModeSelector: main model input size mismatch: "
1210 f"got {len(main_vals)} features, model expects {self.main_input_size}. "
1211 f"The ONNX model and current config must match."
1212 )
1213
1214 for i, v in enumerate(main_vals):
1215 self.main_dataset.m_input[i] = v
1216 main_output = list(self.main_expert.applyMulticlass(self.main_dataset)[0])
1217 self._inference_event_count += 1
1218
1219 # Extract scores
1220 # Main network outputs (139 classes after softmax):
1221 # [0..N_INPUT_IDS-1] = P(true mode is that input_id); class index = input_id
1222 # [N_INPUT_IDS+0] = bad_tag
1223 # [N_INPUT_IDS+1] = cross_deltaC1
1224 # [N_INPUT_IDS+2] = continuum
1225 bp_threshold = config.N_BP_MODES * 2
1226 sign = 1.0 if charged_cat else -1.0
1227 predicted_is_bp = bool(charged_cat)
1228 predicted_input_ids = [
1229 iid for iid in particle_by_input_id
1230 if (iid < bp_threshold) == predicted_is_bp
1231 ]
1232 if charged_cat:
1233 predicted_top_iid = int(np.argmax(main_output[:bp_threshold]))
1234 else:
1235 predicted_top_iid = int(np.argmax(main_output[bp_threshold:config.N_INPUT_IDS])) + bp_threshold
1236 has_predicted_candidate = bool(predicted_input_ids)
1237 if has_predicted_candidate and predicted_top_iid not in particle_by_input_id:
1238 self._missing_top_mode_count += 1
1239 if predicted_input_ids:
1240 max_mode_prob = max(float(main_output[iid]) for iid in predicted_input_ids)
1241 else:
1242 # fallback to the sum of predicted-sector mode outputs + bad_tag
1243 # transforming to 2 * (sum_mode_prob - 0.5) if sum_mode_prob > 0.5, else 0.0
1244 # this is to ensure low values for events when the overall predicted-sector confidence is low
1245 bad_tag_prob = float(main_output[config.N_INPUT_IDS])
1246 if charged_cat:
1247 sum_mode_prob = np.sum(main_output[:bp_threshold]) + bad_tag_prob
1248 else:
1249 sum_mode_prob = np.sum(main_output[bp_threshold:config.N_INPUT_IDS]) + bad_tag_prob
1250 max_mode_prob = np.clip(2 * (sum_mode_prob - 0.5), 0.0, None)
1252 bp_score = sign * max_mode_prob
1253 best_sigprob = max(best_bp_sig, best_b0_sig)
1254 is_high_conf = best_sigprob > config.HIGH_CONF_SIGPROB_MIN
1255 if is_high_conf:
1256 self._high_conf_event_count += 1
1257 if has_predicted_candidate and predicted_top_iid not in particle_by_input_id and is_high_conf:
1259 if not predicted_input_ids and is_high_conf:
1261
1262 # Assign modeSelector_eqSigProb per candidate in the predicted sector only.
1263 # Non-predicted sector candidates are left unset (sigProb-based ranking used downstream).
1264 predicted_rank_pairs = []
1265 non_predicted_rank_pairs = []
1266 for input_id, particle in particle_by_input_id.items():
1267 is_bp_sector = input_id < bp_threshold
1268 if bool(charged_cat) == is_bp_sector:
1269 candidate_score = float(main_output[input_id])
1270 particle.addExtraInfo(f'{self.AUXILIARY_OUTPUT_PREFIX}_eqSigProb', candidate_score)
1271 predicted_rank_pairs.append((particle, candidate_score, input_id))
1272 else:
1273 non_predicted_rank_pairs.append((particle, float(_sigprob_by_input_id[input_id]), input_id))
1274
1275 self._assign_rank_extra_info(predicted_rank_pairs, f'{self.AUXILIARY_OUTPUT_PREFIX}_rank')
1276 self._assign_rank_extra_info(non_predicted_rank_pairs, f'{self.AUXILIARY_OUTPUT_PREFIX}_rank')
1277
1278 # Store event-level outputs in EventExtraInfo
1279 event_extra_info = Belle2.PyStoreObj('EventExtraInfo')
1280 if not event_extra_info.isValid():
1281 event_extra_info.create()
1282 event_extra_info.addExtraInfo(self.output_variable, bp_score)
1283 if self.store_fei_calib_weight:
1284 event_extra_info.addExtraInfo(
1285 'modeSelector_feiCalibWeight', self._compute_fei_calib_weight(best_bp, best_b0)
1286 )
1287 event_extra_info.addExtraInfo(f'{self.AUXILIARY_OUTPUT_PREFIX}_catB0', float(cat_output[0]))
1288 event_extra_info.addExtraInfo(f'{self.AUXILIARY_OUTPUT_PREFIX}_catBp', float(cat_output[1]))
1289 event_extra_info.addExtraInfo(f'{self.AUXILIARY_OUTPUT_PREFIX}_catCont', float(cat_output[2]))
1290
1291 # Debug output (after NN inference so we can save outputs too)
1292 if self.debug:
1293 if self.event_count < self.debug_max_events:
1294 self._print_debug_info(candidates_data, event_features, all_features, max_input_id)
1295
1296 self.event_count += 1
Base class for DBObjPtr and DBArray for easier common treatment.
static void initSupportedInterfaces()
Static function which initializes all supported interfaces, has to be called once before getSupported...
Definition Interface.cc:45
static const std::map< std::string, AbstractInterface * > & getSupportedInterfaces()
Returns interfaces supported by the MVA Interface.
Definition Interface.h:53
Wraps the data of a single event into a Dataset.
Definition Dataset.h:135
static Weightfile loadFromFile(const std::string &filename)
Static function which loads a Weightfile from a file.
a (simplified) python wrapper for StoreObjPtr.
Definition PyStoreObj.h:67
_check_same_training(self, cat_training, main_training)
int _empty_predicted_sector_count
Number of events where predicted sector had no candidate in the event.
_select_contract_version(self, cat_version, main_version)
contract_version
Contract version the loaded models were built for.
deltaM_cut
Cut range for D* delta mass difference.
bool _dstar_veto_checked
Flag to check D* veto prerequisite once.
payload_main_model_given
True if the main payload name was given explicitly instead of derived.
_load_payload(self, accessor, payload_name, label)
int _missing_top_mode_count
Number of events where sector top-output input_id has no candidate in the event.
_onnx_interface
ONNX MVA interface used to create the experts.
_main_accessor
Database accessor for the main payload (None if a local file is used)
payload_cat_model_given
True if the category payload name was given explicitly instead of derived.
_compute_training_event_scalars(self, bp_truth, b0_truth, best_bp_iid, best_bp_dp, best_b0_iid, best_b0_dp)
cat_input_size
Category network input size (number of selected features)
_bp_calib_lookup
Per-decay-mode FEI calibration lookup for the B+ sector.
int _presel_violation_count
Number of candidates failing at least one preselection cut.
_main_local_wf
Main weightfile loaded from a local file (None if taken from the database)
int main_input_size
Main network input size (cat features + cat output + charged flag)
dict _unsupported_experiment_counts
Counts of unsupported raw experiment ids seen during inference.
cat_dataset
SingleDataset for category network inference.
_loaded_checksums
Checksums of the currently loaded payloads, used to detect a change between runs.
__init__(self, particle_lists, payload_cat_model=None, payload_main_model=None, output_variable='BplusScore', cat_model_path=None, main_model_path=None, training_mode=False, skip_nn_evaluation=False, store_fei_calib_weight=False, debug=False, debug_max_events=10)
int _missing_top_mode_high_conf_count
Number of high-confidence events where sector top-output input_id has no candidate.
_check_contract(self, weightfile, identifier, label)
_cat_local_wf
Category weightfile loaded from a local file (None if taken from the database)
_print_debug_info(self, candidates_data, event_features, all_features, max_input_id)
training_mode
Training mode (save features + MC truth, skip NN inference)
bool _event_shape_checked
Flag to check event shape prerequisite once.
_assign_rank_extra_info(particle_score_input_id_triples, variable_name)
bool _gen_calib_weight_missing_warned
Flag to emit the gen-calib-weight missing warning at most once.
store_fei_calib_weight
Store modeSelector_feiCalibWeight in EventExtraInfo (MC only)
int _presel_total_candidates
Number of candidates checked for preselection compliance.
_cat_accessor
Database accessor for the category payload (None if a local file is used)
int _empty_predicted_sector_high_conf_count
Number of high-confidence events where fallback was used.
_b0_calib_lookup
Per-decay-mode FEI calibration lookup for the B0 sector.
_build_feature_array(self, candidates_data, event_features)
main_dataset
SingleDataset for main network inference.
int _inference_event_count
Number of inference events processed.
int _unsupported_experiment_event_count
Number of inference events where unsupported experiment ids were replaced.
list feature_blocks
Feature block names and transformations (from config)
has_inputs
Indices of non-zero features kept after sparsity filtering (None in training mode)
n_input_ids
Number of input_id slots (B+ sector + B0 sector, each split by particle/antiparticle)
int _high_conf_event_count
Number of high-confidence inference events (abs(BplusScore) > threshold)
skip_nn_evaluation
Skip NN inference and fill deterministic placeholder outputs.
_bp_calib_rest
FEI calibration fallback factor for the B+ sector.
_b0_calib_rest
FEI calibration fallback factor for the B0 sector.