| 498 | } |
| 499 | |
| 500 | std::shared_ptr<phi::SelectedRows> PrepareDataForSelectedRows( |
| 501 | const Tensor& input, |
| 502 | const phi::TensorArgDef& target_args_def, |
| 503 | const TransformFlag& transform_flag) { |
| 504 | const auto& tensor_in = input.impl(); |
| 505 | if (tensor_in) { |
| 506 | phi::SelectedRows& selected_rows = |
| 507 | *static_cast<phi::SelectedRows*>(tensor_in.get()); |
| 508 | if ((!transform_flag.NeedTransform() || !selected_rows.initialized() || |
| 509 | (!NeedTransformPlace(selected_rows.place(), |
| 510 | target_args_def.backend, |
| 511 | transform_flag))) && |
| 512 | !NeedTransform2Contiguous( |
| 513 | false, selected_rows.value().meta().is_contiguous())) { |
| 514 | if (NeedTransform2Contiguous( |
| 515 | false, selected_rows.value().meta().is_contiguous()) && |
| 516 | selected_rows.initialized()) { |
| 517 | auto out_new = std::make_shared<phi::SelectedRows>( |
| 518 | selected_rows.rows(), selected_rows.height()); |
| 519 | auto dense_out = Trans2Contiguous(selected_rows.value()); |
| 520 | *out_new->mutable_value() = dense_out; |
| 521 | return out_new; |
| 522 | } |
| 523 | return std::static_pointer_cast<phi::SelectedRows>(tensor_in); |
| 524 | } |
| 525 | |
| 526 | if (selected_rows.place().GetType() == AllocationType::GPUPINNED) { |
| 527 | if (NeedTransform2Contiguous( |
| 528 | false, selected_rows.value().meta().is_contiguous())) { |
| 529 | auto dense_out = Trans2Contiguous(selected_rows.value()); |
| 530 | selected_rows.mutable_value()->ShareDataWith(dense_out); |
| 531 | } |
| 532 | if (transform_flag.NeedTransform() && selected_rows.initialized() && |
| 533 | NeedTransformPlace( |
| 534 | selected_rows.place(), target_args_def.backend, transform_flag)) { |
| 535 | auto dense_out = |
| 536 | TransDataPlace(selected_rows.value(), |
| 537 | phi::TransToPhiPlace(target_args_def.backend)); |
| 538 | selected_rows.mutable_value()->ShareBufferWith(dense_out); |
| 539 | } |
| 540 | return std::static_pointer_cast<phi::SelectedRows>(tensor_in); |
| 541 | } else { |
| 542 | auto out_new = std::make_shared<phi::SelectedRows>( |
| 543 | selected_rows.rows(), selected_rows.height()); |
| 544 | if (NeedTransform2Contiguous( |
| 545 | false, selected_rows.value().meta().is_contiguous())) { |
| 546 | auto dense_out = Trans2Contiguous(selected_rows.value()); |
| 547 | *out_new->mutable_value() = dense_out; |
| 548 | } |
| 549 | if (transform_flag.NeedTransform() && selected_rows.initialized() && |
| 550 | NeedTransformPlace( |
| 551 | selected_rows.place(), target_args_def.backend, transform_flag)) { |
| 552 | auto dense_out = |
| 553 | TransDataPlace(selected_rows.value(), |
| 554 | phi::TransToPhiPlace(target_args_def.backend)); |
| 555 | *out_new->mutable_value() = dense_out; |
| 556 | } |
| 557 | return out_new; |
nothing calls this directly
no test coverage detected