sweep-solver 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- geophyai/__init__.py +4 -0
- sweep/_C.py +35 -0
- sweep/__init__.py +182 -0
- sweep/_jit.py +306 -0
- sweep/backend/__init__.py +5 -0
- sweep/backend/jax/__init__.py +15 -0
- sweep/backend/jax/cuda.py +17 -0
- sweep/backend/torch/__init__.py +16 -0
- sweep/backend/torch/binding.py +53 -0
- sweep/backend/torch/cuda.py +17 -0
- sweep/cli.py +165 -0
- sweep/csrc/CMakeLists.txt +10 -0
- sweep/csrc/bindings/bindings_utils.h +48 -0
- sweep/csrc/bindings/module.cpp +254 -0
- sweep/csrc/cpu/common/cpu_engine.cpp +1181 -0
- sweep/csrc/cpu/common/cpu_engine.h +36 -0
- sweep/csrc/cpu/cpu_binding.cpp +136 -0
- sweep/csrc/cpu/cpu_binding.h +20 -0
- sweep/csrc/cpu/cpu_binding_stub.cpp +54 -0
- sweep/csrc/cpu/equations/acoustic2d/acoustic2d_cpu.cpp +1882 -0
- sweep/csrc/cpu/equations/acoustic2d/acoustic2d_cpu.h +13 -0
- sweep/csrc/cpu/equations/acoustic3d/acoustic3d_cpu.cpp +1645 -0
- sweep/csrc/cpu/equations/acoustic3d/acoustic3d_cpu.h +13 -0
- sweep/csrc/cpu/equations/acoustic_lsrtm2d/acoustic_lsrtm2d_cpu.cpp +1787 -0
- sweep/csrc/cpu/equations/acoustic_lsrtm2d/acoustic_lsrtm2d_cpu.h +13 -0
- sweep/csrc/cpu/equations/acoustic_lsrtm3d/acoustic_lsrtm3d_cpu.cpp +1741 -0
- sweep/csrc/cpu/equations/acoustic_lsrtm3d/acoustic_lsrtm3d_cpu.h +13 -0
- sweep/csrc/cpu/equations/acoustic_vrz2d/acoustic_vrz2d_cpu.cpp +1549 -0
- sweep/csrc/cpu/equations/acoustic_vrz2d/acoustic_vrz2d_cpu.h +13 -0
- sweep/csrc/cpu/equations/acoustic_vrz3d/acoustic_vrz3d_cpu.cpp +1587 -0
- sweep/csrc/cpu/equations/acoustic_vrz3d/acoustic_vrz3d_cpu.h +13 -0
- sweep/csrc/cpu/equations/das2d/das2d_cpu.cpp +752 -0
- sweep/csrc/cpu/equations/das2d/das2d_cpu.h +13 -0
- sweep/csrc/cpu/equations/das3d/das3d_cpu.cpp +938 -0
- sweep/csrc/cpu/equations/das3d/das3d_cpu.h +13 -0
- sweep/csrc/cpu/equations/das_mu2d/das_mu2d_cpu.cpp +1194 -0
- sweep/csrc/cpu/equations/das_mu2d/das_mu2d_cpu.h +13 -0
- sweep/csrc/cpu/equations/das_mu3d/das_mu3d_cpu.cpp +1364 -0
- sweep/csrc/cpu/equations/das_mu3d/das_mu3d_cpu.h +13 -0
- sweep/csrc/cpu/equations/elastic2d/elastic2d_cpu.cpp +826 -0
- sweep/csrc/cpu/equations/elastic2d/elastic2d_cpu.h +13 -0
- sweep/csrc/cpu/equations/elastic3d/elastic3d_cpu.cpp +916 -0
- sweep/csrc/cpu/equations/elastic3d/elastic3d_cpu.h +13 -0
- sweep/csrc/cpu/equations/elastic_tti_sg2d/elastic_tti_sg2d_cpu.cpp +1736 -0
- sweep/csrc/cpu/equations/elastic_tti_sg2d/elastic_tti_sg2d_cpu.h +13 -0
- sweep/csrc/cpu/operators/fd.h +268 -0
- sweep/csrc/cuda/common/acoustic.h +428 -0
- sweep/csrc/cuda/common/acoustic_vrz_fused.cuh +81 -0
- sweep/csrc/cuda/common/boundary/disk_io.cuh +349 -0
- sweep/csrc/cuda/common/boundary/kernels.cuh +193 -0
- sweep/csrc/cuda/common/boundary/runtime.cuh +1590 -0
- sweep/csrc/cuda/common/boundary/saver.cuh +1423 -0
- sweep/csrc/cuda/common/boundary/types.cuh +119 -0
- sweep/csrc/cuda/common/boundary_runtime.cuh +3 -0
- sweep/csrc/cuda/common/boundarysaver.cu +960 -0
- sweep/csrc/cuda/common/boundarysaver.cuh +3 -0
- sweep/csrc/cuda/common/checkpoint_runtime.cuh +347 -0
- sweep/csrc/cuda/common/common.cu +209 -0
- sweep/csrc/cuda/common/common.cuh +55 -0
- sweep/csrc/cuda/common/context.h +96 -0
- sweep/csrc/cuda/common/cudautils.h +147 -0
- sweep/csrc/cuda/common/das.h +403 -0
- sweep/csrc/cuda/common/das_mu.h +543 -0
- sweep/csrc/cuda/common/elastic.h +772 -0
- sweep/csrc/cuda/common/elastic_free_surface.cuh +366 -0
- sweep/csrc/cuda/common/wavetypes.h +3 -0
- sweep/csrc/cuda/equations/acoustic2d/acoustic2d.h +19 -0
- sweep/csrc/cuda/equations/acoustic2d/backward.cu +1156 -0
- sweep/csrc/cuda/equations/acoustic2d/forward.cu +217 -0
- sweep/csrc/cuda/equations/acoustic2d/kernels.cu +170 -0
- sweep/csrc/cuda/equations/acoustic2d/kernels.cuh +454 -0
- sweep/csrc/cuda/equations/acoustic3d/acoustic3d.h +19 -0
- sweep/csrc/cuda/equations/acoustic3d/backward.cu +1284 -0
- sweep/csrc/cuda/equations/acoustic3d/forward.cu +239 -0
- sweep/csrc/cuda/equations/acoustic3d/kernels.cu +194 -0
- sweep/csrc/cuda/equations/acoustic3d/kernels.cuh +547 -0
- sweep/csrc/cuda/equations/acoustic_lsrtm2d/acoustic_lsrtm2d.h +17 -0
- sweep/csrc/cuda/equations/acoustic_lsrtm2d/backward.cu +803 -0
- sweep/csrc/cuda/equations/acoustic_lsrtm2d/forward.cu +219 -0
- sweep/csrc/cuda/equations/acoustic_lsrtm2d/kernels.cu +65 -0
- sweep/csrc/cuda/equations/acoustic_lsrtm2d/kernels.cuh +318 -0
- sweep/csrc/cuda/equations/acoustic_lsrtm3d/acoustic_lsrtm3d.h +17 -0
- sweep/csrc/cuda/equations/acoustic_lsrtm3d/backward.cu +1160 -0
- sweep/csrc/cuda/equations/acoustic_lsrtm3d/forward.cu +223 -0
- sweep/csrc/cuda/equations/acoustic_lsrtm3d/kernels.cu +78 -0
- sweep/csrc/cuda/equations/acoustic_lsrtm3d/kernels.cuh +398 -0
- sweep/csrc/cuda/equations/acoustic_vrz2d/acoustic_vrz2d.h +13 -0
- sweep/csrc/cuda/equations/acoustic_vrz2d/backward.cu +648 -0
- sweep/csrc/cuda/equations/acoustic_vrz2d/forward.cu +201 -0
- sweep/csrc/cuda/equations/acoustic_vrz2d/kernels.cuh +1092 -0
- sweep/csrc/cuda/equations/acoustic_vrz3d/acoustic_vrz3d.h +13 -0
- sweep/csrc/cuda/equations/acoustic_vrz3d/backward.cu +738 -0
- sweep/csrc/cuda/equations/acoustic_vrz3d/forward.cu +206 -0
- sweep/csrc/cuda/equations/acoustic_vrz3d/kernels.cuh +1147 -0
- sweep/csrc/cuda/equations/acoustic_vti_1st_2d/acoustic_vti_1st_2d.h +17 -0
- sweep/csrc/cuda/equations/acoustic_vti_1st_2d/backward.cu +819 -0
- sweep/csrc/cuda/equations/acoustic_vti_1st_2d/forward.cu +353 -0
- sweep/csrc/cuda/equations/acoustic_vti_1st_2d/kernels.cu +4 -0
- sweep/csrc/cuda/equations/acoustic_vti_1st_2d/kernels.cuh +687 -0
- sweep/csrc/cuda/equations/acoustic_vti_1st_3d/acoustic_vti_1st_3d.h +14 -0
- sweep/csrc/cuda/equations/acoustic_vti_1st_3d/backward.cu +780 -0
- sweep/csrc/cuda/equations/acoustic_vti_1st_3d/forward.cu +365 -0
- sweep/csrc/cuda/equations/acoustic_vti_1st_3d/kernels.cu +5 -0
- sweep/csrc/cuda/equations/acoustic_vti_1st_3d/kernels.cuh +777 -0
- sweep/csrc/cuda/equations/das2d/backward.cu +829 -0
- sweep/csrc/cuda/equations/das2d/das2d.h +17 -0
- sweep/csrc/cuda/equations/das2d/forward.cu +255 -0
- sweep/csrc/cuda/equations/das2d/kernels.cuh +673 -0
- sweep/csrc/cuda/equations/das3d/backward.cu +391 -0
- sweep/csrc/cuda/equations/das3d/das3d.h +17 -0
- sweep/csrc/cuda/equations/das3d/forward.cu +179 -0
- sweep/csrc/cuda/equations/das3d/kernels.cuh +639 -0
- sweep/csrc/cuda/equations/das_mu2d/backward.cu +1068 -0
- sweep/csrc/cuda/equations/das_mu2d/das_mu2d.h +17 -0
- sweep/csrc/cuda/equations/das_mu2d/forward.cu +233 -0
- sweep/csrc/cuda/equations/das_mu2d/kernels.cuh +246 -0
- sweep/csrc/cuda/equations/das_mu3d/backward.cu +1171 -0
- sweep/csrc/cuda/equations/das_mu3d/das_mu3d.h +17 -0
- sweep/csrc/cuda/equations/das_mu3d/forward.cu +239 -0
- sweep/csrc/cuda/equations/das_mu3d/kernels.cuh +283 -0
- sweep/csrc/cuda/equations/elastic2d/backward.cu +1485 -0
- sweep/csrc/cuda/equations/elastic2d/elastic2d.h +27 -0
- sweep/csrc/cuda/equations/elastic2d/forward.cu +392 -0
- sweep/csrc/cuda/equations/elastic2d/kernels.cu +0 -0
- sweep/csrc/cuda/equations/elastic2d/kernels.cuh +1853 -0
- sweep/csrc/cuda/equations/elastic3d/backward.cu +1587 -0
- sweep/csrc/cuda/equations/elastic3d/elastic3d.h +24 -0
- sweep/csrc/cuda/equations/elastic3d/forward.cu +466 -0
- sweep/csrc/cuda/equations/elastic3d/kernels.cuh +2658 -0
- sweep/csrc/cuda/equations/elastic_tti_sg2d/backward.cu +784 -0
- sweep/csrc/cuda/equations/elastic_tti_sg2d/elastic_tti_sg2d.h +15 -0
- sweep/csrc/cuda/equations/elastic_tti_sg2d/forward.cu +246 -0
- sweep/csrc/cuda/equations/elastic_tti_sg2d/kernels.cuh +1037 -0
- sweep/csrc/cuda/equations/elastic_tti_sg2d/tensors.h +163 -0
- sweep/csrc/cuda/equations/elastic_vr2d/backward.cu +883 -0
- sweep/csrc/cuda/equations/elastic_vr2d/elastic_vr2d.h +17 -0
- sweep/csrc/cuda/equations/elastic_vr2d/forward.cu +207 -0
- sweep/csrc/cuda/equations/elastic_vr2d/kernels.cuh +1142 -0
- sweep/csrc/cuda/launch/config.h +169 -0
- sweep/csrc/cuda/operators/dim.cuh +12 -0
- sweep/csrc/cuda/operators/gradient.cuh +345 -0
- sweep/csrc/cuda/operators/laplace.cuh +354 -0
- sweep/csrc/cuda/operators/staggered.cuh +540 -0
- sweep/csrc/shared/wavetypes.h +191 -0
- sweep/datasets/__init__.py +115 -0
- sweep/datasets/_benchmarks.py +280 -0
- sweep/datasets/_cache.py +94 -0
- sweep/datasets/_formats.py +264 -0
- sweep/datasets/cli.py +89 -0
- sweep/datasets/marmousi.py +9178 -0
- sweep/datasets/overthrust_2d.py +1595 -0
- sweep/datasets/registry.py +170 -0
- sweep/equations/__init__.py +145 -0
- sweep/equations/_anisotropy_utils.py +89 -0
- sweep/equations/_elastic_step_core.py +234 -0
- sweep/equations/_free_surface.py +324 -0
- sweep/equations/_topography.py +748 -0
- sweep/equations/acoustic.py +175 -0
- sweep/equations/acoustic1st.py +210 -0
- sweep/equations/acoustic3d.py +168 -0
- sweep/equations/acoustic_aniso.py +186 -0
- sweep/equations/acoustic_curvilinear.py +166 -0
- sweep/equations/acoustic_lsrtm.py +175 -0
- sweep/equations/acoustic_lsrtm3d.py +230 -0
- sweep/equations/acoustic_vrr.py +145 -0
- sweep/equations/acoustic_vrz.py +365 -0
- sweep/equations/acoustic_vti_1st.py +626 -0
- sweep/equations/aec.py +61 -0
- sweep/equations/aec_lsrtm.py +100 -0
- sweep/equations/base.py +687 -0
- sweep/equations/cuda_layout.py +42 -0
- sweep/equations/das.py +1520 -0
- sweep/equations/elastic.py +323 -0
- sweep/equations/elastic3d.py +548 -0
- sweep/equations/elasticP.py +84 -0
- sweep/equations/elastic_apm.py +32 -0
- sweep/equations/elastic_curvilinear.py +337 -0
- sweep/equations/elastic_lsrtm.py +90 -0
- sweep/equations/elastic_tti.py +489 -0
- sweep/equations/elastic_tti_sg.py +371 -0
- sweep/equations/elastic_vrr.py +482 -0
- sweep/equations/elasticz.py +53 -0
- sweep/equations/fields.py +115 -0
- sweep/equations/pml.py +302 -0
- sweep/equations/qP_tariq.py +128 -0
- sweep/equations/qP_tti.py +152 -0
- sweep/equations/qP_vti.py +129 -0
- sweep/equations/utils.py +59 -0
- sweep/equations/visco_acoustic.py +191 -0
- sweep/memory/__init__.py +0 -0
- sweep/memory/shape.py +261 -0
- sweep/memory/torch.py +24 -0
- sweep/operators/__init__.py +25 -0
- sweep/operators/factory.py +47 -0
- sweep/operators/general.py +210 -0
- sweep/operators/jax.py +225 -0
- sweep/operators/rsg.py +175 -0
- sweep/operators/torch.py +200 -0
- sweep/propagator/__init__.py +23 -0
- sweep/propagator/_bs_dispatch.py +74 -0
- sweep/propagator/_c.py +1850 -0
- sweep/propagator/_c.pyi +39 -0
- sweep/propagator/_eager_boundary_saving.py +505 -0
- sweep/propagator/_jax_boundary_saving.py +310 -0
- sweep/propagator/_ring_geometry.py +54 -0
- sweep/propagator/_torch_eager.py +420 -0
- sweep/propagator/_torch_eager_custom_grad.py +421 -0
- sweep/propagator/base.py +1007 -0
- sweep/propagator/jax.py +485 -0
- sweep/propagator/jax.pyi +33 -0
- sweep/propagator/options.py +222 -0
- sweep/propagator/options.pyi +111 -0
- sweep/propagator/torch.py +374 -0
- sweep/propagator/torch.pyi +47 -0
- sweep/receivers/__init__.py +0 -0
- sweep/receivers/base.py +6 -0
- sweep/receivers/jax.py +18 -0
- sweep/receivers/torch.py +55 -0
- sweep/scalars.py +150 -0
- sweep/signal.py +89 -0
- sweep/sources/__init__.py +0 -0
- sweep/sources/base.py +24 -0
- sweep/sources/jax.py +54 -0
- sweep/sources/torch.py +70 -0
- sweep/utils/__init__.py +0 -0
- sweep/utils/curvilinear.py +235 -0
- sweep/utils/general.py +121 -0
- sweep/utils/jax.py +50 -0
- sweep/utils/torch.py +41 -0
- sweep_solver-0.1.0.dist-info/LICENSE +21 -0
- sweep_solver-0.1.0.dist-info/METADATA +160 -0
- sweep_solver-0.1.0.dist-info/RECORD +235 -0
- sweep_solver-0.1.0.dist-info/WHEEL +5 -0
- sweep_solver-0.1.0.dist-info/entry_points.txt +3 -0
- sweep_solver-0.1.0.dist-info/top_level.txt +2 -0
geophyai/__init__.py
ADDED
sweep/_C.py
ADDED
|
@@ -0,0 +1,35 @@
|
|
|
1
|
+
"""Lazy JIT entry point for sweep's compiled CUDA/C++ backend.
|
|
2
|
+
|
|
3
|
+
``import sweep._C`` is instant. The extension is compiled against your torch on
|
|
4
|
+
the **first attribute access** (i.e. the first real use of ``impl='c'``), then
|
|
5
|
+
cached — so ``is_torch_binding_available()`` / plain imports never trigger a
|
|
6
|
+
surprise ~3 min compile, and eager/JAX-only users never compile at all. Call
|
|
7
|
+
``sweep.precompile()`` to run that compile up front. See ``sweep/_jit.py``.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from . import _jit
|
|
11
|
+
|
|
12
|
+
_ready = False
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def _load():
|
|
16
|
+
"""Run the one-time JIT compile (cached) and expose the backend's functions
|
|
17
|
+
on this module. Idempotent — used by both ``__getattr__`` (first use) and
|
|
18
|
+
``sweep.precompile()`` (up-front)."""
|
|
19
|
+
global _ready
|
|
20
|
+
if _ready:
|
|
21
|
+
return
|
|
22
|
+
mod = _jit.load()
|
|
23
|
+
_ns = globals()
|
|
24
|
+
for _k in dir(mod):
|
|
25
|
+
if not _k.startswith("__"):
|
|
26
|
+
_ns[_k] = getattr(mod, _k)
|
|
27
|
+
_ready = True
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def __getattr__(name):
|
|
31
|
+
_load() # compile-on-first-use (cached after)
|
|
32
|
+
try:
|
|
33
|
+
return globals()[name]
|
|
34
|
+
except KeyError:
|
|
35
|
+
raise AttributeError(f"module 'sweep._C' has no attribute {name!r}")
|
sweep/__init__.py
ADDED
|
@@ -0,0 +1,182 @@
|
|
|
1
|
+
"""Top-level package helpers for sweep.
|
|
2
|
+
|
|
3
|
+
In addition to the wave-equation engine submodules (`equations`, `propagator`,
|
|
4
|
+
`operators`, …), this package re-exposes the **companion distributions** under
|
|
5
|
+
short namespace aliases::
|
|
6
|
+
|
|
7
|
+
import sweep
|
|
8
|
+
sweep.io.SEGYReader(...) # actually sweep_io.SEGYReader
|
|
9
|
+
from sweep import runner # actually sweep_runner
|
|
10
|
+
from sweep.tasks import TaskRunner # actually sweep_tasks.TaskRunner
|
|
11
|
+
|
|
12
|
+
This works for any companion that's `pip install`'d alongside sweep. Missing
|
|
13
|
+
companions surface a helpful ``AttributeError`` pointing at the right
|
|
14
|
+
``pip install`` command. The companion packages keep their real distribution
|
|
15
|
+
names (`sweep-io`, `sweep-tasks`, …) — the `sweep.<short>` form is purely a
|
|
16
|
+
convenience namespace.
|
|
17
|
+
"""
|
|
18
|
+
|
|
19
|
+
from __future__ import annotations
|
|
20
|
+
|
|
21
|
+
import sys
|
|
22
|
+
from importlib import import_module
|
|
23
|
+
from importlib.abc import MetaPathFinder
|
|
24
|
+
from importlib.util import find_spec
|
|
25
|
+
from pathlib import Path
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
_LAZY_SUBMODULES = {
|
|
29
|
+
"backend",
|
|
30
|
+
"equations",
|
|
31
|
+
"memory",
|
|
32
|
+
"operators",
|
|
33
|
+
"propagator",
|
|
34
|
+
"receivers",
|
|
35
|
+
"signal",
|
|
36
|
+
"sources",
|
|
37
|
+
"utils",
|
|
38
|
+
}
|
|
39
|
+
|
|
40
|
+
|
|
41
|
+
# Companion distributions exposed under the `sweep` namespace.
|
|
42
|
+
# Key = the short name you write as `sweep.<key>` / `from sweep import <key>`
|
|
43
|
+
# Value = the installed distribution's import name (PyPI name with `-` -> `_`)
|
|
44
|
+
_COMPANION_ALIASES: dict[str, str] = {
|
|
45
|
+
"io": "sweep_io",
|
|
46
|
+
"loss": "sweep_loss",
|
|
47
|
+
"nn": "sweep_nn",
|
|
48
|
+
"opt": "sweep_opt",
|
|
49
|
+
"preproc": "sweep_preproc",
|
|
50
|
+
"runner": "sweep_runner",
|
|
51
|
+
"tasks": "sweep_tasks",
|
|
52
|
+
"viz": "sweep_viz",
|
|
53
|
+
"zoo": "sweep_zoo",
|
|
54
|
+
}
|
|
55
|
+
|
|
56
|
+
|
|
57
|
+
def _extend_package_path_with_build_outputs() -> None:
|
|
58
|
+
"""Merge any `build/lib*/sweep` directory into `sweep.__path__`.
|
|
59
|
+
|
|
60
|
+
Lets a `python setup.py build_ext --inplace`-style build show up to
|
|
61
|
+
`import sweep._C` without a separate `pip install -e`.
|
|
62
|
+
"""
|
|
63
|
+
package_dir = Path(__file__).resolve().parent
|
|
64
|
+
repo_root = package_dir.parents[1]
|
|
65
|
+
build_dir = repo_root / "build"
|
|
66
|
+
|
|
67
|
+
if not build_dir.exists():
|
|
68
|
+
return
|
|
69
|
+
|
|
70
|
+
package_path = globals().get("__path__")
|
|
71
|
+
if package_path is None:
|
|
72
|
+
return
|
|
73
|
+
|
|
74
|
+
for candidate in sorted(build_dir.glob("lib*/sweep")):
|
|
75
|
+
candidate_str = str(candidate)
|
|
76
|
+
if candidate.is_dir() and candidate_str not in package_path:
|
|
77
|
+
package_path.append(candidate_str)
|
|
78
|
+
|
|
79
|
+
|
|
80
|
+
_extend_package_path_with_build_outputs()
|
|
81
|
+
|
|
82
|
+
|
|
83
|
+
def is_torch_binding_available() -> bool:
|
|
84
|
+
"""Return True when PyTorch + a CUDA GPU + nvcc are present, so ``sweep._C``
|
|
85
|
+
can be JIT-compiled against your torch on first use. Does NOT trigger the
|
|
86
|
+
compile itself (see ``sweep._jit``)."""
|
|
87
|
+
if find_spec("torch") is None:
|
|
88
|
+
return False
|
|
89
|
+
try:
|
|
90
|
+
from sweep import _jit
|
|
91
|
+
return _jit.can_build()[0]
|
|
92
|
+
except Exception:
|
|
93
|
+
return False
|
|
94
|
+
|
|
95
|
+
|
|
96
|
+
def precompile() -> bool:
|
|
97
|
+
"""Build the compiled CUDA backend (``sweep._C``) now.
|
|
98
|
+
|
|
99
|
+
Runs the one-time, per-GPU-arch JIT compile (~3-5 min) up front — e.g. right
|
|
100
|
+
after ``pip install`` — so it does NOT surprise you on first use of
|
|
101
|
+
``impl='c'``. A no-op once cached. Raises a clear error if PyTorch, a CUDA
|
|
102
|
+
GPU, or a suitable ``nvcc`` (>=12.6) is missing::
|
|
103
|
+
|
|
104
|
+
python -c "import sweep; sweep.precompile()"
|
|
105
|
+
"""
|
|
106
|
+
import sweep._C as _C
|
|
107
|
+
_C._load()
|
|
108
|
+
return True
|
|
109
|
+
|
|
110
|
+
|
|
111
|
+
# ---------------------------------------------------------------------------
|
|
112
|
+
# PEP 562 lazy attribute access — handles:
|
|
113
|
+
# import sweep; sweep.equations (native lazy submodule)
|
|
114
|
+
# import sweep; sweep.io.SEGYReader (companion alias)
|
|
115
|
+
# from sweep import runner (companion alias)
|
|
116
|
+
# ---------------------------------------------------------------------------
|
|
117
|
+
def __getattr__(name: str):
|
|
118
|
+
if name in _LAZY_SUBMODULES:
|
|
119
|
+
module = import_module(f"{__name__}.{name}")
|
|
120
|
+
globals()[name] = module
|
|
121
|
+
return module
|
|
122
|
+
if name in _COMPANION_ALIASES:
|
|
123
|
+
full = _COMPANION_ALIASES[name]
|
|
124
|
+
try:
|
|
125
|
+
module = import_module(full)
|
|
126
|
+
except ImportError as e:
|
|
127
|
+
raise AttributeError(
|
|
128
|
+
f"`sweep.{name}` requires the `{full}` companion package "
|
|
129
|
+
f"(install with `pip install {full.replace('_', '-')}`)."
|
|
130
|
+
) from e
|
|
131
|
+
# Make `from sweep.<name> import X` also work after first access by
|
|
132
|
+
# populating sys.modules under the alias.
|
|
133
|
+
sys.modules[f"sweep.{name}"] = module
|
|
134
|
+
globals()[name] = module
|
|
135
|
+
return module
|
|
136
|
+
raise AttributeError(f"module '{__name__}' has no attribute '{name}'")
|
|
137
|
+
|
|
138
|
+
|
|
139
|
+
def __dir__() -> list[str]:
|
|
140
|
+
return sorted(set(globals()) | _LAZY_SUBMODULES | set(_COMPANION_ALIASES))
|
|
141
|
+
|
|
142
|
+
|
|
143
|
+
# ---------------------------------------------------------------------------
|
|
144
|
+
# Meta-path finder — handles:
|
|
145
|
+
# import sweep.io (resolves to sweep_io)
|
|
146
|
+
# from sweep.io import SEGYReader (ditto)
|
|
147
|
+
# import sweep.io.prefetch (ditto, transitively)
|
|
148
|
+
#
|
|
149
|
+
# Without this, only the PEP-562 attribute paths above work; `import sweep.io`
|
|
150
|
+
# would raise ModuleNotFoundError because there's no `sweep/io/` on disk.
|
|
151
|
+
# ---------------------------------------------------------------------------
|
|
152
|
+
class _CompanionFinder(MetaPathFinder):
|
|
153
|
+
"""Resolve `sweep.<short>` and its descendants to the installed companion."""
|
|
154
|
+
|
|
155
|
+
_PREFIX = "sweep."
|
|
156
|
+
|
|
157
|
+
def find_spec(self, fullname, path=None, target=None): # noqa: D401
|
|
158
|
+
if not fullname.startswith(self._PREFIX):
|
|
159
|
+
return None
|
|
160
|
+
rest = fullname[len(self._PREFIX):]
|
|
161
|
+
head, _, tail = rest.partition(".")
|
|
162
|
+
if head not in _COMPANION_ALIASES:
|
|
163
|
+
return None
|
|
164
|
+
full = _COMPANION_ALIASES[head]
|
|
165
|
+
target_name = full if not tail else f"{full}.{tail}"
|
|
166
|
+
try:
|
|
167
|
+
return find_spec(target_name)
|
|
168
|
+
except (ImportError, ValueError):
|
|
169
|
+
return None
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
# Install once. The check makes a second `import sweep` (e.g. after reload)
|
|
173
|
+
# a no-op rather than registering duplicate finders.
|
|
174
|
+
if not any(isinstance(f, _CompanionFinder) for f in sys.meta_path):
|
|
175
|
+
sys.meta_path.append(_CompanionFinder())
|
|
176
|
+
|
|
177
|
+
|
|
178
|
+
__all__ = [
|
|
179
|
+
"is_torch_binding_available",
|
|
180
|
+
*_LAZY_SUBMODULES,
|
|
181
|
+
*_COMPANION_ALIASES,
|
|
182
|
+
]
|
sweep/_jit.py
ADDED
|
@@ -0,0 +1,306 @@
|
|
|
1
|
+
"""Compile sweep's CUDA/C++ backend against the *user's* torch, on first use.
|
|
2
|
+
|
|
3
|
+
This is why a single ``py3-none`` wheel of sweep works with **any** torch version
|
|
4
|
+
and any Python 3: the compiled extension (``sweep._C``) is not shipped pre-built —
|
|
5
|
+
it is JIT-compiled at runtime via ``torch.utils.cpp_extension.load()`` against
|
|
6
|
+
whatever libtorch is currently imported, then cached. First use of ``impl='c'``
|
|
7
|
+
pays a one-time ~2-5 min compile (only for *this* machine's GPU arch); every run
|
|
8
|
+
after that loads the cached ``.so`` instantly.
|
|
9
|
+
|
|
10
|
+
The C++ sources ship inside the wheel under ``sweep/csrc/`` (package data).
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from __future__ import annotations
|
|
14
|
+
|
|
15
|
+
import glob
|
|
16
|
+
import os
|
|
17
|
+
import shutil
|
|
18
|
+
import sys
|
|
19
|
+
from pathlib import Path
|
|
20
|
+
|
|
21
|
+
_PKG = Path(__file__).resolve().parent
|
|
22
|
+
_CSRC = _PKG / "csrc"
|
|
23
|
+
|
|
24
|
+
_module = None # cached compiled module (process-local)
|
|
25
|
+
|
|
26
|
+
|
|
27
|
+
# --------------------------------------------------------------------------- #
|
|
28
|
+
# CUDA toolkit (nvcc) discovery
|
|
29
|
+
# --------------------------------------------------------------------------- #
|
|
30
|
+
def _nvidia_pip_includes() -> list[str]:
|
|
31
|
+
"""Every ``nvidia/*/include`` dir from the pip CUDA wheels torch pulls in
|
|
32
|
+
(cuda_runtime, cusparse, cublas, cudnn, …) — so nvcc/host cc find the headers
|
|
33
|
+
even when there is no system CUDA toolkit."""
|
|
34
|
+
incs: list[str] = []
|
|
35
|
+
try:
|
|
36
|
+
import nvidia # namespace package from nvidia-*-cu12 wheels
|
|
37
|
+
except Exception:
|
|
38
|
+
return incs
|
|
39
|
+
for base in getattr(nvidia, "__path__", []):
|
|
40
|
+
for inc in sorted(glob.glob(os.path.join(base, "*", "include"))):
|
|
41
|
+
incs.append(inc)
|
|
42
|
+
return incs
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def _torch_cuda_major() -> int | None:
|
|
46
|
+
try:
|
|
47
|
+
import torch
|
|
48
|
+
v = torch.version.cuda # e.g. "12.8"
|
|
49
|
+
return int(v.split(".")[0]) if v else None
|
|
50
|
+
except Exception:
|
|
51
|
+
return None
|
|
52
|
+
|
|
53
|
+
|
|
54
|
+
def _nvcc_version(nvcc: str):
|
|
55
|
+
import re
|
|
56
|
+
import subprocess
|
|
57
|
+
try:
|
|
58
|
+
out = subprocess.run([nvcc, "--version"], capture_output=True,
|
|
59
|
+
text=True, timeout=20).stdout
|
|
60
|
+
m = re.search(r"release (\d+)\.(\d+)", out)
|
|
61
|
+
return (int(m.group(1)), int(m.group(2))) if m else None
|
|
62
|
+
except Exception:
|
|
63
|
+
return None
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
_cuda_home_cache = False # False = not computed; None/str = computed result
|
|
67
|
+
|
|
68
|
+
|
|
69
|
+
def _find_cuda_home() -> str | None:
|
|
70
|
+
"""Return a CUDA_HOME (dir with bin/nvcc) whose CUDA **major matches the
|
|
71
|
+
user's torch**. Priority: explicit CUDA_HOME env, then the pip
|
|
72
|
+
``nvidia-cuda-nvcc-cu12`` wheel (always cu12, what our dep pulls), then nvcc
|
|
73
|
+
on PATH — each version-checked so an old system nvcc (e.g. CUDA 10.1 in
|
|
74
|
+
/usr/bin) is skipped rather than used and failing mid-compile."""
|
|
75
|
+
global _cuda_home_cache
|
|
76
|
+
if _cuda_home_cache is not False:
|
|
77
|
+
return _cuda_home_cache
|
|
78
|
+
|
|
79
|
+
want = _torch_cuda_major()
|
|
80
|
+
allow_old = os.environ.get("SWEEP_JIT_ALLOW_OLD_CUDA", "").strip().lower() \
|
|
81
|
+
in ("1", "true", "yes", "on")
|
|
82
|
+
|
|
83
|
+
def match(nvcc: Path) -> bool:
|
|
84
|
+
if not nvcc.exists():
|
|
85
|
+
return False
|
|
86
|
+
v = _nvcc_version(str(nvcc))
|
|
87
|
+
if v is None:
|
|
88
|
+
return False
|
|
89
|
+
maj, minr = v
|
|
90
|
+
if want is not None and maj != want:
|
|
91
|
+
return False
|
|
92
|
+
# Floor: nvcc 12.4. CUDA 12.0-12.5 ship a <cuda/std> bf16 header whose
|
|
93
|
+
# host-device isnan/isinf call __device__-only half intrinsics; torch's
|
|
94
|
+
# build defines (-D__CUDA_NO_BFLOAT16_CONVERSIONS__ ...) plus
|
|
95
|
+
# --expt-relaxed-constexpr neutralize it from 12.4 up (verified: a clean
|
|
96
|
+
# 12.4 toolkit compiles the whole tree). 12.0-12.3 are untested here, so
|
|
97
|
+
# the guard rejects them; SWEEP_JIT_ALLOW_OLD_CUDA=1 tries one anyway.
|
|
98
|
+
return allow_old or not (maj == 12 and minr < 4)
|
|
99
|
+
|
|
100
|
+
result = None
|
|
101
|
+
# 1. explicit env (respect user config, but only if it matches torch's CUDA)
|
|
102
|
+
for env in ("CUDA_HOME", "CUDA_PATH"):
|
|
103
|
+
h = os.environ.get(env)
|
|
104
|
+
if h and match(Path(h) / "bin" / "nvcc"):
|
|
105
|
+
result = h
|
|
106
|
+
break
|
|
107
|
+
# 2. pip nvidia-cuda-nvcc-cu12 (namespace pkg -> __path__; guaranteed cu12)
|
|
108
|
+
if result is None:
|
|
109
|
+
try:
|
|
110
|
+
import nvidia.cuda_nvcc as _n # type: ignore
|
|
111
|
+
for base in getattr(_n, "__path__", []):
|
|
112
|
+
if match(Path(base) / "bin" / "nvcc"):
|
|
113
|
+
result = str(Path(base))
|
|
114
|
+
break
|
|
115
|
+
except Exception:
|
|
116
|
+
pass
|
|
117
|
+
# 3. nvcc on PATH (version-checked -> skips old /usr/bin/nvcc)
|
|
118
|
+
if result is None:
|
|
119
|
+
p = shutil.which("nvcc")
|
|
120
|
+
if p and match(Path(p)):
|
|
121
|
+
result = str(Path(p).resolve().parent.parent)
|
|
122
|
+
|
|
123
|
+
_cuda_home_cache = result
|
|
124
|
+
return result
|
|
125
|
+
|
|
126
|
+
|
|
127
|
+
def _nvidia_pip_libs() -> list[str]:
|
|
128
|
+
"""``nvidia/*/lib`` dirs so the JIT link step finds libcudart etc. when there
|
|
129
|
+
is no system CUDA toolkit (provided by torch's pip CUDA wheels)."""
|
|
130
|
+
libs: list[str] = []
|
|
131
|
+
try:
|
|
132
|
+
import nvidia
|
|
133
|
+
except Exception:
|
|
134
|
+
return libs
|
|
135
|
+
for base in getattr(nvidia, "__path__", []):
|
|
136
|
+
for lib in sorted(glob.glob(os.path.join(base, "*", "lib"))):
|
|
137
|
+
libs.append(lib)
|
|
138
|
+
return libs
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def _ensure_ninja_on_path() -> None:
|
|
142
|
+
"""torch checks ``ninja --version`` on PATH (not the bundled python pkg)."""
|
|
143
|
+
if shutil.which("ninja"):
|
|
144
|
+
return
|
|
145
|
+
try:
|
|
146
|
+
import ninja # the pip 'ninja' package exposes BIN_DIR
|
|
147
|
+
bindir = getattr(ninja, "BIN_DIR", None)
|
|
148
|
+
if bindir and os.path.isdir(bindir):
|
|
149
|
+
os.environ["PATH"] = bindir + os.pathsep + os.environ.get("PATH", "")
|
|
150
|
+
except Exception:
|
|
151
|
+
pass
|
|
152
|
+
|
|
153
|
+
|
|
154
|
+
def can_build() -> tuple[bool, str]:
|
|
155
|
+
"""(usable, reason) — True when torch+CUDA GPU+nvcc are present so the C
|
|
156
|
+
backend can be JIT-compiled. Does NOT compile. Used by
|
|
157
|
+
``sweep.is_torch_binding_available()`` to avoid a surprise compile."""
|
|
158
|
+
try:
|
|
159
|
+
import torch
|
|
160
|
+
except Exception:
|
|
161
|
+
return False, "PyTorch is not installed"
|
|
162
|
+
if not torch.cuda.is_available():
|
|
163
|
+
return False, "no CUDA GPU is visible"
|
|
164
|
+
if _find_cuda_home() is None:
|
|
165
|
+
return False, (
|
|
166
|
+
"no suitable CUDA toolkit found (need nvcc >=12.4 matching your "
|
|
167
|
+
"torch's CUDA major — 12.0-12.3 ship a broken <cuda/std> bf16 header). "
|
|
168
|
+
"sweep compiles its GPU backend on first use; provide a recent nvcc "
|
|
169
|
+
"via `module load cuda`, a system CUDA Toolkit, or "
|
|
170
|
+
"`conda install -c nvidia cuda-toolkit`. To try an older toolkit "
|
|
171
|
+
"anyway, set SWEEP_JIT_ALLOW_OLD_CUDA=1)")
|
|
172
|
+
return True, "ok"
|
|
173
|
+
|
|
174
|
+
|
|
175
|
+
# --------------------------------------------------------------------------- #
|
|
176
|
+
# source staging (dedupe object basenames)
|
|
177
|
+
# --------------------------------------------------------------------------- #
|
|
178
|
+
def _sources() -> list[str]:
|
|
179
|
+
"""C++/CUDA sources, mirroring build_config.get_sources(). CUDA-only by
|
|
180
|
+
default (fast first compile, what GPU users need); set SWEEP_JIT_FULL=1 to
|
|
181
|
+
also compile the heavy CPU C++ tree."""
|
|
182
|
+
cu = (glob.glob(str(_CSRC / "cuda/common/**/*.cu"), recursive=True)
|
|
183
|
+
+ glob.glob(str(_CSRC / "cuda/equations/**/*.cu"), recursive=True))
|
|
184
|
+
binding = [str(_CSRC / "bindings/module.cpp")]
|
|
185
|
+
if os.environ.get("SWEEP_JIT_FULL", "").lower() in ("1", "true", "yes", "on"):
|
|
186
|
+
cpu = [s for s in glob.glob(str(_CSRC / "cpu/**/*.cpp"), recursive=True)
|
|
187
|
+
if not s.endswith("cpu_binding_stub.cpp")]
|
|
188
|
+
else:
|
|
189
|
+
cpu = [str(_CSRC / "cpu/cpu_binding_stub.cpp")]
|
|
190
|
+
return cpu + cu + binding
|
|
191
|
+
|
|
192
|
+
|
|
193
|
+
def _stage(build_dir: Path) -> tuple[list[str], list[str]]:
|
|
194
|
+
"""cpp_extension.load() flattens object names by basename; sweep has many
|
|
195
|
+
forward.cu / backward.cu / kernels.cu. Copy csrc into a version-stamped
|
|
196
|
+
staging dir with UNIQUE compiled-source basenames (renamed in place so their
|
|
197
|
+
relative #includes still resolve). Idempotent across runs."""
|
|
198
|
+
try:
|
|
199
|
+
from importlib.metadata import version
|
|
200
|
+
_ver = version("sweep-solver")
|
|
201
|
+
except Exception:
|
|
202
|
+
_ver = "dev"
|
|
203
|
+
stage = build_dir / f"csrc_stage_{_ver}"
|
|
204
|
+
done = stage / ".staged"
|
|
205
|
+
if not done.exists():
|
|
206
|
+
shutil.rmtree(stage, ignore_errors=True)
|
|
207
|
+
shutil.copytree(_CSRC, stage)
|
|
208
|
+
for s in _sources():
|
|
209
|
+
rel = Path(s).resolve().relative_to(_CSRC)
|
|
210
|
+
slug = "_".join(rel.with_suffix("").parts)
|
|
211
|
+
os.replace(stage / rel, stage / rel.parent / (slug + rel.suffix))
|
|
212
|
+
done.write_text("ok")
|
|
213
|
+
staged = []
|
|
214
|
+
for s in _sources():
|
|
215
|
+
rel = Path(s).resolve().relative_to(_CSRC)
|
|
216
|
+
slug = "_".join(rel.with_suffix("").parts)
|
|
217
|
+
staged.append(str(stage / rel.parent / (slug + rel.suffix)))
|
|
218
|
+
inc = [str(stage), str(stage / "bindings"), str(stage / "shared"),
|
|
219
|
+
str(stage / "cuda"), str(stage / "cuda/common"), str(stage / "cuda/equations")]
|
|
220
|
+
return staged, inc
|
|
221
|
+
|
|
222
|
+
|
|
223
|
+
def _will_build(build_dir: Path) -> bool:
|
|
224
|
+
"""Whether the next load() will actually *compile* (vs reuse the cached .so).
|
|
225
|
+
|
|
226
|
+
A ``sweep_C.so`` can exist yet still be rebuilt — e.g. after the user upgrades
|
|
227
|
+
torch, whose changed headers make ninja re-link — so "the .so exists" is not a
|
|
228
|
+
reliable signal. Ask ninja (``-n`` dry run) whether any target is stale. This
|
|
229
|
+
drives the one-time "compiling…" notice + verbose output, so a genuine rebuild
|
|
230
|
+
is never a silent 2-5 min hang that looks frozen. When we can't tell, assume a
|
|
231
|
+
build so the user always sees *something*."""
|
|
232
|
+
so = build_dir / "sweep_C.so"
|
|
233
|
+
ninja_file = build_dir / "build.ninja"
|
|
234
|
+
if not so.exists() or not ninja_file.exists():
|
|
235
|
+
return True # never built (no .so / no ninja graph yet)
|
|
236
|
+
_ensure_ninja_on_path()
|
|
237
|
+
ninja = shutil.which("ninja")
|
|
238
|
+
if ninja is None:
|
|
239
|
+
return True # can't check -> assume yes (never hang silently)
|
|
240
|
+
try:
|
|
241
|
+
import subprocess
|
|
242
|
+
r = subprocess.run([ninja, "-n"], cwd=str(build_dir),
|
|
243
|
+
capture_output=True, text=True, timeout=30)
|
|
244
|
+
return "no work to do" not in (r.stdout + r.stderr)
|
|
245
|
+
except Exception:
|
|
246
|
+
return True
|
|
247
|
+
|
|
248
|
+
|
|
249
|
+
# --------------------------------------------------------------------------- #
|
|
250
|
+
# the loader
|
|
251
|
+
# --------------------------------------------------------------------------- #
|
|
252
|
+
def load():
|
|
253
|
+
"""Compile (first call, cached) and return the ``sweep._C`` module."""
|
|
254
|
+
global _module
|
|
255
|
+
if _module is not None:
|
|
256
|
+
return _module
|
|
257
|
+
|
|
258
|
+
import torch
|
|
259
|
+
from torch.utils import cpp_extension
|
|
260
|
+
|
|
261
|
+
ok, why = can_build()
|
|
262
|
+
if not ok:
|
|
263
|
+
raise RuntimeError(
|
|
264
|
+
f"sweep's compiled backend (impl='c') is unavailable: {why}. "
|
|
265
|
+
"Use impl='eager' for a pure-Python (slower) CPU/GPU path.")
|
|
266
|
+
|
|
267
|
+
cuda_home = _find_cuda_home()
|
|
268
|
+
os.environ["CUDA_HOME"] = cuda_home
|
|
269
|
+
os.environ["PATH"] = os.path.join(cuda_home, "bin") + os.pathsep + os.environ.get("PATH", "")
|
|
270
|
+
_ensure_ninja_on_path()
|
|
271
|
+
|
|
272
|
+
build_dir = Path(cpp_extension._get_build_directory("sweep_C", verbose=False))
|
|
273
|
+
build_dir.mkdir(parents=True, exist_ok=True)
|
|
274
|
+
sources, inc = _stage(build_dir)
|
|
275
|
+
# Use ONLY the selected CUDA toolkit's own headers (version-consistent with
|
|
276
|
+
# its nvcc). Do NOT mix in the pip nvidia-*/include dirs: for a torch built
|
|
277
|
+
# against an older CUDA (torch 2.5 = cu121 -> 12.1 headers) those clash with a
|
|
278
|
+
# newer toolkit and break the <cuda/std> bf16 compile.
|
|
279
|
+
inc = inc + [p for p in (os.path.join(cuda_home, "include"),
|
|
280
|
+
os.path.join(cuda_home, "targets", "x86_64-linux", "include"))
|
|
281
|
+
if os.path.isdir(p)]
|
|
282
|
+
|
|
283
|
+
cap = torch.cuda.get_device_capability()
|
|
284
|
+
building = _will_build(build_dir)
|
|
285
|
+
if building:
|
|
286
|
+
print(f"[sweep] compiling the CUDA backend for your GPU (sm_{cap[0]}{cap[1]}) — "
|
|
287
|
+
f"one-time, ~2-5 min, then cached at {build_dir} ...",
|
|
288
|
+
file=sys.stderr, flush=True)
|
|
289
|
+
|
|
290
|
+
_module = cpp_extension.load(
|
|
291
|
+
name="sweep_C",
|
|
292
|
+
sources=sources,
|
|
293
|
+
extra_include_paths=inc,
|
|
294
|
+
extra_cflags=["-O3", "-Wno-attributes", "-fopenmp"],
|
|
295
|
+
# --expt-relaxed-constexpr: lets constexpr __host__ funcs call __device__
|
|
296
|
+
# ones, which some CUDA toolkits' <cuda/std> bf16 headers (e.g. 12.4's
|
|
297
|
+
# nvbf16.h) require to compile. Harmless on toolkits that don't need it.
|
|
298
|
+
extra_cuda_cflags=["-O3", "--use_fast_math", "--expt-relaxed-constexpr",
|
|
299
|
+
"-Xcompiler=-Wno-deprecated-declarations"],
|
|
300
|
+
extra_ldflags=["-fopenmp"] + [f"-L{d}" for d in _nvidia_pip_libs()],
|
|
301
|
+
build_directory=str(build_dir),
|
|
302
|
+
verbose=building,
|
|
303
|
+
)
|
|
304
|
+
if building:
|
|
305
|
+
print("[sweep] CUDA backend compiled and cached.", file=sys.stderr, flush=True)
|
|
306
|
+
return _module
|
|
@@ -0,0 +1,15 @@
|
|
|
1
|
+
"""JAX backend capability helpers."""
|
|
2
|
+
|
|
3
|
+
from . import cuda
|
|
4
|
+
|
|
5
|
+
|
|
6
|
+
def is_available():
|
|
7
|
+
"""Return ``True`` when the JAX backend is importable."""
|
|
8
|
+
try:
|
|
9
|
+
import jax # noqa: F401
|
|
10
|
+
except Exception:
|
|
11
|
+
return False
|
|
12
|
+
return True
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
__all__ = ["cuda", "is_available"]
|
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
"""CUDA capability helpers for the JAX backend."""
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
def is_available():
|
|
5
|
+
"""Return ``True`` when JAX can see at least one GPU device."""
|
|
6
|
+
try:
|
|
7
|
+
import jax
|
|
8
|
+
except Exception:
|
|
9
|
+
return False
|
|
10
|
+
|
|
11
|
+
try:
|
|
12
|
+
return any(device.platform == "gpu" for device in jax.devices())
|
|
13
|
+
except Exception:
|
|
14
|
+
return False
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
__all__ = ["is_available"]
|
|
@@ -0,0 +1,16 @@
|
|
|
1
|
+
"""PyTorch backend capability helpers."""
|
|
2
|
+
|
|
3
|
+
from . import binding
|
|
4
|
+
from . import cuda
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def is_available():
|
|
8
|
+
"""Return ``True`` when the PyTorch backend is importable."""
|
|
9
|
+
try:
|
|
10
|
+
import torch # noqa: F401
|
|
11
|
+
except Exception:
|
|
12
|
+
return False
|
|
13
|
+
return True
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
__all__ = ["binding", "cuda", "is_available"]
|
|
@@ -0,0 +1,53 @@
|
|
|
1
|
+
"""Compiled PyTorch CUDA binding capability helpers.
|
|
2
|
+
|
|
3
|
+
``sweep._C`` is JIT-compiled from source on first use (see ``sweep/_jit.py``), so a
|
|
4
|
+
plain ``import sweep._C`` always succeeds regardless of whether the compile can or
|
|
5
|
+
did happen. These helpers therefore report the real state without triggering a
|
|
6
|
+
compile.
|
|
7
|
+
"""
|
|
8
|
+
|
|
9
|
+
import os
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def is_available() -> bool:
|
|
13
|
+
"""True when the compiled ``sweep._C`` backend is **usable** — i.e. PyTorch,
|
|
14
|
+
a CUDA GPU and a suitable ``nvcc`` (>=12.4) are present, so it can be (or
|
|
15
|
+
already is) JIT-compiled. Does NOT trigger the compile."""
|
|
16
|
+
try:
|
|
17
|
+
from sweep import _jit
|
|
18
|
+
return _jit.can_build()[0]
|
|
19
|
+
except Exception:
|
|
20
|
+
return False
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
def is_compiled() -> bool:
|
|
24
|
+
"""True when the backend is already built — compiled in this process, or a
|
|
25
|
+
cached ``.so`` from a previous run — so the first ``impl='c'`` use is instant."""
|
|
26
|
+
try:
|
|
27
|
+
from sweep import _jit
|
|
28
|
+
if _jit._module is not None:
|
|
29
|
+
return True
|
|
30
|
+
from torch.utils import cpp_extension
|
|
31
|
+
build_dir = cpp_extension._get_build_directory("sweep_C", verbose=False)
|
|
32
|
+
return os.path.exists(os.path.join(build_dir, "sweep_C.so"))
|
|
33
|
+
except Exception:
|
|
34
|
+
return False
|
|
35
|
+
|
|
36
|
+
|
|
37
|
+
def diagnostics() -> dict:
|
|
38
|
+
"""Diagnostics for the compiled backend — usable / why-not / nvcc / built."""
|
|
39
|
+
try:
|
|
40
|
+
from sweep import _jit
|
|
41
|
+
usable, reason = _jit.can_build()
|
|
42
|
+
return {
|
|
43
|
+
"usable": usable, # can impl='c' be used (built now / on first use)?
|
|
44
|
+
"reason": reason, # explanation when usable is False
|
|
45
|
+
"cuda_home": _jit._find_cuda_home(),
|
|
46
|
+
"already_compiled": is_compiled(),
|
|
47
|
+
}
|
|
48
|
+
except Exception as exc: # pragma: no cover
|
|
49
|
+
return {"usable": False, "reason": f"{type(exc).__name__}: {exc}",
|
|
50
|
+
"cuda_home": None, "already_compiled": False}
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
__all__ = ["diagnostics", "is_available", "is_compiled"]
|
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
"""CUDA capability helpers for the PyTorch backend."""
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
def is_available():
|
|
5
|
+
"""Return ``True`` when PyTorch reports CUDA support is available."""
|
|
6
|
+
try:
|
|
7
|
+
import torch
|
|
8
|
+
except Exception:
|
|
9
|
+
return False
|
|
10
|
+
|
|
11
|
+
try:
|
|
12
|
+
return bool(torch.cuda.is_available())
|
|
13
|
+
except Exception:
|
|
14
|
+
return False
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
__all__ = ["is_available"]
|