MCPcopy Create free account
hub / github.com/PaddlePaddle/Paddle / HandleComplexGradToRealGrad

Method HandleComplexGradToRealGrad

paddle/fluid/eager/grad_node_info.cc:895–953  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

893}
894
895void GradNodeBase::HandleComplexGradToRealGrad(
896 paddle::small_vector<std::vector<paddle::Tensor>, kSlotSmallVectorSize>*
897 out_grads) {
898 for (size_t slot_id = 0; slot_id < out_grads->size(); slot_id++) {
899 const std::vector<paddle::Tensor>& slot_out_grads = (*out_grads)[slot_id];
900 for (size_t rank_id = 0; rank_id < slot_out_grads.size(); rank_id++) {
901 if (bwd_out_meta_[slot_id].size() == 0) continue;
902 const GradSlotMeta& slot_meta = bwd_out_meta_[slot_id][rank_id];
903 PADDLE_ENFORCE(
904 slot_meta.HasTensorMeta() > 0,
905 common::errors::Fatal(
906 "We require TensorMeta in GradInputMeta() to obtain forward data "
907 "types."
908 "However, no TensorMeta is detected in bwd_out_meta_."));
909
910 auto fwd_data_type = paddle::framework::TransToProtoVarType(
911 slot_meta.GetTensorMeta().dtype);
912 const paddle::Tensor& grad = slot_out_grads[rank_id];
913
914 if (paddle::framework::IsComplexType(fwd_data_type)) continue;
915 if (!grad.impl()) continue;
916
917 // Only Handle Complex To Real for DenseTensor for now
918 if (phi::DenseTensor::classof(grad.impl().get())) {
919 phi::DenseTensor* grad_dense_tensor =
920 static_cast<phi::DenseTensor*>(grad.impl().get());
921
922 auto curr_data_type =
923 paddle::framework::TransToProtoVarType(grad_dense_tensor->type());
924 if (!paddle::framework::IsComplexType(curr_data_type)) continue;
925
926 // Convert Complex GradOut to Real
927 auto out = std::make_shared<phi::DenseTensor>();
928 paddle::framework::TransComplexToReal(
929 fwd_data_type, curr_data_type, *grad_dense_tensor, out.get());
930
931 (*out_grads)[slot_id][rank_id].set_impl(out);
932 } else if (phi::distributed::DistTensor::classof(grad.impl().get())) {
933 auto grad_dense_tensor =
934 static_cast<phi::distributed::DistTensor*>(grad.impl().get())
935 ->value();
936
937 auto curr_data_type =
938 paddle::framework::TransToProtoVarType(grad_dense_tensor.type());
939 if (!paddle::framework::IsComplexType(curr_data_type)) continue;
940 if (grad_dense_tensor.dims().size() == -1) continue;
941
942 // Convert Complex GradOut to Real
943 auto out = std::make_shared<phi::DenseTensor>();
944 paddle::framework::TransComplexToReal(
945 fwd_data_type, curr_data_type, grad_dense_tensor, out.get());
946
947 *(static_cast<phi::distributed::DistTensor*>(
948 (*out_grads)[slot_id][rank_id].impl().get())
949 ->unsafe_mutable_value()) = *(out.get());
950 }
951 }
952 }

Callers

nothing calls this directly

Calls 12

TransComplexToRealFunction · 0.85
HasTensorMetaMethod · 0.80
set_implMethod · 0.80
TransToProtoVarTypeFunction · 0.50
IsComplexTypeFunction · 0.50
classofFunction · 0.50
sizeMethod · 0.45
implMethod · 0.45
getMethod · 0.45
typeMethod · 0.45
valueMethod · 0.45
dimsMethod · 0.45

Tested by

no test coverage detected