| 633 | } |
| 634 | |
| 635 | phi::distributed::DistTensor* SetKernelDistOutput( |
| 636 | Tensor* out, const phi::distributed::ArgDistAttr& dist_attr) { |
| 637 | PADDLE_ENFORCE_EQ( |
| 638 | paddle::holds_alternative<phi::distributed::TensorDistAttr>(dist_attr), |
| 639 | true, |
| 640 | common::errors::PreconditionNotMet( |
| 641 | "Arg must be a single TensorDistAttr")); |
| 642 | if (out) { |
| 643 | if (out->impl() == nullptr) { |
| 644 | auto dist_t = std::make_shared<phi::distributed::DistTensor>( |
| 645 | phi::DDim(), paddle::get<0>(dist_attr)); |
| 646 | out->set_impl(dist_t); |
| 647 | } |
| 648 | return static_cast<phi::distributed::DistTensor*>(out->impl().get()); |
| 649 | } |
| 650 | return nullptr; |
| 651 | } |
| 652 | |
| 653 | std::vector<phi::distributed::DistTensor*> SetKernelDistOutput( |
| 654 | size_t out_size, std::vector<Tensor>* out) { |
no test coverage detected