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
2 changes: 1 addition & 1 deletion CLAUDE.md
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,7 @@ Unexpected `**kwargs` for any mode trigger a `UserWarning`.
| ---------------------- | ------------------------------------------------------------------------------------------------------- |
| `src/pdex/__init__.py` | `pdex()` entry point and full pipeline logic |
| `src/pdex/_math.py` | Numba JIT-compiled `fold_change()`, `percent_change()`, and `mwu()` wrappers; `pseudobulk()` dispatcher |
| `src/pdex/_utils.py` | `set_numba_threadpool()` — sets Numba thread count before JIT warmup; `_detect_is_log1p()` heuristic |
| `src/pdex/_utils.py` | `set_numba_threadpool()` — sets Numba thread count before JIT warmup; `_available_cpus()` — affinity-aware CPU count (respects cgroup/SLURM limits); `_detect_is_log1p()` heuristic |

### Performance Design

Expand Down
2 changes: 1 addition & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
@@ -1,6 +1,6 @@
[project]
name = "pdex"
version = "0.2.3"
version = "0.2.4"
description = "Parallel differential expression for single-cell perturbation sequencing"
readme = "README.md"
authors = [{ name = "noam teyssier", email = "noam.teyssier@arcinstitute.org" }]
Expand Down
27 changes: 20 additions & 7 deletions src/pdex/_utils.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,5 @@
import logging
import multiprocessing as mp
import os
Comment thread
LeonHafner marked this conversation as resolved.

import numba
Expand All @@ -8,16 +9,28 @@
log = logging.getLogger(__name__)


def _available_cpus() -> int:
"""Return the number of CPUs the current process is allowed to use.

Uses ``os.sched_getaffinity`` on Linux so SLURM/cgroup/taskset limits are
respected; falls back to ``multiprocessing.cpu_count`` on macOS/Windows where
that API is unavailable (those platforms typically run locally without cgroup
caps). This mirrors how Numba derives ``NUMBA_NUM_THREADS``, so the value never
exceeds Numba's cap — unlike ``os.cpu_count()``, which reports every physical
CPU even when affinity or a cgroup restricts the process to far fewer, causing
``numba.set_num_threads`` to raise.
"""
try:
return len(os.sched_getaffinity(0))
except AttributeError:
return mp.cpu_count()
Comment thread
LeonHafner marked this conversation as resolved.


def set_numba_threadpool(threads: int = 0):
available_threads = os.cpu_count()
if available_threads is None:
available_threads = 1
available_threads = _available_cpus()

if threads == 0:
if not available_threads:
threads = 1
else:
threads = available_threads
threads = available_threads
else:
threads = min(threads, available_threads)

Expand Down
9 changes: 3 additions & 6 deletions tests/test_utils.py
Original file line number Diff line number Diff line change
@@ -1,21 +1,18 @@
"""Tests for pdex._utils (set_numba_threadpool)."""

import os

import numba

from pdex._utils import set_numba_threadpool
from pdex._utils import _available_cpus, set_numba_threadpool


class TestSetNumbaThreadpool:
def test_explicit_thread_count(self):
set_numba_threadpool(4)
assert numba.get_num_threads() == 4

def test_zero_uses_all_cpus(self):
def test_zero_uses_all_available_cpus(self):
set_numba_threadpool(0)
expected = os.cpu_count() or 1
assert numba.get_num_threads() == expected
assert numba.get_num_threads() == _available_cpus()

def test_single_thread(self):
set_numba_threadpool(1)
Expand Down
Loading