diff --git a/README.md b/README.md index caef2b2..ba78a5c 100644 --- a/README.md +++ b/README.md @@ -31,7 +31,9 @@ You can get some basic information about the EventNtuple file useing ```checkEve checkEventNtuple file1.root file2.root ``` -This will print the version number (for versions after v6.3.0) and the trigger branches in the file (for versions after v6.8.0) +This will print the version number (for versions after v6.3.0), configured +TrkQual model metadata when available (for versions after v6.13.0), and the trigger branches in the file +(for versions after v6.8.0). ## How to Analyze an EventNtuple diff --git a/bin/checkEventNtuple b/bin/checkEventNtuple index 5899ec4..de82f71 100755 --- a/bin/checkEventNtuple +++ b/bin/checkEventNtuple @@ -3,6 +3,14 @@ import argparse import ROOT +ROOT.gInterpreter.Declare(""" +void printTrkQualMetadata(TH1* metadata) { + for (int bin = 1; bin <= metadata->GetNbinsX(); ++bin) { + std::cout << " - " << metadata->GetXaxis()->GetBinLabel(bin) << std::endl; + } +} +""") + # A main function so that this can be run on the command line def main(): parser = argparse.ArgumentParser( @@ -16,6 +24,7 @@ def main(): foldername = "EventNtuple" histname = "version" treename = "ntuple" + trkqualmetadataname = "trkqual_metadata" for filename in args.filenames: try: f = ROOT.TFile(filename, "READ"); @@ -28,6 +37,12 @@ def main(): else: print("{} histogram does not exist in {} (it is either v06_02_00 or older)".format(histname, filename)) + # Print configured track-quality algorithms, if available. + if folder.GetListOfKeys().Contains(trkqualmetadataname): + metadata = f.Get(foldername+"/"+trkqualmetadataname) + print("\nIt contains the following trkqual branches:", flush=True) + ROOT.printTrkQualMetadata(metadata) + # Print triggers if (folder.GetListOfKeys().Contains(treename)): t = f.Get(foldername+"/"+treename) diff --git a/fcl/README.md b/fcl/README.md index 78deae7..3b2503e 100644 --- a/fcl/README.md +++ b/fcl/README.md @@ -8,6 +8,23 @@ from_tier-type_extra.fcl where ```tier``` is the data tier of the input dataset, ```type``` is the type of dataset (e.g. primary-only, extracted position), and ```extra``` gives some extra information (optional) +## Multiple TrkQual outputs + +Configure each TrkQual result in a track fit with `trkQualLeaves`. `leafname` is +appended to `qual`, so stable output names do not depend on the +order of the configured algorithms: + +``` +trkQualLeaves : [ + { leafname : "" inputTag : "TrkQualAll:ANN" modelVersion : "TrkQual_ANN1_v2" }, + { leafname : "_candidate" inputTag : "TrkQualCandidate:ANN" modelVersion : "TrkQual_ANN1_v3_rc1" } +] +``` + +This writes `trkqual` and `trkqual_candidate`. EventNtuple also records the +track branch, output branch, input tag, and model version in +`EventNtuple/trkqual_metadata`; `checkEventNtuple` prints this metadata. + ## Table of Fcl Files | fcl file | runs on | additional info | @@ -21,7 +38,7 @@ where ```tier``` is the data tier of the input dataset, ```type``` is the type o | from_mcs-ceSimRecoVal.fcl | output of EventNtuple/validation/ceSimReco.fcl | for validating the ```trkhitcalibs``` branch | | from_mcs-mockdata_separateTrkBranches.fcl | mock datasets | example on how to separate the tracks into separate branches again| | from_mcs-mockdata_selectorExample.fcl | mock datasets | example on how to use a selector to select certain types of tracks before putting them into the EventNtuple | -| from_mcs-mixed_trkQualCompare.fcl | reconstructed mixed (i.e. primary+background hits) datasets | shows how to output result of more than one TrkQual | +| from_mcs-mixed_trkQualCompare.fcl | reconstructed mixed (i.e. primary+background hits) datasets | shows explicitly named TrkQual outputs and embedded model-version provenance; requires the listed comparison ONNX models | | from_mcs-primary_addVDSteps.fcl | reconstructed primary (i.e. no background hits) datasets | shows how to add the branch for virtual detector steps | | from_mcs-Run1B.fcl | reconstructed Run-1B (backup plan) datasets | adds the branch for virtual detector steps | | from_mcs-DeMCalib.fcl | reconstructed primary or mixed datasets | only writes one track per event | diff --git a/fcl/from_mcs-mixed_trkQualCompare.fcl b/fcl/from_mcs-mixed_trkQualCompare.fcl index f9a4787..c03300a 100644 --- a/fcl/from_mcs-mixed_trkQualCompare.fcl +++ b/fcl/from_mcs-mixed_trkQualCompare.fcl @@ -1,17 +1,22 @@ #include "EventNtuple/fcl/from_mcs-mockdata.fcl" -# Add another TrackQuality module +# Compare explicitly named TrackQuality branches. The v1.0 and v1.1 ONNX +# files must be installed beside the v2 model before running this example. physics.producers.TrkQualAllV10 : @local::TrkQualAll -physics.producers.TrkQualAllV10.datFilename : "Offline/TrkDiag/data/TrkQual_ANN1_v1.dat" +physics.producers.TrkQualAllV10.onnxFilename : "ArtAnalysis/TrkDiag/data/TrkQual_ANN1_v1.onnx" physics.producers.TrkQualAllV11 : @local::TrkQualAll -physics.producers.TrkQualAllV11.datFilename : "Offline/TrkDiag/data/TrkQual_ANN1_v1.1.dat" +physics.producers.TrkQualAllV11.onnxFilename : "ArtAnalysis/TrkDiag/data/TrkQual_ANN1_v1.1.onnx" physics.producers.TrkQualAllV2 : @local::TrkQualAll -physics.producers.TrkQualAllV2.datFilename : "Offline/TrkDiag/data/TrkQual_ANN1_v2.dat" +physics.producers.TrkQualAllV2.onnxFilename : "ArtAnalysis/TrkDiag/data/TrkQual_ANN1_v2.onnx" physics.EventNtuplePath : [ @sequence::EventNtuple.Path, TrkQualAllV10, TrkQualAllV11, TrkQualAllV2 ] -# Add it to the EventNtuple output -physics.analyzers.EventNtuple.branches[0].trkQualTags : [ "TrkQualAllV10", "TrkQualAllV11", "TrkQualAllV2" ] +# Add explicitly named results and record their model versions in the output. +physics.analyzers.EventNtuple.trk.fits[0].trkQualLeaves : [ + { leafname : "_v1_0" inputTag : "TrkQualAllV10:ANN" modelVersion : "TrkQual_ANN1_v1.0" }, + { leafname : "_v1_1" inputTag : "TrkQualAllV11:ANN" modelVersion : "TrkQual_ANN1_v1.1" }, + { leafname : "_v2" inputTag : "TrkQualAllV2:ANN" modelVersion : "TrkQual_ANN1_v2" } +] -# Removin hits +# Remove hits physics.analyzers.EventNtuple.trk.fillHits : false diff --git a/fcl/from_mcs-reflection.fcl b/fcl/from_mcs-reflection.fcl index f116d4d..3b14237 100644 --- a/fcl/from_mcs-reflection.fcl +++ b/fcl/from_mcs-reflection.fcl @@ -83,14 +83,14 @@ physics.trigger_paths : [ "eTrig", "muTrig"] physics.end_paths : [ "eEnd", "muEnd", "CLPrint" ] physics.analyzers.ENe.trk.fits : [ { @table::ENBranch - trkQualTags : ["TrkQualReflecte"] + trkQualLeaves : [ { leafname : "" inputTag : "TrkQualReflecte:ANN" modelVersion : "TrkQual_ANN1_v2" } ] trkPIDTags : ["TrkPIDReflecte"] input: "Reflecte" } ] physics.analyzers.ENmu.trk.fits : [ { @table::ENBranch - trkQualTags : ["TrkQualReflectmu"] + trkQualLeaves : [ { leafname : "" inputTag : "TrkQualReflectmu:ANN" modelVersion : "TrkQual_ANN1_v2" } ] trkPIDTags : ["TrkPIDReflectmu"] input: "Reflectmu" } diff --git a/fcl/prolog.fcl b/fcl/prolog.fcl index 4b8b68d..1ee3fd1 100644 --- a/fcl/prolog.fcl +++ b/fcl/prolog.fcl @@ -10,6 +10,7 @@ BEGIN_PROLOG TrkQual : { module_type : TrackQuality onnxFilename : "ArtAnalysis/TrkDiag/data/TrkQual_ANN1_v2.onnx" + xgbFilename : "ArtAnalysis/TrkDiag/data/TrkQual_BDT1_v2.0.ubj" debugLevel : 0 } @@ -192,91 +193,94 @@ DeM : { input : "MergeKKDeM" branchname : "dem" fill : true options : { fillMC : true fillHits : false genealogyDepth : -1 matchDepth : -1 } - trkQualTags : ["TrkQualDeM"] + trkQualLeaves : [ { leafname : "" inputTag : "TrkQualDeM:ANN" modelVersion : "ANN1_v2" } ] trkPIDTags : ["TrkPIDDeM"] } UeM : { input : "MergeKKUeM" branchname : "uem" fill : true options : { fillMC : true fillHits : false genealogyDepth : -1 matchDepth : -1 } - trkQualTags : [ ] + trkQualLeaves : [ ] trkPIDTags : [ ] } DmuM : { input : "MergeKKDmuM" branchname : "dmm" fill : true options : { fillMC : true fillHits : false genealogyDepth : -1 matchDepth : -1 } - trkQualTags : [ ] + trkQualLeaves : [ ] trkPIDTags : [ ] } UmuM : { input : "MergeKKUmuM" branchname : "umm" fill : true options : { fillMC : true fillHits : false genealogyDepth : -1 matchDepth : -1 } - trkQualTags : [ ] + trkQualLeaves : [ ] trkPIDTags : [ ] } DeP : { input : "MergeKKDeP" branchname : "dep" fill : true options : { fillMC : true fillHits : false genealogyDepth : -1 matchDepth : -1 } - trkQualTags : [ ] + trkQualLeaves : [ ] trkPIDTags : [ ] } UeP : { input : "MergeKKUeP" branchname : "uep" fill : true options : { fillMC : true fillHits : false genealogyDepth : -1 matchDepth : -1 } - trkQualTags : [ ] + trkQualLeaves : [ ] trkPIDTags : [ ] } DmuP : { input : "MergeKKDmuP" branchname : "dmp" fill : true options : { fillMC : true fillHits : false genealogyDepth : -1 matchDepth : -1 } - trkQualTags : [ ] + trkQualLeaves : [ ] trkPIDTags : [ ] } UmuP : { input : "MergeKKUmuP" branchname : "ump" fill : true options : { fillMC : true fillHits : false genealogyDepth : -1 matchDepth : -1 } - trkQualTags : [ ] + trkQualLeaves : [ ] trkPIDTags : [ ] } Ext : { input : "MergeKKLine" branchname : "trk" fill : true options : { fillMC : true fillHits : true genealogyDepth : -1 matchDepth : -1 } - trkQualTags : [ ] + trkQualLeaves : [ ] trkPIDTags : [ ] } Off : { input : "MergeKKOff" branchname : "trk" fill : true options : { fillMC : true fillHits : true genealogyDepth : -1 matchDepth : -1 } - trkQualTags : [ ] + trkQualLeaves : [ ] trkPIDTags : [ ] } All : { input : "MergeKKAll" branchname : "trk" fill : true options : { fillMC : true fillHits : true genealogyDepth : -1 matchDepth : -1 } - trkQualTags : ["TrkQualAll"] + trkQualLeaves : [ + { leafname : "" inputTag : "TrkQualAll:ANN" modelVersion : "ANN1_v2" }, + { leafname : "_bdt" inputTag : "TrkQualAll:BDT" modelVersion : "BDT1_v2" } + ] trkPIDTags : ["TrkPIDAll"] } DeCalib : { input : "MergeKKDeCalib" branchname : "trk" fill : true options : { fillMC : true fillHits : true genealogyDepth : -1 matchDepth : -1 } - trkQualTags : [ ] + trkQualLeaves : [ ] trkPIDTags : [ ] } De : { input : "MergeKKDe" branchname : "de" fill : true options : { fillMC : true fillHits : false genealogyDepth : -1 matchDepth : -1 } - trkQualTags : ["TrkQualDe"] + trkQualLeaves : [ { leafname : "" inputTag : "TrkQualDe:ANN" modelVersion : "ANN1_v2" } ] trkPIDTags : ["TrkPIDDe"] } @@ -411,7 +415,7 @@ TTMCBranch : { branchname : "trk" fill : true options : { fillMC : true fillHits : true genealogyDepth : -1 matchDepth : -1 } - trkQualTags : [ ] + trkQualLeaves : [ ] trkPIDTags : [ ] } @@ -451,7 +455,7 @@ ENDeCalib.trk.fits : [ branchname : "trk" fill : true options : { fillMC : true fillHits : true genealogyDepth : -1 matchDepth : -1 } - trkQualTags : ["TrkQualDe"] + trkQualLeaves : [ { leafname : "" inputTag : "TrkQualDe:ANN" modelVersion : "ANN1_v2" } ] trkPIDTags : ["TrkPIDDe"] } ] diff --git a/inc/TrkQualMetadata.hh b/inc/TrkQualMetadata.hh new file mode 100644 index 0000000..1a21612 --- /dev/null +++ b/inc/TrkQualMetadata.hh @@ -0,0 +1,14 @@ +#ifndef EventNtuple_TrkQualMetadata_hh +#define EventNtuple_TrkQualMetadata_hh + +#include + +namespace mu2e { + struct TrkQualMetadata { + std::string output_branch; + std::string input_tag; + std::string model_version; + }; +} + +#endif diff --git a/rooutil/README.md b/rooutil/README.md index 54af35a..150c1b7 100644 --- a/rooutil/README.md +++ b/rooutil/README.md @@ -111,6 +111,18 @@ The ```CaloCluster``` class contains all information related to a single calorim Examples: [PlotCaloClusterEnergy.C](./examples/PlotCaloClusterEnergy.C), [PlotCaloClusterEnergy_RecoVsTrue.C](./examples/PlotCaloClusterEnergy_RecoVsTrue.C), [PlotCaloClusterAndHits.C](./examples/PlotCaloClusterAndHits.C), [PlotCaloCluster_SimParticles.C](./examples/PlotCaloCluster_SimParticles.C) +### The ```UserBranch``` Class +There are some branches in EventNtuple that can have their names defined at runtime. For example, additional ```trkqual``` branches can be added in order to compare different trainings or different algorithms. The ```UserBranch``` class handles these sorts of branches. + +Example: [CompareTrkQualTrainings_UserBranches.C](./examples/CompareTrkQualTrainings_UserBranches.C) + +#### TrkQual Branches +RooUtil can bind either output by name with `MakeTrackUserBranch`. +Analysis code can stop on unexpected provenance with: + +``` +util.RequireTrkQualVersion("trkqual", "ANN2_v2"); +``` ### Branches not contained within a class Some branches are not contained in any of the above classes: diff --git a/rooutil/examples/CompareTrkQualTrainings.C b/rooutil/examples/CompareTrkQualTrainings.C deleted file mode 100644 index dbfe4cd..0000000 --- a/rooutil/examples/CompareTrkQualTrainings.C +++ /dev/null @@ -1,170 +0,0 @@ -// -// An example of how to plot the momentum of electrons at the tracker entrance -// This uses cut functions defined in common_cuts.hh -// - -#include "EventNtuple/rooutil/inc/RooUtil.hh" -#include "EventNtuple/rooutil/inc/common_cuts.hh" - -#include "TH2F.h" -#include "TCanvas.h" -#include "TLine.h" -#include "TLatex.h" - -using namespace rooutil; - -void CompareTrkQualTrainings(std::string filename) { - - bool save_plots = false; - std::string plotsdir = "/exp/mu2e/app/users/edmonds/plots/2025-09-25/"; - - // Create the histogram you want to fill - TH2F* hTrkQual_OldVsNew = new TH2F("hTrkQual_OldVsNew", "", 100,0,1, 100,0,1); - TH2F* hTrkQual_OldVsNew_HQ = new TH2F("hTrkQual_OldVsNew_HQ", "", 100,0,1, 100,0,1); - TH2F* hTrkQual_OldVsNew_LQ = new TH2F("hTrkQual_OldVsNew_LQ", "", 100,0,1, 100,0,1); - - // Set up RooUtil - RooUtil util(filename); - // util.Debug(true); - // Loop through the events - for (int i_event = 0; i_event < util.GetNEvents(); ++i_event) { - // Get the next event - auto& event = util.GetEvent(i_event); - - // Get the e_minus tracks from the event - auto e_minus_tracks = event.GetTracks(is_e_minus); - - // Loop through the e_minus tracks - for (auto& track : e_minus_tracks) { - - auto old_trkqual = track.trkqual->result; - auto new_trkqual = track.trkqual_alt->result; - - hTrkQual_OldVsNew->Fill(old_trkqual, new_trkqual); - // Get the track segments at the tracker entrance and has an MC step - auto trk_ent_segments = track.GetSegments([](TrackSegment& segment){ return tracker_entrance(segment) && has_mc_step(segment) && has_reco_step(segment); }); - - // Loop through the tracker entrance track segments - for (auto& segment : trk_ent_segments) { - - auto mom_res = segment.trkseg->mom.R() - segment.trksegmc->mom.R(); - if (std::fabs(mom_res) < 0.25) { - hTrkQual_OldVsNew_HQ->Fill(old_trkqual, new_trkqual); - } - else if (mom_res > 0.70) { - hTrkQual_OldVsNew_LQ->Fill(old_trkqual, new_trkqual); - } - } - } - } - - // Lines for cuts - double old_trkqual_cut_val = 0.93; - TLine* old_trkqual_cut = new TLine(old_trkqual_cut_val, 0, old_trkqual_cut_val, 1.0); - old_trkqual_cut->SetLineWidth(2); - old_trkqual_cut->SetLineStyle(kDashed); - - double new_trkqual_cut_val = 0.95; - TLine* new_trkqual_cut = new TLine(0, new_trkqual_cut_val, 1.0, new_trkqual_cut_val); - new_trkqual_cut->SetLineWidth(2); - new_trkqual_cut->SetLineStyle(kDashed); - - // Draw the histogram - TCanvas* c1 = new TCanvas(); - c1->SetLogz(); - hTrkQual_OldVsNew->SetStats(false); - hTrkQual_OldVsNew->SetTitle("All Tracks"); - hTrkQual_OldVsNew->SetXTitle("TrkQual ANN1 v1"); - hTrkQual_OldVsNew->SetYTitle("TrkQual ANN1 v2"); - hTrkQual_OldVsNew->Draw("COLZ"); - old_trkqual_cut->Draw("LSAME"); - new_trkqual_cut->Draw("LSAME"); - - int old_trkqual_min_bin = hTrkQual_OldVsNew->GetXaxis()->FindBin(old_trkqual_cut_val); - int new_trkqual_min_bin = hTrkQual_OldVsNew->GetYaxis()->FindBin(new_trkqual_cut_val); - // auto integral_old_trkqual = hTrkQual_OldVsNew->Integral(old_trkqual_min_bin, old_trkqual_max_bin, 1, 100); - - double n_total_tracks = hTrkQual_OldVsNew->GetEntries(); - auto integral_old_trkqual = hTrkQual_OldVsNew->Integral(old_trkqual_min_bin, 100, 1, 100); - auto integral_new_trkqual = hTrkQual_OldVsNew->Integral(1, 100, new_trkqual_min_bin, 100); - std::cout << "passes old trkqual (all tracks) " << integral_old_trkqual << " / " << n_total_tracks << " = " << integral_old_trkqual / n_total_tracks << std::endl; - std::cout << "passes new trkqual (all tracks) " << integral_new_trkqual << " / " << n_total_tracks << " = " << integral_new_trkqual / n_total_tracks << std::endl; - - auto integral_fails_both = hTrkQual_OldVsNew->Integral(1, old_trkqual_min_bin, 1, new_trkqual_min_bin); - std::cout << "fails both cuts (all tracks) = " << integral_fails_both << std::endl; - auto integral_passes_old_fails_new = hTrkQual_OldVsNew->Integral(old_trkqual_min_bin, 100, 1, new_trkqual_min_bin); - std::cout << "passes old, fails new (all tracks) = " << integral_passes_old_fails_new << std::endl; - auto integral_fails_old_passes_new = hTrkQual_OldVsNew->Integral(1, old_trkqual_min_bin, new_trkqual_min_bin, 100); - std::cout << "fails old, passes new (all tracks) = " << integral_fails_old_passes_new << std::endl; - auto integral_passes_both = hTrkQual_OldVsNew->Integral(old_trkqual_min_bin, 100, new_trkqual_min_bin, 100); - std::cout << "passes both cuts (all tracks) = " << integral_passes_both << std::endl; - - TLatex* latex = new TLatex(); - latex->SetTextSize(0.06); - latex->SetTextAlign(22); - latex->SetTextColor(kRed); - latex->DrawLatex(0.4, 0.4, Form("#splitline{%.0f}{fail both}", integral_fails_both)); - latex->DrawLatex(0.4, 1.0, Form("%.0f fail old, pass new", integral_fails_old_passes_new)); - latex->DrawLatex(0.95, 0.4, Form("%.0f", integral_passes_old_fails_new)); - latex->DrawLatex(1.0, 1.0, Form("%.0f", integral_passes_both)); - - TCanvas* c2 = new TCanvas(); - c2->SetLogz(); - hTrkQual_OldVsNew_HQ->SetStats(false); - hTrkQual_OldVsNew_HQ->SetTitle("True High-Quality Tracks"); - hTrkQual_OldVsNew_HQ->SetXTitle("TrkQual ANN1 v1"); - hTrkQual_OldVsNew_HQ->SetYTitle("TrkQual ANN1 v2"); - hTrkQual_OldVsNew_HQ->Draw("COLZ"); - old_trkqual_cut->Draw("LSAME"); - new_trkqual_cut->Draw("LSAME"); - - integral_fails_both = hTrkQual_OldVsNew_HQ->Integral(1, old_trkqual_min_bin, 1, new_trkqual_min_bin); - std::cout << "fails both cuts (HQ tracks) = " << integral_fails_both << std::endl; - integral_passes_old_fails_new = hTrkQual_OldVsNew_HQ->Integral(old_trkqual_min_bin, 100, 1, new_trkqual_min_bin); - std::cout << "passes old, fails new (HQ tracks) = " << integral_passes_old_fails_new << std::endl; - integral_fails_old_passes_new = hTrkQual_OldVsNew_HQ->Integral(1, old_trkqual_min_bin, new_trkqual_min_bin, 100); - std::cout << "fails old, passes new (HQ tracks) = " << integral_fails_old_passes_new << std::endl; - integral_passes_both = hTrkQual_OldVsNew_HQ->Integral(old_trkqual_min_bin, 100, new_trkqual_min_bin, 100); - std::cout << "passes both cuts (HQ tracks) = " << integral_passes_both << std::endl; - - latex->DrawLatex(0.4, 0.4, Form("#splitline{%.0f}{fail both}", integral_fails_both)); - latex->DrawLatex(0.4, 1.0, Form("%.0f fail old, pass new", integral_fails_old_passes_new)); - latex->DrawLatex(0.95, 0.4, Form("%.0f", integral_passes_old_fails_new)); - latex->DrawLatex(1.0, 1.0, Form("%.0f", integral_passes_both)); - - TCanvas* c3 = new TCanvas(); - c3->SetLogz(); - hTrkQual_OldVsNew_LQ->SetStats(false); - hTrkQual_OldVsNew_LQ->SetTitle("True Low-Quality Tracks"); - hTrkQual_OldVsNew_LQ->SetXTitle("TrkQual ANN1 v1"); - hTrkQual_OldVsNew_LQ->SetYTitle("TrkQual ANN1 v2"); - hTrkQual_OldVsNew_LQ->Draw("COLZ"); - old_trkqual_cut->Draw("LSAME"); - new_trkqual_cut->Draw("LSAME"); - - - integral_fails_both = hTrkQual_OldVsNew_LQ->Integral(1, old_trkqual_min_bin, 1, new_trkqual_min_bin); - std::cout << "fails both cuts (LQ tracks) = " << integral_fails_both << std::endl; - integral_passes_old_fails_new = hTrkQual_OldVsNew_LQ->Integral(old_trkqual_min_bin, 100, 1, new_trkqual_min_bin); - std::cout << "passes old, fails new (LQ tracks) = " << integral_passes_old_fails_new << std::endl; - integral_fails_old_passes_new = hTrkQual_OldVsNew_LQ->Integral(1, old_trkqual_min_bin, new_trkqual_min_bin, 100); - std::cout << "fails old, passes new (LQ tracks) = " << integral_fails_old_passes_new << std::endl; - integral_passes_both = hTrkQual_OldVsNew_LQ->Integral(old_trkqual_min_bin, 100, new_trkqual_min_bin, 100); - std::cout << "passes both cuts (LQ tracks) = " << integral_passes_both << std::endl; - - latex->DrawLatex(0.4, 0.4, Form("#splitline{%.0f}{fail both}", integral_fails_both)); - latex->DrawLatex(0.4, 1.0, Form("%.0f fail old, pass new", integral_fails_old_passes_new)); - latex->DrawLatex(0.95, 0.4, Form("%.0f", integral_passes_old_fails_new)); - latex->DrawLatex(1.0, 1.0, Form("%.0f", integral_passes_both)); - - if (save_plots) { - std::string pngname = plotsdir + "/Mu2eTrkQual_CompareTrkQualTrainings_All.png"; - c1->SaveAs(pngname.c_str()); - - pngname = plotsdir + "/Mu2eTrkQual_CompareTrkQualTrainings_HighQual.png"; - c2->SaveAs(pngname.c_str()); - - pngname = plotsdir + "/Mu2eTrkQual_CompareTrkQualTrainings_LowQual.png"; - c3->SaveAs(pngname.c_str()); - } -} diff --git a/rooutil/examples/CompareTrkQualTrainings_UserBranches.C b/rooutil/examples/CompareTrkQualTrainings_UserBranches.C new file mode 100644 index 0000000..a0cb276 --- /dev/null +++ b/rooutil/examples/CompareTrkQualTrainings_UserBranches.C @@ -0,0 +1,141 @@ +// +// An example of comparing two track-quality branches using RooUtil's +// typed UserBranch registration instead of hard-coded RooUtil members. +// + +#include "EventNtuple/rooutil/inc/RooUtil.hh" +#include "EventNtuple/rooutil/inc/common_cuts.hh" + +#include "TH2F.h" +#include "TCanvas.h" +#include "TLine.h" +#include "TLatex.h" + +#include +#include + +using namespace rooutil; + +void CompareTrkQualTrainings_UserBranches(std::string filename, + std::string reference_branch = "trkqual", + std::string candidate_branch = "trkqual_bdt", + std::string expected_reference_version = "ANN1_v2", + std::string expected_candidate_version = "BDT1_v2") { + + bool save_plots = false; + std::string plotsdir = "/exp/mu2e/app/users/edmonds/plots/2025-09-25/"; + + TH2F* hTrkQual_ReferenceVsCandidate = new TH2F("hTrkQual_ReferenceVsCandidate", "", 100,0,1, 100,0,1); + TH2F* hTrkQual_ReferenceVsCandidate_HQ = new TH2F("hTrkQual_ReferenceVsCandidate_HQ", "", 100,0,1, 100,0,1); + TH2F* hTrkQual_ReferenceVsCandidate_LQ = new TH2F("hTrkQual_ReferenceVsCandidate_LQ", "", 100,0,1, 100,0,1); + + RooUtil util(filename, true); + // If you want to require a specific version of a model, uncomment the lines below + // util.RequireTrkQualVersion(reference_branch, expected_reference_version); + // util.RequireTrkQualVersion(candidate_branch, expected_candidate_version); + + // In this example, we have commented out reference_trkqual since we will use the standard "trkqual" branch + // auto reference_trkqual = MakeTrackUserBranch(reference_branch); + auto candidate_trkqual = MakeTrackUserBranch(candidate_branch); + util.SetUserBranches({candidate_trkqual}); + + // if (!reference_trkqual->is_bound()) { + // std::cout << "Could not bind reference branch " << reference_branch << std::endl; + // return; + // } + if (!candidate_trkqual->is_bound()) { + std::cout << "Could not bind candidate branch " << candidate_branch << std::endl; + return; + } + + for (int i_event = 0; i_event < util.GetNEvents(); ++i_event) { + auto& event = util.GetEvent(i_event); + auto e_minus_tracks = event.GetTracks(is_e_minus); + + for (auto& track : e_minus_tracks) { + auto* reference_trkqual_result = track.trkqual; + auto* candidate_trkqual_result = track.GetUserBranch(candidate_branch); + if (reference_trkqual_result == nullptr || candidate_trkqual_result == nullptr) { + continue; + } + + const auto reference_value = reference_trkqual_result->result; + const auto candidate_value = candidate_trkqual_result->result; + + hTrkQual_ReferenceVsCandidate->Fill(reference_value, candidate_value); + auto trk_ent_segments = track.GetSegments([](TrackSegment& segment){ return tracker_entrance(segment) && has_mc_step(segment) && has_reco_step(segment); }); + for (auto& segment : trk_ent_segments) { + auto mom_res = segment.trkseg->mom.R() - segment.trksegmc->mom.R(); + if (std::fabs(mom_res) < 0.25) { + hTrkQual_ReferenceVsCandidate_HQ->Fill(reference_value, candidate_value); + } + else if (mom_res > 0.70) { + hTrkQual_ReferenceVsCandidate_LQ->Fill(reference_value, candidate_value); + } + } + } + } + + double reference_cut_val = 0.93; + TLine* reference_cut = new TLine(reference_cut_val, 0, reference_cut_val, 1.0); + reference_cut->SetLineWidth(2); + reference_cut->SetLineStyle(kDashed); + + double candidate_cut_val = 0.95; + TLine* candidate_cut = new TLine(0, candidate_cut_val, 1.0, candidate_cut_val); + candidate_cut->SetLineWidth(2); + candidate_cut->SetLineStyle(kDashed); + + auto draw_summary = [&](TCanvas* canvas, TH2F* hist, const std::string& title) { + canvas->SetLogz(); + hist->SetStats(false); + hist->SetTitle(title.c_str()); + hist->SetXTitle(reference_branch.c_str()); + hist->SetYTitle(candidate_branch.c_str()); + hist->Draw("COLZ"); + reference_cut->Draw("LSAME"); + candidate_cut->Draw("LSAME"); + + int reference_min_bin = hist->GetXaxis()->FindBin(reference_cut_val); + int candidate_min_bin = hist->GetYaxis()->FindBin(candidate_cut_val); + auto fails_both = hist->Integral(1, reference_min_bin, 1, candidate_min_bin); + auto passes_reference_fails_candidate = hist->Integral(reference_min_bin, 100, 1, candidate_min_bin); + auto fails_reference_passes_candidate = hist->Integral(1, reference_min_bin, candidate_min_bin, 100); + auto passes_both = hist->Integral(reference_min_bin, 100, candidate_min_bin, 100); + + std::cout << title << std::endl; + std::cout << " fails both cuts = " << fails_both << std::endl; + std::cout << " passes " << reference_branch << ", fails " << candidate_branch << " = " << passes_reference_fails_candidate << std::endl; + std::cout << " fails " << reference_branch << ", passes " << candidate_branch << " = " << fails_reference_passes_candidate << std::endl; + std::cout << " passes both cuts = " << passes_both << std::endl; + + TLatex* latex = new TLatex(); + latex->SetTextSize(0.06); + latex->SetTextAlign(22); + latex->SetTextColor(kRed); + latex->DrawLatex(0.4, 0.4, Form("#splitline{%.0f}{fail both}", fails_both)); + latex->DrawLatex(0.4, 1.0, Form("%.0f fail %s, pass %s", fails_reference_passes_candidate, reference_branch.c_str(), candidate_branch.c_str())); + latex->DrawLatex(0.95, 0.4, Form("%.0f", passes_reference_fails_candidate)); + latex->DrawLatex(1.0, 1.0, Form("%.0f", passes_both)); + }; + + TCanvas* c1 = new TCanvas(); + draw_summary(c1, hTrkQual_ReferenceVsCandidate, "All Tracks"); + + TCanvas* c2 = new TCanvas(); + draw_summary(c2, hTrkQual_ReferenceVsCandidate_HQ, "True High-Quality Tracks"); + + TCanvas* c3 = new TCanvas(); + draw_summary(c3, hTrkQual_ReferenceVsCandidate_LQ, "True Low-Quality Tracks"); + + if (save_plots) { + std::string pngname = plotsdir + "/Mu2eTrkQual_CompareTrkQualTrainings_UserBranches_All.png"; + c1->SaveAs(pngname.c_str()); + + pngname = plotsdir + "/Mu2eTrkQual_CompareTrkQualTrainings_UserBranches_HighQual.png"; + c2->SaveAs(pngname.c_str()); + + pngname = plotsdir + "/Mu2eTrkQual_CompareTrkQualTrainings_UserBranches_LowQual.png"; + c3->SaveAs(pngname.c_str()); + } +} diff --git a/rooutil/inc/Event.hh b/rooutil/inc/Event.hh index e8789f3..c6f53f0 100644 --- a/rooutil/inc/Event.hh +++ b/rooutil/inc/Event.hh @@ -42,6 +42,7 @@ #include "EventNtuple/inc/MCStepInfo.hh" #include "EventNtuple/rooutil/inc/Track.hh" +#include "EventNtuple/rooutil/inc/UserBranch.hh" #include "EventNtuple/rooutil/inc/TimeCluster.hh" #include "EventNtuple/rooutil/inc/CrvCoinc.hh" #include "EventNtuple/rooutil/inc/CaloCluster.hh" @@ -62,7 +63,6 @@ namespace rooutil { CheckForBranch(ntuple, "trksegs", &this->trksegs); CheckForBranch(ntuple, "trkcalohit", &this->trkcalohit); CheckForBranch(ntuple, "trkqual", &this->trkqual); - CheckForBranch(ntuple, "trkqual3", &this->trkqual_alt); // TODO: un-hardcode those CheckForBranch(ntuple, "crvcoincs", &this->crvcoincs); CheckForBranch(ntuple, "trkpid", &this->trkpid); @@ -100,6 +100,10 @@ namespace rooutil { CheckForBranch(ntuple, "mcsteps_virtualdetector", &this->mcsteps_virtualdetector); } + void SetUserBranches(const std::vector>& branches) { + user_branches = branches; + } + // Check if a branch exists in the TChain, and optionally set its address bool CheckForBranch(TChain* ntuple, const char* branch_name, void* address = nullptr) { if(ntuple->GetBranch(branch_name) == nullptr || ntuple->GetBranchStatus(branch_name) == 0) return false; @@ -142,8 +146,12 @@ namespace rooutil { UpdateObject(track.trkmats, trkmats, i_track, debug); UpdateObject(track.trkhitcalibs, trkhitcalibs, i_track, debug); UpdateObject(track.trkqual, trkqual, i_track, debug); - UpdateObject(track.trkqual_alt, trkqual_alt, i_track, debug); UpdateObject(track.trkpid, trkpid, i_track, debug); + for (const auto& user_branch : user_branches) { + if (user_branch->is_bound() && user_branch->scope() == UserBranchScope::Track) { + track.SetUserBranch(user_branch->name(), user_branch->TrackElementPtr(i_track)); + } + } if (debug) { std::cout << "Event::Update(): Updating Track " << i_track << "... " << std::endl; } track.Update(debug); @@ -287,8 +295,12 @@ namespace rooutil { if (trksegsmc) { trksegsmc->erase(trksegsmc->begin()+trks_to_remove[i_trk]); } if (trkcalohit) { trkcalohit->erase(trkcalohit->begin()+trks_to_remove[i_trk]); } if (trkqual) { trkqual->erase(trkqual->begin()+trks_to_remove[i_trk]); } - if (trkqual_alt) { trkqual_alt->erase(trkqual_alt->begin()+trks_to_remove[i_trk]); } if (trkpid) { trkpid->erase(trkpid->begin()+trks_to_remove[i_trk]); } + for (const auto& user_branch : user_branches) { + if (user_branch->is_bound() && user_branch->scope() == UserBranchScope::Track) { + user_branch->EraseTrack(trks_to_remove[i_trk]); + } + } if (trksegpars_lh) { trksegpars_lh->erase(trksegpars_lh->begin()+trks_to_remove[i_trk]); } if (trksegpars_ch) { trksegpars_ch->erase(trksegpars_ch->begin()+trks_to_remove[i_trk]); } if (trksegpars_kl) { trksegpars_kl->erase(trksegpars_kl->begin()+trks_to_remove[i_trk]); } @@ -420,7 +432,6 @@ namespace rooutil { std::vector* trkcalohit = nullptr; std::vector* trkcalohitmc = nullptr; std::vector* trkqual = nullptr; - std::vector* trkqual_alt = nullptr; // an optional trkqual branch to also use std::vector* trkpid = nullptr; std::vector>* trksegs = nullptr; std::vector>* trksegsmc = nullptr; @@ -431,6 +442,7 @@ namespace rooutil { std::vector>* trkhitsmc = nullptr; std::vector>* trkmats = nullptr; std::vector>* trkhitcalibs = nullptr; + std::vector> user_branches; std::vector* timeclusters = nullptr; diff --git a/rooutil/inc/RooUtil.hh b/rooutil/inc/RooUtil.hh index 0ba2cdd..f18557a 100644 --- a/rooutil/inc/RooUtil.hh +++ b/rooutil/inc/RooUtil.hh @@ -1,15 +1,21 @@ #ifndef RooUtil_hh_ #define RooUtil_hh_ +#include #include +#include +#include #include "TFile.h" #include "TTree.h" #include "TH1I.h" +#include "EventNtuple/inc/TrkQualMetadata.hh" #include "EventNtuple/rooutil/inc/Event.hh" +#include "EventNtuple/rooutil/inc/UserBranch.hh" namespace rooutil { + class RooUtil { public: RooUtil(std::string filename, bool debug = false, std::string treename = "EventNtuple/ntuple") : debug(debug), n_proc_events(0) { @@ -21,6 +27,7 @@ namespace rooutil { ntuple->Add(filename.c_str()); SetVersionNumber(filename); SetNProcessedEvents(filename); + LoadTrkQualMetadata(filename); } else { // assume its a file list std::ifstream filelist(filename); @@ -36,6 +43,7 @@ namespace rooutil { first_line = false; } SetNProcessedEvents(line); + LoadTrkQualMetadata(line); } filelist.close(); } else { @@ -84,6 +92,29 @@ namespace rooutil { int GetNEvents() { return ntuple->GetEntries(); } int GetNProcEvents() { return n_proc_events; } + bool HasTrkQualMetadata(const std::string& output_branch) const { + return trkqual_metadata.find(output_branch) != trkqual_metadata.end(); + } + + const mu2e::TrkQualMetadata& GetTrkQualMetadata(const std::string& output_branch) const { + const auto metadata = trkqual_metadata.find(output_branch); + if (metadata == trkqual_metadata.end()) { + throw std::runtime_error( + "No TrkQual metadata is available for output branch " + output_branch); + } + return metadata->second; + } + + void RequireTrkQualVersion( + const std::string& output_branch, const std::string& expected_model_version) const { + const auto& metadata = GetTrkQualMetadata(output_branch); + if (metadata.model_version != expected_model_version) { + throw std::runtime_error( + "Unexpected TrkQual version for " + output_branch + ": expected " + + expected_model_version + ", got " + metadata.model_version); + } + } + Event& GetEvent(int i_event) { if (debug) { std::cout << "RooUtil::GetEvent(): Getting event " << i_event << std::endl; } ntuple->GetEntry(i_event); @@ -118,6 +149,22 @@ namespace rooutil { TurnOnBranch("*"); } + void SetUserBranches(const std::vector>& branches) { + for (const auto& branch : branches) { + branch->Bind(ntuple); + const auto existing = std::find_if(user_branches.begin(), user_branches.end(), + [&branch](const std::shared_ptr& registered) { + return registered->name() == branch->name(); + }); + if (existing == user_branches.end()) { + user_branches.push_back(branch); + } else { + *existing = branch; + } + } + event->SetUserBranches(user_branches); + } + void CreateOutputEventNtuple(TFile* outfile) { auto dir = outfile->mkdir("EventNtuple"); dir->cd(); @@ -134,7 +181,6 @@ namespace rooutil { if(event->trkcalohit) { output_ntuple->Branch("trkcalohit", event->trkcalohit); } if(event->trkcalohitmc) { output_ntuple->Branch("trkcalohitmc", event->trkcalohitmc); } if(event->trkqual) { output_ntuple->Branch("trkqual", event->trkqual); } - if(event->trkqual_alt) { output_ntuple->Branch("trkqual_alt", event->trkqual_alt); } if(event->trkpid) { output_ntuple->Branch("trkpid", event->trkpid); } if(event->trksegs) { output_ntuple->Branch("trksegs", event->trksegs); } if(event->trksegsmc) { output_ntuple->Branch("trksegsmc", event->trksegsmc); } @@ -170,6 +216,11 @@ namespace rooutil { } if (event->mcsteps_virtualdetector) { output_ntuple->Branch("mcsteps_virtualdetector", event->mcsteps_virtualdetector); } + for (const auto& branch : user_branches) { + if (branch->is_bound()) { + branch->BranchOutput(output_ntuple); + } + } // Write out histograms from input to output hVersion->Write(); @@ -180,12 +231,51 @@ namespace rooutil { } private: + void LoadTrkQualMetadata(const std::string& filename) { + TFile file(filename.c_str(), "READ"); + TH1I* metadata_histogram = nullptr; + file.GetObject("EventNtuple/trkqual_metadata", metadata_histogram); + if (metadata_histogram == nullptr) { + return; + } + + const std::string input_tag_marker = ": input tag = "; + const std::string model_version_marker = "; model version = "; + for (int bin = 1; bin <= metadata_histogram->GetNbinsX(); ++bin) { + const std::string label = metadata_histogram->GetXaxis()->GetBinLabel(bin); + const auto input_tag_pos = label.find(input_tag_marker); + const auto model_version_pos = label.find(model_version_marker); + if (input_tag_pos == std::string::npos || model_version_pos == std::string::npos || + input_tag_pos >= model_version_pos) { + throw std::runtime_error("Invalid TrkQual metadata in " + filename + ": " + label); + } + + mu2e::TrkQualMetadata metadata{ + label.substr(0, input_tag_pos), + label.substr(input_tag_pos + input_tag_marker.size(), + model_version_pos - input_tag_pos - input_tag_marker.size()), + label.substr(model_version_pos + model_version_marker.size()) + }; + const auto existing = trkqual_metadata.find(metadata.output_branch); + if (existing != trkqual_metadata.end() && + (existing->second.input_tag != metadata.input_tag || + existing->second.model_version != metadata.model_version)) { + throw std::runtime_error( + "TrkQual metadata for " + metadata.output_branch + + " differs between input files"); + } + trkqual_metadata[metadata.output_branch] = metadata; + } + } + TChain* ntuple; Event* event; // holds all the variables for SetBranchAddress bool debug; TH1I* hVersion; int n_proc_events; + std::map trkqual_metadata; + std::vector> user_branches; TTree* output_ntuple; // for output }; diff --git a/rooutil/inc/Track.hh b/rooutil/inc/Track.hh index b98763e..76fe6d7 100644 --- a/rooutil/inc/Track.hh +++ b/rooutil/inc/Track.hh @@ -2,6 +2,8 @@ #define Track_hh_ #include +#include +#include #include "EventNtuple/inc/TrkInfo.hh" #include "EventNtuple/inc/TrkInfoMC.hh" #include "EventNtuple/inc/SurfaceStepInfo.hh" @@ -173,6 +175,16 @@ namespace rooutil { return select_hits; } TrackHits hits; + + void SetUserBranch(const std::string& branch_name, void* value) { + user_branches[branch_name] = value; + } + + template + T* GetUserBranch(const std::string& branch_name) { + const auto branch = user_branches.find(branch_name); + return branch == user_branches.end() ? nullptr : static_cast(branch->second); + } // Pointers to the data mu2e::TrkInfo* trk = nullptr; @@ -189,8 +201,11 @@ namespace rooutil { mu2e::TrkCaloHitInfo* trkcalohit = nullptr; std::vector* trkmcsim = nullptr; mu2e::MVAResultInfo* trkqual = nullptr; - mu2e::MVAResultInfo* trkqual_alt = nullptr; // TODO: is there a way to allow for more than two... + mu2e::MVAResultInfo* trkqual_bdt = nullptr; mu2e::MVAResultInfo* trkpid = nullptr; + + private: + std::unordered_map user_branches; }; typedef std::function TrackCut; diff --git a/rooutil/inc/UserBranch.hh b/rooutil/inc/UserBranch.hh new file mode 100644 index 0000000..3bd0321 --- /dev/null +++ b/rooutil/inc/UserBranch.hh @@ -0,0 +1,121 @@ +#ifndef UserBranch_hh_ +#define UserBranch_hh_ + +#include +#include +#include +#include + +#include "TChain.h" +#include "TTree.h" + +namespace rooutil { + enum class UserBranchScope { + Event, + Track + }; + + class UserBranchBase { + public: + UserBranchBase(std::string branch_name, UserBranchScope branch_scope) + : branch_name_(std::move(branch_name)), branch_scope_(branch_scope) {} + virtual ~UserBranchBase() = default; + + const std::string& name() const { return branch_name_; } + UserBranchScope scope() const { return branch_scope_; } + bool is_bound() const { return is_bound_; } + + virtual bool Bind(TChain* ntuple) = 0; + virtual void BranchOutput(TTree* tree) = 0; + virtual void* EventPtr() { return nullptr; } + virtual void* TrackElementPtr(std::size_t) { return nullptr; } + virtual void EraseTrack(std::size_t) {} + + protected: + std::string branch_name_; + UserBranchScope branch_scope_; + bool is_bound_ = false; + }; + + template + class EventUserBranch : public UserBranchBase { + public: + explicit EventUserBranch(const std::string& branch_name) + : UserBranchBase(branch_name, UserBranchScope::Event) {} + + bool Bind(TChain* ntuple) override { + if (ntuple == nullptr || ntuple->GetBranch(branch_name_.c_str()) == nullptr || ntuple->GetBranchStatus(branch_name_.c_str()) == 0) { + is_bound_ = false; + return false; + } + ntuple->SetBranchAddress(branch_name_.c_str(), &value_); + is_bound_ = true; + return true; + } + + void BranchOutput(TTree* tree) override { + if (tree != nullptr && value_ != nullptr) { + tree->Branch(branch_name_.c_str(), value_); + } + } + + void* EventPtr() override { return value_; } + T* value() { return value_; } + + private: + T* value_ = nullptr; + }; + + template + class TrackUserBranch : public UserBranchBase { + public: + explicit TrackUserBranch(const std::string& branch_name) + : UserBranchBase(branch_name, UserBranchScope::Track) {} + + bool Bind(TChain* ntuple) override { + if (ntuple == nullptr || ntuple->GetBranch(branch_name_.c_str()) == nullptr || ntuple->GetBranchStatus(branch_name_.c_str()) == 0) { + is_bound_ = false; + return false; + } + ntuple->SetBranchAddress(branch_name_.c_str(), &values_); + is_bound_ = true; + return true; + } + + void BranchOutput(TTree* tree) override { + if (tree != nullptr && values_ != nullptr) { + tree->Branch(branch_name_.c_str(), values_); + } + } + + void* TrackElementPtr(std::size_t index) override { + if (values_ == nullptr || index >= values_->size()) { + return nullptr; + } + return &(values_->at(index)); + } + + void EraseTrack(std::size_t index) override { + if (values_ != nullptr && index < values_->size()) { + values_->erase(values_->begin() + index); + } + } + + std::vector* values() { return values_; } + + private: + std::vector* values_ = nullptr; + }; + + template + std::shared_ptr> MakeEventUserBranch(const std::string& branch_name) { + return std::make_shared>(branch_name); + } + + template + std::shared_ptr> MakeTrackUserBranch(const std::string& branch_name) { + return std::make_shared>(branch_name); + } +} // namespace rooutil + +#endif diff --git a/src/EventNtupleMaker_module.cc b/src/EventNtupleMaker_module.cc index 4bb25c3..52f727d 100644 --- a/src/EventNtupleMaker_module.cc +++ b/src/EventNtupleMaker_module.cc @@ -80,6 +80,7 @@ #include "EventNtuple/inc/RecoQualInfo.hh" #include "EventNtuple/inc/MVAResultInfo.hh" #include "EventNtuple/inc/BestCrvAssns.hh" +#include "EventNtuple/inc/TrkQualMetadata.hh" #include "EventNtuple/inc/MCStepInfo.hh" #include "EventNtuple/inc/SurfaceStepInfo.hh" #include "EventNtuple/inc/MCStepSummaryInfo.hh" @@ -89,6 +90,7 @@ // C++ includes. #include #include +#include #include using namespace std; @@ -115,13 +117,21 @@ namespace mu2e { fhicl::Atom matchDepth{Name("matchDepth"), Comment("Depth of MC truth matching to keep (-1 = all)")}; }; + struct TrkQualLeafConfig { + using Name=fhicl::Name; + using Comment=fhicl::Comment; + fhicl::Atom leafname{Name("leafname"), Comment("Suffix appended to qual; use an empty string for the canonical TrkQual branch")}; + fhicl::Atom modelVersion{Name("modelVersion"), Comment("TrkQual ML model version recorded in trkqual_metadata")}; + fhicl::Atom inputTag{Name("inputTag"), Comment("Input tag for the TrkQual MVAResultCollection")}; + }; + struct TrkFitConfig { using Name=fhicl::Name; using Comment=fhicl::Comment; fhicl::Atom input{Name("input"), Comment("KalSeedCollection input tag")}; fhicl::Atom branchname{Name("branchname"), Comment("Name of output branch")}; fhicl::Atom fill{Name("fill"), Comment("Set false to skip this branch entirely (no collection reads, no output branches)")}; - fhicl::Sequence trkQualTags{Name("trkQualTags"), Comment("Input tags for MVAResultCollection to use for TrkQuals")}; + fhicl::Sequence> trkQualLeaves{Name("trkQualLeaves"), Comment("TrkQual output branch names, input tags, and model provenance")}; fhicl::Sequence trkPIDTags{Name("trkPIDTags"), Comment("Input tags for MVAResultCollection to use for TrkPID")}; fhicl::Table options{Name("options"), Comment("Per-branch fill options")}; }; @@ -307,6 +317,7 @@ namespace mu2e { Config _conf; std::vector _allTrkFitBranches; // configurations for all track fit branches + std::map> _trkQualMetadata; // main TTree TTree* _ntuple; TH1I* _hVersion; @@ -513,6 +524,14 @@ namespace mu2e { // populate branch list from trk.branches for(const auto& trk_fit_cfg : _conf.trk().fits()){ + const TrkFitBranchIndex i_trk_fit_branch = _allTrkFitBranches.size(); + for (const auto& trkQualConfig : trk_fit_cfg.trkQualLeaves()) { + _trkQualMetadata[i_trk_fit_branch].push_back({ + trk_fit_cfg.branchname() + "qual" + trkQualConfig.leafname(), + trkQualConfig.inputTag(), + trkQualConfig.modelVersion() + }); + } _allTrkFitBranches.push_back(trk_fit_cfg); } @@ -541,7 +560,7 @@ namespace mu2e { _allTSMIs[i_trk_fit_branch] = {}; _allTSHIMCs[i_trk_fit_branch] = {}; - _allTrkQualResults[i_trk_fit_branch].resize(i_trkFitConfig.trkQualTags().size()); + _allTrkQualResults[i_trk_fit_branch].resize(i_trkFitConfig.trkQualLeaves().size()); _allTrkPIDResults[i_trk_fit_branch].resize(i_trkFitConfig.trkPIDTags().size()); for (StepCollIndex ixt = 0; ixt < _extraMCStepTags.size(); ++ixt) { @@ -568,6 +587,16 @@ namespace mu2e { _hVersion->GetXaxis()->SetBinLabel(2, "minor"); _hVersion->SetBinContent(2, 11); _hVersion->GetXaxis()->SetBinLabel(3, "patch"); _hVersion->SetBinContent(3, 2); _hProcEvents = tfs->make("n_proc_events", "number of processed events", 1,0,1); + size_t nTrkQualAlgorithms = 0; + for (const auto& trkFitConfig : _allTrkFitBranches) { + if (trkFitConfig.fill()) nTrkQualAlgorithms += trkFitConfig.trkQualLeaves().size(); + } + TH1I* hTrkQualMetadata = nullptr; + if (nTrkQualAlgorithms > 0) { + hTrkQualMetadata = tfs->make( + "trkqual_metadata", "Configured TrkQual ML algorithms", nTrkQualAlgorithms, 0, nTrkQualAlgorithms); + } + size_t iTrkQualAlgorithm = 0; // event info branch _ntuple->Branch("evtinfo",&_einfo,_buffsize,_splitlevel); if (fillEventMC()) { @@ -596,10 +625,17 @@ namespace mu2e { if(_ftype == KinematicLine) _ntuple->Branch((branch+"segpars_kl.").c_str(),&_allKLIs.at(i_trk_fit_branch),_buffsize,_splitlevel); // TrkCaloHit: currently only 1 _ntuple->Branch((branch+"calohit.").c_str(),&_allTCHIs.at(i_trk_fit_branch)); - for (size_t i_trkQualTag = 0; i_trkQualTag < i_trkFitConfig.trkQualTags().size(); ++i_trkQualTag) { - std::string branchname = "qual"; - if (i_trkQualTag > 0) branchname += std::to_string(i_trkQualTag+1); - _ntuple->Branch((branch+branchname+".").c_str(),&_allTrkQualResults.at(i_trk_fit_branch).at(i_trkQualTag),_buffsize,_splitlevel); + for (size_t i_trkQual = 0; i_trkQual < i_trkFitConfig.trkQualLeaves().size(); ++i_trkQual) { + const auto& trkQualMetadata = _trkQualMetadata.at(i_trk_fit_branch).at(i_trkQual); + const std::string& outputBranch = trkQualMetadata.output_branch; + _ntuple->Branch((outputBranch+".").c_str(),&_allTrkQualResults.at(i_trk_fit_branch).at(i_trkQual),_buffsize,_splitlevel); + char metadataLabel[1024]; + std::snprintf( + metadataLabel, sizeof(metadataLabel), + "%s: input tag = %s; model version = %s", + outputBranch.c_str(), trkQualMetadata.input_tag.c_str(), trkQualMetadata.model_version.c_str()); + hTrkQualMetadata->GetXaxis()->SetBinLabel(++iTrkQualAlgorithm, metadataLabel); + hTrkQualMetadata->SetBinContent(iTrkQualAlgorithm, 1); } for (size_t i_trkPIDTag = 0; i_trkPIDTag < i_trkFitConfig.trkPIDTags().size(); ++i_trkPIDTag) { std::string branchname = "pid"; @@ -636,7 +672,6 @@ namespace mu2e { } } } - // Time clusters if(_conf.timeclusters().fill()) { _ntuple->Branch("timeclusters.",&_tcIs,_buffsize,_splitlevel); @@ -818,9 +853,9 @@ namespace mu2e { _allKSPCHs.push_back(kalSeedPtrCollHandle); std::vector> trkQualCollHandles; - for (const auto& i_trkQualTag : i_trkFitConfig.trkQualTags()) { + for (const auto& i_trkQual : i_trkFitConfig.trkQualLeaves()) { art::Handle trkQualCollHandle; - event.getByLabel(i_trkQualTag,trkQualCollHandle); + event.getByLabel(i_trkQual.inputTag(),trkQualCollHandle); trkQualCollHandles.push_back(trkQualCollHandle); } _allTrkQualCHs.emplace_back(trkQualCollHandles); @@ -913,8 +948,8 @@ namespace mu2e { _allMCVDInfos.at(i_trk_fit_branch).clear(); _allMCSimTIs.at(i_trk_fit_branch).clear(); - for (size_t i_trkQualTag = 0; i_trkQualTag < i_trkFitConfig.trkQualTags().size(); ++i_trkQualTag) { - _allTrkQualResults.at(i_trk_fit_branch).at(i_trkQualTag).clear(); + for (size_t i_trkQual = 0; i_trkQual < i_trkFitConfig.trkQualLeaves().size(); ++i_trkQual) { + _allTrkQualResults.at(i_trk_fit_branch).at(i_trkQual).clear(); } for (size_t i_trkPIDTag = 0; i_trkPIDTag < i_trkFitConfig.trkPIDTags().size(); ++i_trkPIDTag) { _allTrkPIDResults.at(i_trk_fit_branch).at(i_trkPIDTag).clear(); @@ -945,7 +980,6 @@ namespace mu2e { } } } - // Time clusters if(_conf.timeclusters().fill()) { _tcIs.clear();