SWDEV-345571 - Added support for half/half2 data types for warp shuffle functions

Change-Id: I276a55129d6527b73292d75eb18403acd2623c8a


[ROCm/clr commit: 32a9249d3e]
This commit is contained in:
Anusha GodavarthySurya
2023-01-25 08:56:32 +00:00
committed by Anusha Godavarthy Surya
parent 9fdbbe53dc
commit d650f7a50c
@@ -1710,6 +1710,66 @@ THE SOFTWARE.
using half = __half;
using half2 = __half2;
#endif
#if !defined(__HIPCC_RTC__)
#include "amd_device_functions.h"
#include "amd_warp_functions.h"
__device__
inline
__half __shfl(__half var, int src_lane, int width = warpSize) {
union { int i; __half h; } tmp; tmp.h = var;
tmp.i = __shfl(tmp.i, src_lane, width);
return tmp.h;
}
__device__
inline
__half2 __shfl(__half2 var, int src_lane, int width = warpSize) {
union { int i; __half2 h; } tmp; tmp.h = var;
tmp.i = __shfl(tmp.i, src_lane, width);
return tmp.h;
}
__device__
inline
__half __shfl_up(__half var, unsigned int lane_delta, int width = warpSize) {
union { int i; __half h; } tmp; tmp.h = var;
tmp.i = __shfl_up(tmp.i, lane_delta, width);
return tmp.h;
}
__device__
inline
__half2 __shfl_up(__half2 var, unsigned int lane_delta, int width = warpSize) {
union { int i; __half2 h; } tmp; tmp.h = var;
tmp.i = __shfl_up(tmp.i, lane_delta, width);
return tmp.h;
}
__device__
inline
__half __shfl_down(__half var, unsigned int lane_delta, int width = warpSize) {
union { int i; __half h; } tmp; tmp.h = var;
tmp.i = __shfl_down(tmp.i, lane_delta, width);
return tmp.h;
}
__device__
inline
__half2 __shfl_down(__half2 var, unsigned int lane_delta, int width = warpSize) {
union { int i; __half2 h; } tmp; tmp.h = var;
tmp.i = __shfl_down(tmp.i, lane_delta, width);
return tmp.h;
}
__device__
inline
__half __shfl_xor(__half var, int lane_mask, int width = warpSize) {
union { int i; __half h; } tmp; tmp.h = var;
tmp.i = __shfl_xor(tmp.i, lane_mask, width);
return tmp.h;
}
__device__
inline
__half2 __shfl_xor(__half2 var, int lane_mask, int width = warpSize) {
union { int i; __half2 h; } tmp; tmp.h = var;
tmp.i = __shfl_xor(tmp.i, lane_mask, width);
return tmp.h;
}
#endif // !defined(__HIPCC_RTC__)
#endif // defined(__cplusplus)
#elif defined(__GNUC__)
#include "hip_fp16_gcc.h"