11ModeSelector Module for event-level B meson classification.
13This module evaluates a two-stage neural network on FEI B meson candidates
14to compute BplusScore, an event-level score based on all candidates.
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
25from modeSelector
import config
26from ROOT
import Belle2
27from variables
import variables
as vm
32 Event-level B meson classifier using neural networks.
34 This module processes FEI B meson candidates and computes a score using
35 information from all candidates in the event.
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
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.
58 See modeSelector.modeSelector() for the full description of each parameter.
61 AUXILIARY_OUTPUT_PREFIX =
'modeSelector'
66 payload_cat_model=None,
67 payload_main_model=None,
68 output_variable='BplusScore',
72 skip_nn_evaluation=False,
73 store_fei_calib_weight=False,
77 """Initialise module parameters and training data buffers."""
80 self.
particle_lists = particle_lists
if isinstance(particle_lists, list)
else [particle_lists]
93 self.
payload_cat_model = payload_cat_model
if payload_cat_model
is not None else config.DEFAULT_CAT_PAYLOAD
95 self.
payload_main_model = payload_main_model
if payload_main_model
is not None else config.DEFAULT_MAIN_PAYLOAD
147 """Load a basf2 MVA weightfile, raises error for unsupported extensions."""
148 if not path.endswith(
'.root'):
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."
157 """Recover the raw feature indices a model was trained on from its weightfile.
159 Returns None for legacy weightfiles that only carry placeholder variable names,
160 in which case the caller falls back to config.HAS_INPUTS.
162 prefix = config.FEATURE_VAR_PREFIX
164 for name
in variables:
166 if not name.startswith(prefix):
169 indices.append(int(name[len(prefix):]))
172 "ModeSelector: " + label +
" model weightfile has a malformed feature "
173 "name '" + name +
"'. Re-export the model with convert_to_onnx.py."
175 return indices
if indices
else None
178 """Check that SUPPORTED_CONTRACT_VERSIONS is consistent with the current contract version.
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.
183 error = config.contract_consistency_error()
185 b2.B2FATAL(
"ModeSelector: " + error)
188 """Check a model's contract version and return it together with the training name.
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.
197 training = str(identifier)
198 if not weightfile.hasContractVersion():
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."
210 version = weightfile.getContractVersion(-1)
213 "ModeSelector: " + label +
" model weightfile has a malformed contract "
214 "version. Re-export it with convert_to_onnx.py."
219 "ModeSelector: " + label +
" model training '" + training
220 +
"' (contract v" + str(version) +
")"
222 return version, training
225 """Stop unless this release can run models built for the given contract version."""
226 if version
in config.SUPPORTED_CONTRACT_VERSIONS:
228 current = config.MODEL_CONTRACT_VERSION
229 if version > current:
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."
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."
243 """Return the contract version to run with, and warn when it is older than the current one.
245 Both networks are evaluated by one code path, so they must share a contract version.
247 if cat_version != main_version:
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."
253 current = config.MODEL_CONTRACT_VERSION
254 if cat_version < current:
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."
266 """Check that the category and main models come from the same training.
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
274 if cat_training != main_training:
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."
283 """Check the model's output size against what the module's interpretation assumes.
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.
289 n_classes = int(options.m_nClasses)
290 if n_classes != expected:
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."
298 """Warn when the configured globaltags also serve payloads for a newer contract version.
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.
310 current = config.MODEL_CONTRACT_VERSION
312 for version
in range(current + 1, current + 1 + probe_range):
313 name = config.payload_names(version)[0]
315 if accessor.getFilename():
316 newest = (version, name)
317 if newest
is not None:
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."
326 """Called at the beginning of processing."""
333 b2.B2INFO(
"ModeSelector: Running in TRAINING mode (saving features, no NN inference)")
337 bp_calib_map = config.get_fei_calibration_map(521)
342 for _dm, _w
in bp_calib_map.items():
343 if 0 <= _dm < config.N_BP_MODES:
346 b0_calib_map = config.get_fei_calibration_map(511)
351 for _dm, _w
in b0_calib_map.items():
352 if 0 <= _dm < config.N_B0_MODES:
362 b2.B2INFO(
"ModeSelector: Running with NN evaluation disabled")
364 f
"ModeSelector: Using placeholder outputs with {self.cat_input_size} selected features"
375 explicit.append(self.
payload_cat_model +
" (default " + config.DEFAULT_CAT_PAYLOAD +
")")
377 explicit.append(self.
payload_main_model +
" (default " + config.DEFAULT_MAIN_PAYLOAD +
")")
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."
422 """Load the models from the conditions database if the payloads changed."""
433 accessor.getChecksum()
if accessor
is not None else None
447 b2.B2INFO(
"ModeSelector: model payloads changed, reloading the models")
454 """Load a basf2 MVA weightfile from the conditions database payload of the current run."""
455 filename = accessor.getFilename()
458 "ModeSelector: " + label +
" model payload '"
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."
467 """Create the experts for a pair of weightfiles and validate them.
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.
482 cat_opts = basf2_mva.GeneralOptions()
483 cat_wf.getOptions(cat_opts)
484 main_opts = basf2_mva.GeneralOptions()
485 main_wf.getOptions(main_opts)
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')
498 if cat_selection
is None or main_selection
is None:
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."
505 selection_source =
'config.HAS_INPUTS'
506 elif cat_selection != main_selection:
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."
513 selection_source =
'the weightfile'
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})")
528 b2.B2INFO(f
"ModeSelector: Using {len(self.has_inputs)} selected features from {selection_source}")
532 Build the feature index mapping from config.
534 The training used sparse matrices with specific non-zero columns.
535 We need to match that structure.
538 self.
feature_blocks = [(name, transform)
for name, _, transform
in config.FEATURE_BLOCKS]
541 for name, var, _
in config.FEATURE_BLOCKS:
543 vm.addAlias(name, var)
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])
555 for rank, (particle, _, _)
in enumerate(ranked_pairs, start=1):
556 particle.addExtraInfo(variable_name, rank)
560 Compute input_id from decay mode ID, B type, and particle/antiparticle sign.
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
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
578 Extract features from a single particle.
581 tuple: (input_id, feature_dict)
587 val = vm.evaluate(name, particle)
593 for key
in [
'Dst0_deltaMassDiff',
'Dstp_deltaMassDiff']:
599 return input_id, features
603 Build the feature array for neural network input.
606 candidates_data: List of (input_id, features) tuples
607 event_features: Dict of event-level features
610 numpy.ndarray: Feature array ready for NN input
614 feature_matrix = np.zeros((n_blocks, self.
n_input_ids), dtype=np.float32)
617 best_by_input_id = {}
618 for input_id, features
in candidates_data:
621 sig_prob = features.get(
'sigProb')
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)
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)
632 feature_matrix[block_idx, input_id] = transform(val)
635 flat_features = feature_matrix.flatten()
638 event_feat_values = [
643 n_candidates = len(candidates_data)
644 event_feat_values.append(n_candidates / 10.0)
649 max_input_id = max(best_by_input_id.keys(), key=
lambda k: best_by_input_id[k][1])
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}
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])
661 scnd_max_input_id = -10
664 scnd_max_input_id = -10
666 event_feat_values.append(max_input_id / 50.0)
667 event_feat_values.append(scnd_max_input_id / 50.0)
671 event_feat_values.append(experiment / 10.0)
674 all_features = np.concatenate([flat_features, np.array(event_feat_values, dtype=np.float32)])
676 return all_features, max_input_id, n_candidates
679 """Return the raw experiment id."""
681 experiment = int(event_meta.getExperiment())
if event_meta.isValid()
else 0
683 if self.
training_mode or experiment
in config.ALLOWED_EXPERIMENTS:
690 return config.DEFAULT_EXPERIMENT
693 """Extract event-level features from EventShapeContainer."""
698 if event_shape.isValid():
699 obj = event_shape.obj()
703 event_features[
'sphericity'] = 1.5 * (obj.getSphericityEigenvalue(1) + obj.getSphericityEigenvalue(2))
705 event_features[
'thrust'] = obj.getThrust()
707 thrust_axis = obj.getThrustAxis()
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
714 event_features[
'aplanarity'] = 1.5 * obj.getSphericityEigenvalue(2)
717 h0 = obj.getFWMoment(0)
718 h2 = obj.getFWMoment(2)
719 event_features[
'foxWolframR2'] = h2 / h0
if h0 != 0
else 0.0
721 event_features[
'harmonicMomentThrust0'] = obj.getHarmonicMomentThrust(0)
723 event_features[
'harmonicMomentThrust1'] = obj.getHarmonicMomentThrust(1)
725 event_features[
'harmonicMomentThrust2'] = obj.getHarmonicMomentThrust(2)
727 return event_features
731 Compute event-level FEI calibration weight from the best candidates.
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).
740 if best_bp
is not None:
741 v = vm.evaluate(
'extraInfo(SignalProbability)', best_bp)
744 if best_b0
is not None:
745 v = vm.evaluate(
'extraInfo(SignalProbability)', best_b0)
749 best = best_bp
if bp_sig >= b0_sig
else best_b0
751 return config.FEI_CALIB_CONT
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
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)
762 if np.isnan(pdg_v)
or np.isnan(dm_v):
763 return config.FEI_CALIB_CONT
765 abs_pdg = abs(int(pdg_v))
766 if abs_pdg
not in (511, 521):
767 return config.FEI_CALIB_CONT
769 calib_map = config.get_fei_calibration_map(abs_pdg)
770 calib_rest = config.get_fei_calibration_rest(abs_pdg)
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:
778 dm = max(0, min(iid // 2, config.N_BP_MODES - 1))
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)
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
795 """Extract compact MC truth scalars used by training output."""
804 'gen_calib_weight': np.nan,
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),
819 name: np.nan
if np.isnan(value)
else value
820 for name, value
in raw.items()
824 self, bp_truth, b0_truth, best_bp_iid, best_bp_dp, best_b0_iid, best_b0_dp
827 Compute the per-event training scalar fields written to EventExtraInfo.
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().
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'])
851 bp_is_best = 1
if bp_sig >= b0_sig
else 0
852 best_sigprob = max(bp_sig, b0_sig)
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
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))
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)
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
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
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))
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)
885 fei_calib_weight = config.FEI_CALIB_CONT
886 bp_threshold = config.N_BP_MODES * 2
888 if bb_mask_ev
and use_reco_ev:
890 dm = min(max(best_bp_sigprob_iid // 2, 0), config.N_BP_MODES - 1)
893 dm = min(max((best_b0_sigprob_iid - bp_threshold) // 2, 0), config.N_B0_MODES - 1)
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:
903 elif abs_pdg_ev == 511:
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,
930 """Print debug information for preprocessing."""
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}"
939 event_id_str =
"unknown"
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}")
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}")
953 b2.B2DEBUG(10,
"--- Event Features ---")
954 for key, val
in event_features.items():
955 b2.B2DEBUG(10, f
" {key}: {val}")
957 b2.B2DEBUG(10,
"--- Feature Array (non-zero, first 5 blocks) ---")
959 for block_idx
in range(min(5, n_blocks)):
963 block_data = all_features[start:end]
964 non_zero = [(j, v)
for j, v
in enumerate(block_data)
if v != 0]
966 b2.B2DEBUG(10, f
" {block_name}: {non_zero[:10]}...")
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]}")
978 b2.B2DEBUG(10, f
"--- Total feature array shape: {len(all_features)} ---")
981 """Called at the end of processing."""
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}"
996 unsupported_summary =
", ".join(
997 f
"{exp} ({count} event(s))"
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}."
1010 if missing_base_count > 0:
1013 high_conf_missing_frac = 0.0
1014 if high_conf_missing_base_count > 0:
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}%)."
1024 if high_conf_missing_frac > config.MONITOR_WARN_FRACTION:
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}%)."
1032 high_conf_fallback_frac = 0.0
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."
1044 """Called for each event."""
1049 if not _es.isValid():
1051 "ModeSelector: EventShapeContainer not found. "
1052 "Run the EventShapeCalculator module before ModeSelector."
1056 candidates_data = []
1061 particle_by_input_id = {}
1062 _sigprob_by_input_id = {}
1066 if not plist.isValid():
1069 for i
in range(plist.getListSize()):
1070 particle = plist.obj().getParticle(i)
1075 if not particle.hasExtraInfo(
'Dst0_deltaMassDiff'):
1077 "ModeSelector: D* veto ExtraInfo (Dst0_deltaMassDiff) not found on B candidates. "
1078 "Run the DstarVeto module before ModeSelector."
1082 candidates_data.append((input_id, features))
1086 _mbc = features.get(
'Mbc')
1087 _de = features.get(
'deltaE')
1088 _cos = features.get(
'cosTBTO')
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)
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
1102 elif pdg == 511
and sig_prob > best_b0_sig:
1103 best_b0_sig = sig_prob
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
1113 if not candidates_data:
1120 all_features, max_input_id, n_candidates = self.
_build_feature_array(candidates_data, event_features)
1126 if not event_extra_info.isValid():
1127 event_extra_info.create()
1129 for i, value
in enumerate(all_features):
1130 event_extra_info.addExtraInfo(f
'modeSelector_feat_{i:04d}', float(value))
1135 bp_threshold = config.N_BP_MODES * 2
1141 for iid, particle
in particle_by_input_id.items():
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):
1150 if is_cont == 1.0
or not config.truth_tag_matches_pdg(pdg, tag_pdg, dm_id):
1153 if int(iid) < bp_threshold:
1154 if (best_bp_iid < 0)
or (dp < best_bp_dp):
1155 best_bp_iid = int(iid)
1158 if (best_b0_iid < 0)
or (dp < best_b0_dp):
1159 best_b0_iid = int(iid)
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))
1170 bp_truth, b0_truth, best_bp_iid, best_bp_dp, best_b0_iid, best_b0_dp
1172 for name, value
in training_scalars.items():
1173 event_extra_info.addExtraInfo(f
'modeSelector_tr_{name}', float(value))
1182 features = all_features
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."
1194 charged_cat = 1.0
if cat_output[1] > cat_output[0]
else 0.0
1197 for i, v
in enumerate(features.tolist()):
1202 charged_cat = 1.0
if cat_output[1] > cat_output[0]
else 0.0
1205 main_vals = features.tolist() + cat_output + [charged_cat]
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."
1214 for i, v
in enumerate(main_vals):
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
1233 predicted_top_iid = int(np.argmax(main_output[:bp_threshold]))
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:
1239 if predicted_input_ids:
1240 max_mode_prob = max(float(main_output[iid])
for iid
in predicted_input_ids)
1245 bad_tag_prob = float(main_output[config.N_INPUT_IDS])
1247 sum_mode_prob = np.sum(main_output[:bp_threshold]) + bad_tag_prob
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
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:
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))
1273 non_predicted_rank_pairs.append((particle, float(_sigprob_by_input_id[input_id]), input_id))
1280 if not event_extra_info.isValid():
1281 event_extra_info.create()
1284 event_extra_info.addExtraInfo(
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]))
1294 self.
_print_debug_info(candidates_data, event_features, all_features, max_input_id)
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...
static const std::map< std::string, AbstractInterface * > & getSupportedInterfaces()
Returns interfaces supported by the MVA Interface.
Wraps the data of a single event into a Dataset.
static Weightfile loadFromFile(const std::string &filename)
Static function which loads a Weightfile from a file.
a (simplified) python wrapper for StoreObjPtr.
_check_same_training(self, cat_training, main_training)
_check_contract_is_self_consistent(self)
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.
debug_max_events
Max events to debug.
deltaM_cut
Cut range for D* delta mass difference.
payload_cat_model
Payload name for category model.
bool _dstar_veto_checked
Flag to check D* veto prerequisite once.
_load_weightfile(self, path, label)
main_model_path
Path to main model.
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.
event_features
Event-level feature names.
_compute_fei_calib_weight(self, best_bp, best_b0)
_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)
output_variable
Output variable name.
_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.
payload_main_model
Payload name for main model.
_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.
_get_inference_experiment_feature_value(self)
_check_contract(self, weightfile, identifier, label)
cat_model_path
Path to category model.
_cat_local_wf
Category weightfile loaded from a local file (None if taken from the database)
_get_event_features(self)
_get_placeholder_outputs(self, features)
_print_debug_info(self, candidates_data, event_features, all_features, max_input_id)
training_mode
Training mode (save features + MC truth, skip NN inference)
particle_lists
Input particle lists.
bool _event_shape_checked
Flag to check event shape prerequisite once.
_assign_rank_extra_info(particle_score_input_id_triples, variable_name)
_extract_mc_truth_scalars(self, particle)
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)
_feature_selection_from_weightfile(self, variables, label)
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)
int event_count
Event counter for debug.
main_dataset
SingleDataset for main network inference.
cat_expert
Category network expert.
int _inference_event_count
Number of inference events processed.
_warn_if_newer_contract_available(self)
_check_output_classes(self, options, expected, label)
int _unsupported_experiment_event_count
Number of inference events where unsupported experiment ids were replaced.
main_expert
Main network expert.
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)
_get_input_id(self, particle)
_setup_models(self, cat_wf, main_wf)
int _high_conf_event_count
Number of high-confidence inference events (abs(BplusScore) > threshold)
_build_feature_indices(self)
skip_nn_evaluation
Skip NN inference and fill deterministic placeholder outputs.
_bp_calib_rest
FEI calibration fallback factor for the B+ sector.
_extract_particle_features(self, particle)
_require_supported_contract(self, version, label)
_b0_calib_rest
FEI calibration fallback factor for the B0 sector.