MCPcopy Create free account
hub / github.com/Oneflow-Inc/oneflow / AsStridedGrad

Function AsStridedGrad

oneflow/core/framework/tensor_methods.cpp:558–679  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

556}
557
558Maybe<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

Callers 4

AsStridedFunction · 0.85
InplaceAsStridedFunction · 0.85
ApplyMethod · 0.85
operator()Method · 0.85

Calls 15

ConstantClass · 0.85
ReduceSumClass · 0.85
IsOverlappingMemorysFunction · 0.85
MinStorageSizeFunction · 0.85
ReshapeFunction · 0.85
EllipsisIndexClass · 0.85
ndimMethod · 0.80
insertMethod · 0.80
SqueezeFunction · 0.70
ShapeClass · 0.70
AsStridedFunction · 0.70
ExpandFunction · 0.70

Tested by

no test coverage detected