326 Load one input file and apply the per-event preselection and downsampling.
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.
333 args (tuple): (input_file, seed, fraction, cont_fraction, sigprob_thresh).
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.
340 input_file, seed, fraction, cont_fraction, sigprob_thresh = args
344 except Exception
as exc:
345 return None, f
"{input_file}: failed to read input file ({exc})"
347 feats = data[
'features']
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']
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']
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']
376 presel = best_sigprob > sigprob_thresh
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)
382 file_rng = np.random.default_rng(seed)
383 sampled = (file_rng.random(len(best_sigprob)) < sample_prob) & presel
385 keep_events = np.flatnonzero(sampled).astype(np.int64)
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])
395 return np.empty(0, dtype=values.dtype), out_offsets
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
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)
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])
423 sigprob_thresh=config.DEFAULT_FEI_SIGPROB_THRESHOLD,
424 random_state=None, n_workers=None):
426 Load training data and apply optional global downsampling.
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())).
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.
462 if isinstance(input_files, str):
463 input_files = [input_files]
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
476 raise ValueError(
"No input files found.")
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:
484 raise ValueError(f
"fraction must be <= 1.0, got {fraction}.")
486 if cont_fraction > 1.0:
487 raise ValueError(f
"cont_fraction must be <= 1.0, got {cont_fraction}.")
489 rng = np.random.default_rng(random_state)
491 seeds = rng.integers(0, 2**31, size=len(input_files))
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 = []
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)))
515 with ProcessPoolExecutor(max_workers=n_workers)
as executor:
516 results = executor.map(
518 [(f, s, fraction, cont_fraction, sigprob_thresh)
for f, s
in zip(input_files, seeds)],
520 results = tqdm(results, total=len(input_files), desc=
"Loading files")
522 for result, error
in results:
523 if error
is not None:
524 load_errors.append(error)
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)
561 print(f
"\nWARNING: skipped {len(load_errors)}/{len(input_files)} unreadable input file(s):")
562 for err
in load_errors[:20]:
564 if len(load_errors) > 20:
565 print(f
" ... and {len(load_errors) - 20} more files")
567 if not features_list:
568 raise ValueError(
"All input files failed to load. Fix input files before training.")
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)
579 print(f
"\nTotal after concatenation: {features.shape[0]} events")
582 print(
"\nVerifying HAS_INPUTS from feature sparsity...")
584 n_total = features.shape[1]
587 nonzero_cols = set(np.unique(features.nonzero()[1]))
588 all_zero_cols = set(range(n_total)) - nonzero_cols
594 n_input_ids = config.N_INPUT_IDS
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),
600 excluded_cols = set()
601 for start, end
in excluded_block_ranges.values():
602 excluded_cols.update(range(start, end))
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}")
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")
623 for idx
in computed_has_inputs:
626 f.write(
" " +
", ".join(str(x)
for x
in row) +
",\n")
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
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)
640 features = features[:, has_inputs]
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),
648 if sig_input_ids_offsets_list:
652 shifted_offsets = [np.zeros(1, dtype=np.int64)]
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)
659 sig_input_ids_offsets = np.zeros(features.shape[0] + 1, dtype=np.int64)
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)
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)
672 if len(sig_input_ids_offsets) != features.shape[0] + 1:
674 f
"sig_input_ids_offsets length mismatch: got {len(sig_input_ids_offsets)}, "
675 f
"expected {features.shape[0] + 1}"
677 expected_n = int(sig_input_ids_offsets[-1])
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),
684 if len(arr) != expected_n:
686 f
"{name} length mismatch: got {len(arr)}, expected {expected_n} from offsets"
690 sig_input_ids_values,
691 sig_input_ids_offsets,
692 sig_btag_index_values,
696 event_scalars = (is_cont, gen_pdg, bp_is_best, best_sigprob,
697 best_bp_sigprob_iid, best_b0_sigprob_iid)
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),
707 return (features, event_scalars, has_inputs, mc_truth_cand, sig_truth, calib_inputs)
776def train_epoch(model, train_loader, criterion, optimizer, device, disco_lambda=0.0):
778 Train for one epoch with optional distance correlation penalty.
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).
792 tuple: (train_loss, disco_loss) -- average (weighted) classification loss and
793 average distance correlation loss (0.0 if disabled).
800 for batch
in train_loader:
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
809 sample_weights = sample_weights.to(device)
813 sample_weights =
None
815 data, target = data.to(device), target.to(device)
817 optimizer.zero_grad()
821 loss_per_sample = criterion(output, target)
822 if sample_weights
is not None:
823 cls_loss = (loss_per_sample * sample_weights).mean()
825 cls_loss = loss_per_sample.mean()
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)
832 if probs.shape[1] > config.NUM_CAT_LABELS:
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)
841 cont_class = MAIN_BG_CONT
843 cont_mask = target == cont_class
844 if cont_mask.sum() > 1:
846 signal_prob[cont_mask], batch_mbc[cont_mask]
849 loss = cls_loss + disco_loss
853 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
857 total_loss += cls_loss.item() * len(data)
858 total_disco += disco_loss.item()
861 return (total_loss / len(train_loader.dataset),
862 total_disco / n_batches
if n_batches > 0
else 0.0)
959 delta_p_thresh=0.15):
961 Build mode-prediction labels (0 to N_INPUT_IDS+2) from MC truth.
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
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.
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.
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.
996 n_events = len(is_cont)
997 bp_threshold = config.N_BP_MODES * 2
999 (sig_input_ids_values, sig_input_ids_offsets,
1000 sig_btag_index_values, sig_delta_p_values, sig_sigprob_values) = sig_truth
1002 if len(sig_input_ids_offsets) != n_events + 1:
1004 f
"sig_input_ids_offsets length mismatch: got {len(sig_input_ids_offsets)}, expected {n_events + 1}"
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)
1012 labels = np.full(n_events, MAIN_BG_BAD_TAG, dtype=np.int64)
1013 assigned = np.zeros(n_events, dtype=bool)
1017 'multi_same_btag': 0,
1018 'multi_diff_btag': 0,
1022 for evt
in range(n_events):
1023 start = int(sig_input_ids_offsets[evt])
1024 end = int(sig_input_ids_offsets[evt + 1])
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]
1033 if charged_cat[evt]:
1034 sector_mask = evt_iids < bp_threshold
1036 sector_mask = evt_iids >= bp_threshold
1038 if not np.any(sector_mask):
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]
1046 if len(sector_iids) == 1:
1047 labels[evt] = int(sector_iids[0])
1048 assigned[evt] =
True
1049 stats[
'single_signal'] += 1
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
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):
1064 best_label = int(iid_val)
1065 labels[evt] = best_label
1066 assigned[evt] =
True
1067 stats[
'multi_diff_btag'] += 1
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]
1075 no_target = fallback_mask & (~has_target)
1077 sector_abs_pdg = np.where(charged_cat, 521, 511)
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
1087 train_selection = np.ones(n_events, dtype=bool)
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()}"
1093 return labels, train_selection, stats
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)')
1295 args = parser.parse_args()
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))
1304 resolved_num_workers = args.num_workers
1305 print(f
"DataLoader workers: {resolved_num_workers}")
1308 np.random.seed(args.seed)
1309 torch.manual_seed(args.seed)
1312 device = torch.device(
'cuda' if torch.cuda.is_available()
else 'cpu')
1313 print(f
"Using device: {device}")
1316 os.makedirs(args.output, exist_ok=
True)
1319 print(
"\n" +
"=" * 60)
1320 print(
"Loading data")
1322 features, event_scalars, has_inputs, mc_truth_cand, sig_truth, calib_inputs =
load_and_sample_data(
1324 fraction=args.fraction,
1325 cont_fraction=args.cont_fraction,
1326 random_state=args.seed,
1327 n_workers=resolved_num_workers,
1329 is_cont, gen_pdg, bp_is_best, best_sigprob, best_bp_sigprob_iid, best_b0_sigprob_iid = event_scalars
1331 print(
"\nComputing event weights...")
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
1340 print(
"\nBuilding labels...")
1341 if args.network ==
'category':
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}%)")
1349 input_size = features.shape[1]
1351 if not args.use_sparse:
1353 print(
"\nConverting to dense arrays...")
1354 features_dense = features.toarray().astype(np.float32)
1356 print(
"\nUsing sparse data loading (memory-efficient)...")
1357 features_dense =
None
1360 if not args.cat_model:
1361 raise ValueError(
"--cat_model required for main network training")
1364 print(f
"\nLoading category model from {args.cat_model}...")
1365 cat_checkpoint = torch.load(args.cat_model)
1367 input_size=features.shape[1],
1370 cat_model.load_state_dict(cat_checkpoint[
'model_state_dict'])
1371 cat_model = cat_model.to(device)
1373 print(f
" Category model loaded (epoch {cat_checkpoint['epoch']+1})")
1375 if not args.use_sparse:
1376 print(
"\nConverting to dense arrays...")
1377 features_dense = features.toarray().astype(np.float32)
1379 cat_model, features_dense, args.batch_size, device, use_sparse=
False
1382 print(
"\nUsing sparse data loading (memory-efficient)...")
1383 features_dense =
None
1385 cat_model, features, args.batch_size, device, use_sparse=
True
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}%)")
1396 charged_cat_bool = cat_outputs[:, 1] > cat_outputs[:, 0]
1397 charged_cat = charged_cat_bool.astype(np.float32)
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)
1408 if mc_truth_cand
is None or sig_truth
is None:
1410 "Required main-network truth arrays not found in training data. "
1411 "Re-collect training data with produceTrainingInputs.py."
1414 mc_truth_cand, sig_truth, charged_cat_bool, is_cont, gen_pdg,
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())
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']}")
1437 n_before = len(labels)
1438 if features_dense
is not None:
1439 features_dense = features_dense[train_selection]
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())
1449 f
" train_selection: dropped {n_dropped} events "
1450 f
"({n_dropped / n_before * 100:.1f}%); predicted-sector-empty events are kept"
1454 if args.network ==
'category':
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
1463 indices = np.random.permutation(n_events)
1464 train_idx = indices[:n_train]
1465 val_idx = indices[n_train:]
1468 if mbc_values
is not None:
1469 mbc_train = mbc_values[train_idx]
1470 mbc_val = mbc_values[val_idx]
1475 w_train = event_weights[train_idx]
1476 w_val = event_weights[val_idx]
1478 print(f
" Train: {n_train} events")
1479 print(f
" Val: {n_val} events")
1482 if features_dense
is None:
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)
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)
1507 if mbc_train
is not None:
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)
1512 train_dataset = TensorDataset(X_train, y_train, w_train_t)
1513 val_dataset = TensorDataset(X_val, y_val, w_val_t)
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)
1522 print(
"\n" +
"=" * 60)
1523 print(
"Creating model")
1525 model =
MultiClassNet(input_size=input_size, num_labels=num_labels)
1526 model = model.to(device)
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()):,}")
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)
1537 if args.lr_schedule ==
'cosine':
1538 scheduler = optim.lr_scheduler.CosineAnnealingLR(
1539 optimizer, T_max=args.epochs, eta_min=args.eta_min
1543 print(
"\n" +
"=" * 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}")
1557 best_val_loss = float(
'inf')
1559 training_start = time.time()
1560 history = {
'train_loss': [],
'val_loss': [],
'disco_loss': [],
'lr': []}
1561 patience_counter = 0
1562 early_stopping_patience = 5
1564 for epoch
in range(args.epochs):
1565 epoch_start = time.time()
1566 current_lr = optimizer.param_groups[0][
'lr']
1570 model, train_loader, criterion, optimizer, device,
1571 disco_lambda=args.disco_lambda
1574 val_loss =
evaluate(model, val_loader, criterion, device)
1576 epoch_time = time.time() - epoch_start
1579 if scheduler
is not None:
1581 new_lr = optimizer.param_groups[0][
'lr']
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)
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")
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}")
1601 if not (val_loss >= best_val_loss):
1602 best_val_loss = val_loss
1604 patience_counter = 0
1605 model_path = os.path.join(args.output, f
'net_{args.network}.pt')
1608 'model_state_dict': model.state_dict(),
1609 'optimizer_state_dict': optimizer.state_dict(),
1610 'val_loss': val_loss,
1611 'train_loss': train_loss,
1613 'has_inputs': has_inputs,
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,
1621 'lr_schedule': args.lr_schedule,
1622 'eta_min': args.eta_min,
1623 'label_smoothing': args.label_smoothing,
1626 print(f
" -> Saved best model (val_loss={val_loss:.6f})")
1628 patience_counter += 1
1629 if patience_counter >= early_stopping_patience:
1630 print(f
" -> Early stopping at epoch {epoch + 1}")
1633 print(
"\n" +
"=" * 60)
1634 print(
"Training complete")
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')}")
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)
1651 if len(val_loader.dataset) == 0:
1652 print(
"\nNo validation set, skipping evaluation.")
1655 print(
"\n" +
"=" * 60)
1656 print(
"Evaluation on validation set")
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'])
1665 val_labels_list = []
1667 with torch.no_grad():
1668 for batch
in val_loader:
1669 data = batch[0].to(device)
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())
1676 val_probs = np.vstack(val_probs_list)
1677 val_labels = np.concatenate(val_labels_list)
1680 print(f
"\n {'Class':<28} {'N':>8} {'Accuracy':>10} {'Mean P':>10} {'Median P':>10}")
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()
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}")
1696 signal_mask = val_labels < config.N_INPUT_IDS
1697 n_sig = signal_mask.sum()
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}")
1706 (MAIN_BG_BAD_TAG,
'bad_tag'),
1707 (MAIN_BG_CROSS_DC1,
'cross_deltaC1'),
1708 (MAIN_BG_CONT,
'continuum'),
1710 for cls_idx, cls_name
in bg_classes:
1711 cls_mask = val_labels == cls_idx
1712 n_cls = cls_mask.sum()
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}")
1721 overall_acc = (np.argmax(val_probs, axis=1) == val_labels).mean()
1722 print(f
"\n Overall accuracy: {overall_acc:.4f}")