Skip to content

Lock contention inside _Py_Specialize_LoadGlobal under free threading #152075

Description

@hawkinsp

Bug report

Bug description:

While working on improving the run time of the JAX test suite with high thread concurrency under free threading, a lock contention profile pointed me to high lock contention in _Py_Specialize_LoadGlobal

Here's a synthetic benchmark that I generated with the assistance of AI.

#!/usr/bin/env python3
"""Benchmark to demonstrate LOAD_GLOBAL specialization contention in free-threaded Python.

This benchmark spawns multiple threads that frequently read global variables and
builtins (triggering bytecode specialization). An optional background thread
periodically modifies the globals dictionary.
"""

import argparse
import sys
import threading
import time

# Globals to access
G1 = 1
G2 = 2
G3 = 3
G4 = 4
G5 = 5
G6 = 6
G7 = 7
G8 = 8


def reader_worker(num_iters, stop_event):
    # Access globals in a loop to trigger LOAD_GLOBAL
    for _ in range(num_iters):
        if stop_event.is_set():
            break
        a = G1
        b = G2
        c = G3
        d = G4
        e = G5
        f = G6
        g = G7
        h = G8

        # Access builtins
        i = len
        j = sum
        k = abs

        # Trivial use of the loaded values
        _ = a + b + c + d + e + f + g + h


def invalidator_worker(stop_event, interval):
    while not stop_event.is_set():
        # Modify globals to increment keys version and invalidate caches
        globals()["_temp_key"] = 1
        del globals()["_temp_key"]
        if interval > 0:
            time.sleep(interval)


def run_benchmark(num_threads, iters_per_thread, invalidate_interval):
    stop_event = threading.Event()

    # Start invalidator if requested
    invalidator_thread = None
    if invalidate_interval is not None:
        invalidator_thread = threading.Thread(
            target=invalidator_worker,
            args=(stop_event, invalidate_interval),
            daemon=True,
        )
        invalidator_thread.start()

    # Start readers
    threads = []
    start_time = time.perf_counter()

    for _ in range(num_threads):
        t = threading.Thread(target=reader_worker, args=(iters_per_thread, stop_event))
        threads.append(t)
        t.start()

    for t in threads:
        t.join()

    end_time = time.perf_counter()

    # Stop invalidator
    stop_event.set()
    if invalidator_thread:
        invalidator_thread.join()

    return end_time - start_time


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument(
        "--threads", type=int, default=64, help="Number of reader threads (default: 64)"
    )
    parser.add_argument(
        "--iters",
        type=int,
        default=10000000,
        help="Iterations per reader thread (default: 10,000,000)",
    )
    parser.add_argument(
        "--interval",
        type=float,
        default=0.0001,
        help="Invalidation interval in seconds. Use negative value to disable invalidation. (default: 0.0001)",
    )

    args = parser.parse_args()

    invalidate_interval = args.interval if args.interval >= 0 else None

    print(f"Python Executable: {sys.executable}")
    print(f"Python Version: {sys.version}")
    print(f"GIL Enabled: {getattr(sys, '_is_gil_enabled', lambda: 'Unknown')()}")
    print(
        f"Configuration: threads={args.threads}, iters={args.iters:,}, "
        f"invalidate_interval={invalidate_interval}"
    )
    print("Running benchmark...")

    duration = run_benchmark(args.threads, args.iters, invalidate_interval)

    print(f"Finished in {duration:.4f} seconds")


if __name__ == "__main__":
    main()

and I'm running this benchmark on a cloud VM with these characteristics:

Architecture:                x86_64
...
CPU(s):                      128
  On-line CPU(s) list:       0-127
Vendor ID:                   AuthenticAMD
  Model name:                AMD EPYC 7B13
