native 0.0.1
Vectors, masks and wide register packs for C++26
Loading...
Searching...
No Matches
native.arm.bf16.ccm
1// SPDX-License-Identifier: BSD-2-Clause OR Apache-2.0
2module;
3#include "native/isa_import.h"
4#include "native/arm/bf16.h"
5#include "native/arm/bf16_constexpr.h"
6export module native.arm.bf16;
7export import native.arm.features;
8export import native.simd;
9#if NATIVE_HOST_NEON || defined(NATIVE_DOXYGEN)
10namespace native::detail {
11 template<int Lane,std::size_t N,std::size_t M,isa<arm> Arch>
12 consteval simd<float,N,Arch> bfdot_value(simd<float,N,Arch> acc,
13 simd<bf16,2*N,Arch> a,simd<bf16,M,Arch> b) noexcept {
14 std::array<std::uint32_t,N> c{};
15 std::array<std::uint16_t,2*N> av{};
16 std::array<std::uint16_t,M> bv{};
17 acc.store_bits(c.data());a.store_bits(av.data());b.store_bits(bv.data());
18 for(std::size_t i=0;i<N;++i) {
19 auto j=Lane<0?2*i:2*Lane;
20 c[i]=arm_bfdot_bits(c[i],av[2*i],av[2*i+1],bv[j],bv[j+1]);
21 }
22 return simd<float,N,Arch>::load_bits(c.data());
23 }
24 template<isa<arm> Arch>
25 consteval simd<float,4,Arch> bfmmla_value(simd<float,4,Arch> acc,
26 simd<bf16,8,Arch> a,simd<bf16,8,Arch> b) noexcept {
27 std::array<std::uint32_t,4> c{};
28 std::array<std::uint16_t,8> av{},bv{};
29 acc.store_bits(c.data());a.store_bits(av.data());b.store_bits(bv.data());
30 for(unsigned row=0;row<2;++row) for(unsigned column=0;column<2;++column)
31 for(unsigned pair=0;pair<2;++pair) {
32 auto i=4*row+2*pair,j=4*column+2*pair;
33 c[2*row+column]=arm_bfdot_bits(c[2*row+column],av[i],av[i+1],bv[j],bv[j+1]);
34 }
35 return simd<float,4,Arch>::load_bits(c.data());
36 }
37 template<bool Top,int Lane,std::size_t M,isa<arm> Arch>
38 consteval simd<float,4,Arch> bfmlal_value(simd<float,4,Arch> acc,
39 simd<bf16,8,Arch> a,simd<bf16,M,Arch> b) noexcept {
40 std::array<std::uint32_t,4> c{};
41 std::array<std::uint16_t,8> av{};
42 std::array<std::uint16_t,M> bv{};
43 acc.store_bits(c.data());a.store_bits(av.data());b.store_bits(bv.data());
44 for(unsigned i=0;i<4;++i) {
45 auto j=2*i+Top;
46 c[i]=constexpr_float::fma_bits<constexpr_float::binary32>(
47 std::uint32_t(av[j])<<16,std::uint32_t(bv[Lane<0?j:Lane])<<16,c[i]);
48 }
49 return simd<float,4,Arch>::load_bits(c.data());
50 }
51}
52export namespace native {
67
69 template<isa<arm> Arch>
70 requires(Arch.has(arm_feature::neon_bf16))
71 native_nodiscard native_inline __attribute__((target("bf16")))
72 constexpr simd<float, 2, Arch> bfdot(simd<float, 2, Arch> acc, simd<bf16, 4, Arch> a, simd<bf16, 4, Arch> b) noexcept {
73 if consteval { return detail::bfdot_value<-1>(acc,a,b); } else {
74 auto result = detail::arm_bf16::bfdot<Arch>(vget_low_f32(acc.to_storage().to_native()), __builtin_bit_cast(bfloat16x4_t, a.to_native()), __builtin_bit_cast(bfloat16x4_t, b.to_native()));
76 simd<float, 4, Arch>::from_native(vcombine_f32(result, vdup_n_f32(0.f))));
77 }
78 }
79
81 template<isa<arm> Arch>
82 requires(requires { sizeof(simd<float, 2, Arch>); } && requires { sizeof(simd<bf16, 4, Arch>); } && !(Arch.has(arm_feature::neon_bf16)))
83 native_nodiscard consteval
85 return detail::bfdot_value<-1>(acc,a,b);
86 }
87
89 template<isa<arm> Arch, unsigned Lane>
90 requires(Arch.has(arm_feature::neon_bf16) && Lane < 2)
91 native_nodiscard native_inline __attribute__((target("bf16")))
93 if consteval { return detail::bfdot_value<Lane>(acc,a,b); } else {
94 auto result = detail::arm_bf16::bfdot_lane<Arch, Lane>(vget_low_f32(acc.to_storage().to_native()), __builtin_bit_cast(bfloat16x4_t, a.to_native()), __builtin_bit_cast(bfloat16x4_t, b.to_native()));
96 simd<float, 4, Arch>::from_native(vcombine_f32(result, vdup_n_f32(0.f))));
97 }
98 }
99
101 template<isa<arm> Arch, unsigned Lane>
102 requires(requires { sizeof(simd<float, 2, Arch>); } && requires { sizeof(simd<bf16, 4, Arch>); } && !(Arch.has(arm_feature::neon_bf16)) && Lane < 2)
103 native_nodiscard consteval
105 return detail::bfdot_value<Lane>(acc,a,b);
106 }
107
109 template<isa<arm> Arch, unsigned Lane>
110 requires(Arch.has(arm_feature::neon_bf16) && Lane < 4)
111 native_nodiscard native_inline __attribute__((target("bf16")))
113 if consteval { return detail::bfdot_value<Lane>(acc,a,b); } else {
114 auto result = detail::arm_bf16::bfdot_lane<Arch, Lane>(vget_low_f32(acc.to_storage().to_native()), __builtin_bit_cast(bfloat16x4_t, a.to_native()), b.to_native());
116 simd<float, 4, Arch>::from_native(vcombine_f32(result, vdup_n_f32(0.f))));
117 }
118 }
119
121 template<isa<arm> Arch, unsigned Lane>
122 requires(requires { sizeof(simd<float, 2, Arch>); } && requires { sizeof(simd<bf16, 4, Arch>); } && requires { sizeof(simd<bf16, 8, Arch>); } && !(Arch.has(arm_feature::neon_bf16)) && Lane < 4)
123 native_nodiscard consteval
125 return detail::bfdot_value<Lane>(acc,a,b);
126 }
127
129 template<isa<arm> Arch>
130 requires(Arch.has(arm_feature::neon_bf16))
131 native_nodiscard native_inline __attribute__((target("bf16")))
133 if consteval { return detail::bfdot_value<-1>(acc,a,b); } else {
134 auto result = detail::arm_bf16::bfdot<Arch>(acc.to_native(), a.to_native(), b.to_native());
136 }
137 }
138
140 template<isa<arm> Arch>
141 requires(requires { sizeof(simd<float, 4, Arch>); } && requires { sizeof(simd<bf16, 8, Arch>); } && !(Arch.has(arm_feature::neon_bf16)))
142 native_nodiscard consteval
144 return detail::bfdot_value<-1>(acc,a,b);
145 }
146
148 template<isa<arm> Arch, unsigned Lane>
149 requires(Arch.has(arm_feature::neon_bf16) && Lane < 2)
150 native_nodiscard native_inline __attribute__((target("bf16")))
152 if consteval { return detail::bfdot_value<Lane>(acc,a,b); } else {
153 auto result = detail::arm_bf16::bfdot_lane<Arch, Lane>(acc.to_native(), a.to_native(), __builtin_bit_cast(bfloat16x4_t, b.to_native()));
155 }
156 }
157
159 template<isa<arm> Arch, unsigned Lane>
160 requires(requires { sizeof(simd<float, 4, Arch>); } && requires { sizeof(simd<bf16, 8, Arch>); } && requires { sizeof(simd<bf16, 4, Arch>); } && !(Arch.has(arm_feature::neon_bf16)) && Lane < 2)
161 native_nodiscard consteval
163 return detail::bfdot_value<Lane>(acc,a,b);
164 }
165
167 template<isa<arm> Arch, unsigned Lane>
168 requires(Arch.has(arm_feature::neon_bf16) && Lane < 4)
169 native_nodiscard native_inline __attribute__((target("bf16")))
171 if consteval { return detail::bfdot_value<Lane>(acc,a,b); } else {
172 auto result = detail::arm_bf16::bfdot_lane<Arch, Lane>(acc.to_native(), a.to_native(), b.to_native());
174 }
175 }
176
178 template<isa<arm> Arch, unsigned Lane>
179 requires(requires { sizeof(simd<float, 4, Arch>); } && requires { sizeof(simd<bf16, 8, Arch>); } && !(Arch.has(arm_feature::neon_bf16)) && Lane < 4)
180 native_nodiscard consteval
182 return detail::bfdot_value<Lane>(acc,a,b);
183 }
184
186 template<isa<arm> Arch>
187 requires(Arch.has(arm_feature::neon_bf16))
188 native_nodiscard native_inline __attribute__((target("bf16")))
190 if consteval { return detail::bfmmla_value(acc,a,b); } else {
191 auto result = detail::arm_bf16::bfmmla<Arch>(acc.to_native(), a.to_native(), b.to_native());
193 }
194 }
195
197 template<isa<arm> Arch>
198 requires(requires { sizeof(simd<float, 4, Arch>); } && requires { sizeof(simd<bf16, 8, Arch>); } && !(Arch.has(arm_feature::neon_bf16)))
199 native_nodiscard consteval
201 return detail::bfmmla_value(acc,a,b);
202 }
203
205 template<isa<arm> Arch>
206 requires(Arch.has(arm_feature::neon_bf16))
207 native_nodiscard native_inline __attribute__((target("bf16")))
209 if consteval { return detail::bfmlal_value<false,-1>(acc,a,b); } else {
210 auto result = detail::arm_bf16::bfmlalb<Arch>(acc.to_native(), a.to_native(), b.to_native());
212 }
213 }
214
216 template<isa<arm> Arch>
217 requires(requires { sizeof(simd<float, 4, Arch>); } && requires { sizeof(simd<bf16, 8, Arch>); } && !(Arch.has(arm_feature::neon_bf16)))
218 native_nodiscard consteval
220 return detail::bfmlal_value<false,-1>(acc,a,b);
221 }
222
224 template<isa<arm> Arch, unsigned Lane>
225 requires(Arch.has(arm_feature::neon_bf16) && Lane < 4)
226 native_nodiscard native_inline __attribute__((target("bf16")))
228 if consteval { return detail::bfmlal_value<false,Lane>(acc,a,b); } else {
229 auto result = detail::arm_bf16::bfmlalb_lane<Arch, Lane>(acc.to_native(), a.to_native(), __builtin_bit_cast(bfloat16x4_t, b.to_native()));
231 }
232 }
233
235 template<isa<arm> Arch, unsigned Lane>
236 requires(requires { sizeof(simd<float, 4, Arch>); } && requires { sizeof(simd<bf16, 8, Arch>); } && requires { sizeof(simd<bf16, 4, Arch>); } && !(Arch.has(arm_feature::neon_bf16)) && Lane < 4)
237 native_nodiscard consteval
239 return detail::bfmlal_value<false,Lane>(acc,a,b);
240 }
241
243 template<isa<arm> Arch, unsigned Lane>
244 requires(Arch.has(arm_feature::neon_bf16) && Lane < 8)
245 native_nodiscard native_inline __attribute__((target("bf16")))
247 if consteval { return detail::bfmlal_value<false,Lane>(acc,a,b); } else {
248 auto result = detail::arm_bf16::bfmlalb_lane<Arch, Lane>(acc.to_native(), a.to_native(), b.to_native());
250 }
251 }
252
254 template<isa<arm> Arch, unsigned Lane>
255 requires(requires { sizeof(simd<float, 4, Arch>); } && requires { sizeof(simd<bf16, 8, Arch>); } && !(Arch.has(arm_feature::neon_bf16)) && Lane < 8)
256 native_nodiscard consteval
258 return detail::bfmlal_value<false,Lane>(acc,a,b);
259 }
260
262 template<isa<arm> Arch>
263 requires(Arch.has(arm_feature::neon_bf16))
264 native_nodiscard native_inline __attribute__((target("bf16")))
266 if consteval { return detail::bfmlal_value<true,-1>(acc,a,b); } else {
267 auto result = detail::arm_bf16::bfmlalt<Arch>(acc.to_native(), a.to_native(), b.to_native());
269 }
270 }
271
273 template<isa<arm> Arch>
274 requires(requires { sizeof(simd<float, 4, Arch>); } && requires { sizeof(simd<bf16, 8, Arch>); } && !(Arch.has(arm_feature::neon_bf16)))
275 native_nodiscard consteval
277 return detail::bfmlal_value<true,-1>(acc,a,b);
278 }
279
281 template<isa<arm> Arch, unsigned Lane>
282 requires(Arch.has(arm_feature::neon_bf16) && Lane < 4)
283 native_nodiscard native_inline __attribute__((target("bf16")))
285 if consteval { return detail::bfmlal_value<true,Lane>(acc,a,b); } else {
286 auto result = detail::arm_bf16::bfmlalt_lane<Arch, Lane>(acc.to_native(), a.to_native(), __builtin_bit_cast(bfloat16x4_t, b.to_native()));
288 }
289 }
290
292 template<isa<arm> Arch, unsigned Lane>
293 requires(requires { sizeof(simd<float, 4, Arch>); } && requires { sizeof(simd<bf16, 8, Arch>); } && requires { sizeof(simd<bf16, 4, Arch>); } && !(Arch.has(arm_feature::neon_bf16)) && Lane < 4)
294 native_nodiscard consteval
296 return detail::bfmlal_value<true,Lane>(acc,a,b);
297 }
298
300 template<isa<arm> Arch, unsigned Lane>
301 requires(Arch.has(arm_feature::neon_bf16) && Lane < 8)
302 native_nodiscard native_inline __attribute__((target("bf16")))
304 if consteval { return detail::bfmlal_value<true,Lane>(acc,a,b); } else {
305 auto result = detail::arm_bf16::bfmlalt_lane<Arch, Lane>(acc.to_native(), a.to_native(), b.to_native());
307 }
308 }
309
311 template<isa<arm> Arch, unsigned Lane>
312 requires(requires { sizeof(simd<float, 4, Arch>); } && requires { sizeof(simd<bf16, 8, Arch>); } && !(Arch.has(arm_feature::neon_bf16)) && Lane < 8)
313 native_nodiscard consteval
315 return detail::bfmlal_value<true,Lane>(acc,a,b);
316 }
317
318 // Exact deduction rejects unrelated vectors and invalid immediates before
319 // Clang's lax vector conversions can select an overload for another shape.
321 template<isa<arm> Arch, class R, class A, class B> void bfdot(R, A, B) = delete;
322 template<isa<arm> Arch, class R, class A, class B> void bfmmla(R, A, B) = delete;
323 template<isa<arm> Arch, class R, class A, class B> void bfmlalb(R, A, B) = delete;
324 template<isa<arm> Arch, class R, class A, class B> void bfmlalt(R, A, B) = delete;
325 template<isa<arm> Arch, unsigned Lane, class R, class A, class B> void bfdot_lane(R, A, B) = delete;
326 template<isa<arm> Arch, unsigned Lane, class R, class A, class B> void bfmlalb_lane(R, A, B) = delete;
327 template<isa<arm> Arch, unsigned Lane, class R, class A, class B> void bfmlalt_lane(R, A, B) = delete;
328
331}
332#endif
constexpr simd< float, 4, Arch > bfmlalt_lane(simd< float, 4, Arch > acc, simd< bf16, 8, Arch > a, simd< bf16, 4, Arch > b) noexcept
Fused multiply-add of the odd lanes of a with b[Lane] shared by all output lanes.
constexpr simd< float, 2, Arch > bfdot_lane(simd< float, 2, Arch > acc, simd< bf16, 4, Arch > a, simd< bf16, 4, Arch > b) noexcept
BFDOT with the BF16 pair b[2*Lane], b[2*Lane+1] shared by all output lanes.
constexpr simd< float, 4, Arch > bfmlalb(simd< float, 4, Arch > acc, simd< bf16, 8, Arch > a, simd< bf16, 8, Arch > b) noexcept
Fused multiply-add of the even BF16 lanes into the corresponding FP32 lanes.
constexpr simd< float, 4, Arch > bfmlalb_lane(simd< float, 4, Arch > acc, simd< bf16, 8, Arch > a, simd< bf16, 4, Arch > b) noexcept
Fused multiply-add of the even lanes of a with b[Lane] shared by all output lanes.
constexpr simd< float, 4, Arch > bfmmla(simd< float, 4, Arch > acc, simd< bf16, 8, Arch > a, simd< bf16, 8, Arch > b) noexcept
Accumulate a row-major 2x4 matrix times a column-major 4x2 matrix, two BFDOT steps per result.
constexpr simd< float, 4, Arch > bfmlalt(simd< float, 4, Arch > acc, simd< bf16, 8, Arch > a, simd< bf16, 8, Arch > b) noexcept
Fused multiply-add of the odd BF16 lanes into the corresponding FP32 lanes.
constexpr simd< float, 2, Arch > bfdot(simd< float, 2, Arch > acc, simd< bf16, 4, Arch > a, simd< bf16, 4, Arch > b) noexcept
Accumulate each adjacent pair of BF16 products into the corresponding FP32 lane.
#define native_inline
inline [[always_inline]]
Definition attributes.h:212
#define native_nodiscard
C++17 [[nodiscard]].
Definition attributes.h:189
Architecture-tagged vectors, register packs and supporting value types. Native arithmetic follows its...
constexpr int target
First matching requirement, with every later choice checked for shadowing.
Definition isa.h:396
static constexpr simd load_bits(std::uint32_t const *p) noexcept
Load exact binary32 representations from uint32_t words.
Omitted architecture arguments use the native.simd provider's baseline.