mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
change the EP device to default OrtDevice() for memoryType equals CPUInput (#15903)
### Description <!-- Describe your changes. --> change the EP device to default OrtDevice() for memoryType equals CPUInput for cuda, rocm, migraph x and tensorRT EP ### Motivation and Context <!-- - Why is this change required? What problem does it solve? - If it fixes an open issue, please link to the issue here. --> My previous PR (https://github.com/microsoft/onnxruntime/pull/15618) caused random failures on cuda training test GradientCheckerTest.TileGrad (see build https://dev.azure.com/onnxruntime/onnxruntime/_build/results?buildId=986784&view=logs&j=5076e696-f193-5f12-2d8a-703dda41a79b&t=a3824a7c-2162-5e3d-3fdd-8cf808834fbb) and rocm test: root@a59558217e53:/workspace# pytest orttraining/orttraining/test/python/orttraining_test_ortmodule_api.py::test_gradient_correctness_minmax ... E RuntimeError: Error in backward pass execution: Non-zero status code returned while running ATen node. Name:'/_original_module/ATen_Grad/ATen_1' Status Message: Storage size calculation overflowed with sizes=[72340172838076673, 72340172838076673, 128] Potential reason is that if the memType of cuda/tensorRT/rocm/migraphx EP is CPUInput, previously the corresponding device in the IAllocator's memoryInfo is default OrtDevice(), while after my change, it becomes OrtDevice(CPU, xx_PINNED, 0); Changing it back fixed GradientCheckerTest.TileGrad in Win GPU training build.
This commit is contained in:
parent
af6cb2af87
commit
3b8f3a086f
5 changed files with 9 additions and 12 deletions
|
|
@ -1324,6 +1324,7 @@ class PlannerImpl {
|
|||
#endif
|
||||
ORT_RETURN_IF_ERROR(ComputeSingleStreamReusePlan(i));
|
||||
ClearUseCount();
|
||||
freelist_.clear(); // DONOT share freelist across streams
|
||||
}
|
||||
#if !defined(ORT_MINIMAL_BUILD) && defined(ORT_MEMORY_PROFILE)
|
||||
CalculateLifetime(ort_value_usecount);
|
||||
|
|
|
|||
|
|
@ -2532,9 +2532,8 @@ void CUDAExecutionProvider::RegisterStreamHandlers(IStreamCommandHandleRegistry&
|
|||
}
|
||||
|
||||
OrtDevice CUDAExecutionProvider::GetOrtDeviceByMemType(OrtMemType mem_type) const {
|
||||
if (mem_type == OrtMemTypeCPUInput || mem_type == OrtMemTypeCPUOutput) {
|
||||
return OrtDevice(OrtDevice::CPU, OrtDevice::MemType::CUDA_PINNED, 0 /*CPU device id always be 0*/);
|
||||
}
|
||||
if (mem_type == OrtMemTypeCPUInput) return OrtDevice();
|
||||
if (mem_type == OrtMemTypeCPUOutput) return OrtDevice(OrtDevice::CPU, OrtDevice::MemType::CUDA_PINNED, 0 /*CPU device id always be 0*/);
|
||||
return default_device_;
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -1207,9 +1207,8 @@ void MIGraphXExecutionProvider::RegisterStreamHandlers(IStreamCommandHandleRegis
|
|||
}
|
||||
|
||||
OrtDevice MIGraphXExecutionProvider::GetOrtDeviceByMemType(OrtMemType mem_type) const {
|
||||
if (mem_type == OrtMemTypeCPUInput || mem_type == OrtMemTypeCPUOutput) {
|
||||
return OrtDevice(OrtDevice::CPU, OrtDevice::MemType::HIP_PINNED, 0 /*CPU device id always be 0*/);
|
||||
}
|
||||
if (mem_type == OrtMemTypeCPUInput) return OrtDevice();
|
||||
if (mem_type == OrtMemTypeCPUOutput) return OrtDevice(OrtDevice::CPU, OrtDevice::MemType::HIP_PINNED, 0 /*CPU device id always be 0*/);
|
||||
return default_device_;
|
||||
}
|
||||
#ifdef MIGRAPHX_STREAM_SYNC
|
||||
|
|
|
|||
|
|
@ -2337,9 +2337,8 @@ void ROCMExecutionProvider::RegisterStreamHandlers(IStreamCommandHandleRegistry&
|
|||
}
|
||||
|
||||
OrtDevice ROCMExecutionProvider::GetOrtDeviceByMemType(OrtMemType mem_type) const {
|
||||
if (mem_type == OrtMemTypeCPUInput || mem_type == OrtMemTypeCPUOutput) {
|
||||
return OrtDevice(OrtDevice::CPU, OrtDevice::MemType::HIP_PINNED, 0 /*CPU device id always be 0*/);
|
||||
}
|
||||
if (mem_type == OrtMemTypeCPUInput) return OrtDevice();
|
||||
if (mem_type == OrtMemTypeCPUOutput) return OrtDevice(OrtDevice::CPU, OrtDevice::MemType::HIP_PINNED, 0 /*CPU device id always be 0*/);
|
||||
return default_device_;
|
||||
}
|
||||
} // namespace onnxruntime
|
||||
|
|
|
|||
|
|
@ -2843,9 +2843,8 @@ void TensorrtExecutionProvider::RegisterStreamHandlers(IStreamCommandHandleRegis
|
|||
}
|
||||
|
||||
OrtDevice TensorrtExecutionProvider::GetOrtDeviceByMemType(OrtMemType mem_type) const {
|
||||
if (mem_type == OrtMemTypeCPUInput || mem_type == OrtMemTypeCPUOutput) {
|
||||
return OrtDevice(OrtDevice::CPU, OrtDevice::MemType::CUDA_PINNED, 0 /*CPU device id always be 0*/);
|
||||
}
|
||||
if (mem_type == OrtMemTypeCPUInput) return OrtDevice();
|
||||
if (mem_type == OrtMemTypeCPUOutput) return OrtDevice(OrtDevice::CPU, OrtDevice::MemType::CUDA_PINNED, 0 /*CPU device id always be 0*/);
|
||||
return default_device_;
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue