MCPcopy Create free account
hub / github.com/Oneflow-Inc/oneflow / GetTensorDevice4CurrentProcessCtx

Function GetTensorDevice4CurrentProcessCtx

oneflow/core/job/parallel_desc.cpp:170–187  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

168}
169
170Maybe<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
189bool ParallelDesc::TryGetParallelId(int64_t machine_id, int64_t device_id,
190 int64_t* parallel_id) const {

Callers 6

CallMethod · 0.85
CallMethod · 0.85
RawLocalToGlobalFunction · 0.85
InterpretFunction · 0.85
RawRunGlobalNormalOpFunction · 0.85

Calls 3

findMethod · 0.80
endMethod · 0.45

Tested by

no test coverage detected