hipMemcpy[To/From]Symbol(Async) fixes (#1774)

This commit is contained in:
satyanveshd
2020-01-07 08:11:53 +05:30
committed by Maneesh Gupta
parent 7d27247814
commit 6b5ea15dfe
+38 -6
View File
@@ -1273,8 +1273,18 @@ hipError_t hipMemcpyToSymbol(void* dst, const void* src, size_t count,
tprintf(DB_MEM, " symbol '%s' resolved to address:%p\n", symbol_name, dst); tprintf(DB_MEM, " symbol '%s' resolved to address:%p\n", symbol_name, dst);
return ihipLogStatus( if (dst == nullptr) {
hipMemcpy(static_cast<char*>(dst) + offset, src, count, kind)); return ihipLogStatus(hipErrorInvalidSymbol);
}
if (kind == hipMemcpyDeviceToHost || kind == hipMemcpyHostToHost) {
return ihipLogStatus(hipErrorInvalidMemcpyDirection);
} else if (kind == hipMemcpyDeviceToDevice) {
return ihipLogStatus(hipErrorInvalidValue);
}
return ihipLogStatus(hip_internal::memcpySync(static_cast<char*>(dst)+offset, src, count, kind,
hipStreamNull));
} }
hipError_t hipMemcpyFromSymbol(void* dst, const void* src, size_t count, hipError_t hipMemcpyFromSymbol(void* dst, const void* src, size_t count,
@@ -1285,8 +1295,18 @@ hipError_t hipMemcpyFromSymbol(void* dst, const void* src, size_t count,
tprintf(DB_MEM, " symbol '%s' resolved to address:%p\n", symbol_name, dst); tprintf(DB_MEM, " symbol '%s' resolved to address:%p\n", symbol_name, dst);
return ihipLogStatus( if (src == nullptr || dst == nullptr) {
hipMemcpy(dst, static_cast<const char*>(src) + offset, count, kind)); return ihipLogStatus(hipErrorInvalidSymbol);
}
if (kind == hipMemcpyHostToDevice || kind == hipMemcpyHostToHost) {
return ihipLogStatus(hipErrorInvalidMemcpyDirection);
} else if (kind == hipMemcpyDeviceToDevice) {
return ihipLogStatus(hipErrorInvalidValue);
}
return ihipLogStatus(hip_internal::memcpySync(dst, static_cast<const char*>(src)+offset, count, kind,
hipStreamNull));
} }
@@ -1301,11 +1321,17 @@ hipError_t hipMemcpyToSymbolAsync(void* dst, const void* src, size_t count,
if (dst == nullptr) { if (dst == nullptr) {
return ihipLogStatus(hipErrorInvalidSymbol); return ihipLogStatus(hipErrorInvalidSymbol);
} }
if (kind == hipMemcpyDeviceToHost || kind == hipMemcpyHostToHost) {
return ihipLogStatus(hipErrorInvalidMemcpyDirection);
} else if (kind == hipMemcpyDeviceToDevice) {
return ihipLogStatus(hipErrorInvalidValue);
}
hipError_t e = hipSuccess; hipError_t e = hipSuccess;
if (stream) { if (stream) {
try { try {
hip_internal::memcpyAsync((char*)dst+offset, src, count, kind, stream); hip_internal::memcpyAsync(static_cast<char*>(dst)+offset, src, count, kind, stream);
} catch (ihipException& ex) { } catch (ihipException& ex) {
e = ex._code; e = ex._code;
} }
@@ -1327,12 +1353,18 @@ hipError_t hipMemcpyFromSymbolAsync(void* dst, const void* src, size_t count,
if (src == nullptr || dst == nullptr) { if (src == nullptr || dst == nullptr) {
return ihipLogStatus(hipErrorInvalidSymbol); return ihipLogStatus(hipErrorInvalidSymbol);
} }
if (kind == hipMemcpyHostToDevice || kind == hipMemcpyHostToHost) {
return ihipLogStatus(hipErrorInvalidMemcpyDirection);
} else if (kind == hipMemcpyDeviceToDevice) {
return ihipLogStatus(hipErrorInvalidValue);
}
hipError_t e = hipSuccess; hipError_t e = hipSuccess;
stream = ihipSyncAndResolveStream(stream); stream = ihipSyncAndResolveStream(stream);
if (stream) { if (stream) {
try { try {
hip_internal::memcpyAsync(dst, (char*)src+offset, count, kind, stream); hip_internal::memcpyAsync(dst, static_cast<const char*>(src)+offset, count, kind, stream);
} catch (ihipException& ex) { } catch (ihipException& ex) {
e = ex._code; e = ex._code;
} }