| 73 | } |
| 74 | |
| 75 | Maybe<void> GetItemInScalarTensor(const std::shared_ptr<Tensor>& scalar_tensor, void* scalar_ptr, |
| 76 | size_t size) { |
| 77 | CHECK_EQ_OR_RETURN(GetSizeOfDataType(scalar_tensor->dtype()->data_type()), size) |
| 78 | << "invalid size"; |
| 79 | CHECK_OR_RETURN(scalar_tensor->is_eager()) << "Only eager scalar tensor support GetItem."; |
| 80 | CHECK_EQ_OR_RETURN(scalar_tensor->nelement(), 1) |
| 81 | << "can only convert a tensor of size 1 to a Python scalar"; |
| 82 | std::shared_ptr<LocalTensor> local_tensor; |
| 83 | { |
| 84 | auto tensor = scalar_tensor; |
| 85 | if (tensor->is_global()) { |
| 86 | Symbol<ParallelDesc> parallel_desc; |
| 87 | { |
| 88 | const ParallelConf parallel_conf = GenParallelConfOfCpuOnAllRanks(); |
| 89 | JUST(PhysicalRun( |
| 90 | [¶llel_desc, ¶llel_conf](InstructionsBuilder* builder) -> Maybe<void> { |
| 91 | parallel_desc = SymbolOf(*JUST(builder->GetParallelDescSymbol(parallel_conf))); |
| 92 | return Maybe<void>::Ok(); |
| 93 | })); |
| 94 | } |
| 95 | const auto& broadcast_sbp = JUST(MakeBroadcastSbpParallel()); |
| 96 | tensor = JUST(functional::ToGlobal(tensor, parallel_desc, {broadcast_sbp}, /*grad_sbp=*/{}, |
| 97 | /*check_meta=*/false, /*copy=*/false)); |
| 98 | tensor = JUST(functional::GlobalToLocal(tensor, /*copy=*/false)); |
| 99 | } |
| 100 | local_tensor = JUST(tensor->AsLocalTensor()); |
| 101 | } |
| 102 | JUST(SyncReadSmallMem(reinterpret_cast<char*>(scalar_ptr), size, local_tensor)); |
| 103 | return Maybe<void>::Ok(); |
| 104 | } |
| 105 | |
| 106 | } // namespace one |
| 107 | } // namespace oneflow |
no test coverage detected