Belle II Software light-2609-luna
train Namespace Reference

Classes

class  MultiClassNet
 
class  SparseDataset
 

Functions

 _validate_trees (f, input_file)
 
 _load_root_training_file (input_file)
 
 concat_file_data (data_list)
 
 save_npz_shard (path, data)
 
 _load_npz_shard (input_file)
 
 load_training_file (input_file)
 
 _process_one_file (args)
 
 load_and_sample_data (input_files, fraction=1.0, cont_fraction=1.0, sigprob_thresh=config.DEFAULT_FEI_SIGPROB_THRESHOLD, random_state=None, n_workers=None)
 
 train_epoch (model, train_loader, criterion, optimizer, device, disco_lambda=0.0)
 
 evaluate (model, loader, criterion, device)
 
 distance_corr (var_1, var_2, normedweight=None, power=1)
 
 build_mode_labels (mc_truth_cand, sig_truth, charged_cat, is_cont, gen_pdg, delta_p_thresh=0.15)
 
 sparse_collate_fn (batch)
 
 build_category_labels (is_cont, gen_pdg)
 
 compute_event_weights (event_scalars, calib_inputs)
 
 generate_category_outputs (cat_model, features, batch_size, device, use_sparse)
 
 main ()
 

Variables

 _config_path = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), 'config.py')
 Path to config.py in the package directory above, loaded directly instead of through the package.
 
 _spec = importlib.util.spec_from_file_location('modeSelector_config', _config_path)
 Import spec used to load config.py as a standalone module.
 
 config = importlib.util.module_from_spec(_spec)
 ModeSelector configuration, loaded without importing the basf2-dependent package.
 
 MAIN_BG_BAD_TAG = config.N_INPUT_IDS
 Main network label index for bad-tag background.
 
int MAIN_BG_CROSS_DC1 = config.N_INPUT_IDS + 1
 Main network label index for cross-deltaC1 background.
 
int MAIN_BG_CONT = config.N_INPUT_IDS + 2
 Main network label index for continuum background.
 
int MAIN_NUM_LABELS = config.N_INPUT_IDS + 3
 Total number of main network output labels.
 
tuple EVENT_TRUTH_FIELDS
 Per-event truth/label branches written to the 'events' tree, aliased as ms_tr_<name>
 
dict EVENT_TRUTH_DTYPES
 dtype used to cast each ms_tr_<name> branch back from the ROOT double it was stored as
 
 REQUIRED_EVENT_BRANCHES
 Branches required in the 'events' tree of a training input file.
 
tuple SIG_CANDIDATE_BRANCHES
 Branches read from (and required in) the 'sig_candidates' tree.
 
int NPZ_SHARD_VERSION = 1
 Version tag written into every .npz shard, checked on load.
 
tuple SIG_PACKED_VALUE_FIELDS
 Packed ragged per-candidate arrays, all indexed by 'sig_input_ids_offsets'.
 

Detailed Description

Training script for ModeSelector neural networks.

This script trains the two-stage ModeSelector:
1. Category network (B0 vs B+ vs continuum classification)
2. Main network (signal vs background classification using category output)

Inputs are either produceTrainingInputs.py ROOT files or the .npz shards written by
convert_training_inputs.py. Prefer the shards for anything beyond a few hundred ROOT
files: reading the raw ROOT inputs costs ~11 s per file and is paid again on every
invocation, i.e. twice per full training (category, then main).

Usage:
    # Train category network (single file)
    python3 train.py --input modeSelector_training.root --network category --output networks/

    # Train category network (multiple files)
    python3 train.py --input modeSelector_training_*.root --network category --output networks/

    # Train on pre-converted .npz shards (recommended for large samples)
    python3 train.py --input 'converted/*.npz' --network category --use_sparse

    # Train main network (requires trained category network)
    python3 train.py --input modeSelector_training.root --network main \
        --cat_model networks/net_category.pt --output networks/

Function Documentation

◆ _load_npz_shard()

_load_npz_shard ( input_file)
protected
Read a .npz shard written by convert_training_inputs.py back into the same dict
layout _load_root_training_file() returns.

Definition at line 281 of file train.py.

281def _load_npz_shard(input_file):
282 """
283 Read a .npz shard written by convert_training_inputs.py back into the same dict
284 layout _load_root_training_file() returns.
285 """
286 with np.load(input_file) as f:
287 version = int(f['format_version'])
288 if version != NPZ_SHARD_VERSION:
289 raise ValueError(
290 f"{input_file}: .npz shard format version {version}, "
291 f"expected {NPZ_SHARD_VERSION} -- regenerate with convert_training_inputs.py"
292 )
293 data = {
294 'features': sparse.csr_matrix(
295 (f['feat_data'], f['feat_indices'], f['feat_indptr']),
296 shape=tuple(f['feat_shape']),
297 )
298 }
299 for name in EVENT_TRUTH_FIELDS:
300 data[name] = f[name].astype(EVENT_TRUTH_DTYPES[name], copy=False)
301 for name in SIG_PACKED_VALUE_FIELDS:
302 data[name] = f[name]
303 data['sig_input_ids_offsets'] = f['sig_input_ids_offsets']
304
305 return data
306
307

◆ _load_root_training_file()

_load_root_training_file ( input_file)
protected
Read one produceTrainingInputs.py ROOT output file ('events' + 'sig_candidates'
trees, written via variablesToNtuple) into the same key layout the rest of
load_and_sample_data() used to get from the legacy npz format.

Definition at line 131 of file train.py.

131def _load_root_training_file(input_file):
132 """
133 Read one produceTrainingInputs.py ROOT output file ('events' + 'sig_candidates'
134 trees, written via variablesToNtuple) into the same key layout the rest of
135 load_and_sample_data() used to get from the legacy npz format.
136 """
137 with uproot.open(input_file) as f:
138 _validate_trees(f, input_file)
139 events = f['events']
140 # Pass the branch list explicitly. With expressions=None, uproot re-derives the
141 # full branch-name table once per branch; on this tree (1644 ms_feat_* branches
142 # plus truth/index branches, 1667 total) that name resolution costs ~17 s per
143 # file on top of the ~5 s of actual basket I/O.
144 ev_arrays = events.arrays(list(events.keys()), library='np')
145 sig_arrays = f['sig_candidates'].arrays(list(SIG_CANDIDATE_BRANCHES), library='np')
146
147 # The 'events' tree has one row per *processed* event, including events with no
148 # B+/B0 candidate at all -- ModeSelectorModule.event() returns before creating
149 # EventExtraInfo for those, so every ms_tr_*/ms_feat_* value comes back as NaN.
150 # Drop them here to match the legacy pipeline, which never appended such events.
151 valid_events = ~np.isnan(ev_arrays['ms_tr_is_cont'])
152 if not np.all(valid_events):
153 ev_arrays = {key: values[valid_events] for key, values in ev_arrays.items()}
154
155 n_events = len(ev_arrays['__event__'])
156 feat_indices = sorted(int(k[len('ms_feat_'):]) for k in ev_arrays if k.startswith('ms_feat_'))
157 feature_matrix = np.column_stack(
158 [ev_arrays[f'ms_feat_{i:04d}'] for i in feat_indices]
159 ).astype(np.float32)
160 feats = sparse.csr_matrix(feature_matrix)
161
162 data = {'features': feats}
163 for name in EVENT_TRUTH_FIELDS:
164 data[name] = ev_arrays[f'ms_tr_{name}'].astype(EVENT_TRUTH_DTYPES[name])
165
166 # Join sig_candidates rows to their event row via the (experiment, run, event)
167 # triplet written to both trees, mirroring the per-event dedup the module used
168 # to perform before writing the legacy CSR-style sig_* arrays directly.
169 sig_input_id = sig_arrays['ms_sig_input_id']
170 valid = ~np.isnan(sig_input_id)
171 # np.rec, not np.core.records: the latter is private in numpy 2 and raises there.
172 # np.rec is public in both 1.x and 2.x, so this works under the basf2 externals
173 # (numpy 1.26) and in a standalone venv with numpy 2.
174 event_key = np.rec.fromarrays(
175 [ev_arrays['__experiment__'], ev_arrays['__run__'], ev_arrays['__event__']],
176 names='exp,run,evt',
177 )
178 cand_key = np.rec.fromarrays(
179 [sig_arrays['__experiment__'][valid], sig_arrays['__run__'][valid], sig_arrays['__event__'][valid]],
180 names='exp,run,evt',
181 )
182 order = np.argsort(event_key)
183 sorted_event_key = event_key[order]
184 idx_in_sorted = np.searchsorted(sorted_event_key, cand_key)
185 if len(cand_key) and not np.array_equal(sorted_event_key[idx_in_sorted], cand_key):
186 raise ValueError(f"{input_file}: could not join sig_candidates rows to events by (experiment, run, event)")
187 cand_event_idx = order[idx_in_sorted]
188
189 row_order = np.argsort(cand_event_idx, kind='stable')
190 cand_event_idx = cand_event_idx[row_order]
191 data['sig_input_ids_values'] = sig_input_id[valid][row_order].astype(np.int16)
192 data['sig_btag_index_values'] = np.nan_to_num(
193 sig_arrays['ms_sig_btag_index'][valid][row_order], nan=-1
194 ).astype(np.int16)
195 data['sig_delta_p_values'] = sig_arrays['ms_sig_delta_p'][valid][row_order].astype(np.float32)
196 data['sig_sigprob_values'] = sig_arrays['ms_sig_sigprob'][valid][row_order].astype(np.float32)
197
198 counts = np.bincount(cand_event_idx, minlength=n_events)
199 sig_input_ids_offsets = np.zeros(n_events + 1, dtype=np.int32)
200 np.cumsum(counts, out=sig_input_ids_offsets[1:])
201 data['sig_input_ids_offsets'] = sig_input_ids_offsets
202
203 return data
204
205

◆ _process_one_file()

_process_one_file ( args)
protected
Load one input file and apply the per-event preselection and downsampling.

Module-level (rather than a closure) so it can be dispatched to a ProcessPoolExecutor.
Loading is GIL-bound on the numpy/scipy side, so a thread pool scales negatively here:
32 threads measured slower than a single thread on the v7 inputs.

Parameters:
    args (tuple): (input_file, seed, fraction, cont_fraction, sigprob_thresh).

Returns:
    tuple: (result, error). On success `result` is the tuple of sampled arrays and
    `error` is None; on failure `result` is None and `error` is a message naming the
    file, so one unreadable input among thousands does not abort the run.

Definition at line 324 of file train.py.

