MCPcopy Create free account
hub / github.com/DeepRec-AI/DeepRec / SelectAndScatter4DGePlus

Method SelectAndScatter4DGePlus

tensorflow/compiler/xla/reference_util.cc:322–396  ·  view source on GitHub ↗

static */

Source from the content-addressed store, hash-verified

320}
321
322/* static */ std::unique_ptr<Array4D<float>>
323ReferenceUtil::SelectAndScatter4DGePlus(const Array4D<float>& operand,
324 const Array4D<float>& source,
325 float init,
326 absl::Span<const int64> window,
327 absl::Span<const int64> stride,
328 bool same_padding) {
329 Padding padding = same_padding ? Padding::kSame : Padding::kValid;
330 auto result = absl::make_unique<Array4D<float>>(operand.n1(), operand.n2(),
331 operand.n3(), operand.n4());
332 std::vector<int64> dim_lengths{operand.n1(), operand.n2(), operand.n3(),
333 operand.n4()};
334 auto padding_both = xla::MakePadding(dim_lengths, window, stride, padding);
335 // Fill the output, with the initial value.
336 result->Fill(init);
337
338 std::vector<int64> window_counts(window.size(), 0);
339 std::vector<int64> pad_low(window.size(), 0);
340 for (int64 i = 0; i < window.size(); ++i) {
341 window_counts[i] =
342 WindowCount(dim_lengths[i], window[i], stride[i], padding);
343 pad_low[i] = padding_both[i].first;
344 }
345 CHECK_EQ(window_counts[0], source.n1());
346 CHECK_EQ(window_counts[1], source.n2());
347 CHECK_EQ(window_counts[2], source.n3());
348 CHECK_EQ(window_counts[3], source.n4());
349
350 // Do a full 4D select and Scatter.
351 for (int64 i0 = 0; i0 < window_counts[0]; ++i0) {
352 for (int64 i1 = 0; i1 < window_counts[1]; ++i1) {
353 for (int64 i2 = 0; i2 < window_counts[2]; ++i2) {
354 for (int64 i3 = 0; i3 < window_counts[3]; ++i3) {
355 // Now we are inside a window and need to find the max and the argmax.
356 int64 i0_base = i0 * stride[0] - pad_low[0];
357 int64 i1_base = i1 * stride[1] - pad_low[1];
358 int64 i2_base = i2 * stride[2] - pad_low[2];
359 int64 i3_base = i3 * stride[3] - pad_low[3];
360 int64 scatter_0 = (i0_base >= 0) ? i0_base : 0;
361 int64 scatter_1 = (i1_base >= 0) ? i1_base : 0;
362 int64 scatter_2 = (i2_base >= 0) ? i2_base : 0;
363 int64 scatter_3 = (i3_base >= 0) ? i3_base : 0;
364 float val = operand(scatter_0, scatter_1, scatter_2, scatter_3);
365 for (int64 i0_win = 0; i0_win < window[0]; ++i0_win) {
366 for (int64 i1_win = 0; i1_win < window[1]; ++i1_win) {
367 for (int64 i2_win = 0; i2_win < window[2]; ++i2_win) {
368 for (int64 i3_win = 0; i3_win < window[3]; ++i3_win) {
369 if (i0_base + i0_win >= 0 && i1_base + i1_win >= 0 &&
370 i2_base + i2_win >= 0 && i3_base + i3_win >= 0 &&
371 i0_base + i0_win < operand.n1() &&
372 i1_base + i1_win < operand.n2() &&
373 i2_base + i2_win < operand.n3() &&
374 i3_base + i3_win < operand.n4()) {
375 float tmp = operand(i0_base + i0_win, i1_base + i1_win,
376 i2_base + i2_win, i3_base + i3_win);
377 if (tmp > val) {
378 val = tmp;
379 scatter_0 = i0_base + i0_win;

Callers

nothing calls this directly

Calls 7

MakePaddingFunction · 0.85
n4Method · 0.80
n1Method · 0.45
n2Method · 0.45
n3Method · 0.45
FillMethod · 0.45
sizeMethod · 0.45

Tested by

no test coverage detected