parx 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.
- parx/__init__.py +111 -0
- parx/_check.py +14 -0
- parx/_julia_init.py +106 -0
- parx/_lp.py +59 -0
- parx/analysis.py +441 -0
- parx/diagnostics.py +67 -0
- parx/io.py +120 -0
- parx/io_partition.py +152 -0
- parx/julia/Manifest.toml +325 -0
- parx/julia/Project.toml +14 -0
- parx/julia/src/LinearRegions.jl +13 -0
- parx/julia/src/bridge.jl +24 -0
- parx/julia/src/exact.jl +177 -0
- parx/julia/src/lp.jl +91 -0
- parx/julia/src/self_test.jl +29 -0
- parx/julia/src/sparse.jl +75 -0
- parx/juliapkg.json +13 -0
- parx/methods/__init__.py +87 -0
- parx/methods/exact_julia.py +50 -0
- parx/methods/exact_julia_fast.py +56 -0
- parx/methods/exact_python.py +230 -0
- parx/methods/sparse_julia.py +43 -0
- parx/methods/sparse_python.py +79 -0
- parx/network.py +159 -0
- parx/partition.py +210 -0
- parx/precompile.py +65 -0
- parx/region.py +37 -0
- parx/verify.py +179 -0
- parx/viz.py +2009 -0
- parx-0.1.0.dist-info/METADATA +376 -0
- parx-0.1.0.dist-info/RECORD +34 -0
- parx-0.1.0.dist-info/WHEEL +5 -0
- parx-0.1.0.dist-info/licenses/LICENSE +21 -0
- parx-0.1.0.dist-info/top_level.txt +1 -0
parx/__init__.py
ADDED
|
@@ -0,0 +1,111 @@
|
|
|
1
|
+
"""
|
|
2
|
+
parx — POLyhedral Activation Region Xplorer
|
|
3
|
+
|
|
4
|
+
Exactly enumerates the linear regions of ReLU neural networks.
|
|
5
|
+
"""
|
|
6
|
+
|
|
7
|
+
from __future__ import annotations
|
|
8
|
+
|
|
9
|
+
from importlib.metadata import PackageNotFoundError
|
|
10
|
+
from importlib.metadata import version as _pkg_version
|
|
11
|
+
|
|
12
|
+
import numpy as np
|
|
13
|
+
|
|
14
|
+
from parx import (
|
|
15
|
+
methods, # noqa: F401 — trigger method registration
|
|
16
|
+
viz, # noqa: F401 — expose parx.viz
|
|
17
|
+
)
|
|
18
|
+
from parx._check import check_julia
|
|
19
|
+
from parx._julia_init import ensure_julia # noqa: F401
|
|
20
|
+
from parx.analysis import (
|
|
21
|
+
always_active_neurons,
|
|
22
|
+
complexity_profile,
|
|
23
|
+
dead_neurons,
|
|
24
|
+
neuron_activity_rates,
|
|
25
|
+
region_size_summary,
|
|
26
|
+
)
|
|
27
|
+
from parx.io import iter_state_dicts
|
|
28
|
+
from parx.io_partition import load_partition, save_partition
|
|
29
|
+
from parx.methods import get_method, list_methods
|
|
30
|
+
from parx.network import extract_features, load_network
|
|
31
|
+
from parx.partition import Partition
|
|
32
|
+
from parx.precompile import precompile
|
|
33
|
+
from parx.region import Region
|
|
34
|
+
from parx.viz import (
|
|
35
|
+
animate_epochs,
|
|
36
|
+
animate_epochs_video,
|
|
37
|
+
plot_feature_embedding,
|
|
38
|
+
region_palette,
|
|
39
|
+
)
|
|
40
|
+
|
|
41
|
+
check_julia()
|
|
42
|
+
|
|
43
|
+
try:
|
|
44
|
+
__version__ = _pkg_version("parx")
|
|
45
|
+
except PackageNotFoundError: # pragma: no cover - source tree, not installed
|
|
46
|
+
__version__ = "0.0.0+unknown"
|
|
47
|
+
|
|
48
|
+
__all__ = [
|
|
49
|
+
"__version__",
|
|
50
|
+
"always_active_neurons",
|
|
51
|
+
"animate_epochs",
|
|
52
|
+
"animate_epochs_video",
|
|
53
|
+
"complexity_profile",
|
|
54
|
+
"compute_partition",
|
|
55
|
+
"dead_neurons",
|
|
56
|
+
"ensure_julia",
|
|
57
|
+
"extract_features",
|
|
58
|
+
"iter_state_dicts",
|
|
59
|
+
"list_methods",
|
|
60
|
+
"load_network",
|
|
61
|
+
"load_partition",
|
|
62
|
+
"neuron_activity_rates",
|
|
63
|
+
"Partition",
|
|
64
|
+
"plot_feature_embedding",
|
|
65
|
+
"precompile",
|
|
66
|
+
"Region",
|
|
67
|
+
"region_palette",
|
|
68
|
+
"region_size_summary",
|
|
69
|
+
"save_partition",
|
|
70
|
+
]
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def compute_partition(
|
|
74
|
+
source,
|
|
75
|
+
data: np.ndarray,
|
|
76
|
+
*,
|
|
77
|
+
method: str = "sparse_julia",
|
|
78
|
+
include_output_layer: bool = False,
|
|
79
|
+
**method_kwargs,
|
|
80
|
+
) -> Partition:
|
|
81
|
+
"""Find the linear regions of a ReLU network.
|
|
82
|
+
|
|
83
|
+
Parameters
|
|
84
|
+
----------
|
|
85
|
+
source:
|
|
86
|
+
A single network's parameters. Accepted forms: PyTorch ``state_dict``,
|
|
87
|
+
``nn.Module``, or path to a ``.pth`` / ``.h5`` file. For per-epoch
|
|
88
|
+
analyses, iterate over state dicts at the call site (see
|
|
89
|
+
:func:`parx.io.iter_state_dicts`).
|
|
90
|
+
data:
|
|
91
|
+
Input points, shape ``(N, input_dim)``. Sparse methods scan the array
|
|
92
|
+
for activation patterns; exact methods use ``data[0]`` as the DFS
|
|
93
|
+
starting point.
|
|
94
|
+
method:
|
|
95
|
+
Name of a registered region-finding method. Built-ins:
|
|
96
|
+
``"sparse_julia"`` (default), ``"exact_julia"``, ``"sparse_python"``,
|
|
97
|
+
``"exact_python"``. See :func:`parx.list_methods`.
|
|
98
|
+
include_output_layer:
|
|
99
|
+
Include the final linear layer in the partition. Defaults to ``False``
|
|
100
|
+
because only hidden ReLU layers define the polyhedral partition.
|
|
101
|
+
**method_kwargs:
|
|
102
|
+
Forwarded verbatim to the chosen method's function.
|
|
103
|
+
|
|
104
|
+
Returns
|
|
105
|
+
-------
|
|
106
|
+
Partition
|
|
107
|
+
"""
|
|
108
|
+
weights, biases = load_network(source, include_output_layer=include_output_layer)
|
|
109
|
+
fn = get_method(method)
|
|
110
|
+
result = fn(weights, biases, data, **method_kwargs)
|
|
111
|
+
return Partition.from_result(result, weights, biases)
|
parx/_check.py
ADDED
|
@@ -0,0 +1,14 @@
|
|
|
1
|
+
import shutil
|
|
2
|
+
|
|
3
|
+
|
|
4
|
+
def check_julia() -> None:
|
|
5
|
+
"""Raise a helpful error if Julia is not found on PATH."""
|
|
6
|
+
if shutil.which("julia") is None:
|
|
7
|
+
raise RuntimeError(
|
|
8
|
+
"\n\nparx requires Julia to be installed and available on PATH.\n"
|
|
9
|
+
"Install Julia via juliaup (recommended):\n\n"
|
|
10
|
+
" macOS/Linux: curl -fsSL https://install.julialang.org | sh\n"
|
|
11
|
+
" Windows: winget install julia -s msstore\n\n"
|
|
12
|
+
"After installing, restart your terminal and try again.\n"
|
|
13
|
+
"See https://github.com/Johanmkr/parx for more details.\n"
|
|
14
|
+
)
|
parx/_julia_init.py
ADDED
|
@@ -0,0 +1,106 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Handles Julia runtime initialization.
|
|
3
|
+
Import this module once; subsequent calls to ensure_julia() are no-ops.
|
|
4
|
+
"""
|
|
5
|
+
|
|
6
|
+
import json
|
|
7
|
+
import os
|
|
8
|
+
import shutil
|
|
9
|
+
from pathlib import Path
|
|
10
|
+
|
|
11
|
+
_julia_initialized = False
|
|
12
|
+
_jl = None
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def _resolve_juliaup_shim() -> str | None:
|
|
16
|
+
"""Work around a juliapkg bug when Julia was installed via juliaup.
|
|
17
|
+
|
|
18
|
+
juliaup puts a launcher shim (not a real Julia binary) at
|
|
19
|
+
``~/.juliaup/bin/julia``; the actual per-version installs live under
|
|
20
|
+
``~/.julia/juliaup/julia-<version>+.../``. juliapkg tries to
|
|
21
|
+
opportunistically upgrade to the newest Julia release on every resolve;
|
|
22
|
+
when that install attempt fails, it silently falls back to the shim path
|
|
23
|
+
itself instead of one of the already-installed real binaries. juliacall
|
|
24
|
+
then derives the system-image path relative to the shim's directory,
|
|
25
|
+
which never has a ``lib/julia/sys.so`` next to it, and the first Julia
|
|
26
|
+
call crashes with "could not load library ... sys.so ... No such file or
|
|
27
|
+
directory". See CONTRIBUTING.md's Troubleshooting section.
|
|
28
|
+
|
|
29
|
+
Returns the absolute path to the real Julia binary juliaup's default
|
|
30
|
+
channel points at, or ``None`` if ``julia`` isn't on PATH, isn't a
|
|
31
|
+
juliaup shim, or the real binary can't be resolved for any reason — in
|
|
32
|
+
which case juliapkg's normal resolution is left alone.
|
|
33
|
+
"""
|
|
34
|
+
julia_on_path = shutil.which("julia")
|
|
35
|
+
if julia_on_path is None:
|
|
36
|
+
return None
|
|
37
|
+
|
|
38
|
+
# juliaup's shim is typically a symlink to a binary literally named
|
|
39
|
+
# "julialauncher"; a plain (non-juliaup) Julia install is not. Some
|
|
40
|
+
# installs may ship a separate "julia" launcher alongside "julialauncher",
|
|
41
|
+
# so also accept a sibling julialauncher in the same directory.
|
|
42
|
+
real_target = os.path.realpath(julia_on_path)
|
|
43
|
+
launcher_names = ("julialauncher", "julialauncher.exe")
|
|
44
|
+
if os.path.basename(real_target) not in launcher_names:
|
|
45
|
+
shim_dir = os.path.dirname(julia_on_path)
|
|
46
|
+
ext = ".exe" if os.name == "nt" else ""
|
|
47
|
+
candidate = os.path.join(shim_dir, "julialauncher" + ext)
|
|
48
|
+
if not os.path.isfile(candidate):
|
|
49
|
+
return None
|
|
50
|
+
real_target = candidate
|
|
51
|
+
|
|
52
|
+
# juliaup does not follow JULIA_DEPOT_PATH, but defines its own
|
|
53
|
+
# override for ~/.julia (matches juliapkg's own lookup).
|
|
54
|
+
depot = os.environ.get("JULIAUP_DEPOT_PATH") or os.path.join(
|
|
55
|
+
os.path.expanduser("~"), ".julia"
|
|
56
|
+
)
|
|
57
|
+
judir = os.path.join(depot, "juliaup")
|
|
58
|
+
try:
|
|
59
|
+
with open(os.path.join(judir, "juliaup.json")) as f:
|
|
60
|
+
meta = json.load(f)
|
|
61
|
+
version = meta["InstalledChannels"][meta["Default"]]["Version"]
|
|
62
|
+
info = meta["InstalledVersions"][version]
|
|
63
|
+
if "BinaryPath" in info:
|
|
64
|
+
exe = os.path.join(judir, info["BinaryPath"])
|
|
65
|
+
else:
|
|
66
|
+
ext = ".exe" if os.name == "nt" else ""
|
|
67
|
+
exe = os.path.join(judir, info["Path"], "bin", "julia" + ext)
|
|
68
|
+
exe = os.path.abspath(exe)
|
|
69
|
+
except (OSError, KeyError, ValueError, json.JSONDecodeError):
|
|
70
|
+
return None
|
|
71
|
+
|
|
72
|
+
return exe if os.path.isfile(exe) else None
|
|
73
|
+
|
|
74
|
+
|
|
75
|
+
def ensure_julia():
|
|
76
|
+
"""Initialize the Julia runtime and load the LinearRegions module.
|
|
77
|
+
|
|
78
|
+
juliacall/juliapkg manages its own Julia project environment, so we load
|
|
79
|
+
our LinearRegions module via include() rather than registering it as a
|
|
80
|
+
package. Safe to call multiple times — only runs once per process.
|
|
81
|
+
"""
|
|
82
|
+
global _julia_initialized, _jl
|
|
83
|
+
|
|
84
|
+
if _julia_initialized:
|
|
85
|
+
return _jl
|
|
86
|
+
|
|
87
|
+
os.environ.setdefault("JULIA_NUM_THREADS", "auto")
|
|
88
|
+
# Let Julia own signal handling so its GC threads don't conflict with
|
|
89
|
+
# Python's signal machinery. Must be set before juliacall is imported.
|
|
90
|
+
os.environ.setdefault("PYTHON_JULIACALL_HANDLE_SIGNALS", "yes")
|
|
91
|
+
|
|
92
|
+
# Work around the juliaup-shim juliapkg bug (see _resolve_juliaup_shim's
|
|
93
|
+
# docstring). setdefault so an explicit user override always wins.
|
|
94
|
+
juliaup_exe = _resolve_juliaup_shim()
|
|
95
|
+
if juliaup_exe is not None:
|
|
96
|
+
os.environ.setdefault("PYTHON_JULIAPKG_EXE", juliaup_exe)
|
|
97
|
+
|
|
98
|
+
from juliacall import Main as jl
|
|
99
|
+
|
|
100
|
+
julia_file = Path(__file__).parent / "julia" / "src" / "LinearRegions.jl"
|
|
101
|
+
jl.seval(f'include("{julia_file.as_posix()}")')
|
|
102
|
+
|
|
103
|
+
_jl = jl
|
|
104
|
+
_julia_initialized = True
|
|
105
|
+
|
|
106
|
+
return _jl
|
parx/_lp.py
ADDED
|
@@ -0,0 +1,59 @@
|
|
|
1
|
+
"""Shared LP primitives used by region finders and partition verification."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
from scipy.optimize import linprog
|
|
7
|
+
|
|
8
|
+
_ZERO_NORM_TOL = 1e-10
|
|
9
|
+
_SLACK_TOL = 1e-8
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def chebyshev_center(
|
|
13
|
+
D: np.ndarray,
|
|
14
|
+
g: np.ndarray,
|
|
15
|
+
*,
|
|
16
|
+
max_radius: float = 1e3,
|
|
17
|
+
) -> tuple[np.ndarray | None, float]:
|
|
18
|
+
"""Chebyshev centre of ``{x : D x ≤ g}``.
|
|
19
|
+
|
|
20
|
+
Returns ``(x_interior, radius)`` or ``(None, 0.0)`` when empty / degenerate.
|
|
21
|
+
|
|
22
|
+
Zero-norm rows are filtered (they would make the LP unbounded); a
|
|
23
|
+
zero-norm row with strictly negative ``g[i]`` makes the system infeasible
|
|
24
|
+
(returns ``(None, 0.0)``). Unbounded polytopes are detected by the radius
|
|
25
|
+
hitting ``max_radius``; callers can compare against that value to decide
|
|
26
|
+
whether to treat the region as unbounded.
|
|
27
|
+
"""
|
|
28
|
+
m, n = D.shape
|
|
29
|
+
if m == 0:
|
|
30
|
+
return np.zeros(n), float("inf")
|
|
31
|
+
|
|
32
|
+
row_norms = np.linalg.norm(D, axis=1)
|
|
33
|
+
valid = row_norms > _ZERO_NORM_TOL
|
|
34
|
+
|
|
35
|
+
if np.any(g[~valid] < -_ZERO_NORM_TOL):
|
|
36
|
+
return None, 0.0
|
|
37
|
+
if not valid.any():
|
|
38
|
+
return np.zeros(n), float("inf")
|
|
39
|
+
|
|
40
|
+
D_v = D[valid]
|
|
41
|
+
g_v = g[valid]
|
|
42
|
+
norms_v = row_norms[valid]
|
|
43
|
+
|
|
44
|
+
# Variables = (x_1, …, x_n, r). Constraint i: D_v[i] · x + ‖D_v[i]‖ · r ≤ g_v[i].
|
|
45
|
+
A_ub = np.hstack([D_v, norms_v[:, None]])
|
|
46
|
+
b_ub = g_v
|
|
47
|
+
c = np.zeros(n + 1)
|
|
48
|
+
c[-1] = -1.0 # maximise r ↔ minimise -r
|
|
49
|
+
bounds = [(None, None)] * n + [(0.0, max_radius)]
|
|
50
|
+
|
|
51
|
+
res = linprog(c, A_ub=A_ub, b_ub=b_ub, bounds=bounds, method="highs")
|
|
52
|
+
if not res.success or res.x is None:
|
|
53
|
+
return None, 0.0
|
|
54
|
+
|
|
55
|
+
x = res.x[:n]
|
|
56
|
+
r = float(res.x[-1])
|
|
57
|
+
if r < _SLACK_TOL:
|
|
58
|
+
return None, 0.0
|
|
59
|
+
return x, r
|