324def _process_one_file(args):
325 """
326 Load one input file and apply the per-event preselection and downsampling.
327
328 Module-level (rather than a closure) so it can be dispatched to a ProcessPoolExecutor.
329 Loading is GIL-bound on the numpy/scipy side, so a thread pool scales negatively here:
330 32 threads measured slower than a single thread on the v7 inputs.
331
332 Parameters:
333 args (tuple): (input_file, seed, fraction, cont_fraction, sigprob_thresh).
334
335 Returns:
336 tuple: (result, error). On success `result` is the tuple of sampled arrays and
337 `error` is None; on failure `result` is None and `error` is a message naming the
338 file, so one unreadable input among thousands does not abort the run.
339 """
340 input_file, seed, fraction, cont_fraction, sigprob_thresh = args
341
342 try:
343 data = load_training_file(input_file)
344 except Exception as exc:
345 return None, f"{input_file}: failed to read input file ({exc})"
346
347 feats = data['features']
348
349 is_cont = data['is_cont']
350 gen_pdg = data['gen_pdg']
351 bp_is_best = data['bp_is_best']
352 best_sigprob = data['best_sigprob']
353 best_bp_sigprob_iid = data['best_bp_sigprob_iid']
354 best_b0_sigprob_iid = data['best_b0_sigprob_iid']
355
356 bp_tag_is_gen = data['bp_tag_is_gen']
357 b0_tag_is_gen = data['b0_tag_is_gen']
358 bp_gen_decay_mode_id = data['bp_gen_decay_mode_id']
359 b0_gen_decay_mode_id = data['b0_gen_decay_mode_id']
360 bp_gen_fei_calib_weight = data['bp_gen_fei_calib_weight']
361 b0_gen_fei_calib_weight = data['b0_gen_fei_calib_weight']
362 fei_calib_weight = data['fei_calib_weight']
363
364 best_bp_iid = data['best_bp_iid']
365 best_bp_dp = data['best_bp_dp']
366 best_b0_iid = data['best_b0_iid']
367 best_b0_dp = data['best_b0_dp']
368 sig_input_ids_values = data['sig_input_ids_values']
369 sig_input_ids_offsets = data['sig_input_ids_offsets']
370 sig_btag_index_values = data['sig_btag_index_values']
371 sig_delta_p_values = data['sig_delta_p_values']
372 sig_sigprob_values = data['sig_sigprob_values']
373
374 # Keep events where the best candidate passes the sigProb threshold.
375 # Applies to all events including continuum.
376 presel = best_sigprob > sigprob_thresh
377
378 sample_prob = np.full(len(best_sigprob), fraction, dtype=np.float32)
379 sample_prob[is_cont == 1] *= cont_fraction
380 sample_prob = np.minimum(sample_prob, 1.0)
381
382 file_rng = np.random.default_rng(seed)
383 sampled = (file_rng.random(len(best_sigprob)) < sample_prob) & presel
384
385 keep_events = np.flatnonzero(sampled).astype(np.int64)
386
387 def _subset_packed(values, offsets, keep):
388 starts = offsets[keep]
389 ends = offsets[keep + 1]
390 lengths = (ends - starts).astype(np.int64)
391 out_offsets = np.zeros(len(keep) + 1, dtype=np.int64)
392 np.cumsum(lengths, out=out_offsets[1:])
393 total_len = int(out_offsets[-1])
394 if total_len == 0:
395 return np.empty(0, dtype=values.dtype), out_offsets
396 # Gather the kept slices without a Python-level loop over events: build the
397 # source index of every kept element from the per-slice start and a within-slice
398 # ramp (arange minus the offset each slice begins at in the output).
399 idx = np.arange(total_len, dtype=np.int64)
400 slice_of = np.repeat(np.arange(len(keep), dtype=np.int64), lengths)
401 src = idx - out_offsets[slice_of] + starts[slice_of]
402 return values[src], out_offsets
403
404 sig_iid_v_s, sig_off_s = _subset_packed(sig_input_ids_values, sig_input_ids_offsets, keep_events)
405 sig_btag_v_s, _ = _subset_packed(sig_btag_index_values, sig_input_ids_offsets, keep_events)
406 sig_dp_v_s, _ = _subset_packed(sig_delta_p_values, sig_input_ids_offsets, keep_events)
407 sig_sigprob_v_s, _ = _subset_packed(sig_sigprob_values, sig_input_ids_offsets, keep_events)
408
409 result = (feats[sampled],
410 is_cont[sampled], gen_pdg[sampled], bp_is_best[sampled], best_sigprob[sampled],
411 best_bp_sigprob_iid[sampled], best_b0_sigprob_iid[sampled],
412 best_bp_iid[sampled], best_bp_dp[sampled],
413 best_b0_iid[sampled], best_b0_dp[sampled],
414 sig_iid_v_s, sig_off_s, sig_btag_v_s, sig_dp_v_s, sig_sigprob_v_s,
415 bp_tag_is_gen[sampled], b0_tag_is_gen[sampled],
416 bp_gen_decay_mode_id[sampled], b0_gen_decay_mode_id[sampled],
417 bp_gen_fei_calib_weight[sampled], b0_gen_fei_calib_weight[sampled],
418 fei_calib_weight[sampled])
419 return result, None
420
421

◆ _validate_trees()

_validate_trees ( f,
input_file )
protected
Check that an open training input file has the trees and branches
_load_root_training_file() needs, raising ValueError naming what is missing.

Called on the already-open file rather than as a separate preflight pass: at grid
scale (>10k files) a standalone pass costs one extra open per file, which dominated
the old serial preflight loop (~1.1 s/file).

Definition at line 107 of file train.py.

107def _validate_trees(f, input_file):
108 """
109 Check that an open training input file has the trees and branches
110 _load_root_training_file() needs, raising ValueError naming what is missing.
111
112 Called on the already-open file rather than as a separate preflight pass: at grid
113 scale (>10k files) a standalone pass costs one extra open per file, which dominated
114 the old serial preflight loop (~1.1 s/file).
115 """
116 if 'events' not in f or 'sig_candidates' not in f:
117 raise ValueError(f"{input_file}: missing 'events' or 'sig_candidates' tree")
118
119 event_keys = set(f['events'].keys())
120 sig_keys = set(f['sig_candidates'].keys())
121 missing_event = [k for k in REQUIRED_EVENT_BRANCHES if k not in event_keys]
122 if not any(k.startswith('ms_feat_') for k in event_keys):
123 missing_event.append('ms_feat_*')
124 missing_sig = [k for k in SIG_CANDIDATE_BRANCHES if k not in sig_keys]
125 if missing_event or missing_sig:
126 raise ValueError(
127 f"{input_file}: missing branches events={missing_event} sig_candidates={missing_sig}"
128 )
129
130

◆ build_category_labels()

build_category_labels ( is_cont,
gen_pdg )
Build category labels (B0=0, B+=1, continuum=2) from compact MC truth scalars.

Parameters:
    is_cont (ndarray of int8): 1 if continuum event, 0 otherwise.
    gen_pdg (ndarray of int16): mostcommonBTagPDG of the best-sigProb candidate
        (0 for continuum).

Returns:
    ndarray of int64: Category labels (0=B0, 1=B+, 2=continuum).

Definition at line 1168 of file train.py.

1168def build_category_labels(is_cont, gen_pdg):
1169 """
1170 Build category labels (B0=0, B+=1, continuum=2) from compact MC truth scalars.
1171
1172 Parameters:
1173 is_cont (ndarray of int8): 1 if continuum event, 0 otherwise.
1174 gen_pdg (ndarray of int16): mostcommonBTagPDG of the best-sigProb candidate
1175 (0 for continuum).
1176
1177 Returns:
1178 ndarray of int64: Category labels (0=B0, 1=B+, 2=continuum).
1179 """
1180 abs_gen_pdg = np.abs(gen_pdg.astype(np.int32))
1181 known_tag = (abs_gen_pdg == 511) | (abs_gen_pdg == 521)
1182 return np.where(
1183 is_cont | (~known_tag),
1184 2,
1185 np.where(abs_gen_pdg == 521, 1, 0)
1186 ).astype(np.int64)
1187
1188

◆ build_mode_labels()

build_mode_labels ( mc_truth_cand,
sig_truth,
charged_cat,
is_cont,
gen_pdg,
delta_p_thresh = 0.15 )
Build mode-prediction labels (0 to N_INPUT_IDS+2) from MC truth.

Classes:
- 0 to N_INPUT_IDS-1: signal mode; class index equals the candidate's input_id
- N_INPUT_IDS+0 (136): bad_tag
- N_INPUT_IDS+1 (137): cross_deltaC1 (|mostcommonBTagPDG| != |PDG sector|)
- N_INPUT_IDS+2 (138): continuum

Signal label assignment per event:
1. Use charged_cat to select predicted B sector (B+ or B0).
2. If exactly one candidate in that sector has isSignal==1: choose it.
3. If multiple isSignal==1 candidates:
   - If all have same mostcommonBTagIndex: choose largest input_id.
   - If they have different mostcommonBTagIndex: choose smallest deltaP.
     Tie-break by larger sigProb, then larger input_id.
4. If no isSignal==1 candidate in predicted sector, fall back to current
   deltaP-based logic from compact truth scalars.

Parameters:
    mc_truth_cand (tuple): (best_bp_iid, best_bp_dp, best_b0_iid, best_b0_dp) --
        per-event arrays from load_and_sample_data. best_bp_iid/b0_iid are int16
        with sentinel -1; best_bp_dp/b0_dp are float32 with sentinel inf.
    sig_truth (tuple): (sig_input_ids_values, sig_input_ids_offsets,
        sig_btag_index_values, sig_delta_p_values, sig_sigprob_values) --
        packed ragged isSignal==1 candidate metadata.
    charged_cat (ndarray of bool): True if category network predicts B+, False for B0.
    is_cont (ndarray of int8): 1 if continuum event (from best overall candidate), else 0.
    gen_pdg (ndarray of int16): mostcommonBTagPDG of best overall candidate (0 for continuum).
    delta_p_thresh (float): Threshold for mostcommonBTagDeltaP. Default: 0.15.

Returns:
    tuple: (labels, train_selection, stats) where labels is ndarray of int64 with mode
        labels (0 to N_INPUT_IDS+2), train_selection is ndarray of bool for events kept
        for main-network training, and stats is a dict of label assignment counters.

Definition at line 958 of file train.py.

