Skip to content
Open
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
57 changes: 50 additions & 7 deletions include/deep_jit/cache/memory.hpp
Original file line number Diff line number Diff line change
@@ -1,23 +1,66 @@
#pragma once

#include <exception>
#include <future>
#include <memory>
#include <mutex>
#include <unordered_map>

namespace deep_jit {

// In-memory cache keyed by value type.
template <typename Key, typename Value>
class MemCache {
using ValuePtr = std::shared_ptr<Value>;

std::unordered_map<Key, std::shared_future<ValuePtr>> pending;
mutable std::mutex mutex;

public:
std::unordered_map<Key, std::shared_ptr<Value>> cache;
// Kept public for source compatibility. Concurrent callers must use
// get_or_create() rather than accessing the container directly.
std::unordered_map<Key, ValuePtr> cache;

template <typename Factory>
std::shared_ptr<Value> get_or_create(const Key& key, Factory&& factory) {
if (const auto iterator = cache.find(key); iterator != cache.end())
return iterator->second;
auto value = factory();
cache.emplace(key, value);
return value;
ValuePtr get_or_create(const Key& key, Factory&& factory) {
std::shared_future<ValuePtr> future;
std::shared_ptr<std::promise<ValuePtr>> producer;
{
std::lock_guard lock(mutex);
if (const auto iterator = cache.find(key); iterator != cache.end())
return iterator->second;
if (const auto iterator = pending.find(key); iterator != pending.end()) {
future = iterator->second;
} else {
producer = std::make_shared<std::promise<ValuePtr>>();
future = producer->get_future().share();
pending.emplace(key, future);
}
}

// A miss for a different key must not wait for this factory. Callers
// that lost the race for the same key wait on the shared result.
if (producer == nullptr)
return future.get();

try {
auto value = factory();
{
std::lock_guard lock(mutex);
cache.emplace(key, value);
pending.erase(key);
}
producer->set_value(value);
return value;
} catch (...) {
const auto exception = std::current_exception();
producer->set_exception(exception);
{
std::lock_guard lock(mutex);
pending.erase(key);
}
std::rethrow_exception(exception);
}
}
};

Expand Down
4 changes: 4 additions & 0 deletions include/deep_jit/runtime/runtime.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,10 @@ class Runtime {
std::shared_ptr<Kernel> compile(const std::string& name, const std::string& source, const CompilerOptions& override_options = {}) {
const auto options = default_compiler_options.override_with(override_options);
const auto key = cache_key(source, options);
// A waiter for the same key must not retain the Python GIL while the
// producer needs to reacquire it after compiling. Cover the complete
// single-flight operation; nested releases in the backend are no-ops.
GilScopedRelease gil_release;
return mem_cache.get_or_create(key, [&] {
return Backend::load(compile(name, source, key, options), env);
});
Expand Down
75 changes: 75 additions & 0 deletions tests/test_cuda.py
Original file line number Diff line number Diff line change
Expand Up @@ -283,6 +283,80 @@ def worker():
print(f'validated compile-time GIL release ({progress[0]} observer iterations)', flush=True)


def validate_concurrent_compile_single_flight(module_path, temporary_dir):
cache_root = temporary_dir / 'gil_single_flight_cache'
compiler_barrier_dir = temporary_dir / 'gil_single_flight_compiler_barrier'
compiler_barrier_dir.mkdir()
code = '''
import importlib.util
import sys
import threading

import torch

spec = importlib.util.spec_from_file_location('deep_jit_cuda_test', sys.argv[1])
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
module.prepare_gil_runtime()

start = threading.Barrier(3)
kernel_addresses = []
errors = []

def compile_same_key():
try:
torch.cuda.set_device(0)
start.wait()
kernel_addresses.append(module.compile_cached_for_gil_test())
except BaseException as exception:
errors.append(repr(exception))

threads = [threading.Thread(target=compile_same_key) for _ in range(2)]
for thread in threads:
thread.start()
start.wait()
for thread in threads:
thread.join()

assert not errors, errors
assert len(kernel_addresses) == 2, kernel_addresses
assert len(set(kernel_addresses)) == 1, kernel_addresses
print(f'KERNEL_ADDRESS={kernel_addresses[0]}')
'''
env = os.environ.copy()
env.update({
'GIL_TEST_JIT_CACHE_DIR': str(cache_root),
'GIL_TEST_JIT_NVCC_COMPILER': str(TEST_CUDA_PROJECT / 'scripts' / 'nvcc_barrier.py'),
'DEEP_JIT_TEST_REAL_NVCC': str(Path(CUDA_HOME) / 'bin' / 'nvcc'),
'DEEP_JIT_TEST_NVCC_BARRIER_DIR': str(compiler_barrier_dir),
'DEEP_JIT_TEST_NVCC_BARRIER_SIZE': '1',
'DEEP_JIT_TEST_NVCC_DELAY_SECONDS': '1',
})
process = register_process_group(subprocess.Popen(
[sys.executable, '-c', code, str(module_path)],
env=env,
stdout=subprocess.PIPE,
stderr=subprocess.STDOUT,
text=True,
start_new_session=True,
))
try:
try:
output, _ = process.communicate(timeout=120)
except subprocess.TimeoutExpired as exception:
raise AssertionError('same-key compile deadlocked while waiting for the GIL') from exception
assert process.returncode == 0, output
address_lines = [line for line in output.splitlines() if line.startswith('KERNEL_ADDRESS=')]
assert len(address_lines) == 1, output
assert int(address_lines[0].removeprefix('KERNEL_ADDRESS=')) != 0, output
finally:
terminate_process_group(process)

compiler_invocations = list(compiler_barrier_dir.iterdir())
assert len(compiler_invocations) == 1, compiler_invocations
print('validated same-key compile single-flight without GIL deadlock', flush=True)


def validate_diagnostic_output(module_path, temporary_dir):
code = '''
import importlib.util
Expand Down Expand Up @@ -811,6 +885,7 @@ def run_worker():
raise AssertionError('invalid CUDA source unexpectedly compiled')
assert module.run_registered_jit(33) == 34
validate_compile_releases_gil(module, temporary_dir)
validate_concurrent_compile_single_flight(module.__file__, temporary_dir)
validate_direct_disk_cache_publication(module.__file__, temporary_dir)
validate_crashed_disk_cache_writer(module.__file__, temporary_dir)
module.run_tests(module)
Expand Down
7 changes: 7 additions & 0 deletions tests/test_cuda_proj/main.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3373,6 +3373,12 @@ std::string compile_for_gil_test() {
return gil_test_runtime->compile_without_load("gil_compile", get_template_source(81)).string();
}

std::uintptr_t compile_cached_for_gil_test() {
DJ_HOST_ASSERT(gil_test_runtime != nullptr, "GIL test runtime was not prepared");
return reinterpret_cast<std::uintptr_t>(
gil_test_runtime->compile("gil_single_flight", get_template_source(83)).get());
}

void init_python_api_jit(const std::string& library_root) {
const auto root = fs::absolute(library_root).lexically_normal();
python_api_jit = deep_jit::create_lazy_jit<deep_jit::CUDA>(deep_jit::Config(
Expand Down Expand Up @@ -3412,4 +3418,5 @@ PYBIND11_MODULE(TORCH_EXTENSION_NAME, module) {
module.def("publish_disk_cache_entry", &publish_disk_cache_entry);
module.def("prepare_gil_runtime", &prepare_gil_runtime);
module.def("compile_for_gil_test", &compile_for_gil_test);
module.def("compile_cached_for_gil_test", &compile_cached_for_gil_test);
}
159 changes: 159 additions & 0 deletions tests/test_memory_cache.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,159 @@
import os
import subprocess
import tempfile
from pathlib import Path


ROOT = Path(__file__).resolve().parents[1]


CPP_SOURCE = r'''
#include <atomic>
#include <barrier>
#include <chrono>
#include <cstdlib>
#include <memory>
#include <stdexcept>
#include <string>
#include <thread>
#include <vector>

#include <deep_jit/cache/memory.hpp>

namespace {

void require(bool condition) {
if (not condition)
std::abort();
}

void test_same_key_single_flight() {
deep_jit::MemCache<std::string, int> cache;
constexpr int num_threads = 16;
std::barrier start(num_threads);
std::atomic<int> factory_calls = 0;
std::vector<std::shared_ptr<int>> values(num_threads);
std::vector<std::thread> threads;

for (int index = 0; index < num_threads; ++index) {
threads.emplace_back([&, index] {
start.arrive_and_wait();
values[index] = cache.get_or_create("same", [&] {
++factory_calls;
std::this_thread::sleep_for(std::chrono::milliseconds(20));
return std::make_shared<int>(17);
});
});
}
for (auto& thread : threads)
thread.join();

require(factory_calls == 1);
for (const auto& value : values)
require(value == values.front() and *value == 17);
}

void test_different_keys_remain_parallel() {
deep_jit::MemCache<int, int> cache;
std::barrier start(2);
std::atomic<int> active = 0;
std::atomic<int> max_active = 0;

auto worker = [&](int key) {
start.arrive_and_wait();
return cache.get_or_create(key, [&] {
const int current = ++active;
int observed = max_active.load();
while (observed < current and
not max_active.compare_exchange_weak(observed, current)) {}
std::this_thread::sleep_for(std::chrono::milliseconds(30));
--active;
return std::make_shared<int>(key);
});
};

std::shared_ptr<int> first;
std::shared_ptr<int> second;
std::thread thread_a([&] { first = worker(1); });
std::thread thread_b([&] { second = worker(2); });
thread_a.join();
thread_b.join();

require(max_active == 2);
require(*first == 1 and *second == 2);
}

void test_failure_is_shared_and_retryable() {
deep_jit::MemCache<std::string, int> cache;
constexpr int num_threads = 8;
std::barrier start(num_threads);
std::atomic<int> factory_calls = 0;
std::atomic<int> failures = 0;
std::vector<std::thread> threads;

for (int index = 0; index < num_threads; ++index) {
threads.emplace_back([&] {
start.arrive_and_wait();
try {
(void)cache.get_or_create("retry", [&]() -> std::shared_ptr<int> {
++factory_calls;
std::this_thread::sleep_for(std::chrono::milliseconds(20));
throw std::runtime_error("expected failure");
});
} catch (const std::runtime_error&) {
++failures;
}
});
}
for (auto& thread : threads)
thread.join();

require(factory_calls == 1);
require(failures == num_threads);
const auto recovered = cache.get_or_create("retry", [&] {
++factory_calls;
return std::make_shared<int>(41);
});
require(factory_calls == 2 and *recovered == 41);
}

} // namespace

int main() {
test_same_key_single_flight();
test_different_keys_remain_parallel();
test_failure_is_shared_and_retryable();
}
'''


def main() -> None:
compiler = os.environ.get('CXX', 'c++')
with tempfile.TemporaryDirectory(prefix='deep-jit-memory-cache-') as directory:
directory = Path(directory)
source = directory / 'test.cpp'
executable = directory / 'test'
source.write_text(CPP_SOURCE, encoding='utf-8')
subprocess.run(
[
compiler,
'-std=c++20',
'-O2',
'-pthread',
'-Wall',
'-Wextra',
'-Werror',
'-I',
str(ROOT / 'include'),
str(source),
'-o',
str(executable),
],
check=True,
timeout=60,
)
subprocess.run([str(executable)], check=True, timeout=30)


if __name__ == '__main__':
main()