Skip to content

Commit 278e877

Browse files
committed
Replace g_launch_mutex with GlobalLock
Improve safety via static typing techniques. Signed-off-by: Greg Bonik <gbonik@nvidia.com>
1 parent ced90c9 commit 278e877

10 files changed

Lines changed: 374 additions & 262 deletions

File tree

‎cext/CMakeLists.txt‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -126,7 +126,7 @@ endfunction()
126126

127127

128128
add_test_executable(test_stream_buffer
129-
test/test_stream_buffer.cpp cuda_loader.cpp cuda_helper.cpp memory.cpp)
129+
test/test_stream_buffer.cpp cuda_loader.cpp cuda_helper.cpp memory.cpp py.cpp)
130130

131131
add_test_executable(test_hash_map test/test_hash_map.cpp memory.cpp)
132132

‎cext/compiled_host.cpp‎

Lines changed: 8 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -188,15 +188,16 @@ int CompiledHostProgram_init(PyObject* self, PyObject* args, PyObject* kwargs) {
188188

189189
PyObject* invoke_host_entry(
190190
CompiledHostProgram& program,
191-
void** arguments) {
191+
void** arguments,
192+
GlobalLock& lock) {
192193
int32_t result = program.executable.entry(arguments, &program.runtime);
193194
if (result < 0) {
194195
if (PyErr_Occurred()) return nullptr;
195196
raise(PyExc_RuntimeError, "compiled host code failed with status ", result);
196197
return nullptr;
197198
}
198199
if (result != CUDA_SUCCESS) {
199-
Result<const DriverApi*> driver_result = get_driver_api();
200+
Result<const DriverApi*> driver_result = get_driver_api(lock);
200201
if (!driver_result.is_ok()) return nullptr;
201202
const DriverApi* driver = *driver_result;
202203
raise(PyExc_RuntimeError, "cuda error occurred: ",
@@ -212,6 +213,7 @@ PyObject* CompiledHostProgram_invoke(PyObject* self, PyObject* argument_addresse
212213
raise(PyExc_TypeError, "compiled host argument addresses must be a tuple");
213214
return nullptr;
214215
}
216+
GlobalLock lock;
215217
Py_ssize_t count = PyTuple_GET_SIZE(argument_addresses);
216218
Vec<void*> arguments;
217219
arguments.reserve(count);
@@ -221,7 +223,7 @@ PyObject* CompiledHostProgram_invoke(PyObject* self, PyObject* argument_addresse
221223
if (PyErr_Occurred()) return nullptr;
222224
arguments.push_back(address);
223225
}
224-
return compiled_host_program_invoke(self, arguments.data());
226+
return compiled_host_program_invoke(self, arguments.data(), lock);
225227
}
226228

227229

@@ -286,7 +288,8 @@ bool compiled_host_program_check(PyObject* object) {
286288
}
287289

288290

289-
PyObject* compiled_host_program_invoke(PyObject* program_object, void** arguments) {
291+
PyObject* compiled_host_program_invoke(PyObject* program_object, void** arguments,
292+
GlobalLock& lock) {
290293
if (!compiled_host_program_check(program_object)) {
291294
raise(
292295
PyExc_TypeError,
@@ -295,7 +298,7 @@ PyObject* compiled_host_program_invoke(PyObject* program_object, void** argument
295298
return nullptr;
296299
}
297300
CompiledHostProgram& program = py_unwrap<CompiledHostProgram>(program_object);
298-
return invoke_host_entry(program, arguments);
301+
return invoke_host_entry(program, arguments, lock);
299302
}
300303

301304

‎cext/compiled_host.h‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -8,5 +8,5 @@
88

99

1010
bool compiled_host_program_check(PyObject* object);
11-
PyObject* compiled_host_program_invoke(PyObject* program, void** arguments);
11+
PyObject* compiled_host_program_invoke(PyObject* program, void** arguments, GlobalLock& lock);
1212
Status compiled_host_init(PyObject* module);

‎cext/cuda_helper.cpp‎

Lines changed: 18 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -37,7 +37,8 @@ PyObject* get_max_grid_size(PyObject *self, PyObject *args) {
3737
if (!PyArg_ParseTuple(args, "i", &device_id))
3838
return nullptr;
3939

40-
Result<const DriverApi*> driver = get_driver_api();
40+
GlobalLock lock;
41+
Result<const DriverApi*> driver = get_driver_api(lock);
4142
if (!driver.is_ok()) return nullptr;
4243

4344
CUdevice dev;
@@ -123,7 +124,8 @@ PyObject* get_compute_capability(PyObject *self, PyObject *args) {
123124
int device_id = 0;
124125
if (!PyArg_ParseTuple(args, "|i", &device_id)) return nullptr;
125126

126-
Result<const DriverApi*> driver_result = get_driver_api();
127+
GlobalLock lock;
128+
Result<const DriverApi*> driver_result = get_driver_api(lock);
127129
if (!driver_result.is_ok()) return nullptr;
128130

129131
Result<ComputeCapability> computeCapability =
@@ -135,7 +137,8 @@ PyObject* get_compute_capability(PyObject *self, PyObject *args) {
135137
PyObject* get_driver_version(PyObject *self, PyObject *Py_UNUSED(ignored)) {
136138
int major, minor;
137139

138-
Result<const DriverApi*> driver_result = get_driver_api();
140+
GlobalLock lock;
141+
Result<const DriverApi*> driver_result = get_driver_api(lock);
139142
if (!driver_result.is_ok()) return nullptr;
140143
const DriverApi* d = *driver_result;
141144

@@ -152,7 +155,8 @@ PyObject* get_driver_version(PyObject *self, PyObject *Py_UNUSED(ignored)) {
152155
// ========== Context helpers ==========
153156

154157
PyObject* synchronize_context(PyObject* self, PyObject* Py_UNUSED(ignored)) {
155-
Result<const DriverApi*> driver_result = get_driver_api();
158+
GlobalLock lock;
159+
Result<const DriverApi*> driver_result = get_driver_api(lock);
156160
if (!driver_result.is_ok()) return nullptr;
157161
const DriverApi* d = *driver_result;
158162

@@ -167,7 +171,8 @@ PyObject* synchronize_context(PyObject* self, PyObject* Py_UNUSED(ignored)) {
167171
// ========== Stream helpers ==========
168172

169173
PyObject* create_stream(PyObject* self, PyObject* Py_UNUSED(ignored)) {
170-
Result<const DriverApi*> driver_result = get_driver_api();
174+
GlobalLock lock;
175+
Result<const DriverApi*> driver_result = get_driver_api(lock);
171176
if (!driver_result.is_ok()) return nullptr;
172177
const DriverApi* d = *driver_result;
173178

@@ -181,10 +186,11 @@ PyObject* create_stream(PyObject* self, PyObject* Py_UNUSED(ignored)) {
181186
}
182187

183188
PyObject* destroy_stream(PyObject* self, PyObject* arg) {
189+
GlobalLock lock;
184190
CUstream stream = static_cast<CUstream>(PyLong_AsVoidPtr(arg));
185191
if (PyErr_Occurred()) return nullptr;
186192

187-
Result<const DriverApi*> driver_result = get_driver_api();
193+
Result<const DriverApi*> driver_result = get_driver_api(lock);
188194
if (!driver_result.is_ok()) return nullptr;
189195
const DriverApi* d = *driver_result;
190196

@@ -226,15 +232,14 @@ static CUresult shim_cuLaunchKernelEx(
226232
}
227233

228234
static PyObject* spy_on_cuLaunchKernel_begin(PyObject* self, PyObject* arg) {
229-
#ifdef Py_GIL_DISABLED
230-
PyCriticalSectionGuard guard(&g_spy_mutex);
231-
#endif
235+
GlobalLock lock;
236+
232237
if (g_real_cuLaunchKernelEx) {
233238
raise(PyExc_RuntimeError, "Already spying");
234239
return nullptr;
235240
}
236241

237-
Result<const DriverApi*> driver_result = get_driver_api();
242+
Result<const DriverApi*> driver_result = get_driver_api(lock);
238243
if (!driver_result.is_ok()) return nullptr;
239244

240245
DriverApi* api = const_cast<DriverApi*>(*driver_result);
@@ -245,15 +250,14 @@ static PyObject* spy_on_cuLaunchKernel_begin(PyObject* self, PyObject* arg) {
245250
}
246251

247252
static PyObject* spy_on_cuLaunchKernel_end(PyObject* self, PyObject* arg) {
248-
#ifdef Py_GIL_DISABLED
249-
PyCriticalSectionGuard guard(&g_spy_mutex);
250-
#endif
253+
GlobalLock lock;
254+
251255
if (!g_real_cuLaunchKernelEx) {
252256
raise(PyExc_RuntimeError, "Not spying");
253257
return nullptr;
254258
}
255259

256-
Result<const DriverApi*> driver_result = get_driver_api();
260+
Result<const DriverApi*> driver_result = get_driver_api(lock);
257261
if (!driver_result.is_ok()) return nullptr;
258262

259263
DriverApi* api = const_cast<DriverApi*>(*driver_result);

‎cext/cuda_loader.cpp‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -73,7 +73,7 @@ static Result<cuGetProcAddress_v2_t> get_cuGetProcAddress_from_python() {
7373

7474
static constexpr int MIN_DRIVER_VERSION = 13000;
7575

76-
Result<const DriverApi*> get_driver_api() {
76+
Result<const DriverApi*> get_driver_api(GlobalLock& lock) {
7777
static bool initialized;
7878
static DriverApi instance;
7979
if (!initialized) {

‎cext/cuda_loader.h‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -82,7 +82,7 @@ typedef CUresult (*cuGetProcAddress_v2_t)
8282

8383
Status driver_api_init(DriverApi* driver_api, cuGetProcAddress_v2_t _cuGetProcAddress);
8484

85-
Result<const DriverApi*> get_driver_api();
85+
Result<const DriverApi*> get_driver_api(GlobalLock& lock);
8686

8787

8888
class CudaContextGuard {

‎cext/py.cpp‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -47,3 +47,7 @@ void log_python_error(const char* filename, int line, const char* level, SavedEx
4747
PyErr_SetExcInfo(old_excinfo_type, old_excinfo_value, old_excinfo_tb);
4848
}
4949

50+
#ifdef Py_GIL_DISABLED
51+
PyMutex GlobalLock::mutex_ = {0};
52+
#endif
53+

‎cext/py.h‎

Lines changed: 72 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -601,3 +601,75 @@ class PyCriticalSectionGuard {
601601
PyCriticalSection _py_cs;
602602
};
603603
#endif
604+
605+
606+
// In a GIL build, asserts that we are holding the GIL.
607+
// In a free-threaded build, enters the global critical section.
608+
//
609+
// Passing a reference to an object of this class to a function serves as a "proof"
610+
// that the lock is being held.
611+
//
612+
// For example:
613+
//
614+
// // Helper function that mutates the global state in a thread-unsafe way.
615+
// // Takes a `GlobalLock&` reference to indicate that either the GIL or the global
616+
// // critical section must be held.
617+
// static int next_number(GlobalLock&) {
618+
// static int counter = 0;
619+
// return counter++;
620+
// }
621+
//
622+
// // Method implementation
623+
// static PyObject* foo(PyObject* self, PyObject* args) {
624+
// // In a GIL build, we know we're holding the GIL since this is a method, and
625+
// // thus it's safe to instantiate a GlobalLock.
626+
// //
627+
// // In a free-threaded build, this will enter the global critical section
628+
// // to emulate the GIL.
629+
// GlobalLock lock;
630+
//
631+
// // Call the thread-unsafe helper.
632+
// int number = next_number(lock);
633+
//
634+
// // ...
635+
// }
636+
class GlobalLock {
637+
public:
638+
GlobalLock()
639+
#ifdef Py_GIL_DISABLED
640+
: guard_(&mutex_)
641+
#endif
642+
{
643+
#ifndef Py_GIL_DISABLED
644+
// Will crash with a fatal error if we aren't holding the GIL
645+
PyThreadState_Get();
646+
#endif // Py_GIL_DISABLED
647+
}
648+
649+
// Use this constructor when you are holding a GILGuard
650+
// but you don't necessarily have an attached thread state.
651+
explicit GlobalLock(GILGuard& gil_guard)
652+
#ifdef Py_GIL_DISABLED
653+
: guard_(&mutex_)
654+
#endif
655+
{}
656+
657+
GlobalLock(const GlobalLock&) = delete;
658+
void operator==(const GlobalLock&) = delete;
659+
660+
private:
661+
#ifdef Py_GIL_DISABLED
662+
PyCriticalSectionGuard guard_;
663+
static PyMutex mutex_;
664+
#endif
665+
};
666+
667+
668+
template <typename T>
669+
class ProtectedByGlobalLock {
670+
static_assert(std::is_trivially_destructible_v<T>);
671+
T object_;
672+
public:
673+
T& get(GlobalLock&) { return object_; }
674+
};
675+

0 commit comments

Comments
 (0)