959 delta_p_thresh=0.15):
960 """
961 Build mode-prediction labels (0 to N_INPUT_IDS+2) from MC truth.
962
963 Classes:
964 - 0 to N_INPUT_IDS-1: signal mode; class index equals the candidate's input_id
965 - N_INPUT_IDS+0 (136): bad_tag
966 - N_INPUT_IDS+1 (137): cross_deltaC1 (|mostcommonBTagPDG| != |PDG sector|)
967 - N_INPUT_IDS+2 (138): continuum
968
969 Signal label assignment per event:
970 1. Use charged_cat to select predicted B sector (B+ or B0).
971 2. If exactly one candidate in that sector has isSignal==1: choose it.
972 3. If multiple isSignal==1 candidates:
973 - If all have same mostcommonBTagIndex: choose largest input_id.
974 - If they have different mostcommonBTagIndex: choose smallest deltaP.
975 Tie-break by larger sigProb, then larger input_id.
976 4. If no isSignal==1 candidate in predicted sector, fall back to current
977 deltaP-based logic from compact truth scalars.
978
979 Parameters:
980 mc_truth_cand (tuple): (best_bp_iid, best_bp_dp, best_b0_iid, best_b0_dp) --
981 per-event arrays from load_and_sample_data. best_bp_iid/b0_iid are int16
982 with sentinel -1; best_bp_dp/b0_dp are float32 with sentinel inf.
983 sig_truth (tuple): (sig_input_ids_values, sig_input_ids_offsets,
984 sig_btag_index_values, sig_delta_p_values, sig_sigprob_values) --
985 packed ragged isSignal==1 candidate metadata.
986 charged_cat (ndarray of bool): True if category network predicts B+, False for B0.
987 is_cont (ndarray of int8): 1 if continuum event (from best overall candidate), else 0.
988 gen_pdg (ndarray of int16): mostcommonBTagPDG of best overall candidate (0 for continuum).
989 delta_p_thresh (float): Threshold for mostcommonBTagDeltaP. Default: 0.15.
990
991 Returns:
992 tuple: (labels, train_selection, stats) where labels is ndarray of int64 with mode
993 labels (0 to N_INPUT_IDS+2), train_selection is ndarray of bool for events kept
994 for main-network training, and stats is a dict of label assignment counters.
995 """
996 n_events = len(is_cont)
997 bp_threshold = config.N_BP_MODES * 2
998
999 (sig_input_ids_values, sig_input_ids_offsets,
1000 sig_btag_index_values, sig_delta_p_values, sig_sigprob_values) = sig_truth
1001
1002 if len(sig_input_ids_offsets) != n_events + 1:
1003 raise ValueError(
1004 f"sig_input_ids_offsets length mismatch: got {len(sig_input_ids_offsets)}, expected {n_events + 1}"
1005 )
1006
1007 # --- Per-sector best candidate lookup ---
1008 best_bp_iid, best_bp_dp, best_b0_iid, best_b0_dp = mc_truth_cand
1009 best_iid = np.where(charged_cat, best_bp_iid, best_b0_iid).astype(np.int64)
1010 best_dp = np.where(charged_cat, best_bp_dp, best_b0_dp)
1011
1012 labels = np.full(n_events, MAIN_BG_BAD_TAG, dtype=np.int64)
1013 assigned = np.zeros(n_events, dtype=bool)
1014
1015 stats = {
1016 'single_signal': 0,
1017 'multi_same_btag': 0,
1018 'multi_diff_btag': 0,
1019 'fallback': 0,
1020 }
1021
1022 for evt in range(n_events):
1023 start = int(sig_input_ids_offsets[evt])
1024 end = int(sig_input_ids_offsets[evt + 1])
1025 if start >= end:
1026 continue
1027
1028 evt_iids = sig_input_ids_values[start:end]
1029 evt_btag_idx = sig_btag_index_values[start:end]
1030 evt_dp = sig_delta_p_values[start:end]
1031 evt_sigprob = sig_sigprob_values[start:end]
1032
1033 if charged_cat[evt]:
1034 sector_mask = evt_iids < bp_threshold
1035 else:
1036 sector_mask = evt_iids >= bp_threshold
1037
1038 if not np.any(sector_mask):
1039 continue
1040
1041 sector_iids = evt_iids[sector_mask]
1042 sector_btag_idx = evt_btag_idx[sector_mask]
1043 sector_dp = evt_dp[sector_mask]
1044 sector_sigprob = evt_sigprob[sector_mask]
1045
1046 if len(sector_iids) == 1:
1047 labels[evt] = int(sector_iids[0])
1048 assigned[evt] = True
1049 stats['single_signal'] += 1
1050 continue
1051
1052 if len(np.unique(sector_btag_idx)) == 1:
1053 labels[evt] = int(np.max(sector_iids))
1054 assigned[evt] = True
1055 stats['multi_same_btag'] += 1
1056 continue
1057
1058 best_tuple = None
1059 best_label = None
1060 for iid_val, dp_val, sigprob_val in zip(sector_iids, sector_dp, sector_sigprob):
1061 key = (float(dp_val), -float(sigprob_val), -int(iid_val))
1062 if (best_tuple is None) or (key < best_tuple):
1063 best_tuple = key
1064 best_label = int(iid_val)
1065 labels[evt] = best_label
1066 assigned[evt] = True
1067 stats['multi_diff_btag'] += 1
1068
1069 # --- Fallback labels for events without an isSignal==1 target ---
1070 fallback_mask = ~assigned
1071 stats['fallback'] = int(fallback_mask.sum())
1072 has_target = fallback_mask & (best_iid >= 0) & (best_dp < delta_p_thresh)
1073 labels[has_target] = best_iid[has_target]
1074
1075 no_target = fallback_mask & (~has_target)
1076 if no_target.any():
1077 sector_abs_pdg = np.where(charged_cat, 521, 511)
1078
1079 nt = no_target
1080 labels[nt & (is_cont == 1)] = MAIN_BG_CONT
1081 cross_dc1 = nt & (is_cont == 0) & (np.abs(gen_pdg.astype(np.int32)) != sector_abs_pdg)
1082 labels[cross_dc1] = MAIN_BG_CROSS_DC1
1083 # remaining no_target events keep bad_tag (set as default above)
1084
1085 # Keep all preselected events. Predicted-sector-empty events are retained and
1086 # trained against the fallback background labels above.
1087 train_selection = np.ones(n_events, dtype=bool)
1088
1089 assert labels.min() >= 0 and labels.max() <= MAIN_BG_CONT, (
1090 f"Labels out of range [0, {MAIN_BG_CONT}]: "
1091 f"min={labels.min()}, max={labels.max()}"
1092 )
1093 return labels, train_selection, stats
1094
1095

◆ compute_event_weights()

compute_event_weights ( event_scalars,
calib_inputs )
Return the per-event FEI calibration weights for the training loss.

The weights are computed by ModeSelectorModule when the training inputs are produced
and stored as fei_calib_weight: the calibration weight for the highest-sigProb
candidate's decay mode when that candidate has a truth-compatible tag PDG and its own
DeltaP is below config.DELTA_P_THRESH, the generated decay mode weight otherwise, and
config.FEI_CALIB_CONT for continuum.

Parameters:
    event_scalars (tuple): From load_and_sample_data: (is_cont, gen_pdg, bp_is_best,
        best_sigprob, best_bp_sigprob_iid, best_b0_sigprob_iid).
    calib_inputs (tuple): From load_and_sample_data: (bp_tag_is_gen, b0_tag_is_gen,
        bp_gen_dm_id, b0_gen_dm_id, bp_gen_calib_w, b0_gen_calib_w, stored_fei_calib_w).

Returns:
    ndarray of float32: Per-event FEI calibration weights, shape (n_events,).

Definition at line 1189 of file train.py.

1189def compute_event_weights(event_scalars, calib_inputs):
1190 """
1191 Return the per-event FEI calibration weights for the training loss.
1192
1193 The weights are computed by ModeSelectorModule when the training inputs are produced
1194 and stored as fei_calib_weight: the calibration weight for the highest-sigProb
1195 candidate's decay mode when that candidate has a truth-compatible tag PDG and its own
1196 DeltaP is below config.DELTA_P_THRESH, the generated decay mode weight otherwise, and
1197 config.FEI_CALIB_CONT for continuum.
1198
1199 Parameters:
1200 event_scalars (tuple): From load_and_sample_data: (is_cont, gen_pdg, bp_is_best,
1201 best_sigprob, best_bp_sigprob_iid, best_b0_sigprob_iid).
1202 calib_inputs (tuple): From load_and_sample_data: (bp_tag_is_gen, b0_tag_is_gen,
1203 bp_gen_dm_id, b0_gen_dm_id, bp_gen_calib_w, b0_gen_calib_w, stored_fei_calib_w).
1204
1205 Returns:
1206 ndarray of float32: Per-event FEI calibration weights, shape (n_events,).
1207 """
1208 is_cont = event_scalars[0]
1209 weights = np.asarray(calib_inputs[-1], dtype=np.float32)
1210
1211 n_bad = int((~np.isfinite(weights) | (weights <= 0)).sum())
1212 if n_bad > 0:
1213 raise ValueError(f"compute_event_weights: {n_bad}/{len(weights)} events have a non-finite or non-positive weight")
1214 cont_mask = (is_cont == 1)
1215 n_cont_mismatch = int((weights[cont_mask] != np.float32(config.FEI_CALIB_CONT)).sum())
1216 if n_cont_mismatch > 0:
1217 print(
1218 f" [WARNING] compute_event_weights: {n_cont_mismatch}/{int(cont_mask.sum())} continuum events "
1219 f"do not have the continuum weight {config.FEI_CALIB_CONT}"
1220 )
1221
1222 print(
1223 f" Event weights: mean={weights.mean():.4f}, "
1224 f"min={weights.min():.4f}, max={weights.max():.4f}"
1225 )
1226 print(f" Continuum: {int(cont_mask.sum())} events")
1227 return weights
1228
1229

◆ concat_file_data()

concat_file_data ( data_list)
Concatenate per-file dicts as returned by _load_root_training_file() into one dict
of the same layout.

Parameters:
    data_list (list of dict): Per-file data dicts, in the order to concatenate.

Returns:
    dict: Single dict with the feature matrices vertically stacked, the per-event
    truth arrays concatenated, and the packed ragged sig_* arrays concatenated with
    'sig_input_ids_offsets' rebased onto the combined value arrays.

Definition at line 215 of file train.py.

215def concat_file_data(data_list):
216 """
217 Concatenate per-file dicts as returned by _load_root_training_file() into one dict
218 of the same layout.
219
220 Parameters:
221 data_list (list of dict): Per-file data dicts, in the order to concatenate.
222
223 Returns:
224 dict: Single dict with the feature matrices vertically stacked, the per-event
225 truth arrays concatenated, and the packed ragged sig_* arrays concatenated with
226 'sig_input_ids_offsets' rebased onto the combined value arrays.
227 """
228 if not data_list:
229 raise ValueError("concat_file_data() got no input dicts.")
230 if len(data_list) == 1:
231 return data_list[0]
232
233 out = {'features': sparse.vstack([d['features'] for d in data_list], format='csr')}
234 for name in EVENT_TRUTH_FIELDS:
235 out[name] = np.concatenate([d[name] for d in data_list])
236 for name in SIG_PACKED_VALUE_FIELDS:
237 out[name] = np.concatenate([d[name] for d in data_list])
238
239 # Each file's offsets start at 0; shift every file's tail by the running total so the
240 # combined offsets index the concatenated value arrays.
241 offsets = [np.zeros(1, dtype=np.int32)]
242 shift = 0
243 for d in data_list:
244 file_offsets = d['sig_input_ids_offsets']
245 offsets.append(file_offsets[1:].astype(np.int64) + shift)
246 shift += int(file_offsets[-1])
247 out['sig_input_ids_offsets'] = np.concatenate(offsets).astype(np.int64)
248
249 return out
250
251

◆ distance_corr()

distance_corr ( var_1,
var_2,
normedweight = None,
power = 1 )
Computes distance correlation between var_1 and var_2.

The distance correlation is a measure of dependence between two random variables.
It is zero if and only if the variables are independent.

Parameters:
    var_1 (torch.Tensor): First variable, shape (n,) (e.g., Mbc).
    var_2 (torch.Tensor): Second variable, shape (n,) (e.g., classifier output).
    normedweight (torch.Tensor, optional): Per-example weight, shape (n,), should sum to n.
        If None, uses uniform weights.
    power (int): Exponent for distance correlation (default 1).

Returns:
    torch.Tensor: Distance correlation coefficient (scalar).

Definition at line 892 of file train.py.

