SWDEV-351969 - TLS Optimization

- Aggregate all TLS(Thread Local Storage) variables into a single class
- This is to improve cache accesses per thread

Change-Id: Ic8361eaeae290fff00254684e309471958365eb9


[ROCm/clr commit: 8b391ef18c]
This commit is contained in:
Rakesh Roy
2022-09-26 15:59:27 +05:30
committed by Rakesh Roy
parent 2be45a82ec
commit f149b21399
9 changed files with 145 additions and 115 deletions
+18 -21
View File
@@ -28,10 +28,7 @@
std::vector<hip::Device*> g_devices;
namespace hip {
thread_local Device* g_device = nullptr;
thread_local std::stack<Device*> g_ctxtStack;
thread_local hipError_t g_lastError = hipSuccess;
thread_local TlsAggregator tls;
Device* host_device = nullptr;
//init() is only to be called from the HIP_INIT macro only once
@@ -84,13 +81,13 @@ bool init() {
}
Device* getCurrentDevice() {
return g_device;
return tls.device_;
}
void setCurrentDevice(unsigned int index) {
assert(index<g_devices.size());
g_device = g_devices[index];
uint32_t preferredNumaNode = g_device->devices()[0]->getPreferredNumaNode();
tls.device_ = g_devices[index];
uint32_t preferredNumaNode = (tls.device_)->devices()[0]->getPreferredNumaNode();
amd::Os::setPreferredNumaNode(preferredNumaNode);
}
@@ -160,7 +157,7 @@ hipError_t hipCtxCreate(hipCtx_t *ctx, unsigned int flags, hipDevice_t device)
// Increment ref count for device primary context
g_devices[device]->retain();
g_ctxtStack.push(g_devices[device]);
tls.ctxt_stack_.push(g_devices[device]);
HIP_RETURN(hipSuccess);
}
@@ -169,15 +166,15 @@ hipError_t hipCtxSetCurrent(hipCtx_t ctx) {
HIP_INIT_API(hipCtxSetCurrent, ctx);
if (ctx == nullptr) {
if(!g_ctxtStack.empty()) {
g_ctxtStack.pop();
if(!tls.ctxt_stack_.empty()) {
tls.ctxt_stack_.pop();
}
} else {
hip::g_device = reinterpret_cast<hip::Device*>(ctx);
if(!g_ctxtStack.empty()) {
g_ctxtStack.pop();
hip::tls.device_ = reinterpret_cast<hip::Device*>(ctx);
if(!tls.ctxt_stack_.empty()) {
tls.ctxt_stack_.pop();
}
g_ctxtStack.push(hip::getCurrentDevice());
tls.ctxt_stack_.push(hip::getCurrentDevice());
}
HIP_RETURN(hipSuccess);
@@ -221,8 +218,8 @@ hipError_t hipCtxDestroy(hipCtx_t ctx) {
}
// Need to remove the ctx of calling thread if its the top one
if (!g_ctxtStack.empty() && g_ctxtStack.top() == dev) {
g_ctxtStack.pop();
if (!tls.ctxt_stack_.empty() && tls.ctxt_stack_.top() == dev) {
tls.ctxt_stack_.pop();
}
// Remove context from global context list
@@ -240,11 +237,11 @@ hipError_t hipCtxPopCurrent(hipCtx_t* ctx) {
HIP_INIT_API(hipCtxPopCurrent, ctx);
hip::Device** dev = reinterpret_cast<hip::Device**>(ctx);
if (!g_ctxtStack.empty()) {
if (!tls.ctxt_stack_.empty()) {
if (dev != nullptr) {
*dev = g_ctxtStack.top();
*dev = tls.ctxt_stack_.top();
}
g_ctxtStack.pop();
tls.ctxt_stack_.pop();
} else {
DevLogError("Context Stack empty \n");
HIP_RETURN(hipErrorInvalidContext);
@@ -261,8 +258,8 @@ hipError_t hipCtxPushCurrent(hipCtx_t ctx) {
HIP_RETURN(hipErrorInvalidContext);
}
hip::g_device = dev;
g_ctxtStack.push(hip::getCurrentDevice());
hip::tls.device_ = dev;
tls.ctxt_stack_.push(hip::getCurrentDevice());
HIP_RETURN(hipSuccess);
}