# With no invalidation:
$ python benchmark_contention.py  --interval -1
...
Python Version: 3.15.0b2+dev free-threading build (heads/3.15:ba0cae13cea, Jun 24 2026, 02:50:32) [GCC 15.2.0]
GIL Enabled: False
Configuration: threads=64, iters=10,000,000, invalidate_interval=None
Running benchmark...
Finished in 2.1970 seconds

# With invalidation:
$ python benchmark_contention.py
Python Version: 3.15.0b2+dev free-threading build (heads/3.15:ba0cae13cea, Jun 24 2026, 02:50:32) [GCC 15.2.0]
GIL Enabled: False
Configuration: threads=64, iters=10,000,000, invalidate_interval=0.0001
Running benchmark...
Finished in 23.2520 seconds

However I thought of trying this patch, which immediately abandons the attempt to specialize if acquiring globals/builtins critical section would block:

diff --git a/Include/critical_section.h b/Include/critical_section.h
index 732bfab7ecf..b90bc7fda2a 100644
--- a/Include/critical_section.h
+++ b/Include/critical_section.h
@@ -68,6 +68,9 @@ PyCriticalSection2_Begin(PyCriticalSection2 *c, PyObject *a, PyObject *b);
 PyAPI_FUNC(void)
 PyCriticalSection2_End(PyCriticalSection2 *c);
 
+PyAPI_FUNC(int)
+PyCriticalSection2_TryBegin(PyCriticalSection2 *c, PyObject *a, PyObject *b);
+
 // These are definitions for the stable ABI. For GIL-ful builds they're
 // conditionally redefined as no-ops in cpython/critical_section.h.
 
diff --git a/Include/internal/pycore_critical_section.h b/Include/internal/pycore_critical_section.h
index 51d99d74ca1..d32fc2fdf67 100644
--- a/Include/internal/pycore_critical_section.h
+++ b/Include/internal/pycore_critical_section.h
@@ -193,6 +193,56 @@ _PyCriticalSection2_Begin(PyThreadState *tstate, PyCriticalSection2 *c, PyObject
     _PyCriticalSection2_BeginMutex(tstate, c, &a->ob_mutex, &b->ob_mutex);
 }
 
+static inline int
+_PyCriticalSection_TryBeginMutex(PyThreadState *tstate, PyCriticalSection *c, PyMutex *m)
+{
+    if (PyMutex_LockFast(m)) {
+        c->_cs_mutex = m;
+        c->_cs_prev = tstate->critical_section;
+        tstate->critical_section = (uintptr_t)c;
+        return 1;
+    }
+    return 0;
+}
+
+static inline int
+_PyCriticalSection2_TryBeginMutex(PyThreadState *tstate, PyCriticalSection2 *c, PyMutex *m1, PyMutex *m2)
+{
+    if (m1 == m2) {
+        c->_cs_mutex2 = NULL;
+        return _PyCriticalSection_TryBeginMutex(tstate, &c->_cs_base, m1);
+    }
+
+    if ((uintptr_t)m2 < (uintptr_t)m1) {
+        PyMutex *tmp = m1;
+        m1 = m2;
+        m2 = tmp;
+    }
+
+    if (PyMutex_LockFast(m1)) {
+        if (PyMutex_LockFast(m2)) {
+            c->_cs_base._cs_mutex = m1;
+            c->_cs_mutex2 = m2;
+            c->_cs_base._cs_prev = tstate->critical_section;
+
+            uintptr_t p = (uintptr_t)c | _Py_CRITICAL_SECTION_TWO_MUTEXES;
+            tstate->critical_section = p;
+            return 1;
+        }
+        else {
+            PyMutex_Unlock(m1);
+            return 0;
+        }
+    }
+    return 0;
+}
+
+static inline int
+_PyCriticalSection2_TryBegin(PyThreadState *tstate, PyCriticalSection2 *c, PyObject *a, PyObject *b)
+{
+    return _PyCriticalSection2_TryBeginMutex(tstate, c, &a->ob_mutex, &b->ob_mutex);
+}
+
 static inline void
 _PyCriticalSection2_End(PyThreadState *tstate, PyCriticalSection2 *c)
 {
diff --git a/Python/critical_section.c b/Python/critical_section.c
index dbee6f236a7..2d81425cee2 100644
--- a/Python/critical_section.c
+++ b/Python/critical_section.c
@@ -217,3 +217,14 @@ PyCriticalSection2_End(PyCriticalSection2 *c)
     _PyCriticalSection2_End(_PyThreadState_GET(), c);
 #endif
 }
