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 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