Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 11 additions & 2 deletions src/backend/common/DefaultMemoryManager.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -140,10 +140,19 @@ float DefaultMemoryManager::getMemoryPressure() {
}
}

bool DefaultMemoryManager::jitTreeExceedsMemoryPressure(size_t bytes) {
bool DefaultMemoryManager::jitTreeExceedsMemoryPressure(
size_t jit_tree_buffer_bytes) {
lock_guard_t lock(this->memory_mutex);
memory_info &current = this->getCurrentMemoryInfo();
return 2 * bytes > current.lock_bytes;
if (current.lock_bytes > 0.25f * current.max_bytes) {
/// Evaluate JIT if half of all locked buffers are locked by this JIT
/// tree
return jit_tree_buffer_bytes > current.lock_bytes * 0.5f;
} else {
/// Evaluate if this JIT Tree accounts for 10% of total memory on the
/// device
return jit_tree_buffer_bytes > 0.10f * current.max_bytes;
}
}

void *DefaultMemoryManager::alloc(bool user_lock, const unsigned ndims,
Expand Down
11 changes: 11 additions & 0 deletions src/backend/cuda/Array.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
********************************************************/

#include <Array.hpp>
#include <common/Logger.hpp>
#include <common/half.hpp>
#include <common/jit/NodeIterator.hpp>
#include <copy.hpp>
Expand Down Expand Up @@ -253,8 +254,13 @@ Node_ptr Array<T>::getNode() const {
template<typename T>
kJITHeuristics passesJitHeuristics(span<Node *> root_nodes) {
if (!evalFlag()) { return kJITHeuristics::Pass; }
static auto getLogger = [&] { return spdlog::get("jit"); };
for (Node *n : root_nodes) {
if (n->getHeight() > static_cast<int>(getMaxJitSize())) {
AF_TRACE(
"JIT tree evaluated because of tree height exceeds limit: {} > "
"{}",
n->getHeight(), getMaxJitSize());
return kJITHeuristics::TreeHeight;
}
}
Expand Down Expand Up @@ -313,9 +319,14 @@ kJITHeuristics passesJitHeuristics(span<Node *> root_nodes) {
// should be checking the amount of memory available to guard
// this eval
if (param_size >= max_param_size) {
AF_TRACE(
"JIT tree evaluated because of kernel parameter size: {} >= {}",
param_size, max_param_size);
return kJITHeuristics::KernelParameterSize;
}
if (jitTreeExceedsMemoryPressure(info.total_buffer_size)) {
AF_TRACE("JIT tree evaluated because of memory pressure: {}",
info.total_buffer_size);
return kJITHeuristics::MemoryPressure;
}
}
Expand Down
20 changes: 17 additions & 3 deletions src/backend/oneapi/Array.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@
#include <Array.hpp>

#include <Param.hpp>
#include <common/Logger.hpp>
#include <common/MemoryManagerBase.hpp>
#include <common/half.hpp>
#include <common/jit/NodeIterator.hpp>
Expand Down Expand Up @@ -312,8 +313,13 @@ Node_ptr Array<T>::getNode() const {
template<typename T>
kJITHeuristics passesJitHeuristics(span<Node *> root_nodes) {
if (!evalFlag()) { return kJITHeuristics::Pass; }
static auto getLogger = [&] { return common::loggerFactory("jit"); };
for (const Node *n : root_nodes) {
if (n->getHeight() > static_cast<int>(getMaxJitSize())) {
AF_TRACE(
"JIT tree evaluated because of tree height exceeds limit: {} > "
"{}",
n->getHeight(), getMaxJitSize());
return kJITHeuristics::TreeHeight;
}
}
Expand Down Expand Up @@ -386,9 +392,17 @@ kJITHeuristics passesJitHeuristics(span<Node *> root_nodes) {

bool isParamLimit = param_size >= max_param_size;

if (isParamLimit) { return kJITHeuristics::KernelParameterSize; }
// TODO(umar): check buffer limit for JIT kernel generation
// if (isBufferLimit) { return kJITHeuristics::MemoryPressure; }
if (isParamLimit) {
AF_TRACE(
"JIT tree evaluated because of kernel parameter size: {} >= {}",
param_size, max_param_size);
return kJITHeuristics::KernelParameterSize;
}
if (isBufferLimit) {
AF_TRACE("JIT tree evaluated because of memory pressure: {}",
info.total_buffer_size);
return kJITHeuristics::MemoryPressure;
}
}
return kJITHeuristics::Pass;
}
Expand Down
14 changes: 14 additions & 0 deletions src/backend/oneapi/jit.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
#include <kernel_headers/jit.hpp>

#include <Array.hpp>
#include <Kernel.hpp>
#include <common/dispatch.hpp>
#include <common/half.hpp>
#include <common/jit/ModdimNode.hpp>
Expand Down Expand Up @@ -591,6 +592,19 @@ void evalNodes(vector<Param<T>>& outputs, const vector<Node*>& output_nodes) {
(size_t)ap[0].dims[2]};
ndims = 3;
}

{
using namespace oneapi::kernel_logger;
AF_TRACE(
"Launching {}: Dims: [{},{},{},{}] Global: "
"[{},{},{}] threads: {}",
funcName, ap[0].dims[0], ap[0].dims[1],
ap[0].dims[2], ap[0].dims[3], global[0], global[1],
global[2],
global[0] * std::max<size_t>(1, global[1]) *
std::max<size_t>(1, global[2]));
}

cl_event kernel_event;
CL_CHECK(clEnqueueNDRangeKernel(
q, kernel, ndims, offset.data(), global.data(), nullptr,
Expand Down
19 changes: 17 additions & 2 deletions src/backend/opencl/Array.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@

#include <Array.hpp>

#include <common/Logger.hpp>
#include <common/half.hpp>
#include <common/jit/NodeIterator.hpp>
#include <common/jit/ScalarNode.hpp>
Expand Down Expand Up @@ -301,8 +302,13 @@ Node_ptr Array<T>::getNode() const {
template<typename T>
kJITHeuristics passesJitHeuristics(span<Node *> root_nodes) {
if (!evalFlag()) { return kJITHeuristics::Pass; }
static auto getLogger = [&] { return common::loggerFactory("jit"); };
for (const Node *n : root_nodes) {
if (n->getHeight() > static_cast<int>(getMaxJitSize())) {
AF_TRACE(
"JIT tree evaluated because of tree height exceeds limit: {} > "
"{}",
n->getHeight(), getMaxJitSize());
return kJITHeuristics::TreeHeight;
}
}
Expand Down Expand Up @@ -377,8 +383,17 @@ kJITHeuristics passesJitHeuristics(span<Node *> root_nodes) {

bool isParamLimit = param_size >= max_param_size;

if (isParamLimit) { return kJITHeuristics::KernelParameterSize; }
if (isBufferLimit) { return kJITHeuristics::MemoryPressure; }
if (isParamLimit) {
AF_TRACE(
"JIT tree evaluated because of kernel parameter size: {} >= {}",
param_size, max_param_size);
return kJITHeuristics::KernelParameterSize;
}
if (isBufferLimit) {
AF_TRACE("JIT tree evaluated because of memory pressure: {}",
info.total_buffer_size);
return kJITHeuristics::MemoryPressure;
}
}
return kJITHeuristics::Pass;
}
Expand Down