892def distance_corr(var_1, var_2, normedweight=None, power=1):
893 """
894 Computes distance correlation between var_1 and var_2.
895
896 The distance correlation is a measure of dependence between two random variables.
897 It is zero if and only if the variables are independent.
898
899 Parameters:
900 var_1 (torch.Tensor): First variable, shape (n,) (e.g., Mbc).
901 var_2 (torch.Tensor): Second variable, shape (n,) (e.g., classifier output).
902 normedweight (torch.Tensor, optional): Per-example weight, shape (n,), should sum to n.
903 If None, uses uniform weights.
904 power (int): Exponent for distance correlation (default 1).
905
906 Returns:
907 torch.Tensor: Distance correlation coefficient (scalar).
908 """
909 n = len(var_1)
910
911 # Handle normedweight - create uniform weights if not provided
912 if normedweight is None:
913 normedweight = torch.ones(n, device=var_1.device, dtype=var_1.dtype)
914 elif isinstance(normedweight, (int, float)):
915 normedweight = normedweight * torch.ones(n, device=var_1.device, dtype=var_1.dtype)
916
917 # Compute pairwise distance matrices
918 amat = torch.abs(var_1.unsqueeze(1) - var_1.unsqueeze(0))
919 bmat = torch.abs(var_2.unsqueeze(1) - var_2.unsqueeze(0))
920
921 # Compute row averages (weighted by normedweight across columns)
922 amatavg = torch.mean(amat * normedweight.unsqueeze(0), dim=1)
923 bmatavg = torch.mean(bmat * normedweight.unsqueeze(0), dim=1)
924
925 # Double centering
926 Amat = (amat
927 - amatavg.unsqueeze(1)
928 - amatavg.unsqueeze(0)
929 + torch.mean(amatavg * normedweight))
930
931 Bmat = (bmat
932 - bmatavg.unsqueeze(1)
933 - bmatavg.unsqueeze(0)
934 + torch.mean(bmatavg * normedweight))
935
936 # Compute final statistics
937 ABavg = torch.mean(Amat * Bmat * normedweight.unsqueeze(0), dim=1)
938 AAavg = torch.mean(Amat * Amat * normedweight.unsqueeze(0), dim=1)
939 BBavg = torch.mean(Bmat * Bmat * normedweight.unsqueeze(0), dim=1)
940
941 # Compute correlation based on power
942 if power == 1:
943 dCorr = (torch.mean(ABavg * normedweight) /
944 torch.sqrt(torch.mean(AAavg * normedweight) *
945 torch.mean(BBavg * normedweight)))
946 elif power == 2:
947 dCorr = (torch.mean(ABavg * normedweight)**2 /
948 (torch.mean(AAavg * normedweight) *
949 torch.mean(BBavg * normedweight)))
950 else:
951 dCorr = ((torch.mean(ABavg * normedweight) /
952 torch.sqrt(torch.mean(AAavg * normedweight) *
953 torch.mean(BBavg * normedweight)))**power)
954
955 return dCorr
956
957

◆ evaluate()

evaluate ( model,
loader,
criterion,
device )
Evaluate model, applying per-sample weights when present.

Expects criterion with reduction='none'. Weights are read from the last
batch element when the batch has more than 2 elements.

Definition at line 865 of file train.py.

865def evaluate(model, loader, criterion, device):
866 """Evaluate model, applying per-sample weights when present.
867
868 Expects criterion with reduction='none'. Weights are read from the last
869 batch element when the batch has more than 2 elements.
870 """
871 model.eval()
872 total_loss = 0
873
874 with torch.no_grad():
875 for batch in loader:
876 data = batch[0].to(device)
877 target = batch[1].to(device)
878 output = model(data)
879 loss_per_sample = criterion(output, target)
880 if len(batch) > 2:
881 sample_weights = batch[-1].to(device)
882 loss = (loss_per_sample * sample_weights).mean()
883 else:
884 loss = loss_per_sample.mean()
885 total_loss += loss.item() * len(data)
886
887 if len(loader.dataset) == 0:
888 return float('nan')
889 return total_loss / len(loader.dataset)
890
891

◆ generate_category_outputs()

generate_category_outputs ( cat_model,
features,
batch_size,
device,
use_sparse )
Generate category softmax outputs for all samples.

Parameters:
    cat_model (nn.Module): Trained category model in eval mode.
    features (scipy.sparse matrix or np.ndarray): Input feature matrix.
    batch_size (int): Batch size for inference.
    device (torch.device): Device to use.
    use_sparse (bool): Whether features is sparse and should be densified batch-wise.

Returns:
    np.ndarray: Category softmax outputs with shape (n_samples, 3).

Definition at line 1230 of file train.py.

1230def generate_category_outputs(cat_model, features, batch_size, device, use_sparse):
1231 """
1232 Generate category softmax outputs for all samples.
1233
1234 Parameters:
1235 cat_model (nn.Module): Trained category model in eval mode.
1236 features (scipy.sparse matrix or np.ndarray): Input feature matrix.
1237 batch_size (int): Batch size for inference.
1238 device (torch.device): Device to use.
1239 use_sparse (bool): Whether features is sparse and should be densified batch-wise.
1240
1241 Returns:
1242 np.ndarray: Category softmax outputs with shape (n_samples, 3).
1243 """
1244 print("\nGenerating category outputs...")
1245 n_samples = features.shape[0]
1246 cat_outputs = []
1247 with torch.no_grad():
1248 for i in tqdm(range(0, n_samples, batch_size), desc=" Processing batches"):
1249 batch_slice = slice(i, i + batch_size)
1250 if use_sparse:
1251 batch_np = features[batch_slice].toarray().astype(np.float32)
1252 else:
1253 batch_np = features[batch_slice]
1254 batch = torch.from_numpy(batch_np).to(device)
1255 cat_out = torch.softmax(cat_model(batch), dim=1)
1256 cat_outputs.append(cat_out.cpu().numpy())
1257 return np.vstack(cat_outputs)
1258
1259

◆ load_and_sample_data()

load_and_sample_data ( input_files,
fraction = 1.0,
cont_fraction = 1.0,
sigprob_thresh = config.DEFAULT_FEI_SIGPROB_THRESHOLD,
random_state = None,
n_workers = None )
Load training data and apply optional global downsampling.

Parameters:
    input_files (str or list of str): Path(s) to modeSelector_training.root file(s)
        produced by produceTrainingInputs.py, or to .npz shards written by
        convert_training_inputs.py. The two may be mixed.
    fraction (float): Uniform BB sampling fraction applied after loading (default 1.0).
    cont_fraction (float): Additional continuum downscale relative to fraction (default 1.0).
    sigprob_thresh (float): Minimum signal probability threshold. Default: 0.001.
    random_state (int): Random seed.
    n_workers (int): Loader processes (default: min(16, cpu_count())).

Returns:
    features (sparse matrix): Sampled and filtered feature matrix (only has_inputs columns).
    event_scalars (tuple): (is_cont, gen_pdg, bp_is_best, best_sigprob,
        best_bp_sigprob_iid, best_b0_sigprob_iid) -- per-event compact MC truth scalars;
        int8/int16/float32 arrays of shape (n_events,).
    has_inputs (list of int): Selected feature indices (non-zero features, excluding Mbc).
    mc_truth_cand (tuple): (best_bp_iid, best_bp_dp, best_b0_iid, best_b0_dp) --
        per-event arrays of shape (n_events,). best_bp_iid/b0_iid are int16 with
        sentinel -1 when no qualifying candidate exists; best_bp_dp/b0_dp are float32
        with sentinel inf. Pre-filtered: truth-compatible tag PDG and is_cont != 1.
    sig_truth (tuple): (sig_input_ids_values, sig_input_ids_offsets,
        sig_btag_index_values, sig_delta_p_values, sig_sigprob_values) --
        packed ragged arrays of per-event isSignal==1 candidates on deduplicated
        input_ids. Event i slice is values[offsets[i]:offsets[i+1]] and aligned
        across all *_values arrays.
    calib_inputs (tuple): (bp_tag_is_gen, b0_tag_is_gen, bp_gen_dm_id, b0_gen_dm_id,
        bp_gen_calib_w, b0_gen_calib_w, stored_fei_calib_w) -- per-event arrays for FEI
        calibration weight computation. bp/b0_gen_dm_id are int16 with sentinel -1
        (missing) or 999 (rest calibration). stored_fei_calib_w is the event-level
        weight computed by ModeSelectorModule and used as the training weight in
        compute_event_weights. bp/b0_gen_calib_w are float32 stored calibration weights
        from generatedDecayWeights.

Definition at line 422 of file train.py.

