| 893 | } |
| 894 | |
| 895 | void 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 | } |
nothing calls this directly
no test coverage detected