commkit 1.0.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.
- commkit/__init__.py +74 -0
- commkit/_cuda/__init__.py +321 -0
- commkit/_cuda/compiler.py +88 -0
- commkit/_cuda/src/bps_min_d2.cu +104 -0
- commkit/_cuda/src/cs_block.cu +119 -0
- commkit/_cuda/src/selftest.cu +14 -0
- commkit/analysis/__init__.py +55 -0
- commkit/analysis/_common.py +236 -0
- commkit/analysis/allan.py +108 -0
- commkit/analysis/drift.py +213 -0
- commkit/analysis/interferometry.py +887 -0
- commkit/analysis/linewidth.py +480 -0
- commkit/analysis/trajectory.py +91 -0
- commkit/backend.py +507 -0
- commkit/coding/__init__.py +23 -0
- commkit/coding/base.py +17 -0
- commkit/coding/bch.py +6 -0
- commkit/coding/convolutional.py +7 -0
- commkit/coding/crc.py +7 -0
- commkit/coding/galois.py +8 -0
- commkit/coding/hamming.py +6 -0
- commkit/coding/interleaving.py +7 -0
- commkit/coding/ldpc.py +8 -0
- commkit/coding/polar.py +8 -0
- commkit/coding/ratematch.py +6 -0
- commkit/coding/reed_solomon.py +6 -0
- commkit/coding/turbo.py +8 -0
- commkit/core/__init__.py +32 -0
- commkit/core/frame.py +992 -0
- commkit/core/generation.py +581 -0
- commkit/core/signal.py +725 -0
- commkit/equalization/__init__.py +49 -0
- commkit/equalization/_block.py +1855 -0
- commkit/equalization/_common.py +606 -0
- commkit/equalization/_kernels_jax.py +1720 -0
- commkit/equalization/_kernels_numba.py +1704 -0
- commkit/equalization/blind.py +223 -0
- commkit/equalization/linear.py +365 -0
- commkit/equalization/polarization.py +790 -0
- commkit/equalization/result.py +191 -0
- commkit/equalization/sequential.py +2805 -0
- commkit/filtering.py +1120 -0
- commkit/frequency.py +1191 -0
- commkit/helpers.py +489 -0
- commkit/impairments/__init__.py +43 -0
- commkit/impairments/channel/__init__.py +20 -0
- commkit/impairments/channel/linear.py +310 -0
- commkit/impairments/channel/nonlinear.py +11 -0
- commkit/impairments/frontend.py +229 -0
- commkit/impairments/noise.py +105 -0
- commkit/impairments/source.py +219 -0
- commkit/io.py +308 -0
- commkit/logger.py +103 -0
- commkit/mapping/__init__.py +46 -0
- commkit/mapping/bits.py +240 -0
- commkit/mapping/constellation.py +153 -0
- commkit/mapping/gray.py +429 -0
- commkit/mapping/llr.py +253 -0
- commkit/mapping/shaping.py +218 -0
- commkit/metrics.py +949 -0
- commkit/multirate.py +476 -0
- commkit/plotting/__init__.py +78 -0
- commkit/plotting/analysis.py +627 -0
- commkit/plotting/constellation.py +483 -0
- commkit/plotting/equalizer.py +390 -0
- commkit/plotting/eye.py +388 -0
- commkit/plotting/spectral.py +575 -0
- commkit/plotting/sync.py +953 -0
- commkit/plotting/theme.py +203 -0
- commkit/plotting/waveform.py +200 -0
- commkit/py.typed +0 -0
- commkit/recovery/__init__.py +51 -0
- commkit/recovery/bps.py +337 -0
- commkit/recovery/corrections.py +751 -0
- commkit/recovery/pilots.py +803 -0
- commkit/recovery/pll.py +482 -0
- commkit/recovery/tikhonov.py +424 -0
- commkit/recovery/viterbi_viterbi.py +227 -0
- commkit/spectral.py +560 -0
- commkit/timing.py +841 -0
- commkit-1.0.0.dist-info/METADATA +145 -0
- commkit-1.0.0.dist-info/RECORD +84 -0
- commkit-1.0.0.dist-info/WHEEL +4 -0
- commkit-1.0.0.dist-info/licenses/LICENSE +21 -0
commkit/backend.py
ADDED
|
@@ -0,0 +1,507 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Computational backend management and device orchestration.
|
|
3
|
+
|
|
4
|
+
This module provides the infrastructure for backend-agnostic execution across
|
|
5
|
+
CPU (NumPy), GPU (CuPy), and JAX. It implements a data-driven dispatch mechanism
|
|
6
|
+
that allows the library to automatically adjust its internal logic based on where
|
|
7
|
+
the input data resides.
|
|
8
|
+
|
|
9
|
+
The backend system is designed to be stateless and transparent, requiring
|
|
10
|
+
minimal explicit device management from the user.
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
import types
|
|
14
|
+
from functools import cache, lru_cache
|
|
15
|
+
from typing import Any
|
|
16
|
+
|
|
17
|
+
import numpy as np
|
|
18
|
+
|
|
19
|
+
from .logger import logger
|
|
20
|
+
|
|
21
|
+
# Try to import CuPy and verify functionality
|
|
22
|
+
try:
|
|
23
|
+
import cupy as cp
|
|
24
|
+
|
|
25
|
+
# Aggressive check: try to allocate and run a simple operation.
|
|
26
|
+
# This catches cases where CuPy is installed but shared libraries (nvrtc, cublas) are missing.
|
|
27
|
+
try:
|
|
28
|
+
cp.arange(1)
|
|
29
|
+
_CUPY_AVAILABLE = True
|
|
30
|
+
logger.info("CuPy is available and functional, defaulting Signals to GPU.")
|
|
31
|
+
except Exception:
|
|
32
|
+
# Fallback if functional check fails
|
|
33
|
+
_CUPY_AVAILABLE = False
|
|
34
|
+
cp = None
|
|
35
|
+
logger.warning(
|
|
36
|
+
"CuPy has problems with shared libraries, falling back to NumPy."
|
|
37
|
+
)
|
|
38
|
+
|
|
39
|
+
except ImportError:
|
|
40
|
+
_CUPY_AVAILABLE = False
|
|
41
|
+
cp = None
|
|
42
|
+
logger.debug("CuPy is not available, falling back to NumPy.")
|
|
43
|
+
|
|
44
|
+
# Any for CuPy array to avoid a hard dependency in the type hint if not installed
|
|
45
|
+
ArrayType = np.ndarray | Any
|
|
46
|
+
|
|
47
|
+
# JAX lazy loading cache
|
|
48
|
+
_JAX_CACHE: dict[str, Any] = {}
|
|
49
|
+
|
|
50
|
+
|
|
51
|
+
def _get_jax() -> tuple[types.ModuleType | None, types.ModuleType | None, Any | None]:
|
|
52
|
+
"""
|
|
53
|
+
Lazy loader for JAX modules and its DLPack interface.
|
|
54
|
+
|
|
55
|
+
Returns
|
|
56
|
+
-------
|
|
57
|
+
jax : module or None
|
|
58
|
+
The base `jax` module if installed, else None.
|
|
59
|
+
jnp : module or None
|
|
60
|
+
The `jax.numpy` namespace if installed, else None.
|
|
61
|
+
dlpack : module or None
|
|
62
|
+
The `jax.dlpack` interface for zero-copy transfers, else None.
|
|
63
|
+
"""
|
|
64
|
+
if "jax" not in _JAX_CACHE:
|
|
65
|
+
try:
|
|
66
|
+
import jax
|
|
67
|
+
import jax.numpy as jnp
|
|
68
|
+
from jax import dlpack
|
|
69
|
+
|
|
70
|
+
_JAX_CACHE["jax"] = jax
|
|
71
|
+
_JAX_CACHE["jnp"] = jnp
|
|
72
|
+
_JAX_CACHE["dlpack"] = dlpack
|
|
73
|
+
except ImportError:
|
|
74
|
+
_JAX_CACHE["jax"] = None
|
|
75
|
+
|
|
76
|
+
return _JAX_CACHE.get("jax"), _JAX_CACHE.get("jnp"), _JAX_CACHE.get("dlpack")
|
|
77
|
+
|
|
78
|
+
|
|
79
|
+
@lru_cache(maxsize=8)
|
|
80
|
+
def _get_jax_device(platform: str) -> Any | None:
|
|
81
|
+
"""
|
|
82
|
+
Retrieves a specific JAX device by platform name.
|
|
83
|
+
|
|
84
|
+
Parameters
|
|
85
|
+
----------
|
|
86
|
+
platform : {"cpu", "gpu", "tpu"}
|
|
87
|
+
The target hardware platform identifier.
|
|
88
|
+
|
|
89
|
+
Returns
|
|
90
|
+
-------
|
|
91
|
+
device : Device or None
|
|
92
|
+
The first discovered device for the specified platform, or None
|
|
93
|
+
if JAX is missing or the platform is unsupported.
|
|
94
|
+
"""
|
|
95
|
+
jax, _, _ = _get_jax()
|
|
96
|
+
if jax is None:
|
|
97
|
+
return None
|
|
98
|
+
try:
|
|
99
|
+
# Map our common names to JAX platform names
|
|
100
|
+
platform_map = {"cpu": "cpu", "gpu": "cuda", "tpu": "tpu"}
|
|
101
|
+
jax_platform = platform_map.get(platform, platform)
|
|
102
|
+
return jax.devices(jax_platform)[0]
|
|
103
|
+
except (RuntimeError, IndexError):
|
|
104
|
+
return None
|
|
105
|
+
|
|
106
|
+
|
|
107
|
+
_FORCE_CPU = False
|
|
108
|
+
|
|
109
|
+
|
|
110
|
+
def use_cpu_only(force: bool = True) -> None:
|
|
111
|
+
"""
|
|
112
|
+
Enforces a CPU-only execution path, disabling GPU discovery.
|
|
113
|
+
|
|
114
|
+
This function effectively hides CuPy from the library, even if a
|
|
115
|
+
functional NVIDIA GPU and CuPy installation are present.
|
|
116
|
+
|
|
117
|
+
Parameters
|
|
118
|
+
----------
|
|
119
|
+
force : bool, default True
|
|
120
|
+
If True, blocks all CUDA-accelerated operations.
|
|
121
|
+
"""
|
|
122
|
+
global _FORCE_CPU
|
|
123
|
+
_FORCE_CPU = force
|
|
124
|
+
|
|
125
|
+
|
|
126
|
+
def is_cupy_available() -> bool:
|
|
127
|
+
"""
|
|
128
|
+
Checks if NVIDIA GPU acceleration is functional via CuPy.
|
|
129
|
+
|
|
130
|
+
Returns
|
|
131
|
+
-------
|
|
132
|
+
bool
|
|
133
|
+
True if CuPy is installed, functional, and not explicitly disabled
|
|
134
|
+
via `use_cpu_only`.
|
|
135
|
+
"""
|
|
136
|
+
if _FORCE_CPU:
|
|
137
|
+
return False
|
|
138
|
+
return _CUPY_AVAILABLE
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
def get_array_module(data: Any) -> types.ModuleType:
|
|
142
|
+
"""
|
|
143
|
+
Infers the array module (NumPy or CuPy) for the given data.
|
|
144
|
+
|
|
145
|
+
The decision is made by inspecting the **actual type of the data**, not the
|
|
146
|
+
global availability/force flags. A CuPy array is therefore always reported
|
|
147
|
+
as CuPy - even under :func:`use_cpu_only` - because that flag governs the
|
|
148
|
+
default *placement of new* arrays, not the module of data that already lives
|
|
149
|
+
on the GPU. Reporting NumPy for a CuPy array would route GPU data into
|
|
150
|
+
NumPy code paths and raise ``TypeError`` (or silently mis-dispatch).
|
|
151
|
+
|
|
152
|
+
Parameters
|
|
153
|
+
----------
|
|
154
|
+
data : array_like or list
|
|
155
|
+
The input data to inspect.
|
|
156
|
+
|
|
157
|
+
Returns
|
|
158
|
+
-------
|
|
159
|
+
module
|
|
160
|
+
`cupy` if the data is a CuPy device array, otherwise `numpy`
|
|
161
|
+
(CPU arrays, lists, and scalars).
|
|
162
|
+
"""
|
|
163
|
+
# `cp is not None` <=> CuPy imported and passed the functional check at import
|
|
164
|
+
# time (it is set to None otherwise), so no CuPy array can exist when it is
|
|
165
|
+
# None. This is intentionally independent of `is_cupy_available()`, which
|
|
166
|
+
# also returns False under `use_cpu_only()`.
|
|
167
|
+
if cp is not None and isinstance(data, cp.ndarray):
|
|
168
|
+
return cp
|
|
169
|
+
return np
|
|
170
|
+
|
|
171
|
+
|
|
172
|
+
@cache
|
|
173
|
+
def get_scipy_module(xp: types.ModuleType) -> types.ModuleType:
|
|
174
|
+
"""
|
|
175
|
+
Returns the signal processing library compatible with the given array module.
|
|
176
|
+
|
|
177
|
+
Parameters
|
|
178
|
+
----------
|
|
179
|
+
xp : module
|
|
180
|
+
The array module (typically `numpy` or `cupy`).
|
|
181
|
+
|
|
182
|
+
Returns
|
|
183
|
+
-------
|
|
184
|
+
sp : module
|
|
185
|
+
The corresponding signal processing module (`scipy` or `cupyx.scipy`).
|
|
186
|
+
"""
|
|
187
|
+
# Match sp to the actual array module, independent of the force-CPU flag:
|
|
188
|
+
# if xp is CuPy we must return cupyx.scipy so dispatch() stays internally
|
|
189
|
+
# consistent (xp/sp paired) for GPU arrays passed under use_cpu_only().
|
|
190
|
+
if cp is not None and xp is cp:
|
|
191
|
+
import cupyx.scipy
|
|
192
|
+
import cupyx.scipy.ndimage
|
|
193
|
+
import cupyx.scipy.signal
|
|
194
|
+
import cupyx.scipy.special
|
|
195
|
+
|
|
196
|
+
return cupyx.scipy
|
|
197
|
+
|
|
198
|
+
import scipy
|
|
199
|
+
import scipy.ndimage
|
|
200
|
+
import scipy.signal
|
|
201
|
+
import scipy.special
|
|
202
|
+
|
|
203
|
+
return scipy
|
|
204
|
+
|
|
205
|
+
|
|
206
|
+
def to_device(data: Any, device: str) -> ArrayType:
|
|
207
|
+
"""
|
|
208
|
+
Moves data between CPU and GPU devices.
|
|
209
|
+
|
|
210
|
+
Parameters
|
|
211
|
+
----------
|
|
212
|
+
data : array_like
|
|
213
|
+
The data to move.
|
|
214
|
+
device : {"CPU", "GPU"}
|
|
215
|
+
Target device name (case-insensitive).
|
|
216
|
+
|
|
217
|
+
Returns
|
|
218
|
+
-------
|
|
219
|
+
array_like
|
|
220
|
+
The data residing on the target device.
|
|
221
|
+
|
|
222
|
+
Raises
|
|
223
|
+
------
|
|
224
|
+
ImportError
|
|
225
|
+
If "GPU" is requested but CuPy is not available.
|
|
226
|
+
ValueError
|
|
227
|
+
If an unsupported device name is provided.
|
|
228
|
+
|
|
229
|
+
Notes
|
|
230
|
+
-----
|
|
231
|
+
If the data is already on the target device, this operation
|
|
232
|
+
typically returns a view or the original array to avoid
|
|
233
|
+
unnecessary copies.
|
|
234
|
+
"""
|
|
235
|
+
logger.debug("Moving data to %s.", device.upper())
|
|
236
|
+
device = device.lower()
|
|
237
|
+
if device == "cpu":
|
|
238
|
+
# Dispatch by the *actual array type*, independent of the force-CPU flag
|
|
239
|
+
# (mirrors get_array_module/get_scipy_module). An array that already lives
|
|
240
|
+
# on the GPU must always be brought to host; gating the ``.get()`` on
|
|
241
|
+
# is_cupy_available() means use_cpu_only() leaves a CuPy array unfetchable
|
|
242
|
+
# and the np.asarray() fallback raises "Implicit conversion ... use .get()".
|
|
243
|
+
if cp is not None and isinstance(data, cp.ndarray):
|
|
244
|
+
return data.get()
|
|
245
|
+
if isinstance(data, np.ndarray):
|
|
246
|
+
return data
|
|
247
|
+
return np.asarray(data)
|
|
248
|
+
|
|
249
|
+
elif device == "gpu":
|
|
250
|
+
if not is_cupy_available():
|
|
251
|
+
raise ImportError("CuPy is not available.")
|
|
252
|
+
if isinstance(data, cp.ndarray):
|
|
253
|
+
return data
|
|
254
|
+
return cp.asarray(data)
|
|
255
|
+
|
|
256
|
+
else:
|
|
257
|
+
raise ValueError(f"Unknown device: {device.upper()}")
|
|
258
|
+
|
|
259
|
+
|
|
260
|
+
def dispatch(
|
|
261
|
+
data: Any,
|
|
262
|
+
) -> tuple[ArrayType, types.ModuleType, types.ModuleType]:
|
|
263
|
+
"""
|
|
264
|
+
Inspects data and returns appropriate backend modules.
|
|
265
|
+
|
|
266
|
+
This helper is used throughout the library to implement backend-agnostic
|
|
267
|
+
functional logic.
|
|
268
|
+
|
|
269
|
+
Parameters
|
|
270
|
+
----------
|
|
271
|
+
data : array_like
|
|
272
|
+
The input data to analyze.
|
|
273
|
+
|
|
274
|
+
Returns
|
|
275
|
+
-------
|
|
276
|
+
data_array : array_like
|
|
277
|
+
The input data forced to an array on its current device.
|
|
278
|
+
xp : module
|
|
279
|
+
The array module (`numpy` or `cupy`).
|
|
280
|
+
sp : module
|
|
281
|
+
The signal processing module (`scipy` or `cupyx.scipy`).
|
|
282
|
+
"""
|
|
283
|
+
xp = get_array_module(data)
|
|
284
|
+
sp = get_scipy_module(xp)
|
|
285
|
+
|
|
286
|
+
if not isinstance(data, (np.ndarray, getattr(cp, "ndarray", type(None)))):
|
|
287
|
+
data = xp.asarray(data)
|
|
288
|
+
|
|
289
|
+
return data, xp, sp
|
|
290
|
+
|
|
291
|
+
|
|
292
|
+
def to_jax(data: Any, device: str | None = None, dtype: Any | None = None) -> Any:
|
|
293
|
+
"""
|
|
294
|
+
Converts data to a JAX array with optimized device placement.
|
|
295
|
+
|
|
296
|
+
This function supports zero-copy transfers from CuPy using DLPack
|
|
297
|
+
when moving data between CUDA-managed memories.
|
|
298
|
+
|
|
299
|
+
Parameters
|
|
300
|
+
----------
|
|
301
|
+
data : array_like
|
|
302
|
+
Input data (NumPy array, CuPy array, list, or scalar).
|
|
303
|
+
device : {"CPU", "GPU", "TPU"}, optional
|
|
304
|
+
Target JAX device platform. If None, the function attempts to
|
|
305
|
+
preserve the device of the original data.
|
|
306
|
+
dtype : dtype, optional
|
|
307
|
+
Target data type. If None (default), implicit casting logic is applied:
|
|
308
|
+
complex128 -> complex64 and float64 -> float32 are enforced to avoid
|
|
309
|
+
backend bottlenecks, unless JAX x64 mode is explicitly enabled.
|
|
310
|
+
|
|
311
|
+
Returns
|
|
312
|
+
-------
|
|
313
|
+
jax_array : jax.Array
|
|
314
|
+
A JAX array residing on the specified or inferred device.
|
|
315
|
+
|
|
316
|
+
Raises
|
|
317
|
+
------
|
|
318
|
+
ImportError
|
|
319
|
+
If the `jax` library is not installed.
|
|
320
|
+
ValueError
|
|
321
|
+
If the requested `device` platform is not available in the
|
|
322
|
+
local JAX environment.
|
|
323
|
+
"""
|
|
324
|
+
jax, jnp, jax_dlpack = _get_jax()
|
|
325
|
+
if jax is None or jnp is None:
|
|
326
|
+
raise ImportError("JAX is not installed.")
|
|
327
|
+
|
|
328
|
+
# Check for JAX x64 mode
|
|
329
|
+
try:
|
|
330
|
+
from jax import config
|
|
331
|
+
|
|
332
|
+
x64_enabled = config.read("jax_enable_x64")
|
|
333
|
+
except (ImportError, AttributeError):
|
|
334
|
+
x64_enabled = False
|
|
335
|
+
|
|
336
|
+
# Resolution of target dtype
|
|
337
|
+
# If explicit dtype is None, we apply the "DSP Design" heuristic:
|
|
338
|
+
# Downgrade 64-bit to 32-bit for performance unless x64 is strictly requested.
|
|
339
|
+
target_dtype = None
|
|
340
|
+
if dtype is not None:
|
|
341
|
+
target_dtype = dtype
|
|
342
|
+
elif not x64_enabled:
|
|
343
|
+
# Auto-cast logic
|
|
344
|
+
if hasattr(data, "dtype"):
|
|
345
|
+
dt = data.dtype
|
|
346
|
+
if dt == "complex128":
|
|
347
|
+
target_dtype = "complex64"
|
|
348
|
+
elif dt == "float64":
|
|
349
|
+
target_dtype = "float32"
|
|
350
|
+
|
|
351
|
+
# Apply cast if needed (before transfer if possible/efficient)
|
|
352
|
+
# For NumPy: cast on CPU before transfer/conversion
|
|
353
|
+
if (
|
|
354
|
+
isinstance(data, np.ndarray)
|
|
355
|
+
and target_dtype is not None
|
|
356
|
+
and data.dtype != target_dtype
|
|
357
|
+
):
|
|
358
|
+
data = data.astype(target_dtype)
|
|
359
|
+
|
|
360
|
+
# For CuPy: cast on GPU before DLPack
|
|
361
|
+
if (
|
|
362
|
+
is_cupy_available()
|
|
363
|
+
and isinstance(data, cp.ndarray)
|
|
364
|
+
and target_dtype is not None
|
|
365
|
+
and data.dtype != target_dtype
|
|
366
|
+
):
|
|
367
|
+
data = data.astype(target_dtype)
|
|
368
|
+
|
|
369
|
+
target_device = None
|
|
370
|
+
if device is not None:
|
|
371
|
+
target_device = _get_jax_device(device.lower())
|
|
372
|
+
if target_device is None:
|
|
373
|
+
raise ValueError(f"Requested JAX device '{device}' is not available.")
|
|
374
|
+
|
|
375
|
+
# --- Conversion paths (all funnel to `result`) ---
|
|
376
|
+
result = None
|
|
377
|
+
|
|
378
|
+
# 1. Handle CuPy -> JAX (GPU)
|
|
379
|
+
if is_cupy_available() and isinstance(data, cp.ndarray):
|
|
380
|
+
try:
|
|
381
|
+
# DLPack requires contiguous memory and proper alignment.
|
|
382
|
+
# Enforce contiguous layout and 16-byte alignment (JAX/XLA requirement).
|
|
383
|
+
needs_copy = not data.flags.c_contiguous
|
|
384
|
+
if not needs_copy:
|
|
385
|
+
# Check for 16-byte alignment (common requirement for vectorized loads)
|
|
386
|
+
if data.data.ptr % 16 != 0:
|
|
387
|
+
needs_copy = True
|
|
388
|
+
|
|
389
|
+
if needs_copy:
|
|
390
|
+
data = cp.array(data, copy=True, order="C")
|
|
391
|
+
|
|
392
|
+
if jax_dlpack is not None:
|
|
393
|
+
jax_arr = jax_dlpack.from_dlpack(data)
|
|
394
|
+
if target_device and jax_arr.device != target_device:
|
|
395
|
+
result = jax.device_put(jax_arr, target_device)
|
|
396
|
+
else:
|
|
397
|
+
result = jax_arr
|
|
398
|
+
|
|
399
|
+
except Exception as e:
|
|
400
|
+
logger.debug(
|
|
401
|
+
"DLPack transfer from CuPy to JAX failed: %s. Falling back to explicit conversion.",
|
|
402
|
+
e,
|
|
403
|
+
)
|
|
404
|
+
|
|
405
|
+
# 2. Optimized Placement
|
|
406
|
+
# If a target device is specified, use device_put directly.
|
|
407
|
+
# This is more efficient than jnp.asarray(data) + device_put because it avoids
|
|
408
|
+
# an intermediate placement on the JAX default device.
|
|
409
|
+
if result is None and target_device:
|
|
410
|
+
result = jax.device_put(data, target_device)
|
|
411
|
+
|
|
412
|
+
# 3. Preservation Logic (No target device specified)
|
|
413
|
+
if result is None and isinstance(data, np.ndarray):
|
|
414
|
+
# Default for NumPy is CPU; ensure it stays there to preserve device origin.
|
|
415
|
+
# JAX might otherwise default to placing it on GPU if available.
|
|
416
|
+
cpu_dev = _get_jax_device("cpu")
|
|
417
|
+
if cpu_dev:
|
|
418
|
+
result = jax.device_put(data, cpu_dev)
|
|
419
|
+
|
|
420
|
+
# 4. General case (lists, scalars, or existing JAX arrays)
|
|
421
|
+
if result is None:
|
|
422
|
+
result = jnp.asarray(data)
|
|
423
|
+
|
|
424
|
+
# --- Post-conversion dtype guard ---
|
|
425
|
+
# Ensures the returned array matches the requested dtype, catching edge cases
|
|
426
|
+
# where DLPack, device_put, or JAX x64 mode silently preserve the original precision.
|
|
427
|
+
if target_dtype is not None and hasattr(result, "dtype"):
|
|
428
|
+
jax_target = jnp.dtype(target_dtype)
|
|
429
|
+
if result.dtype != jax_target:
|
|
430
|
+
logger.debug(
|
|
431
|
+
"to_jax: post-conversion dtype mismatch (%s != %s), casting.",
|
|
432
|
+
result.dtype,
|
|
433
|
+
jax_target,
|
|
434
|
+
)
|
|
435
|
+
result = result.astype(jax_target)
|
|
436
|
+
|
|
437
|
+
return result
|
|
438
|
+
|
|
439
|
+
|
|
440
|
+
def from_jax(data: Any) -> ArrayType:
|
|
441
|
+
"""
|
|
442
|
+
Converts a JAX array to a backend-compatible array (NumPy or CuPy).
|
|
443
|
+
|
|
444
|
+
Standardizes on NumPy for CPU/TPU arrays and CuPy for GPU arrays
|
|
445
|
+
to maintain compatibility with the rest of the library. Uses zero-copy
|
|
446
|
+
DLPack transfers for GPU arrays when available.
|
|
447
|
+
|
|
448
|
+
Parameters
|
|
449
|
+
----------
|
|
450
|
+
data : jax.Array
|
|
451
|
+
Input JAX array to convert.
|
|
452
|
+
|
|
453
|
+
Returns
|
|
454
|
+
-------
|
|
455
|
+
array : array_like
|
|
456
|
+
A NumPy array (if on CPU/TPU) or a CuPy array (if on GPU).
|
|
457
|
+
"""
|
|
458
|
+
# Detect platform
|
|
459
|
+
platform = "cpu"
|
|
460
|
+
try:
|
|
461
|
+
# Standard JAX 0.4.x+ device inspection
|
|
462
|
+
if hasattr(data, "device"):
|
|
463
|
+
platform = data.device.platform
|
|
464
|
+
elif hasattr(data, "devices"):
|
|
465
|
+
platform = list(data.devices())[0].platform
|
|
466
|
+
except Exception:
|
|
467
|
+
pass
|
|
468
|
+
|
|
469
|
+
is_gpu = platform in ("cuda", "gpu")
|
|
470
|
+
|
|
471
|
+
if is_gpu and is_cupy_available():
|
|
472
|
+
# Try zero-copy via DLPack to CuPy
|
|
473
|
+
try:
|
|
474
|
+
return cp.from_dlpack(data)
|
|
475
|
+
except Exception as e:
|
|
476
|
+
logger.debug(
|
|
477
|
+
"DLPack transfer from JAX to CuPy failed: %s. Falling back to NumPy conversion.",
|
|
478
|
+
e,
|
|
479
|
+
)
|
|
480
|
+
|
|
481
|
+
if is_gpu and not is_cupy_available():
|
|
482
|
+
logger.warning(
|
|
483
|
+
"JAX array is on GPU, but CuPy is not available. Falling back to NumPy (CPU)."
|
|
484
|
+
)
|
|
485
|
+
|
|
486
|
+
# Convert to numpy (will copy from GPU/TPU if needed)
|
|
487
|
+
return np.asarray(data)
|
|
488
|
+
|
|
489
|
+
|
|
490
|
+
def is_jax_array(data: Any) -> bool:
|
|
491
|
+
"""
|
|
492
|
+
Checks if the given data is a JAX array without eagerly importing JAX.
|
|
493
|
+
|
|
494
|
+
Parameters
|
|
495
|
+
----------
|
|
496
|
+
data : any
|
|
497
|
+
The object to check.
|
|
498
|
+
|
|
499
|
+
Returns
|
|
500
|
+
-------
|
|
501
|
+
bool
|
|
502
|
+
True if `data` is a `jax.Array` instance.
|
|
503
|
+
"""
|
|
504
|
+
jax, _, _ = _get_jax()
|
|
505
|
+
if jax is None:
|
|
506
|
+
return False
|
|
507
|
+
return isinstance(data, jax.Array)
|
|
@@ -0,0 +1,23 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Channel coding / forward error correction (FEC).
|
|
3
|
+
|
|
4
|
+
Scaffold only - no algorithms are implemented yet. This package reserves the
|
|
5
|
+
public layout for the bits-layer neighbour of :mod:`commkit.mapping`:
|
|
6
|
+
encoders turn information bits -> coded bits, which ``mapping.map_bits`` maps to
|
|
7
|
+
symbols; soft decoders consume LLRs produced by ``mapping.compute_llr`` /
|
|
8
|
+
``metrics`` (see :mod:`commkit.coding.base` for the shared hard/soft
|
|
9
|
+
interface conventions).
|
|
10
|
+
|
|
11
|
+
Each module below is an importable placeholder carrying only its scope
|
|
12
|
+
docstring. The package is intentionally **absent from the top-level
|
|
13
|
+
``commkit`` public surface** until at least one real encode/decode entry
|
|
14
|
+
point exists; do not add ``from . import coding`` to ``commkit/__init__.py``
|
|
15
|
+
before then.
|
|
16
|
+
|
|
17
|
+
Promotion path (apply the §7.4 size+cohesion trigger as real code lands):
|
|
18
|
+
``ldpc`` and ``polar`` graduate to ``{construction,decode}`` subpackages, and
|
|
19
|
+
the algebraic block codes (``hamming``/``bch``/``reed_solomon``) collect under a
|
|
20
|
+
``block/`` subpackage once they share enough ``galois`` machinery.
|
|
21
|
+
"""
|
|
22
|
+
|
|
23
|
+
__all__: list[str] = []
|
commkit/coding/base.py
ADDED
|
@@ -0,0 +1,17 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Shared coding interfaces and conventions (scaffold).
|
|
3
|
+
|
|
4
|
+
Defines the contracts every code in this package will share - ``Encoder`` /
|
|
5
|
+
``Decoder`` protocols, a ``CodewordResult`` dataclass, and the hard/soft
|
|
6
|
+
interface conventions - so that algebraic, convolutional, and modern
|
|
7
|
+
capacity-approaching codes present a uniform surface.
|
|
8
|
+
|
|
9
|
+
Soft (LLR) convention - agreed with :func:`commkit.mapping.compute_llr`:
|
|
10
|
+
|
|
11
|
+
LLR_k = log[ P(b_k = 0 | r) / P(b_k = 1 | r) ]
|
|
12
|
+
|
|
13
|
+
so a **positive** LLR means bit 0 is more likely, a **negative** LLR means bit
|
|
14
|
+
1, and the magnitude is the confidence. Soft-input decoders here consume LLRs
|
|
15
|
+
in exactly this sign/scale; soft-output decoders (BCJR, belief propagation)
|
|
16
|
+
emit them in the same convention. No implementation yet.
|
|
17
|
+
"""
|
commkit/coding/bch.py
ADDED
commkit/coding/crc.py
ADDED
commkit/coding/galois.py
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Finite-field arithmetic (scaffold).
|
|
3
|
+
|
|
4
|
+
GF(2) / GF(2^m) element and polynomial arithmetic - the algebraic foundation
|
|
5
|
+
shared by the BCH and Reed-Solomon codes. GF lookup tables are small and stay
|
|
6
|
+
host-side (NumPy); backend dispatch is reserved for the hot array math only.
|
|
7
|
+
No implementation yet.
|
|
8
|
+
"""
|
commkit/coding/ldpc.py
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
1
|
+
"""
|
|
2
|
+
LDPC codes (scaffold).
|
|
3
|
+
|
|
4
|
+
Low-density parity-check codes: parity-check-matrix construction
|
|
5
|
+
(Gallager / PEG / quasi-cyclic) and iterative belief-propagation decoding
|
|
6
|
+
(sum-product / min-sum / layered). Promote to ``ldpc/{construction,decode}``
|
|
7
|
+
once both halves outgrow a single file. No implementation yet.
|
|
8
|
+
"""
|
commkit/coding/polar.py
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Polar codes (scaffold).
|
|
3
|
+
|
|
4
|
+
Polar codes: frozen-bit-set construction (Bhattacharyya / Gaussian
|
|
5
|
+
approximation) and successive-cancellation (SC) / SC-list (SCL, CRC-aided)
|
|
6
|
+
decoding. Promote to ``polar/{construction,decode}`` once both halves outgrow
|
|
7
|
+
a single file. No implementation yet.
|
|
8
|
+
"""
|
commkit/coding/turbo.py
ADDED
|
@@ -0,0 +1,8 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Turbo codes (scaffold).
|
|
3
|
+
|
|
4
|
+
Parallel-concatenated convolutional codes with iterative BCJR decoding, reusing
|
|
5
|
+
the constituent soft-output decoder from :mod:`commkit.coding.convolutional`
|
|
6
|
+
and the permutation from :mod:`commkit.coding.interleaving`. No
|
|
7
|
+
implementation yet.
|
|
8
|
+
"""
|
commkit/core/__init__.py
ADDED
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Core data structures and signal factories for commkit.
|
|
3
|
+
|
|
4
|
+
Re-exports the primary containers and generation factories so existing imports
|
|
5
|
+
(``from commkit.core import Signal, Preamble, SingleCarrierFrame``) keep
|
|
6
|
+
working after the split of the former monolithic ``core.py`` into a package.
|
|
7
|
+
|
|
8
|
+
``signal`` is the thin container (no leaf-module dependencies); ``frame`` and
|
|
9
|
+
``generation`` build on it. The generation factories are also re-exported at
|
|
10
|
+
the package top level (``commkit.generate_qam(...)`` etc.).
|
|
11
|
+
"""
|
|
12
|
+
|
|
13
|
+
from .frame import Preamble, SingleCarrierFrame
|
|
14
|
+
from .generation import (
|
|
15
|
+
generate,
|
|
16
|
+
generate_pam,
|
|
17
|
+
generate_psk,
|
|
18
|
+
generate_psqam,
|
|
19
|
+
generate_qam,
|
|
20
|
+
)
|
|
21
|
+
from .signal import Signal
|
|
22
|
+
|
|
23
|
+
__all__ = [
|
|
24
|
+
"Preamble",
|
|
25
|
+
"Signal",
|
|
26
|
+
"SingleCarrierFrame",
|
|
27
|
+
"generate",
|
|
28
|
+
"generate_pam",
|
|
29
|
+
"generate_psk",
|
|
30
|
+
"generate_psqam",
|
|
31
|
+
"generate_qam",
|
|
32
|
+
]
|