424 random_state=None, n_workers=None):
425 """
426 Load training data and apply optional global downsampling.
427
428 Parameters:
429 input_files (str or list of str): Path(s) to modeSelector_training.root file(s)
430 produced by produceTrainingInputs.py, or to .npz shards written by
431 convert_training_inputs.py. The two may be mixed.
432 fraction (float): Uniform BB sampling fraction applied after loading (default 1.0).
433 cont_fraction (float): Additional continuum downscale relative to fraction (default 1.0).
434 sigprob_thresh (float): Minimum signal probability threshold. Default: 0.001.
435 random_state (int): Random seed.
436 n_workers (int): Loader processes (default: min(16, cpu_count())).
437
438 Returns:
439 features (sparse matrix): Sampled and filtered feature matrix (only has_inputs columns).
440 event_scalars (tuple): (is_cont, gen_pdg, bp_is_best, best_sigprob,
441 best_bp_sigprob_iid, best_b0_sigprob_iid) -- per-event compact MC truth scalars;
442 int8/int16/float32 arrays of shape (n_events,).
443 has_inputs (list of int): Selected feature indices (non-zero features, excluding Mbc).
444 mc_truth_cand (tuple): (best_bp_iid, best_bp_dp, best_b0_iid, best_b0_dp) --
445 per-event arrays of shape (n_events,). best_bp_iid/b0_iid are int16 with
446 sentinel -1 when no qualifying candidate exists; best_bp_dp/b0_dp are float32
447 with sentinel inf. Pre-filtered: truth-compatible tag PDG and is_cont != 1.
448 sig_truth (tuple): (sig_input_ids_values, sig_input_ids_offsets,
449 sig_btag_index_values, sig_delta_p_values, sig_sigprob_values) --
450 packed ragged arrays of per-event isSignal==1 candidates on deduplicated
451 input_ids. Event i slice is values[offsets[i]:offsets[i+1]] and aligned
452 across all *_values arrays.
453 calib_inputs (tuple): (bp_tag_is_gen, b0_tag_is_gen, bp_gen_dm_id, b0_gen_dm_id,
454 bp_gen_calib_w, b0_gen_calib_w, stored_fei_calib_w) -- per-event arrays for FEI
455 calibration weight computation. bp/b0_gen_dm_id are int16 with sentinel -1
456 (missing) or 999 (rest calibration). stored_fei_calib_w is the event-level
457 weight computed by ModeSelectorModule and used as the training weight in
458 compute_event_weights. bp/b0_gen_calib_w are float32 stored calibration weights
459 from generatedDecayWeights.
460 """
461
462 if isinstance(input_files, str):
463 input_files = [input_files]
464
465 # Expand any glob patterns (e.g. grid downloads: many small *.root files).
466 # recursive=True is required for ** to descend into subdirectories; without it
467 # glob silently returns no matches and the pattern would fall through as a
468 # literal (nonexistent) file path below.
469 expanded = []
470 for f in input_files:
471 matches = glob.glob(f, recursive=True)
472 expanded.extend(matches if matches else [f])
473 input_files = expanded
474
475 if not input_files:
476 raise ValueError("No input files found.")
477
478 parent_dirs = sorted({os.path.dirname(os.path.abspath(f)) for f in input_files})
479 print("Resolved input parent directories:")
480 for d in parent_dirs:
481 print(f" {d}")
482
483 if fraction > 1.0:
484 raise ValueError(f"fraction must be <= 1.0, got {fraction}.")
485
486 if cont_fraction > 1.0:
487 raise ValueError(f"cont_fraction must be <= 1.0, got {cont_fraction}.")
488
489 rng = np.random.default_rng(random_state)
490 # Pre-generate per-file seeds so parallel workers are independent
491 seeds = rng.integers(0, 2**31, size=len(input_files))
492
493 features_list = []
494 is_cont_list, gen_pdg_list, bp_is_best_list, best_sigprob_list = [], [], [], []
495 best_bp_sigprob_iid_list, best_b0_sigprob_iid_list = [], []
496 best_bp_iid_list, best_bp_dp_list = [], []
497 best_b0_iid_list, best_b0_dp_list = [], []
498 sig_input_ids_values_list, sig_input_ids_offsets_list = [], []
499 sig_btag_index_values_list, sig_delta_p_values_list, sig_sigprob_values_list = [], [], []
500 bp_tag_is_gen_list, b0_tag_is_gen_list = [], []
501 bp_gen_dm_id_list, b0_gen_dm_id_list = [], []
502 bp_gen_calib_w_list, b0_gen_calib_w_list = [], []
503 stored_fei_calib_w_list = []
504
505 if n_workers is None:
506 n_workers = min(16, os.cpu_count() or 1)
507 n_workers = max(1, min(n_workers, len(input_files)))
508
509 # Process pool, not threads: loading is dominated by numpy/scipy work that holds the
510 # GIL, so threads scale negatively here (32 threads measured ~2x slower per file than
511 # a single thread on the v7 inputs). Results stream in submission order and are
512 # appended straight to the per-field lists, so only one result is held beyond the
513 # lists themselves.
514 load_errors = []
515 with ProcessPoolExecutor(max_workers=n_workers) as executor:
516 results = executor.map(
517 _process_one_file,
518 [(f, s, fraction, cont_fraction, sigprob_thresh) for f, s in zip(input_files, seeds)],
519 )
520 results = tqdm(results, total=len(input_files), desc="Loading files")
521
522 for result, error in results:
523 if error is not None:
524 load_errors.append(error)
525 continue
526 (feats,
527 is_cont_r, gen_pdg_r, bp_is_best_r, best_sigprob_r,
528 best_bp_sigprob_iid_r, best_b0_sigprob_iid_r,
529 bp_iid, bp_dp, b0_iid, b0_dp,
530 sig_iid_v, sig_off, sig_btag_v, sig_dp_v, sig_sigprob_v,
531 bp_tg, b0_tg, bp_gd, b0_gd, bp_gcw, b0_gcw, stored_fcw) = result
532 features_list.append(feats)
533 is_cont_list.append(is_cont_r)
534 gen_pdg_list.append(gen_pdg_r)
535 bp_is_best_list.append(bp_is_best_r)
536 best_sigprob_list.append(best_sigprob_r)
537 best_bp_sigprob_iid_list.append(best_bp_sigprob_iid_r)
538 best_b0_sigprob_iid_list.append(best_b0_sigprob_iid_r)
539 best_bp_iid_list.append(bp_iid)
540 best_bp_dp_list.append(bp_dp)
541 best_b0_iid_list.append(b0_iid)
542 best_b0_dp_list.append(b0_dp)
543 sig_input_ids_values_list.append(sig_iid_v)
544 sig_input_ids_offsets_list.append(sig_off)
545 sig_btag_index_values_list.append(sig_btag_v)
546 sig_delta_p_values_list.append(sig_dp_v)
547 sig_sigprob_values_list.append(sig_sigprob_v)
548 bp_tag_is_gen_list.append(bp_tg)
549 b0_tag_is_gen_list.append(b0_tg)
550 bp_gen_dm_id_list.append(bp_gd)
551 b0_gen_dm_id_list.append(b0_gd)
552 bp_gen_calib_w_list.append(bp_gcw)
553 b0_gen_calib_w_list.append(b0_gcw)
554 stored_fei_calib_w_list.append(stored_fcw)
555
556 # Files that fail to open/read (e.g. corrupted or incomplete grid downloads) or are
557 # missing required branches are skipped with a warning rather than aborting the whole
558 # run -- expected at grid scale, where a handful of bad job outputs among thousands is
559 # common.
560 if load_errors:
561 print(f"\nWARNING: skipped {len(load_errors)}/{len(input_files)} unreadable input file(s):")
562 for err in load_errors[:20]:
563 print(f" {err}")
564 if len(load_errors) > 20:
565 print(f" ... and {len(load_errors) - 20} more files")
566
567 if not features_list:
568 raise ValueError("All input files failed to load. Fix input files before training.")
569
570 # Concatenate across all files
571 features = sparse.vstack(features_list, format='csr')
572 is_cont = np.concatenate(is_cont_list)
573 gen_pdg = np.concatenate(gen_pdg_list)
574 bp_is_best = np.concatenate(bp_is_best_list)
575 best_sigprob = np.concatenate(best_sigprob_list)
576 best_bp_sigprob_iid = np.concatenate(best_bp_sigprob_iid_list)
577 best_b0_sigprob_iid = np.concatenate(best_b0_sigprob_iid_list)
578
579 print(f"\nTotal after concatenation: {features.shape[0]} events")
580
581 # Compute has_inputs dynamically and verify against config.HAS_INPUTS
582 print("\nVerifying HAS_INPUTS from feature sparsity...")
583
584 n_total = features.shape[1]
585
586 # Stage 1: All-zero columns
587 nonzero_cols = set(np.unique(features.nonzero()[1]))
588 all_zero_cols = set(range(n_total)) - nonzero_cols
589
590 # Stage 2: Exclude blocks that should not be active network inputs.
591 # Mbc is excluded to avoid output correlation. With skipTreeFit=True,
592 # Dst0_chiProb and Dstp_chiProb do not add independent information beyond
593 # Bdaughter_chiProb, so they are excluded as well.
594 n_input_ids = config.N_INPUT_IDS # 136
595 excluded_block_ranges = {
596 'Dst0_chiProb': (n_input_ids * 7, n_input_ids * 8),
597 'Dstp_chiProb': (n_input_ids * 8, n_input_ids * 9),
598 'Mbc': (n_input_ids * 10, n_input_ids * 11),
599 }
600 excluded_cols = set()
601 for start, end in excluded_block_ranges.values():
602 excluded_cols.update(range(start, end))
603
604 computed_has_inputs = sorted(set(range(n_total)) - (all_zero_cols | excluded_cols))
605 n_computed_remove = n_total - len(computed_has_inputs)
606 print(f" All-zero columns: {len(all_zero_cols)}")
607 for block_name, (start, end) in excluded_block_ranges.items():
608 print(f" {block_name} block [{start}, {end}): {end - start} columns")
609 print(f" Total to remove: {n_computed_remove}")
610
611 # Verify against hardcoded list in config
612 if computed_has_inputs != config.HAS_INPUTS:
613 print(" WARNING: HAS_INPUTS mismatch!")
614 print(f" Computed from data: {len(computed_has_inputs)} kept ({n_computed_remove} removed)")
615 print(f" config.HAS_INPUTS: {len(config.HAS_INPUTS)} kept ({n_total - len(config.HAS_INPUTS)} removed)")
616 print(" Data may not cover all input_ids.")
617 out_path = "has_inputs_recomputed.txt"
618 n_kept = len(computed_has_inputs)
619 n_removed = n_total - n_kept
620 with open(out_path, "w") as f:
621 f.write("HAS_INPUTS = [\n")
622 row = []
623 for idx in computed_has_inputs:
624 row.append(idx)
625 if len(row) == 17:
626 f.write(" " + ", ".join(str(x) for x in row) + ",\n")
627 row = []
628 if row:
629 f.write(" " + ", ".join(str(x) for x in row) + ",\n")
630 f.write(f"] # {n_kept} indices kept ({n_total} - {n_removed} removed)\n")
631 print(f" Recomputed list written to {out_path}")
632 print(" Using recomputed HAS_INPUTS for this run.")
633 has_inputs = computed_has_inputs
634 else:
635 print(f" HAS_INPUTS verified OK ({len(config.HAS_INPUTS)} kept, "
636 f"{n_total - len(config.HAS_INPUTS)} removed)")
637 has_inputs = list(config.HAS_INPUTS)
638
639 # Apply feature selection
640 features = features[:, has_inputs]
641
642 mc_truth_cand = (
643 np.concatenate(best_bp_iid_list),
644 np.concatenate(best_bp_dp_list),
645 np.concatenate(best_b0_iid_list),
646 np.concatenate(best_b0_dp_list),
647 )
648 if sig_input_ids_offsets_list:
649 # Rebase each file's offsets (which start at 0) onto the concatenated value
650 # arrays. Kept in numpy rather than building a Python list of ~1e8 ints, which
651 # costs several GB of boxed integers on the full v7 sample.
652 shifted_offsets = [np.zeros(1, dtype=np.int64)]
653 shift = 0
654 for file_offsets in sig_input_ids_offsets_list:
655 shifted_offsets.append(file_offsets[1:].astype(np.int64) + shift)
656 shift += int(file_offsets[-1])
657 sig_input_ids_offsets = np.concatenate(shifted_offsets)
658 else:
659 sig_input_ids_offsets = np.zeros(features.shape[0] + 1, dtype=np.int64)
660
661 if sig_input_ids_values_list:
662 sig_input_ids_values = np.concatenate(sig_input_ids_values_list).astype(np.int16, copy=False)
663 sig_btag_index_values = np.concatenate(sig_btag_index_values_list).astype(np.int16, copy=False)
664 sig_delta_p_values = np.concatenate(sig_delta_p_values_list).astype(np.float32, copy=False)
665 sig_sigprob_values = np.concatenate(sig_sigprob_values_list).astype(np.float32, copy=False)
666 else:
667 sig_input_ids_values = np.empty(0, dtype=np.int16)
668 sig_btag_index_values = np.empty(0, dtype=np.int16)
669 sig_delta_p_values = np.empty(0, dtype=np.float32)
670 sig_sigprob_values = np.empty(0, dtype=np.float32)
671
672 if len(sig_input_ids_offsets) != features.shape[0] + 1:
673 raise ValueError(
674 f"sig_input_ids_offsets length mismatch: got {len(sig_input_ids_offsets)}, "
675 f"expected {features.shape[0] + 1}"
676 )
677 expected_n = int(sig_input_ids_offsets[-1])
678 for name, arr in (
679 ('sig_input_ids_values', sig_input_ids_values),
680 ('sig_btag_index_values', sig_btag_index_values),
681 ('sig_delta_p_values', sig_delta_p_values),
682 ('sig_sigprob_values', sig_sigprob_values),
683 ):
684 if len(arr) != expected_n:
685 raise ValueError(
686 f"{name} length mismatch: got {len(arr)}, expected {expected_n} from offsets"
687 )
688
689 sig_truth = (
690 sig_input_ids_values,
691 sig_input_ids_offsets,
692 sig_btag_index_values,
693 sig_delta_p_values,
694 sig_sigprob_values,
695 )
696 event_scalars = (is_cont, gen_pdg, bp_is_best, best_sigprob,
697 best_bp_sigprob_iid, best_b0_sigprob_iid)
698 calib_inputs = (
699 np.concatenate(bp_tag_is_gen_list),
700 np.concatenate(b0_tag_is_gen_list),
701 np.concatenate(bp_gen_dm_id_list),
702 np.concatenate(b0_gen_dm_id_list),
703 np.concatenate(bp_gen_calib_w_list),
704 np.concatenate(b0_gen_calib_w_list),
705 np.concatenate(stored_fei_calib_w_list),
706 )
707 return (features, event_scalars, has_inputs, mc_truth_cand, sig_truth, calib_inputs)
708
709

◆ load_training_file()

load_training_file ( input_file)
Load one training input, dispatching on file extension.

Parameters:
    input_file (str): Either a produceTrainingInputs.py ROOT file or a .npz shard
        written by convert_training_inputs.py.

Returns:
    dict: Per-file data dict (see _load_root_training_file()).

Definition at line 308 of file train.py.

308def load_training_file(input_file):
309 """
310 Load one training input, dispatching on file extension.
311
312 Parameters:
313 input_file (str): Either a produceTrainingInputs.py ROOT file or a .npz shard
314 written by convert_training_inputs.py.
315
316 Returns:
317 dict: Per-file data dict (see _load_root_training_file()).
318 """
319 if input_file.endswith('.npz'):
320 return _load_npz_shard(input_file)
321 return _load_root_training_file(input_file)
322
323