+
+#undef PyCriticalSection2_TryBegin
+int
+PyCriticalSection2_TryBegin(PyCriticalSection2 *c, PyObject *a, PyObject *b)
+{
+#ifdef Py_GIL_DISABLED
+    return _PyCriticalSection2_TryBegin(_PyThreadState_GET(), c, a, b);
+#else
+    return 1;
+#endif
+}
diff --git a/Python/specialize.c b/Python/specialize.c
index 459e69de570..53a126d9f48 100644
--- a/Python/specialize.c
+++ b/Python/specialize.c
@@ -1447,9 +1447,18 @@ _Py_Specialize_LoadGlobal(
     PyObject *globals, PyObject *builtins,
     _Py_CODEUNIT *instr, PyObject *name)
 {
-    Py_BEGIN_CRITICAL_SECTION2(globals, builtins);
+#ifdef Py_GIL_DISABLED
+    PyCriticalSection2 cs;
+    if (PyCriticalSection2_TryBegin(&cs, globals, builtins)) {
+        specialize_load_global_lock_held(globals, builtins, instr, name);
+        PyCriticalSection2_End(&cs);
+    }
+    else {
+        unspecialize(instr);
+    }
+#else
     specialize_load_global_lock_held(globals, builtins, instr, name);
-    Py_END_CRITICAL_SECTION2();
+#endif
 }
 
 static int

and with that change I get these timings:

$ python benchmark_contention.py  --interval -1
Python Version: 3.15.0b2+dev free-threading build (heads/3.15-dirty:ba0cae13cea, Jun 24 2026, 01:28:11) [GCC 15.2.0]
GIL Enabled: False
Configuration: threads=64, iters=10,000,000, invalidate_interval=None
Running benchmark...
Finished in 1.9692 seconds

$ python benchmark_contention.py
Python Version: 3.15.0b2+dev free-threading build (heads/3.15-dirty:ba0cae13cea, Jun 24 2026, 01:28:11) [GCC 15.2.0]
GIL Enabled: False
Configuration: threads=64, iters=10,000,000, invalidate_interval=0.0001
Running benchmark...
Finished in 11.0028 seconds

Might we land something like that?

CPython versions tested on:

3.15

Operating systems tested on:

Linux

Linked PRs

