static */
| 320 | } |
| 321 | |
| 322 | /* static */ std::unique_ptr<Array4D<float>> |
| 323 | ReferenceUtil::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; |