diff --git a/TrkDiag/data/TrackPID_v1.dat b/TrkDiag/data/TrackPID_v1.dat index 9ec23e9..ba8617c 100644 --- a/TrkDiag/data/TrackPID_v1.dat +++ b/TrkDiag/data/TrackPID_v1.dat @@ -1,16 +1,16 @@ tensor_sequential1dense1BiasAddReadVariableOp0 5 -3.95117021 0 -4.19677019 -0.242374584 0 +0.67720902 -0.069714874 -0.694897473 -0.390859604 -0.520138502 tensor_sequential1dense12CastReadVariableOp0 50 --0.705076814 -0.432341874 -0.553150296 0.21540308 -0.351793796 0.476246953 -0.106644973 -0.417871863 -0.0440464132 0.177883551 -0.16141963 -0.47916162 0.606347501 0.366820812 -0.230519354 -0.472491771 0.313038588 0.386520803 0.371591628 -0.224177033 0.457869917 -0.171457827 0.519692957 -0.278755456 -0.38365525 -0.329625219 0.141675428 -0.350409567 0.103589013 -0.0581277721 0.432253748 -0.301287413 -0.615136921 0.0705423355 0.448004067 -2.0370698 -0.893573642 -0.030412674 0.35889414 -0.297698051 -0.367111266 0.301380754 -0.571814477 -0.362872422 0.073617816 0.193396747 0.336626947 0.388768613 -0.599104285 0.534656346 +0.423877656 0.00515525788 -0.0492455959 0.0618200749 0.954318166 -0.604066908 0.720992923 0.797571242 0.188263789 0.821694672 0.0712498724 -1.47822523 -0.0212162137 -0.74622488 -0.798371017 -0.768375695 -0.531521499 -0.0368613712 -0.179116294 0.0123269027 0.983402967 0.564578056 -0.296996564 -1.52724409 -0.7494542 0.185713574 0.663779736 0.236173511 -1.05086386 0.494367659 -0.369194329 0.0828162059 -0.342357725 0.472442687 1.30193877 0.348389268 -0.699540198 -1.2993387 0.42568168 0.99106276 -0.429250211 -1.27230895 -0.330463171 -0.248296201 1.03471065 -0.0554589368 -0.583253503 -0.0872679353 0.451856405 -1.0176909 tensor_sequential1dense12BiasAddReadVariableOp0 10 --2.48500657 0 -3.64792609 0 0 4.17723703 -2.36129212 0 2.68298984 2.38057661 +0.360056847 0.393950015 0 0.914983153 -0.501553297 -0.233833224 0.253284663 0.0172187667 0.731160939 0.163516104 tensor_sequential1dense21BiasAddReadVariableOp0 5 -3.48705363 -3.67133808 -0.0374668688 -3.27729893 -3.68775105 +0.347450912 0.387974828 -0.184288666 -0.10895285 0.235439435 tensor_sequential1dense1CastReadVariableOp0 20 -0.836482465 -0.212255061 0.0658833757 -0.730508149 0.558905005 0.485133708 -0.564192295 0.523953795 -0.0750864968 -0.673678458 1.23919535 0.296814561 -0.67768687 -3.03104234 -0.0644903779 -5.03836155 0.609472513 3.84570956 -5.23402023 -0.525624335 +1.28559923 2.32776809 2.48449874 -0.832492709 -0.711716473 -0.254594088 0.0412935913 -0.0256965607 -0.28075251 0.273529559 0.619013488 -0.114196911 0.06435287 -4.25883675 -5.56657028 -0.356506407 -0.542718291 0.180514976 -2.02528548 0.602497101 tensor_sequential1dense21CastReadVariableOp0 50 -0.115136102 -0.436833829 0.30147469 0.27063179 0.270309329 0.197524309 -0.497536033 0.0962924361 0.301537275 -0.469570994 -0.300404757 0.271288723 -0.328934014 0.600233972 0.0958410427 -0.342013776 0.460635483 -0.492932588 0.0162852407 0.40738827 0.369479835 -0.610442698 0.217112422 -0.558332384 0.355603456 0.753055334 -0.311895877 0.441717386 0.203458861 -0.0115984064 -0.00314931269 -0.259888023 -0.360559404 0.32812199 0.161191747 -0.248590976 0.301618695 0.0615119934 -0.583952665 -0.0536812544 0.148783773 0.110628158 -0.348669648 -0.313134849 0.324738473 0.23614575 0.276106864 -0.546891332 -0.116325177 -0.168601051 +0.782625318 0.205530673 -0.229301631 -0.530559421 1.08883691 -0.696652055 -0.829344809 0.399743646 -0.416170895 -1.41417921 -0.576976597 0.318252385 0.397290647 -0.38962242 -0.126767218 3.39918113 0.797734439 0.49670127 -0.0584926009 2.68833351 -0.91740644 -2.2402432 -0.672245562 0.0128413225 -0.707842648 0.307995051 -0.036977753 0.519744575 0.22346288 -0.360910714 0.824952066 0.168633327 -0.547800422 0.0286999047 -0.00679685175 0.0468300357 -0.18069458 -0.427938163 -0.396098584 0.859228551 -1.05733216 -2.06048298 -0.579410791 -0.375878483 -0.973983347 0.0930191353 0.818828642 0.291658878 0.354209214 0.44286707 tensor_sequential1dense31AddReadVariableOp0 1 -3.68641305 +-8.2496891 tensor_sequential1dense31CastReadVariableOp0 5 -0.151621565 -0.436080575 -0.779650092 -0.385260344 -0.245328888 +2.08103108 1.69884634 -0.421817899 -0.895168364 1.36597729 diff --git a/TrkDiag/inc/TrackPID.hxx b/TrkDiag/inc/TrackPID_v0.hxx similarity index 98% rename from TrkDiag/inc/TrackPID.hxx rename to TrkDiag/inc/TrackPID_v0.hxx index be8c174..b07f320 100644 --- a/TrkDiag/inc/TrackPID.hxx +++ b/TrkDiag/inc/TrackPID_v0.hxx @@ -1,7 +1,7 @@ //Code generated automatically by TMVA for Inference of Model file [TrackPID.onnx] at [Fri Sep 26 15:44:22 2025] -#ifndef ROOT_TMVA_SOFIE_TRACKPID -#define ROOT_TMVA_SOFIE_TRACKPID +#ifndef ROOT_TMVA_SOFIE_TRACKPID_V0 +#define ROOT_TMVA_SOFIE_TRACKPID_V0 #include #include @@ -9,7 +9,7 @@ #include "TMVA/SOFIE_common.hxx" #include -namespace TMVA_SOFIE_TrackPID{ +namespace TMVA_SOFIE_TrackPID_v0{ namespace BLAS{ extern "C" void sgemv_(const char * trans, const int * m, const int * n, const float * alpha, const float * A, const int * lda, const float * X, const int * incx, const float * beta, const float * Y, const int * incy); @@ -62,7 +62,7 @@ std::vector fTensor_sequential1dense1BiasAddReadVariableOp0bcast = std::v float * tensor_sequential1dense1BiasAddReadVariableOp0bcast = fTensor_sequential1dense1BiasAddReadVariableOp0bcast.data(); -Session(std::string filename ="TrackPID.dat") { +Session(std::string filename ="TrackPID_v0.dat") { //--- reading weights from file std::ifstream f; @@ -263,6 +263,6 @@ std::vector infer(float* tensor_inputlayer){ return fTensor_output; } }; -} //TMVA_SOFIE_TrackPID +} //TMVA_SOFIE_TrackPID_V0 -#endif // ROOT_TMVA_SOFIE_TRACKPID +#endif // ROOT_TMVA_SOFIE_TRACKPID_V0 diff --git a/TrkDiag/inc/TrackPID_v1.hxx b/TrkDiag/inc/TrackPID_v1.hxx new file mode 100644 index 0000000..81405ce --- /dev/null +++ b/TrkDiag/inc/TrackPID_v1.hxx @@ -0,0 +1,268 @@ +//Code generated automatically by TMVA for Inference of Model file [TrackPID_v1.onnx] at [Wed Jul 22 20:07:11 2026] + +#ifndef ROOT_TMVA_SOFIE_TRACKPID_V1 +#define ROOT_TMVA_SOFIE_TRACKPID_V1 + +#include +#include +#include +#include "TMVA/SOFIE_common.hxx" +#include + +namespace TMVA_SOFIE_TrackPID_v1{ +namespace BLAS{ + extern "C" void sgemv_(const char * trans, const int * m, const int * n, const float * alpha, const float * A, + const int * lda, const float * X, const int * incx, const float * beta, const float * Y, const int * incy); + extern "C" void sgemm_(const char * transa, const char * transb, const int * m, const int * n, const int * k, + const float * alpha, const float * A, const int * lda, const float * B, const int * ldb, + const float * beta, float * C, const int * ldc); +}//BLAS +struct Session { +std::vector fTensor_sequential1dense1BiasAddReadVariableOp0 = std::vector(5); +float * tensor_sequential1dense1BiasAddReadVariableOp0 = fTensor_sequential1dense1BiasAddReadVariableOp0.data(); +std::vector fTensor_sequential1dense12CastReadVariableOp0 = std::vector(50); +float * tensor_sequential1dense12CastReadVariableOp0 = fTensor_sequential1dense12CastReadVariableOp0.data(); +std::vector fTensor_sequential1dense12BiasAddReadVariableOp0 = std::vector(10); +float * tensor_sequential1dense12BiasAddReadVariableOp0 = fTensor_sequential1dense12BiasAddReadVariableOp0.data(); +std::vector fTensor_sequential1dense21BiasAddReadVariableOp0 = std::vector(5); +float * tensor_sequential1dense21BiasAddReadVariableOp0 = fTensor_sequential1dense21BiasAddReadVariableOp0.data(); +std::vector fTensor_sequential1dense1CastReadVariableOp0 = std::vector(20); +float * tensor_sequential1dense1CastReadVariableOp0 = fTensor_sequential1dense1CastReadVariableOp0.data(); +std::vector fTensor_sequential1dense21CastReadVariableOp0 = std::vector(50); +float * tensor_sequential1dense21CastReadVariableOp0 = fTensor_sequential1dense21CastReadVariableOp0.data(); +std::vector fTensor_sequential1dense31AddReadVariableOp0 = std::vector(1); +float * tensor_sequential1dense31AddReadVariableOp0 = fTensor_sequential1dense31AddReadVariableOp0.data(); +std::vector fTensor_sequential1dense31CastReadVariableOp0 = std::vector(5); +float * tensor_sequential1dense31CastReadVariableOp0 = fTensor_sequential1dense31CastReadVariableOp0.data(); + +//--- declare and allocate the intermediate tensors +std::vector fTensor_sequential1dense31AddReadVariableOp0bcast = std::vector(32); +float * tensor_sequential1dense31AddReadVariableOp0bcast = fTensor_sequential1dense31AddReadVariableOp0bcast.data(); +std::vector fTensor_sequential1dense21Relu0 = std::vector(160); +float * tensor_sequential1dense21Relu0 = fTensor_sequential1dense21Relu0.data(); +std::vector fTensor_sequential1dense21BiasAddReadVariableOp0bcast = std::vector(160); +float * tensor_sequential1dense21BiasAddReadVariableOp0bcast = fTensor_sequential1dense21BiasAddReadVariableOp0bcast.data(); +std::vector fTensor_sequential1dense12Relu0 = std::vector(320); +float * tensor_sequential1dense12Relu0 = fTensor_sequential1dense12Relu0.data(); +std::vector fTensor_sequential1dense12MatMulGemm80 = std::vector(320); +float * tensor_sequential1dense12MatMulGemm80 = fTensor_sequential1dense12MatMulGemm80.data(); +std::vector fTensor_sequential1dense12BiasAddReadVariableOp0bcast = std::vector(320); +float * tensor_sequential1dense12BiasAddReadVariableOp0bcast = fTensor_sequential1dense12BiasAddReadVariableOp0bcast.data(); +std::vector fTensor_sequential1dense21MatMulGemm90 = std::vector(160); +float * tensor_sequential1dense21MatMulGemm90 = fTensor_sequential1dense21MatMulGemm90.data(); +std::vector fTensor_sequential1dense1Relu0 = std::vector(160); +float * tensor_sequential1dense1Relu0 = fTensor_sequential1dense1Relu0.data(); +std::vector fTensor_output = std::vector(32); +float * tensor_output = fTensor_output.data(); +std::vector fTensor_sequential1dense31MatMulGemm60 = std::vector(32); +float * tensor_sequential1dense31MatMulGemm60 = fTensor_sequential1dense31MatMulGemm60.data(); +std::vector fTensor_sequential1dense1MatMulGemm70 = std::vector(160); +float * tensor_sequential1dense1MatMulGemm70 = fTensor_sequential1dense1MatMulGemm70.data(); +std::vector fTensor_sequential1dense1BiasAddReadVariableOp0bcast = std::vector(160); +float * tensor_sequential1dense1BiasAddReadVariableOp0bcast = fTensor_sequential1dense1BiasAddReadVariableOp0bcast.data(); + + +Session(std::string filename ="TrackPID_v1.dat") { + +//--- reading weights from file + std::ifstream f; + f.open(filename); + if (!f.is_open()) { + throw std::runtime_error("tmva-sofie failed to open file " + filename + " for input weights"); + } + std::string tensor_name; + size_t length; + f >> tensor_name >> length; + if (tensor_name != "tensor_sequential1dense1BiasAddReadVariableOp0" ) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor name; expected name is tensor_sequential1dense1BiasAddReadVariableOp0 , read " + tensor_name; + throw std::runtime_error(err_msg); + } + if (length != 5) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor size; expected size is 5 , read " + std::to_string(length) ; + throw std::runtime_error(err_msg); + } + for (size_t i = 0; i < length; ++i) + f >> tensor_sequential1dense1BiasAddReadVariableOp0[i]; + f >> tensor_name >> length; + if (tensor_name != "tensor_sequential1dense12CastReadVariableOp0" ) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor name; expected name is tensor_sequential1dense12CastReadVariableOp0 , read " + tensor_name; + throw std::runtime_error(err_msg); + } + if (length != 50) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor size; expected size is 50 , read " + std::to_string(length) ; + throw std::runtime_error(err_msg); + } + for (size_t i = 0; i < length; ++i) + f >> tensor_sequential1dense12CastReadVariableOp0[i]; + f >> tensor_name >> length; + if (tensor_name != "tensor_sequential1dense12BiasAddReadVariableOp0" ) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor name; expected name is tensor_sequential1dense12BiasAddReadVariableOp0 , read " + tensor_name; + throw std::runtime_error(err_msg); + } + if (length != 10) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor size; expected size is 10 , read " + std::to_string(length) ; + throw std::runtime_error(err_msg); + } + for (size_t i = 0; i < length; ++i) + f >> tensor_sequential1dense12BiasAddReadVariableOp0[i]; + f >> tensor_name >> length; + if (tensor_name != "tensor_sequential1dense21BiasAddReadVariableOp0" ) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor name; expected name is tensor_sequential1dense21BiasAddReadVariableOp0 , read " + tensor_name; + throw std::runtime_error(err_msg); + } + if (length != 5) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor size; expected size is 5 , read " + std::to_string(length) ; + throw std::runtime_error(err_msg); + } + for (size_t i = 0; i < length; ++i) + f >> tensor_sequential1dense21BiasAddReadVariableOp0[i]; + f >> tensor_name >> length; + if (tensor_name != "tensor_sequential1dense1CastReadVariableOp0" ) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor name; expected name is tensor_sequential1dense1CastReadVariableOp0 , read " + tensor_name; + throw std::runtime_error(err_msg); + } + if (length != 20) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor size; expected size is 20 , read " + std::to_string(length) ; + throw std::runtime_error(err_msg); + } + for (size_t i = 0; i < length; ++i) + f >> tensor_sequential1dense1CastReadVariableOp0[i]; + f >> tensor_name >> length; + if (tensor_name != "tensor_sequential1dense21CastReadVariableOp0" ) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor name; expected name is tensor_sequential1dense21CastReadVariableOp0 , read " + tensor_name; + throw std::runtime_error(err_msg); + } + if (length != 50) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor size; expected size is 50 , read " + std::to_string(length) ; + throw std::runtime_error(err_msg); + } + for (size_t i = 0; i < length; ++i) + f >> tensor_sequential1dense21CastReadVariableOp0[i]; + f >> tensor_name >> length; + if (tensor_name != "tensor_sequential1dense31AddReadVariableOp0" ) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor name; expected name is tensor_sequential1dense31AddReadVariableOp0 , read " + tensor_name; + throw std::runtime_error(err_msg); + } + if (length != 1) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor size; expected size is 1 , read " + std::to_string(length) ; + throw std::runtime_error(err_msg); + } + for (size_t i = 0; i < length; ++i) + f >> tensor_sequential1dense31AddReadVariableOp0[i]; + f >> tensor_name >> length; + if (tensor_name != "tensor_sequential1dense31CastReadVariableOp0" ) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor name; expected name is tensor_sequential1dense31CastReadVariableOp0 , read " + tensor_name; + throw std::runtime_error(err_msg); + } + if (length != 5) { + std::string err_msg = "TMVA-SOFIE failed to read the correct tensor size; expected size is 5 , read " + std::to_string(length) ; + throw std::runtime_error(err_msg); + } + for (size_t i = 0; i < length; ++i) + f >> tensor_sequential1dense31CastReadVariableOp0[i]; + f.close(); + +//---- allocate the intermediate dynamic tensors +//--- broadcast bias tensor sequential1dense1BiasAddReadVariableOp0for Gemm op + { + float * data = TMVA::Experimental::SOFIE::UTILITY::UnidirectionalBroadcast(tensor_sequential1dense1BiasAddReadVariableOp0,{ 5 }, { 32 , 5 }); + std::copy(data, data + 160, tensor_sequential1dense1BiasAddReadVariableOp0bcast); + delete [] data; + } +//--- broadcast bias tensor sequential1dense12BiasAddReadVariableOp0for Gemm op + { + float * data = TMVA::Experimental::SOFIE::UTILITY::UnidirectionalBroadcast(tensor_sequential1dense12BiasAddReadVariableOp0,{ 10 }, { 32 , 10 }); + std::copy(data, data + 320, tensor_sequential1dense12BiasAddReadVariableOp0bcast); + delete [] data; + } +//--- broadcast bias tensor sequential1dense21BiasAddReadVariableOp0for Gemm op + { + float * data = TMVA::Experimental::SOFIE::UTILITY::UnidirectionalBroadcast(tensor_sequential1dense21BiasAddReadVariableOp0,{ 5 }, { 32 , 5 }); + std::copy(data, data + 160, tensor_sequential1dense21BiasAddReadVariableOp0bcast); + delete [] data; + } +//--- broadcast bias tensor sequential1dense31AddReadVariableOp0for Gemm op + { + float * data = TMVA::Experimental::SOFIE::UTILITY::UnidirectionalBroadcast(tensor_sequential1dense31AddReadVariableOp0,{ 1 }, { 32 , 1 }); + std::copy(data, data + 32, tensor_sequential1dense31AddReadVariableOp0bcast); + delete [] data; + } +} + +std::vector infer(float* tensor_kerastensor){ + +//--------- Gemm + char op_0_transA = 'n'; + char op_0_transB = 'n'; + int op_0_m = 32; + int op_0_n = 5; + int op_0_k = 4; + float op_0_alpha = 1; + float op_0_beta = 1; + int op_0_lda = 4; + int op_0_ldb = 5; + std::copy(tensor_sequential1dense1BiasAddReadVariableOp0bcast, tensor_sequential1dense1BiasAddReadVariableOp0bcast + 160, tensor_sequential1dense1MatMulGemm70); + BLAS::sgemm_(&op_0_transB, &op_0_transA, &op_0_n, &op_0_m, &op_0_k, &op_0_alpha, tensor_sequential1dense1CastReadVariableOp0, &op_0_ldb, tensor_kerastensor, &op_0_lda, &op_0_beta, tensor_sequential1dense1MatMulGemm70, &op_0_n); + +//------ RELU + for (int id = 0; id < 160 ; id++){ + tensor_sequential1dense1Relu0[id] = ((tensor_sequential1dense1MatMulGemm70[id] > 0 )? tensor_sequential1dense1MatMulGemm70[id] : 0); + } + +//--------- Gemm + char op_2_transA = 'n'; + char op_2_transB = 'n'; + int op_2_m = 32; + int op_2_n = 10; + int op_2_k = 5; + float op_2_alpha = 1; + float op_2_beta = 1; + int op_2_lda = 5; + int op_2_ldb = 10; + std::copy(tensor_sequential1dense12BiasAddReadVariableOp0bcast, tensor_sequential1dense12BiasAddReadVariableOp0bcast + 320, tensor_sequential1dense12MatMulGemm80); + BLAS::sgemm_(&op_2_transB, &op_2_transA, &op_2_n, &op_2_m, &op_2_k, &op_2_alpha, tensor_sequential1dense12CastReadVariableOp0, &op_2_ldb, tensor_sequential1dense1Relu0, &op_2_lda, &op_2_beta, tensor_sequential1dense12MatMulGemm80, &op_2_n); + +//------ RELU + for (int id = 0; id < 320 ; id++){ + tensor_sequential1dense12Relu0[id] = ((tensor_sequential1dense12MatMulGemm80[id] > 0 )? tensor_sequential1dense12MatMulGemm80[id] : 0); + } + +//--------- Gemm + char op_4_transA = 'n'; + char op_4_transB = 'n'; + int op_4_m = 32; + int op_4_n = 5; + int op_4_k = 10; + float op_4_alpha = 1; + float op_4_beta = 1; + int op_4_lda = 10; + int op_4_ldb = 5; + std::copy(tensor_sequential1dense21BiasAddReadVariableOp0bcast, tensor_sequential1dense21BiasAddReadVariableOp0bcast + 160, tensor_sequential1dense21MatMulGemm90); + BLAS::sgemm_(&op_4_transB, &op_4_transA, &op_4_n, &op_4_m, &op_4_k, &op_4_alpha, tensor_sequential1dense21CastReadVariableOp0, &op_4_ldb, tensor_sequential1dense12Relu0, &op_4_lda, &op_4_beta, tensor_sequential1dense21MatMulGemm90, &op_4_n); + +//------ RELU + for (int id = 0; id < 160 ; id++){ + tensor_sequential1dense21Relu0[id] = ((tensor_sequential1dense21MatMulGemm90[id] > 0 )? tensor_sequential1dense21MatMulGemm90[id] : 0); + } + +//--------- Gemm + char op_6_transA = 'n'; + char op_6_transB = 'n'; + int op_6_m = 32; + int op_6_n = 1; + int op_6_k = 5; + float op_6_alpha = 1; + float op_6_beta = 1; + int op_6_lda = 5; + int op_6_ldb = 1; + std::copy(tensor_sequential1dense31AddReadVariableOp0bcast, tensor_sequential1dense31AddReadVariableOp0bcast + 32, tensor_sequential1dense31MatMulGemm60); + BLAS::sgemm_(&op_6_transB, &op_6_transA, &op_6_n, &op_6_m, &op_6_k, &op_6_alpha, tensor_sequential1dense31CastReadVariableOp0, &op_6_ldb, tensor_sequential1dense21Relu0, &op_6_lda, &op_6_beta, tensor_sequential1dense31MatMulGemm60, &op_6_n); + for (int id = 0; id < 32 ; id++){ + tensor_output[id] = 1 / (1 + std::exp( - tensor_sequential1dense31MatMulGemm60[id])); + } + return fTensor_output; +} +}; +} //TMVA_SOFIE_TrackPID_v1 + +#endif // ROOT_TMVA_SOFIE_TRACKPID_V1 diff --git a/TrkDiag/src/TrackPID_module.cc b/TrkDiag/src/TrackPID_module.cc index 53d5c09..f0d1648 100644 --- a/TrkDiag/src/TrackPID_module.cc +++ b/TrkDiag/src/TrackPID_module.cc @@ -1,28 +1,39 @@ // // Reconstruction-level PID determination from a track. The current implementation -// just uses the calorimeter matching information (TrkCaloHit), eventually it will -// also process dE/dx, etc. +// uses the calorimeter matching information (TrkCaloHit) in addition to basic track-level +// information in an MVA. +// +// Produces a collection of PID scores per input track collection, aligned by index // // Original author: Dave Brown (LBNL) // // framework #include "art/Framework/Principal/Event.h" -#include "fhiclcpp/ParameterSet.h" #include "art/Framework/Principal/Handle.h" #include "art/Framework/Core/EDProducer.h" #include "art_root_io/TFileService.h" #include "art/Utilities/make_tool.h" +#include "fhiclcpp/ParameterSet.h" +#include "fhiclcpp/types/Atom.h" +#include "fhiclcpp/types/Sequence.h" + // utilities #include "Offline/ProditionsService/inc/ProditionsHandle.hh" #include "Offline/Mu2eUtilities/inc/MVATools.hh" #include "Offline/ConfigTools/inc/ConfigFileLookupPolicy.hh" + // data #include "Offline/RecoDataProducts/inc/KalSeed.hh" +#include "Offline/RecoDataProducts/inc/KalSeedDtDt.hh" #include "Offline/RecoDataProducts/inc/MVAResult.hh" -#include "ArtAnalysis/TrkDiag/inc/TrackPID.hxx" #include "Offline/GeometryService/inc/GeomHandle.hh" #include "Offline/CalorimeterGeom/inc/DiskCalorimeter.hh" + +// models +#include "ArtAnalysis/TrkDiag/inc/TrackPID_v0.hxx" +#include "ArtAnalysis/TrkDiag/inc/TrackPID_v1.hxx" + // C++ #include #include @@ -32,7 +43,10 @@ #include using namespace std; -namespace TMVA_SOFIE_TrackPID { +namespace TMVA_SOFIE_TrackPID_v0 { + class Session; +} +namespace TMVA_SOFIE_TrackPID_v1 { class Session; } @@ -43,13 +57,13 @@ namespace mu2e { using Name=fhicl::Name; using Comment=fhicl::Comment; - fhicl::Atom MaxDE{Name("MaxDE"), Comment("Maximum difference between calorimeter cluster EDep energy and the track energy (assuming electron mass)")}; - fhicl::Atom DT{ Name("DeltaTOffset"), - Comment("Track - Calorimeter time offset")}; // this should be a condition FIXME! - fhicl::Atom kalSeedPtrTag{Name("KalSeedPtrCollection"), Comment("Input tag for KalSeedPtrCollection")}; - fhicl::Atom printMVA{Name("printMVA"), Comment("print the MVA used"), false}; - fhicl::Atom datFilename{Name("datFilename"), Comment("Filename for the .dat file to use")}; - fhicl::Atom debug{Name("debugLevel"), Comment("Debug printout Level"), 0}; + fhicl::Atom maxDE { Name("MaxDE") , Comment("Maximum E(calo) - P(trk)")}; + fhicl::Sequence kalSeeds { Name("KalSeeds") , Comment("KalSeed (Ptr) collection names") }; + fhicl::OptionalSequence kalSeedDtDts{ Name("KalSeedDtDts"), Comment("KalSeedDtDt collection names") }; + fhicl::Atom datFilename { Name("DatFilename") , Comment("Filename for the .dat file to use")}; + fhicl::Atom MVAVersion { Name("MVAVersion") , Comment("MVA version ID to determine the input features")}; + fhicl::Atom printMVA { Name("PrintMVA") , Comment("Print the MVA used"), false}; + fhicl::Atom debug { Name("DebugLevel") , Comment("Debug printout Level"), 0}; }; using Parameters = art::EDProducer::Table; @@ -58,85 +72,220 @@ namespace mu2e { private: void produce(art::Event& event) override; void initializeMVA(std::string xmlfilename); + float evaluateMVA(const KalSeed& seed, const KalSeedDtDt* dtdt); + XYZVectorD momentumAtCalo(const KalSeed& seed); + + // MVA version-specific calls + float evaluateMVA_v0(const KalSeed& seed); + float evaluateMVA_v1(const KalSeed& seed, const KalSeedDtDt* dtdt); + + float maxDe_; + std::vector kalSeeds_; + std::vector kalSeedDtDts_; + int MVAVersion_; + bool printMVA_; + int debugLevel_; - float _maxde, _dtoffset; - art::InputTag _kalSeedPtrTag; - bool _printMVA; - int _debug; + const mu2e::Calorimeter* calo_; - std::shared_ptr mva_; + std::shared_ptr mva_v0_; + std::shared_ptr mva_v1_; }; - TrackPID::TrackPID(const Parameters& conf) : - art::EDProducer(conf), - _maxde(conf().MaxDE()), - _dtoffset(conf().DT()), - _kalSeedPtrTag(conf().kalSeedPtrTag()), - _printMVA(conf().printMVA()), - _debug(conf().debug()) + //================================================================ + TrackPID::TrackPID(const Parameters& conf) + : art::EDProducer(conf) + , maxDe_(conf().maxDE()) + , kalSeeds_(conf().kalSeeds()) + , MVAVersion_(conf().MVAVersion()) + , printMVA_(conf().printMVA()) + , debugLevel_(conf().debug()) { - produces(); + if(MVAVersion_ > 1) throw cet::exception("RECO") << "Unknown MVA feature version " << MVAVersion_; + + if(conf().kalSeedDtDts(kalSeedDtDts_)) { + if(kalSeeds_.size() != kalSeedDtDts_.size()) throw cet::exception("RECO") << "KalSeed and KalSeedDtDt lists must match"; + } else if(MVAVersion_ == 1) { + throw cet::exception("RECO") << "KalSeedDtDt is not provided but is required for MVA feature version 2"; + } + + // Produce a PID collection per track collection + for(const auto& name : kalSeeds_) { + produces(name); + } + + // Define the MVA ConfigFileLookupPolicy configFile; - mva_ = std::make_shared(configFile(conf().datFilename())); + if (MVAVersion_ == 0) mva_v0_ = std::make_shared(configFile(conf().datFilename())); + else if(MVAVersion_ == 1) mva_v1_ = std::make_shared(configFile(conf().datFilename())); + } + + //================================================================ + XYZVectorD TrackPID::momentumAtCalo(const KalSeed& seed) { + if(!seed.hasCaloCluster()) return XYZVectorD(0.,0.,0.); // no cluster + + // Find the segment that is closest to the cluster position + auto const& tchs = seed.caloHit(); + return seed.nearestSegment(tchs._rptoca)->momentum3(); + } + + //================================================================ + float TrackPID::evaluateMVA_v0(const KalSeed& seed) { + std::array features{-9999,-9999,-9999,-9999}; // features used for training + double score = -999.; // < -1 = invalid + constexpr float dtoffset = -1.15; // value used for this version of the MVA + + // Require a calo cluster + if(!seed.hasCaloCluster() || !seed.caloHit()._flag.hasAllProperties(StrawHitFlag::active)) { + if(debugLevel_ > 1) printf("[TrackPID::%s] Input track has no calo cluster\n", __func__); + return score; + } + + // Fill the features + auto const& tchs = seed.caloHit(); + auto const& cc = tchs.caloCluster(); + auto trkmom = momentumAtCalo(seed); + features[0] = cc->energyDep() - sqrt(trkmom.Mag2()); + // move into detector coordinates. Yikes!! + XYZVectorF cpos = XYZVectorF(calo_->geomUtil().mu2eToTracker(calo_->geomUtil().diskFFToMu2e( cc->diskID(), cc->cog3Vector()))); + features[1] = sqrt(cpos.Perp2()); + // compute transverse direction WRT position + cpos.SetZ(0.0); + trkmom.SetZ(0.0); + features[2] = cpos.Dot(trkmom)/sqrt(cpos.Mag2()*trkmom.Mag2()); + // the following includes the (Calibrated) light-propagation time delay. It should eventually be put in the reconstruction FIXME! + // This velocity should come from conditions FIXME! + features[3] = tchs.t0().t0()-tchs.time()- std::min((float)200.0,std::max((float)0.0,tchs.hitLen()))*0.005 - dtoffset; + // hard cut on the energy difference. This rejects cosmic rays which hit the calo and produce an upstream-going track that is then + // reconstructed as a downstream particle associated to this cluster + if(features[0] < maxDe_) { + // evaluate the MVA + auto mvaout = mva_v0_->infer(features.data()); + score = mvaout[0]; + } else if(debugLevel_ > 0) { + printf("[TrackPID::%s] Energy difference at %.3f, above threshold %.3f", __func__, features[0], maxDe_); + } + + if(debugLevel_ > 0) { + printf("[TrackPID::%s] Input features: {%.3f, %.3f, %.3f, %.3f} output: %.5f\n", __func__, + features[0], features[1], features[2], features[3], score); + } + + return score; + } + + //================================================================ + float TrackPID::evaluateMVA_v1(const KalSeed& seed, const KalSeedDtDt* dtdt) { + double score = -999.; // < -1 = invalid + + if(!dtdt) return score; + + // Require a calo cluster + if(!seed.hasCaloCluster() || !seed.caloHit()._flag.hasAllProperties(StrawHitFlag::active)) { + if(debugLevel_ > 1) printf("[TrackPID::%s] Input track has no calo cluster\n", __func__); + return score; + } + + // Retrieve the calo cluster + auto const& tchs = seed.caloHit(); + auto const& cc = tchs.caloCluster(); + const float edep = cc->energyDep(); + + // Retrieve the track momentum at the calo + const auto momVec = momentumAtCalo(seed); + const float mom = std::sqrt(momVec.Mag2()); + + // Check if this track should be skipped + if(edep - mom > maxDe_) return score; // Selection provided by the user + if(mom == 0.) return score; // E / P not defined + + // Define the features + const float ep = edep / mom; + const float dt = tchs._udt; // unbiased time difference + const float fitcon = seed.fitConsistency(); + const float slope = dtdt->slope(); + + // Assign the features + std::array features = {ep, dt, fitcon, slope}; + + // Evaluate the MVA + const auto mvaout = mva_v1_->infer(features.data()); + score = mvaout[0]; + if(debugLevel_ > 0) { + printf("[TrackPID::%s] Input features: {%.3f, %.3f, %.3f, %.3f} output: %.5f\n", __func__, + features[0], features[1], features[2], features[3], score); + } + + // Return the score + return score; + } + + //================================================================ + float TrackPID::evaluateMVA(const KalSeed& seed, const KalSeedDtDt* dtdt) { + float score = -999.; // < -1 = invalid + + // Check if it's an acceptable track fit + static TrkFitFlag goodfit(TrkFitFlag::kalmanOK); + if(!seed.status().hasAllProperties(goodfit)) return score; + + // Evaluate the MVA using the proper call for the given version + switch(MVAVersion_) { + case 0: score = evaluateMVA_v0(seed); break; + case 1: score = evaluateMVA_v1(seed, dtdt); break; + default: throw cet::exception("RECO") << "Unknown MVA feature version " << MVAVersion_; + } + + // Return the score + return score; } void TrackPID::produce(art::Event& event ) { mu2e::GeomHandle calo; - // create output - unique_ptr mvacol(new MVAResultCollection()); - // get the KalSeedsPtrs - art::Handle kalSeedPtrHandle; - event.getByLabel(_kalSeedPtrTag, kalSeedPtrHandle); - const auto& kalSeedPtrs = *kalSeedPtrHandle; - - // Go through the tracks and calculate the track PID - for (const auto& kalSeedPtr : kalSeedPtrs) { - const auto& kalSeed = *kalSeedPtr; - std::array features{-9999,-9999,-9999,-9999}; // features used for training - double mvaval = -1; - - // Fill the features - static TrkFitFlag goodfit(TrkFitFlag::kalmanOK); - if (kalSeed.status().hasAllProperties(goodfit)){ - if(kalSeed.hasCaloCluster() && kalSeed.caloHit()._flag.hasAllProperties(StrawHitFlag::active)){ - auto const& tchs = kalSeed.caloHit(); - auto const& cc = tchs.caloCluster(); - XYZVectorD trkmom = kalSeed.nearestSegment(tchs._rptoca)->momentum3(); - features[0] = cc->energyDep() - sqrt(trkmom.Mag2()); - // move into detector coordinates. Yikes!! - XYZVectorF cpos = XYZVectorF(calo->geomUtil().mu2eToTracker(calo->geomUtil().diskFFToMu2e( cc->diskID(), cc->cog3Vector()))); - features[1] = sqrt(cpos.Perp2()); - // compute transverse direction WRT position - cpos.SetZ(0.0); - trkmom.SetZ(0.0); - features[2] = cpos.Dot(trkmom)/sqrt(cpos.Mag2()*trkmom.Mag2()); - // the following includes the (Calibrated) light-propagation time delay. It should eventually be put in the reconstruction FIXME! - // This velocity should come from conditions FIXME! - features[3] = tchs.t0().t0()-tchs.time()- std::min((float)200.0,std::max((float)0.0,tchs.hitLen()))*0.005 - _dtoffset; - // hard cut on the energy difference. This rejects cosmic rays which hit the calo and produce an upstream-going track that is then - // reconstructed as a downstream particle associated to this cluster - if(features[0] < _maxde){ - // evaluate the MVA - auto mvaout = mva_->infer(features.data()); - mvaval = mvaout[0]; - } - else if (_debug > 0) { - printf("energy difference at %f, above threshold %f", features[0], _maxde); - } + calo_ = calo.get(); + + // Loop over all KalSeed collections + for(size_t index = 0; index < kalSeeds_.size(); ++index) { + const auto& name = kalSeeds_.at(index); + + // Retrieve the KalSeed collection from the event, checking if it is a collection of KalSeed or KalSeedPtr + art::Handle seedHandle; + event.getByLabel(name, seedHandle); + art::Handle seedPtrHandle; + const bool isSeedCollection = seedHandle.isValid(); + if(!isSeedCollection) { + event.getByLabel(name, seedPtrHandle); + if(!seedPtrHandle.isValid()) { + throw cet::exception("RECO") << "TrackPID: No KalSeed or KalSeedPtr collection with label " << name << std::endl; } } - if (_debug > 0) { - printf("TrackPID ; input features: %f ; %f ; %f ; %f ; output: %f", - features[0], features[1], features[2], features[3], mvaval); + const auto nseeds = (isSeedCollection) ? seedHandle->size() : seedPtrHandle->size(); + + // Retrieve the KalSeedDtDt collection, if requested + const KalSeedDtDtCollection* dtdtCollection = nullptr; + if(!kalSeedDtDts_.empty()) { + auto dtdt_handle = event.getValidHandle(kalSeedDtDts_.at(index)); + dtdtCollection = &(*dtdt_handle); + if(dtdtCollection->size() != nseeds) throw cet::exception("RECO") << "KalSeed and KalSeedDtDt collections must have matching size. N(seeds) = " + << nseeds << " N(dtdts) = " << dtdtCollection->size() + << " Seeds tag = " << name << " DtDts tag = " << kalSeedDtDts_.at(index); } - mvacol->push_back(MVAResult(mvaval)); - } - if (mvacol->size() != kalSeedPtrs.size()) { - throw cet::exception("TrackPID") << "KalSeedPtr and MVAResult sizes are inconsistent: KalSeedPtr.size() = " << kalSeedPtrs.size() << " ; MVAResult.size() = " << mvacol->size(); + + + // Create the output collection + std::unique_ptr mvaCol(new MVAResultCollection()); + + // Loop over all seeds in the collection and evaluate the PID score + if(debugLevel_ > 0) std::cout << "[TrackPID::" << __func__ << "] Processing " << nseeds << " seeds from collection " << name << std::endl; + for(size_t iseed = 0; iseed < nseeds; ++iseed) { + const auto& seed = (isSeedCollection) ? seedHandle->at(iseed) : *seedPtrHandle->at(iseed); + if(debugLevel_ > 1) std::cout << "[TrackPID::" << __func__ << "] Processing seed " << iseed << " with " << seed.hits().size() << " hits" << std::endl; + mvaCol->emplace_back(MVAResult(evaluateMVA(seed, (dtdtCollection) ? &dtdtCollection->at(iseed) : nullptr))); + } + + // Put the results into the event + event.put(std::move(mvaCol), name); } - // put the output products into the event - event.put(move(mvacol)); } }// mu2e