Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
TensorBroadcasting.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2014 Benoit Steiner <benoit.steiner.goog@gmail.com>
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11#ifndef EIGEN_TENSOR_TENSOR_BROADCASTING_H
12#define EIGEN_TENSOR_TENSOR_BROADCASTING_H
13
14// IWYU pragma: private
15#include "./InternalHeaderCheck.h"
16
17namespace Eigen {
18
19namespace internal {
20template <typename Broadcast, typename XprType>
21struct traits<TensorBroadcastingOp<Broadcast, XprType>> : public traits<XprType> {
22 typedef typename XprType::Scalar Scalar;
23 typedef traits<XprType> XprTraits;
24 typedef typename XprTraits::StorageKind StorageKind;
25 typedef typename XprTraits::Index Index;
26 static constexpr int NumDimensions = XprTraits::NumDimensions;
27 static constexpr int Layout = XprTraits::Layout;
28 typedef typename XprTraits::PointerType PointerType;
29 enum {
30 // Broadcast is read-only.
31 Flags = traits<XprType>::Flags & ~LvalueBit
32 };
33};
34
35template <typename Broadcast, typename XprType>
36struct eval<TensorBroadcastingOp<Broadcast, XprType>, Eigen::Dense> {
37 typedef const TensorBroadcastingOp<Broadcast, XprType> EIGEN_DEVICE_REF type;
38};
39
40template <typename Dims>
41struct is_input_scalar : std::false_type {};
42template <>
43struct is_input_scalar<Sizes<>> : std::true_type {};
44template <std::ptrdiff_t... Indices>
45struct is_input_scalar<Sizes<Indices...>> : bool_constant<Sizes<Indices...>::total_size == 1> {};
46
47} // end namespace internal
48
52template <typename Broadcast, typename XprType>
53class TensorBroadcastingOp : public TensorBase<TensorBroadcastingOp<Broadcast, XprType>, ReadOnlyAccessors> {
54 public:
55 typedef typename Eigen::internal::traits<TensorBroadcastingOp>::Scalar Scalar;
56 typedef typename Eigen::NumTraits<Scalar>::Real RealScalar;
57 typedef typename XprType::CoeffReturnType CoeffReturnType;
58 typedef typename Eigen::internal::ref_selector<TensorBroadcastingOp>::type Nested;
59 typedef typename Eigen::internal::traits<TensorBroadcastingOp>::StorageKind StorageKind;
60 typedef typename Eigen::internal::traits<TensorBroadcastingOp>::Index Index;
61
62 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorBroadcastingOp(const XprType& expr, const Broadcast& broadcast)
63 : m_xpr(expr), m_broadcast(broadcast) {}
64
65 EIGEN_DEVICE_FUNC const Broadcast& broadcast() const { return m_broadcast; }
66
67 EIGEN_DEVICE_FUNC const internal::remove_all_t<typename XprType::Nested>& expression() const { return m_xpr; }
68
69 protected:
70 typename XprType::Nested m_xpr;
71 const Broadcast m_broadcast;
72};
73
74// Eval as rvalue
75template <typename Broadcast, typename ArgType, typename Device>
76struct TensorEvaluator<const TensorBroadcastingOp<Broadcast, ArgType>, Device> {
78 typedef typename XprType::Index Index;
79 static constexpr int NumDims = internal::array_size<typename TensorEvaluator<ArgType, Device>::Dimensions>::value;
80 typedef DSizes<Index, NumDims> Dimensions;
81 typedef typename XprType::Scalar Scalar;
82 typedef typename TensorEvaluator<ArgType, Device>::Dimensions InputDimensions;
83 typedef typename XprType::CoeffReturnType CoeffReturnType;
84 typedef typename PacketType<CoeffReturnType, Device>::type PacketReturnType;
85 static constexpr int PacketSize = PacketType<CoeffReturnType, Device>::size;
86
87 protected: // all the non-static fields must have the same access control, otherwise the TensorEvaluator won't be
88 // standard layout;
89 bool isCopy, nByOne, oneByN;
90
91 public:
92 typedef StorageMemory<CoeffReturnType, Device> Storage;
93 typedef typename Storage::Type EvaluatorPointerType;
94
95 enum {
96 IsAligned = TensorEvaluator<ArgType, Device>::IsAligned,
97 PacketAccess = TensorEvaluator<ArgType, Device>::PacketAccess,
99 PreferBlockAccess = true,
100 RawAccess = false
101 };
102 static constexpr int Layout = TensorEvaluator<ArgType, Device>::Layout;
103
104 typedef std::remove_const_t<Scalar> ScalarNoConst;
105
106 // We do block based broadcasting using a trick with 2x tensor rank and 0
107 // strides. See block method implementation for details.
108 typedef DSizes<Index, 2 * NumDims> BroadcastDimensions;
109
110 //===- Tensor block evaluation strategy (see TensorBlock.h) -------------===//
111 typedef internal::TensorBlockDescriptor<NumDims, Index> TensorBlockDesc;
112 typedef internal::TensorBlockScratchAllocator<Device> TensorBlockScratch;
113
114 typedef typename TensorEvaluator<const ArgType, Device>::TensorBlock ArgTensorBlock;
115
116 typedef typename internal::TensorMaterializedBlock<ScalarNoConst, NumDims, Layout, Index> TensorBlock;
117 //===--------------------------------------------------------------------===//
118
119 EIGEN_STRONG_INLINE TensorEvaluator(const XprType& op, const Device& device)
120 : isCopy(false),
121 nByOne(false),
122 oneByN(false),
123 m_device(device),
124 m_broadcast(op.broadcast()),
125 m_impl(op.expression(), device) {
126 // The broadcasting op doesn't change the rank of the tensor. One can't broadcast a scalar
127 // and store the result in a scalar. Instead one should reshape the scalar into a N-D
128 // tensor with N >= 1 of 1 element first and then broadcast.
129 EIGEN_STATIC_ASSERT((NumDims > 0), YOU_MADE_A_PROGRAMMING_MISTAKE);
130 const InputDimensions& input_dims = m_impl.dimensions();
131 isCopy = true;
132 for (int i = 0; i < NumDims; ++i) {
133 eigen_assert(input_dims[i] > 0);
134 m_dimensions[i] = input_dims[i] * m_broadcast[i];
135 if (m_broadcast[i] != 1) {
136 isCopy = false;
137 }
138 }
139
140 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
141 m_inputStrides[0] = 1;
142 m_outputStrides[0] = 1;
143 for (int i = 1; i < NumDims; ++i) {
144 m_inputStrides[i] = m_inputStrides[i - 1] * input_dims[i - 1];
145 m_outputStrides[i] = m_outputStrides[i - 1] * m_dimensions[i - 1];
146 }
147 } else {
148 m_inputStrides[NumDims - 1] = 1;
149 m_outputStrides[NumDims - 1] = 1;
150 for (int i = NumDims - 2; i >= 0; --i) {
151 m_inputStrides[i] = m_inputStrides[i + 1] * input_dims[i + 1];
152 m_outputStrides[i] = m_outputStrides[i + 1] * m_dimensions[i + 1];
153 }
154 }
155
156 if (input_dims[0] == 1) {
157 oneByN = true;
158 for (int i = 1; i < NumDims; ++i) {
159 if (m_broadcast[i] != 1) {
160 oneByN = false;
161 break;
162 }
163 }
164 } else if (input_dims[NumDims - 1] == 1) {
165 nByOne = true;
166 for (int i = 0; i < NumDims - 1; ++i) {
167 if (m_broadcast[i] != 1) {
168 nByOne = false;
169 break;
170 }
171 }
172 }
173
174 // Handle special format like NCHW, its input shape is '[1, N..., 1]' and
175 // broadcast shape is '[N, 1..., N]'
176 if (!oneByN && !nByOne) {
177 if (input_dims[0] == 1 && input_dims[NumDims - 1] == 1 && NumDims > 2) {
178 nByOne = true;
179 oneByN = true;
180 for (int i = 1; i < NumDims - 1; ++i) {
181 if (m_broadcast[i] != 1) {
182 nByOne = false;
183 oneByN = false;
184 break;
185 }
186 }
187 }
188 }
189 }
190
191 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE const Dimensions& dimensions() const { return m_dimensions; }
192
193 EIGEN_STRONG_INLINE bool evalSubExprsIfNeeded(EvaluatorPointerType) {
194 m_impl.evalSubExprsIfNeeded(nullptr);
195 return true;
196 }
197
198#ifdef EIGEN_USE_THREADS
199 template <typename EvalSubExprsCallback>
200 EIGEN_STRONG_INLINE void evalSubExprsIfNeededAsync(EvaluatorPointerType, EvalSubExprsCallback done) {
201 m_impl.evalSubExprsIfNeededAsync(nullptr, [done](bool) { done(true); });
202 }
203#endif // EIGEN_USE_THREADS
204
205 EIGEN_STRONG_INLINE void cleanup() { m_impl.cleanup(); }
206
207 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE CoeffReturnType coeff(Index index) const {
208 EIGEN_IF_CONSTEXPR ((internal::is_input_scalar<internal::remove_all_t<InputDimensions>>::value)) {
209 return m_impl.coeff(0);
210 }
211
212 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
213 if (isCopy) {
214 return m_impl.coeff(index);
215 } else {
216 return coeffColMajor(index);
217 }
218 } else {
219 if (isCopy) {
220 return m_impl.coeff(index);
221 } else {
222 return coeffRowMajor(index);
223 }
224 }
225 }
226
227 // The per-dim integer div/mod look expensive but amortize well: packet paths
228 // call this once per PacketSize outputs, and modern x86 hardware div is
229 // ~20 cycles. Prototyped replacing div/mod with TensorIntDivisor (the
230 // pattern used in TensorShuffling et al.); measured net-negative on Intel
231 // Raptor Lake across bench_broadcasting (more shapes regress 3-5% than
232 // improve). Left as hardware div.
233 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index indexColMajor(Index index) const {
234 Index inputIndex = 0;
235 EIGEN_UNROLL_LOOP
236 for (int i = NumDims - 1; i > 0; --i) {
237 const Index idx = index / m_outputStrides[i];
238 if (internal::index_statically_eq<Broadcast>(i, 1)) {
239 eigen_assert(idx < m_impl.dimensions()[i]);
240 inputIndex += idx * m_inputStrides[i];
241 } else {
242 if (internal::index_statically_eq<InputDimensions>(i, 1)) {
243 eigen_assert(idx % m_impl.dimensions()[i] == 0);
244 } else {
245 inputIndex += (idx % m_impl.dimensions()[i]) * m_inputStrides[i];
246 }
247 }
248 index -= idx * m_outputStrides[i];
249 }
250 EIGEN_IF_CONSTEXPR (internal::index_statically_eq<Broadcast>(0, 1)) {
251 eigen_assert(index < m_impl.dimensions()[0]);
252 inputIndex += index;
253 } else {
254 EIGEN_IF_CONSTEXPR (internal::index_statically_eq<InputDimensions>(0, 1)) {
255 eigen_assert(index % m_impl.dimensions()[0] == 0);
256 } else {
257 inputIndex += (index % m_impl.dimensions()[0]);
258 }
259 }
260 return inputIndex;
261 }
262
263 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeffColMajor(Index index) const {
264 return m_impl.coeff(indexColMajor(index));
265 }
266
267 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index indexRowMajor(Index index) const {
268 Index inputIndex = 0;
269 EIGEN_UNROLL_LOOP
270 for (int i = 0; i < NumDims - 1; ++i) {
271 const Index idx = index / m_outputStrides[i];
272 if (internal::index_statically_eq<Broadcast>(i, 1)) {
273 eigen_assert(idx < m_impl.dimensions()[i]);
274 inputIndex += idx * m_inputStrides[i];
275 } else {
276 if (internal::index_statically_eq<InputDimensions>(i, 1)) {
277 eigen_assert(idx % m_impl.dimensions()[i] == 0);
278 } else {
279 inputIndex += (idx % m_impl.dimensions()[i]) * m_inputStrides[i];
280 }
281 }
282 index -= idx * m_outputStrides[i];
283 }
284 EIGEN_IF_CONSTEXPR (internal::index_statically_eq<Broadcast>(NumDims - 1, 1)) {
285 eigen_assert(index < m_impl.dimensions()[NumDims - 1]);
286 inputIndex += index;
287 } else {
288 EIGEN_IF_CONSTEXPR (internal::index_statically_eq<InputDimensions>(NumDims - 1, 1)) {
289 eigen_assert(index % m_impl.dimensions()[NumDims - 1] == 0);
290 } else {
291 inputIndex += (index % m_impl.dimensions()[NumDims - 1]);
292 }
293 }
294 return inputIndex;
295 }
296
297 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE CoeffReturnType coeffRowMajor(Index index) const {
298 return m_impl.coeff(indexRowMajor(index));
299 }
300
301 template <int LoadMode>
302 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE PacketReturnType packet(Index index) const {
303 EIGEN_IF_CONSTEXPR ((internal::is_input_scalar<internal::remove_all_t<InputDimensions>>::value)) {
304 return internal::pset1<PacketReturnType>(m_impl.coeff(0));
305 }
306
307 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
308 if (isCopy) {
309#ifdef EIGEN_GPU_COMPILE_PHASE
310 // See PR 437: on NVIDIA P100 and K20m we observed a x3-4 speed up by enforcing
311 // unaligned loads here. The reason is unclear though.
312 return m_impl.template packet<Unaligned>(index);
313#else
314 return m_impl.template packet<LoadMode>(index);
315#endif
316 } else if (oneByN && !nByOne) {
317 return packetNByOne<LoadMode>(index);
318 } else if (!oneByN && nByOne) {
319 return packetOneByN<LoadMode>(index);
320 } else if (oneByN && nByOne) {
321 return packetOneByNByOne<LoadMode>(index);
322 } else {
323 return packetColMajor<LoadMode>(index);
324 }
325 } else {
326 if (isCopy) {
327#ifdef EIGEN_GPU_COMPILE_PHASE
328 // See above.
329 return m_impl.template packet<Unaligned>(index);
330#else
331 return m_impl.template packet<LoadMode>(index);
332#endif
333 } else if (oneByN && !nByOne) {
334 return packetOneByN<LoadMode>(index);
335 } else if (!oneByN && nByOne) {
336 return packetNByOne<LoadMode>(index);
337 } else if (oneByN && nByOne) {
338 return packetOneByNByOne<LoadMode>(index);
339 } else {
340 return packetRowMajor<LoadMode>(index);
341 }
342 }
343 }
344
345 template <int LoadMode>
346 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketReturnType packetOneByNByOne(Index index) const {
347 eigen_assert(index + PacketSize - 1 < dimensions().TotalSize());
348
349 EIGEN_ALIGN_TO_BOUNDARY(internal::unpacket_traits<PacketReturnType>::alignment)
350 std::remove_const_t<CoeffReturnType> values[PacketSize];
351 Index startDim, endDim;
352 Index inputIndex, outputOffset, batchedIndex;
353
354 EIGEN_IF_CONSTEXPR (static_cast<int>(Layout) == static_cast<int>(ColMajor)) {
355 startDim = NumDims - 1;
356 endDim = 1;
357 } else {
358 startDim = 0;
359 endDim = NumDims - 2;
360 }
361
362 batchedIndex = index % m_outputStrides[startDim];
363 inputIndex = batchedIndex / m_outputStrides[endDim];
364 outputOffset = batchedIndex % m_outputStrides[endDim];
365
366 if (outputOffset + PacketSize <= m_outputStrides[endDim]) {
367 values[0] = m_impl.coeff(inputIndex);
368 return internal::pload1<PacketReturnType>(values);
369 } else {
370 EIGEN_UNROLL_LOOP
371 for (int i = 0, cur = 0; i < PacketSize; ++i, ++cur) {
372 if (outputOffset + cur < m_outputStrides[endDim]) {
373 values[i] = m_impl.coeff(inputIndex);
374 } else {
375 ++inputIndex;
376 inputIndex = (inputIndex == m_inputStrides[startDim] ? 0 : inputIndex);
377 values[i] = m_impl.coeff(inputIndex);
378 outputOffset = 0;
379 cur = 0;
380 }
381 }
382 return internal::pload<PacketReturnType>(values);
383 }
384 }
385
386 template <int LoadMode>
387 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketReturnType packetOneByN(Index index) const {
388 // Consider the flattened tensor [v0, ..., vN],
389 // Concatenates m_broadcast[dim] copies,
390 // [v0, ..., vN, v0, ..., vN, ... ]
391 // with dim == NumDims - 1 for col-major, dim == 0 for row-major.
392 eigen_assert(index + PacketSize - 1 < dimensions().TotalSize());
393
394 // Size of flattened tensor.
395 const Index M =
396 (static_cast<int>(Layout) == static_cast<int>(ColMajor)) ? m_inputStrides[NumDims - 1] : m_inputStrides[0];
397 Index inputIndex = index % M;
398 if (inputIndex + PacketSize <= M) {
399 return m_impl.template packet<Unaligned>(inputIndex);
400 } else {
401 EIGEN_ALIGN_TO_BOUNDARY(internal::unpacket_traits<PacketReturnType>::alignment)
402 std::remove_const_t<CoeffReturnType> values[PacketSize];
403 EIGEN_UNROLL_LOOP
404 for (int i = 0; i < PacketSize; ++i) {
405 if (inputIndex > M - 1) {
406 inputIndex = 0;
407 }
408 values[i] = m_impl.coeff(inputIndex++);
409 }
410 return internal::pload<PacketReturnType>(values);
411 }
412 }
413
414 template <int LoadMode>
415 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketReturnType packetNByOne(Index index) const {
416 // Consider the flattened tensor [v0, ..., vN],
417 // Interleaves m_broadcast[dim] copies,
418 // [v0, v0, ..., v1, v1, ..., vN, vN, ... ]
419 // with dim == 0 for col-major, dim == NumDims - 1 for row-major.
420 eigen_assert(index + PacketSize - 1 < dimensions().TotalSize());
421
422 const Index M =
423 (static_cast<int>(Layout) == static_cast<int>(ColMajor)) ? m_broadcast[0] : m_broadcast[NumDims - 1];
424
425 Index inputIndex = index / M;
426 Index outputOffset = index % M;
427 if (outputOffset + PacketSize <= M) {
428 return internal::pset1<PacketReturnType>(m_impl.coeff(inputIndex));
429 } else {
430 EIGEN_ALIGN_TO_BOUNDARY(internal::unpacket_traits<PacketReturnType>::alignment)
431 std::remove_const_t<CoeffReturnType> values[PacketSize];
432 EIGEN_UNROLL_LOOP
433 for (int i = 0; i < PacketSize; ++i) {
434 if (outputOffset < M) {
435 values[i] = m_impl.coeff(inputIndex);
436 ++outputOffset;
437 } else {
438 values[i] = m_impl.coeff(++inputIndex);
439 outputOffset = 1; // Next offset.
440 }
441 }
442 return internal::pload<PacketReturnType>(values);
443 }
444 }
445
446 // Ignore the LoadMode and always use unaligned loads since we can't guarantee
447 // the alignment at compile time.
448 template <int LoadMode>
449 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketReturnType packetColMajor(Index index) const {
450 eigen_assert(index + PacketSize - 1 < dimensions().TotalSize());
451
452 const Index originalIndex = index;
453
454 Index inputIndex = 0;
455 EIGEN_UNROLL_LOOP
456 for (int i = NumDims - 1; i > 0; --i) {
457 const Index idx = index / m_outputStrides[i];
458 if (internal::index_statically_eq<Broadcast>(i, 1)) {
459 eigen_assert(idx < m_impl.dimensions()[i]);
460 inputIndex += idx * m_inputStrides[i];
461 } else {
462 if (internal::index_statically_eq<InputDimensions>(i, 1)) {
463 eigen_assert(idx % m_impl.dimensions()[i] == 0);
464 } else {
465 inputIndex += (idx % m_impl.dimensions()[i]) * m_inputStrides[i];
466 }
467 }
468 index -= idx * m_outputStrides[i];
469 }
470 Index innermostLoc;
471 EIGEN_IF_CONSTEXPR (internal::index_statically_eq<Broadcast>(0, 1)) {
472 eigen_assert(index < m_impl.dimensions()[0]);
473 innermostLoc = index;
474 } else {
475 EIGEN_IF_CONSTEXPR (internal::index_statically_eq<InputDimensions>(0, 1)) {
476 eigen_assert(index % m_impl.dimensions()[0] == 0);
477 innermostLoc = 0;
478 } else {
479 innermostLoc = index % m_impl.dimensions()[0];
480 }
481 }
482 inputIndex += innermostLoc;
483
484 // TODO: This could be extended to the second dimension if we're not
485 // broadcasting alongside the first dimension, and so on.
486 if (innermostLoc + PacketSize <= m_impl.dimensions()[0]) {
487 return m_impl.template packet<Unaligned>(inputIndex);
488 } else {
489 EIGEN_ALIGN_TO_BOUNDARY(internal::unpacket_traits<PacketReturnType>::alignment)
490 std::remove_const_t<CoeffReturnType> values[PacketSize];
491 values[0] = m_impl.coeff(inputIndex);
492 EIGEN_UNROLL_LOOP
493 for (int i = 1; i < PacketSize; ++i) {
494 if (innermostLoc + i < m_impl.dimensions()[0]) {
495 values[i] = m_impl.coeff(inputIndex + i);
496 } else {
497 values[i] = coeffColMajor(originalIndex + i);
498 }
499 }
500 PacketReturnType rslt = internal::pload<PacketReturnType>(values);
501 return rslt;
502 }
503 }
504
505 template <int LoadMode>
506 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE PacketReturnType packetRowMajor(Index index) const {
507 eigen_assert(index + PacketSize - 1 < dimensions().TotalSize());
508
509 const Index originalIndex = index;
510
511 Index inputIndex = 0;
512 EIGEN_UNROLL_LOOP
513 for (int i = 0; i < NumDims - 1; ++i) {
514 const Index idx = index / m_outputStrides[i];
515 if (internal::index_statically_eq<Broadcast>(i, 1)) {
516 eigen_assert(idx < m_impl.dimensions()[i]);
517 inputIndex += idx * m_inputStrides[i];
518 } else {
519 if (internal::index_statically_eq<InputDimensions>(i, 1)) {
520 eigen_assert(idx % m_impl.dimensions()[i] == 0);
521 } else {
522 inputIndex += (idx % m_impl.dimensions()[i]) * m_inputStrides[i];
523 }
524 }
525 index -= idx * m_outputStrides[i];
526 }
527 Index innermostLoc;
528 EIGEN_IF_CONSTEXPR (internal::index_statically_eq<Broadcast>(NumDims - 1, 1)) {
529 eigen_assert(index < m_impl.dimensions()[NumDims - 1]);
530 innermostLoc = index;
531 } else {
532 EIGEN_IF_CONSTEXPR (internal::index_statically_eq<InputDimensions>(NumDims - 1, 1)) {
533 eigen_assert(index % m_impl.dimensions()[NumDims - 1] == 0);
534 innermostLoc = 0;
535 } else {
536 innermostLoc = index % m_impl.dimensions()[NumDims - 1];
537 }
538 }
539 inputIndex += innermostLoc;
540
541 // TODO: This could be extended to the second dimension if we're not
542 // broadcasting alongside the first dimension, and so on.
543 if (innermostLoc + PacketSize <= m_impl.dimensions()[NumDims - 1]) {
544 return m_impl.template packet<Unaligned>(inputIndex);
545 } else {
546 EIGEN_ALIGN_TO_BOUNDARY(internal::unpacket_traits<PacketReturnType>::alignment)
547 std::remove_const_t<CoeffReturnType> values[PacketSize];
548 values[0] = m_impl.coeff(inputIndex);
549 EIGEN_UNROLL_LOOP
550 for (int i = 1; i < PacketSize; ++i) {
551 if (innermostLoc + i < m_impl.dimensions()[NumDims - 1]) {
552 values[i] = m_impl.coeff(inputIndex + i);
553 } else {
554 values[i] = coeffRowMajor(originalIndex + i);
555 }
556 }
557 PacketReturnType rslt = internal::pload<PacketReturnType>(values);
558 return rslt;
559 }
560 }
561
562 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorOpCost costPerCoeff(bool vectorized) const {
563 double compute_cost = TensorOpCost::AddCost<Index>();
564 EIGEN_IF_CONSTEXPR (NumDims > 0) {
565 if (!isCopy) {
566 EIGEN_UNROLL_LOOP
567 for (int i = NumDims - 1; i > 0; --i) {
568 compute_cost += TensorOpCost::DivCost<Index>();
569 if (internal::index_statically_eq<Broadcast>(i, 1)) {
570 compute_cost += TensorOpCost::MulCost<Index>() + TensorOpCost::AddCost<Index>();
571 } else {
572 if (!internal::index_statically_eq<InputDimensions>(i, 1)) {
573 compute_cost +=
574 TensorOpCost::MulCost<Index>() + TensorOpCost::ModCost<Index>() + TensorOpCost::AddCost<Index>();
575 }
576 }
577 compute_cost += TensorOpCost::MulCost<Index>() + TensorOpCost::AddCost<Index>();
578 }
579 }
580 }
581 return m_impl.costPerCoeff(vectorized) + TensorOpCost(0, 0, compute_cost, vectorized, PacketSize);
582 }
583
584 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE internal::TensorBlockResourceRequirements getResourceRequirements() const {
585 // Target the L1 cache: measured ~30% faster than targeting the last-level
586 // cache on large tensors.
587 const size_t target_size = m_device.firstLevelCacheSize();
588 return internal::TensorBlockResourceRequirements::merge(
589 m_impl.getResourceRequirements(), internal::TensorBlockResourceRequirements::skewed<Scalar>(target_size));
590 }
591
592 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorBlock block(TensorBlockDesc& desc, TensorBlockScratch& scratch,
593 bool /*root_of_expr_ast*/ = false) const {
594 BlockBroadcastingParams params = blockBroadcastingParams(desc);
595
596 if (params.inner_dim_size == 0 || params.bcast_dim_size == 0) {
597 return emptyBlock();
598 }
599
600 // Prepare storage for the materialized broadcasting result.
601 const typename TensorBlock::Storage block_storage = TensorBlock::prepareStorage(desc, scratch);
602 ScalarNoConst* materialized_output = block_storage.data();
603
604 // We potentially will need to materialize input blocks.
605 size_t materialized_input_size = 0;
606 ScalarNoConst* materialized_input = nullptr;
607
608 // Initialize block broadcasting iterator state for outer dimensions (outer
609 // with regard to bcast dimension). Dimensions in this array are always in
610 // inner_most -> outer_most order (col major layout).
611 array<BlockBroadcastingIteratorState, NumDims> it;
612 int idx = 0;
613
614 for (int i = params.inner_dim_count + 1; i < NumDims; ++i) {
615 const Index dim = IsColMajor ? i : NumDims - 1 - i;
616 it[idx].size = params.output_dims[dim];
617 it[idx].count = 0;
618 it[idx].output_stride = m_outputStrides[dim];
619 it[idx].output_span = it[idx].output_stride * (it[idx].size - 1);
620 idx++;
621 }
622
623 // Write output into the beginning of `materialized_output`.
624 Index output_offset = 0;
625
626 // We will fill output block by broadcasting along the bcast dim, and
627 // iterating over outer dimension.
628 const Index output_size = NumDims == 0 ? 1 : params.output_dims.TotalSize();
629
630 for (Index num_output_coeffs = 0; num_output_coeffs < output_size;) {
631 ScalarNoConst* bcast_output = materialized_output + num_output_coeffs;
632 Index bcast_offset = desc.offset() + output_offset;
633
634 // Broadcast along the bcast dimension.
635 num_output_coeffs += BroadcastBlockAlongBcastDim(params, bcast_offset, scratch, bcast_output, &materialized_input,
636 &materialized_input_size);
637
638 // Switch to the next outer dimension.
639 for (int j = 0; j < idx; ++j) {
640 if (++it[j].count < it[j].size) {
641 output_offset += it[j].output_stride;
642 break;
643 }
644 it[j].count = 0;
645 output_offset -= it[j].output_span;
646 }
647 }
648
649 return block_storage.AsTensorMaterializedBlock();
650 }
651
652 EIGEN_DEVICE_FUNC EvaluatorPointerType data() const { return nullptr; }
653
654 const TensorEvaluator<ArgType, Device>& impl() const { return m_impl; }
655
656 Broadcast functor() const { return m_broadcast; }
657
658 private:
659 static constexpr bool IsColMajor = static_cast<int>(Layout) == static_cast<int>(ColMajor);
660
661 // We will build a general case block broadcasting on top of broadcasting
662 // primitive that will do broadcasting only for the inner dimension(s) along
663 // the first dimension smaller than the input size (it's called `bcast_dim`).
664 //
665 // Example:
666 // dim: 0 1 2 (ColMajor)
667 // input size: [9, 3, 6]
668 // block size: [9, 2, 6]
669 //
670 // We will compute broadcasted block by iterating over the outer dimensions
671 // before `bcast_dim` (only dimension `2` in this example) and computing
672 // broadcasts along the `bcast_dim` (dimension `1` in this example).
673
674 // BlockBroadcastingParams holds precomputed parameters for broadcasting a
675 // single block along the broadcasting dimension. Sizes and strides along the
676 // `bcast_dim` might be invalid, they will be adjusted later in
677 // `BroadcastBlockAlongBcastDim`.
678 struct BlockBroadcastingParams {
679 Dimensions input_dims; // input expression dimensions
680 Dimensions output_dims; // output block sizes
681 Dimensions output_strides; // output block strides
682
683 int inner_dim_count; // count inner dimensions matching in size
684 int bcast_dim; // broadcasting dimension index
685 Index bcast_dim_size; // broadcasting dimension size
686 Index inner_dim_size; // inner dimensions size
687
688 // Block sizes and strides for the input block where all dimensions before
689 // `bcast_dim` are equal to `1`.
690 Dimensions input_block_sizes;
691 Dimensions input_block_strides;
692
693 // Block sizes and strides for blocks with extra dimensions and strides `0`.
694 BroadcastDimensions bcast_block_sizes;
695 BroadcastDimensions bcast_block_strides;
696 BroadcastDimensions bcast_input_strides;
697 };
698
699 struct BlockBroadcastingIteratorState {
700 Index size;
701 Index count;
702 Index output_stride;
703 Index output_span;
704 };
705
706 EIGEN_DEVICE_FUNC EIGEN_ALWAYS_INLINE BlockBroadcastingParams blockBroadcastingParams(TensorBlockDesc& desc) const {
707 BlockBroadcastingParams params;
708
709 params.input_dims = Dimensions(m_impl.dimensions());
710
711 // Output block sizes and strides.
712 params.output_dims = desc.dimensions();
713 params.output_strides = internal::strides<Layout>(params.output_dims);
714
715 // Find the broadcasting dimension (first dimension with output size smaller
716 // than the input size).
717 params.bcast_dim = 0;
718 params.bcast_dim_size = 1;
719 params.inner_dim_size = 1;
720
721 // Count the number of inner dimensions that have the same size in the block
722 // and in the broadcast expression.
723 params.inner_dim_count = 0;
724
725 for (int i = 0; i < NumDims; ++i) {
726 const int dim = IsColMajor ? i : NumDims - i - 1;
727
728 if (params.output_dims[dim] == m_dimensions[dim]) {
729 params.inner_dim_size *= params.output_dims[dim];
730 ++params.inner_dim_count;
731 continue;
732 }
733
734 // First non-matching dimension is the broadcasting dimension.
735 eigen_assert(params.output_dims[dim] < m_dimensions[dim]);
736 params.bcast_dim = dim;
737 params.bcast_dim_size = params.output_dims[dim];
738 break;
739 }
740
741 // Calculate the input block size for looking into the input.
742 for (int i = 0; i < params.inner_dim_count; ++i) {
743 const int dim = IsColMajor ? i : NumDims - i - 1;
744 params.input_block_sizes[dim] = params.input_dims[dim];
745 }
746 for (int i = params.inner_dim_count; i < NumDims; ++i) {
747 const int dim = IsColMajor ? i : NumDims - i - 1;
748 params.input_block_sizes[dim] = 1;
749 }
750 params.input_block_strides = internal::strides<Layout>(params.input_block_sizes);
751
752 // Broadcast with the 0-stride trick: Create 1 extra dim for each
753 // broadcast, set the input stride to 0.
754 //
755 // When ColMajor:
756 //
757 // - bcast_block_sizes:
758 // [d_0, b_0, d_1, b_1, ...]
759 //
760 // - bcast_block_strides:
761 // [output_block_strides[0], output_block_strides[0] * d_0,
762 // output_block_strides[1], output_block_strides[1] * d_1,
763 // ...]
764 //
765 // - bcast_input_strides:
766 // [input_block_strides[0], 0,
767 // input_block_strides[1], 0,
768 // ...].
769 //
770 for (int i = 0; i < params.inner_dim_count; ++i) {
771 const int dim = IsColMajor ? i : NumDims - i - 1;
772
773 const int copy_dim = IsColMajor ? 2 * i : 2 * NumDims - 2 * i - 1;
774 const int broadcast_dim = IsColMajor ? copy_dim + 1 : copy_dim - 1;
775
776 params.bcast_block_sizes[copy_dim] = params.input_dims[dim];
777 params.bcast_block_sizes[broadcast_dim] = m_broadcast[dim];
778 params.bcast_block_strides[copy_dim] = params.output_strides[dim];
779 params.bcast_block_strides[broadcast_dim] = params.output_strides[dim] * params.input_dims[dim];
780 params.bcast_input_strides[copy_dim] = params.input_block_strides[dim];
781 params.bcast_input_strides[broadcast_dim] = 0;
782 }
783
784 for (int i = 2 * params.inner_dim_count; i < 2 * NumDims; ++i) {
785 const int dim = IsColMajor ? i : 2 * NumDims - i - 1;
786 params.bcast_block_sizes[dim] = 1;
787 params.bcast_block_strides[dim] = 0;
788 params.bcast_input_strides[dim] = 0;
789 }
790
791 return params;
792 }
793
794 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE TensorBlock emptyBlock() const {
795 DSizes<Index, NumDims> dimensions;
796 for (int i = 0; i < NumDims; ++i) dimensions[i] = 0;
797 return TensorBlock(internal::TensorBlockKind::kView, nullptr, dimensions);
798 }
799
800 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index BroadcastBlockAlongBcastDim(
801 BlockBroadcastingParams params, Index bcast_offset, TensorBlockScratch& scratch,
802 ScalarNoConst* materialized_output, ScalarNoConst** materialized_input, size_t* materialized_input_size) const {
803 if (params.bcast_dim_size == 1) {
804 // We just need one block read using the ready-set values above.
805 return BroadcastBlock(params.input_block_sizes, params.input_block_strides, params.bcast_block_sizes,
806 params.bcast_block_strides, params.bcast_input_strides, bcast_offset, 0, scratch,
807 materialized_output, materialized_input, materialized_input_size);
808
809 } else if (params.input_dims[params.bcast_dim] == 1) {
810 // Broadcast bcast dimension (< NumDims) by bcast_dim_size.
811 const int broadcast_bcast_dim =
812 IsColMajor ? 2 * params.inner_dim_count + 1 : 2 * NumDims - 2 * params.inner_dim_count - 2;
813
814 params.bcast_block_sizes[broadcast_bcast_dim] = params.bcast_dim_size;
815 params.bcast_input_strides[broadcast_bcast_dim] = 0;
816 params.bcast_block_strides[broadcast_bcast_dim] = params.output_strides[params.bcast_dim];
817
818 return BroadcastBlock(params.input_block_sizes, params.input_block_strides, params.bcast_block_sizes,
819 params.bcast_block_strides, params.bcast_input_strides, bcast_offset, 0, scratch,
820 materialized_output, materialized_input, materialized_input_size);
821
822 } else {
823 // Keep track of the total number of the coefficients written to the
824 // output block.
825 Index num_output_coeffs = 0;
826
827 // The general case. Let's denote the output block as
828 //
829 // x[..., a:a+bcast_dim_size, :, ..., :]
830 //
831 // where a:a+bcast_dim_size is a slice on the bcast_dim dimension
832 // (< NumDims). We need to split the a:a+bcast_dim_size into possibly 3
833 // sub-blocks:
834 //
835 // (1) a:b, where b is the smallest multiple of
836 // input_dims[bcast_dim_start] in [a, a+bcast_dim_size].
837 //
838 // (2) b:c, where c is the largest multiple of input_dims[bcast_dim_start]
839 // in [a, a+bcast_dim_size].
840 //
841 // (3) c:a+bcast_dim_size .
842 //
843 // Or, when b and c do not exist, we just need to process the whole block
844 // together.
845
846 // Find a.
847 const Index bcast_dim_left_index = bcast_offset / m_outputStrides[params.bcast_dim];
848
849 // Find b and c.
850 const Index input_bcast_dim_size = params.input_dims[params.bcast_dim];
851
852 // First multiple after a. This is b when <= bcast_dim_left_index +
853 // bcast_dim_size.
854 const Index first_multiple =
855 numext::div_ceil<Index>(bcast_dim_left_index, input_bcast_dim_size) * input_bcast_dim_size;
856
857 if (first_multiple <= bcast_dim_left_index + params.bcast_dim_size) {
858 // b exists, so does c. Find it.
859 const Index last_multiple =
860 (bcast_dim_left_index + params.bcast_dim_size) / input_bcast_dim_size * input_bcast_dim_size;
861 const int copy_bcast_dim =
862 IsColMajor ? 2 * params.inner_dim_count : 2 * NumDims - 2 * params.inner_dim_count - 1;
863 const int broadcast_bcast_dim =
864 IsColMajor ? 2 * params.inner_dim_count + 1 : 2 * NumDims - 2 * params.inner_dim_count - 2;
865
866 if (first_multiple > bcast_dim_left_index) {
867 const Index head_size = first_multiple - bcast_dim_left_index;
868 params.input_block_sizes[params.bcast_dim] = head_size;
869 params.bcast_block_sizes[copy_bcast_dim] = head_size;
870 params.bcast_input_strides[copy_bcast_dim] = params.input_block_strides[params.bcast_dim];
871 params.bcast_block_strides[copy_bcast_dim] = params.output_strides[params.bcast_dim];
872 params.bcast_block_sizes[broadcast_bcast_dim] = 1;
873 params.bcast_input_strides[broadcast_bcast_dim] = 0;
874 params.bcast_block_strides[broadcast_bcast_dim] =
875 params.output_strides[params.bcast_dim] * params.input_dims[params.bcast_dim];
876
877 num_output_coeffs +=
878 BroadcastBlock(params.input_block_sizes, params.input_block_strides, params.bcast_block_sizes,
879 params.bcast_block_strides, params.bcast_input_strides, bcast_offset, 0, scratch,
880 materialized_output, materialized_input, materialized_input_size);
881 }
882 if (first_multiple < last_multiple) {
883 params.input_block_sizes[params.bcast_dim] = input_bcast_dim_size;
884 params.bcast_block_sizes[copy_bcast_dim] = input_bcast_dim_size;
885 params.bcast_input_strides[copy_bcast_dim] = params.input_block_strides[params.bcast_dim];
886 params.bcast_block_strides[copy_bcast_dim] = params.output_strides[params.bcast_dim];
887 params.bcast_block_sizes[broadcast_bcast_dim] = (last_multiple - first_multiple) / input_bcast_dim_size;
888 params.bcast_input_strides[broadcast_bcast_dim] = 0;
889 params.bcast_block_strides[broadcast_bcast_dim] =
890 params.output_strides[params.bcast_dim] * params.input_dims[params.bcast_dim];
891 const Index offset = (first_multiple - bcast_dim_left_index) * m_outputStrides[params.bcast_dim];
892
893 num_output_coeffs +=
894 BroadcastBlock(params.input_block_sizes, params.input_block_strides, params.bcast_block_sizes,
895 params.bcast_block_strides, params.bcast_input_strides, bcast_offset, offset, scratch,
896 materialized_output, materialized_input, materialized_input_size);
897 }
898 if (last_multiple < bcast_dim_left_index + params.bcast_dim_size) {
899 const Index tail_size = bcast_dim_left_index + params.bcast_dim_size - last_multiple;
900 params.input_block_sizes[params.bcast_dim] = tail_size;
901 params.bcast_block_sizes[copy_bcast_dim] = tail_size;
902 params.bcast_input_strides[copy_bcast_dim] = params.input_block_strides[params.bcast_dim];
903 params.bcast_block_strides[copy_bcast_dim] = params.output_strides[params.bcast_dim];
904 params.bcast_block_sizes[broadcast_bcast_dim] = 1;
905 params.bcast_input_strides[broadcast_bcast_dim] = 0;
906 params.bcast_block_strides[broadcast_bcast_dim] =
907 params.output_strides[params.bcast_dim] * params.input_dims[params.bcast_dim];
908 const Index offset = (last_multiple - bcast_dim_left_index) * m_outputStrides[params.bcast_dim];
909
910 num_output_coeffs +=
911 BroadcastBlock(params.input_block_sizes, params.input_block_strides, params.bcast_block_sizes,
912 params.bcast_block_strides, params.bcast_input_strides, bcast_offset, offset, scratch,
913 materialized_output, materialized_input, materialized_input_size);
914 }
915 } else {
916 // b and c do not exist.
917 const int copy_bcast_dim =
918 IsColMajor ? 2 * params.inner_dim_count : 2 * NumDims - 2 * params.inner_dim_count - 1;
919 params.input_block_sizes[params.bcast_dim] = params.bcast_dim_size;
920 params.bcast_block_sizes[copy_bcast_dim] = params.bcast_dim_size;
921 params.bcast_input_strides[copy_bcast_dim] = params.input_block_strides[params.bcast_dim];
922 params.bcast_block_strides[copy_bcast_dim] = params.output_strides[params.bcast_dim];
923
924 num_output_coeffs +=
925 BroadcastBlock(params.input_block_sizes, params.input_block_strides, params.bcast_block_sizes,
926 params.bcast_block_strides, params.bcast_input_strides, bcast_offset, 0, scratch,
927 materialized_output, materialized_input, materialized_input_size);
928 }
929
930 return num_output_coeffs;
931 }
932 }
933
934 EIGEN_DEVICE_FUNC EIGEN_STRONG_INLINE Index BroadcastBlock(
935 const Dimensions& input_block_sizes, const Dimensions& input_block_strides,
936 const BroadcastDimensions& bcast_block_sizes, const BroadcastDimensions& bcast_block_strides,
937 const BroadcastDimensions& bcast_input_strides, Index bcast_offset, Index offset, TensorBlockScratch& scratch,
938 ScalarNoConst* materialized_output, ScalarNoConst** materialized_input, size_t* materialized_input_size) const {
939 // ---------------------------------------------------------------------- //
940 // Tensor block descriptor for reading block from the input.
941 const Index input_offset = bcast_offset + offset;
942 TensorBlockDesc input_desc(IsColMajor ? indexColMajor(input_offset) : indexRowMajor(input_offset),
943 input_block_sizes);
944
945 ArgTensorBlock input_block = m_impl.block(input_desc, scratch);
946
947 // ---------------------------------------------------------------------- //
948 // Materialize input block into a temporary memory buffer only if it's not
949 // already available in the arg block.
950 const ScalarNoConst* input_buffer = nullptr;
951
952 if (input_block.data() != nullptr) {
953 // Input block already has raw data, there is no need to materialize it.
954 input_buffer = input_block.data();
955
956 } else {
957 // Otherwise we have to do block assignment into a temporary buffer.
958
959 // Maybe reuse previously allocated buffer, or allocate a new one with a
960 // scratch allocator.
961 const size_t input_total_size = input_block_sizes.TotalSize();
962 if (*materialized_input == nullptr || *materialized_input_size < input_total_size) {
963 *materialized_input_size = input_total_size;
964 void* mem = scratch.allocate(*materialized_input_size * sizeof(Scalar));
965 *materialized_input = static_cast<ScalarNoConst*>(mem);
966 }
967
968 typedef internal::TensorBlockAssignment<ScalarNoConst, NumDims, typename ArgTensorBlock::XprType, Index>
969 TensorBlockAssignment;
970
971 TensorBlockAssignment::Run(
972 TensorBlockAssignment::target(input_block_sizes, input_block_strides, *materialized_input),
973 input_block.expr());
974
975 input_buffer = *materialized_input;
976 }
977
978 // ---------------------------------------------------------------------- //
979 // Copy data from materialized input block to the materialized output, using
980 // given broadcast strides (strides with zeroes).
981 typedef internal::TensorBlockIO<ScalarNoConst, Index, 2 * NumDims, Layout> TensorBlockIO;
982
983 typename TensorBlockIO::Src src(bcast_input_strides, input_buffer);
984 typename TensorBlockIO::Dst dst(bcast_block_sizes, bcast_block_strides, materialized_output + offset);
985
986 return TensorBlockIO::Copy(dst, src);
987 }
988
989 protected:
990 const Device EIGEN_DEVICE_REF m_device;
991 const std::remove_reference_t<Broadcast> m_broadcast;
992 Dimensions m_dimensions;
993 array<Index, NumDims> m_outputStrides;
994 array<Index, NumDims> m_inputStrides;
995 TensorEvaluator<ArgType, Device> m_impl;
996};
997
998} // end namespace Eigen
999
1000#endif // EIGEN_TENSOR_TENSOR_BROADCASTING_H
The tensor base class.
Definition TensorForwardDeclarations.h:69
Definition TensorBroadcasting.h:53
Namespace containing all symbols from the Eigen library.
The tensor evaluator class.
Definition TensorEvaluator.h:47