diff --git a/ml/networks/training/ivysaurus_pid/CreateTrainArrays.py b/ml/networks/training/ivysaurus_pid/CreateTrainArrays.py new file mode 100644 index 0000000..b8f8d82 --- /dev/null +++ b/ml/networks/training/ivysaurus_pid/CreateTrainArrays.py @@ -0,0 +1,194 @@ +import argparse +import numpy as np +import math + +import FileHelper + +######################################################################################################## + +detector_boundaries = { +"dune_hd" : { + "MinX":(-360.0 + 10.0), + "MaxX":(360.0 - 10.0), + "MinY":(-600.0 + 10.0), + "MaxY":(600.0 - 10.0), + "MinZ":(0 + 10.0), + "MaxZ":(1394.0 - 10.0) + } +} + +######################################################################################################## + +def main(args): + # Split training sample into contained and exiting + this_detector_boundaries = detector_boundaries.get(args.detector) + + (startGridU, startGridU_valid, + startGridV, startGridV_valid, + startGridW, startGridW_valid, + endGridU, endGridU_valid, + endGridV, endGridV_valid, + endGridW, endGridW_valid, + pfpVars, trackVars, showerVars, + y) = FileHelper.readTree(args, this_detector_boundaries) + + nEntries = startGridU.shape[0] + + if this_detector_boundaries is None: + raise ValueError(f"Unknown detector type {args.detector}. Available options: {list(detector_boundaries.keys())}") + + contained_mask = (pfpVars[:,0] > this_detector_boundaries["MinX"]) & (pfpVars[:,0] < this_detector_boundaries["MaxX"]) &\ + (pfpVars[:,1] > this_detector_boundaries["MinY"]) & (pfpVars[:,1] < this_detector_boundaries["MaxY"]) &\ + (pfpVars[:,2] > this_detector_boundaries["MinZ"]) & (pfpVars[:,2] < this_detector_boundaries["MaxZ"]) + + print('n_contained:', np.sum(contained_mask)) + print('n_exiting:', np.sum(~contained_mask)) + + for is_contained in [True, False] : + target_mask = (contained_mask == is_contained) + n_entries = np.sum(target_mask) + ntest = math.floor(n_entries * 0.1) + ntrain = math.floor(n_entries * 0.9) + + print('Contained' if is_contained else 'Exiting') + print('n_test:', ntest) + print('n_train:', ntrain) + + indices = np.flatnonzero(target_mask) + np.random.shuffle(indices) + train_idx = indices[:ntrain] + test_idx = indices[ntrain:ntrain + ntest] + + startGridU_train = startGridU[train_idx] + startGridV_train = startGridV[train_idx] + startGridW_train = startGridW[train_idx] + startGridU_test = startGridU[test_idx] + startGridV_test = startGridV[test_idx] + startGridW_test = startGridW[test_idx] + startGridU_valid_train = startGridU_valid[train_idx] + startGridV_valid_train = startGridV_valid[train_idx] + startGridW_valid_train = startGridW_valid[train_idx] + startGridU_valid_test = startGridU_valid[test_idx] + startGridV_valid_test = startGridV_valid[test_idx] + startGridW_valid_test = startGridW_valid[test_idx] + + endGridU_train = endGridU[train_idx] + endGridV_train = endGridV[train_idx] + endGridW_train = endGridW[train_idx] + endGridU_test = endGridU[test_idx] + endGridV_test = endGridV[test_idx] + endGridW_test = endGridW[test_idx] + endGridU_valid_train = endGridU_valid[train_idx] + endGridV_valid_train = endGridV_valid[train_idx] + endGridW_valid_train = endGridW_valid[train_idx] + endGridU_valid_test = endGridU_valid[test_idx] + endGridV_valid_test = endGridV_valid[test_idx] + endGridW_valid_test = endGridW_valid[test_idx] + + pfpVars_train = pfpVars[:,0][train_idx] + pfpVars_test = pfpVars[:,0][test_idx] + + trackVars_train = trackVars[train_idx] + trackVars_test = trackVars[test_idx] + + showerVars_train = showerVars[train_idx] + showerVars_test = showerVars[test_idx] + + y_train = y[train_idx] + y_test = y[test_idx] + + print('--------------------------------------') + print('startGridU_train', startGridU_train.shape) + print('startGridV_train', startGridV_train.shape) + print('startGridW_train', startGridW_train.shape) + print('startGridU_valid_train', startGridU_valid_train.shape) + print('startGridV_valid_train', startGridV_valid_train.shape) + print('startGridW_valid_train', startGridW_valid_train.shape) + print('endGridU_train', endGridU_train.shape) + print('endGridV_train', endGridV_train.shape) + print('endGridW_train', endGridW_train.shape) + print('endGridU_valid_train', endGridU_valid_train.shape) + print('endGridV_valid_train', endGridV_valid_train.shape) + print('endGridW_valid_train', endGridW_valid_train.shape) + print('pfpVars_train', pfpVars_train.shape) + print('trackVars_train', trackVars_train.shape) + print('showerVars_train', showerVars_train.shape) + print('y_train', y_train.shape) + + class_labels = ['Muon', 'Proton', 'Pion', 'Electron', 'Photon'] + particleType_train = np.argmax(y_train, axis=1) + class_counts_train = np.bincount(particleType_train, minlength=len(class_labels)) + + for i in range(len(class_labels)) : + print(f'n_{class_labels[i]}_train: {class_counts_train[i]}') + + print('--------------------------------------') + print('startGridU_test', startGridU_test.shape) + print('startGridV_test', startGridV_test.shape) + print('startGridW_test', startGridW_test.shape) + print('startGridU_valid_test', startGridU_valid_test.shape) + print('startGridV_valid_test', startGridV_valid_test.shape) + print('startGridW_valid_test', startGridW_valid_test.shape) + print('endGridU_test', endGridU_test.shape) + print('endGridV_test', endGridV_test.shape) + print('endGridW_test', endGridW_test.shape) + print('endGridU_valid_test', endGridU_valid_test.shape) + print('endGridV_valid_test', endGridV_valid_test.shape) + print('endGridW_valid_test', endGridW_valid_test.shape) + print('pfpVars_test', pfpVars_test.shape) + print('trackVars_test', trackVars_test.shape) + print('showerVars_test', showerVars_test.shape) + print('y_test', y_test.shape) + + particleType_test = np.argmax(y_test, axis=1) + class_counts_test = np.bincount(particleType_test, minlength=len(class_labels)) + + for i in range(len(class_labels)) : + print(f'n_{class_labels[i]}_test: {class_counts_test[i]}') + + print('--------------------------------------') + + prefix = f'{"Contained" if is_contained else "Exiting"}' + output_file_name = f"{args.output_dir}/{args.file_name.replace('.root', '')}_{prefix}.npz" + + np.savez(output_file_name, + startU_train=startGridU_train, startU_mask_train=startGridU_valid_train, + startV_train=startGridV_train, startV_mask_train=startGridV_valid_train, + startW_train=startGridW_train, startW_mask_train=startGridW_valid_train, + startU_test=startGridU_test, startU_mask_test=startGridU_valid_test, + startV_test=startGridV_test, startV_mask_test=startGridV_valid_test, + startW_test=startGridW_test, startW_mask_test=startGridW_valid_test, + endU_train=endGridU_train, endU_mask_train=endGridU_valid_train, + endV_train=endGridV_train, endV_mask_train=endGridV_valid_train, + endW_train=endGridW_train, endW_mask_train=endGridW_valid_train, + endU_test=endGridU_test, endU_mask_test=endGridU_valid_test, + endV_test=endGridV_test, endV_mask_test=endGridV_valid_test, + endW_test=endGridW_test, endW_mask_test=endGridW_valid_test, + trackVars_train=trackVars_train, + trackVars_test=trackVars_test, + showerVars_test=showerVars_test, + showerVars_train=showerVars_train, + y_train=y_train, + y_test=y_test) + +######################################################################################################## + +def parse_cli(): + parser = argparse.ArgumentParser(description="Ivysaurus PID") + + parser.add_argument("--file_name", type=str, required=True, help="Name of file to process") + parser.add_argument("--input_dir", type=str, required=True, help="Input file directory") + parser.add_argument("--output_dir", type=str, required=True, help="Dir to save processed files") + parser.add_argument("--pdgs", type=int, nargs="+", default=[13, 2212, 211, 11, 22], help="PDG IDs to collect - ATTN: order matters!") + parser.add_argument("--detector", type=str, default="dune_hd", help="Detector: dune_hd, ") + parser.add_argument("--dimensions", type=int, default=24, help="Grid dimensions") + parser.add_argument("--is_contained", action="store_true", help="Training for contained particles?") + + return parser.parse_args() + +######################################################################################################## + +if __name__ == "__main__": + args = parse_cli() + main(args) + diff --git a/ml/networks/training/ivysaurus_pid/FileHelper.py b/ml/networks/training/ivysaurus_pid/FileHelper.py new file mode 100644 index 0000000..4882195 --- /dev/null +++ b/ml/networks/training/ivysaurus_pid/FileHelper.py @@ -0,0 +1,352 @@ +import numpy as np +import uproot + +from tensorflow.keras.utils import to_categorical + +import Normalisation + +######################################################################################################## + +def readTree(args, detector) : + + print('Reading trees. This may take a while...') + + branch_names = [ + "StartGridU", "StartGridV", "StartGridW", + "EndGridU", "EndGridV", "EndGridW", + "PFPN2DHits", + "RecoEndX", "RecoEndY", "RecoEndZ", + "TrueEndX", "TrueEndY", "TrueEndZ", + "NTrackChildren", "NShowerChildren", + "NGrandChildren", "NChildHits", + "ChildEnergy", "ChildTrackScore", + "TrackLength", "TrackWobble", + "PFPTrackShowerScore", + "TrackMomComparison", + "ShowerDisplacement", + "ShowerDCA", + "ShowerTrackStubLength", + "ShowerNuVertexAvSeparation", + "ShowerNuVertexChargeAsymmetry", + "TruePDG", + "IsPrimary" + ] + + file_name = f'{args.input_dir}/{args.file_name}' + with uproot.open(file_name) as treeFile: + tree = treeFile["ivyTrain/ivysaur"] + #tree = treeFile["ivysaur"] + branches = tree.arrays(expressions=branch_names, library="np") + + # Grid lists + startGridU = branches['StartGridU'] + startGridV = branches['StartGridV'] + startGridW = branches['StartGridW'] + endGridU = branches['EndGridU'] + endGridV = branches['EndGridV'] + endGridW = branches['EndGridW'] + # PFPVar lists + nHits2D = branches['PFPN2DHits'] + trackScore = branches['PFPTrackShowerScore'] + # TrackVar lists + nTrackChildren = branches['NTrackChildren'] + nShowerChildren = branches['NShowerChildren'] + nGrandChildren = branches['NGrandChildren'] + nChildHits = branches['NChildHits'] + childEnergy = branches['ChildEnergy'] + childTrackScore = branches['ChildTrackScore'] + trackLength = branches['TrackLength'] + trackWobble = branches['TrackWobble'] + momComparison = branches['TrackMomComparison'] + # ShowerVar lists + displacement = branches['ShowerDisplacement'] + dca = branches['ShowerDCA'] + trackStubLength = branches['ShowerTrackStubLength'] + nuVertexAvSeparation = branches['ShowerNuVertexAvSeparation'] + nuVertexChargeAsymmetry = branches['ShowerNuVertexChargeAsymmetry'] + # Truth + particlePDG = branches['TruePDG'] + # Misc + endX = branches['RecoEndX'] + endY = branches['RecoEndY'] + endZ = branches['RecoEndZ'] + isPrimary = branches['IsPrimary'] + del branches + distToEdge = np.min( + np.stack([ + np.abs(endX - detector['MinX']), + np.abs(endX - detector['MaxX']), + np.abs(endY - detector['MinY']), + np.abs(endY - detector['MaxY']), + np.abs(endZ - detector['MinZ']), + np.abs(endZ - detector['MaxZ']) + ], axis=0), + axis=0) + nEntries = particlePDG.shape[0] + + ################################### + # Only collect taget PDGs + ################################### + target_mask = np.isin(np.abs(particlePDG), args.pdgs) + + startGridU = startGridU[target_mask] + startGridV = startGridV[target_mask] + startGridW = startGridW[target_mask] + endGridU = endGridU[target_mask] + endGridV = endGridV[target_mask] + endGridW = endGridW[target_mask] + nHits2D = nHits2D[target_mask] + trackScore = trackScore[target_mask] + distToEdge = distToEdge[target_mask] + endX = endX[target_mask] + endY = endY[target_mask] + endZ = endZ[target_mask] + nTrackChildren = nTrackChildren[target_mask] + nShowerChildren = nShowerChildren[target_mask] + nGrandChildren = nGrandChildren[target_mask] + nChildHits = nChildHits[target_mask] + childEnergy = childEnergy[target_mask] + childTrackScore = childTrackScore[target_mask] + trackLength = trackLength[target_mask] + trackWobble = trackWobble[target_mask] + momComparison = momComparison[target_mask] + displacement = displacement[target_mask] + dca = dca[target_mask] + trackStubLength = trackStubLength[target_mask] + nuVertexAvSeparation = nuVertexAvSeparation[target_mask] + nuVertexChargeAsymmetry = nuVertexChargeAsymmetry[target_mask] + particlePDG = particlePDG[target_mask] + isPrimary = isPrimary[target_mask] + nEntries = len(particlePDG) + + # Refinement of the particlePDG vector + print('We have ', str(nEntries), ' PFParticles overall!') + print('nMuons: ', np.count_nonzero(abs(particlePDG) == 13)) + print('nProtons: ', np.count_nonzero(abs(particlePDG) == 2212)) + print('nPions: ', np.count_nonzero(abs(particlePDG) == 211)) + print('nKaons: ', np.count_nonzero(abs(particlePDG) == 321)) + print('nElectrons: ', np.count_nonzero(abs(particlePDG) == 11)) + print('nPhotons: ', np.count_nonzero(abs(particlePDG) == 22)) + + # Handle grids + # Work out validity (invalid = 0) + startGridU_valid = startGridU > 1e-7 + startGridV_valid = startGridV > 1e-7 + startGridW_valid = startGridW > 1e-7 + endGridU_valid = endGridU > 1e-7 + endGridV_valid = endGridV > 1e-7 + endGridW_valid = endGridW > 1e-7 + + # Log energy values + startGridU[startGridU_valid] = np.log1p(startGridU[startGridU_valid]) + startGridV[startGridV_valid] = np.log1p(startGridV[startGridV_valid]) + startGridW[startGridW_valid] = np.log1p(startGridW[startGridW_valid]) + endGridU[endGridU_valid] = np.log1p(endGridU[endGridU_valid]) + endGridV[endGridV_valid] = np.log1p(endGridV[endGridV_valid]) + endGridW[endGridW_valid] = np.log1p(endGridW[endGridW_valid]) + + # Normalise them + print('--------------------------------------------------') + print(f'startGridU mean: {np.mean(startGridU[startGridU_valid]):.4f}') + print(f'startGridU std: {np.std(startGridU[startGridU_valid]):.4f}') + print('--------------------------------------------------') + print(f'startGridV mean: {np.mean(startGridV[startGridV_valid]):.4f}') + print(f'startGridV std: {np.std(startGridV[startGridV_valid]):.4f}') + print('--------------------------------------------------') + print(f'startGridW mean: {np.mean(startGridW[startGridW_valid]):.4f}') + print(f'startGridW std: {np.std(startGridW[startGridW_valid]):.4f}') + print('--------------------------------------------------') + print(f'endGridU mean: {np.mean(endGridU[endGridU_valid]):.4f}') + print(f'endGridU std: {np.std(endGridU[endGridU_valid]):.4f}') + print('--------------------------------------------------') + print(f'endGridV mean: {np.mean(endGridV[endGridV_valid]):.4f}') + print(f'endGridV std: {np.std(endGridV[endGridV_valid]):.4f}') + print('--------------------------------------------------') + print(f'endGridW mean: {np.mean(endGridW[endGridW_valid]):.4f}') + print(f'endGridW std: {np.std(endGridW[endGridW_valid]):.4f}') + print('--------------------------------------------------') + startGridU[startGridU_valid] = (startGridU[startGridU_valid] - Normalisation.grid_mean) / Normalisation.grid_std + startGridV[startGridV_valid] = (startGridV[startGridV_valid] - Normalisation.grid_mean) / Normalisation.grid_std + startGridW[startGridW_valid] = (startGridW[startGridW_valid] - Normalisation.grid_mean) / Normalisation.grid_std + endGridU[endGridU_valid] = (endGridU[endGridU_valid] - Normalisation.grid_mean) / Normalisation.grid_std + endGridV[endGridV_valid] = (endGridV[endGridV_valid] - Normalisation.grid_mean) / Normalisation.grid_std + endGridW[endGridW_valid] = (endGridW[endGridW_valid] - Normalisation.grid_mean) / Normalisation.grid_std + + # PFP vars + # Normalise them + print('--------------------------------------------------') + print(f'nHits2D mean: {np.mean(nHits2D):.4f}') + print(f'nHits2D std: {np.std(nHits2D):.4f}') + print('--------------------------------------------------') + print(f'trackScore mean: {np.mean(trackScore):.4f}') + print(f'trackScore std: {np.std(trackScore):.4f}') + print('--------------------------------------------------') + print(f'distToEdge mean: {np.mean(distToEdge):.4f}') + print(f'distToEdge std: {np.std(distToEdge):.4f}') + print('--------------------------------------------------') + + nHits2D = (nHits2D - Normalisation.nHits2D_mean) / Normalisation.nHits2D_std + trackScore = (trackScore - Normalisation.trackScore_mean) / Normalisation.trackScore_std + distToEdge = (distToEdge - Normalisation.distToEdge_mean) / Normalisation.distToEdge_std + + # Track vars + # Work out validity (invalid = -1) + nTrackChildren_valid = nTrackChildren > -0.5 + nShowerChildren_valid = nShowerChildren > -0.5 + nGrandChildren_valid = nGrandChildren > -0.5 + nChildHits_valid = nChildHits > -0.5 + childEnergy_valid = childEnergy > -0.5 + childTrackScore_valid = childTrackScore > -0.5 + trackLength_valid = trackLength > -0.5 + trackWobble_valid = trackWobble > -0.5 + momComparison_valid = momComparison > -0.5 + + # Normalise + print('--------------------------------------------------') + print(f'nTrackChildren mean: {np.mean(nTrackChildren[nTrackChildren_valid]):.4f}') + print(f'nTrackChildren std: {np.std(nTrackChildren[nTrackChildren_valid]):.4f}') + print('--------------------------------------------------') + print(f'nShowerChildren mean: {round(float(np.mean(nShowerChildren[nShowerChildren_valid])), 4)}') + print(f'nShowerChildren std: {round(float(np.std(nShowerChildren[nShowerChildren_valid])), 4)}') + print('--------------------------------------------------') + print(f'nGrandChildren mean: {round(float(np.mean(nGrandChildren[nGrandChildren_valid])), 4)}') + print(f'nGrandChildren std: {round(float(np.std(nGrandChildren[nGrandChildren_valid])), 4)}') + print('--------------------------------------------------') + print(f'nChildHits mean: {round(float(np.mean(nChildHits[nChildHits_valid])), 4)}') + print(f'nChildHits std: {round(float(np.std(nChildHits[nChildHits_valid])), 4)}') + print('--------------------------------------------------') + print(f'childEnergy mean: {round(float(np.mean(childEnergy[childEnergy_valid])), 4)}') + print(f'childEnergy std: {round(float(np.std(childEnergy[childEnergy_valid])), 4)}') + print('--------------------------------------------------') + print(f'childTrackScore mean: {round(float(np.mean(childTrackScore[childTrackScore_valid])), 4)}') + print(f'childTrackScore std: {round(float(np.std(childTrackScore[childTrackScore_valid])), 4)}') + print('--------------------------------------------------') + print(f'trackLength mean: {round(float(np.mean(trackLength[trackLength_valid])), 4)}') + print(f'trackLength std: {round(float(np.std(trackLength[trackLength_valid])), 4)}') + print('--------------------------------------------------') + print(f'trackWobble mean: {round(float(np.mean(trackWobble[trackWobble_valid])), 4)}') + print(f'trackWobble std: {round(float(np.std(trackWobble[trackWobble_valid])), 4)}') + print('--------------------------------------------------') + print(f'momComparison mean: {round(float(np.mean(momComparison[momComparison_valid])), 4)}') + print(f'momComparison std: {round(float(np.std(momComparison[momComparison_valid])), 4)}') + print('--------------------------------------------------') + + nTrackChildren[nTrackChildren_valid] = (nTrackChildren[nTrackChildren_valid] - Normalisation.nTrackChildren_mean) / Normalisation.nTrackChildren_std + nShowerChildren[nShowerChildren_valid] = (nShowerChildren[nShowerChildren_valid] - Normalisation.nShowerChildren_mean) / Normalisation.nShowerChildren_std + nGrandChildren[nGrandChildren_valid] = (nGrandChildren[nGrandChildren_valid] - Normalisation.nGrandChildren_mean) / Normalisation.nGrandChildren_std + nChildHits[nChildHits_valid] = (nChildHits[nChildHits_valid] - Normalisation.nChildHits_mean) / Normalisation.nChildHits_std + childEnergy[childEnergy_valid] = (childEnergy[childEnergy_valid] - Normalisation.childEnergy_mean) / Normalisation.childEnergy_std + childTrackScore[childTrackScore_valid] = (childTrackScore[childTrackScore_valid] - Normalisation.childTrackScore_mean) / Normalisation.childTrackScore_std + trackLength[trackLength_valid] = (trackLength[trackLength_valid] - Normalisation.trackLength_mean) / Normalisation.trackLength_std + trackWobble[trackWobble_valid] = (trackWobble[trackWobble_valid] - Normalisation.trackWobble_mean) / Normalisation.trackWobble_std + momComparison[momComparison_valid] = (momComparison[momComparison_valid] - Normalisation. momComparison_mean) / Normalisation.momComparison_std + + # Shower vars + # Work out validity (invalid = -1) + displacement_valid = displacement > -0.5 + dca_valid = dca > -0.5 + trackStubLength_valid = trackStubLength > -0.5 + nuVertexAvSeparation_valid = nuVertexAvSeparation > -0.5 + nuVertexChargeAsymmetry_valid = nuVertexChargeAsymmetry > -0.5 + + # Normalise + print(f'displacement mean: {round(float(np.mean(displacement[displacement_valid])), 4)}') + print(f'displacement std: {round(float(np.std(displacement[displacement_valid])), 4)}') + print('--------------------------------------------------') + print(f'dca mean: {round(float(np.mean(dca[dca_valid])), 4)}') + print(f'dca std: {round(float(np.std(dca[dca_valid])), 4)}') + print('--------------------------------------------------') + print(f'trackStubLength mean: {round(float(np.mean(trackStubLength[trackStubLength_valid])), 4)}') + print(f'trackStubLength std: {round(float(np.std(trackStubLength[trackStubLength_valid])), 4)}') + print('--------------------------------------------------') + print(f'nuVertexAvSeparation mean: {round(float(np.mean(nuVertexAvSeparation[nuVertexAvSeparation_valid])), 4)}') + print(f'nuVertexAvSeparation std: {round(float(np.std(nuVertexAvSeparation[nuVertexAvSeparation_valid])), 4)}') + print('--------------------------------------------------') + print(f'nuVertexChargeAsymmetry mean: {round(float(np.mean(nuVertexChargeAsymmetry[nuVertexChargeAsymmetry_valid])), 4)}') + print(f'nuVertexChargeAsymmetry std: {round(float(np.std(nuVertexChargeAsymmetry[nuVertexChargeAsymmetry_valid])), 4)}') + print('--------------------------------------------------') + + displacement[displacement_valid] = (displacement[displacement_valid] - Normalisation.displacement_mean) / Normalisation.displacement_std + dca[dca_valid] = (dca[dca_valid] - Normalisation.dca_mean) / Normalisation.dca_std + trackStubLength[trackStubLength_valid] = (trackStubLength[trackStubLength_valid] - Normalisation.trackStubLength_mean) / Normalisation.trackStubLength_std + nuVertexAvSeparation[nuVertexAvSeparation_valid] = (nuVertexAvSeparation[nuVertexAvSeparation_valid] - Normalisation.nuVertexAvSeparation_mean) / Normalisation.nuVertexAvSeparation_std + nuVertexChargeAsymmetry[nuVertexChargeAsymmetry_valid] = (nuVertexChargeAsymmetry[nuVertexChargeAsymmetry_valid] - Normalisation.nuVertexChargeAsymmetry_mean) / Normalisation.nuVertexChargeAsymmetry_std + + # Convert to expected format + dimensions = args.dimensions + startGridU = startGridU.reshape((nEntries, dimensions, dimensions, 1)) + startGridV = startGridV.reshape((nEntries, dimensions, dimensions, 1)) + startGridW = startGridW.reshape((nEntries, dimensions, dimensions, 1)) + endGridU = endGridU.reshape((nEntries, dimensions, dimensions, 1)) + endGridV = endGridV.reshape((nEntries, dimensions, dimensions, 1)) + endGridW = endGridW.reshape((nEntries, dimensions, dimensions, 1)) + startGridU_valid = startGridU_valid.astype(np.float32).reshape((nEntries, dimensions, dimensions, 1)) + startGridV_valid = startGridV_valid.astype(np.float32).reshape((nEntries, dimensions, dimensions, 1)) + startGridW_valid = startGridW_valid.astype(np.float32).reshape((nEntries, dimensions, dimensions, 1)) + endGridU_valid = endGridU_valid.astype(np.float32).reshape((nEntries, dimensions, dimensions, 1)) + endGridV_valid = endGridV_valid.astype(np.float32).reshape((nEntries, dimensions, dimensions, 1)) + endGridW_valid = endGridW_valid.astype(np.float32).reshape((nEntries, dimensions, dimensions, 1)) + nHits2D = nHits2D.reshape((nEntries, 1)) + trackScore = trackScore.reshape((nEntries, 1)) + nTrackChildren = nTrackChildren.reshape((nEntries, 1)) + distToEdge = distToEdge.reshape((nEntries,1)) + isPrimary = isPrimary.astype(np.float32).reshape((nEntries, 1)) + nShowerChildren = nShowerChildren.reshape((nEntries, 1)) + nGrandChildren = nGrandChildren.reshape((nEntries, 1)) + nChildHits = nChildHits.reshape((nEntries, 1)) + childEnergy = childEnergy.reshape((nEntries, 1)) + childTrackScore = childTrackScore.reshape((nEntries, 1)) + trackLength = trackLength.reshape((nEntries, 1)) + trackWobble = trackWobble.reshape((nEntries, 1)) + momComparison = momComparison.reshape((nEntries, 1)) + nTrackChildren_valid = nTrackChildren_valid.astype(np.float32).reshape((nEntries, 1)) + nShowerChildren_valid = nShowerChildren_valid.astype(np.float32).reshape((nEntries, 1)) + nGrandChildren_valid = nGrandChildren_valid.astype(np.float32).reshape((nEntries, 1)) + nChildHits_valid = nChildHits_valid.astype(np.float32).reshape((nEntries, 1)) + childEnergy_valid = childEnergy_valid.astype(np.float32).reshape((nEntries, 1)) + childTrackScore_valid = childTrackScore_valid.astype(np.float32).reshape((nEntries, 1)) + trackLength_valid = trackLength_valid.astype(np.float32).reshape((nEntries, 1)) + trackWobble_valid = trackWobble_valid.astype(np.float32).reshape((nEntries, 1)) + momComparison_valid = momComparison_valid.astype(np.float32).reshape((nEntries, 1)) + displacement = displacement.reshape((nEntries, 1)) + dca = dca.reshape((nEntries, 1)) + trackStubLength = trackStubLength.reshape((nEntries, 1)) + nuVertexAvSeparation = nuVertexAvSeparation.reshape((nEntries, 1)) + nuVertexChargeAsymmetry = nuVertexChargeAsymmetry.reshape((nEntries, 1)) + displacement_valid = displacement_valid.astype(np.float32).reshape((nEntries, 1)) + dca_valid = dca_valid.astype(np.float32).reshape((nEntries, 1)) + trackStubLength_valid = trackStubLength_valid.astype(np.float32).reshape((nEntries, 1)) + nuVertexAvSeparation_valid = nuVertexAvSeparation_valid.astype(np.float32).reshape((nEntries, 1)) + nuVertexChargeAsymmetry_valid = nuVertexChargeAsymmetry_valid.astype(np.float32).reshape((nEntries, 1)) + particlePDG = particlePDG.reshape((nEntries, 1)) + endX = endX.reshape((nEntries, 1)) + endY = endY.reshape((nEntries, 1)) + endZ = endZ.reshape((nEntries, 1)) + + pfpVars = np.concatenate((endX, endY, endZ), axis=1) + trackVars = np.concatenate((nTrackChildren, nTrackChildren_valid, + nShowerChildren, nShowerChildren_valid, + nGrandChildren, nGrandChildren_valid, + nChildHits, nChildHits_valid, + childEnergy, childEnergy_valid, + childTrackScore, childTrackScore_valid, + trackLength, trackLength_valid, + trackWobble, trackWobble_valid, + momComparison, momComparison_valid, nHits2D, trackScore, distToEdge, isPrimary), axis=1) + showerVars = np.concatenate((displacement, displacement_valid, + dca, dca_valid, + trackStubLength, trackStubLength_valid, + nuVertexAvSeparation, nuVertexAvSeparation_valid, + nuVertexChargeAsymmetry, nuVertexChargeAsymmetry_valid), axis=1) + + # muons = 0, protons = 1, pions = 2, kaons = 3, electrons = 4, photons = 5 + particlePDG[abs(particlePDG) == 13] = 0 + particlePDG[abs(particlePDG) == 2212] = 1 + particlePDG[abs(particlePDG) == 211] = 2 + # particlePDG[abs(particlePDG) == 321] = 3 + particlePDG[abs(particlePDG) == 11] = 3 + particlePDG[abs(particlePDG) == 22] = 4 + y = to_categorical(particlePDG, 5) + + return startGridU, startGridU_valid, startGridV, startGridV_valid, startGridW, startGridW_valid, endGridU, endGridU_valid, endGridV, endGridV_valid, endGridW, endGridW_valid, pfpVars, trackVars, showerVars, y + diff --git a/ml/networks/training/ivysaurus_pid/Normalisation.py b/ml/networks/training/ivysaurus_pid/Normalisation.py new file mode 100644 index 0000000..5d8d64b --- /dev/null +++ b/ml/networks/training/ivysaurus_pid/Normalisation.py @@ -0,0 +1,41 @@ +# Grid +grid_mean = 0.0021 +grid_std = 0.0027 +# PFPVars +nHits2D_mean = 592.2099 +nHits2D_std = 1215.6621 +trackScore_mean = 0.5912 +trackScore_std = 0.1593 +distToEdge_mean = 104.0723 +distToEdge_std = 91.5137 +# TrackVars +nTrackChildren_mean = 0.1725 +nTrackChildren_std = 0.4825 +nShowerChildren_mean = 0.1007 +nShowerChildren_std = 0.3474 +nGrandChildren_mean = 0.0283 +nGrandChildren_std = 0.2066 +nChildHits_mean = 271.294 +nChildHits_std = 656.6655 +childEnergy_mean = 0.2317 +childEnergy_std = 0.8004 +childTrackScore_mean = 0.56 +childTrackScore_std = 0.1305 +trackLength_mean = 189.9965 +trackLength_std = 1135.1893 +trackWobble_mean = 15.5898 +trackWobble_std = 12.2151 +momComparison_mean = 5.5677 +momComparison_std = 4.2623 +# ShowerVars +displacement_mean = 35.7004 +displacement_std = 82.8258 +dca_mean = 23.7716 +dca_std = 722.3403 +trackStubLength_mean = 38.2381 +trackStubLength_std = 56.1333 +nuVertexAvSeparation_mean = 22.6225 +nuVertexAvSeparation_std = 49.5967 +nuVertexChargeAsymmetry_mean = 0.9123 +nuVertexChargeAsymmetry_std = 0.1745 + diff --git a/ml/networks/training/ivysaurus_pid/ivysaurus_model.py b/ml/networks/training/ivysaurus_pid/ivysaurus_model.py new file mode 100644 index 0000000..2fb75cd --- /dev/null +++ b/ml/networks/training/ivysaurus_pid/ivysaurus_model.py @@ -0,0 +1,184 @@ +import torch +import torch.nn as nn +import torch.nn.functional as F + +######################################################################################################## + +class ViewScaler(nn.Module): + def __init__(self): + super().__init__() + self.scale = nn.Parameter(torch.ones(1)) + self.shift = nn.Parameter(torch.zeros(1)) + + def forward(self, x): + return x * self.scale + self.shift + +######################################################################################################## + +class ResidualBlock(nn.Module): + def __init__(self, in_channels, filters, kernel_size=3): + super().__init__() + pad = kernel_size // 2 + self.conv1 = nn.Conv2d(in_channels, filters, kernel_size, padding=pad, bias=False) + self.bn1 = nn.BatchNorm2d(filters) + self.conv2 = nn.Conv2d(filters, filters, kernel_size, padding=pad, bias=False) + self.bn2 = nn.BatchNorm2d(filters) + self.dropout = nn.Dropout2d(0.1) + + # Match channels if needed + self.project = None + if in_channels != filters: + self.project = nn.Sequential( + nn.Conv2d(in_channels, filters, 1, bias=False), + nn.BatchNorm2d(filters), + ) + + def forward(self, x): + shortcut = x + out = F.relu(self.bn1(self.conv1(x))) + out = self.bn2(self.conv2(out)) + + if self.project is not None: + shortcut = self.project(shortcut) + + out = out + shortcut + out = F.relu(out) + out = self.dropout(out) + + return out + +######################################################################################################## + +class SpatialAttention(nn.Module): + def __init__(self, kernel_size=3): + super().__init__() + padding = kernel_size // 2 + self.conv = nn.Conv2d(2, 1, kernel_size, padding=padding, bias=False) + self.sigmoid = nn.Sigmoid() + + def forward(self, x): + avg_out = torch.mean(x, dim=1, keepdim=True) + max_out = torch.amax(x, dim=1, keepdim=True) + combined = torch.cat([avg_out, max_out], dim=1) + attention_map = self.sigmoid(self.conv(combined)) + + return attention_map + +######################################################################################################## + +class SharedEncoder(nn.Module): + def __init__(self): + super().__init__() + + self.stem = nn.Sequential( + nn.Conv2d(2, 32, 3, padding=1, bias=False), + nn.BatchNorm2d(32), + nn.ReLU(inplace=True), + ) + + self.block1 = nn.Sequential( + ResidualBlock(32, 32), + ) + self.pool1 = nn.MaxPool2d(2) + + self.block2 = nn.Sequential( + ResidualBlock(32, 64), + ) + + self.block3 = nn.Sequential( + ResidualBlock(64, 128), + ) + + self.spatial_attention = SpatialAttention() + + def forward(self, x): + x = self.stem(x) + x = self.block1(x) + x = self.pool1(x) + x = self.block2(x) + x = self.block3(x) + + attention_weights = self.spatial_attention(x) + x = x * (1 + attention_weights) + + gap = x.mean(dim=(2,3)) + gmp = x.amax(dim=(2,3)) + std = x.std((2,3)) + x = torch.cat([gap, gmp, std], dim=1) + + return x + +######################################################################################################## + +class IvysaurusModel(nn.Module): + def __init__(self, dimensions, nclasses, nTrackVars, nShowerVars): + super().__init__() + + self.encoder = SharedEncoder() # shared across all views & start/end + + # One scaler per view + self.scalerU = ViewScaler() + self.scalerV = ViewScaler() + self.scalerW = ViewScaler() + + # Each branch: 2 input channels for start and end grid in each of three views + encoder_out = 128 * 2 * 2 * 3 + combined_feat = encoder_out + nTrackVars + nShowerVars + + self.head = nn.Sequential( + nn.Linear(combined_feat, 256, bias=False), + nn.BatchNorm1d(256), + nn.ReLU(inplace=True), + nn.Dropout(0.3), + nn.Linear(256, 128, bias=False), + nn.BatchNorm1d(128), + nn.ReLU(inplace=True), + nn.Dropout(0.3)) + self.out = nn.Linear(128, nclasses) + + def _branch(self, start_scaled, start_mask, end_scaled, end_mask): + start_combined = torch.cat([start_scaled, start_mask], dim=1) # (N, 2, H, W) + end_combined = torch.cat([end_scaled, end_mask], dim=1) # (N, 2, H, W) + + start_feat = self.encoder(start_combined) # (N, 128) + end_feat = self.encoder(end_combined) # (N, 128) + + return torch.cat([start_feat, end_feat], dim=1) # (N, 256) + + def forward(self, + startU, startU_mask, endU, endU_mask, + startV, startV_mask, endV, endV_mask, + startW, startW_mask, endW, endW_mask, + trackVars, showerVars): + + startU = startU.permute(0, 3, 1, 2) + startV = startV.permute(0, 3, 1, 2) + startW = startW.permute(0, 3, 1, 2) + endU = endU.permute(0, 3, 1, 2) + endV = endV.permute(0, 3, 1, 2) + endW = endW.permute(0, 3, 1, 2) + startU_mask = startU_mask.permute(0, 3, 1, 2) + startV_mask = startV_mask.permute(0, 3, 1, 2) + startW_mask = startW_mask.permute(0, 3, 1, 2) + endU_mask = endU_mask.permute(0, 3, 1, 2) + endV_mask = endV_mask.permute(0, 3, 1, 2) + endW_mask = endW_mask.permute(0, 3, 1, 2) + + startU = self.scalerU(startU) + endU = self.scalerU(endU) + startV = self.scalerV(startV) + endV = self.scalerV(endV) + startW = self.scalerW(startW) + endW = self.scalerW(endW) + + branchU = self._branch(startU, startU_mask, endU, endU_mask) + branchV = self._branch(startV, startV_mask, endV, endV_mask) + branchW = self._branch(startW, startW_mask, endW, endW_mask) + + combined = torch.cat([branchU, branchV, branchW, + trackVars, showerVars], dim=1) + + combined = self.head(combined) + logits = self.out(combined) + return logits + diff --git a/ml/networks/training/ivysaurus_pid/train.py b/ml/networks/training/ivysaurus_pid/train.py new file mode 100644 index 0000000..7c7e07b --- /dev/null +++ b/ml/networks/training/ivysaurus_pid/train.py @@ -0,0 +1,238 @@ +import argparse +import glob +import os + +import numpy as np +import torch +import torch.nn as nn +from torch.utils.data import Dataset, DataLoader, WeightedRandomSampler +from sklearn.metrics import balanced_accuracy_score +from sklearn.metrics import classification_report +from sklearn.metrics import confusion_matrix + + +from ivysaurus_model import IvysaurusModel + +# RUN ON GPU:1 +os.environ["CUDA_VISIBLE_DEVICES"] = "1" + +######################################################################################################## + +class IvysaurusDataset(Dataset): + def __init__(self, grids, trackVars, showerVars, y): + # grids: dict of name -> (N, H, W, 1) numpy arrays + self.grids = {k: torch.from_numpy(np.asarray(v, dtype=np.float32)) for k, v in grids.items()} # + self.trackVars = torch.from_numpy(trackVars.astype(np.float32)) + self.showerVars = torch.from_numpy(showerVars.astype(np.float32)) + # CrossEntropyLoss wants class indices, not one-hot + self.labels = torch.from_numpy(np.argmax(y, axis=1).astype(np.int64)) + + def __len__(self): + return self.labels.shape[0] + + def __getitem__(self, i): + sample = {k: v[i] for k, v in self.grids.items()} + sample["trackVars"] = self.trackVars[i] + sample["showerVars"] = self.showerVars[i] + return sample, self.labels[i] + +######################################################################################################## + +GRID_KEYS = [ + "startU", "startU_mask", "endU", "endU_mask", + "startV", "startV_mask", "endV", "endV_mask", + "startW", "startW_mask", "endW", "endW_mask", +] + +######################################################################################################## + +def run_model(model, batch, device): + args = [batch[key].to(device) for key in GRID_KEYS] + args.append(batch["trackVars"].to(device)) + args.append(batch["showerVars"].to(device)) + return model(*args) + +######################################################################################################## + +def main(args): + torch.manual_seed(42) + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + print("Using device:", device) + + # Load data + suffix = "Contained" if args.is_contained else "Exiting" if args.is_exiting else "" + trainFileNames = glob.glob(f'{args.input_dir}/*_{suffix}*.npz') + + grids_train = {key: [] for key in GRID_KEYS} + grids_test = {key: [] for key in GRID_KEYS} + trackVars_train = [] + showerVars_train = [] + trackVars_test = [] + showerVars_test = [] + y_train = [] + y_test = [] + + for fname in trainFileNames: + print(f"Reading file: {fname}") + + with np.load(fname) as data: + for key in GRID_KEYS: + grids_train[key].append(data[f"{key}_train"]) + grids_test[key].append(data[f"{key}_test"]) + + trackVars_train.append(data['trackVars_train']) + trackVars_test.append(data['trackVars_test']) + showerVars_train.append(data['showerVars_train']) + showerVars_test.append(data['showerVars_test']) + y_train.append(data['y_train']) + y_test.append(data['y_test']) + + for key in GRID_KEYS: + grids_train[key] = np.concatenate(grids_train[key], axis=0) + grids_test[key] = np.concatenate(grids_test[key], axis=0) + + trackVars_train = np.concatenate(trackVars_train, axis=0) + trackVars_test = np.concatenate(trackVars_test, axis=0) + showerVars_train = np.concatenate(showerVars_train, axis=0) + showerVars_test = np.concatenate(showerVars_test, axis=0) + y_train = np.concatenate(y_train, axis=0) + y_test = np.concatenate(y_test, axis=0) + + # Work out some network shapes + n_classes = y_train.shape[1] + n_track_vars = trackVars_train.shape[1] + n_shower_vars = showerVars_train.shape[1] + dimensions = grids_train[GRID_KEYS[0]].shape[1] + + print('n_classes:', n_classes) + print('n_track_vars:', n_track_vars) + print('n_shower_vars:', n_shower_vars) + print("y_train:", y_train.shape, "y_test:", y_test.shape) + print('Train') + print(np.unique(np.argmax(y_train, axis=1), return_counts=True)) + print('Test') + print(np.unique(np.argmax(y_test, axis=1), return_counts=True)) + + train_ds = IvysaurusDataset(grids_train, trackVars_train, showerVars_train, y_train) + test_ds = IvysaurusDataset(grids_test, trackVars_test, showerVars_test, y_test) + train_loader = DataLoader(train_ds, batch_size=args.batch_size, shuffle=True, num_workers=4, pin_memory=True) + test_loader = DataLoader(test_ds, batch_size=args.batch_size, shuffle=False, num_workers=4, pin_memory=True) + + # Class weights + particle_type = np.argmax(y_train, axis=1) + counts = [np.count_nonzero(particle_type == c) for c in range(n_classes)] + print("Class Counts:", counts) + maxParticle = max(counts) + classWeights = np.sqrt(np.array([maxParticle / c for c in counts], dtype=np.float32)) + print("Class Weights:", classWeights) + class_weights_t = torch.from_numpy(classWeights).to(device) + + # Model stuff + model = IvysaurusModel(dimensions, n_classes, n_track_vars, n_shower_vars).to(device) + optimiser = torch.optim.AdamW(model.parameters(), lr=args.learning_rate, weight_decay=1e-4) + criterion = nn.CrossEntropyLoss(weight=class_weights_t) + scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimiser, mode='max', factor=0.5, patience=1) + model_path = f'{args.output_dir}/my_model_{"contained" if args.is_contained else "exiting" if args.is_exiting else "all"}.pt' + + # Training loop + best_val_acc = -1.0 + for epoch in range(args.n_epochs): + # Train + model.train() + all_preds_train = [] + all_labels_train = [] + train_loss, train_correct, train_total = 0.0, 0, 0 + + for batch, labels in train_loader: + labels = labels.to(device) + optimiser.zero_grad() + logits = run_model(model, batch, device) + loss = criterion(logits, labels) + loss.backward() + optimiser.step() + + preds = logits.argmax(1) + all_preds_train.append(preds.cpu().numpy()) + all_labels_train.append(labels.cpu().numpy()) + + train_loss += loss.item() * labels.size(0) + train_correct += (logits.argmax(1) == labels).sum().item() + train_total += labels.size(0) + + all_preds_train = np.concatenate(all_preds_train) + all_labels_train = np.concatenate(all_labels_train) + train_bal_acc = balanced_accuracy_score(all_preds_train, all_labels_train) + train_loss /= train_total + train_acc = train_correct / train_total + + # Validate + model.eval() + all_preds = [] + all_labels = [] + val_loss, val_correct, val_total = 0.0, 0, 0 + + with torch.no_grad(): + for batch, labels in test_loader: + labels = labels.to(device) + logits = run_model(model, batch, device) + loss = criterion(logits, labels) + + preds = logits.argmax(1) + all_preds.append(preds.cpu().numpy()) + all_labels.append(labels.cpu().numpy()) + + val_loss += loss.item() * labels.size(0) + val_correct += (preds == labels).sum().item() + val_total += labels.size(0) + + all_preds = np.concatenate(all_preds) + all_labels = np.concatenate(all_labels) + val_bal_acc = balanced_accuracy_score(all_labels, all_preds) + val_loss /= val_total + val_acc = val_correct / val_total + + print(f"Epoch {epoch+1}/{args.n_epochs} - " + f"loss: {train_loss:.4f} - acc: {train_acc:.4f} - train_bal_acc: {train_bal_acc:.4f} - " + f"val_loss: {val_loss:.4f} - val_acc: {val_acc:.4f} - val_bal_acc: {val_bal_acc:.4f}") + class_names = ["muon", "proton", "pion", "electron", "photon"] + print(classification_report(all_labels, all_preds, target_names=class_names, digits=4)) + cm = confusion_matrix(all_labels, all_preds) + + print("Confusion Matrix:") + print(" Predicted") + print(" " + " ".join(f"{name:>9}" for name in class_names)) + for i, row in enumerate(cm): + print(f"{class_names[i]:>10} " + " ".join(f"{x:9d}" for x in row)) + + scheduler.step(val_bal_acc) + + # checkpoint: save best on val_acc + if val_bal_acc > best_val_acc: + best_val_acc = val_bal_acc + model_cpu = model.cpu() + model_cpu.eval() # just to make sure + scripted = torch.jit.script(model_cpu) + scripted.save(model_path) + model.to(device) + print(f" val_acc improved to {val_bal_acc:.4f}, saved model to {model_path}") + +######################################################################################################## + +def parse_cli(): + parser = argparse.ArgumentParser(description="Ivysaurus PID") + parser.add_argument("--input_dir", type=str, required=True, help="Input file directory") + parser.add_argument("--output_dir", type=str, required=True, help="Dir to save model") + parser.add_argument("--is_contained", action="store_true", help="Training for contained particles?") + parser.add_argument("--is_exiting", action="store_true", help="Training for exiting particles?") + parser.add_argument("--n_epochs", type=int, default=10, help="Number of epochs, default=10") + parser.add_argument("--batch_size", type=int, default=64, help="Batch size, default=64") + parser.add_argument("--learning_rate", type=float, default=1e-3, help="Learning rate, default=1e-3") + + return parser.parse_args() + +######################################################################################################## + +if __name__ == "__main__": + args = parse_cli() + main(args) +