rocr: refactor VMemorySetAccess function

Refactor VMemorySetAccess so that it can be re-used in the following
patch.

Change-Id: I341241da7a59724bb3611172f0d26b0689d7bb46
This commit is contained in:
David Yat Sin
2024-09-12 18:50:43 +00:00
والد 588a5a2fd3
کامیت 8f1b05660a
3فایلهای تغییر یافته به همراه77 افزوده شده و 51 حذف شده
@@ -868,6 +868,11 @@ class Runtime {
};
std::map<const void*, MappedHandle> mapped_handle_map_; // Indexed by VA
hsa_status_t
VMemorySetAccessPerHandle(void *va, MappedHandle &MappedHandle,
const hsa_amd_memory_access_desc_t *desc,
const size_t desc_cnt);
// Frees runtime memory when the runtime library is unloaded if safe to do so.
// Failure to release the runtime indicates an incorrect application but is
// common (example: calls library routines at process exit).
@@ -1328,6 +1328,7 @@ hsa_status_t hsa_amd_vmem_set_access(void* va, size_t size,
TRY;
IS_OPEN();
IS_BAD_PTR(va);
IS_ZERO(size);
IS_BAD_PTR(desc);
IS_ZERO(desc_cnt);
@@ -3367,6 +3367,71 @@ hsa_status_t Runtime::MappedHandleAllowedAgent::RemoveAccess() {
return (ret) ? HSA_STATUS_ERROR : HSA_STATUS_SUCCESS;
}
// Note: VMemorySetAccessPerHandle should be called with &memory_lock_ held
hsa_status_t
Runtime::VMemorySetAccessPerHandle(void *va, MappedHandle &mappedHandle,
const hsa_amd_memory_access_desc_t *desc,
const size_t desc_cnt) {
for (int i = 0; i < desc_cnt; i++) {
Agent *targetAgent = Agent::Convert(desc[i].agent_handle);
const size_t &size = mappedHandle.size;
const hsa_access_permission_t &perm = desc[i].permissions;
auto agentPermsIt = mappedHandle.allowed_agents.find(targetAgent);
if (agentPermsIt == mappedHandle.allowed_agents.end()) {
/* Agent not previously allowed, we need a new entry */
mappedHandle.allowed_agents.emplace(
std::piecewise_construct, std::forward_as_tuple(targetAgent),
std::forward_as_tuple(&mappedHandle, targetAgent, va, size, perm));
if (mappedHandle.allowed_agents[targetAgent].EnableAccess(perm) !=
HSA_STATUS_SUCCESS) {
mappedHandle.allowed_agents.erase(targetAgent);
return HSA_STATUS_ERROR;
}
} else {
/* Previous permissions are same as current permission */
if (agentPermsIt->second.permissions == perm)
continue;
/* Permissions are different - update access */
if (agentPermsIt->second.RemoveAccess() != HSA_STATUS_SUCCESS)
throw AMD::hsa_exception(HSA_STATUS_ERROR, "Failed to remove access for memory handle.");
if (agentPermsIt->second.EnableAccess(perm) != HSA_STATUS_SUCCESS) {
mappedHandle.allowed_agents.erase(agentPermsIt);
return HSA_STATUS_ERROR;
}
}
}
// Remove agents that were previously allowed but not included in current list
for (auto agentPermsIt = mappedHandle.allowed_agents.begin();
agentPermsIt != mappedHandle.allowed_agents.end();) {
bool agent_removed = true;
for (int i = 0; i < desc_cnt; i++) {
Agent *checkAgent = Agent::Convert(desc[i].agent_handle);
if (agentPermsIt->first == checkAgent) {
agent_removed = false;
break;
}
}
if (agent_removed) {
if (agentPermsIt->second.RemoveAccess() != HSA_STATUS_SUCCESS)
throw AMD::hsa_exception(HSA_STATUS_ERROR, "Failed to remove access for memory handle.");
agentPermsIt = mappedHandle.allowed_agents.erase(agentPermsIt);
} else {
++agentPermsIt;
}
}
return HSA_STATUS_SUCCESS;
}
hsa_status_t Runtime::VMemorySetAccess(void* va, size_t size,
const hsa_amd_memory_access_desc_t* desc,
const size_t desc_cnt) {
@@ -3406,57 +3471,12 @@ hsa_status_t Runtime::VMemorySetAccess(void* va, size_t size,
va_chunk += mappedHandleIt->second.size;
}
for (int i = 0; i < desc_cnt; i++) {
Agent* targetAgent = Agent::Convert(desc[i].agent_handle);
for (auto mappedHandleIt : mappedHandles) {
auto agentPermsIt = mappedHandleIt.second->allowed_agents.find(targetAgent);
if (agentPermsIt == mappedHandleIt.second->allowed_agents.end()) {
/* Agent not previously allowed, we need a new entry */
mappedHandleIt.second->allowed_agents.emplace(
std::piecewise_construct, std::forward_as_tuple(targetAgent),
std::forward_as_tuple(mappedHandleIt.second, targetAgent, mappedHandleIt.first, size,
desc[i].permissions));
if (mappedHandleIt.second->allowed_agents[targetAgent].EnableAccess(desc[i].permissions) !=
HSA_STATUS_SUCCESS) {
mappedHandleIt.second->allowed_agents.erase(targetAgent);
return HSA_STATUS_ERROR;
}
} else {
/* Previous permissions are same as current permission */
if (agentPermsIt->second.permissions == desc[i].permissions) continue;
/* Permissions are different - update access */
if (agentPermsIt->second.RemoveAccess() != HSA_STATUS_SUCCESS) return HSA_STATUS_ERROR;
if (agentPermsIt->second.EnableAccess(desc[i].permissions) != HSA_STATUS_SUCCESS) {
mappedHandleIt.second->allowed_agents.erase(agentPermsIt);
return HSA_STATUS_ERROR;
}
// Remove agents that were previously allowed but not included in current list
for (auto agentPermsIt = mappedHandleIt.second->allowed_agents.begin();
agentPermsIt != mappedHandleIt.second->allowed_agents.end();) {
bool agent_removed = true;
for (int i = 0; i < desc_cnt; i++) {
if (agentPermsIt->first == Agent::Convert(desc[i].agent_handle)) {
agent_removed = false;
continue;
}
}
if (agent_removed) {
assert(agentPermsIt->second.va == va);
if (agentPermsIt->second.RemoveAccess() != HSA_STATUS_SUCCESS) return HSA_STATUS_ERROR;
agentPermsIt = mappedHandleIt.second->allowed_agents.erase(agentPermsIt);
} else {
++agentPermsIt;
}
}
}
}
hsa_status_t status;
for (auto mappedHandleIt : mappedHandles) {
status = VMemorySetAccessPerHandle(mappedHandleIt.first,
*mappedHandleIt.second, desc, desc_cnt);
if (status != HSA_STATUS_SUCCESS)
return status;
}
return HSA_STATUS_SUCCESS;
}