diff --git a/onnxruntime/python/dlpack_convertor.cc b/onnxruntime/python/dlpack_convertor.cc index a44442f1bb..07373dbaf5 100644 --- a/onnxruntime/python/dlpack_convertor.cc +++ b/onnxruntime/python/dlpack_convertor.cc @@ -84,7 +84,11 @@ DLContext get_dlpack_context(const OrtValue& ort_value, const int64_t& device_id ctx.device_type = DLDeviceType::kDLCPU; break; case OrtDevice::GPU: +#ifdef USE_ROCM + ctx.device_type = DLDeviceType::kDLROCM; +#else ctx.device_type = DLDeviceType::kDLGPU; +#endif break; default: ORT_THROW("Cannot pack tensors on this device."); diff --git a/orttraining/orttraining/training_ops/rocm/rocm_training_kernels.cc b/orttraining/orttraining/training_ops/rocm/rocm_training_kernels.cc index 74788a50ad..616a98964d 100644 --- a/orttraining/orttraining/training_ops/rocm/rocm_training_kernels.cc +++ b/orttraining/orttraining/training_ops/rocm/rocm_training_kernels.cc @@ -124,6 +124,7 @@ class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, Recv class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, RecordEvent); class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, WaitEvent); +class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, YieldOp); #ifdef ORT_USE_NCCL class ONNX_OPERATOR_KERNEL_CLASS_NAME(kRocmExecutionProvider, kMSDomain, 1, NcclAllReduce); @@ -247,6 +248,7 @@ Status RegisterRocmTrainingKernels(KernelRegistry& kernel_registry) { // BuildKernelCreateInfo, // BuildKernelCreateInfo, + BuildKernelCreateInfo, #ifdef ORT_USE_NCCL BuildKernelCreateInfo,