diff --git a/include/deep_jit/cache/memory.hpp b/include/deep_jit/cache/memory.hpp index 70a987f..a045352 100644 --- a/include/deep_jit/cache/memory.hpp +++ b/include/deep_jit/cache/memory.hpp @@ -1,6 +1,9 @@ #pragma once +#include +#include #include +#include #include namespace deep_jit { @@ -8,16 +11,56 @@ namespace deep_jit { // In-memory cache keyed by value type. template class MemCache { + using ValuePtr = std::shared_ptr; + + std::unordered_map> pending; + mutable std::mutex mutex; + public: - std::unordered_map> cache; + // Kept public for source compatibility. Concurrent callers must use + // get_or_create() rather than accessing the container directly. + std::unordered_map cache; template - std::shared_ptr 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 future; + std::shared_ptr> 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>(); + 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); + } } }; diff --git a/include/deep_jit/runtime/runtime.hpp b/include/deep_jit/runtime/runtime.hpp index 8ca176c..ea83112 100644 --- a/include/deep_jit/runtime/runtime.hpp +++ b/include/deep_jit/runtime/runtime.hpp @@ -56,6 +56,10 @@ class Runtime { std::shared_ptr 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); }); diff --git a/tests/test_cuda.py b/tests/test_cuda.py index 360594a..9d4318e 100644 --- a/tests/test_cuda.py +++ b/tests/test_cuda.py @@ -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 @@ -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) diff --git a/tests/test_cuda_proj/main.cpp b/tests/test_cuda_proj/main.cpp index 4baf23e..34818ec 100644 --- a/tests/test_cuda_proj/main.cpp +++ b/tests/test_cuda_proj/main.cpp @@ -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( + 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::Config( @@ -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); } diff --git a/tests/test_memory_cache.py b/tests/test_memory_cache.py new file mode 100644 index 0000000..88c0749 --- /dev/null +++ b/tests/test_memory_cache.py @@ -0,0 +1,159 @@ +import os +import subprocess +import tempfile +from pathlib import Path + + +ROOT = Path(__file__).resolve().parents[1] + + +CPP_SOURCE = r''' +#include +#include +#include +#include +#include +#include +#include +#include +#include + +#include + +namespace { + +void require(bool condition) { + if (not condition) + std::abort(); +} + +void test_same_key_single_flight() { + deep_jit::MemCache cache; + constexpr int num_threads = 16; + std::barrier start(num_threads); + std::atomic factory_calls = 0; + std::vector> values(num_threads); + std::vector 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(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 cache; + std::barrier start(2); + std::atomic active = 0; + std::atomic 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(key); + }); + }; + + std::shared_ptr first; + std::shared_ptr 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 cache; + constexpr int num_threads = 8; + std::barrier start(num_threads); + std::atomic factory_calls = 0; + std::atomic failures = 0; + std::vector 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 { + ++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(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()