dl
dl_array_simd.h
浏览该文件的文档.
1
13#pragma once
14
15#include "base/dl_array.h"
16#include "base/dl_type_op.h"
17
18#if (DL_CPU_ARCH == DL_CPU_ARCH_AMD64 || DL_CPU_ARCH == DL_CPU_ARCH_X86)
19#include <simde/x86/sse.h>
20#else
21namespace dl
22{
23 using simde__m128 = Array<float, 4>;
24
25 inline simde__m128 simde_mm_load_ps(float arr[4])
26 {
27 simde__m128 ret;
28 ret[0] = arr[0];
29 ret[1] = arr[1];
30 ret[2] = arr[2];
31 ret[3] = arr[3];
32 return ret;
33 }
34 inline simde__m128 simde_mm_set1_ps(float a)
35 {
36 simde__m128 ret;
37 ret[0] = a;
38 ret[1] = a;
39 ret[2] = a;
40 ret[3] = a;
41 return ret;
42 }
43 inline void simde_mm_store_ps(simde__m128& ret, const simde__m128& a)
44 {
45 ret = a;
46 }
47
48 inline simde__m128 simde_mm_add_ps(const simde__m128& a, const simde__m128& b)
49 {
50 return a + b;
51 }
52 inline simde__m128 simde_mm_sub_ps(const simde__m128& a, const simde__m128& b)
53 {
54 return a - b;
55 }
56 inline simde__m128 simde_mm_mul_ps(const simde__m128& a, const simde__m128& b)
57 {
58 return a * b;
59 }
60 inline simde__m128 simde_mm_div_ps(const simde__m128& a, const simde__m128& b)
61 {
62 return a / b;
63 }
64}
65#endif
66
67namespace dl
68{
69 template<typename T, size_t N>
70 requires std::is_arithmetic_v<T>
71 class alignas(16) ArrayS2;
72
73 template<size_t N>
74 ArrayS2<float, N> simd_m128u_op(simde__m128(*op)(simde__m128, simde__m128),
75 const ArrayS2<float, N>& a, const Array<float, 2>& b)
76 {
78 if constexpr (N == 4)
79 {
80 simde__m128 vec_ax = simde_mm_load_ps(a.x);
81 simde__m128 vec_bx = simde_mm_set1_ps(b.x);
82 vec_ax = op(vec_ax, vec_bx);
83 simde_mm_store_ps(ret.x, vec_ax);
84
85 simde__m128 vec_ay = simde_mm_load_ps(a.y);
86 simde__m128 vec_by = simde_mm_set1_ps(b.y);
87 vec_ay = op(vec_ay, vec_by);
88 simde_mm_store_ps(ret.y, vec_ay);
89 }
90 else
91 {
92 static_assert(N == 4);
93 }
94 return ret;
95 }
96}
97
98namespace dl
99{
100template<typename T, size_t N>
101 requires std::is_arithmetic_v<T>
102class alignas(16) ArrayS2
103{
104public:
105 union
106 {
107 T _data[N * 2];
108 struct
109 {
110 T x[N];
111 T y[N];
112 };
113 };
114
115
116 ArrayS2() = default;
117
118 // slow
119 constexpr Array<T, 2> Get(size_t i) const
120 {
121 return { x[i], y[i] };
122 }
123 constexpr Array<T, 2> Sum(size_t i, size_t j) const
124 {
125 Array<T, 2> ret;
126 ret[0] = x[i] + x[j];
127 ret[1] = y[i] + y[j];
128 return ret;
129 }
130 /*constexpr const Array<T, 2> operator[](size_t i) const
131 {
132 return { x[i], y[i] };
133 }*/
134
135 constexpr ArrayS2(std::initializer_list<Array<T, 2>> data)
136 {
137 size_t N_LIST = std::size(data);
138 for (size_t i = 0; i < N; ++i)
139 x[i] = i < N_LIST ? data.x : T{};
140 for (size_t i = 0; i < N; ++i)
141 y[i] = i < N_LIST ? data.y : T{};
142 }
143};
144
145template<typename T, std::size_t N>
147{
148 if constexpr (N == 4)
149 a = a + b;
150 else
151 {
152 for (size_t i = 0; i < N; ++i)
153 {
154 a.x[i] += b[0];
155 a.y[i] += b[1];
156 }
157 }
158 return a;
159}
160template<typename T, std::size_t N>
162{
163 if constexpr (N == 4)
164 return simd_m128u_op(simde_mm_add_ps, a, b);
165 else
166 {
167 ArrayS2<T, N> ret = a;
168 ret += b;
169 return ret;
170 }
171}
172
173template<typename T, std::size_t N>
175{
176 if constexpr (N == 4)
177 a = a - b;
178 else
179 {
180 for (size_t i = 0; i < N; ++i)
181 {
182 a.x[i] -= b[0];
183 a.y[i] -= b[1];
184 }
185 }
186 return a;
187}
188template<typename T, std::size_t N>
190{
191 if constexpr (N == 4)
192 return simd_m128u_op(simde_mm_sub_ps, a, b);
193 else
194 {
195 ArrayS2<T, N> ret = a;
196 ret -= b;
197 return ret;
198 }
199}
200
201template<typename T, std::size_t N>
203{
204 if constexpr (N == 4)
205 a = a * b;
206 else
207 {
208 for (size_t i = 0; i < N; ++i)
209 {
210 a.x[i] *= b[0];
211 a.y[i] *= b[1];
212 }
213 }
214 return a;
215}
216template<typename T, std::size_t N>
218{
219 if constexpr (N == 4)
220 return simd_m128u_op(simde_mm_mul_ps, a, b);
221 else
222 {
223 ArrayS2<T, N> ret = a;
224 ret *= b;
225 return ret;
226 }
227}
228
229template<typename T, std::size_t N>
231{
232 if constexpr (N == 4)
233 a = a / b;
234 else
235 {
236 for (size_t i = 0; i < N; ++i)
237 {
238 a.x[i] /= b[0];
239 a.y[i] /= b[1];
240 }
241 }
242 return a;
243}
244template<typename T, std::size_t N>
246{
247 if constexpr (N == 4)
248 return simd_m128u_op(simde_mm_div_ps, a, b);
249 else
250 {
251 ArrayS2<T, N> ret = a;
252 ret /= b;
253 return ret;
254 }
255}
256
257
258template<size_t M = 1, typename T, size_t N>
259constexpr ArrayS2<T, N - M> toShorter(const ArrayS2<T, N>& a)
260{
261 static_assert(M <= N);
262 ArrayS2<T, N - M> ret;
263 for (size_t i = 0; i < N - M; ++i)
264 {
265 ret.x[i] = a.x[i];
266 ret.y[i] = a.y[i];
267 }
268 return ret;
269}
270
271}
ArrayS2()=default
constexpr Array< T, 2 > Get(size_t i) const
constexpr Array< T, 2 > Sum(size_t i, size_t j) const
constexpr ArrayS2(std::initializer_list< Array< T, 2 > > data)
类似std::array,增加xyzw、wh等成员访问
通用类型运算
ArrayS2< float, N > simd_m128u_op(simde__m128(*op)(simde__m128, simde__m128), const ArrayS2< float, N > &a, const Array< float, 2 > &b)
constexpr ArrayS2< T, N > & operator/=(ArrayS2< T, N > &a, const Array< T, 2 > &b)
constexpr ArrayS2< T, N - M > toShorter(const ArrayS2< T, N > &a)
constexpr ArrayS2< T, N > & operator*=(ArrayS2< T, N > &a, const Array< T, 2 > &b)
constexpr ArrayS2< T, N > & operator+=(ArrayS2< T, N > &a, const Array< T, 2 > &b)
constexpr ArrayS2< T, N > operator+(const ArrayS2< T, N > &a, const Array< T, 2 > &b)
constexpr ArrayS2< T, N > operator-(const ArrayS2< T, N > &a, const Array< T, 2 > &b)
constexpr ArrayS2< T, N > operator/(const ArrayS2< T, N > &a, const Array< T, 2 > &b)
constexpr ArrayS2< T, N > & operator-=(ArrayS2< T, N > &a, const Array< T, 2 > &b)
constexpr ArrayS2< T, N > operator*(const ArrayS2< T, N > &a, const Array< T, 2 > &b)