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
1308 np.random.seed(args.seed)
1309 torch.manual_seed(args.seed)
1310
1311
1312 device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
1313 print(f"Using device: {device}")
1314
1315
1316 os.makedirs(args.output, exist_ok=True)
1317
1318
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
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
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
1358
1359 else:
1360 if not args.cat_model:
1361 raise ValueError("--cat_model required for main network training")
1362
1363
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
1396 charged_cat_bool = cat_outputs[:, 1] > cat_outputs[:, 0]
1397 charged_cat = charged_cat_bool.astype(np.float32)
1398
1399
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
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
1436
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
1454 if args.network == 'category':
1455 mbc_values = None
1456
1457
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
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
1482 if features_dense is None:
1483
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
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
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
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
1533
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
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
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
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
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
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
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
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
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
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
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
1721 overall_acc = (np.argmax(val_probs, axis=1) == val_labels).mean()
1722 print(f"\n Overall accuracy: {overall_acc:.4f}")
1723
1724