◆ main()

main ( )
Parse command-line arguments and run the requested training.

Definition at line 1260 of file train.py.

1260def main():
1261 """Parse command-line arguments and run the requested training."""
1262 parser = argparse.ArgumentParser(description='Train ModeSelector networks')
1263 parser.add_argument('--input', required=True, nargs='+',
1264 help='One or more modeSelector_training.root paths, or .npz shards from '
1265 'convert_training_inputs.py (shell glob or space-separated list)')
1266 parser.add_argument('--network', choices=['category', 'main'], required=True,
1267 help='Which network to train')
1268 parser.add_argument('--cat_model', help='Trained category model (required for main network)')
1269 parser.add_argument('--output', default='networks/', help='Output directory for trained models')
1270 parser.add_argument('--fraction', type=float, default=1.0,
1271 help='Optional uniform BB downsampling fraction after loading inputs')
1272 parser.add_argument('--cont_fraction', type=float, default=1.0,
1273 help='Additional continuum downscale relative to --fraction at training time')
1274 parser.add_argument('--batch_size', type=int, default=None,
1275 help='Batch size (default: 16384 for category network, 32768 for main network)')
1276 parser.add_argument('--num_workers', type=int, default=None,
1277 help='Number of worker processes for input loading and the DataLoader '
1278 '(default: auto = min(8, max(1, cpu_count//2)))')
1279 parser.add_argument('--epochs', type=int, default=50, help='Number of epochs')
1280 parser.add_argument('--lr', type=float, default=5e-4, help='Initial learning rate')
1281 parser.add_argument('--lr_schedule', choices=['constant', 'cosine'], default='cosine',
1282 help='Learning rate schedule (default: cosine)')
1283 parser.add_argument('--eta_min', type=float, default=1e-5,
1284 help='Minimum learning rate for cosine schedule (default 1e-5)')
1285 parser.add_argument('--weight_decay', type=float, default=2e-4, help='Weight decay for AdamW')
1286 parser.add_argument('--val_split', type=float, default=0.3, help='Validation split')
1287 parser.add_argument('--seed', type=int, default=42, help='Random seed')
1288 parser.add_argument('--disco_lambda', type=float, default=0.0,
1289 help='Distance correlation penalty coefficient (0=disabled)')
1290 parser.add_argument('--label_smoothing', type=float, default=0,
1291 help='Label smoothing for CrossEntropyLoss (0=disabled, default)')
1292 parser.add_argument('--use_sparse', action='store_true',
1293 help='Use sparse data loading (memory-efficient but slower)')
1294
1295 args = parser.parse_args()
1296
1297 if args.batch_size is None:
1298 args.batch_size = 2**14 if args.network == 'category' else 2**15
1299 print(f"Batch size: {args.batch_size}")
1300 if args.num_workers is None:
1301 cpu_count = os.cpu_count() or 1
1302 resolved_num_workers = min(8, max(1, cpu_count // 2))
1303 else:
1304 resolved_num_workers = args.num_workers
1305 print(f"DataLoader workers: {resolved_num_workers}")
1306
1307 # Set random seeds
1308 np.random.seed(args.seed)
1309 torch.manual_seed(args.seed)
1310
1311 # Device
1312 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
1313 print(f"Using device: {device}")
1314
1315 # Create output directory
1316 os.makedirs(args.output, exist_ok=True)
1317
1318 # Load input data and apply optional post-production downsampling
1319 print("\n" + "=" * 60)
1320 print("Loading data")
1321 print("=" * 60)
1322 features, event_scalars, has_inputs, mc_truth_cand, sig_truth, calib_inputs = load_and_sample_data(
1323 args.input,
1324 fraction=args.fraction,
1325 cont_fraction=args.cont_fraction,
1326 random_state=args.seed,
1327 n_workers=resolved_num_workers,
1328 )
1329 is_cont, gen_pdg, bp_is_best, best_sigprob, best_bp_sigprob_iid, best_b0_sigprob_iid = event_scalars
1330
1331 print("\nComputing event weights...")
1332 event_weights = compute_event_weights(event_scalars, calib_inputs)
1333
1334 print(f"\nFeature matrix shape: {features.shape}")
1335 print(f" Sparse matrix memory: {features.data.nbytes / 1024**2:.1f} MB")
1336 print(f" Selected features (has_inputs): {len(has_inputs)}")
1337 sparse_extra_features = None
1338
1339 # Build labels
1340 print("\nBuilding labels...")
1341 if args.network == 'category':
1342 labels = build_category_labels(is_cont, gen_pdg)
1343 num_labels = 3
1344 print(" Category distribution:")
1345 print(f" B0: {(labels == 0).sum()} ({(labels == 0).mean() * 100:.1f}%)")
1346 print(f" B+: {(labels == 1).sum()} ({(labels == 1).mean() * 100:.1f}%)")
1347 print(f" Continuum: {(labels == 2).sum()} ({(labels == 2).mean() * 100:.1f}%)")
1348
1349 input_size = features.shape[1]
1350
1351 if not args.use_sparse:
1352 # Convert to dense for training
1353 print("\nConverting to dense arrays...")
1354 features_dense = features.toarray().astype(np.float32)
1355 else:
1356 print("\nUsing sparse data loading (memory-efficient)...")
1357 features_dense = None # Will use SparseDataset below
1358
1359 else: # main network
1360 if not args.cat_model:
1361 raise ValueError("--cat_model required for main network training")
1362
1363 # Load trained category model
1364 print(f"\nLoading category model from {args.cat_model}...")
1365 cat_checkpoint = torch.load(args.cat_model)
1366 cat_model = MultiClassNet(
1367 input_size=features.shape[1],
1368 num_labels=3
1369 )
1370 cat_model.load_state_dict(cat_checkpoint['model_state_dict'])
1371 cat_model = cat_model.to(device)
1372 cat_model.eval()
1373 print(f" Category model loaded (epoch {cat_checkpoint['epoch']+1})")
1374
1375 if not args.use_sparse:
1376 print("\nConverting to dense arrays...")
1377 features_dense = features.toarray().astype(np.float32)
1378 cat_outputs = generate_category_outputs(
1379 cat_model, features_dense, args.batch_size, device, use_sparse=False
1380 )
1381 else:
1382 print("\nUsing sparse data loading (memory-efficient)...")
1383 features_dense = None
1384 cat_outputs = generate_category_outputs(
1385 cat_model, features, args.batch_size, device, use_sparse=True
1386 )
1387
1388 print(f" Category outputs shape: {cat_outputs.shape}")
1389 print(" Category predictions:")
1390 cat_preds = np.argmax(cat_outputs, axis=1)
1391 print(f" B0: {(cat_preds == 0).sum()} ({(cat_preds == 0).mean() * 100:.1f}%)")
1392 print(f" B+: {(cat_preds == 1).sum()} ({(cat_preds == 1).mean() * 100:.1f}%)")
1393 print(f" Continuum: {(cat_preds == 2).sum()} ({(cat_preds == 2).mean() * 100:.1f}%)")
1394
1395 # Compute charged_cat flag: 1 if B+ is predicted over B0 (matches ModeSelectorModule inference)
1396 charged_cat_bool = cat_outputs[:, 1] > cat_outputs[:, 0]
1397 charged_cat = charged_cat_bool.astype(np.float32)
1398
1399 # Add category outputs + charged_cat flag to features
1400 print("\nAppending category outputs to input features...")
1401 cat_augments = np.hstack([cat_outputs, charged_cat.reshape(-1, 1)]).astype(np.float32)
1402 input_size = features.shape[1] + cat_augments.shape[1]
1403 print(f" New feature size: {input_size}")
1404 if features_dense is not None:
1405 features_dense = np.hstack([features_dense, cat_augments]).astype(np.float32)
1406
1407 # Build mode-prediction labels using category network prediction
1408 if mc_truth_cand is None or sig_truth is None:
1409 raise ValueError(
1410 "Required main-network truth arrays not found in training data. "
1411 "Re-collect training data with produceTrainingInputs.py."
1412 )
1413 labels, train_selection, label_stats = build_mode_labels(
1414 mc_truth_cand, sig_truth, charged_cat_bool, is_cont, gen_pdg,
1415 )
1416 num_labels = MAIN_NUM_LABELS
1417 n_signal = int((labels < config.N_INPUT_IDS).sum())
1418 bg_bad = int((labels == MAIN_BG_BAD_TAG).sum())
1419 bg_dc1 = int((labels == MAIN_BG_CROSS_DC1).sum())
1420 bg_cont = int((labels == MAIN_BG_CONT).sum())
1421 n_tot = len(labels)
1422 print(" Main network label distribution (before train_selection):")
1423 print(f" signal modes (0-{config.N_INPUT_IDS - 1}): {n_signal} ({n_signal / n_tot * 100:.1f}%)")
1424 print(f" bad_tag ({MAIN_BG_BAD_TAG}): {bg_bad} ({bg_bad / n_tot * 100:.1f}%)")
1425 print(f" cross_deltaC1 ({MAIN_BG_CROSS_DC1}): {bg_dc1} ({bg_dc1 / n_tot * 100:.1f}%)")
1426 print(f" continuum ({MAIN_BG_CONT}): {bg_cont} ({bg_cont / n_tot * 100:.1f}%)")
1427 print(" Label assignment branches:")
1428 print(f" isSignal single candidate: {label_stats['single_signal']}")
1429 print(f" isSignal multi, same btag index: {label_stats['multi_same_btag']}")
1430 print(f" isSignal multi, different btag index: {label_stats['multi_diff_btag']}")
1431 print(f" fallback deltaP/background logic: {label_stats['fallback']}")
1432
1433 mbc_values = None
1434
1435 # Apply train_selection. The main network currently keeps all preselected
1436 # events, including predicted-sector-empty fallback-background cases.
1437 n_before = len(labels)
1438 if features_dense is not None:
1439 features_dense = features_dense[train_selection]
1440 else:
1441 features = features[train_selection]
1442 sparse_extra_features = cat_augments[train_selection]
1443 labels = labels[train_selection]
1444 event_weights = event_weights[train_selection]
1445 if mbc_values is not None:
1446 mbc_values = mbc_values[train_selection]
1447 n_dropped = n_before - int(train_selection.sum())
1448 print(
1449 f" train_selection: dropped {n_dropped} events "
1450 f"({n_dropped / n_before * 100:.1f}%); predicted-sector-empty events are kept"
1451 )
1452
1453 # Mbc values for DisCo loss (category network; main network handled above)
1454 if args.network == 'category':
1455 mbc_values = None
1456
1457 # Train/val split
1458 print(f"\nSplitting train/val (val_split={args.val_split})...")
1459 n_events = len(labels) if features_dense is None else len(features_dense)
1460 n_val = int(n_events * args.val_split)
1461 n_train = n_events - n_val
1462
1463 indices = np.random.permutation(n_events)
1464 train_idx = indices[:n_train]
1465 val_idx = indices[n_train:]
1466
1467 # Split Mbc values and event weights
1468 if mbc_values is not None:
1469 mbc_train = mbc_values[train_idx]
1470 mbc_val = mbc_values[val_idx]
1471 else:
1472 mbc_train = None
1473 mbc_val = None
1474
1475 w_train = event_weights[train_idx]
1476 w_val = event_weights[val_idx]
1477
1478 print(f" Train: {n_train} events")
1479 print(f" Val: {n_val} events")
1480
1481 # Create dataloaders (sparse or dense)
1482 if features_dense is None:
1483 # Sparse loading (category and main networks)
1484 mbc_train_np = mbc_train if mbc_train is not None else None
1485 mbc_val_np = mbc_val if mbc_val is not None else None
1486 train_extra = sparse_extra_features[train_idx] if sparse_extra_features is not None else None
1487 val_extra = sparse_extra_features[val_idx] if sparse_extra_features is not None else None
1488 train_dataset = SparseDataset(features[train_idx], labels[train_idx],
1489 mbc_train_np, train_extra, w_train)
1490 val_dataset = SparseDataset(features[val_idx], labels[val_idx],
1491 mbc_val_np, val_extra, w_val)
1492 drop_last = len(train_dataset) > args.batch_size
1493 train_loader = DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True,
1494 collate_fn=sparse_collate_fn, num_workers=resolved_num_workers,
1495 drop_last=drop_last)
1496 val_loader = DataLoader(val_dataset, batch_size=args.batch_size, shuffle=False,
1497 collate_fn=sparse_collate_fn, num_workers=resolved_num_workers)
1498 else:
1499 # Dense loading (default)
1500 X_train = torch.from_numpy(features_dense[train_idx])
1501 y_train = torch.from_numpy(labels[train_idx])
1502 X_val = torch.from_numpy(features_dense[val_idx])
1503 y_val = torch.from_numpy(labels[val_idx])
1504 w_train_t = torch.from_numpy(w_train)
1505 w_val_t = torch.from_numpy(w_val)
1506
1507 if mbc_train is not None:
1508 # Include mbc in dataset so shuffling aligns correctly with DisCo
1509 train_dataset = TensorDataset(X_train, y_train, mbc_train, w_train_t)
1510 val_dataset = TensorDataset(X_val, y_val, mbc_val, w_val_t)
1511 else:
1512 train_dataset = TensorDataset(X_train, y_train, w_train_t)
1513 val_dataset = TensorDataset(X_val, y_val, w_val_t)
1514
1515 drop_last = len(train_dataset) > args.batch_size
1516 train_loader = DataLoader(train_dataset, batch_size=args.batch_size, shuffle=True,
1517 num_workers=resolved_num_workers, drop_last=drop_last)
1518 val_loader = DataLoader(val_dataset, batch_size=args.batch_size, shuffle=False,
1519 num_workers=resolved_num_workers)
1520
1521 # Create model
1522 print("\n" + "=" * 60)
1523 print("Creating model")
1524 print("=" * 60)
1525 model = MultiClassNet(input_size=input_size, num_labels=num_labels)
1526 model = model.to(device)
1527
1528 print(f" Input size: {input_size}")
1529 print(f" Num labels: {num_labels}")
1530 print(f" Parameters: {sum(p.numel() for p in model.parameters()):,}")
1531
1532 # Loss and optimizer
1533 # reduction='none' gives per-sample losses; weights are applied in train_epoch and evaluate
1534 criterion = nn.CrossEntropyLoss(label_smoothing=args.label_smoothing, reduction='none')
1535 optimizer = optim.AdamW(model.parameters(), lr=args.lr, weight_decay=args.weight_decay)
1536 scheduler = None
1537 if args.lr_schedule == 'cosine':
1538 scheduler = optim.lr_scheduler.CosineAnnealingLR(
1539 optimizer, T_max=args.epochs, eta_min=args.eta_min
1540 )
1541
1542 # Training loop
1543 print("\n" + "=" * 60)
1544 print("Training")
1545 print("=" * 60)
1546 if args.disco_lambda > 0:
1547 print(f"Using distance correlation with lambda={args.disco_lambda}")
1548 print("Training configuration:")
1549 print(f" epochs: {args.epochs}")
1550 print(f" batch_size: {args.batch_size}")
1551 print(f" lr: {args.lr:.2e}")
1552 print(f" lr_schedule: {args.lr_schedule}")
1553 if args.lr_schedule == 'cosine':
1554 print(f" eta_min: {args.eta_min:.2e}")
1555 print(f" label_smoothing: {args.label_smoothing}")
1556
1557 best_val_loss = float('inf')
1558 best_epoch = 0
1559 training_start = time.time()
1560 history = {'train_loss': [], 'val_loss': [], 'disco_loss': [], 'lr': []}
1561 patience_counter = 0
1562 early_stopping_patience = 5
1563
1564 for epoch in range(args.epochs):
1565 epoch_start = time.time()
1566 current_lr = optimizer.param_groups[0]['lr']
1567
1568 # Train
1569 train_loss, disco_loss = train_epoch(
1570 model, train_loader, criterion, optimizer, device,
1571 disco_lambda=args.disco_lambda
1572 )
1573
1574 val_loss = evaluate(model, val_loader, criterion, device)
1575
1576 epoch_time = time.time() - epoch_start
1577
1578 # Scheduler step
1579 if scheduler is not None:
1580 scheduler.step()
1581 new_lr = optimizer.param_groups[0]['lr']
1582
1583 history['train_loss'].append(train_loss)
1584 history['val_loss'].append(val_loss)
1585 history['disco_loss'].append(disco_loss)
1586 history['lr'].append(current_lr)
1587
1588 # Print progress
1589 if args.disco_lambda > 0:
1590 print(f"Epoch {epoch+1:3d}/{args.epochs}: "
1591 f"train_loss={train_loss:.6f}, disco_loss={disco_loss:.6f}, "
1592 f"val_loss={val_loss:.6f}, lr={current_lr:.2e}, time={epoch_time:.1f}s")
1593 else:
1594 print(f"Epoch {epoch+1:3d}/{args.epochs}: "
1595 f"train_loss={train_loss:.6f}, val_loss={val_loss:.6f}, "
1596 f"lr={current_lr:.2e}, time={epoch_time:.1f}s")
1597 if new_lr < current_lr:
1598 print(f" -> LR reduced: {current_lr:.2e} -> {new_lr:.2e}")
1599
1600 # Save best model (treat nan val_loss as always saving, for no-val-set runs)
1601 if not (val_loss >= best_val_loss):
1602 best_val_loss = val_loss
1603 best_epoch = epoch
1604 patience_counter = 0
1605 model_path = os.path.join(args.output, f'net_{args.network}.pt')
1606 torch.save({
1607 'epoch': epoch,
1608 'model_state_dict': model.state_dict(),
1609 'optimizer_state_dict': optimizer.state_dict(),
1610 'val_loss': val_loss,
1611 'train_loss': train_loss,
1612 'history': history,
1613 'has_inputs': has_inputs,
1614 'config': {
1615 'input_size': input_size,
1616 'num_labels': num_labels,
1617 'fraction': args.fraction,
1618 'cont_fraction': args.cont_fraction,
1619 'network_type': args.network,
1620 'lr': args.lr,
1621 'lr_schedule': args.lr_schedule,
1622 'eta_min': args.eta_min,
1623 'label_smoothing': args.label_smoothing,
1624 }
1625 }, model_path)
1626 print(f" -> Saved best model (val_loss={val_loss:.6f})")
1627 else:
1628 patience_counter += 1
1629 if patience_counter >= early_stopping_patience:
1630 print(f" -> Early stopping at epoch {epoch + 1}")
1631 break
1632
1633 print("\n" + "=" * 60)
1634 print("Training complete")
1635 print("=" * 60)
1636 training_time = time.time() - training_start
1637 minutes = int(training_time // 60)
1638 seconds = int(training_time % 60)
1639 print(f"Training time: {minutes}m {seconds}s")
1640 print(f"Best epoch: {best_epoch+1}")
1641 print(f"Best val loss: {best_val_loss:.6f}")
1642 print(f"Model saved to: {os.path.join(args.output, f'net_{args.network}.pt')}")
1643
1644 # Update checkpoint with full history (best-epoch save only captured history up to that point)
1645 model_path = os.path.join(args.output, f'net_{args.network}.pt')
1646 final_ckpt = torch.load(model_path)
1647 final_ckpt['history'] = history
1648 torch.save(final_ckpt, model_path)
1649
1650 # Evaluation on validation set
1651 if len(val_loader.dataset) == 0:
1652 print("\nNo validation set, skipping evaluation.")
1653 return
1654
1655 print("\n" + "=" * 60)
1656 print("Evaluation on validation set")
1657 print("=" * 60)
1658
1659 # Load best model for evaluation
1660 best_checkpoint = torch.load(os.path.join(args.output, f'net_{args.network}.pt'))
1661 model.load_state_dict(best_checkpoint['model_state_dict'])
1662 model.eval()
1663
1664 val_probs_list = []
1665 val_labels_list = []
1666
1667 with torch.no_grad():
1668 for batch in val_loader:
1669 data = batch[0].to(device)
1670 target = batch[1]
1671 output = model(data)
1672 probs = torch.softmax(output, dim=1)
1673 val_probs_list.append(probs.cpu().numpy())
1674 val_labels_list.append(target.numpy())
1675
1676 val_probs = np.vstack(val_probs_list)
1677 val_labels = np.concatenate(val_labels_list)
1678
1679 # Per-class metrics
1680 print(f"\n {'Class':<28} {'N':>8} {'Accuracy':>10} {'Mean P':>10} {'Median P':>10}")
1681 print(f" {'-'*66}")
1682
1683 if args.network == 'category':
1684 eval_classes = [(i, name) for i, name in enumerate(['B0', 'B+', 'continuum'])]
1685 for cls_idx, cls_name in eval_classes:
1686 cls_mask = val_labels == cls_idx
1687 n_cls = cls_mask.sum()
1688 if n_cls > 0:
1689 cls_probs = val_probs[cls_mask, cls_idx]
1690 pred_cls = np.argmax(val_probs[cls_mask], axis=1)
1691 accuracy = (pred_cls == cls_idx).mean()
1692 print(f" {cls_name:<28} {n_cls:>8,} {accuracy:>10.4f} "
1693 f"{cls_probs.mean():>10.4f} {np.median(cls_probs):>10.4f}")
1694 else:
1695 # Main network: show signal modes aggregate + 3 background classes
1696 signal_mask = val_labels < config.N_INPUT_IDS
1697 n_sig = signal_mask.sum()
1698 if n_sig > 0:
1699 sig_true = val_labels[signal_mask]
1700 sig_probs = val_probs[signal_mask][np.arange(n_sig), sig_true]
1701 sig_acc = (np.argmax(val_probs[signal_mask], axis=1) == sig_true).mean()
1702 cls_name = f'signal modes (0-{config.N_INPUT_IDS - 1})'
1703 print(f" {cls_name:<28} {n_sig:>8,} {sig_acc:>10.4f} "
1704 f"{sig_probs.mean():>10.4f} {np.median(sig_probs):>10.4f}")
1705 bg_classes = [
1706 (MAIN_BG_BAD_TAG, 'bad_tag'),
1707 (MAIN_BG_CROSS_DC1, 'cross_deltaC1'),
1708 (MAIN_BG_CONT, 'continuum'),
1709 ]
1710 for cls_idx, cls_name in bg_classes:
1711 cls_mask = val_labels == cls_idx
1712 n_cls = cls_mask.sum()
1713 if n_cls > 0:
1714 cls_probs = val_probs[cls_mask, cls_idx]
1715 pred_cls = np.argmax(val_probs[cls_mask], axis=1)
1716 accuracy = (pred_cls == cls_idx).mean()
1717 print(f" {cls_name:<28} {n_cls:>8,} {accuracy:>10.4f} "
1718 f"{cls_probs.mean():>10.4f} {np.median(cls_probs):>10.4f}")
1719
1720 # Overall accuracy
1721 overall_acc = (np.argmax(val_probs, axis=1) == val_labels).mean()
1722 print(f"\n Overall accuracy: {overall_acc:.4f}")
1723
1724
Definition main.py:1

◆ save_npz_shard()

save_npz_shard ( path,
data )
Write a per-file/merged data dict to a .npz shard readable by _load_npz_shard().

The CSR feature matrix is stored as its three component arrays plus its shape; all
other entries are stored as-is. Uncompressed (np.savez): the features are ~1.6%
dense, so a shard is small already and load speed matters more than size.

Parameters:
    path (str): Output .npz path.
    data (dict): Data dict as returned by _load_root_training_file()/concat_file_data().

Definition at line 252 of file train.py.

252def save_npz_shard(path, data):
253 """
254 Write a per-file/merged data dict to a .npz shard readable by _load_npz_shard().
255
256 The CSR feature matrix is stored as its three component arrays plus its shape; all
257 other entries are stored as-is. Uncompressed (np.savez): the features are ~1.6%
258 dense, so a shard is small already and load speed matters more than size.
259
260 Parameters:
261 path (str): Output .npz path.
262 data (dict): Data dict as returned by _load_root_training_file()/concat_file_data().
263 """
264 feats = data['features'].tocsr()
265 arrays = {
266 'format_version': np.asarray(NPZ_SHARD_VERSION, dtype=np.int32),
267 'feat_data': feats.data.astype(np.float32, copy=False),
268 'feat_indices': feats.indices.astype(np.int32, copy=False),
269 'feat_indptr': feats.indptr.astype(np.int64, copy=False),
270 'feat_shape': np.asarray(feats.shape, dtype=np.int64),
271 }
272 for name in EVENT_TRUTH_FIELDS:
273 arrays[name] = data[name]
274 for name in SIG_PACKED_VALUE_FIELDS:
275 arrays[name] = data[name]
276 arrays['sig_input_ids_offsets'] = data['sig_input_ids_offsets']
277
278 np.savez(path, **arrays)
279
280

◆ sparse_collate_fn()

sparse_collate_fn ( batch)
Collate function for SparseDataset with optional Mbc and weight tensors.

Definition at line 1155 of file train.py.

1155def sparse_collate_fn(batch):
1156 """Collate function for SparseDataset with optional Mbc and weight tensors."""
1157 if len(batch[0]) == 4:
1158 features, labels, mbc, weights = zip(*batch)
1159 return (torch.stack(features), torch.stack(labels),
1160 torch.stack(mbc), torch.stack(weights))
1161 if len(batch[0]) == 3:
1162 features, labels, third = zip(*batch)
1163 return torch.stack(features), torch.stack(labels), torch.stack(third)
1164 features, labels = zip(*batch)
1165 return torch.stack(features), torch.stack(labels)
1166
1167

◆ train_epoch()

train_epoch ( model,
train_loader,
criterion,
optimizer,
device,
disco_lambda = 0.0 )
Train for one epoch with optional distance correlation penalty.

Parameters:
    model (nn.Module): Model to train.
    train_loader (DataLoader): Training data loader. Batch tuple formats supported:
        2-tuple (features, labels): no weights, no DisCo;
        3-tuple (features, labels, weights): weighted loss, no DisCo;
        4-tuple (features, labels, mbc, weights): weighted loss + DisCo.
    criterion (nn.Module): Loss function with reduction='none' (per-sample losses required).
    optimizer (torch.optim.Optimizer): Optimizer.
    device (torch.device): Device to use.
    disco_lambda (float): Coefficient for distance correlation loss (0 = disabled).

Returns:
    tuple: (train_loss, disco_loss) -- average (weighted) classification loss and
        average distance correlation loss (0.0 if disabled).

Definition at line 776 of file train.py.

776def train_epoch(model, train_loader, criterion, optimizer, device, disco_lambda=0.0):
777 """
778 Train for one epoch with optional distance correlation penalty.
779
780 Parameters:
781 model (nn.Module): Model to train.
782 train_loader (DataLoader): Training data loader. Batch tuple formats supported:
783 2-tuple (features, labels): no weights, no DisCo;
784 3-tuple (features, labels, weights): weighted loss, no DisCo;
785 4-tuple (features, labels, mbc, weights): weighted loss + DisCo.
786 criterion (nn.Module): Loss function with reduction='none' (per-sample losses required).
787 optimizer (torch.optim.Optimizer): Optimizer.
788 device (torch.device): Device to use.
789 disco_lambda (float): Coefficient for distance correlation loss (0 = disabled).
790
791 Returns:
792 tuple: (train_loss, disco_loss) -- average (weighted) classification loss and
793 average distance correlation loss (0.0 if disabled).
794 """
795 model.train()
796 total_loss = 0
797 total_disco = 0
798 n_batches = 0
799
800 for batch in train_loader:
801 # Unpack batch based on tuple length
802 if len(batch) == 4:
803 data, target, batch_mbc, sample_weights = batch
804 batch_mbc = batch_mbc.to(device)
805 sample_weights = sample_weights.to(device)
806 elif len(batch) == 3:
807 data, target, sample_weights = batch
808 batch_mbc = None
809 sample_weights = sample_weights.to(device)
810 else:
811 data, target = batch
812 batch_mbc = None
813 sample_weights = None
814
815 data, target = data.to(device), target.to(device)
816
817 optimizer.zero_grad()
818 output = model(data)
819
820 # Classification loss (per-sample when criterion has reduction='none')
821 loss_per_sample = criterion(output, target)
822 if sample_weights is not None:
823 cls_loss = (loss_per_sample * sample_weights).mean()
824 else:
825 cls_loss = loss_per_sample.mean()
826
827 # Distance correlation loss (if enabled)
828 disco_loss = torch.tensor(0.0, device=device)
829 if disco_lambda > 0 and batch_mbc is not None:
830 probs = torch.softmax(output, dim=1)
831
832 if probs.shape[1] > config.NUM_CAT_LABELS: # main network
833 # Signal proxy: max softmax prob over predicted sector's signal modes.
834 # charged_cat is appended as the last feature in the batch.
835 bp_threshold = config.N_BP_MODES * 2
836 charged_cat_batch = data[:, -1] > 0.5
837 bp_max = probs[:, :bp_threshold].max(dim=1).values
838 b0_max = probs[:, bp_threshold:config.N_INPUT_IDS].max(dim=1).values
839 signal_prob = torch.where(charged_cat_batch, bp_max, b0_max)
840
841 cont_class = MAIN_BG_CONT
842
843 cont_mask = target == cont_class
844 if cont_mask.sum() > 1:
845 disco_loss = disco_loss + disco_lambda * distance_corr(
846 signal_prob[cont_mask], batch_mbc[cont_mask]
847 )
848
849 loss = cls_loss + disco_loss
850 loss.backward()
851
852 # Gradient clipping for stability
853 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
854
855 optimizer.step()
856
857 total_loss += cls_loss.item() * len(data)
858 total_disco += disco_loss.item()
859 n_batches += 1
860
861 return (total_loss / len(train_loader.dataset),
862 total_disco / n_batches if n_batches > 0 else 0.0)
863
864

Variable Documentation

◆ _config_path

_config_path = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), 'config.py')
protected

