|
49 | 49 | #include <TRandom.h> |
50 | 50 | #include <TString.h> |
51 | 51 |
|
| 52 | +#include <algorithm> |
52 | 53 | #include <array> |
53 | 54 | #include <chrono> |
54 | 55 | #include <cstddef> |
@@ -570,11 +571,12 @@ class pidTPCModule |
570 | 571 | std::unique_ptr<float[]> networkPrediction(new float[predictionSize * NParticleTypes]); // For each mass hypotheses |
571 | 572 |
|
572 | 573 | const float nNclNormalization = response->GetNClNormalization(); |
573 | | - float durationNetwork = 0; |
| 574 | + float durationNetwork = 0.f; |
574 | 575 |
|
575 | 576 | std::vector<float> trackProperties(trackPropSize); |
| 577 | + std::vector<float> outputNetwork; // output buffer, allocation is reused for all mass hypotheses |
576 | 578 | uint64_t counterTrackProps = 0; |
577 | | - int loopCounter = 0; |
| 579 | + uint64_t loopCounter = 0; |
578 | 580 |
|
579 | 581 | // To load the Hadronic rate once for each collision |
580 | 582 | std::vector<float> hadronicRateForCollision(collisions.size(), 0.0f); |
@@ -626,14 +628,13 @@ class pidTPCModule |
626 | 628 | } |
627 | 629 |
|
628 | 630 | const auto startNetworkEval = std::chrono::high_resolution_clock::now(); |
629 | | - const float* const outputNetwork = network.evalModel(trackProperties); |
| 631 | + network.evalModel(trackProperties, outputNetwork); |
630 | 632 | const auto stopNetworkEval = std::chrono::high_resolution_clock::now(); |
631 | 633 | durationNetwork += std::chrono::duration<float, std::ratio<1, NanoToOne>>(stopNetworkEval - startNetworkEval).count(); |
632 | | - for (uint64_t kPrediction = 0; kPrediction < predictionSize; kPrediction += outputDimensions) { |
633 | | - for (int lOutputDim = 0; lOutputDim < outputDimensions; ++lOutputDim) { |
634 | | - networkPrediction[kPrediction + lOutputDim + predictionSize * loopCounter] = outputNetwork[kPrediction + lOutputDim]; |
635 | | - } |
| 634 | + if (outputNetwork.size() != predictionSize) { |
| 635 | + LOG(fatal) << "Network output size (" << outputNetwork.size() << ") does not match the expected prediction size (" << predictionSize << ")"; |
636 | 636 | } |
| 637 | + std::copy(outputNetwork.begin(), outputNetwork.end(), networkPrediction.get() + predictionSize * loopCounter); |
637 | 638 |
|
638 | 639 | counterTrackProps = 0; |
639 | 640 | ++loopCounter; |
|
0 commit comments