| 556 | } |
| 557 | |
| 558 | Maybe<Tensor> AsStridedGrad(const std::shared_ptr<one::Tensor>& dy, |
| 559 | const std::shared_ptr<one::Tensor>& input, |
| 560 | const std::vector<int64_t>& sizes, const std::vector<int64_t>& strides, |
| 561 | const int64_t storage_offset) { |
| 562 | CHECK_OR_RETURN(input->is_local()) << "input must be local tensor."; |
| 563 | // reference: torch/csrc/autograd/FunctionsManual.cpp |
| 564 | const size_t odim = dy->ndim(); |
| 565 | std::vector<int64_t> out_sizes_, out_strides_; |
| 566 | out_sizes_.reserve(odim); |
| 567 | out_strides_.reserve(odim); |
| 568 | auto grad = dy; |
| 569 | for (int64_t i = odim - 1; i >= 0; i--) { |
| 570 | auto size_i = sizes[i]; |
| 571 | auto stride_i = strides[i]; |
| 572 | if (size_i == 0) { |
| 573 | return functional::Constant(*dy->shape(), 0, grad->dtype(), JUST(grad->device())); |
| 574 | } else if (size_i == 1) { |
| 575 | grad = JUST(functional::Squeeze(grad, std::vector<int32_t>{int(i)})); |
| 576 | } else if (stride_i == 0) { |
| 577 | grad = JUST(functional::ReduceSum(grad, std::vector<int32_t>{int(i)}, false, NullOpt)); |
| 578 | } else { |
| 579 | out_sizes_.insert(out_sizes_.begin(), size_i); |
| 580 | out_strides_.insert(out_strides_.begin(), stride_i); |
| 581 | } |
| 582 | } |
| 583 | |
| 584 | // Step (2)~(4) for the algorithm in NOTE [ Detecting Memory Overlap Within A |
| 585 | // Strided Tensor ] |
| 586 | // on output geometry |
| 587 | const bool out_maybe_overlap = IsOverlappingMemorys(out_sizes_, out_strides_); |
| 588 | |
| 589 | // For input geometry, |
| 590 | // check for size 0 dimensions, |
| 591 | // skip size 1 dimensions, |
| 592 | // Step (0)~(1) for the algorithm in NOTE [ Detecting Memory Overlap Within A |
| 593 | // Strided Tensor ] |
| 594 | // on input geometry |
| 595 | auto idim = input->ndim(); |
| 596 | std::vector<int64_t> inp_sizes(input->shape()->begin(), input->shape()->end()); |
| 597 | std::vector<int64_t> inp_strides(JUST(input->stride())->begin(), JUST(input->stride())->end()); |
| 598 | std::vector<int64_t> inp_sizes_, inp_strides_; |
| 599 | inp_sizes_.reserve(idim); |
| 600 | inp_strides_.reserve(idim); |
| 601 | for (int64_t i = idim - 1; i >= 0; i--) { |
| 602 | auto size_i = inp_sizes[i]; |
| 603 | auto stride_i = inp_strides[i]; |
| 604 | if (size_i == 0) { |
| 605 | return functional::Constant(*input->shape(), 0, grad->dtype(), JUST(grad->device())); |
| 606 | } else if (size_i != 1) { |
| 607 | inp_sizes_.insert(inp_sizes_.begin(), size_i); |
| 608 | inp_strides_.insert(inp_strides_.begin(), stride_i); |
| 609 | } |
| 610 | } |
| 611 | // Step (1)~(4) for the algorithm in NOTE [ Detecting Memory Overlap Within A |
| 612 | // Strided Tensor ] |
| 613 | // on input geometry |
| 614 | const bool inp_maybe_overlap = IsOverlappingMemorys(inp_sizes_, inp_strides_); |
| 615 |
no test coverage detected