Path to config.py in the package directory above, loaded directly instead of through the package.

Definition at line 59 of file train.py.

◆ _spec

_spec = importlib.util.spec_from_file_location('modeSelector_config', _config_path)
protected

Import spec used to load config.py as a standalone module.

Definition at line 61 of file train.py.

◆ config

config = importlib.util.module_from_spec(_spec)

ModeSelector configuration, loaded without importing the basf2-dependent package.

Definition at line 63 of file train.py.

◆ EVENT_TRUTH_DTYPES

dict EVENT_TRUTH_DTYPES
Initial value:
1= {
2 'is_cont': np.int8, 'gen_pdg': np.int16,
3 'bp_gen_decay_mode_id': np.int16, 'b0_gen_decay_mode_id': np.int16,
4 'bp_gen_fei_calib_weight': np.float32, 'b0_gen_fei_calib_weight': np.float32,
5 'bp_is_best': np.int8, 'best_sigprob': np.float32,
6 'best_bp_sigprob_iid': np.int16, 'best_b0_sigprob_iid': np.int16,
7 'bp_tag_is_gen': np.int8, 'b0_tag_is_gen': np.int8, 'fei_calib_weight': np.float32,
8 'best_bp_iid': np.int16, 'best_bp_dp': np.float32,
9 'best_b0_iid': np.int16, 'best_b0_dp': np.float32,
10}

