Skip to content

Commit 5e1b1f9

Browse files
committed
DetectorsBase: keep the Propagator's double off Metal
The double overloads of getFieldXYZ and getBz were already guarded where they are declared, but not where they are defined. The explicit double in the crossing-point helper becomes GPUdoubleValue, and the differences of nearly equal crossing and centre coordinates that feed atan2 become GPUdoubleCalc, matching the phiCross and dphi next to them. Both aliases are double everywhere except Metal, so the preprocessed source is unchanged for every other backend and for the host.
1 parent 3e299bb commit 5e1b1f9

1 file changed

Lines changed: 9 additions & 5 deletions

File tree

‎Detectors/Base/src/Propagator.cxx‎

Lines changed: 9 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -586,20 +586,20 @@ GPUd() bool PropagatorImpl<value_T>::propagateToR(track_T& track, value_type r,
586586
GPUdoubleCalc phiCross[2] = {}, dphi[2] = {};
587587
auto curv = track.getCurvature(bz);
588588
bool clockwise = curv < 0; // q+ in B+ or q- in B- goes clockwise
589-
auto phiLoc = math_utils::detail::asin<double>(track.getSnp());
589+
auto phiLoc = math_utils::detail::asin<o2::gpu::GPUdoubleValue>(track.getSnp());
590590
auto phi0 = phiLoc + track.getAlpha();
591591
o2::math_utils::detail::bringTo02Pi(phi0);
592592
for (int i = 0; i < cross.nDCA; i++) {
593593
// track pT direction angle at crossing points:
594594
// == angle of the tangential to track circle at the crossing point X,Y
595595
// == normal to the radial vector from the track circle center {X-cX, Y-cY}
596596
// i.e. the angle of the vector {Y-cY, -(X-cx)}
597-
auto normX = double(cross.yDCA[i]) - double(traux.yC), normY = -(double(cross.xDCA[i]) - double(traux.xC));
597+
auto normX = o2::gpu::GPUdoubleCalc(cross.yDCA[i]) - o2::gpu::GPUdoubleCalc(traux.yC), normY = -(o2::gpu::GPUdoubleCalc(cross.xDCA[i]) - o2::gpu::GPUdoubleCalc(traux.xC));
598598
if (!clockwise) {
599599
normX = -normX;
600600
normY = -normY;
601601
}
602-
phiCross[i] = math_utils::detail::atan2<double>(normY, normX);
602+
phiCross[i] = math_utils::detail::atan2<o2::gpu::GPUdoubleValue>(normY, normX);
603603
o2::math_utils::detail::bringTo02Pi(phiCross[i]);
604604
dphi[i] = phiCross[i] - phi0;
605605
if (dphi[i] > o2::constants::math::PI) {
@@ -615,7 +615,7 @@ GPUd() bool PropagatorImpl<value_T>::propagateToR(track_T& track, value_type r,
615615
auto phiLocFin = phiLoc + deltaPhi;
616616
// case1
617617
if (math_utils::detail::abs<value_type>(phiLocFin) < MaxPhiLocSafe) { // just 1 step propagation
618-
auto deltaX = (math_utils::detail::sin<double>(phiLocFin) - track.getSnp()) / track.getCurvature(bz);
618+
auto deltaX = (math_utils::detail::sin<o2::gpu::GPUdoubleValue>(phiLocFin) - track.getSnp()) / track.getCurvature(bz);
619619
if (!propagateTo(track, track.getX() + deltaX, bzOnly, maxSnp, maxStep, matCorr, tofInfo, signCorr)) {
620620
return false;
621621
}
@@ -638,7 +638,7 @@ GPUd() bool PropagatorImpl<value_T>::propagateToR(track_T& track, value_type r,
638638

639639
// propagate to phiLoc = +-MaxPhiLocSafe
640640
auto tgtPhiLoc = deltaPhi > 0 ? MaxPhiLocSafe : -MaxPhiLocSafe;
641-
auto deltaX = (math_utils::detail::sin<double>(tgtPhiLoc) - track.getSnp()) / track.getCurvature(bz);
641+
auto deltaX = (math_utils::detail::sin<o2::gpu::GPUdoubleValue>(tgtPhiLoc) - track.getSnp()) / track.getCurvature(bz);
642642
if (!propagateTo(track, track.getX() + deltaX, bzOnly, maxSnp, maxStep, matCorr, tofInfo, signCorr)) {
643643
return false;
644644
}
@@ -1094,23 +1094,27 @@ GPUd() void PropagatorImpl<value_T>::getFieldXYZ(const math_utils::Point3D<float
10941094
getFieldXYZImpl<float>(xyz, bxyz);
10951095
}
10961096

1097+
#ifndef __METAL__ // MSL has no double; the float twin remains
10971098
template <typename value_T>
10981099
GPUd() void PropagatorImpl<value_T>::getFieldXYZ(const math_utils::Point3D<double> xyz, double* bxyz) const
10991100
{
11001101
getFieldXYZImpl<double>(xyz, bxyz);
11011102
}
1103+
#endif
11021104

11031105
template <typename value_T>
11041106
GPUd() float PropagatorImpl<value_T>::getBz(const math_utils::Point3D<float> xyz) const
11051107
{
11061108
return getBzImpl<float>(xyz);
11071109
}
11081110

1111+
#ifndef __METAL__ // MSL has no double; the float twin remains
11091112
template <typename value_T>
11101113
GPUd() double PropagatorImpl<value_T>::getBz(const math_utils::Point3D<double> xyz) const
11111114
{
11121115
return getBzImpl<double>(xyz);
11131116
}
1117+
#endif
11141118

11151119
namespace o2::base
11161120
{

0 commit comments

Comments
 (0)