| 168 | } |
| 169 | |
| 170 | Maybe<Symbol<Device>> GetTensorDevice4CurrentProcessCtx(Symbol<ParallelDesc> parallel_desc, |
| 171 | Optional<int64_t>* parallel_id) { |
| 172 | static thread_local HashMap<Symbol<ParallelDesc>, Optional<int64_t>> parallel_desc2parallel_id; |
| 173 | static thread_local HashMap<Symbol<ParallelDesc>, Symbol<Device>> parallel_desc2device; |
| 174 | auto parallel_id_iter = parallel_desc2parallel_id.find(parallel_desc); |
| 175 | auto device_iter = parallel_desc2device.find(parallel_desc); |
| 176 | if (device_iter == parallel_desc2device.end()) { |
| 177 | CHECK_OR_RETURN(parallel_id_iter == parallel_desc2parallel_id.end()); |
| 178 | Optional<int64_t> id_val; |
| 179 | const auto& device_symbol = JUST(parallel_desc->GetTensorDevice4CurrentProcessCtx(&id_val)); |
| 180 | parallel_id_iter = parallel_desc2parallel_id.emplace(parallel_desc, id_val).first; |
| 181 | device_iter = parallel_desc2device.emplace(parallel_desc, device_symbol).first; |
| 182 | } else { |
| 183 | CHECK_OR_RETURN(parallel_id_iter != parallel_desc2parallel_id.end()); |
| 184 | } |
| 185 | *parallel_id = parallel_id_iter->second; |
| 186 | return device_iter->second; |
| 187 | } |
| 188 | |
| 189 | bool ParallelDesc::TryGetParallelId(int64_t machine_id, int64_t device_id, |
| 190 | int64_t* parallel_id) const { |
no test coverage detected