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.
Files changed (235) hide show
  1. geophyai/__init__.py +4 -0
  2. sweep/_C.py +35 -0
  3. sweep/__init__.py +182 -0
  4. sweep/_jit.py +306 -0
  5. sweep/backend/__init__.py +5 -0
  6. sweep/backend/jax/__init__.py +15 -0
  7. sweep/backend/jax/cuda.py +17 -0
  8. sweep/backend/torch/__init__.py +16 -0
  9. sweep/backend/torch/binding.py +53 -0
  10. sweep/backend/torch/cuda.py +17 -0
  11. sweep/cli.py +165 -0
  12. sweep/csrc/CMakeLists.txt +10 -0
  13. sweep/csrc/bindings/bindings_utils.h +48 -0
  14. sweep/csrc/bindings/module.cpp +254 -0
  15. sweep/csrc/cpu/common/cpu_engine.cpp +1181 -0
  16. sweep/csrc/cpu/common/cpu_engine.h +36 -0
  17. sweep/csrc/cpu/cpu_binding.cpp +136 -0
  18. sweep/csrc/cpu/cpu_binding.h +20 -0
  19. sweep/csrc/cpu/cpu_binding_stub.cpp +54 -0
  20. sweep/csrc/cpu/equations/acoustic2d/acoustic2d_cpu.cpp +1882 -0
  21. sweep/csrc/cpu/equations/acoustic2d/acoustic2d_cpu.h +13 -0
  22. sweep/csrc/cpu/equations/acoustic3d/acoustic3d_cpu.cpp +1645 -0
  23. sweep/csrc/cpu/equations/acoustic3d/acoustic3d_cpu.h +13 -0
  24. sweep/csrc/cpu/equations/acoustic_lsrtm2d/acoustic_lsrtm2d_cpu.cpp +1787 -0
  25. sweep/csrc/cpu/equations/acoustic_lsrtm2d/acoustic_lsrtm2d_cpu.h +13 -0
  26. sweep/csrc/cpu/equations/acoustic_lsrtm3d/acoustic_lsrtm3d_cpu.cpp +1741 -0
  27. sweep/csrc/cpu/equations/acoustic_lsrtm3d/acoustic_lsrtm3d_cpu.h +13 -0
  28. sweep/csrc/cpu/equations/acoustic_vrz2d/acoustic_vrz2d_cpu.cpp +1549 -0
  29. sweep/csrc/cpu/equations/acoustic_vrz2d/acoustic_vrz2d_cpu.h +13 -0
  30. sweep/csrc/cpu/equations/acoustic_vrz3d/acoustic_vrz3d_cpu.cpp +1587 -0
  31. sweep/csrc/cpu/equations/acoustic_vrz3d/acoustic_vrz3d_cpu.h +13 -0
  32. sweep/csrc/cpu/equations/das2d/das2d_cpu.cpp +752 -0
  33. sweep/csrc/cpu/equations/das2d/das2d_cpu.h +13 -0
  34. sweep/csrc/cpu/equations/das3d/das3d_cpu.cpp +938 -0
  35. sweep/csrc/cpu/equations/das3d/das3d_cpu.h +13 -0
  36. sweep/csrc/cpu/equations/das_mu2d/das_mu2d_cpu.cpp +1194 -0
  37. sweep/csrc/cpu/equations/das_mu2d/das_mu2d_cpu.h +13 -0
  38. sweep/csrc/cpu/equations/das_mu3d/das_mu3d_cpu.cpp +1364 -0
  39. sweep/csrc/cpu/equations/das_mu3d/das_mu3d_cpu.h +13 -0
  40. sweep/csrc/cpu/equations/elastic2d/elastic2d_cpu.cpp +826 -0
  41. sweep/csrc/cpu/equations/elastic2d/elastic2d_cpu.h +13 -0
  42. sweep/csrc/cpu/equations/elastic3d/elastic3d_cpu.cpp +916 -0
  43. sweep/csrc/cpu/equations/elastic3d/elastic3d_cpu.h +13 -0
  44. sweep/csrc/cpu/equations/elastic_tti_sg2d/elastic_tti_sg2d_cpu.cpp +1736 -0
  45. sweep/csrc/cpu/equations/elastic_tti_sg2d/elastic_tti_sg2d_cpu.h +13 -0
  46. sweep/csrc/cpu/operators/fd.h +268 -0
  47. sweep/csrc/cuda/common/acoustic.h +428 -0
  48. sweep/csrc/cuda/common/acoustic_vrz_fused.cuh +81 -0
  49. sweep/csrc/cuda/common/boundary/disk_io.cuh +349 -0
  50. sweep/csrc/cuda/common/boundary/kernels.cuh +193 -0
  51. sweep/csrc/cuda/common/boundary/runtime.cuh +1590 -0
  52. sweep/csrc/cuda/common/boundary/saver.cuh +1423 -0
  53. sweep/csrc/cuda/common/boundary/types.cuh +119 -0
  54. sweep/csrc/cuda/common/boundary_runtime.cuh +3 -0
  55. sweep/csrc/cuda/common/boundarysaver.cu +960 -0
  56. sweep/csrc/cuda/common/boundarysaver.cuh +3 -0
  57. sweep/csrc/cuda/common/checkpoint_runtime.cuh +347 -0
  58. sweep/csrc/cuda/common/common.cu +209 -0
  59. sweep/csrc/cuda/common/common.cuh +55 -0
  60. sweep/csrc/cuda/common/context.h +96 -0
  61. sweep/csrc/cuda/common/cudautils.h +147 -0
  62. sweep/csrc/cuda/common/das.h +403 -0
  63. sweep/csrc/cuda/common/das_mu.h +543 -0
  64. sweep/csrc/cuda/common/elastic.h +772 -0
  65. sweep/csrc/cuda/common/elastic_free_surface.cuh +366 -0
  66. sweep/csrc/cuda/common/wavetypes.h +3 -0
  67. sweep/csrc/cuda/equations/acoustic2d/acoustic2d.h +19 -0
  68. sweep/csrc/cuda/equations/acoustic2d/backward.cu +1156 -0
  69. sweep/csrc/cuda/equations/acoustic2d/forward.cu +217 -0
  70. sweep/csrc/cuda/equations/acoustic2d/kernels.cu +170 -0
  71. sweep/csrc/cuda/equations/acoustic2d/kernels.cuh +454 -0
  72. sweep/csrc/cuda/equations/acoustic3d/acoustic3d.h +19 -0
  73. sweep/csrc/cuda/equations/acoustic3d/backward.cu +1284 -0
  74. sweep/csrc/cuda/equations/acoustic3d/forward.cu +239 -0
  75. sweep/csrc/cuda/equations/acoustic3d/kernels.cu +194 -0
  76. sweep/csrc/cuda/equations/acoustic3d/kernels.cuh +547 -0
  77. sweep/csrc/cuda/equations/acoustic_lsrtm2d/acoustic_lsrtm2d.h +17 -0
  78. sweep/csrc/cuda/equations/acoustic_lsrtm2d/backward.cu +803 -0
  79. sweep/csrc/cuda/equations/acoustic_lsrtm2d/forward.cu +219 -0
  80. sweep/csrc/cuda/equations/acoustic_lsrtm2d/kernels.cu +65 -0
  81. sweep/csrc/cuda/equations/acoustic_lsrtm2d/kernels.cuh +318 -0
  82. sweep/csrc/cuda/equations/acoustic_lsrtm3d/acoustic_lsrtm3d.h +17 -0
  83. sweep/csrc/cuda/equations/acoustic_lsrtm3d/backward.cu +1160 -0
  84. sweep/csrc/cuda/equations/acoustic_lsrtm3d/forward.cu +223 -0
  85. sweep/csrc/cuda/equations/acoustic_lsrtm3d/kernels.cu +78 -0
  86. sweep/csrc/cuda/equations/acoustic_lsrtm3d/kernels.cuh +398 -0
  87. sweep/csrc/cuda/equations/acoustic_vrz2d/acoustic_vrz2d.h +13 -0
  88. sweep/csrc/cuda/equations/acoustic_vrz2d/backward.cu +648 -0
  89. sweep/csrc/cuda/equations/acoustic_vrz2d/forward.cu +201 -0
  90. sweep/csrc/cuda/equations/acoustic_vrz2d/kernels.cuh +1092 -0
  91. sweep/csrc/cuda/equations/acoustic_vrz3d/acoustic_vrz3d.h +13 -0
  92. sweep/csrc/cuda/equations/acoustic_vrz3d/backward.cu +738 -0
  93. sweep/csrc/cuda/equations/acoustic_vrz3d/forward.cu +206 -0
  94. sweep/csrc/cuda/equations/acoustic_vrz3d/kernels.cuh +1147 -0
  95. sweep/csrc/cuda/equations/acoustic_vti_1st_2d/acoustic_vti_1st_2d.h +17 -0
  96. sweep/csrc/cuda/equations/acoustic_vti_1st_2d/backward.cu +819 -0
  97. sweep/csrc/cuda/equations/acoustic_vti_1st_2d/forward.cu +353 -0
  98. sweep/csrc/cuda/equations/acoustic_vti_1st_2d/kernels.cu +4 -0
  99. sweep/csrc/cuda/equations/acoustic_vti_1st_2d/kernels.cuh +687 -0
  100. sweep/csrc/cuda/equations/acoustic_vti_1st_3d/acoustic_vti_1st_3d.h +14 -0
  101. sweep/csrc/cuda/equations/acoustic_vti_1st_3d/backward.cu +780 -0
  102. sweep/csrc/cuda/equations/acoustic_vti_1st_3d/forward.cu +365 -0
  103. sweep/csrc/cuda/equations/acoustic_vti_1st_3d/kernels.cu +5 -0
  104. sweep/csrc/cuda/equations/acoustic_vti_1st_3d/kernels.cuh +777 -0
  105. sweep/csrc/cuda/equations/das2d/backward.cu +829 -0
  106. sweep/csrc/cuda/equations/das2d/das2d.h +17 -0
  107. sweep/csrc/cuda/equations/das2d/forward.cu +255 -0
  108. sweep/csrc/cuda/equations/das2d/kernels.cuh +673 -0
  109. sweep/csrc/cuda/equations/das3d/backward.cu +391 -0
  110. sweep/csrc/cuda/equations/das3d/das3d.h +17 -0
  111. sweep/csrc/cuda/equations/das3d/forward.cu +179 -0
  112. sweep/csrc/cuda/equations/das3d/kernels.cuh +639 -0
  113. sweep/csrc/cuda/equations/das_mu2d/backward.cu +1068 -0
  114. sweep/csrc/cuda/equations/das_mu2d/das_mu2d.h +17 -0
  115. sweep/csrc/cuda/equations/das_mu2d/forward.cu +233 -0
  116. sweep/csrc/cuda/equations/das_mu2d/kernels.cuh +246 -0
  117. sweep/csrc/cuda/equations/das_mu3d/backward.cu +1171 -0
  118. sweep/csrc/cuda/equations/das_mu3d/das_mu3d.h +17 -0
  119. sweep/csrc/cuda/equations/das_mu3d/forward.cu +239 -0
  120. sweep/csrc/cuda/equations/das_mu3d/kernels.cuh +283 -0
  121. sweep/csrc/cuda/equations/elastic2d/backward.cu +1485 -0
  122. sweep/csrc/cuda/equations/elastic2d/elastic2d.h +27 -0
  123. sweep/csrc/cuda/equations/elastic2d/forward.cu +392 -0
  124. sweep/csrc/cuda/equations/elastic2d/kernels.cu +0 -0
  125. sweep/csrc/cuda/equations/elastic2d/kernels.cuh +1853 -0
  126. sweep/csrc/cuda/equations/elastic3d/backward.cu +1587 -0
  127. sweep/csrc/cuda/equations/elastic3d/elastic3d.h +24 -0
  128. sweep/csrc/cuda/equations/elastic3d/forward.cu +466 -0
  129. sweep/csrc/cuda/equations/elastic3d/kernels.cuh +2658 -0
  130. sweep/csrc/cuda/equations/elastic_tti_sg2d/backward.cu +784 -0
  131. sweep/csrc/cuda/equations/elastic_tti_sg2d/elastic_tti_sg2d.h +15 -0
  132. sweep/csrc/cuda/equations/elastic_tti_sg2d/forward.cu +246 -0
  133. sweep/csrc/cuda/equations/elastic_tti_sg2d/kernels.cuh +1037 -0
  134. sweep/csrc/cuda/equations/elastic_tti_sg2d/tensors.h +163 -0
  135. sweep/csrc/cuda/equations/elastic_vr2d/backward.cu +883 -0
  136. sweep/csrc/cuda/equations/elastic_vr2d/elastic_vr2d.h +17 -0
  137. sweep/csrc/cuda/equations/elastic_vr2d/forward.cu +207 -0
  138. sweep/csrc/cuda/equations/elastic_vr2d/kernels.cuh +1142 -0
  139. sweep/csrc/cuda/launch/config.h +169 -0
  140. sweep/csrc/cuda/operators/dim.cuh +12 -0
  141. sweep/csrc/cuda/operators/gradient.cuh +345 -0
  142. sweep/csrc/cuda/operators/laplace.cuh +354 -0
  143. sweep/csrc/cuda/operators/staggered.cuh +540 -0
  144. sweep/csrc/shared/wavetypes.h +191 -0
  145. sweep/datasets/__init__.py +115 -0
  146. sweep/datasets/_benchmarks.py +280 -0
  147. sweep/datasets/_cache.py +94 -0
  148. sweep/datasets/_formats.py +264 -0
  149. sweep/datasets/cli.py +89 -0
  150. sweep/datasets/marmousi.py +9178 -0
  151. sweep/datasets/overthrust_2d.py +1595 -0
  152. sweep/datasets/registry.py +170 -0
  153. sweep/equations/__init__.py +145 -0
  154. sweep/equations/_anisotropy_utils.py +89 -0
  155. sweep/equations/_elastic_step_core.py +234 -0
  156. sweep/equations/_free_surface.py +324 -0
  157. sweep/equations/_topography.py +748 -0
  158. sweep/equations/acoustic.py +175 -0
  159. sweep/equations/acoustic1st.py +210 -0
  160. sweep/equations/acoustic3d.py +168 -0
  161. sweep/equations/acoustic_aniso.py +186 -0
  162. sweep/equations/acoustic_curvilinear.py +166 -0
  163. sweep/equations/acoustic_lsrtm.py +175 -0
  164. sweep/equations/acoustic_lsrtm3d.py +230 -0
  165. sweep/equations/acoustic_vrr.py +145 -0
  166. sweep/equations/acoustic_vrz.py +365 -0
  167. sweep/equations/acoustic_vti_1st.py +626 -0
  168. sweep/equations/aec.py +61 -0
  169. sweep/equations/aec_lsrtm.py +100 -0
  170. sweep/equations/base.py +687 -0
  171. sweep/equations/cuda_layout.py +42 -0
  172. sweep/equations/das.py +1520 -0
  173. sweep/equations/elastic.py +323 -0
  174. sweep/equations/elastic3d.py +548 -0
  175. sweep/equations/elasticP.py +84 -0
  176. sweep/equations/elastic_apm.py +32 -0
  177. sweep/equations/elastic_curvilinear.py +337 -0
  178. sweep/equations/elastic_lsrtm.py +90 -0
  179. sweep/equations/elastic_tti.py +489 -0
  180. sweep/equations/elastic_tti_sg.py +371 -0
  181. sweep/equations/elastic_vrr.py +482 -0
  182. sweep/equations/elasticz.py +53 -0
  183. sweep/equations/fields.py +115 -0
  184. sweep/equations/pml.py +302 -0
  185. sweep/equations/qP_tariq.py +128 -0
  186. sweep/equations/qP_tti.py +152 -0
  187. sweep/equations/qP_vti.py +129 -0
  188. sweep/equations/utils.py +59 -0
  189. sweep/equations/visco_acoustic.py +191 -0
  190. sweep/memory/__init__.py +0 -0
  191. sweep/memory/shape.py +261 -0
  192. sweep/memory/torch.py +24 -0
  193. sweep/operators/__init__.py +25 -0
  194. sweep/operators/factory.py +47 -0
  195. sweep/operators/general.py +210 -0
  196. sweep/operators/jax.py +225 -0
  197. sweep/operators/rsg.py +175 -0
  198. sweep/operators/torch.py +200 -0
  199. sweep/propagator/__init__.py +23 -0
  200. sweep/propagator/_bs_dispatch.py +74 -0
  201. sweep/propagator/_c.py +1850 -0
  202. sweep/propagator/_c.pyi +39 -0
  203. sweep/propagator/_eager_boundary_saving.py +505 -0
  204. sweep/propagator/_jax_boundary_saving.py +310 -0
  205. sweep/propagator/_ring_geometry.py +54 -0
  206. sweep/propagator/_torch_eager.py +420 -0
  207. sweep/propagator/_torch_eager_custom_grad.py +421 -0
  208. sweep/propagator/base.py +1007 -0
  209. sweep/propagator/jax.py +485 -0
  210. sweep/propagator/jax.pyi +33 -0
  211. sweep/propagator/options.py +222 -0
  212. sweep/propagator/options.pyi +111 -0
  213. sweep/propagator/torch.py +374 -0
  214. sweep/propagator/torch.pyi +47 -0
  215. sweep/receivers/__init__.py +0 -0
  216. sweep/receivers/base.py +6 -0
  217. sweep/receivers/jax.py +18 -0
  218. sweep/receivers/torch.py +55 -0
  219. sweep/scalars.py +150 -0
  220. sweep/signal.py +89 -0
  221. sweep/sources/__init__.py +0 -0
  222. sweep/sources/base.py +24 -0
  223. sweep/sources/jax.py +54 -0
  224. sweep/sources/torch.py +70 -0
  225. sweep/utils/__init__.py +0 -0
  226. sweep/utils/curvilinear.py +235 -0
  227. sweep/utils/general.py +121 -0
  228. sweep/utils/jax.py +50 -0
  229. sweep/utils/torch.py +41 -0
  230. sweep_solver-0.1.0.dist-info/LICENSE +21 -0
  231. sweep_solver-0.1.0.dist-info/METADATA +160 -0
  232. sweep_solver-0.1.0.dist-info/RECORD +235 -0
  233. sweep_solver-0.1.0.dist-info/WHEEL +5 -0
  234. sweep_solver-0.1.0.dist-info/entry_points.txt +3 -0
  235. sweep_solver-0.1.0.dist-info/top_level.txt +2 -0
geophyai/__init__.py ADDED
@@ -0,0 +1,4 @@
1
+ import sys
2
+ import sweep
3
+
4
+ sys.modules['geophyai'] = sweep
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,5 @@
1
+ """Backend capability helpers exposed under ``sweep.backend``."""
2
+
3
+ from . import jax, torch
4
+
5
+ __all__ = ["jax", "torch"]
@@ -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"]