已合并
expm1 simt代码整改 #2707
ligen75创建于 5月12日
expm1 simt代码整改 #2707
已合并
ligen75创建于 5月12日
已删除 :master合入到cann/ops-mathmaster
1 个文件变更+21-16
@@ -21,6 +21,11 @@
21#include "atvoss/util/placeholder.h"21#include "atvoss/util/placeholder.h"
22#include <limits>22#include <limits>
23 23 
24+#ifdef __CCE_AICORE__
25+#include "simt_api/math_functions.h"
26+#include "simt_api/asc_simt.h"
27+#endif
28+ 
24namespace Expm1Op {29namespace Expm1Op {
25using namespace Ops::Base;30using namespace Ops::Base;
26const int CAST_MODE_NONE = 0;31const int CAST_MODE_NONE = 0;
@@ -46,36 +51,36 @@ template <typename T>
46__simt_vf__ __aicore__51__simt_vf__ __aicore__
47LAUNCH_BOUND(THREAD_NUM) inline void Expm1SimtCompute(__ubuf__ T* x, __ubuf__ T* y, const int64_t totalNum)52LAUNCH_BOUND(THREAD_NUM) inline void Expm1SimtCompute(__ubuf__ T* x, __ubuf__ T* y, const int64_t totalNum)
48{53{
49- for (int64_t i = Simt::GetThreadIdx(); i < totalNum; i += Simt::GetThreadNum()) {54+ for (int64_t i = threadIdx.x; i < totalNum; i += blockDim.x) {
50 float f1 = x[i];55 float f1 = x[i];
51- float f0 = Simt::Expm1(f1);56+ float f0 = expm1f(f1);
52 float f2 = f1 * INV_LN2_APPROX;57 float f2 = f1 * INV_LN2_APPROX;
53- float f3 = Simt::Round(f2);58+ float f3 = roundf(f2);
54- float f4 = Simt::Abs(f1);59+ float f4 = fabsf(f1);
55 bool p1 = f4 < LN2_HALF_APPROX;60 bool p1 = f4 < LN2_HALF_APPROX;
56 float f5 = p1 ? 0.0f : f3;61 float f5 = p1 ? 0.0f : f3;
57 float f6 = -f5;62 float f6 = -f5;
58 float f7 = LN2_APPROX;63 float f7 = LN2_APPROX;
59- float f8 = Simt::Fma(f6, f7, f1);64+ float f8 = fmaf(f6, f7, f1);
60 float f9 = ONE_MINUS_LN2_APPROX;65 float f9 = ONE_MINUS_LN2_APPROX;
61- float f10 = Simt::Fma(f6, f9, f8);66+ float f10 = fmaf(f6, f9, f8);
62 bool p2 = f5 == FLOAT_128;67 bool p2 = f5 == FLOAT_128;
63 float f11 = f5 + FLOAT_NEG_ONE;68 float f11 = f5 + FLOAT_NEG_ONE;
64 float f12 = p2 ? f11 : f5;69 float f12 = p2 ? f11 : f5;
65 float f13 = C5;70 float f13 = C5;
66 float f14 = C4;71 float f14 = C4;
67- float f15 = Simt::Fma(f14, f10, f13);72+ float f15 = fmaf(f14, f10, f13);
68 float f16 = C3;73 float f16 = C3;
69- float f17 = Simt::Fma(f15, f10, f16);74+ float f17 = fmaf(f15, f10, f16);
70 float f18 = C2;75 float f18 = C2;
71- float f19 = Simt::Fma(f17, f10, f18);76+ float f19 = fmaf(f17, f10, f18);
72 float f20 = C1;77 float f20 = C1;
73- float f21 = Simt::Fma(f19, f10, f20);78+ float f21 = fmaf(f19, f10, f20);
74 float f22 = f10 * f21;79 float f22 = f10 * f21;
75- float f23 = Simt::Fma(f22, f10, f10);80+ float f23 = fmaf(f22, f10, f10);
76- float f24 = Simt::Exp2(f12);81+ float f24 = exp2f(f12);
77 float f25 = f24 + FLOAT_NEG_ONE;82 float f25 = f24 + FLOAT_NEG_ONE;
78- float f26 = Simt::Fma(f23, f24, f25);83+ float f26 = fmaf(f23, f24, f25);
79 float f27 = f26 + f26;84 float f27 = f26 + f26;
80 float f28 = p2 ? f27 : f26;85 float f28 = p2 ? f27 : f26;
81 bool p3 = f12 > FLOAT_128;86 bool p3 = f12 > FLOAT_128;
@@ -85,7 +90,7 @@ LAUNCH_BOUND(THREAD_NUM) inline void Expm1SimtCompute(__ubuf__ T* x, __ubuf__ T*
85 bool p5 = f1 == 0.0f;90 bool p5 = f1 == 0.0f;
86 float f31 = f1 + f1;91 float f31 = f1 + f1;
87 float f32 = p5 ? f31 : f30;92 float f32 = p5 ? f31 : f30;
88- y[i] = Simt::Abs(x[i]) > FLOAT_2 ? f0 : f32;93+ y[i] = fabsf(x[i]) > FLOAT_2 ? f0 : f32;
89 }94 }
90}95}
91#endif96#endif
@@ -97,7 +102,7 @@ struct Expm1Custom : public Vec::ElemwiseUnaryOP<T, T> {
97#ifdef __CCE_AICORE__102#ifdef __CCE_AICORE__
98 __ubuf__ T* srcAddr = (__ubuf__ T*)src.GetPhyAddr();103 __ubuf__ T* srcAddr = (__ubuf__ T*)src.GetPhyAddr();
99 __ubuf__ T* dstAddr = (__ubuf__ T*)dst.GetPhyAddr();104 __ubuf__ T* dstAddr = (__ubuf__ T*)dst.GetPhyAddr();
100- Simt::VF_CALL<Expm1SimtCompute<T>>(Simt::Dim3(THREAD_NUM), srcAddr, dstAddr, count);105+ asc_vf_call<Expm1SimtCompute<T>>(dim3(THREAD_NUM), srcAddr, dstAddr, count);
101#endif106#endif
102 }107 }
103};108};
@@ -115,4 +120,4 @@ struct Expm1DAG {
115 using OpDag = DAGSch<Outputs, void, MemCfg>;120 using OpDag = DAGSch<Outputs, void, MemCfg>;
116};121};
117} // namespace Expm1Op122} // namespace Expm1Op
118-#endif // OPS_MATH_EXPM1_DAG_H123+#endif // OPS_MATH_EXPM1_DAG_H