Belle II Software light-2609-luna
train.py
1#!/usr/bin/env python3
2"""
3Training script for ModeSelector neural networks.
4
5This script trains the two-stage ModeSelector:
61. Category network (B0 vs B+ vs continuum classification)
72. Main network (signal vs background classification using category output)
8
9Inputs are either produceTrainingInputs.py ROOT files or the .npz shards written by
10convert_training_inputs.py. Prefer the shards for anything beyond a few hundred ROOT
11files: reading the raw ROOT inputs costs ~11 s per file and is paid again on every
12invocation, i.e. twice per full training (category, then main).
13
14Usage:
15 # Train category network (single file)
16 python3 train.py --input modeSelector_training.root --network category --output networks/
17
18 # Train category network (multiple files)
19 python3 train.py --input modeSelector_training_*.root --network category --output networks/
20
21 # Train on pre-converted .npz shards (recommended for large samples)
22 python3 train.py --input 'converted/*.npz' --network category --use_sparse
23
24 # Train main network (requires trained category network)
25 python3 train.py --input modeSelector_training.root --network main \
26 --cat_model networks/net_category.pt --output networks/
27"""
28
29import argparse
30import glob
31import os
32import time
33from concurrent.futures import ProcessPoolExecutor
34
35import numpy as np
36import torch
37import torch.nn as nn
38import torch.optim as optim
39import uproot
40from scipy import sparse
41from torch.utils.data import DataLoader, TensorDataset
42from tqdm import tqdm
43
44try:
45 from modeSelector import config
46except Exception:
47 # Training needs nothing from basf2 -- only the constants in config.py, which itself
48 # imports nothing. The package __init__ does pull in basf2 and ROOT (via
49 # ModeSelectorModule/dstarVeto/generatedDecayWeights), so importing config through the
50 # package would force a basf2 setup. Load it straight from the file instead, which is
51 # what lets train.py run in a standalone venv with a CUDA torch build (the basf2
52 # externals ship torch CPU-only).
53 #
54 # Deliberately broad: if a basf2 PYTHONPATH is inherited (HTCondor getenv=true) but the
55 # interpreter is the venv's, pybasf2's compiled module is found and fails to initialise
56 # with SystemError rather than ImportError.
57 import importlib.util
58
59 _config_path = os.path.join(os.path.dirname(os.path.dirname(os.path.abspath(__file__))), 'config.py')
60
61 _spec = importlib.util.spec_from_file_location('modeSelector_config', _config_path)
62
63 config = importlib.util.module_from_spec(_spec)
64 _spec.loader.exec_module(config)
65
66
67MAIN_BG_BAD_TAG = config.N_INPUT_IDS
68
69MAIN_BG_CROSS_DC1 = config.N_INPUT_IDS + 1
70
71MAIN_BG_CONT = config.N_INPUT_IDS + 2
72
73MAIN_NUM_LABELS = config.N_INPUT_IDS + 3
74
75
76EVENT_TRUTH_FIELDS = (
77 'is_cont', 'gen_pdg', 'bp_gen_decay_mode_id', 'b0_gen_decay_mode_id',
78 'bp_gen_fei_calib_weight', 'b0_gen_fei_calib_weight', 'bp_is_best',
79 'best_sigprob', 'best_bp_sigprob_iid', 'best_b0_sigprob_iid',
80 'bp_tag_is_gen', 'b0_tag_is_gen', 'fei_calib_weight',
81 'best_bp_iid', 'best_bp_dp', 'best_b0_iid', 'best_b0_dp',
82)
83
84EVENT_TRUTH_DTYPES = {
85 'is_cont': np.int8, 'gen_pdg': np.int16,
86 'bp_gen_decay_mode_id': np.int16, 'b0_gen_decay_mode_id': np.int16,
87 'bp_gen_fei_calib_weight': np.float32, 'b0_gen_fei_calib_weight': np.float32,
88 'bp_is_best': np.int8, 'best_sigprob': np.float32,
89 'best_bp_sigprob_iid': np.int16, 'best_b0_sigprob_iid': np.int16,
90 'bp_tag_is_gen': np.int8, 'b0_tag_is_gen': np.int8, 'fei_calib_weight': np.float32,
91 'best_bp_iid': np.int16, 'best_bp_dp': np.float32,
92 'best_b0_iid': np.int16, 'best_b0_dp': np.float32,
93}
94
95
96
97REQUIRED_EVENT_BRANCHES = tuple(f'ms_tr_{name}' for name in EVENT_TRUTH_FIELDS) + (
98 '__experiment__', '__run__', '__event__',
99)
100
101SIG_CANDIDATE_BRANCHES = (
102 '__experiment__', '__run__', '__event__',
103 'ms_sig_input_id', 'ms_sig_btag_index', 'ms_sig_delta_p', 'ms_sig_sigprob',
104)
105
106
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
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
206
207NPZ_SHARD_VERSION = 1
208
209SIG_PACKED_VALUE_FIELDS = (
210 'sig_input_ids_values', 'sig_btag_index_values',
211 'sig_delta_p_values', 'sig_sigprob_values',
212)
213
214
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
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
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
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
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
422def load_and_sample_data(input_files, fraction=1.0, cont_fraction=1.0,
423 sigprob_thresh=config.DEFAULT_FEI_SIGPROB_THRESHOLD,
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
710class MultiClassNet(nn.Module):
711 """Multi-class classification network with fully connected layers.
712
713 Two architectures selected automatically by num_labels:
714 - Deep (num_labels <= 10, e.g. category network with 3 classes):
715 256 -> 128 (dropout 0.2) -> 64 (dropout 0.15) -> 32 (dropout 0.1) -> 16 -> num_labels
716 - Shallow (num_labels > 10, e.g. main network with 139 classes):
717 256 -> 128 (dropout 0.1) -> 128 (dropout 0.1) -> 64 (dropout 0.05) -> num_labels
718 - ReLU activation, Xavier initialization (gain=0.5, bias=0.01)
719 """
720
721 def __init__(self, input_size, num_labels=3):
722 """Build the network graph for the given input size and number of output classes."""
723 super().__init__()
724
725
726 self.activation = nn.ReLU()
727
728 if num_labels <= 10:
729
730 self.network = nn.Sequential(
731 nn.Linear(input_size, 128),
732 self.activation,
733 nn.Linear(128, 128),
734 nn.Dropout(0.2),
735 self.activation,
736 nn.Linear(128, 64),
737 nn.Dropout(0.15),
738 self.activation,
739 nn.Linear(64, 32),
740 nn.Dropout(0.1),
741 self.activation,
742 nn.Linear(32, 16),
743 self.activation,
744 nn.Linear(16, num_labels)
745 )
746 else:
747 self.network = nn.Sequential(
748 nn.Linear(input_size, 256),
749 self.activation,
750 nn.Linear(256, 128),
751 nn.Dropout(0.1),
752 self.activation,
753 nn.Linear(128, 128),
754 nn.Dropout(0.1),
755 self.activation,
756 nn.Linear(128, 64),
757 nn.Dropout(0.05),
758 self.activation,
759 nn.Linear(64, num_labels)
760 )
761
762 self._init_weights()
763
764 def _init_weights(self):
765 """Initialize weights with Xavier uniform (gain=0.5) and bias=0.01."""
766 for m in self.modules():
767 if isinstance(m, nn.Linear):
768 nn.init.xavier_uniform_(m.weight, gain=0.5)
769 nn.init.constant_(m.bias, 0.01)
770
771 def forward(self, x):
772 """Run a forward pass through the network."""
773 return self.network(x)
774
775
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
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
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
958def build_mode_labels(mc_truth_cand, sig_truth, charged_cat, is_cont, gen_pdg,
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
1096class SparseDataset(torch.utils.data.Dataset):
1097 """Dataset for sparse feature matrices.
1098
1099 Converts each row to dense on-the-fly during batching.
1100 Slower than dense loading but significantly more memory-efficient.
1101 """
1102
1103 def __init__(self, features_sparse, labels, mbc_values=None, extra_features=None,
1104 weights=None):
1105 """
1106 Parameters:
1107 features_sparse (scipy.sparse matrix): Sparse feature matrix of shape
1108 (n_samples, n_features). CSR format recommended.
1109 labels (np.ndarray): Target labels, shape (n_samples,).
1110 mbc_values (np.ndarray, optional): Mbc values per event for DisCo loss,
1111 shape (n_samples,).
1112 extra_features (np.ndarray, optional): Additional dense features concatenated
1113 per sample, shape (n_samples, n_extra).
1114 weights (np.ndarray, optional): Per-event loss weights, shape (n_samples,)
1115 (e.g. FEI calibration weights).
1116 """
1117
1118 self.features = features_sparse.tocsr()
1119
1120 self.labels = labels
1121
1122 self.mbc_values = mbc_values
1123
1124 self.extra_features = extra_features
1125
1126 self.weights = weights
1127
1128 def __len__(self):
1129 """Return the number of samples in the dataset."""
1130 return self.features.shape[0]
1131
1132 def __getitem__(self, idx):
1133 """Return one sample (features, label, [mbc, [weight]]) for the given index."""
1134 # Convert sparse row to dense 1D tensor
1135 feature_row = self.features[idx].toarray().astype(np.float32).squeeze()
1136 if self.extra_features is not None:
1137 feature_row = np.concatenate([feature_row, self.extra_features[idx]]).astype(np.float32)
1138 feature_row = torch.from_numpy(feature_row)
1139 label = torch.tensor(self.labels[idx], dtype=torch.long)
1140
1141 if self.mbc_values is not None:
1142 mbc = torch.tensor(self.mbc_values[idx], dtype=torch.float32)
1143 if self.weights is not None:
1144 w = torch.tensor(self.weights[idx], dtype=torch.float32)
1145 return feature_row, label, mbc, w
1146 return feature_row, label, mbc
1147
1148 if self.weights is not None:
1149 w = torch.tensor(self.weights[idx], dtype=torch.float32)
1150 return feature_row, label, w
1151
1152 return feature_row, label
1153
1154
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
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
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
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
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
1725if __name__ == '__main__':
1726 main()
network
Fully connected network layers.
Definition train.py:730
__init__(self, input_size, num_labels=3)
Definition train.py:721
forward(self, x)
Definition train.py:771
_init_weights(self)
Definition train.py:764
activation
Activation function applied between all linear layers.
Definition train.py:726
weights
Per-event loss weights, or None.
Definition train.py:1126
extra_features
Additional dense features concatenated at retrieval time, or None.
Definition train.py:1124
mbc_values
Mbc values per event for DisCo loss, or None.
Definition train.py:1122
__getitem__(self, idx)
Definition train.py:1132
__init__(self, features_sparse, labels, mbc_values=None, extra_features=None, weights=None)
Definition train.py:1104
features
Feature matrix in CSR format.
Definition train.py:1118
labels
Target class labels.
Definition train.py:1120
Definition main.py:1
_load_root_training_file(input_file)
Definition train.py:131
evaluate(model, loader, criterion, device)
Definition train.py:865
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)
Definition train.py:424
_process_one_file(args)
Definition train.py:324
compute_event_weights(event_scalars, calib_inputs)
Definition train.py:1189
_load_npz_shard(input_file)
Definition train.py:281
main()
Definition train.py:1260
sparse_collate_fn(batch)
Definition train.py:1155
distance_corr(var_1, var_2, normedweight=None, power=1)
Definition train.py:892
_validate_trees(f, input_file)
Definition train.py:107
concat_file_data(data_list)
Definition train.py:215
generate_category_outputs(cat_model, features, batch_size, device, use_sparse)
Definition train.py:1230
train_epoch(model, train_loader, criterion, optimizer, device, disco_lambda=0.0)
Definition train.py:776
build_category_labels(is_cont, gen_pdg)
Definition train.py:1168
build_mode_labels(mc_truth_cand, sig_truth, charged_cat, is_cont, gen_pdg, delta_p_thresh=0.15)
Definition train.py:959
save_npz_shard(path, data)
Definition train.py:252
load_training_file(input_file)
Definition train.py:308