dtype used to cast each ms_tr_<name> branch back from the ROOT double it was stored as

Definition at line 84 of file train.py.

◆ EVENT_TRUTH_FIELDS

tuple EVENT_TRUTH_FIELDS
Initial value:
1= (
2 'is_cont', 'gen_pdg', 'bp_gen_decay_mode_id', 'b0_gen_decay_mode_id',
3 'bp_gen_fei_calib_weight', 'b0_gen_fei_calib_weight', 'bp_is_best',
4 'best_sigprob', 'best_bp_sigprob_iid', 'best_b0_sigprob_iid',
5 'bp_tag_is_gen', 'b0_tag_is_gen', 'fei_calib_weight',
6 'best_bp_iid', 'best_bp_dp', 'best_b0_iid', 'best_b0_dp',
7)

Per-event truth/label branches written to the 'events' tree, aliased as ms_tr_<name>

Definition at line 76 of file train.py.

◆ MAIN_BG_BAD_TAG

MAIN_BG_BAD_TAG = config.N_INPUT_IDS

Main network label index for bad-tag background.

Definition at line 67 of file train.py.

◆ MAIN_BG_CONT

int MAIN_BG_CONT = config.N_INPUT_IDS + 2

Main network label index for continuum background.

Definition at line 71 of file train.py.

◆ MAIN_BG_CROSS_DC1

int MAIN_BG_CROSS_DC1 = config.N_INPUT_IDS + 1

Main network label index for cross-deltaC1 background.

Definition at line 69 of file train.py.

◆ MAIN_NUM_LABELS

int MAIN_NUM_LABELS = config.N_INPUT_IDS + 3

Total number of main network output labels.

Definition at line 73 of file train.py.

◆ NPZ_SHARD_VERSION

int NPZ_SHARD_VERSION = 1

Version tag written into every .npz shard, checked on load.

Definition at line 207 of file train.py.

◆ REQUIRED_EVENT_BRANCHES

REQUIRED_EVENT_BRANCHES
Initial value:
1= tuple(f'ms_tr_{name}' for name in EVENT_TRUTH_FIELDS) + (
2 '__experiment__', '__run__', '__event__',
3)

Branches required in the 'events' tree of a training input file.

Definition at line 97 of file train.py.

◆ SIG_CANDIDATE_BRANCHES

tuple SIG_CANDIDATE_BRANCHES
Initial value:
1= (
2 '__experiment__', '__run__', '__event__',
3 'ms_sig_input_id', 'ms_sig_btag_index', 'ms_sig_delta_p', 'ms_sig_sigprob',
4)

Branches read from (and required in) the 'sig_candidates' tree.

Definition at line 101 of file train.py.

◆ SIG_PACKED_VALUE_FIELDS

tuple SIG_PACKED_VALUE_FIELDS
Initial value:
1= (
2 'sig_input_ids_values', 'sig_btag_index_values',
3 'sig_delta_p_values', 'sig_sigprob_values',
4)

Packed ragged per-candidate arrays, all indexed by 'sig_input_ids_offsets'.

Definition at line 209 of file train.py.