Skip to content

Commit 4fc2ede

Browse files
committed
GPU: Template GPUChainTracking::DoTRDGPUTracking to support both o2track and gputrack types
1 parent a9b5ba5 commit 4fc2ede

6 files changed

Lines changed: 37 additions & 26 deletions

File tree

GPU/GPUTracking/Base/GPUConstantMem.h

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -81,8 +81,32 @@ struct GPUConstantMem {
8181
#ifdef GPUCA_KERNEL_DEBUGGER_OUTPUT
8282
GPUKernelDebugOutput debugOutput;
8383
#endif
84+
85+
#if defined(GPUCA_HAVE_O2HEADERS) && defined(GPUCA_NOCOMPAT)
86+
template <int I>
87+
GPUd() auto& getTRDTracker();
88+
#else // GPUCA_HAVE_O2HEADERS
89+
template <int I>
90+
GPUdi() GPUTRDTrackerGPU& getTRDTracker()
91+
{
92+
return trdTrackerGPU;
93+
}
94+
#endif // !GPUCA_HAVE_O2HEADERS
8495
};
8596

97+
#if defined(GPUCA_HAVE_O2HEADERS) && defined(GPUCA_NOCOMPAT)
98+
template <>
99+
GPUdi() auto& GPUConstantMem::getTRDTracker<0>()
100+
{
101+
return trdTrackerGPU;
102+
}
103+
template <>
104+
GPUdi() auto& GPUConstantMem::getTRDTracker<1>()
105+
{
106+
return trdTrackerO2;
107+
}
108+
#endif
109+
86110
#ifdef GPUCA_NOCOMPAT
87111
union GPUConstantMemCopyable {
88112
GPUConstantMemCopyable() {} // NOLINT: We want an empty constructor, not a default one

GPU/GPUTracking/Global/GPUChainTracking.h

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -158,7 +158,8 @@ class GPUChainTracking : public GPUChain, GPUReconstructionHelpers::helperDelega
158158
int RunTPCTrackingSlices();
159159
int RunTPCTrackingMerger(bool synchronizeOutput = true);
160160
int RunTRDTracking();
161-
int DoTRDGPUTracking();
161+
template <int I>
162+
int DoTRDGPUTracking(GPUTRDTracker* externalInstance = nullptr);
162163
int RunTPCCompression();
163164
int RunTPCDecompression();
164165
int RunRefit();

GPU/GPUTracking/Global/GPUChainTrackingTRD.cxx

Lines changed: 6 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -55,7 +55,7 @@ int GPUChainTracking::RunTRDTracking()
5555
}
5656
}
5757

58-
Tracker.DoTracking(this);
58+
DoTRDGPUTracking<GPUTRDTrackerKernels::gpuVersion>();
5959

6060
mIOPtrs.nTRDTracks = Tracker.NTracks();
6161
mIOPtrs.trdTracks = Tracker.Tracks();
@@ -64,7 +64,8 @@ int GPUChainTracking::RunTRDTracking()
6464
return 0;
6565
}
6666

67-
int GPUChainTracking::DoTRDGPUTracking()
67+
template <int I>
68+
int GPUChainTracking::DoTRDGPUTracking(GPUTRDTracker* externalInstance)
6869
{
6970
#ifdef GPUCA_HAVE_O2HEADERS
7071
bool doGPU = GetRecoStepsGPU() & RecoStep::TRDTracking;
@@ -88,3 +89,6 @@ int GPUChainTracking::DoTRDGPUTracking()
8889
#endif
8990
return (0);
9091
}
92+
93+
template int GPUChainTracking::DoTRDGPUTracking<0>(GPUTRDTracker*);
94+
template int GPUChainTracking::DoTRDGPUTracking<1>(GPUTRDTracker*);

GPU/GPUTracking/TRDTracking/GPUTRDTracker.cxx

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -260,7 +260,7 @@ void GPUTRDTracker_t<TRDTRK, PROP>::DoTracking(GPUChainTracking* chainTracking)
260260
auto timeStart = std::chrono::high_resolution_clock::now();
261261

262262
if (mRec->GetRecoStepsGPU() & GPUDataTypes::RecoStep::TRDTracking) {
263-
chainTracking->DoTRDGPUTracking();
263+
chainTracking->DoTRDGPUTracking<0>();
264264
} else {
265265
#ifdef WITH_OPENMP
266266
#pragma omp parallel for num_threads(mRec->GetProcessingSettings().ompThreads)

GPU/GPUTracking/TRDTracking/GPUTRDTrackerKernels.cxx

Lines changed: 1 addition & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -21,31 +21,10 @@
2121

2222
using namespace GPUCA_NAMESPACE::gpu;
2323

24-
#ifdef GPUCA_HAVE_O2HEADERS
25-
template <int I>
26-
GPUd() auto& getTracker(GPUTRDTrackerKernels::processorType& processors);
27-
template <>
28-
GPUdi() auto& getTracker<0>(GPUTRDTrackerKernels::processorType& processors)
29-
{
30-
return processors.trdTrackerGPU;
31-
}
32-
template <>
33-
GPUdi() auto& getTracker<1>(GPUTRDTrackerKernels::processorType& processors)
34-
{
35-
return processors.trdTrackerO2;
36-
}
37-
#else
38-
template <int I>
39-
GPUdi() GPUTRDTrackerGPU& getTracker(GPUTRDTrackerKernels::processorType& processors)
40-
{
41-
return processors.trdTrackerGPU;
42-
}
43-
#endif
44-
4524
template <int I>
4625
GPUdii() void GPUTRDTrackerKernels::Thread(int nBlocks, int nThreads, int iBlock, int iThread, GPUsharedref() GPUSharedMemory& smem, processorType& processors)
4726
{
48-
auto& trdTracker = getTracker<I>(processors);
27+
auto& trdTracker = processors.getTRDTracker<I>();
4928
GPUCA_OPENMP(parallel for if(!trdTracker.GetRec().GetProcessingSettings().ompKernels) num_threads(trdTracker.GetRec().GetProcessingSettings().ompThreads))
5029
for (int i = get_global_id(0); i < trdTracker.NTracks(); i += get_global_size(0)) {
5130
trdTracker.DoTrackingThread(i, get_global_id(0));

GPU/GPUTracking/TRDTracking/GPUTRDTrackerKernels.h

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -25,6 +25,9 @@ namespace gpu
2525
class GPUTRDTrackerKernels : public GPUKernelTemplate
2626
{
2727
public:
28+
enum K { defaultKernel = 0,
29+
gpuVersion = 0,
30+
o2Version = 1 };
2831
GPUhdi() CONSTEXPR static GPUDataTypes::RecoStep GetRecoStep() { return GPUCA_RECO_STEP::TRDTracking; }
2932
template <int iKernel = defaultKernel>
3033
GPUd() static void Thread(int nBlocks, int nThreads, int iBlock, int iThread, GPUsharedref() GPUSharedMemory& smem, processorType& processors);

0 commit comments

Comments
 (0)