Activity

  1. ZeroIntensity commented on Jun 24, 2026

    @ZeroIntensity
    Member

    We probably want to find a way to make this lockless instead. The index lookups don't seem too hard as long as we can atomically validate them later, but making dict versions thread-safe looks a little more complicated.

    cc @mpage @colesbury

  2. mpage commented on Jul 2, 2026

    @mpage
    Contributor

    @hawkinsp - Thanks for the report. Your approach looks good to me. Do you want to put up a PR? I'm happy to do so if you're not interested.

    @ZeroIntensity - I don't think we need to make this lock free. Specializations of LOAD_GLOBAL should be pretty stable. Under that assumption threads that attempt to specialize the same LOAD_GLOBAL should make forward progress and ultimately quiesce.

  3. hawkinsp commented on Jul 14, 2026

    @hawkinsp
    ContributorAuthor

    Sorry for the slow response, I was on vacation. Sure, I'll send a PR.

  4. added 3 commits that reference this issue on Jul 14, 2026
  5. added a commit that references this issue on Jul 28, 2026
  6. added a commit that references this issue on Jul 31, 2026
  7. added a commit that references this issue on Aug 14, 2026
  8. cakedev0 commented on Oct 5, 2026

    @cakedev0

    Just a FYI:

    I hit this with joblib's threading backend, which scikit-learn uses for random forests: many threads start running the same code at the same time, and with thread-local bytecode each of them specializes its own copy. A simple reproducer:

    import sys
    import time
    
    from joblib import Parallel, delayed
    
    x = None
    exec("def f():\n" + "    x\n" * 200_000)  # 200k LOAD_GLOBALs, like a big chunk of library code
    
    
    def task():
        f()
        f()
    
    
    print(sys.version, "| tlbc:", sys._xoptions.get("tlbc", "default"))
    for n_jobs in [1, 2, 4, 8, 14]:
        times = []
        for _ in range(3):  # the same function, 3 Parallel calls
            tic = time.perf_counter()
            Parallel(n_jobs=n_jobs, backend="threading")(delayed(task)() for _ in range(n_jobs))
            times.append(time.perf_counter() - tic)
        print(f"n_jobs={n_jobs:2d}: " + ", ".join(f"{1e3 * t:6.1f} ms" for t in times))

    On a 14-core laptop, 3.14.6t:

    n_jobs= 1:    6.4 ms,    0.7 ms,    0.7 ms
    n_jobs= 2:   34.3 ms,   13.7 ms,   13.0 ms
    n_jobs= 4:   70.5 ms,   14.2 ms,   12.5 ms
    n_jobs= 8:   86.8 ms,   12.6 ms,   13.2 ms
    n_jobs=14:  119.1 ms,   16.2 ms,   17.1 ms
    

    The first call gets slower with the number of threads, the next ones only pay joblib's own overhead (~13 ms).

    On a bigger machine (2 x 86-core Xeon 6787P, 344 hardware threads, 3.14.7t), the first call falls off a cliff above 16 threads, even with 10x fewer loads (" x\n" * 20_000):

    n_jobs=  1:      1.1 ms,    0.2 ms,    0.1 ms
    n_jobs= 16:     77.9 ms,   23.3 ms,   19.4 ms
    n_jobs= 43:  86562.1 ms,   53.8 ms,   42.7 ms
    n_jobs= 86: 139015.8 ms,  159.5 ms,   79.6 ms
    n_jobs=172: 284912.6 ms,  236.2 ms,  216.7 ms
    

    A first call of almost 5 minutes, instead of ~0.2 s. With 2k loads, the first calls take 9.0 s / 14.7 s / 28.8 s at 43 / 86 / 172 threads.

    The reproducer is a worst case, with 20k global loads in a row running in lockstep; real code collides less
    but still: on the same machine fitting a scikit-learn ExtraTreesClassifier(n_estimators=688, n_jobs=-1) (joblib's threading backend, 344 threads, fresh process) on make_classification(n_samples=100_000, n_features=20) takes 11.8 s, vs 4.2-4.4 s with PYTHON_TLBC=0.

  9. hawkinsp commented on Oct 5, 2026

    @hawkinsp
    ContributorAuthor

    I wonder if we might consider backporting the fix to 3.14 and/or 3.15?

  10. ZeroIntensity commented on Oct 5, 2026

    @ZeroIntensity
    Member

    Backporting seems reasonable to me, but the decision should be up to @hugovk.

  11. added a commit that references this issue on Oct 6, 2026
  12. cakedev0 commented on Oct 6, 2026

    @cakedev0

    Thanks for the quick answers/actions, I must say backporting to 3.14 & 3.15 should make our life easier, thanks!

    Note that the cliff I observed on the reproducer above (~80ms at 16 threads to 86s at 43), is a product of two thing:

  13. hugovk commented on Oct 9, 2026

    @hugovk
    Member

    This can wait until 3.15.1.

  14. added a commit that references this issue on Oct 10, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions