diff --git a/.github/workflows/docs.yml b/.github/workflows/docs.yml index 8fde7b70..e127756d 100644 --- a/.github/workflows/docs.yml +++ b/.github/workflows/docs.yml @@ -12,6 +12,7 @@ on: pull_request: branches: - develop + - epic/324-pyodide-browser paths: - 'docs/**' - 'src/**' @@ -39,9 +40,9 @@ jobs: - name: Install dependencies run: | uv python install 3.12 - # The torch extra is required: gallery examples call ICLabel, which - # raises ImportError without it rather than degrading to a skip. - uv sync --python 3.12 --extra docs --extra torch + # ICLabel runs through the packaged ONNX network and raises ImportError + # without its runtime extra rather than degrading to a skip. + uv sync --python 3.12 --extra docs --extra iclabel - name: Build HTML documentation run: | diff --git a/.github/workflows/release.yml b/.github/workflows/release.yml index 9a6f2c05..332b1e89 100644 --- a/.github/workflows/release.yml +++ b/.github/workflows/release.yml @@ -135,8 +135,32 @@ jobs: fi # Packaged data the runtime needs. - if [ "$(unzip -l "$whl" | grep -c 'netICL.mat' || true)" = "0" ]; then - echo "::error::ICLabel weights (netICL.mat) missing from the wheel"; exit 1 + if [ "$(unzip -l "$whl" | grep -c 'eegprep/plugins/ICLabel/iclabel.onnx' || true)" = "0" ]; then + echo "::error::ICLabel network (iclabel.onnx) missing from the wheel"; exit 1 + fi + if [ "$(unzip -l "$whl" | grep -c 'netICL.mat' || true)" != "0" ]; then + echo "::error::ICLabel weights (netICL.mat) should not ship in the wheel; export iclabel.onnx instead"; exit 1 + fi + packaged_model_sha256=$(unzip -p "$whl" eegprep/plugins/ICLabel/iclabel.onnx | sha256sum | cut -d' ' -f1) + source_model_sha256=$(sha256sum src/eegprep/plugins/ICLabel/iclabel.onnx | cut -d' ' -f1) + if [ "$packaged_model_sha256" != "$source_model_sha256" ]; then + echo "::error::Packaged ICLabel network differs from the source artifact"; exit 1 + fi + sdist_model_path=$(tar tzf "$sdist" | grep '/src/eegprep/plugins/ICLabel/iclabel.onnx$' | head -n 1) + if [ -z "$sdist_model_path" ]; then + echo "::error::ICLabel network (iclabel.onnx) missing from the sdist"; exit 1 + fi + if [ "$(tar tzf "$sdist" | grep -c '/src/eegprep/plugins/ICLabel/netICL.mat$' || true)" != "0" ]; then + echo "::error::ICLabel weights (netICL.mat) should not ship in the sdist; export iclabel.onnx instead"; exit 1 + fi + sdist_model_sha256=$(tar xOf "$sdist" "$sdist_model_path" | sha256sum | cut -d' ' -f1) + if [ "$sdist_model_sha256" != "$source_model_sha256" ]; then + echo "::error::Sdist ICLabel network differs from the source artifact"; exit 1 + fi + report_model_sha256=$(python3 -c 'import json; print(json.load(open("tools/iclabel/parity_report.json"))["candidates"]["weight_only"]["sha256"])') + report_default_artifact=$(python3 -c 'import json; print(json.load(open("tools/iclabel/parity_report.json"))["default_artifact"])') + if [ "$report_default_artifact" != "iclabel_int8_weight_only.onnx" ] || [ "$source_model_sha256" != "$report_model_sha256" ]; then + echo "::error::Source ICLabel network does not match the committed parity report"; exit 1 fi if [ "$(unzip -l "$whl" | grep -c 'resources/help/.*\.md' || true)" = "0" ]; then echo "::error::Packaged help resources missing from the wheel"; exit 1 @@ -149,19 +173,69 @@ jobs: - name: Smoke test the built wheel run: | + # shellcheck disable=SC1087 set -euo pipefail uv venv --python 3.12 /tmp/relcheck - VIRTUAL_ENV=/tmp/relcheck uv pip install dist/*.whl - VIRTUAL_ENV=/tmp/relcheck uv run --no-project python -c " + wheel=$(ls dist/*.whl) + VIRTUAL_ENV=/tmp/relcheck uv pip install "${wheel}[iclabel]" "onnxruntime==1.30.0" + VIRTUAL_ENV=/tmp/relcheck uv run --no-project python - <<'PY' import eegprep + from importlib.resources import files + import numpy as np + print('installed', eegprep.__version__) + model = files('eegprep.plugins.ICLabel').joinpath('iclabel.onnx') + assert model.is_file(), 'installed ICLabel model is missing' + assert model.read_bytes(), 'installed ICLabel model is empty' + from eegprep.functions.popfunc.pop_loadset import pop_loadset + from eegprep.plugins.ICLabel.pop_iclabel import pop_iclabel + + eeg = pop_loadset('sample_data/eeglab_data_with_ica_tmp.set') + labeled = pop_iclabel(eeg, 'default') + probabilities = np.asarray(labeled['etc']['ic_classification']['ICLabel']['classifications']) + assert probabilities.ndim == 2 and probabilities.shape.__getitem__(1) == 7 + assert np.isfinite(probabilities).all() + print('installed-wheel ICLabel smoke test passed') + from eegprep.plugins.clean_rawdata.private.yulewalk import yulewalk + f = np.concatenate((np.array((0,2,3,13,16,40,80))*2.0/258, (1.0,))) + B, A = yulewalk(8, f, np.array((3,.75,.33,.33,1,1,3,3), float)) + assert np.all(np.abs(np.roots(A)) < 1), 'unstable ASR shaping filter' + print('ASR shaping filter designs and is stable at 258 Hz') + PY + + - name: Smoke test the built sdist + run: | + # shellcheck disable=SC1087 + set -euo pipefail + uv venv --python 3.12 /tmp/relcheck-sdist + sdist=$(ls dist/*.tar.gz) + VIRTUAL_ENV=/tmp/relcheck-sdist uv pip install "${sdist}[iclabel]" "onnxruntime==1.30.0" + VIRTUAL_ENV=/tmp/relcheck-sdist uv run --no-project python - <<'PY' + import eegprep + from importlib.resources import files import numpy as np - f = np.concatenate((np.array([0,2,3,13,16,40,80])*2.0/258, [1.0])) - B, A = yulewalk(8, f, np.array([3,.75,.33,.33,1,1,3,3], float)) + + print('installed from sdist', eegprep.__version__) + model = files('eegprep.plugins.ICLabel').joinpath('iclabel.onnx') + assert model.is_file(), 'installed ICLabel model is missing' + assert model.read_bytes(), 'installed ICLabel model is empty' + from eegprep.functions.popfunc.pop_loadset import pop_loadset + from eegprep.plugins.ICLabel.pop_iclabel import pop_iclabel + + eeg = pop_loadset('sample_data/eeglab_data_with_ica_tmp.set') + labeled = pop_iclabel(eeg, 'default') + probabilities = np.asarray(labeled['etc']['ic_classification']['ICLabel']['classifications']) + assert probabilities.ndim == 2 and probabilities.shape.__getitem__(1) == 7 + assert np.isfinite(probabilities).all() + print('installed-sdist ICLabel smoke test passed') + + from eegprep.plugins.clean_rawdata.private.yulewalk import yulewalk + f = np.concatenate((np.array((0,2,3,13,16,40,80))*2.0/258, (1.0,))) + B, A = yulewalk(8, f, np.array((3,.75,.33,.33,1,1,3,3), float)) assert np.all(np.abs(np.roots(A)) < 1), 'unstable ASR shaping filter' print('ASR shaping filter designs and is stable at 258 Hz') - " + PY - uses: actions/upload-artifact@v4 with: diff --git a/.github/workflows/test.yml b/.github/workflows/test.yml index 60afe629..6fbe9be9 100644 --- a/.github/workflows/test.yml +++ b/.github/workflows/test.yml @@ -2,9 +2,9 @@ name: Tests on: push: - branches: [ master, develop ] + branches: [ master, develop, epic/324-pyodide-browser ] pull_request: - branches: [ master, develop ] + branches: [ master, develop, epic/324-pyodide-browser ] env: # Set MATLAB batch licensing token for private repos or when using MATLAB Engine @@ -142,7 +142,7 @@ jobs: - name: Install dependencies run: | - uv sync --python ${{ matrix.python-version }} --all-extras --group dev + uv sync --python "${{ matrix.python-version }}" --all-extras --group dev - name: Set up MATLAB if: steps.matlab-availability.outputs.available == 'true' @@ -165,8 +165,8 @@ jobs: print(candidates[-1]) PY ) - echo "MATLAB_ROOT=$MATLAB_ROOT" >> $GITHUB_ENV - echo "$MATLAB_ROOT/bin" >> $GITHUB_PATH + echo "MATLAB_ROOT=$MATLAB_ROOT" >> "$GITHUB_ENV" + echo "$MATLAB_ROOT/bin" >> "$GITHUB_PATH" - name: Install Python MATLAB Engine if: steps.matlab-availability.outputs.available == 'true' @@ -229,3 +229,156 @@ jobs: if: always() run: | uv pip list + + pyodide: + name: Pyodide harness, ICA benchmark, and ICLabel browser parity + runs-on: ubuntu-22.04 + steps: + - uses: actions/checkout@v4 + + - name: Install uv + uses: astral-sh/setup-uv@v5 + with: + enable-cache: true + cache-dependency-glob: "uv.lock" + + - name: Install Node.js + uses: actions/setup-node@v4 + with: + node-version: 22 + + - name: Install Python 3.12 + run: | + uv python install 3.12 + + - name: Install benchmark dependencies + run: | + uv sync --python 3.12 --group dev --extra iclabel --extra torch + + - name: Build EEGPrep wheel from the working tree + run: | + mkdir -p pyodide-artifacts + uv build --wheel --out-dir pyodide-artifacts + + - name: Regenerate and validate ICLabel quantized artifacts + run: | + set -euo pipefail + candidate_dir="$RUNNER_TEMP/iclabel-quantized" + report_path="$candidate_dir/parity_report.json" + uv run --no-sync python tools/iclabel/quantize_iclabel_onnx.py \ + --output-dir "$candidate_dir" \ + --report "$report_path" \ + > pyodide-artifacts/iclabel-quantization.log + uv run --no-sync python - "$report_path" "$candidate_dir/iclabel_int8_weight_only.onnx" <<'PY' + import json + from pathlib import Path + import sys + + report = json.loads(Path(sys.argv[1]).read_text()) + selected = report["default_artifact"] + if selected != "iclabel_int8_weight_only.onnx": + raise SystemExit(f"unexpected selected ICLabel artifact: {selected}") + if not report["candidates"]["weight_only"]["gate_pass"]: + raise SystemExit("regenerated weight-only ICLabel artifact failed its parity gate") + generated = Path(sys.argv[2]).read_bytes() + shipped = Path("src/eegprep/plugins/ICLabel/iclabel.onnx").read_bytes() + if generated != shipped: + raise SystemExit("shipped iclabel.onnx differs from the regenerated selected artifact") + committed = json.loads(Path("tools/iclabel/parity_report.json").read_text()) + if report["candidates"]["weight_only"]["sha256"] != committed["candidates"]["weight_only"]["sha256"]: + raise SystemExit("committed ICLabel parity report is stale relative to regenerated artifact") + PY + + - name: Build the locked docopt wheel + run: | + uv run --no-sync python tools/pyodide/prepare_docopt_wheel.py \ + --output-dir pyodide-artifacts/docopt + + - name: Validate Node runner argument handling + run: | + if output=$(node tools/pyodide/run_pyodide.mjs --not-a-real-option 2>&1); then + echo "The runner accepted an invalid option" >&2 + exit 1 + fi + grep -q "Usage:" <<< "$output" + + - name: Run Pyodide sample-data smoke test + run: | + set -o pipefail + tools/pyodide/run_pyodide.sh \ + --wheel "$(find pyodide-artifacts -maxdepth 1 -type f -name '*.whl' -print -quit)" \ + --docopt-wheel "$(find pyodide-artifacts/docopt -maxdepth 1 -type f -name 'docopt-0.6.2-*.whl' -print -quit)" \ + --script tools/pyodide/smoke.py \ + --sample-data-dir sample_data \ + 2>&1 | tee pyodide-artifacts/pyodide-smoke.log + + - name: Run native ICLabel parity reference + env: + MNE_CONFIGDIR: ${{ runner.temp }}/mne + run: | + uv run --no-sync python tools/pyodide/iclabel_parity.py \ + --platform native \ + --sample-data-dir sample_data \ + > pyodide-artifacts/iclabel-native.json + + - name: Run browser ICLabel parity through Pyodide + run: | + set -o pipefail + tools/pyodide/run_pyodide.sh \ + --iclabel-web \ + --wheel "$(find pyodide-artifacts -maxdepth 1 -type f -name '*.whl' -print -quit)" \ + --docopt-wheel "$(find pyodide-artifacts/docopt -maxdepth 1 -type f -name 'docopt-0.6.2-*.whl' -print -quit)" \ + --script tools/pyodide/iclabel_parity.py \ + --sample-data-dir sample_data \ + --output pyodide-artifacts/iclabel-pyodide.json \ + -- --platform pyodide \ + 2>&1 | tee pyodide-artifacts/iclabel-pyodide.log + + - name: Compare ICLabel browser parity + run: | + set -o pipefail + uv run --no-sync python tools/pyodide/compare_iclabel.py \ + --native pyodide-artifacts/iclabel-native.json \ + --pyodide pyodide-artifacts/iclabel-pyodide.json \ + --report pyodide-artifacts/iclabel-parity.md \ + | tee pyodide-artifacts/iclabel-parity.json + + - name: Run native ICA and matmul benchmark + env: + MNE_CONFIGDIR: ${{ runner.temp }}/mne + OMP_NUM_THREADS: "1" + OPENBLAS_NUM_THREADS: "1" + MKL_NUM_THREADS: "1" + VECLIB_MAXIMUM_THREADS: "1" + NUMEXPR_NUM_THREADS: "1" + run: | + uv run --no-sync python tools/pyodide/benchmark.py \ + --platform native \ + > pyodide-artifacts/benchmark-native.json + + - name: Run Pyodide ICA and matmul benchmark + run: | + set -o pipefail + tools/pyodide/run_pyodide.sh \ + --wheel "$(find pyodide-artifacts -maxdepth 1 -type f -name '*.whl' -print -quit)" \ + --docopt-wheel "$(find pyodide-artifacts/docopt -maxdepth 1 -type f -name 'docopt-0.6.2-*.whl' -print -quit)" \ + --script tools/pyodide/benchmark.py \ + --output pyodide-artifacts/benchmark-pyodide.json \ + -- --platform pyodide \ + 2>&1 | tee pyodide-artifacts/benchmark-pyodide.log + + - name: Compare benchmark gates + run: | + set -o pipefail + uv run --no-sync python tools/pyodide/compare_benchmarks.py \ + --native pyodide-artifacts/benchmark-native.json \ + --pyodide pyodide-artifacts/benchmark-pyodide.json \ + --report pyodide-artifacts/benchmark-comparison.md \ + | tee pyodide-artifacts/benchmark-comparison.json + + - name: Upload Pyodide raw output + if: always() + uses: actions/upload-artifact@v4 + with: + name: pyodide-phase-2-through-phase-6-output + path: pyodide-artifacts diff --git a/AGENTS.md b/AGENTS.md index 3c8ca38f..d0168132 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -14,7 +14,7 @@ Primary references: - `src/eegprep/functions/sigprocfunc/`: EEGLAB-style low-level signal processing functions such as `runica.py`, `runamica.py`, `topoplot.py`, `epoch.py`, and `eegrej.py`. - `src/eegprep/plugins/clean_rawdata/`: Python ports of the EEGLAB clean_rawdata plugin, including `clean_*` and ASR modules. - `src/eegprep/plugins/clean_rawdata/private/`: ports of clean_rawdata private helpers such as `fit_eeg_distribution`, `geometric_median`, FIR helpers, covariance helpers, and spherical-spline interpolation. -- `src/eegprep/plugins/ICLabel/`: Python ports of the EEGLAB ICLabel plugin and bundled `netICL.mat`. +- `src/eegprep/plugins/ICLabel/`: Python ports of the EEGLAB ICLabel plugin and the bundled `iclabel.onnx` network (exported from `netICL.mat` via `tools/iclabel/export_iclabel_onnx.py`). - `src/eegprep/plugins/firfilt/`: Python ports of the EEGLAB firfilt plugin helpers. - `src/eegprep/functions/miscfunc/`: EEGLAB-style miscellaneous helpers, including format conversion and numerical utilities. - `src/eegprep/functions/eegobj/`: Python counterpart to EEGLAB's `functions/@eegobj/`. diff --git a/README.md b/README.md index 1afb579e..9f9ca6cc 100644 --- a/README.md +++ b/README.md @@ -50,8 +50,8 @@ lean package with: uv add eegprep ``` -To install all optional extras, including GUI, console, docs, and classifier -dependencies, use: +To install all optional extras, including GUI, console, docs, classifier, +MATLAB/Octave parity, and system-aware job scheduling dependencies, use: ```bash uv add "eegprep[all]" diff --git a/docs/parity/eeglab_final_parity_matrix.json b/docs/parity/eeglab_final_parity_matrix.json index 0fc63901..cae14064 100644 --- a/docs/parity/eeglab_final_parity_matrix.json +++ b/docs/parity/eeglab_final_parity_matrix.json @@ -441,7 +441,7 @@ "status": "implemented", "responsible_phase": "phase_6", "phase_issue": "#163 ICLabel, viewprops, and visual diagnostics completion", - "rationale": "EEGPrep has standalone ICLabel feature extraction, default netICL classification, ICLabel flagging, label statistics, menu wiring, and history-capable pop wrappers. The default Python network is packaged; EEGLAB lite/beta artifacts are explicit MATLAB/Octave passthrough choices and raise a clear standalone limitation instead of being silently ignored.", + "rationale": "EEGPrep has standalone ICLabel feature extraction, default network classification, ICLabel flagging, label statistics, menu wiring, and history-capable pop wrappers. The default network is packaged as iclabel.onnx and classified through onnxruntime; EEGLAB lite/beta artifacts are explicit MATLAB/Octave passthrough choices and raise a clear standalone limitation instead of being silently ignored.", "user_facing_surface": [ "Tools > Classify components using ICLabel", "iclabel", diff --git a/docs/source/api/ica_and_components.rst b/docs/source/api/ica_and_components.rst index 36abac43..1060dc8d 100644 --- a/docs/source/api/ica_and_components.rst +++ b/docs/source/api/ica_and_components.rst @@ -34,6 +34,7 @@ return zero-based component and sample indices. eegprep.icaproj eegprep.icavar eegprep.iclabel + eegprep.iclabel_async eegprep.optimal_kmeans eegprep.picard eegprep.posact diff --git a/docs/source/api/interactive_pop_workflows.rst b/docs/source/api/interactive_pop_workflows.rst index 193624d7..3f4aea70 100644 --- a/docs/source/api/interactive_pop_workflows.rst +++ b/docs/source/api/interactive_pop_workflows.rst @@ -111,6 +111,7 @@ ICA and Components eegprep.pop_expica eegprep.pop_icathresh eegprep.pop_iclabel + eegprep.pop_iclabel_async eegprep.pop_prop eegprep.pop_prop_extended eegprep.pop_runica diff --git a/docs/source/changelog.rst b/docs/source/changelog.rst index 55956b04..91acbfc7 100644 --- a/docs/source/changelog.rst +++ b/docs/source/changelog.rst @@ -10,6 +10,50 @@ the `GitHub Releases `_ page. Unreleased ========== +- Added the asynchronous Pyodide/Emscripten ICLabel path through pinned ONNX + Runtime Web. ``iclabel_async`` and ``pop_iclabel_async`` preserve the native + post-processing, console history, and GUI/console session synchronization; + browser execution is gated by native-to-browser classification parity. +- Added a pinned Pyodide 0.29.5 harness and CI gate that installs the working-tree + wheel, runs a continuous sample-data smoke pipeline, and records one-thread + ``runica``/Picard and runica-shaped BLAS benchmarks. See :doc:`pyodide_benchmark`; + Pyodide Web Workers are documented as the browser responsiveness and independent-job + concurrency boundary, not as intra-ICA threading. +- Installing ``eegprep`` no longer pulls ``oct2py``, ``psutil``, or ``pyedflib``. + The Octave parity engine now needs ``eegprep[eeglab]`` and the system-RAM helper in + ``num_jobs_from_reservation`` now needs ``eegprep[sys]``; both raise an ``ImportError`` + naming the extra when it is missing, and ``eegprep[all]`` still installs everything. + ``pyedflib`` was unused by the package and is gone from every published install. + This removes the last base dependencies that have no WebAssembly build, so the base + requirement set can resolve under Pyodide. +- ICLabel now ships a gate-selected weight-only int8 ONNX artifact (2,932,897 + bytes) while retaining a reproducible float32 reference and calibrated int8 + candidate under ``tools/iclabel/artifacts/``. On the frozen, subject-disjoint + 217-component real-data evaluation set, the shipped artifact matched the + float32 teacher on 100% of top-1 labels and 100% of existing + ``pop_icflag`` keep-or-reject decisions, with a maximum probability drift of + 0.01346 and a mean drift of 0.00080. The calibrated candidate measured + 98.1567% and 100%, respectively, but exceeded the 0.015 maximum probability + drift gate. Feature extraction, normalization, augmentation, softmax, class + set, and rejection thresholds are unchanged. +- ``asr_process`` now resolves ``max_mem=None`` to a fixed 64 MB instead of probing free + system RAM through ``psutil``. + This matches the ``maxmem=64`` default that ``asr_calibrate`` and ``clean_asr`` already + use, so the whole ASR pipeline assumes one memory budget and block sizes no longer vary + with the machine's free memory. + ``clean_asr`` already passed 64, so the standard cleaning pipeline is unchanged; only + direct ``asr_process(..., max_mem=None)`` calls see different block splitting, and + because the reconstruction matrix is refreshed on a per-block grid their output changes + accordingly. + Pass ``max_mem`` explicitly to pin the previous behavior. +- ICLabel (``iclabel``/``pop_iclabel``) now classifies the default network through + ``onnxruntime`` instead of torch. The package ships ``iclabel.onnx`` in place of + ``netICL.mat``; install the new ``iclabel`` extra (``eegprep[iclabel]``) to run + classification. torch is now needed only to regenerate the ONNX artifact from + ``netICL.mat`` (``tools/iclabel/export_iclabel_onnx.py``), not to run ICLabel. + The preserved float32 reference remains the probability-parity artifact for the + previous torch and MATLAB paths; the shipped int8 artifact is validated separately + by the frozen top-1 and keep-or-reject semantic gate. - Added a source-driven developer audit for the current EEGLAB test-port project. It discovers all MATLAB wrapper and regression methods from the pinned suite, collects pytest provenance, resolves same-directory scenario diff --git a/docs/source/contributing.rst b/docs/source/contributing.rst index 708939e3..5ff6acca 100644 --- a/docs/source/contributing.rst +++ b/docs/source/contributing.rst @@ -66,7 +66,7 @@ If you only need documentation dependencies, sync the docs extra: .. code-block:: bash - uv sync --extra docs --group dev + uv sync --group dev --extra docs --extra iclabel ``uv sync`` installs: diff --git a/docs/source/development.rst b/docs/source/development.rst index f65b80ad..f09a7ec5 100644 --- a/docs/source/development.rst +++ b/docs/source/development.rst @@ -68,7 +68,7 @@ Install Documentation Dependencies .. code-block:: bash - uv sync --extra docs --group dev + uv sync --group dev --extra docs --extra iclabel This installs: @@ -377,7 +377,7 @@ acceptance criteria: .. code-block:: bash - uv sync --group dev --extra docs + uv sync --group dev --extra docs --extra iclabel uv run --no-sync sphinx-build -b html docs/source docs/_build/html The ``docs/Makefile`` target remains available for local iteration: @@ -578,7 +578,7 @@ Documentation Build Errors .. code-block:: bash - uv sync --extra docs --group dev + uv sync --group dev --extra docs --extra iclabel Git Conflicts ------------- diff --git a/docs/source/examples/plot_reject_artifacts.py b/docs/source/examples/plot_reject_artifacts.py index 57ecd72e..bd41f532 100644 --- a/docs/source/examples/plot_reject_artifacts.py +++ b/docs/source/examples/plot_reject_artifacts.py @@ -149,9 +149,9 @@ def _sample_data_dir() -> Path: figures = pop_topoplot(EEG_ica, 0, [1, 2, 3], "Component maps", [1, 3], 0, "electrodes", "off") print("scalp-map figures:", len(figures)) -# ICLabel requires the torch extra: pip install "eegprep[torch]". Without it -# pop_iclabel raises ImportError rather than degrading, so this example fails -# loudly instead of silently documenting a skipped step. +# ICLabel requires the iclabel extra: pip install "eegprep[iclabel]". Without +# it pop_iclabel raises ImportError rather than degrading, so this example +# fails loudly instead of silently documenting a skipped step. CLASSES = ("Brain", "Muscle", "Eye", "Heart", "Line Noise", "Channel Noise", "Other") EEG_ica, label_com = pop_iclabel(EEG_ica, "default", return_com=True) print(label_com) diff --git a/docs/source/index.rst b/docs/source/index.rst index a4812dae..69e6df8a 100644 --- a/docs/source/index.rst +++ b/docs/source/index.rst @@ -90,6 +90,7 @@ Manual development releasing changelog + pyodide_benchmark Five-Minute Script ================== diff --git a/docs/source/pyodide_benchmark.md b/docs/source/pyodide_benchmark.md new file mode 100644 index 00000000..4e02164c --- /dev/null +++ b/docs/source/pyodide_benchmark.md @@ -0,0 +1,138 @@ +# Pyodide browser execution, Phase 2/3 benchmark, and Phase 5 ICLabel parity + +This page records the Phase 2 gate and the Phase 3 runica backend decision for +the [browser epic](https://github.com/sccn/eegprep/issues/324). +The harness installs a wheel built from the working tree under Pyodide, runs a +real continuous-data smoke pipeline, and measures the ICA and matrix products +that determine whether later browser work is worthwhile. + +## Runtime boundary + +The harness pins [Pyodide 0.29.5](https://pyodide.org/en/stable/usage/faq.html) +and uses Pyodide's `emfs:` local-wheel transport. The package-resolution path +pre-installs a universal `docopt==0.6.2` wheel built from a hash-checked sdist +because the PyPI release is an sdist. The live harness also pins the +Pyodide-compatible pure Python versions `mne==1.10.0`, `sympy==1.14.0`, and +`threadpoolctl==3.6.0` while +leaving EEGPrep's published dependency ranges unchanged. +The smoke install completed successfully with this closure; no additional +sdist-only blocker surfaced in the live harness. + +EEGPrep's ICA implementations do not use MNE. In particular, importing and +running `runica` does not eagerly import MNE. The current published dependency +closure still includes MNE for EEG file and interoperability paths; removing +that browser-install dependency is a separate packaging task. + +Pyodide does not provide pthread support, so one Pyodide instance cannot run +four or eight Python threads for a single ICA call. The benchmark therefore +limits native BLAS to one thread as well. A browser [Web Worker](https://pyodide.org/en/stable/usage/index.html) +is still the correct execution boundary: it keeps the UI responsive, but it +does not make one ICA solve faster. + +The scalable design is a bounded worker pool for independent jobs. EEGPrep +should own the worker-safe operation contract and serialized progress/result +messages. The NEMAR/OSA host should own worker lifecycle, queueing, resource +limits, cancellation policy, and whether a request uses one worker or a pool. +Splitting one existing `runica` or Picard solve across workers is out of scope: +its iterative weights are global state and would require a new synchronized +distributed algorithm and a separate parity gate. + +## Reproducing the gate + +The harness requires `uv`, Node.js 22 or newer, and npm. From the repository +root, build the EEGPrep wheel and run the sample-data smoke test with: + +```bash +uv build --wheel --out-dir pyodide-artifacts +tools/pyodide/run_pyodide.sh \ + --wheel pyodide-artifacts/eegprep-*.whl \ + --script tools/pyodide/smoke.py \ + --sample-data-dir sample_data +``` + +The CI job runs the same smoke test, then benchmarks a fixed `(64, 15000)` +continuous array (64 channels, 60 seconds at 250 Hz). Each algorithm has one +warm-up and three measured runs, with a maximum of 512 ICA iterations and a +fixed seed. The matrix benchmark uses the two runica products: + +* `(64, 64) @ (64, 49)` for block activations; +* `(64, 64) @ (64, 64)` for the square weight update. + +Picard is measured through its underlying `picard(..., return_n_iter=True)` API +with the same algorithmic options as EEGPrep's `eeg_picard` wrapper; benchmark +output is quiet and includes the underlying iteration telemetry. The benchmark +does not change either production ICA implementation. `runica` casts its input +to float64, so Phase 3 selects one matrix-product helper at import time: native +platforms retain NumPy `@`, while Emscripten uses SciPy `dgemm`. Picard remains +unchanged. + +## Measured results + +The following tables are generated from the native and Pyodide JSON reports +with `tools/pyodide/compare_benchmarks.py`. The raw reports and loader logs are +also uploaded by CI as the `pyodide-phase-2-through-phase-6-output` artifact. +That bundle includes the Phase 3 backend-selection benchmark and the Phase 6 +frozen ICLabel quantization report in addition to the Phase 2 harness and Phase +5 browser parity evidence. + +## Phase 5 browser ICLabel + +The same harness can run the packaged default ICLabel model through ONNX +Runtime Web: + +```bash +tools/pyodide/run_pyodide.sh \ + --iclabel-web \ + --wheel pyodide-artifacts/eegprep-*.whl \ + --script tools/pyodide/iclabel_parity.py \ + --sample-data-dir sample_data \ + --output pyodide-artifacts/iclabel-pyodide.json \ + -- --platform pyodide +``` + +The Node host loads the pinned `onnxruntime-web` package, registers the +asynchronous bridge, and passes the packaged `iclabel.onnx` bytes into the +Pyodide runtime. The Python `iclabel_async`/`pop_iclabel_async` APIs never +import native `onnxruntime` on Emscripten. CI compares the browser result with +a native `onnxruntime` reference using `rtol=1e-4` and `atol=1e-5`, the +established ICLabel parity tolerance. The bridge defaults to ONNX Runtime +Web's WASM execution provider with one thread; concurrency across independent +jobs remains a host/Web Worker concern. + +### ICA + +| Algorithm | Native median (s) | Native median iterations | Native converged | Pyodide median (s) | Pyodide median iterations | Pyodide converged | +| --- | ---: | ---: | :---: | ---: | ---: | :---: | +| runica | 84.859703 | 512.0 | false | 204.447027 | 512.0 | false | +| picard | 6.825637 | 511.0 | false | 65.385375 | 511.0 | false | + +### Matrix multiplication + +| runica product | dtype | Native NumPy `@` (s) | Native BLAS (s) | Native speedup | Pyodide NumPy `@` (s) | Pyodide BLAS (s) | Pyodide speedup | +| --- | --- | ---: | ---: | ---: | ---: | ---: | ---: | +| activation | float64 | 0.000012 | 0.000017 | 0.69x | 0.000244 | 0.000099 | 2.46x | +| activation | float32 | 0.000006 | 0.000014 | 0.40x | 0.000242 | 0.000264 | 0.91x | +| weight_update | float64 | 0.000008 | 0.000018 | 0.46x | 0.000332 | 0.000125 | 2.66x | +| weight_update | float32 | 0.000005 | 0.000015 | 0.34x | 0.000317 | 0.000154 | 2.06x | + +### Decisions + +The comparison script records whether Picard remains a safe browser default +and whether SciPy BLAS is at least 1.2x faster than NumPy `@` for every tested +runica product and dtype. Those decisions are intentionally based on measured +convergence and timings, not on the code path alone. + +For this run, Picard was much faster in wall-clock time but emitted its +non-convergence warning at 511 iterations; runica also reached the 512-step +cap. The conservative browser-default gate therefore returns **false** for +Picard. The original all-dtypes BLAS gate also returns **false**: the float32 +activation product was 0.91x with `sgemm`, below the 1.2x threshold. + +Phase 3 adopts the narrower production-relevant decision. Because `runica` +already operates on float64 data, its Emscripten hot-loop products use +`scipy.linalg.blas.dgemm`; the measured float64 speedups were 2.46x for +activation and 2.66x for the weight update. Native execution retains NumPy +`@`, and Picard is not monkeypatched or distributed across workers. The +comparison script intentionally retains the universal all-dtypes gate, so its +`phase3_recommended` field remains **false** even though the float64-only +implementation is now covered and benchmarked. diff --git a/docs/source/releasing.rst b/docs/source/releasing.rst index fca2f58d..1669460f 100644 --- a/docs/source/releasing.rst +++ b/docs/source/releasing.rst @@ -72,9 +72,11 @@ What the workflow does 1. **verify** — ``ruff check``, ``ruff format --check``, ``ty check``, and the full test suite with ``EEGPREP_SKIP_MATLAB=1``. 2. **build** — ``uv build``, then checks that the tag matches ``__version__``, that - the vendored EEGLAB checkout is absent from both artifacts, that packaged data - (ICLabel weights, help resources) is present, that ``twine check --strict`` - passes, and that the built wheel installs and imports. + the vendored EEGLAB checkout is absent from both artifacts, that the wheel and + sdist contain the selected ICLabel network (and no source ``netICL.mat``), and + that both artifacts match the committed Phase 6 parity-report hash. It also + runs ``twine check --strict`` and installs both artifacts with the ICLabel + extra for runtime classification and the ASR filter smoke test. 3. **publish** — uploads to PyPI using Trusted Publishing, so no API token is stored in the repository. 4. **github-release** — creates the GitHub release with the artifacts attached. diff --git a/docs/source/user_guide/ica_rejection.rst b/docs/source/user_guide/ica_rejection.rst index fc4807fa..18e9f3b6 100644 --- a/docs/source/user_guide/ica_rejection.rst +++ b/docs/source/user_guide/ica_rejection.rst @@ -146,9 +146,23 @@ preserves event display state through the ``scroll_event`` option. Network Availability ==================== -The standalone Python engine ships the default ICLabel network -(``netICL.mat``). EEGLAB ``lite`` and ``beta`` network artifacts are not -bundled in the Python package; requesting them with ``engine=None`` raises a -clear limitation. They can still be requested through ``engine="matlab"`` or -``engine="octave"`` when that runtime has an EEGLAB ICLabel checkout with those -artifacts. +The standalone Python engine ships the gate-selected weight-only int8 ICLabel +network as ``iclabel.onnx`` and classifies components through `onnxruntime`; +install the ``iclabel`` extra (``eegprep[iclabel]``) to run it. Feature +extraction, float32 input normalization, four-way augmentation, and output +softmax remain unchanged. On the frozen 217-component, subject-disjoint +evaluation set recorded in ``tools/iclabel/evaluation_manifest.json``, the +shipped artifact agreed with the preserved float32 teacher on all 217 top-1 +labels and all 217 keep-or-reject decisions under the existing +``pop_icflag`` thresholds (100% and 100%), with a maximum probability drift of +0.01346. The calibrated int8 candidate measured 98.1567% top-1 and 100% +keep-or-reject but exceeded the fixed 0.015 probability-drift gate, so the +smaller weight-only candidate is shipped. These are measured parity results on +the frozen set, not a general accuracy claim. The reproducible float32 +reference, both int8 candidates, and the complete report are retained under +``tools/iclabel/``. + +EEGLAB ``lite`` and ``beta`` network artifacts are not bundled in the Python +package; requesting them with ``engine=None`` raises a clear limitation. They +can still be requested through ``engine="matlab"`` or ``engine="octave"`` +when that runtime has an EEGLAB ICLabel checkout with those artifacts. diff --git a/docs/source/user_guide/installation.rst b/docs/source/user_guide/installation.rst index ab4e22ad..399e968e 100644 --- a/docs/source/user_guide/installation.rst +++ b/docs/source/user_guide/installation.rst @@ -56,24 +56,45 @@ To install eegprep from source for development: git clone https://github.com/sccn/eegprep.git cd eegprep - uv sync --group dev + uv sync --group dev --extra iclabel ``uv sync`` creates the project environment, installs EEGPrep in editable mode, and uses ``uv.lock`` for reproducible dependency resolution. The development -environment includes the GUI and console runtime dependencies, so a fresh -checkout can immediately launch ``uv run eegprep-console --full``. +environment includes the GUI, console, and ICLabel runtime dependencies, so a +fresh checkout can immediately launch ``uv run eegprep-console --full`` and +run the ICLabel quickstart. To develop or build documentation from source, include the docs extra: .. code-block:: bash - uv sync --group dev --extra docs + uv sync --group dev --extra docs --extra iclabel Optional Dependencies ===================== eegprep has several optional dependencies that enable additional functionality: +ICLabel (component classification) +----------------------------------- + +``pop_iclabel``/``iclabel`` classify independent components through a +packaged ONNX network run with `onnxruntime`. Install the ``iclabel`` extra +to enable it: + +.. code-block:: bash + + uv add "eegprep[iclabel]" + +PyTorch is not required to run ICLabel classification; it is only needed to +regenerate the preserved float32 reference +``tools/iclabel/artifacts/iclabel_float32.onnx`` from ``netICL.mat`` during +development (see ``tools/iclabel/export_iclabel_onnx.py``). The package's +``iclabel.onnx`` is the Phase 6 gate-selected weight-only int8 artifact (about +2.93 MB); the quantization and frozen-set evaluation workflow is documented in +``tools/iclabel/quantize_iclabel_onnx.py`` and +``tools/iclabel/evaluation_manifest.json``. + PyTorch (for GPU acceleration) ------------------------------ @@ -95,6 +116,38 @@ For CPU-only PyTorch: uv add torch --index-url https://download.pytorch.org/whl/cpu +MATLAB/Octave Parity Helpers +----------------------------- + +``eeglab_eeg_checkset``, ``clean_drifts``, and other helpers in +``eeglabcompat`` can call into a real EEGLAB installation through Octave for +development and parity testing. This path needs ``oct2py``, which is not +installed by default because it pulls in several packages with no +WebAssembly build. Install it with the ``eeglab`` extra: + +.. code-block:: bash + + uv add "eegprep[eeglab]" + +Without this extra, calling ``get_eeglab('OCT')`` raises an ``ImportError`` +naming this extra. + +System-Aware Parallel Jobs +--------------------------- + +``bids_preproc``'s ``ReservePerJob`` argument can size the number of +parallel jobs from available system RAM (for example ``"4GB"``). This needs +``psutil``, which is not installed by default. Install it with the ``sys`` +extra: + +.. code-block:: bash + + uv add "eegprep[sys]" + +Without this extra, a memory-based ``ReservePerJob`` reservation raises an +``ImportError`` naming this extra. CPU-based reservations (for example +``"2CPU"``) do not need it. + AMICA ----- @@ -125,7 +178,7 @@ To build the documentation locally: .. code-block:: bash - uv sync --group dev --extra docs + uv sync --group dev --extra docs --extra iclabel uv run --no-sync sphinx-build -b html docs/source docs/_build/html The ``docs/Makefile`` target is also available: @@ -147,7 +200,7 @@ Or with specific extras: .. code-block:: bash - uv add "eegprep[torch,gui,docs]" + uv add "eegprep[iclabel,gui,docs,eeglab,sys]" Verification ============ diff --git a/docs/source/user_guide/interactive_console.rst b/docs/source/user_guide/interactive_console.rst index 2da7345c..96c2a036 100644 --- a/docs/source/user_guide/interactive_console.rst +++ b/docs/source/user_guide/interactive_console.rst @@ -82,6 +82,32 @@ Assignment-style calls also work: This console behavior is specific to ``eegprep-console``. Normal Python imports keep standard Python semantics, where returned values must be assigned manually. +Browser ICLabel +=============== + +When the console is hosted inside a Pyodide/Emscripten browser runtime, ICLabel +uses the asynchronous ONNX Runtime Web path. Run it with ``await`` so the +result is committed to the same GUI/console session: + +.. code-block:: python + + EEG = await pop_iclabel_async(EEG, "default") + +The browser-specific synchronous calls ``pop_iclabel`` and ``iclabel`` fail +fast with an instruction to use their async counterparts. A browser host must +register the ONNX Runtime Web bridge; placing Pyodide in a Web Worker is +recommended for UI responsiveness. The async history command is recorded with +``await`` and can be replayed with: + +.. code-block:: python + + await eegh(1) + +If the selected dataset state or selection changes while classification is +running, EEGPrep raises ``RuntimeError`` and discards the stale result instead +of committing it to a different dataset. History-only commands do not +invalidate the in-flight classification. + In-place Workspace Edits ======================== diff --git a/docs/source/user_guide/plugins.rst b/docs/source/user_guide/plugins.rst index a0fec8ec..c8dbe65e 100644 --- a/docs/source/user_guide/plugins.rst +++ b/docs/source/user_guide/plugins.rst @@ -126,9 +126,10 @@ Use the lower-level FIRFilt helpers for custom order/window work: ICLabel ======= -The standalone Python engine ships the default ICLabel network. ``lite`` and -``beta`` network requests require MATLAB or Octave with an EEGLAB ICLabel -checkout that provides those artifacts. +The standalone Python engine ships the default ICLabel network as an ONNX +artifact and classifies through `onnxruntime` (``eegprep[iclabel]``). ``lite`` +and ``beta`` network requests require MATLAB or Octave with an EEGLAB +ICLabel checkout that provides those artifacts. .. code-block:: python @@ -136,6 +137,17 @@ checkout that provides those artifacts. stats = eeg_icalabelstat(EEG, threshold=0.9, verbose=False) EEG, com = pop_icflag(EEG, return_com=True) +In a Pyodide/Emscripten browser runtime, use the asynchronous entry point and +await the ONNX Runtime Web-backed operation: + +.. code-block:: python + + EEG, com = await pop_iclabel_async(EEG, "default", return_com=True) + +The browser host supplies the ONNX Runtime Web bridge. The synchronous +``iclabel`` and ``pop_iclabel`` entry points are intentionally unavailable in +that runtime; native Python behavior is unchanged. + Review components visually with ``pop_viewprops`` before removing them. DIPFIT diff --git a/docs/source/user_guide/quickstart.rst b/docs/source/user_guide/quickstart.rst index 71e44dde..49f1037a 100644 --- a/docs/source/user_guide/quickstart.rst +++ b/docs/source/user_guide/quickstart.rst @@ -34,6 +34,10 @@ Five-Minute Python Workflow Run this from the repository root after installing EEGPrep or syncing the source checkout. +The ICA/ICLabel example near the end of this page also needs the optional +runtime dependency. For a source checkout, run ``uv sync --group dev --extra iclabel``; +for an installed package, use ``pip install eegprep[iclabel]``. + .. code-block:: python from pathlib import Path diff --git a/pre-commit.py b/pre-commit.py index 37d49f05..96d7b8fa 100755 --- a/pre-commit.py +++ b/pre-commit.py @@ -106,6 +106,7 @@ "**/*.mp4", "**/*.npy", "**/*.npz", + "**/*.onnx", "**/*.pdf", "**/*.png", "**/*.set", diff --git a/pyproject.toml b/pyproject.toml index cff92f9d..ce88e76c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -38,11 +38,8 @@ dependencies = [ "mne>=1.10.0", "neo>=0.14.2", "numpy>=2.1.0", - "oct2py>=5.5.0", "packaging>=23.0", - "psutil>=7.0.0", "pybids>=0.4", - "pyedflib>=0.1.42", "pyyaml>=6.0", "python-picard>=0.8,<0.9", # 1.14.1 is what Pyodide 0.29.5 ships, so the floor cannot go above it without @@ -54,8 +51,15 @@ dependencies = [ ] [project.optional-dependencies] +# Export-only: regenerates the float32 ICLabel reference from netICL.mat via +# tools/iclabel/export_iclabel_onnx.py. Not required to run ICLabel +# classification; see the iclabel extra for that. torch = [ - "torch>=2.0" + "torch>=2.0", + "onnx>=1.14", +] +iclabel = [ + "onnxruntime>=1.18", ] gui = [ "pyqtgraph>=0.13.7", @@ -76,11 +80,20 @@ docs = [ "sphinx-togglebutton>=0.3.0", "sphinxcontrib-spelling>=7.1.0" ] +eeglab = [ + "oct2py>=5.5.0", +] +sys = [ + "psutil>=7.0.0", +] all = [ "eegprep[torch]", + "eegprep[iclabel]", "eegprep[gui]", "eegprep[console]", "eegprep[docs]", + "eegprep[eeglab]", + "eegprep[sys]", ] [project.scripts] @@ -95,6 +108,9 @@ dev = [ # Published installs still keep GUI/console dependencies behind extras. "ipython>=8.0", "pytest>=8.0", + # Used only to write synthetic EDF fixtures in tests/test_bids_load_frombids_helpers.py; + # eegprep itself reads EDF through mne/neo and never imports pyedflib. + "pyedflib>=0.1.42", "pyqtgraph>=0.13.7", "PySide6>=6.6", "ruff>=0.15.14", @@ -215,15 +231,15 @@ exclude = ["eegprep.eeglab*"] "resources/**/*.xyz", "plugins/**/*.json", "plugins/**/*.locs", - "plugins/**/*.mat", "plugins/**/*.md", + "plugins/**/*.onnx", "plugins/**/*.sfp", "plugins/**/*.tsv", "plugins/**/*.txt", ] [tool.setuptools.exclude-package-data] -"eegprep" = ["bin/*", "eeglab/**"] +"eegprep" = ["bin/*", "eeglab/**", "plugins/ICLabel/netICL.mat"] [tool.setuptools.dynamic] version = { attr = "eegprep.__version__" } diff --git a/src/eegprep/__init__.py b/src/eegprep/__init__.py index 7880350d..9abdd884 100644 --- a/src/eegprep/__init__.py +++ b/src/eegprep/__init__.py @@ -214,6 +214,7 @@ "icaproj": ("eegprep.functions.sigprocfunc.ica_helpers", "icaproj"), "icavar": ("eegprep.functions.sigprocfunc.ica_helpers", "icavar"), "iclabel": ("eegprep.plugins.ICLabel.iclabel", "iclabel"), + "iclabel_async": ("eegprep.plugins.ICLabel.iclabel", "iclabel_async"), "eeg_icflag": ("eegprep.plugins.ICLabel.eeg_icflag", "eeg_icflag"), "inputdlg2": ("eegprep.functions.guifunc.inputdlg2", "inputdlg2"), "importevent": ("eegprep.functions.popfunc.importevent", "importevent"), @@ -339,6 +340,7 @@ "pop_headplot": ("eegprep.functions.popfunc.pop_headplot", "pop_headplot"), "pop_icflag": ("eegprep.plugins.ICLabel.pop_icflag", "pop_icflag"), "pop_iclabel": ("eegprep.plugins.ICLabel.pop_iclabel", "pop_iclabel"), + "pop_iclabel_async": ("eegprep.plugins.ICLabel.pop_iclabel", "pop_iclabel_async"), "pop_icathresh": ("eegprep.functions.popfunc.pop_icathresh", "pop_icathresh"), "pop_jointprob": ("eegprep.functions.popfunc.pop_jointprob", "pop_jointprob"), "pop_kaiserbeta": ("eegprep.plugins.firfilt.pop_kaiserbeta", "pop_kaiserbeta"), diff --git a/src/eegprep/functions/adminfunc/console.py b/src/eegprep/functions/adminfunc/console.py index 5678bf0c..445a1b6e 100644 --- a/src/eegprep/functions/adminfunc/console.py +++ b/src/eegprep/functions/adminfunc/console.py @@ -4,6 +4,7 @@ import argparse import ast +import builtins import contextvars import importlib import inspect @@ -13,14 +14,17 @@ import threading import warnings from collections.abc import Callable, Iterator, Mapping +from dataclasses import dataclass from typing import Any import matplotlib.pyplot as plt +import numpy as np import eegprep from eegprep.functions.adminfunc.eegh import eegh, eegh_find from eegprep.extension_runtime import ExtensionRuntime from eegprep.functions.adminfunc.eeglab import gui +from eegprep.functions.adminfunc.storage import _MutationTrackedArray from eegprep.functions.guifunc.session import EEGPrepSession, follow_dataset_selection, normalize_dataset_indices from eegprep.functions.popfunc.pop_eegplot import eegplot_accept_creates_dataset from eegprep.functions.popfunc.pop_newset import pop_newset @@ -39,6 +43,26 @@ _ALLEEG_ASSIGNMENT_PATTERN = re.compile(r"^\s*ALLEEG\s*=") _PYTHON_IDENTIFIER_PATTERN = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*$") _CONSOLE_COMMAND_EXPORTS = {"pop_newset": pop_newset} +_TRACKED_MISSING = object() +_CANONICAL_ASYNC_IMPORTS = { + "eegprep.plugins.ICLabel.iclabel": {"iclabel_async": "iclabel_async"}, + "eegprep.plugins.ICLabel.pop_iclabel": {"pop_iclabel_async": "pop_iclabel_async"}, +} +_CANONICAL_MODULE_CHILDREN = { + "eegprep": { + "plugins": ("eegprep.plugins", {}), + }, + "eegprep.plugins": { + "ICLabel": ("eegprep.plugins.ICLabel", {}), + }, + "eegprep.plugins.ICLabel": { + "iclabel": ("eegprep.plugins.ICLabel.iclabel", _CANONICAL_ASYNC_IMPORTS["eegprep.plugins.ICLabel.iclabel"]), + "pop_iclabel": ( + "eegprep.plugins.ICLabel.pop_iclabel", + _CANONICAL_ASYNC_IMPORTS["eegprep.plugins.ICLabel.pop_iclabel"], + ), + }, +} _BROWSER_ACCEPT_POP_FUNCTIONS = { "pop_autorej", "pop_eegthresh", @@ -48,6 +72,7 @@ "pop_rejspec", "pop_rejtrend", } +_IN_PLACE_DATASET_POP_FUNCTIONS = frozenset({"pop_saveset"}) class ConsolePopResult: @@ -194,6 +219,298 @@ def __repr__(self) -> str: return f"" +class _TrackedList(list[Any]): + """List view that keeps nested console edits connected to the dataset.""" + + def __init__(self, source: list[Any], on_mutation: Callable[[], None]) -> None: + super().__init__(source) + self._source = source + self._on_mutation = on_mutation + + @property + def source(self) -> list[Any]: + return self._source + + def __getitem__(self, index: Any) -> Any: + value = super().__getitem__(index) + if isinstance(index, slice): + return _TrackedList(value, self._on_mutation) + return _tracked_value(value, self._on_mutation) + + def copy(self) -> _TrackedList: + return _TrackedList(list(self), self._on_mutation) + + def __iter__(self) -> Iterator[Any]: + for value in super().__iter__(): + yield _tracked_value(value, self._on_mutation) + + def _sync(self) -> None: + self._source[:] = list.__iter__(self) + self._on_mutation() + + def __setitem__(self, index: Any, value: Any) -> None: + super().__setitem__(index, _unwrap_tracked(value)) + self._sync() + + def __delitem__(self, index: Any) -> None: + super().__delitem__(index) + self._sync() + + def append(self, value: Any) -> None: + super().append(_unwrap_tracked(value)) + self._sync() + + def extend(self, values: Any) -> None: + super().extend(_unwrap_tracked(value) for value in values) + self._sync() + + def insert(self, index: Any, value: Any) -> None: + super().insert(index, _unwrap_tracked(value)) + self._sync() + + def pop(self, index: Any = -1) -> Any: + value = super().pop(index) + self._sync() + return _tracked_value(value, self._on_mutation) + + def remove(self, value: Any) -> None: + super().remove(_unwrap_tracked(value)) + self._sync() + + def clear(self) -> None: + super().clear() + self._sync() + + def reverse(self) -> None: + super().reverse() + self._sync() + + def sort(self, *args: Any, **kwargs: Any) -> None: + super().sort(*args, **kwargs) + self._sync() + + def __iadd__(self, values: Any) -> _TrackedList: + super().__iadd__([_unwrap_tracked(value) for value in values]) + self._sync() + return self + + def __imul__(self, count: Any) -> Any: + super().__imul__(count) + self._sync() + return self + + +class _TrackedDataset(dict[str, Any]): + """Console-only dataset view that advances freshness on mutable edits.""" + + _NON_DATASET_FIELDS = frozenset({"history"}) + + def __init__(self, source: dict[str, Any], on_mutation: Callable[[], None]) -> None: + super().__init__(source) + self._source = source + self._on_mutation = on_mutation + + @property + def source(self) -> dict[str, Any]: + return self._source + + def __getitem__(self, key: str) -> Any: + return _tracked_value(super().__getitem__(key), self._on_mutation) + + def get(self, key: str, default: Any = None) -> Any: + return _tracked_value(super().get(key, default), self._on_mutation) + + def copy(self) -> _TrackedDataset: + return _TrackedDataset(dict(self), self._on_mutation) + + def items(self) -> Any: + for key, value in super().items(): + yield key, _tracked_value(value, self._on_mutation) + + def values(self) -> Any: + for value in super().values(): + yield _tracked_value(value, self._on_mutation) + + def __setitem__(self, key: str, value: Any) -> None: + value = _unwrap_tracked(value) + super().__setitem__(key, value) + self._source[key] = value + if key not in self._NON_DATASET_FIELDS: + self._on_mutation() + + def __delitem__(self, key: str) -> None: + super().__delitem__(key) + del self._source[key] + if key not in self._NON_DATASET_FIELDS: + self._on_mutation() + + def update(self, *args: Any, **kwargs: Any) -> None: + updates = dict(*args, **kwargs) + for key, value in updates.items(): + self[key] = value + + def clear(self) -> None: + if not self: + return + super().clear() + self._source.clear() + self._on_mutation() + + def pop(self, key: str, default: Any = _TRACKED_MISSING) -> Any: + present = key in self + if default is _TRACKED_MISSING: + value = super().pop(key) + self._source.pop(key) + else: + value = super().pop(key, default) + self._source.pop(key, default) + if present and key not in self._NON_DATASET_FIELDS: + self._on_mutation() + return value + + def popitem(self) -> tuple[str, Any]: + key, value = super().popitem() + self._source.pop(key) + if key not in self._NON_DATASET_FIELDS: + self._on_mutation() + return key, value + + def __ior__(self, other: Any) -> Any: + self.update(other) + return self + + def setdefault(self, key: str, default: Any = None) -> Any: + if key in self: + return self[key] + self[key] = default + return self[key] + + def __or__(self, other: Any) -> _TrackedDataset: + copied = self.copy() + copied.update(other) + return copied + + def __ror__(self, other: Any) -> _TrackedDataset: + copied = _TrackedDataset(dict(other), self._on_mutation) + copied.update(self) + return copied + + +def _tracked_value(value: Any, on_mutation: Callable[[], None]) -> Any: + if isinstance(value, dict) and not isinstance(value, _TrackedDataset): + return _TrackedDataset(value, on_mutation) + if isinstance(value, list) and not isinstance(value, _TrackedList): + return _TrackedList(value, on_mutation) + if isinstance(value, np.ndarray) and not isinstance(value, _MutationTrackedArray): + tracked = value.view(_MutationTrackedArray) + tracked._on_mutation = on_mutation + return tracked + return value + + +def _unwrap_tracked(value: Any) -> Any: + if isinstance(value, _TrackedDataset): + return value.source + if isinstance(value, _TrackedList): + return value.source + if isinstance(value, _MutationTrackedArray): + return value.view(np.ndarray) + return value + + +@dataclass +class _AsyncDatasetWatch: + """Shared wrappers and reference count for concurrent async calls.""" + + eeg_source: dict[str, Any] | list[dict[str, Any]] + eeg_tracked: _TrackedDataset | _TrackedList + alleeg_source: list[dict[str, Any]] + alleeg_tracked: _TrackedList + active_calls: int = 1 + + +class _ConsoleImportedModule: + """Module proxy that replaces canonical async exports with console wrappers.""" + + def __init__(self, module: Any, bridge: EEGPrepConsoleWorkspace, exports: Mapping[str, str]) -> None: + self._module = module + self._bridge = bridge + self._exports = exports + + def __getattr__(self, name: str) -> Any: + export = self._exports.get(name) + if export == "iclabel_async": + return self._bridge.async_wrapper(export) + if export == "pop_iclabel_async": + return self._bridge.pop_wrapper(export) + child = _CANONICAL_MODULE_CHILDREN.get(getattr(self._module, "__name__", ""), {}).get(name) + if child is not None: + module_name, exports = child + return _ConsoleImportedModule(importlib.import_module(module_name), self._bridge, exports) + return getattr(self._module, name) + + def __dir__(self) -> list[str]: + children = _CANONICAL_MODULE_CHILDREN.get(getattr(self._module, "__name__", ""), {}) + return sorted(set(dir(self._module)) | set(self._exports) | set(children)) + + +class _ConsoleImportlib: + """Importlib proxy that preserves console wrappers for canonical modules.""" + + def __init__(self, bridge: EEGPrepConsoleWorkspace) -> None: + self._bridge = bridge + + def import_module(self, name: str, package: str | None = None) -> Any: + module = importlib.import_module(name, package) + exports = _CANONICAL_ASYNC_IMPORTS.get(name) + if exports: + return _ConsoleImportedModule(module, self._bridge, exports) + children = _CANONICAL_MODULE_CHILDREN.get(name) + if children: + return _ConsoleImportedModule(module, self._bridge, {}) + return module + + def __getattr__(self, name: str) -> Any: + return getattr(importlib, name) + + +class ConsoleAsyncFunction(LazyWorkspaceExport): + """Console wrapper that rejects stale results from public async functions.""" + + def __init__( + self, + name: str, + bridge: EEGPrepConsoleWorkspace, + value: Any | None = None, + resolver: Callable[[], Any] | None = None, + ) -> None: + super().__init__(name, value=value, resolver=resolver) + self.bridge = bridge + + def __call__(self, *args: Any, **kwargs: Any) -> Any: + target_state = self.bridge.session.dataset_state_token() + return self._call_async(target_state, *args, **kwargs) + + async def _call_async( + self, + target_state: tuple[Any, ...], + *args: Any, + **kwargs: Any, + ) -> Any: + watch = self.bridge.begin_async_dataset_watch(args, kwargs) + args, kwargs = self.bridge.unwrap_async_call(args, kwargs) + try: + result = await self.resolve()(*args, **kwargs) + if not self.bridge.session.dataset_state_unchanged(target_state): + raise RuntimeError("ICLabel result discarded because the session changed while it was running.") + return result + finally: + self.bridge.end_async_dataset_watch(watch) + + def __repr__(self) -> str: + return f"" + + class ConsolePopFunction(LazyWorkspaceExport): """Console wrapper that stores ``pop_*`` EEG outputs back into the session.""" @@ -208,6 +525,9 @@ def __init__( self.bridge = bridge def __call__(self, *args: Any, **kwargs: Any) -> Any: + if self.name == "pop_iclabel_async": + target_state = self.bridge.session.dataset_state_token() + return self._call_async(target_state, *args, **kwargs) function = self.resolve() call_kwargs = dict(kwargs) recorded_commands: set[str] = set() @@ -256,8 +576,35 @@ def accept_browser_result(eeg_out: Any, command: str) -> None: self.bridge._pop_updated_session = True self.bridge.pull_from_session() return ConsolePopResult(self.bridge.session.EEG, command, updated=False) + eeg_result, result_command = _extract_pop_eeg_and_command(result) + if ( + self.name in _IN_PLACE_DATASET_POP_FUNCTIONS + and not result_command + and eeg_result is self.bridge.session.EEG + ): + self.bridge.session.notify_changed(dataset_changed=True) return self.bridge.accept_pop_result(result, args, kwargs) + async def _call_async( + self, + target_state: tuple[Any, ...], + *args: Any, + **kwargs: Any, + ) -> Any: + watch = self.bridge.begin_async_dataset_watch(args, kwargs) + args, kwargs = self.bridge.unwrap_async_call(args, kwargs) + try: + function = self.resolve() + call_kwargs = dict(kwargs) + if _accepts_return_com(function) and "return_com" not in call_kwargs: + call_kwargs["return_com"] = True + result = await function(*args, **call_kwargs) + if not self.bridge.session.dataset_state_unchanged(target_state): + raise RuntimeError("ICLabel result discarded because the session changed while it was running.") + return self.bridge.accept_pop_result(result, args, kwargs) + finally: + self.bridge.end_async_dataset_watch(watch) + def __repr__(self) -> str: return f"" @@ -271,8 +618,14 @@ def __init__(self, bridge: EEGPrepConsoleWorkspace) -> None: def __getattr__(self, name: str) -> Any: if name.startswith("pop_"): return self._bridge.pop_wrapper(name) + if name == "iclabel_async": + return self._bridge.async_wrapper(name) if name == "eegh": return self._bridge.namespace["eegh"] + child = _CANONICAL_MODULE_CHILDREN.get("eegprep", {}).get(name) + if child is not None: + module_name, exports = child + return _ConsoleImportedModule(importlib.import_module(module_name), self._bridge, exports) return getattr(eegprep, name) def __dir__(self) -> list[str]: @@ -288,7 +641,7 @@ class ConsoleEegh: def __init__(self, bridge: EEGPrepConsoleWorkspace) -> None: self.bridge = bridge - def __call__(self, command: Any = None, *args: Any) -> str: + def __call__(self, command: Any = None, *args: Any) -> Any: if command is None: return eegh(None, self.bridge.session.ALLCOM) if isinstance(command, str) and command.strip().lower() == "find": @@ -309,6 +662,8 @@ def __call__(self, command: Any = None, *args: Any) -> str: return "" history_command = session.history_command_at(value) if history_command: + if "await " in history_command: + return self.bridge.execute_history_command_async(history_command) self.bridge.execute_history_command(history_command) return history_command session = self.bridge.session @@ -318,7 +673,7 @@ def __call__(self, command: Any = None, *args: Any) -> str: else: session.clear_last_command() if args and isinstance(args[0], dict): - eegh(normalized, args[0]) + eegh(normalized, _unwrap_tracked(args[0])) self.bridge.pull_from_session() return normalized @@ -345,8 +700,13 @@ def __init__( self.command_echo = command_echo self.extension_runtime = extension_runtime or ExtensionRuntime.empty() self.namespace: dict[str, Any] = {} + self._builtin_import = builtins.__import__ self._eegprep_proxy = ConsoleEEGPrepModule(self) + self._importlib_proxy = _ConsoleImportlib(self) self._wrapped_pop_exports: dict[str, ConsolePopFunction] = {} + self._wrapped_async_exports: dict[str, ConsoleAsyncFunction] = {} + self._export_values = dict(exports or {}) + self._async_dataset_watch: _AsyncDatasetWatch | None = None self._syncing = False self._pop_updated_session = False self._pop_needs_source_history = False @@ -363,6 +723,8 @@ def __init__( def close(self) -> None: """Detach this workspace from session notifications.""" + self._async_dataset_watch = None + self.pull_from_session() self.session.remove_change_listener(self._session_changed) self.session.remove_command_echo_listener(self._echo_session_command) self.session.remove_gui_action_listener(self._gui_action_event) @@ -372,8 +734,16 @@ def pull_from_session(self) -> None: """Mirror session state into the console namespace.""" self.namespace["eegprep"] = self._eegprep_proxy self.namespace.update(self._wrapped_pop_exports) - self.namespace["EEG"] = self.session.EEG - self.namespace["ALLEEG"] = self.session.ALLEEG + self.namespace.update(self._wrapped_async_exports) + watch = self._async_dataset_watch + if watch is not None and self.session.EEG is watch.eeg_source: + self.namespace["EEG"] = watch.eeg_tracked + else: + self.namespace["EEG"] = self.session.EEG + if watch is not None and self.session.ALLEEG is watch.alleeg_source: + self.namespace["ALLEEG"] = watch.alleeg_tracked + else: + self.namespace["ALLEEG"] = self.session.ALLEEG self.namespace["CURRENTSET"] = self.session.current_set_value() self.namespace["ALLCOM"] = self.session.ALLCOM self.namespace["LASTCOM"] = self.session.LASTCOM @@ -404,7 +774,7 @@ def after_execute(self, source: str, *, success: bool = True) -> None: changed = False if "ALLEEG" in targets or "CURRENTSET" in targets: - alleeg = self.namespace.get("ALLEEG", []) + alleeg = _unwrap_tracked(self.namespace.get("ALLEEG", [])) if not isinstance(alleeg, list): raise ValueError("ALLEEG must be a list of EEG datasets") if "CURRENTSET" in targets: @@ -423,6 +793,7 @@ def after_execute(self, source: str, *, success: bool = True) -> None: changed = True if eeg_changed: + eeg = _unwrap_tracked(eeg) if not _is_eeg_selection(eeg): raise ValueError("EEG must be an EEG dataset dictionary or a list of EEG dataset dictionaries") self._store_eeg(eeg, history_command) @@ -520,7 +891,15 @@ def accept_pop_result(self, result: Any, args: tuple[Any, ...], kwargs: Mapping[ return ConsolePopResult(self.session.EEG, command, updated=True) def _bind_base_namespace(self) -> None: - self.namespace.update({"eegprep": self._eegprep_proxy, "session": self.session, "window": self.window}) + self.namespace.update( + { + "eegprep": self._eegprep_proxy, + "session": self.session, + "importlib": self._importlib_proxy, + "window": self.window, + "__builtins__": {**vars(builtins), "__import__": self._console_import}, + } + ) self.namespace.update(_CONSOLE_COMMAND_EXPORTS) self.namespace["eegh"] = ConsoleEegh(self) @@ -535,6 +914,10 @@ def _bind_exports(self, exports: Mapping[str, Any] | None) -> None: wrapped = ConsolePopFunction(name, self, None if exports is None else exports[name]) self._wrapped_pop_exports[name] = wrapped self.namespace[name] = wrapped + elif name == "iclabel_async": + wrapped = ConsoleAsyncFunction(name, self, None if exports is None else exports[name]) + self._wrapped_async_exports[name] = wrapped + self.namespace[name] = wrapped else: self.namespace[name] = LazyWorkspaceExport(name, None if exports is None else exports[name]) @@ -556,12 +939,59 @@ def pop_wrapper(self, name: str) -> ConsolePopFunction: self._wrapped_pop_exports[name] = wrapped return wrapped + def async_wrapper(self, name: str) -> ConsoleAsyncFunction: + """Return the console-aware wrapper for a public async function.""" + wrapped = self._wrapped_async_exports.get(name) + if wrapped is None: + value = self._export_values.get(name) + wrapped = ConsoleAsyncFunction(name, self, value=value) + self._wrapped_async_exports[name] = wrapped + return wrapped + + def begin_async_dataset_watch(self, args: tuple[Any, ...], kwargs: Mapping[str, Any]) -> _AsyncDatasetWatch | None: + source = _unwrap_tracked(_first_call_eeg_argument(args, kwargs)) + if ( + source is not self.session.EEG + or not isinstance(source, (dict, list)) + or _unwrap_tracked(self.namespace.get("EEG")) is not source + ): + return None + watch = self._async_dataset_watch + if watch is not None and watch.eeg_source is source and watch.alleeg_source is self.session.ALLEEG: + watch.active_calls += 1 + self.pull_from_session() + return watch + tracked = _tracked_value(source, self.session.mark_dataset_changed) + alleeg_tracked = _tracked_value(self.session.ALLEEG, self.session.mark_dataset_changed) + if not isinstance(tracked, (_TrackedDataset, _TrackedList)) or not isinstance(alleeg_tracked, _TrackedList): + return None + watch = _AsyncDatasetWatch(source, tracked, self.session.ALLEEG, alleeg_tracked) + self._async_dataset_watch = watch + self.pull_from_session() + return watch + + def unwrap_async_call( + self, args: tuple[Any, ...], kwargs: Mapping[str, Any] + ) -> tuple[tuple[Any, ...], dict[str, Any]]: + return tuple(_unwrap_tracked(arg) for arg in args), { + key: _unwrap_tracked(value) for key, value in kwargs.items() + } + + def end_async_dataset_watch(self, watch: _AsyncDatasetWatch | None) -> None: + if watch is None or self._async_dataset_watch is not watch: + return + watch.active_calls -= 1 + if watch.active_calls: + return + self._async_dataset_watch = None + self.pull_from_session() + def _session_changed(self, _session: EEGPrepSession) -> None: if not self._syncing: self.pull_from_session() def _namespace_eeg_changed(self, targets: set[str]) -> bool: - return "EEG" in targets or self.namespace.get("EEG") is not self.session.EEG + return "EEG" in targets or _unwrap_tracked(self.namespace.get("EEG")) is not self.session.EEG def _history_command_for_source(self, source: str, targets: set[str]) -> str: lastcom = str(self.namespace.get("LASTCOM") or "").strip() @@ -575,8 +1005,42 @@ def _restore_imported_wrappers(self, source: str) -> None: for local_name, export_name in _eegprep_import_aliases(source).items(): if export_name == "eegprep": self.namespace[local_name] = self._eegprep_proxy + elif export_name == "importlib": + self.namespace[local_name] = self._importlib_proxy + elif export_name.startswith("__module__:"): + self.namespace[local_name] = self._importlib_proxy.import_module( + export_name.removeprefix("__module__:") + ) elif export_name.startswith("pop_"): self.namespace[local_name] = self.pop_wrapper(export_name) + elif export_name == "iclabel_async": + self.namespace[local_name] = self.async_wrapper(export_name) + + def _console_import( + self, + name: str, + globals: dict[str, Any] | None = None, + locals: dict[str, Any] | None = None, + fromlist: tuple[str, ...] = (), + level: int = 0, + ) -> Any: + module = self._builtin_import(name, globals, locals, fromlist, level) + if name == "importlib": + return self._importlib_proxy + eegprep_exports = set(getattr(eegprep, "__all__", ())) + if name == "eegprep" and ( + not fromlist or (fromlist and all(item == "*" or item in eegprep_exports for item in fromlist)) + ): + return self._eegprep_proxy + canonical_exports = _CANONICAL_ASYNC_IMPORTS.get(name) + if canonical_exports and (not fromlist or any(item == "*" or item in canonical_exports for item in fromlist)): + return _ConsoleImportedModule(module, self, canonical_exports) + canonical_children = _CANONICAL_MODULE_CHILDREN.get(name) + if canonical_children and fromlist and any(item == "*" or item in canonical_children for item in fromlist): + return _ConsoleImportedModule(module, self, {}) + if canonical_exports and not fromlist: + return _ConsoleImportedModule(module, self, canonical_exports) + return module def _store_eeg(self, eeg: Any, command: str, *, new: bool = False, index: int | list[int] | None = None) -> None: self._syncing = True @@ -591,6 +1055,16 @@ def execute_history_command(self, command: str) -> None: exec(source, self.namespace) self.after_execute(source) + async def execute_history_command_async(self, command: str) -> str: + """Execute an awaitable EEGLAB history command through the console namespace.""" + source = _console_python_command(command) + code = compile(source, "", "exec", flags=ast.PyCF_ALLOW_TOP_LEVEL_AWAIT) + result = eval(code, self.namespace) # noqa: S307 - executing trusted session history + if inspect.isawaitable(result): + await result + self.after_execute(source) + return command + def _refresh(self) -> None: if self.refresh is not None: self.refresh() @@ -1530,10 +2004,36 @@ def _eegprep_import_aliases(source: str) -> dict[str, str]: for alias in node.names: if alias.name == "eegprep": aliases[alias.asname or "eegprep"] = "eegprep" + elif alias.name == "importlib": + aliases[alias.asname or "importlib"] = "importlib" + elif alias.name in _CANONICAL_ASYNC_IMPORTS: + aliases[alias.asname or alias.name.split(".")[0]] = ( + "__module__:" + alias.name if alias.asname else "eegprep" + ) elif isinstance(node, ast.ImportFrom) and node.module == "eegprep": for alias in node.names: - if alias.name.startswith("pop_"): + if alias.name.startswith("pop_") or alias.name == "iclabel_async": aliases[alias.asname or alias.name] = alias.name + elif isinstance(node, ast.ImportFrom): + export = _CANONICAL_ASYNC_IMPORTS.get(node.module or "") + if export: + for alias in node.names: + if alias.name == "*": + aliases.update({name: value for name, value in export.items()}) + elif alias.name in export: + aliases[alias.asname or alias.name] = export[alias.name] + children = _CANONICAL_MODULE_CHILDREN.get(node.module or "") + if children: + for alias in node.names: + if alias.name == "*": + aliases.update( + { + child_name: "__module__:" + module_name + for child_name, (module_name, _exports) in children.items() + } + ) + elif alias.name in children: + aliases[alias.asname or alias.name] = "__module__:" + children[alias.name][0] return aliases diff --git a/src/eegprep/functions/adminfunc/eeglabcompat.py b/src/eegprep/functions/adminfunc/eeglabcompat.py index a9147b77..ccf975f1 100644 --- a/src/eegprep/functions/adminfunc/eeglabcompat.py +++ b/src/eegprep/functions/adminfunc/eeglabcompat.py @@ -295,7 +295,13 @@ def get_eeglab(runtime: str = default_runtime, *, auto_file_roundtrip: bool = Tr # not yet loaded, do so now if rt == 'oct': - from oct2py import Oct2Py, get_log + try: + from oct2py import Oct2Py, get_log + except ImportError: + raise ImportError( + "oct2py is required to use the Octave runtime. Install it with 'pip install eegprep[eeglab]' " + "(or 'uv pip install eegprep[eeglab]')." + ) engine = Oct2Py(logger=get_log()) engine.logger = get_log("new_log") diff --git a/src/eegprep/functions/adminfunc/storage.py b/src/eegprep/functions/adminfunc/storage.py index 08c6af02..e14432f1 100644 --- a/src/eegprep/functions/adminfunc/storage.py +++ b/src/eegprep/functions/adminfunc/storage.py @@ -3,6 +3,7 @@ from __future__ import annotations from copy import deepcopy +from collections.abc import Callable from pathlib import Path import shutil import tempfile @@ -15,6 +16,102 @@ FDT_DTYPE = np.dtype(" None: + self._flat = array.flat + self._on_mutation = on_mutation + + def __getitem__(self, key: Any) -> Any: + return self._flat[key] + + def __setitem__(self, key: Any, value: Any) -> None: + self._flat[key] = value + self._on_mutation() + + def __iter__(self) -> Any: + return iter(self._flat) + + def __len__(self) -> int: + return len(self._flat) + + +class _MutationTrackedArray(np.ndarray): + """Array view that invokes a callback after an in-place write.""" + + _on_mutation: Callable[[], None] | None + + def __array_finalize__(self, source: Any) -> None: + self._on_mutation = getattr(source, "_on_mutation", None) + + def _notify(self) -> None: + if self._on_mutation is not None: + self._on_mutation() + + def __setitem__(self, key: Any, value: Any) -> None: + super().__setitem__(key, value) + self._notify() + + def fill(self, value: Any) -> None: + super().fill(value) + self._notify() + + @property + def flat(self) -> _MutationTrackedFlat: + """Return a one-dimensional view that preserves write tracking.""" + return _MutationTrackedFlat(self, self._notify) + + def sort(self, *args: Any, **kwargs: Any) -> None: + super().sort(*args, **kwargs) + self._notify() + + def partition(self, *args: Any, **kwargs: Any) -> None: + super().partition(*args, **kwargs) + self._notify() + + def byteswap(self, inplace: bool = False) -> np.ndarray: + result = super().byteswap(inplace=inplace) + if inplace: + self._notify() + return result + + def __array_ufunc__(self, ufunc: Any, method: str, *inputs: Any, **kwargs: Any) -> Any: + outputs = kwargs.get("out") + mutating_input = inputs[0] if method == "at" and inputs else None + tracked_outputs = outputs and any(isinstance(output, _MutationTrackedArray) for output in outputs) + if outputs: + kwargs["out"] = tuple( + np.asarray(output) if isinstance(output, _MutationTrackedArray) else output for output in outputs + ) + inputs = tuple(np.asarray(value) if isinstance(value, _MutationTrackedArray) else value for value in inputs) + result = getattr(ufunc, method)(*inputs, **kwargs) + if method == "at" and isinstance(mutating_input, _MutationTrackedArray): + mutating_input._notify() + elif method == "__call__" and tracked_outputs: + for output in outputs: + if isinstance(output, _MutationTrackedArray): + output._notify() + return result + + def __array_function__(self, function: Any, types: Any, args: Any, kwargs: Any) -> Any: # ty: ignore[invalid-method-override] + if function is np.copyto and args and isinstance(args[0], _MutationTrackedArray): + converted_args = tuple( + np.asarray(value) if isinstance(value, _MutationTrackedArray) else value for value in args + ) + result = np.copyto(*converted_args, **kwargs) + args[0]._notify() + return result + if function is np.put and args and isinstance(args[0], _MutationTrackedArray): + converted_args = tuple( + np.asarray(value) if isinstance(value, _MutationTrackedArray) else value for value in args + ) + result = np.put(*converted_args, **kwargs) + args[0]._notify() + return result + return super().__array_function__(function, types, args, kwargs) + + class _BackingFile: """Reference-counted backing path shared by copy-on-write handles.""" @@ -32,6 +129,22 @@ def release(self) -> None: self.path.unlink(missing_ok=True) +def _unwrap_memmap_data(value: Any) -> Any: + """Replace MemmapData handles with their mapped arrays, including inside containers. + + Nested handles have to be unwrapped too. Leaving one inside a list would make numpy + dispatch straight back into ``MemmapData.__array_function__`` and recurse forever on + calls like ``np.concatenate([mapped_a, mapped_b])``. + """ + if isinstance(value, MemmapData): + return value._memmap() + if isinstance(value, dict): + return {key: _unwrap_memmap_data(item) for key, item in value.items()} + if isinstance(value, (list, tuple)): + return type(value)(_unwrap_memmap_data(item) for item in value) + return value + + class MemmapData: """NumPy-compatible handle for channel-major EEG data stored on disk. @@ -76,6 +189,7 @@ def __init__( self._backing.acquire() self._released = False self._array: np.memmap | None = None + self._mutation_revision = 0 self._validate_backing_file() @classmethod @@ -189,10 +303,25 @@ def size(self) -> int: """Return the total sample count.""" return int(np.prod(self._shape)) + @property + def mutation_revision(self) -> int: + """Return the in-process write revision for this mapped dataset.""" + return self._mutation_revision + @property def T(self) -> np.ndarray: """Return a transposed array view.""" - return self._memmap().T + return self._tracked_view(self._memmap().T) + + @property + def flat(self) -> _MutationTrackedFlat: + """Return a one-dimensional view that preserves write tracking.""" + return _MutationTrackedFlat(self._memmap(), self._mark_mutated) + + def fill(self, value: Any) -> None: + """Fill the mapped data and advance its mutation revision.""" + self._memmap().fill(value) + self._mark_mutated() def flush(self) -> None: """Flush pending writes to the backing file.""" @@ -290,7 +419,33 @@ def change_file(self, filename: str | Path, *, writable: bool = False) -> None: def reshape(self, *shape: Any, **kwargs: Any) -> np.ndarray: """Return a reshaped view using NumPy's reshape semantics.""" - return self._memmap().reshape(*shape, **kwargs) + return self._tracked_view(self._memmap().reshape(*shape, **kwargs)) + + def transpose(self, *axes: Any) -> np.ndarray: + """Return a transposed view that preserves mutation tracking.""" + return self._tracked_view(self._memmap().transpose(*axes)) + + def ravel(self, *args: Any, **kwargs: Any) -> np.ndarray: + """Return a flattened view or copy using NumPy's ravel semantics.""" + return self._tracked_view(self._memmap().ravel(*args, **kwargs)) + + def sort(self, *args: Any, **kwargs: Any) -> None: + """Sort mapped data in place and advance its mutation revision.""" + self._memmap().sort(*args, **kwargs) + self._mark_mutated() + + def partition(self, *args: Any, **kwargs: Any) -> None: + """Partition mapped data in place and advance its mutation revision.""" + self._memmap().partition(*args, **kwargs) + self._mark_mutated() + + def byteswap(self, inplace: bool = False) -> np.ndarray: + """Byte-swap mapped data while preserving mutation tracking.""" + result = self._memmap().byteswap(inplace=inplace) + if inplace: + self._mark_mutated() + return self._tracked_view(result) + return result def astype(self, *args: Any, **kwargs: Any) -> np.ndarray: """Return a typed array using NumPy's astype semantics.""" @@ -305,27 +460,46 @@ def mean(self, *args: Any, **kwargs: Any) -> Any: return self._memmap().mean(*args, **kwargs) def __array__(self, dtype: np.dtype | str | None = None, copy: bool | None = None) -> np.ndarray: - array = np.asarray(self._memmap()) + array = self._memmap() if dtype is not None: - return array.astype(dtype, copy=bool(copy)) + if np.dtype(dtype) != self._dtype or copy: + return np.asarray(array, dtype=dtype).copy() + return self._tracked_view(array) if copy: - return array.copy() - return array + return np.asarray(array).copy() + return self._tracked_view(array) def __getitem__(self, key: Any) -> Any: - return self._memmap()[key] + value = self._memmap()[key] + return self._tracked_view(value) if isinstance(value, np.ndarray) else value def __setitem__(self, key: Any, value: Any) -> None: if not self.writable: raise ValueError("assignment destination is read-only") self._detach_for_write() self._memmap()[key] = value + self._mark_mutated() + + def __array_function__(self, function: Any, types: Any, args: Any, kwargs: Any) -> Any: + if function in {np.copyto, np.put} and args and args[0] is self: + converted_args = (self._memmap(), *args[1:]) + result = function(*converted_args, **kwargs) + self._mark_mutated() + return result + # Defining this method opts the type into numpy's dispatch protocol, which means + # NotImplemented here does not fall back, it raises: every numpy function other than + # the two above would fail on a MemmapData. That defeats a handle whose whole purpose + # is to stand in for an array, and np.size(data) in pop_reref is one caller of many. + # Delegate to the mapped array instead. The two mutating calls this class tracks are + # handled above, and writes through a returned view are tracked by _tracked_view. + return function(*_unwrap_memmap_data(args), **_unwrap_memmap_data(kwargs)) def __len__(self) -> int: return len(self._memmap()) def __iter__(self) -> Any: - return iter(self._memmap()) + for index in range(len(self)): + yield self[index] def __getattr__(self, name: str) -> Any: return getattr(self._memmap(), name) @@ -413,6 +587,16 @@ def _replace_backing(self, backing: _BackingFile, *, mode: str) -> None: self._array = None old_backing.release() + def _mark_mutated(self) -> None: + self._mutation_revision += 1 + + def _tracked_view(self, array: np.ndarray) -> np.ndarray: + if not np.shares_memory(array, self._memmap()): + return array + tracked: _MutationTrackedArray = array.view(_MutationTrackedArray) + tracked._on_mutation = self._mark_mutated + return tracked + def _physical_shape(shape: tuple[int, ...], transposed: bool) -> tuple[int, ...]: if not transposed: diff --git a/src/eegprep/functions/guifunc/main_window.py b/src/eegprep/functions/guifunc/main_window.py index ea358f80..903ddeed 100644 --- a/src/eegprep/functions/guifunc/main_window.py +++ b/src/eegprep/functions/guifunc/main_window.py @@ -2,8 +2,11 @@ from __future__ import annotations +import asyncio import ctypes import ctypes.util +import inspect +import logging import sys from typing import Any @@ -33,6 +36,7 @@ _MACOS_MENU_BRANDING_RETRY_MS = 100 COMING_SOON_SUFFIX = " (coming soon)" COMING_SOON_TOOLTIP = "This workflow is not available in EEGPrep yet." +logger = logging.getLogger(__name__) def _require_qt() -> tuple[Any, Any, Any]: @@ -87,6 +91,7 @@ def __init__( native_file_dialogs=native_file_dialogs, extension_runtime=self.extension_runtime, ) + self._async_menu_tasks: set[asyncio.Future[Any]] = set() self._build_central_widget() self.refresh() @@ -221,10 +226,35 @@ def _queue_application_branding(self) -> None: def _dispatch_menu_action(self, action_id: str) -> None: try: - self.dispatcher.dispatch_gui(action_id, self.window) + result = self.dispatcher.dispatch_gui(action_id, self.window) + if inspect.isawaitable(result): + self._schedule_async_menu_action(result) finally: self._queue_application_branding() + def _schedule_async_menu_action(self, awaitable: Any) -> None: + """Schedule an async menu action on the active browser event loop.""" + try: + task = asyncio.get_running_loop().create_task(awaitable) + except RuntimeError: + if inspect.iscoroutine(awaitable): + awaitable.close() + raise RuntimeError("Async GUI menu actions require an active event loop.") from None + self._async_menu_tasks.add(task) + task.add_done_callback(self._finish_async_menu_action) + + def _finish_async_menu_action(self, task: asyncio.Future[Any]) -> None: + self._async_menu_tasks.discard(task) + if task.cancelled(): + return + exception = task.exception() + if exception is not None: + logger.error( + "Asynchronous EEGPrep GUI menu action failed: %s", + exception, + exc_info=(type(exception), exception, exception.__traceback__), + ) + def _current_menu_specs(self) -> tuple[MenuItemSpec, ...]: specs = [] for spec in eeglab_menus( diff --git a/src/eegprep/functions/guifunc/menu_actions.py b/src/eegprep/functions/guifunc/menu_actions.py index 1fd4f6bb..c4683647 100644 --- a/src/eegprep/functions/guifunc/menu_actions.py +++ b/src/eegprep/functions/guifunc/menu_actions.py @@ -4,6 +4,7 @@ import inspect import logging +import sys import webbrowser from collections.abc import Callable, Mapping from functools import wraps @@ -22,6 +23,7 @@ logger = logging.getLogger(__name__) +_IS_EMSCRIPTEN = sys.platform == "emscripten" def _with_file_dialog_override(method: Callable[..., Any]) -> Callable[..., Any]: @@ -275,8 +277,10 @@ def __init__( self.extension_runtime = extension_runtime or ExtensionRuntime.empty() self._long_tasks: list[LongTaskHandle] = [] - def dispatch_gui(self, action: str, parent: Any | None = None) -> None: + def dispatch_gui(self, action: str, parent: Any | None = None) -> Any: """Run a menu action from Qt and show user-facing errors.""" + if _IS_EMSCRIPTEN and action.partition(":")[0] == "pop_iclabel": + return self.dispatch_gui_async(action, parent) with self.session.gui_action(action): try: self.dispatch(action, parent) @@ -286,6 +290,26 @@ def dispatch_gui(self, action: str, parent: Any | None = None) -> None: raise self._warn(parent, str(exc)) + async def dispatch_gui_async( + self, + action: str, + parent: Any | None = None, + *, + renderer: Any | None = None, + ) -> None: + """Run the browser ICLabel menu action while awaiting its Promise backend.""" + if action.partition(":")[0] != "pop_iclabel": + self.dispatch_gui(action, parent) + return + with self.session.gui_action(action): + try: + await self._run_pop_iclabel_async(parent, renderer=renderer) + except Exception as exc: + logger.exception("EEGPrep GUI menu action failed: %s", action) + if parent is None: + raise + self._warn(parent, str(exc)) + @_with_file_dialog_override def dispatch(self, action: str, parent: Any | None = None) -> None: """Run a menu action.""" @@ -1319,6 +1343,25 @@ def commit_component_rejection(eeg_out: Any, _states: dict[int, bool]) -> None: self._store_current_from_gui(eeg_out, command=command) self._refresh() + async def _run_pop_iclabel_async(self, parent: Any | None, *, renderer: Any | None = None) -> None: + selection = self._current_selection_or_warn(parent, allow_multiple=True) + if selection is None: + return + from eegprep.plugins.ICLabel.pop_iclabel import pop_iclabel_async + + target_state = self.session.dataset_state_token() + out = await pop_iclabel_async(selection, renderer=renderer, return_com=True) + if not self.session.dataset_state_unchanged(target_state): + raise RuntimeError("ICLabel result discarded because the session changed while it was running.") + eeg_out, command = out[0], out[1] if isinstance(out, tuple) and len(out) > 1 else "" + if command: + store_kwargs = {"command": command} + target_indices = list(self.session.selected_dataset_indices()) + if target_indices: + store_kwargs["index"] = target_indices + self._store_current_from_gui(eeg_out, **store_kwargs) + self._refresh() + def _run_pop_runica_long_task( self, selection: Any, @@ -1419,7 +1462,7 @@ def _select_multiple_datasets(self, parent: Any | None) -> None: if command: self.session.echo_command(command) self.session.add_history(command, notify=False) - self.session.notify_changed() + self.session.notify_changed(dataset_changed=True) self._refresh() def _select_study_set(self, parent: Any | None) -> None: diff --git a/src/eegprep/functions/guifunc/session.py b/src/eegprep/functions/guifunc/session.py index cefdbdc2..6b8f18e0 100644 --- a/src/eegprep/functions/guifunc/session.py +++ b/src/eegprep/functions/guifunc/session.py @@ -2,10 +2,12 @@ from __future__ import annotations +import hashlib from collections.abc import Iterable from contextlib import contextmanager from copy import deepcopy from dataclasses import dataclass, field +from pathlib import Path from typing import Any, Callable, Iterator import numpy as np @@ -14,7 +16,7 @@ from eegprep.functions.adminfunc.eeg_retrieve import eeg_retrieve from eegprep.functions.adminfunc.eeg_store import eeg_store from eegprep.functions.adminfunc.pop_delset import pop_delset -from eegprep.functions.adminfunc.storage import offload_storedisk_datasets +from eegprep.functions.adminfunc.storage import MemmapData, offload_storedisk_datasets from eegprep.functions.popfunc.eeg_emptyset import eeg_emptyset @@ -35,6 +37,96 @@ def has_eeg_data(eeg: Any) -> bool: return True +def _dataset_storage_token(dataset: Any) -> tuple[Any, ...] | None: + """Return a cheap mutation token for file-backed EEG data.""" + if not isinstance(dataset, dict): + return None + data = dataset.get("data") + if isinstance(data, MemmapData): + path = data.path + revision = data.mutation_revision + elif isinstance(data, np.memmap): + filename = data.filename + path = Path(filename) if filename else None + revision = 0 + else: + return None + try: + stat = path.stat() if path is not None else None + except OSError: + stat = None + return ( + type(data).__name__, + str(path) if path is not None else None, + tuple(int(size) for size in data.shape), + np.dtype(data.dtype).str, + revision, + None if stat is None else stat.st_mtime_ns, + None if stat is None else stat.st_size, + ) + + +def _dataset_content_token(dataset: Any) -> bytes: + """Return an authoritative token for mutable dataset content.""" + digest = hashlib.blake2b(digest_size=16) + _update_dataset_digest(dataset, digest, set()) + return digest.digest() + + +def _update_dataset_digest(value: Any, digest: Any, seen: set[int]) -> None: + if isinstance(value, dict): + value_id = id(value) + if value_id in seen: + digest.update(b"") + return + seen.add(value_id) + digest.update(b"dict[") + for key in sorted(value, key=str): + if key == "history": + continue + _update_dataset_digest(str(key), digest, seen) + _update_dataset_digest(value[key], digest, seen) + digest.update(b"]") + seen.remove(value_id) + return + if isinstance(value, (list, tuple)): + value_id = id(value) + if value_id in seen: + digest.update(b"") + return + seen.add(value_id) + digest.update(b"list[") + for item in value: + _update_dataset_digest(item, digest, seen) + digest.update(b"]") + seen.remove(value_id) + return + if isinstance(value, np.ndarray): + digest.update(b"ndarray") + digest.update(np.dtype(value.dtype).str.encode()) + digest.update(repr(tuple(value.shape)).encode()) + if value.dtype.hasobject: + _update_dataset_digest(value.tolist(), digest, seen) + else: + digest.update(np.ascontiguousarray(value).tobytes()) + return + if isinstance(value, MemmapData): + digest.update(b"MemmapData") + digest.update(str(value.path).encode()) + digest.update(repr(value.mutation_revision).encode()) + try: + _update_dataset_digest(np.asarray(value), digest, seen) + except RuntimeError: + digest.update(repr(value).encode()) + return + if value is None or isinstance(value, (bool, int, float, complex, str, bytes)): + digest.update(type(value).__name__.encode()) + digest.update(repr(value).encode()) + return + digest.update(type(value).__name__.encode()) + digest.update(repr(value).encode()) + + def normalize_dataset_indices(indices: Any, *, allow_empty: bool = True) -> list[int]: """Normalize EEGLAB-facing 1-based dataset indices for session state. @@ -141,6 +233,7 @@ class EEGPrepSession: _listeners: list[Callable[["EEGPrepSession"], None]] = field(default_factory=list, init=False, repr=False) _command_echo_listeners: list[Callable[[str], None]] = field(default_factory=list, init=False, repr=False) _gui_action_listeners: list[Callable[[str, str], None]] = field(default_factory=list, init=False, repr=False) + _dataset_revision: int = field(default=0, init=False, repr=False) def add_change_listener(self, listener: Callable[["EEGPrepSession"], None]) -> None: """Register a callback that runs after session state changes.""" @@ -198,11 +291,46 @@ def echo_command(self, command: str | None) -> None: for listener in list(self._command_echo_listeners): listener(command) - def notify_changed(self) -> None: + def notify_changed(self, *, dataset_changed: bool = False) -> None: """Notify listeners that session-backed state changed.""" + if dataset_changed: + self._mark_dataset_changed() for listener in list(self._listeners): listener(self) + def _mark_dataset_changed(self) -> None: + self._dataset_revision += 1 + + def mark_dataset_changed(self) -> None: + """Advance the dataset freshness revision after an observed mutation.""" + self._mark_dataset_changed() + + def dataset_state_token(self) -> tuple[Any, ...]: + """Return a token for rejecting stale asynchronous dataset results.""" + current = self.EEG if isinstance(self.EEG, list) else [self.EEG] + selected_slots = tuple( + id(self.ALLEEG[index - 1]) if 1 <= index <= len(self.ALLEEG) else 0 for index in self.CURRENTSET + ) + selected_datasets = tuple(self.ALLEEG[index - 1] for index in self.CURRENTSET if 1 <= index <= len(self.ALLEEG)) + tracked_ids: set[int] = set() + tracked_datasets: list[Any] = [] + for dataset in (*current, *selected_datasets): + if id(dataset) not in tracked_ids: + tracked_ids.add(id(dataset)) + tracked_datasets.append(dataset) + return ( + self._dataset_revision, + tuple(self.CURRENTSET), + selected_slots, + tuple(id(dataset) for dataset in current), + tuple(_dataset_storage_token(dataset) for dataset in tracked_datasets), + tuple(_dataset_content_token(dataset) for dataset in tracked_datasets), + ) + + def dataset_state_unchanged(self, token: tuple[Any, ...]) -> bool: + """Return whether dataset state still matches a captured async token.""" + return self.dataset_state_token() == token + def current_eeg(self) -> dict[str, Any] | list[dict[str, Any]]: """Return the current EEG selection.""" return self.EEG @@ -257,7 +385,7 @@ def store_current( if mark_saved: self.mark_current_saved() self.add_history(command, notify=False) - self.notify_changed() + self.notify_changed(dataset_changed=True) return stored_index def retrieve(self, indices: int | list[int]) -> dict[str, Any] | list[dict[str, Any]]: @@ -267,7 +395,7 @@ def retrieve(self, indices: int | list[int]) -> dict[str, Any] | list[dict[str, eeg, self.ALLEEG, current = eeg_retrieve(self.ALLEEG, selection if use_vector else selection[0]) self.EEG = eeg self.CURRENTSET = normalize_dataset_indices(current, allow_empty=False) - self.notify_changed() + self.notify_changed(dataset_changed=True) return eeg def apply_workspace_state( @@ -329,7 +457,7 @@ def apply_workspace_state( self.CURRENTSTUDY = int(currentstudy or 0) self.add_history(command, notify=False) - self.notify_changed() + self.notify_changed(dataset_changed=dataset_changed) def delete_current(self) -> None: """Delete the current dataset selection from memory. @@ -348,7 +476,7 @@ def delete_current(self) -> None: return self.CURRENTSET = [] self.EEG = eeg_emptyset() - self.notify_changed() + self.notify_changed(dataset_changed=True) def clear_all(self) -> None: """Clear all datasets and study state.""" @@ -357,7 +485,8 @@ def clear_all(self) -> None: self.CURRENTSET = [] self.STUDY = None self.CURRENTSTUDY = 0 - self.add_history("STUDY = []; CURRENTSTUDY = 0; ALLEEG = []; EEG=[]; CURRENTSET=[];") + self.add_history("STUDY = []; CURRENTSTUDY = 0; ALLEEG = []; EEG=[]; CURRENTSET=[];", notify=False) + self.notify_changed(dataset_changed=True) def set_study( self, @@ -385,7 +514,7 @@ def set_study( self.EEG = eeg_emptyset() offload_storedisk_datasets(self.ALLEEG, set(self.CURRENTSET)) self.add_history(command, notify=False) - self.notify_changed() + self.notify_changed(dataset_changed=alleeg is not None) def _resolve_workspace_eeg( self, @@ -473,6 +602,7 @@ def mark_current_saved(self) -> None: if 1 <= index <= len(self.ALLEEG): self.ALLEEG[index - 1]["saved"] = "yes" offload_storedisk_datasets(self.ALLEEG, set(self.CURRENTSET)) + self._mark_dataset_changed() def menu_statuses(self) -> set[str]: """Return EEGLAB-style menu status tokens for the current state.""" diff --git a/src/eegprep/functions/miscfunc/misc.py b/src/eegprep/functions/miscfunc/misc.py index 2bda712c..03518759 100644 --- a/src/eegprep/functions/miscfunc/misc.py +++ b/src/eegprep/functions/miscfunc/misc.py @@ -207,7 +207,8 @@ def num_jobs_from_reservation(ReservePerJob: str) -> int: import psutil except ImportError: raise ImportError( - "psutil is required to determine available system RAM. Please install it with 'uv pip install psutil'." + "psutil is required to determine available system RAM. " + "Install it with 'pip install eegprep[sys]' (or 'uv pip install eegprep[sys]')." ) avail_amt = psutil.virtual_memory().available unit = reserve[-2:].upper() diff --git a/src/eegprep/functions/sigprocfunc/runica.py b/src/eegprep/functions/sigprocfunc/runica.py index 712caa84..aecdb062 100644 --- a/src/eegprep/functions/sigprocfunc/runica.py +++ b/src/eegprep/functions/sigprocfunc/runica.py @@ -26,6 +26,7 @@ from scipy.linalg import sqrtm, pinv, eig from ...plugins.clean_rawdata.private.ransac import rand_permutation from ..miscfunc.misc import finite_pinv +from . import runica_matmul as _runica_matmul logger = logging.getLogger(__name__) @@ -685,7 +686,7 @@ def runica(data, **kwargs): for t in range(0, lastt, block): # Extract and process block (MATLAB line 846) # MATLAB: u = weights*double(data(:,timeperm(t:t+block-1))) + bias*onesrow - u = weights @ data[:, timeperm[t : t + block]] + bias + u = _runica_matmul.runica_matmul(weights, data[:, timeperm[t : t + block]]) + bias # Apply tanh nonlinearity (MATLAB line 848) y = np.tanh(u) @@ -693,7 +694,9 @@ def runica(data, **kwargs): # Extended-ICA natural gradient weight update (MATLAB line 849) # weights = weights + lrate*(BI-signs*y*u'-u*u')*weights signed_y = signs[:, np.newaxis] * y - weights = weights + lrate * ((BI - (signed_y + u) @ u.T) @ weights) + weights = weights + lrate * _runica_matmul.runica_matmul( + BI - _runica_matmul.runica_matmul(signed_y + u, u.T), weights + ) # Bias update for tanh (MATLAB line 850) # bias = bias + lrate*sum((-2*y)')'; @@ -719,10 +722,10 @@ def runica(data, **kwargs): # Pick random subset (MATLAB lines 869-876) # Use randint to avoid index overflow (rand() * datalength could equal datalength) rp = rng.randint(1, datalength, size=kurtsize) - partact = weights @ data[:, rp[:kurtsize]] + partact = _runica_matmul.runica_matmul(weights, data[:, rp[:kurtsize]]) else: # For small data sets, use whole data (MATLAB lines 877-878) - partact = weights @ data + partact = _runica_matmul.runica_matmul(weights, data) # Compute kurtosis (MATLAB lines 880-882) m2 = np.mean(partact**2, axis=1) ** 2 @@ -872,7 +875,7 @@ def runica(data, **kwargs): # Extract and process block (MATLAB line 1021) # MATLAB: u = weights*double(data(:,timeperm(t:t+block-1))) + bias*onesrow # Note: MATLAB uses 1-based indexing, so t:t+block-1 means t to t+block - u = weights @ data[:, timeperm[t : t + block]] + bias + u = _runica_matmul.runica_matmul(weights, data[:, timeperm[t : t + block]]) + bias # Apply logistic nonlinearity (MATLAB line 1022) # Clip u to prevent overflow in exp @@ -883,7 +886,9 @@ def runica(data, **kwargs): # Natural gradient weight update (MATLAB line 1023) # weights = weights + lrate*(BI+(1-2*y)*u')*weights y_update = 1.0 - 2.0 * y - weights = weights + lrate * ((BI + y_update @ u.T) @ weights) + weights = weights + lrate * _runica_matmul.runica_matmul( + BI + _runica_matmul.runica_matmul(y_update, u.T), weights + ) # Bias update (MATLAB line 1024) # bias = bias + lrate*sum((1-2*y)')'; @@ -1015,14 +1020,16 @@ def runica(data, **kwargs): with np.errstate(divide="ignore", over="ignore", invalid="ignore"): for t in range(0, lastt, block): # Extract and process block - NO BIAS (MATLAB line 1145) - u = weights @ data[:, timeperm[t : t + block]] + u = _runica_matmul.runica_matmul(weights, data[:, timeperm[t : t + block]]) # Apply tanh nonlinearity (MATLAB line 1146) y = np.tanh(u) # Extended-ICA natural gradient weight update (MATLAB line 1147) signed_y = signs[:, np.newaxis] * y - weights = weights + lrate * ((BI - (signed_y + u) @ u.T) @ weights) + weights = weights + lrate * _runica_matmul.runica_matmul( + BI - _runica_matmul.runica_matmul(signed_y + u, u.T), weights + ) # NO BIAS UPDATE for no-bias variant @@ -1043,9 +1050,9 @@ def runica(data, **kwargs): if kurtsize < frames: # Use randint to avoid index overflow (rand() * datalength could equal datalength) rp = rng.randint(1, datalength, size=kurtsize) - partact = weights @ data[:, rp[:kurtsize]] + partact = _runica_matmul.runica_matmul(weights, data[:, rp[:kurtsize]]) else: - partact = weights @ data + partact = _runica_matmul.runica_matmul(weights, data) m2 = np.mean(partact**2, axis=1) ** 2 m4 = np.mean(partact**4, axis=1) @@ -1176,7 +1183,7 @@ def runica(data, **kwargs): with np.errstate(divide="ignore", over="ignore", invalid="ignore"): for t in range(0, lastt, block): # Extract and process block - NO BIAS (MATLAB line 1315) - u = weights @ data[:, timeperm[t : t + block]] + u = _runica_matmul.runica_matmul(weights, data[:, timeperm[t : t + block]]) # Apply logistic nonlinearity (MATLAB line 1316) u = np.maximum(u, -MAX_WEIGHT) @@ -1185,7 +1192,9 @@ def runica(data, **kwargs): # Natural gradient weight update (MATLAB line 1317) y_update = 1.0 - 2.0 * y - weights = weights + lrate * ((BI + y_update @ u.T) @ weights) + weights = weights + lrate * _runica_matmul.runica_matmul( + BI + _runica_matmul.runica_matmul(y_update, u.T), weights + ) # NO BIAS UPDATE for no-bias variant diff --git a/src/eegprep/functions/sigprocfunc/runica_matmul.py b/src/eegprep/functions/sigprocfunc/runica_matmul.py new file mode 100644 index 00000000..3ce2818c --- /dev/null +++ b/src/eegprep/functions/sigprocfunc/runica_matmul.py @@ -0,0 +1,31 @@ +"""Platform-selected matrix products for the float64 runica hot path. + +The runica training loops provide the existing ``np.errstate`` suppression +around these calls. Keeping that context at the loop boundary avoids paying +for a new context on every matrix product. +""" + +from __future__ import annotations + +import sys + +import numpy as np +from scipy.linalg import blas + +_dgemm = getattr(blas, "dgemm") + + +def _numpy_matmul(left: np.ndarray, right: np.ndarray) -> np.ndarray: + return left @ right + + +def _blas_matmul(left: np.ndarray, right: np.ndarray) -> np.ndarray: + return _dgemm(alpha=1.0, a=left, b=right) + + +if sys.platform == "emscripten": + BACKEND = "scipy.linalg.blas.dgemm" + runica_matmul = _blas_matmul +else: + BACKEND = "numpy.matmul" + runica_matmul = _numpy_matmul diff --git a/src/eegprep/plugins/ICLabel/iclabel.onnx b/src/eegprep/plugins/ICLabel/iclabel.onnx new file mode 100644 index 00000000..0b136301 Binary files /dev/null and b/src/eegprep/plugins/ICLabel/iclabel.onnx differ diff --git a/src/eegprep/plugins/ICLabel/iclabel.py b/src/eegprep/plugins/ICLabel/iclabel.py index 1961018b..093b3b1e 100644 --- a/src/eegprep/plugins/ICLabel/iclabel.py +++ b/src/eegprep/plugins/ICLabel/iclabel.py @@ -1,12 +1,17 @@ """ICLabel module for classifying independent components in EEG data.""" from copy import deepcopy -import os +import sys import numpy as np _SUPPORTED_ALGORITHMS = ('default', 'lite', 'beta') +_IS_EMSCRIPTEN = sys.platform == 'emscripten' +SYNC_UNAVAILABLE_MESSAGE = ( + 'ICLabel synchronous entry points are unavailable under Emscripten; ' + 'use await iclabel_async(...) or await pop_iclabel_async(...).' +) def iclabel(EEG, algorithm='default', engine=None): @@ -30,90 +35,110 @@ def iclabel(EEG, algorithm='default', engine=None): EEG : dict EEGLAB EEG structure with ICLabel classifications added """ + if _IS_EMSCRIPTEN: + raise RuntimeError(SYNC_UNAVAILABLE_MESSAGE) + return _iclabel_sync(EEG, algorithm=algorithm, engine=engine) + + +async def iclabel_async(EEG, algorithm='default', engine=None): + """Apply ICLabel through an awaitable native or browser backend. + + The browser backend is selected once when the ICLabel module is loaded and + awaits the ONNX Runtime Web Promise. Native callers may use this entry + point as an asynchronous spelling of the same pure processing operation. + """ + algorithm = _normalize_algorithm(algorithm) + EEG = deepcopy(EEG) + + if engine in ['matlab', 'octave']: + if _IS_EMSCRIPTEN: + raise RuntimeError(SYNC_UNAVAILABLE_MESSAGE) + return _iclabel_sync(EEG, algorithm=algorithm, engine=engine) + if engine is not None: + raise ValueError(f"Unsupported engine: {engine}. Should be None, 'matlab', or 'octave'") + if algorithm != 'default': + raise NotImplementedError( + "EEGPrep standalone Python ICLabel only ships the default network (iclabel.onnx). " + f"The '{algorithm}' network is available only with engine='matlab' or engine='octave' " + "and an EEGLAB ICLabel checkout that provides that artifact." + ) + + from eegprep.plugins.ICLabel.iclabel_net_onnx import run_iclabel_net_async + + image, psdmed, autocorr = _prepare_features(EEG) + output_np = await run_iclabel_net_async(image, psdmed, autocorr) + return _attach_classification(EEG, output_np, algorithm) + + +def _iclabel_sync(EEG, algorithm='default', engine=None): algorithm = _normalize_algorithm(algorithm) EEG = deepcopy(EEG) - # Check if using MATLAB or Octave implementation if engine in ['matlab', 'octave']: from eegprep.functions.adminfunc.eeglabcompat import get_eeglab - # Determine which engine to use runtime = 'MAT' if engine == 'matlab' else 'OCT' eeglab = get_eeglab(runtime=runtime) - - # Run ICLabel using MATLAB/Octave, passing the algorithm parameter if algorithm == 'default': return eeglab.iclabel(EEG) - else: - return eeglab.iclabel(EEG, algorithm) - - # Default Python implementation - elif engine is None: - if algorithm != 'default': - raise NotImplementedError( - "EEGPrep standalone Python ICLabel only ships the default network (netICL.mat). " - f"The '{algorithm}' network is available only with engine='matlab' or engine='octave' " - "and an EEGLAB ICLabel checkout that provides that artifact." - ) - try: - import torch - except ImportError as e: - raise ImportError( - f"PyTorch is not installed in your environment ({e}). " - f"To include torch, install eegprep as eegprep[all] or " - f"install the torch package manually (see Getting Started " - f"on pytorch.org for specifics for your platform)." - ) from e - - from eegprep.plugins.ICLabel.iclabel_net import ICLabelNet - from eegprep import ICL_feature_extractor - - # ICLABEL Extract ICLabel features from an EEG dataset. - features = ICL_feature_extractor(EEG, True) - - # Equivalent of MATLAB code reshaping - features[0] = np.single( - np.concatenate([features[0], -features[0], features[0][:, ::-1, :, :], -features[0][:, ::-1, :, :]], axis=3) - ) - features[1] = np.single(np.tile(features[1], (1, 1, 1, 4))) - features[2] = np.single(np.tile(features[2], (1, 1, 1, 4))) - # print('Feature 0 shape:', features[0].shape) - # print('Feature 1 shape:', features[1].shape) - # print('Feature 2 shape:', features[2].shape) - - # Load the ICLabelNet model - base_dir = os.path.dirname(os.path.abspath(__file__)) - data_path = os.path.join(base_dir, 'netICL.mat') - model = ICLabelNet(data_path) - - # Convert the features to torch tensors - image = torch.tensor(features[0]).permute(-1, 2, 0, 1) - psdmed = torch.tensor(features[1]).permute(-1, 2, 0, 1) - autocorr = torch.tensor(features[2]).permute(-1, 2, 0, 1) - - # Get the output from the model - output = model(image, psdmed, autocorr) - output_np = output.detach().numpy() - output_np = output_np.T # Transpose the array - output_np = np.reshape(output_np, (-1, 4), order='F') # Reshape to have 4 columns - output_np = np.mean(output_np, axis=1) # Compute the mean along the second axis (columns) - output_np = np.reshape(output_np, (7, -1), order='F') # Reshape to have 7 rows - output_np = output_np.T # Transpose back - - if 'ic_classification' not in EEG['etc']: - EEG['etc']['ic_classification'] = {} - if 'ICLabel' not in EEG['etc']['ic_classification']: - EEG['etc']['ic_classification']['ICLabel'] = {} - - EEG['etc']['ic_classification']['ICLabel']['classes'] = np.array( - ['Brain', 'Muscle', 'Eye', 'Heart', 'Line Noise', 'Channel Noise', 'Other'], dtype=object + return eeglab.iclabel(EEG, algorithm) + if engine is not None: + raise ValueError(f"Unsupported engine: {engine}. Should be None, 'matlab', or 'octave'") + if algorithm != 'default': + raise NotImplementedError( + "EEGPrep standalone Python ICLabel only ships the default network (iclabel.onnx). " + f"The '{algorithm}' network is available only with engine='matlab' or engine='octave' " + "and an EEGLAB ICLabel checkout that provides that artifact." ) - EEG['etc']['ic_classification']['ICLabel']['classifications'] = output_np - EEG['etc']['ic_classification']['ICLabel']['version'] = algorithm - return EEG - else: - raise ValueError(f"Unsupported engine: {engine}. Should be None, 'matlab', or 'octave'") + from eegprep.plugins.ICLabel.iclabel_net_onnx import run_iclabel_net + + image, psdmed, autocorr = _prepare_features(EEG) + output_np = run_iclabel_net(image, psdmed, autocorr) + return _attach_classification(EEG, output_np, algorithm) + + +def _prepare_features(EEG): + from eegprep import ICL_feature_extractor + + return _prepare_network_inputs(ICL_feature_extractor(EEG, True)) + + +def _prepare_network_inputs(features): + """Apply ICLabel augmentation and convert feature arrays to network inputs.""" + topo, psdmed, autocorr = features + topo = np.single(np.concatenate([topo, -topo, topo[:, ::-1, :, :], -topo[:, ::-1, :, :]], axis=3)) + psdmed = np.single(np.tile(psdmed, (1, 1, 1, 4))) + autocorr = np.single(np.tile(autocorr, (1, 1, 1, 4))) + return ( + np.transpose(topo, (3, 2, 0, 1)), + np.transpose(psdmed, (3, 2, 0, 1)), + np.transpose(autocorr, (3, 2, 0, 1)), + ) + + +def _postprocess_network_output(output_np): + """Average the four augmented network outputs into component probabilities.""" + output_np = output_np.T + output_np = np.reshape(output_np, (-1, 4), order='F') + output_np = np.mean(output_np, axis=1) + return np.reshape(output_np, (7, -1), order='F').T + + +def _attach_classification(EEG, output_np, algorithm): + output_np = _postprocess_network_output(output_np) + + if 'ic_classification' not in EEG['etc']: + EEG['etc']['ic_classification'] = {} + if 'ICLabel' not in EEG['etc']['ic_classification']: + EEG['etc']['ic_classification']['ICLabel'] = {} + + EEG['etc']['ic_classification']['ICLabel']['classes'] = np.array( + ['Brain', 'Muscle', 'Eye', 'Heart', 'Line Noise', 'Channel Noise', 'Other'], dtype=object + ) + EEG['etc']['ic_classification']['ICLabel']['classifications'] = output_np + EEG['etc']['ic_classification']['ICLabel']['version'] = algorithm + return EEG def _normalize_algorithm(algorithm): diff --git a/src/eegprep/plugins/ICLabel/iclabel_net_onnx.py b/src/eegprep/plugins/ICLabel/iclabel_net_onnx.py new file mode 100644 index 00000000..2ff53bcf --- /dev/null +++ b/src/eegprep/plugins/ICLabel/iclabel_net_onnx.py @@ -0,0 +1,76 @@ +"""onnxruntime backend for the packaged ICLabel default network. + +This is the runtime counterpart to :mod:`iclabel_net` (the torch definition +used only to build and export the network). It loads ``iclabel.onnx``, the +selected packaged artifact produced by the Phase 6 quantization pipeline, and +runs it through onnxruntime so ICLabel classification does not require torch. +""" + +import os +import sys + +import numpy as np + +_INPUT_NAMES = ('image', 'psdmed', 'autocorr') +_OUTPUT_NAME = 'output' +_IS_EMSCRIPTEN = sys.platform == 'emscripten' + +_session = None + + +def _get_session(): + global _session + if _session is not None: + return _session + try: + import onnxruntime as ort + except ImportError as e: + raise ImportError( + f"onnxruntime is not installed in your environment ({e}). " + f"To include onnxruntime, install eegprep as eegprep[iclabel] or " + f"eegprep[all]." + ) from e + base_dir = os.path.dirname(os.path.abspath(__file__)) + model_path = os.path.join(base_dir, 'iclabel.onnx') + _session = ort.InferenceSession(model_path, providers=['CPUExecutionProvider']) + return _session + + +def run_iclabel_net(image, psdmed, autocorr): + """Run the packaged ICLabel network through onnxruntime. + + Parameters + ---------- + image, psdmed, autocorr : numpy.ndarray + NCHW float32 arrays, matching the torch ``ICLabelNet.forward`` inputs. + + Returns + ------- + numpy.ndarray + Network output, shaped like the torch model's output. + """ + if _IS_EMSCRIPTEN: + raise RuntimeError( + "ICLabel synchronous ONNX execution is unavailable under Emscripten; use await run_iclabel_net_async(...)." + ) + session = _get_session() + inputs = { + _INPUT_NAMES[0]: np.asarray(image, dtype=np.float32), + _INPUT_NAMES[1]: np.asarray(psdmed, dtype=np.float32), + _INPUT_NAMES[2]: np.asarray(autocorr, dtype=np.float32), + } + (output,) = session.run([_OUTPUT_NAME], inputs) + return output + + +if _IS_EMSCRIPTEN: + from eegprep.plugins.ICLabel.iclabel_net_onnx_web import run_iclabel_net_async as _run_iclabel_net_async +else: + + async def _run_iclabel_net_async(image, psdmed, autocorr): + return run_iclabel_net(image, psdmed, autocorr) + + +async def run_iclabel_net_async(image, psdmed, autocorr): + """Run the platform-selected ICLabel backend asynchronously.""" + return await _run_iclabel_net_async(image, psdmed, autocorr) diff --git a/src/eegprep/plugins/ICLabel/iclabel_net_onnx_web.py b/src/eegprep/plugins/ICLabel/iclabel_net_onnx_web.py new file mode 100644 index 00000000..ffe5fa80 --- /dev/null +++ b/src/eegprep/plugins/ICLabel/iclabel_net_onnx_web.py @@ -0,0 +1,29 @@ +"""ONNX Runtime Web adapter used by the Pyodide/Emscripten backend. + +The JavaScript host registers ``eegprep_iclabel_web`` before importing this +module. Its ``run`` method accepts flat Float32Array inputs and their shapes, +then returns a Promise resolving to the flat output Float32Array. Keeping the +adapter at this boundary leaves the ICLabel processing code independent of the +browser host and avoids importing native ``onnxruntime`` in Pyodide. +""" + +import numpy as np + + +async def run_iclabel_net_async(image, psdmed, autocorr): + """Run ICLabel through the host's asynchronous ONNX Runtime Web bridge.""" + from js import eegprep_iclabel_web # ty: ignore[unresolved-import] + from pyodide.ffi import to_js # ty: ignore[unresolved-import] + + image = np.ascontiguousarray(image, dtype=np.float32) + psdmed = np.ascontiguousarray(psdmed, dtype=np.float32) + autocorr = np.ascontiguousarray(autocorr, dtype=np.float32) + output = await eegprep_iclabel_web.run( + to_js(image.reshape(-1)), + to_js(list(image.shape)), + to_js(psdmed.reshape(-1)), + to_js(list(psdmed.shape)), + to_js(autocorr.reshape(-1)), + to_js(list(autocorr.shape)), + ) + return np.asarray(output.to_py(), dtype=np.float32).reshape((-1, 7, 1, 1)) diff --git a/src/eegprep/plugins/ICLabel/pop_iclabel.py b/src/eegprep/plugins/ICLabel/pop_iclabel.py index 092286aa..e67d5733 100644 --- a/src/eegprep/plugins/ICLabel/pop_iclabel.py +++ b/src/eegprep/plugins/ICLabel/pop_iclabel.py @@ -6,7 +6,7 @@ from eegprep.functions.guifunc.inputgui import inputgui from eegprep.functions.guifunc.spec import ControlSpec, DialogSpec -from eegprep.plugins.ICLabel.iclabel import iclabel +from eegprep.plugins.ICLabel.iclabel import SYNC_UNAVAILABLE_MESSAGE, _IS_EMSCRIPTEN, iclabel, iclabel_async _VERSIONS = ("default", "lite", "beta") @@ -22,6 +22,59 @@ def pop_iclabel( return_com: bool = False, ): """Classify independent components using ICLabel.""" + if _IS_EMSCRIPTEN: + raise RuntimeError(SYNC_UNAVAILABLE_MESSAGE) + return _pop_iclabel_sync( + EEG, + icversion=icversion, + gui=gui, + renderer=renderer, + engine=engine, + return_com=return_com, + ) + + +async def pop_iclabel_async( + EEG, + icversion: str | None = None, + *, + gui: bool | None = None, + renderer=None, + engine=None, + return_com: bool = False, +): + """Classify independent components through an awaitable ICLabel backend.""" + if EEG is None: + return (None, "") if return_com else None + if gui is None: + gui = icversion is None + if gui: + result = _run_gui(renderer=renderer) + if result is None: + return (EEG, "") if return_com else EEG + icversion = result["icversion"] + icversion = "default" if icversion is None else str(icversion).lower() + if icversion not in _VERSIONS: + raise ValueError("icversion must be one of 'default', 'lite', or 'beta'") + if isinstance(EEG, list): + output = [await pop_iclabel_async(item, icversion, gui=False, engine=engine) for item in EEG] + command = _history_command(icversion, asynchronous=True) + return (output, command) if return_com else output + _require_ica(EEG) + output = await iclabel_async(EEG, algorithm=icversion, engine=engine) + command = _history_command(icversion, asynchronous=True) + return (output, command) if return_com else output + + +def _pop_iclabel_sync( + EEG, + *, + icversion: str | None, + gui: bool | None, + renderer, + engine, + return_com: bool, +): if EEG is None: return (None, "") if return_com else None if gui is None: @@ -75,5 +128,6 @@ def _require_ica(EEG): raise ValueError("ICLabel requires an ICA decomposition. Run pop_runica first.") -def _history_command(icversion): - return f"EEG = pop_iclabel(EEG, '{icversion}');" +def _history_command(icversion, *, asynchronous=False): + await_prefix = "await " if asynchronous else "" + return f"EEG = {await_prefix}pop_iclabel{'' if not asynchronous else '_async'}(EEG, '{icversion}');" diff --git a/src/eegprep/plugins/clean_rawdata/asr_process.py b/src/eegprep/plugins/clean_rawdata/asr_process.py index c77b0134..572142e3 100644 --- a/src/eegprep/plugins/clean_rawdata/asr_process.py +++ b/src/eegprep/plugins/clean_rawdata/asr_process.py @@ -10,6 +10,12 @@ logger = logging.getLogger(__name__) +# Fixed memory budget (MB) used when max_mem is not given. Matches the maxmem=64 default +# already used by asr_calibrate() and clean_asr() elsewhere in this plugin, so the whole ASR +# pipeline assumes the same memory budget without probing the runtime environment (psutil is +# an optional dependency and has no WebAssembly build). +DEFAULT_MAX_MEM_MB = 64 + def asr_process( data, srate, state, window_len=0.5, lookahead=None, step_size=32, max_dims=0.66, max_mem=None, use_gpu=False @@ -37,8 +43,8 @@ def asr_process( Max: window_len * srate. Default: 32. max_dims (float or int, optional): Maximum dimensions/fraction of dimensions to remove. Default: 0.66 (fraction). - max_mem (int, optional): Maximum memory in MB for processing large chunks. Process in one block if None. - Default: None. + max_mem (int, optional): Maximum memory in MB for processing large chunks. + Default: None, which resolves to DEFAULT_MAX_MEM_MB (64 MB). use_gpu (bool, optional): Whether to use GPU (not implemented). Default: False. Returns @@ -60,10 +66,7 @@ def asr_process( if lookahead is None: lookahead = window_len / 2 if max_mem is None: - # use at most half of available memory - import psutil - - max_mem = psutil.virtual_memory().free / 1024**2 / 2 + max_mem = DEFAULT_MAX_MEM_MB # Ensure window length is adequate window_len = max(window_len, 1.5 * C / srate) @@ -108,15 +111,7 @@ def asr_process( # Calculate number of splits for memory management if max_mem * 1024 * 1024 - C * C * P * 8 * 3 < 0: - logger.warning( - "Memory too low, increasing it (rejection block size now " - "depends on available memory so it might not be fully reproducible)..." - ) - import psutil - - max_mem = psutil.virtual_memory().free / 1024**2 / 2 - if max_mem * 1024 * 1024 - C * C * P * 8 * 3 < 0: - raise RuntimeError('Not enough memory') + raise RuntimeError('Not enough memory') # Calculate memory bytes needed (following reference implementation formula) bytes_needed = C * C * S * 8 * 8 + C * C * 8 * S / step_size + C * S * 8 * 2 + S * 8 * 5 diff --git a/src/eegprep/plugins/clean_rawdata/private/ransac.py b/src/eegprep/plugins/clean_rawdata/private/ransac.py index 7dfa0ebe..91bba772 100644 --- a/src/eegprep/plugins/clean_rawdata/private/ransac.py +++ b/src/eegprep/plugins/clean_rawdata/private/ransac.py @@ -4,10 +4,16 @@ import numpy as np -from ....functions.adminfunc.eeglabcompat import get_eeglab from .sphericalSplineInterpolate import sphericalSplineInterpolate +def get_eeglab(*args, **kwargs): + """Load the EEGLAB bridge only when a MATLAB or Octave path needs it.""" + from ....functions.adminfunc.eeglabcompat import get_eeglab as load_eeglab + + return load_eeglab(*args, **kwargs) + + def rand_sample(n: int, m: int, stream: np.random.RandomState) -> np.ndarray: """Random sampling without replacement using Fisher-Yates shuffle. diff --git a/src/eegprep/resources/help/pop_iclabel.md b/src/eegprep/resources/help/pop_iclabel.md index bdbfc670..f87515eb 100644 --- a/src/eegprep/resources/help/pop_iclabel.md +++ b/src/eegprep/resources/help/pop_iclabel.md @@ -5,6 +5,8 @@ Usage: EEG = pop_iclabel(EEG) EEG = pop_iclabel(EEG, 'default') EEG, command = pop_iclabel(EEG, 'default', return_com=True) + EEG = await pop_iclabel_async(EEG, 'default') + EEG, command = await pop_iclabel_async(EEG, 'default', return_com=True) Inputs: @@ -22,10 +24,27 @@ Behavior: - Results are stored in `EEG.etc.ic_classification.ICLabel`. - The result includes the ICLabel class names, per-component class probabilities, and the selected version string. - Lists of datasets are processed one dataset at a time with the same selected version. -- Standalone Python EEGPrep ships the default ICLabel network (`netICL.mat`). - The EEGLAB `lite` and `beta` network artifacts are explicit MATLAB/Octave - passthrough choices and raise a clear limitation when requested with the - standalone Python engine. +- In Pyodide/Emscripten, use `await iclabel_async(EEG)` or + `await pop_iclabel_async(EEG, 'default')`. The synchronous `iclabel` and + `pop_iclabel` entry points fail fast because ONNX Runtime Web is asynchronous. +- The browser host must register EEGPrep's ONNX Runtime Web bridge before + calling the async entry point. A host may run the Pyodide environment in a + Web Worker to keep the UI responsive. +- Standalone Python EEGPrep ships the gate-selected weight-only int8 ICLabel + network as an ONNX artifact (`iclabel.onnx`) and classifies through + `onnxruntime`; install the `iclabel` extra (`eegprep[iclabel]`) to run it. + Feature extraction, float32 input normalization, augmentation, and output + softmax are unchanged. On the frozen 217-component, subject-disjoint set in + `tools/iclabel/evaluation_manifest.json`, the shipped artifact agreed with + the preserved float32 teacher on 100% of top-1 labels and 100% of existing + `pop_icflag` keep-or-reject decisions, with a maximum probability drift of + 0.01346. The calibrated candidate measured 98.1567% top-1 and 100% + keep-or-reject but exceeded the frozen 0.015 probability-drift gate; the + smaller weight-only candidate is the shipped default. These are measured + frozen-set parity results, not a general accuracy claim. The EEGLAB `lite` + and `beta` network artifacts are + explicit MATLAB/Octave passthrough choices and raise a clear limitation when + requested with the standalone Python engine. Example: @@ -36,3 +55,6 @@ Notes: - ICLabel class probabilities are ordered as Brain, Muscle, Eye, Heart, Line Noise, Channel Noise, and Other. - This wrapper uses EEGPrep's packaged ICLabel implementation and does not require an EEGLAB checkout at runtime. +- Async history commands are replayable from `eegprep-console`; for example, + use `await eegh(1)` when the selected history entry contains + `pop_iclabel_async`. diff --git a/tests/conftest.py b/tests/conftest.py index 13c1d9bd..934434e3 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -68,7 +68,6 @@ def _preload_matlab_libstdcxx() -> None: "tests/test_eeg_rpsd_parity.py", "tests/test_eegfindboundaries.py", "tests/test_envtopo_parity.py", - "tests/test_iclabel.py", "tests/test_iclabel_features.py", "tests/test_parity_rng.py", "tests/test_pinv.py", @@ -95,6 +94,7 @@ def _preload_matlab_libstdcxx() -> None: "tests/test_eeglabcompat.py::TestGetEeglab::", "tests/test_eeglabcompat.py::TestPopEegfiltnew::", "tests/test_epoch.py::TestEpochParity::", + "tests/test_iclabel.py::TestICLabelEngines::", "tests/test_matlab_path.py::TestMatlabPath::test_get_eeglab_mat", "tests/test_matlab_path.py::TestMatlabPath::test_python_matlab_engine", "tests/test_matlab_path.py::TestMatlabPath::test_start_matlab_engine", diff --git a/tests/test_console_workspace.py b/tests/test_console_workspace.py index 1645f895..17196f7b 100644 --- a/tests/test_console_workspace.py +++ b/tests/test_console_workspace.py @@ -1,6 +1,7 @@ from __future__ import annotations import ast +import asyncio from contextlib import redirect_stderr import io import importlib @@ -90,6 +91,13 @@ def _fake_pop_without_command(eeg, *, return_com=False): return (output, "") if return_com else output +async def _fake_pop_iclabel_async(eeg, icversion="default", *, return_com=False): + del icversion + output = dict(eeg, setname="async-iclabel") + command = "EEG = await pop_iclabel_async(EEG, 'default');" + return (output, command) if return_com else output + + def _fake_pop_topoplot(eeg, *, return_com=False): command = "pop_topoplot(EEG, typeplot=1, items=[0])" return (["figure"], command) if return_com else ["figure"] @@ -224,6 +232,653 @@ def test_storedisk_session_retrieve_and_console_pop_call_stay_synchronized(tmp_p EEG_OPTIONS.update(old_options) +def test_console_async_iclabel_pop_updates_shared_session_once(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": _fake_pop_iclabel_async}) + + result = asyncio.run(workspace.namespace["pop_iclabel_async"](workspace.namespace["EEG"])) + workspace.after_execute("await pop_iclabel_async(EEG)") + + assert result.eeg is session.EEG + assert session.EEG["setname"] == "async-iclabel" + assert session.CURRENTSET == [1] + assert session.ALLCOM == ["EEG = await pop_iclabel_async(EEG, 'default');"] + assert workspace.namespace["LASTCOM"] == session.LASTCOM + workspace.close() + + +def test_console_async_iclabel_captures_state_before_coroutine_starts(): + session = EEGPrepSession() + session.store_current(_demo_eeg("first"), new=True) + session.store_current(_demo_eeg("second"), new=True) + session.retrieve(1) + + async def fake_pop(eeg, *, return_com=False): + output = dict(eeg, setname="stale-result") + command = "EEG = await pop_iclabel_async(EEG, 'default');" + return (output, command) if return_com else output + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": fake_pop}) + pending = workspace.namespace["pop_iclabel_async"](workspace.namespace["EEG"]) + workspace.after_execute("pending = pop_iclabel_async(EEG)") + session.retrieve(2) + + with pytest.raises(RuntimeError, match="session changed"): + asyncio.run(pending) + + assert session.EEG["setname"] == "second" + assert session.ALLCOM == [] + workspace.close() + + +def test_console_iclabel_async_export_captures_state_before_coroutine_starts(): + session = EEGPrepSession() + session.store_current(_demo_eeg("first"), new=True) + session.store_current(_demo_eeg("second"), new=True) + session.retrieve(1) + + async def fake_iclabel(eeg, *, algorithm="default", engine=None): + del algorithm, engine + return dict(eeg, setname="stale-result") + + workspace = EEGPrepConsoleWorkspace(session, exports={"iclabel_async": fake_iclabel}) + assert workspace.namespace["eegprep"].iclabel_async is workspace.namespace["iclabel_async"] + pending = workspace.namespace["iclabel_async"](workspace.namespace["EEG"]) + session.retrieve(2) + + with pytest.raises(RuntimeError, match="session changed"): + asyncio.run(pending) + + assert session.EEG["setname"] == "second" + assert session.ALLCOM == [] + workspace.close() + + +@pytest.mark.parametrize( + ("import_line", "local_name"), + [ + ("from eegprep import iclabel_async", "iclabel_async"), + ("from eegprep import iclabel_async as classify", "classify"), + ], +) +def test_console_imported_iclabel_async_uses_freshness_wrapper(import_line, local_name): + session = EEGPrepSession() + session.store_current(_demo_eeg("first"), new=True) + session.store_current(_demo_eeg("second"), new=True) + session.retrieve(1) + + async def fake_iclabel(eeg, *, algorithm="default", engine=None): + del algorithm, engine + return dict(eeg, setname="stale-result") + + workspace = EEGPrepConsoleWorkspace(session, exports={"iclabel_async": fake_iclabel}) + exec(f"{import_line}\npending = {local_name}(EEG)", workspace.namespace) + assert isinstance(workspace.namespace[local_name], console_module.ConsoleAsyncFunction) + session.retrieve(2) + + with pytest.raises(RuntimeError, match="session changed"): + asyncio.run(workspace.namespace["pending"]) + + assert session.EEG["setname"] == "second" + workspace.close() + + +def test_console_imported_eegprep_module_uses_freshness_wrapper(): + session = EEGPrepSession() + session.store_current(_demo_eeg("first"), new=True) + session.store_current(_demo_eeg("second"), new=True) + session.retrieve(1) + + async def fake_iclabel(eeg, *, algorithm="default", engine=None): + del algorithm, engine + return dict(eeg, setname="stale-result") + + workspace = EEGPrepConsoleWorkspace(session, exports={"iclabel_async": fake_iclabel}) + exec("import eegprep as ep\npending = ep.iclabel_async(EEG)", workspace.namespace) + assert isinstance(workspace.namespace["ep"], console_module.ConsoleEEGPrepModule) + session.retrieve(2) + + with pytest.raises(RuntimeError, match="session changed"): + asyncio.run(workspace.namespace["pending"]) + + assert session.EEG["setname"] == "second" + workspace.close() + + +def test_console_imported_iclabel_submodule_uses_freshness_wrapper(): + session = EEGPrepSession() + session.store_current(_demo_eeg("first"), new=True) + session.store_current(_demo_eeg("second"), new=True) + session.retrieve(1) + + async def fake_iclabel(eeg, *, algorithm="default", engine=None): + del algorithm, engine + return dict(eeg, setname="stale-result") + + workspace = EEGPrepConsoleWorkspace(session, exports={"iclabel_async": fake_iclabel}) + exec( + "from eegprep.plugins.ICLabel.iclabel import iclabel_async as classify\npending = classify(EEG)", + workspace.namespace, + ) + assert isinstance(workspace.namespace["classify"], console_module.ConsoleAsyncFunction) + session.retrieve(2) + + with pytest.raises(RuntimeError, match="session changed"): + asyncio.run(workspace.namespace["pending"]) + + assert session.EEG["setname"] == "second" + workspace.close() + + +def test_console_imported_pop_iclabel_submodule_uses_freshness_wrapper(): + session = EEGPrepSession() + session.store_current(_demo_eeg("first"), new=True) + session.store_current(_demo_eeg("second"), new=True) + session.retrieve(1) + + async def fake_pop(eeg, *, return_com=False): + output = dict(eeg, setname="stale-result") + command = "EEG = await pop_iclabel_async(EEG, 'default');" + return (output, command) if return_com else output + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": fake_pop}) + exec( + "from eegprep.plugins.ICLabel.pop_iclabel import pop_iclabel_async as classify\npending = classify(EEG)", + workspace.namespace, + ) + assert isinstance(workspace.namespace["classify"], console_module.ConsolePopFunction) + session.retrieve(2) + + with pytest.raises(RuntimeError, match="session changed"): + asyncio.run(workspace.namespace["pending"]) + + assert session.EEG["setname"] == "second" + workspace.close() + + +def test_console_wildcard_eegprep_import_uses_freshness_wrapper(): + session = EEGPrepSession() + session.store_current(_demo_eeg("first"), new=True) + session.store_current(_demo_eeg("second"), new=True) + session.retrieve(1) + + async def fake_iclabel(eeg, *, algorithm="default", engine=None): + del algorithm, engine + return dict(eeg, setname="stale-result") + + workspace = EEGPrepConsoleWorkspace(session, exports={"iclabel_async": fake_iclabel}) + exec("from eegprep import *\npending = iclabel_async(EEG)", workspace.namespace) + assert isinstance(workspace.namespace["iclabel_async"], console_module.ConsoleAsyncFunction) + session.retrieve(2) + + with pytest.raises(RuntimeError, match="session changed"): + asyncio.run(workspace.namespace["pending"]) + + assert session.EEG["setname"] == "second" + workspace.close() + + +@pytest.mark.parametrize( + "source", + [ + "import eegprep.plugins.ICLabel.iclabel as mod\npending = mod.iclabel_async(EEG)", + "from eegprep.plugins.ICLabel import iclabel\npending = iclabel.iclabel_async(EEG)", + "import importlib\nmod = importlib.import_module('eegprep.plugins.ICLabel.iclabel')\npending = mod.iclabel_async(EEG)", + "from eegprep.plugins.ICLabel.iclabel import *\npending = iclabel_async(EEG)", + "import eegprep.plugins.ICLabel.iclabel\npending = eegprep.plugins.ICLabel.iclabel.iclabel_async(EEG)", + ], +) +def test_console_canonical_iclabel_import_forms_use_freshness_wrapper(source): + session = EEGPrepSession() + session.store_current(_demo_eeg("first"), new=True) + session.store_current(_demo_eeg("second"), new=True) + session.retrieve(1) + + async def fake_iclabel(eeg, *, algorithm="default", engine=None): + del algorithm, engine + return dict(eeg, setname="stale-result") + + workspace = EEGPrepConsoleWorkspace(session, exports={"iclabel_async": fake_iclabel}) + try: + exec(source, workspace.namespace) + assert asyncio.iscoroutine(workspace.namespace["pending"]) + session.retrieve(2) + + with pytest.raises(RuntimeError, match="session changed"): + asyncio.run(workspace.namespace["pending"]) + + assert session.EEG["setname"] == "second" + finally: + workspace.close() + + +def test_console_async_iclabel_rejects_raw_session_alias_edit(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + + async def scenario(): + started = asyncio.Event() + release = asyncio.Event() + + async def classify(eeg): + started.set() + await release.wait() + return dict(eeg, setname="stale-result") + + workspace = EEGPrepConsoleWorkspace(session, exports={"iclabel_async": classify}) + task = asyncio.create_task(workspace.namespace["iclabel_async"](workspace.namespace["EEG"])) + await started.wait() + session.EEG["setname"] = "edited-through-session" + release.set() + try: + with pytest.raises(RuntimeError, match="session changed"): + await task + finally: + workspace.close() + + asyncio.run(scenario()) + assert session.EEG["setname"] == "edited-through-session" + + +def test_console_async_iclabel_tracks_nested_copy_slice_and_setdefault_edits(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + + async def scenario(): + started = asyncio.Event() + release = asyncio.Event() + + async def classify(eeg): + started.set() + await release.wait() + return dict(eeg, setname="stale-result") + + workspace = EEGPrepConsoleWorkspace(session, exports={"iclabel_async": classify}) + task = asyncio.create_task(workspace.namespace["iclabel_async"](workspace.namespace["EEG"])) + await started.wait() + workspace.namespace["EEG"]["event"] = [{"latency": 1.0}] + workspace.namespace["EEG"]["event"][0:1][0]["latency"] = 2.0 + workspace.namespace["EEG"].copy()["event"][0]["latency"] = 3.0 + workspace.namespace["EEG"].setdefault("etc", {})["edited"] = True + release.set() + try: + with pytest.raises(RuntimeError, match="session changed"): + await task + finally: + workspace.close() + + asyncio.run(scenario()) + assert session.EEG["event"] == [{"latency": 3.0}] + assert session.EEG["etc"] == {"edited": True} + + +def test_console_async_iclabel_discards_result_after_in_place_console_edit(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + + async def scenario(): + started = asyncio.Event() + release = asyncio.Event() + + async def stale_pop(eeg, *, return_com=False): + started.set() + await release.wait() + output = dict(eeg, setname="stale-result") + command = "EEG = await pop_iclabel_async(EEG, 'default');" + return (output, command) if return_com else output + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": stale_pop}) + task = asyncio.create_task(workspace.namespace["pop_iclabel_async"](workspace.namespace["EEG"])) + await started.wait() + workspace.namespace["EEG"]["data"][0, 0] = 99.0 + workspace.namespace["EEG"]["markers"] = [] + workspace.namespace["EEG"]["markers"].append({"latency": 1.0}) + release.set() + try: + with pytest.raises(RuntimeError, match="session changed"): + await task + finally: + workspace.close() + + asyncio.run(scenario()) + assert session.EEG["data"][0, 0] == 99.0 + assert session.EEG["markers"] == [{"latency": 1.0}] + assert session.ALLCOM == [] + + +def test_console_concurrent_async_iclabel_calls_keep_shared_mutation_watch(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + + async def scenario(): + started = asyncio.Event() + first_release = asyncio.Event() + second_release = asyncio.Event() + call_number = 0 + + async def classify(eeg): + nonlocal call_number + call_number += 1 + current_call = call_number + if current_call == 2: + started.set() + await (first_release if current_call == 1 else second_release).wait() + return dict(eeg) + + workspace = EEGPrepConsoleWorkspace(session, exports={"iclabel_async": classify}) + first = asyncio.create_task(workspace.namespace["iclabel_async"](workspace.namespace["EEG"])) + second = asyncio.create_task(workspace.namespace["iclabel_async"](workspace.namespace["EEG"])) + await started.wait() + first_release.set() + await first + workspace.namespace["EEG"]["setname"] = "edited-during-second-call" + second_release.set() + try: + with pytest.raises(RuntimeError, match="session changed"): + await second + finally: + workspace.close() + + asyncio.run(scenario()) + assert session.EEG["setname"] == "edited-during-second-call" + + +def test_console_async_iclabel_rejects_alleeg_alias_edit(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + + async def scenario(): + started = asyncio.Event() + release = asyncio.Event() + + async def stale_pop(eeg, *, return_com=False): + started.set() + await release.wait() + output = dict(eeg, setname="stale-result") + command = "EEG = await pop_iclabel_async(EEG, 'default');" + return (output, command) if return_com else output + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": stale_pop}) + task = asyncio.create_task(workspace.namespace["pop_iclabel_async"](workspace.namespace["EEG"])) + await started.wait() + workspace.namespace["ALLEEG"][0]["setname"] = "edited-through-alleeg" + release.set() + try: + with pytest.raises(RuntimeError, match="session changed"): + await task + finally: + workspace.close() + + asyncio.run(scenario()) + assert session.EEG["setname"] == "edited-through-alleeg" + assert session.ALLCOM == [] + + +def test_console_async_iclabel_history_sync_keeps_mutation_watch(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + + async def scenario(): + started = asyncio.Event() + release = asyncio.Event() + + async def stale_pop(eeg, *, return_com=False): + started.set() + await release.wait() + output = dict(eeg, setname="stale-result") + command = "EEG = await pop_iclabel_async(EEG, 'default');" + return (output, command) if return_com else output + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": stale_pop}) + task = asyncio.create_task(workspace.namespace["pop_iclabel_async"](workspace.namespace["EEG"])) + await started.wait() + workspace.namespace["eegh"]("plot(EEG);", workspace.namespace["EEG"]) + workspace.namespace["EEG"]["setname"] = "edited-after-history" + release.set() + try: + with pytest.raises(RuntimeError, match="session changed"): + await task + finally: + workspace.close() + + asyncio.run(scenario()) + assert session.EEG["setname"] == "edited-after-history" + assert session.ALLCOM == ["plot(EEG);"] + + +def test_console_async_iclabel_tracks_in_place_multi_dataset_edit(): + session = EEGPrepSession() + session.store_current(_demo_eeg("first"), new=True) + session.store_current(_demo_eeg("second"), new=True) + session.retrieve([1, 2]) + + async def scenario(): + started = asyncio.Event() + release = asyncio.Event() + + async def stale_pop(eeg, *, return_com=False): + started.set() + await release.wait() + output = [dict(item, setname="stale-result") for item in eeg] + command = "EEG = await pop_iclabel_async(EEG, 'default');" + return (output, command) if return_com else output + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": stale_pop}) + task = asyncio.create_task(workspace.namespace["pop_iclabel_async"](workspace.namespace["EEG"])) + await started.wait() + workspace.namespace["EEG"][1]["setname"] = "edited" + release.set() + try: + with pytest.raises(RuntimeError, match="session changed"): + await task + finally: + workspace.close() + + asyncio.run(scenario()) + assert [item["setname"] for item in session.EEG] == ["first", "edited"] + assert session.ALLCOM == [] + + +def test_console_in_place_pop_result_marks_dataset_changed(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + token = session.dataset_state_token() + + def save_in_place(eeg): + eeg["saved"] = "yes" + return eeg + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_saveset": save_in_place}) + workspace.namespace["pop_saveset"](workspace.namespace["EEG"]) + + assert not session.dataset_state_unchanged(token) + assert session._dataset_revision == 2 + assert session.EEG["saved"] == "yes" + workspace.close() + + +def test_console_cancelled_pop_does_not_mark_dataset_changed(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + token = session.dataset_state_token() + + def cancelled_pop(eeg, *, return_com=False): + return (eeg, "") if return_com else eeg + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_rmbase": cancelled_pop}) + workspace.namespace["pop_rmbase"](workspace.namespace["EEG"]) + + assert session.dataset_state_unchanged(token) + workspace.close() + + +def test_console_async_iclabel_ignores_noop_dataset_pop(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + + async def scenario(): + started = asyncio.Event() + release = asyncio.Event() + + async def classify(eeg, *, return_com=False): + started.set() + await release.wait() + output = dict(eeg, setname="classified") + command = "EEG = await pop_iclabel_async(EEG, 'default');" + return (output, command) if return_com else output + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": classify}) + task = asyncio.create_task(workspace.namespace["pop_iclabel_async"](workspace.namespace["EEG"])) + await started.wait() + token = session.dataset_state_token() + workspace.namespace["EEG"].pop("missing", None) + assert session.dataset_state_unchanged(token) + release.set() + try: + await task + finally: + workspace.close() + + asyncio.run(scenario()) + assert session.EEG["setname"] == "classified" + + +def test_console_async_iclabel_pop_preserves_ordered_dataset_selection(): + session = EEGPrepSession() + session.store_current(_demo_eeg("first"), new=True) + session.store_current(_demo_eeg("second"), new=True) + session.retrieve([2, 1]) + + async def fake_pop(eeg, *, return_com=False): + output = [dict(item, setname=f"updated-{item['setname']}") for item in eeg] + command = "EEG = await pop_iclabel_async(EEG, 'default');" + return (output, command) if return_com else output + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": fake_pop}) + asyncio.run(workspace.namespace["pop_iclabel_async"](workspace.namespace["EEG"])) + workspace.after_execute("await pop_iclabel_async(EEG)") + + assert session.CURRENTSET == [2, 1] + assert [session.ALLEEG[index - 1]["setname"] for index in session.CURRENTSET] == ["updated-second", "updated-first"] + assert session.ALLCOM == ["EEG = await pop_iclabel_async(EEG, 'default');"] + workspace.close() + + +def test_console_async_iclabel_history_replay_awaits_without_duplicate_history(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + session.add_history("EEG = await pop_iclabel_async(EEG, 'default');") + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": _fake_pop_iclabel_async}) + + replayed = asyncio.run(workspace.namespace["eegh"](1)) + + assert replayed == "EEG = await pop_iclabel_async(EEG, 'default');" + assert session.EEG["setname"] == "async-iclabel" + assert session.ALLCOM == ["EEG = await pop_iclabel_async(EEG, 'default');"] + workspace.close() + + +def test_console_async_iclabel_failure_does_not_mutate_session(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + + async def failing_pop(_eeg, *, return_com=False): + del return_com + raise RuntimeError("browser inference failed") + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": failing_pop}) + with pytest.raises(RuntimeError, match="browser inference failed"): + asyncio.run(workspace.namespace["pop_iclabel_async"](workspace.namespace["EEG"])) + + assert session.EEG["setname"] == "demo" + assert session.CURRENTSET == [1] + assert session.ALLCOM == [] + workspace.close() + + +def test_console_async_iclabel_discards_result_after_session_changes(): + session = EEGPrepSession() + session.store_current(_demo_eeg("first"), new=True) + + async def stale_pop(eeg, *, return_com=False): + session.store_current(_demo_eeg("second"), new=True) + output = dict(eeg, setname="stale-result") + command = "EEG = await pop_iclabel_async(EEG, 'default');" + return (output, command) if return_com else output + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": stale_pop}) + with pytest.raises(RuntimeError, match="session changed"): + asyncio.run(workspace.namespace["pop_iclabel_async"](workspace.namespace["EEG"])) + + assert session.EEG["setname"] == "second" + assert session.CURRENTSET == [2] + assert session.ALLCOM == [] + workspace.close() + + +def test_console_async_iclabel_discards_result_after_in_place_dataset_edit(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + + async def scenario(): + started = asyncio.Event() + release = asyncio.Event() + + async def stale_pop(eeg, *, return_com=False): + started.set() + await release.wait() + output = dict(eeg, setname="stale-result") + command = "EEG = await pop_iclabel_async(EEG, 'default');" + return (output, command) if return_com else output + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": stale_pop}) + task = asyncio.create_task(workspace.namespace["pop_iclabel_async"](workspace.namespace["EEG"])) + await started.wait() + session.mark_current_saved() + release.set() + try: + with pytest.raises(RuntimeError, match="session changed"): + await task + finally: + workspace.close() + + asyncio.run(scenario()) + assert session.EEG["setname"] == "demo" + assert session.EEG["saved"] == "yes" + assert session.ALLCOM == [] + + +def test_console_async_iclabel_allows_history_change_during_inference(): + session = EEGPrepSession() + session.store_current(_demo_eeg(), new=True) + + async def scenario(): + started = asyncio.Event() + release = asyncio.Event() + + async def classify(eeg, *, return_com=False): + started.set() + await release.wait() + output = dict(eeg, setname="classified") + command = "EEG = await pop_iclabel_async(EEG, 'default');" + return (output, command) if return_com else output + + workspace = EEGPrepConsoleWorkspace(session, exports={"pop_iclabel_async": classify}) + task = asyncio.create_task(workspace.namespace["pop_iclabel_async"](workspace.namespace["EEG"])) + await started.wait() + workspace.namespace["eegh"]("plot(EEG);", workspace.namespace["EEG"]) + release.set() + try: + await task + finally: + workspace.close() + + asyncio.run(scenario()) + assert session.EEG["setname"] == "classified" + assert session.ALLCOM == ["plot(EEG);", "EEG = await pop_iclabel_async(EEG, 'default');"] + + def test_console_currentset_reassignment_preserves_both_datasets(): session = EEGPrepSession() session.store_current(_demo_eeg("first"), new=True) @@ -1298,6 +1953,27 @@ def test_console_restores_pop_wrappers_after_from_import(): assert workspace.namespace["pop_reref"] is original_wrapper +@pytest.mark.parametrize( + ("import_line", "local_name"), + [ + ("from eegprep import iclabel_async", "iclabel_async"), + ("from eegprep import iclabel_async as classify", "classify"), + ], +) +def test_console_restores_imported_iclabel_async_wrappers(import_line, local_name): + session = EEGPrepSession() + + async def fake_iclabel(eeg): + return eeg + + workspace = EEGPrepConsoleWorkspace(session, exports={"iclabel_async": fake_iclabel}) + workspace.namespace[local_name] = fake_iclabel + workspace.after_execute(import_line) + + assert isinstance(workspace.namespace[local_name], console_module.ConsoleAsyncFunction) + workspace.close() + + def test_console_restores_aliased_pop_wrapper_after_from_import(): session = EEGPrepSession() session.store_current(_demo_eeg(), new=True) diff --git a/tests/test_gui_main_window.py b/tests/test_gui_main_window.py index 592362ba..66592a40 100644 --- a/tests/test_gui_main_window.py +++ b/tests/test_gui_main_window.py @@ -1,4 +1,5 @@ import ast +import asyncio import inspect import os import logging @@ -2017,3 +2018,38 @@ def dispatch_gui(self, action_id, parent): self.assertEqual(dispatcher.actions, [("pop_loadset", window.window)]) self.assertEqual(branding_calls, ["branding"]) window.window.close() + + def test_gui_main_window_schedules_async_menu_action(self): + from eegprep.functions.guifunc.main_window import EEGPrepMainWindow + + async def run(): + class AsyncDispatcher: + def __init__(self): + self.calls = [] + self.completed = False + + def dispatch_gui(self, action_id, parent): + self.calls.append((action_id, parent)) + + async def complete(): + self.completed = True + + return complete() + + window = EEGPrepMainWindow.__new__(EEGPrepMainWindow) + dispatcher = AsyncDispatcher() + window.window = object() + window.dispatcher = dispatcher + window._queue_application_branding = lambda: None + window._async_menu_tasks = set() + try: + window._dispatch_menu_action("pop_iclabel") + + self.assertEqual(dispatcher.calls, [("pop_iclabel", window.window)]) + self.assertFalse(dispatcher.completed) + await asyncio.sleep(0) + self.assertTrue(dispatcher.completed) + finally: + window._async_menu_tasks.clear() + + asyncio.run(run()) diff --git a/tests/test_gui_pop_iclabel.py b/tests/test_gui_pop_iclabel.py index 4b180fd7..10cfca2b 100644 --- a/tests/test_gui_pop_iclabel.py +++ b/tests/test_gui_pop_iclabel.py @@ -1,18 +1,34 @@ +import asyncio import unittest from unittest import mock import numpy as np -from eegprep.plugins.ICLabel.pop_iclabel import pop_iclabel, pop_iclabel_dialog_spec +import eegprep.plugins.ICLabel.pop_iclabel as pop_iclabel_module +import eegprep.functions.guifunc.menu_actions as menu_actions_module +from eegprep.functions.guifunc.menu_actions import MenuActionDispatcher +from eegprep.functions.guifunc.session import EEGPrepSession +from eegprep.plugins.ICLabel.pop_iclabel import pop_iclabel, pop_iclabel_async, pop_iclabel_dialog_spec -def _eeg(): +def _eeg(setname="demo"): return { "data": np.zeros((2, 20), dtype=np.float32), "nbchan": 2, "pnts": 20, "trials": 1, "srate": 100, + "xmin": 0.0, + "xmax": 0.19, + "times": np.arange(20) / 100, + "event": [], + "urevent": [], + "epoch": [], + "chanlocs": [], + "chaninfo": {}, + "setname": setname, + "history": "", + "ref": "", "icaweights": np.eye(2), "icasphere": np.eye(2), "icawinv": np.eye(2), @@ -75,6 +91,98 @@ def test_missing_ica_raises_clear_error(self): with self.assertRaisesRegex(ValueError, "requires an ICA decomposition"): pop_iclabel(eeg, "default") + def test_async_result_has_replayable_history_command(self): + async def run(): + updated = dict(_eeg(), etc={"ic_classification": {"ICLabel": {"version": "default"}}}) + with mock.patch.object(pop_iclabel_module, "iclabel_async", new=mock.AsyncMock(return_value=updated)): + return await pop_iclabel_async(_eeg(), "default", gui=False, return_com=True) + + out, command = asyncio.run(run()) + + self.assertEqual(out["etc"]["ic_classification"]["ICLabel"]["version"], "default") + self.assertEqual(command, "EEG = await pop_iclabel_async(EEG, 'default');") + + def test_async_list_recursion_preserves_input_order(self): + first = _eeg() + second = _eeg() + updated_first = dict(first, setname="first") + updated_second = dict(second, setname="second") + + async def run(): + with mock.patch.object( + pop_iclabel_module, + "iclabel_async", + new=mock.AsyncMock(side_effect=[updated_first, updated_second]), + ): + return await pop_iclabel_async([first, second], "default", gui=False, return_com=True) + + output, command = asyncio.run(run()) + + self.assertEqual([item["setname"] for item in output], ["first", "second"]) + self.assertEqual(command, "EEG = await pop_iclabel_async(EEG, 'default');") + + def test_sync_entry_point_fails_fast_under_emscripten(self): + with mock.patch.object(pop_iclabel_module, "_IS_EMSCRIPTEN", True): + with self.assertRaisesRegex(RuntimeError, r'or await pop_iclabel_async\(\.\.\.\)'): + pop_iclabel(None) + + def test_emscripten_gui_dispatch_awaits_and_commits_to_original_slot(self): + session = EEGPrepSession() + session.store_current(_eeg(), new=True) + updated = dict(session.EEG, setname="classified") + command = "EEG = await pop_iclabel_async(EEG, 'default');" + + async def classify(selection, *, renderer=None, return_com=False): + self.assertIs(selection, session.EEG) + self.assertIsNone(renderer) + return (updated, command) if return_com else updated + + dispatcher = MenuActionDispatcher(session) + with ( + mock.patch.object(menu_actions_module, "_IS_EMSCRIPTEN", True), + mock.patch.object(pop_iclabel_module, "pop_iclabel_async", side_effect=classify), + ): + asyncio.run(dispatcher.dispatch_gui("pop_iclabel")) + + self.assertEqual(session.EEG["setname"], "classified") + self.assertEqual(session.CURRENTSET, [1]) + self.assertEqual(session.ALLCOM, [command]) + + def test_emscripten_gui_dispatch_discards_stale_result(self): + session = EEGPrepSession() + session.store_current(_eeg("first"), new=True) + original = session.ALLEEG[0] + + dispatcher = MenuActionDispatcher(session) + + async def scenario(): + started = asyncio.Event() + release = asyncio.Event() + + async def classify(_selection, *, renderer=None, return_com=False): + started.set() + await release.wait() + return (dict(original, setname="classified"), "EEG = await pop_iclabel_async(EEG, 'default');") + + with ( + mock.patch.object(menu_actions_module, "_IS_EMSCRIPTEN", True), + mock.patch.object(pop_iclabel_module, "pop_iclabel_async", side_effect=classify), + ): + task = asyncio.create_task(dispatcher.dispatch_gui("pop_iclabel")) + await started.wait() + session.mark_current_saved() + session.store_current(_eeg("second"), new=True) + release.set() + with self.assertRaisesRegex(RuntimeError, "session changed"): + await task + + asyncio.run(scenario()) + + self.assertIs(session.ALLEEG[0], original) + self.assertEqual(session.ALLEEG[0]["saved"], "yes") + self.assertEqual(session.ALLEEG[1]["setname"], "second") + self.assertEqual(session.ALLCOM, []) + if __name__ == "__main__": unittest.main() diff --git a/tests/test_iclabel.py b/tests/test_iclabel.py index e538291a..1561e269 100644 --- a/tests/test_iclabel.py +++ b/tests/test_iclabel.py @@ -1,21 +1,113 @@ import os import unittest +from unittest import mock + import numpy as np from eegprep import ICL_feature_extractor, iclabel, pop_loadset from eegprep.utils.testing import has_optional_dependency +from eegprep.plugins.ICLabel.eeg_icflag import eeg_icflag +from eegprep.plugins.ICLabel.pop_icflag import DEFAULT_ICFLAG_THRESHOLDS +import eegprep.plugins.ICLabel.iclabel as iclabel_module +import eegprep.plugins.ICLabel.iclabel_net_onnx as iclabel_onnx_module + # where the test resources local_url = os.path.join(os.path.dirname(__file__), '../sample_data/') +def _async_eeg(): + return { + 'data': np.zeros((2, 20), dtype=np.float32), + 'nbchan': 2, + 'pnts': 20, + 'trials': 1, + 'srate': 100, + 'icaweights': np.eye(2), + 'icasphere': np.eye(2), + 'icawinv': np.eye(2), + 'icachansind': np.arange(2), + 'etc': {}, + } + + +class TestICLabelAsync(unittest.IsolatedAsyncioTestCase): + async def test_native_async_entry_point_uses_shared_postprocessing(self): + eeg = _async_eeg() + network_output = np.arange(28, dtype=np.float32).reshape(4, 7, 1, 1) + features = tuple(np.zeros((4, 1), dtype=np.float32) for _ in range(3)) + + with ( + mock.patch.object(iclabel_module, '_prepare_features', return_value=features) as prepare, + mock.patch( + 'eegprep.plugins.ICLabel.iclabel_net_onnx.run_iclabel_net_async', + new=mock.AsyncMock(return_value=network_output), + ) as run, + ): + output = await iclabel_module.iclabel_async(eeg) + + prepare.assert_called_once() + run.assert_awaited_once_with(*features) + classification = output['etc']['ic_classification']['ICLabel']['classifications'] + self.assertEqual(classification.shape, (1, 7)) + self.assertEqual(output['etc']['ic_classification']['ICLabel']['version'], 'default') + + def test_sync_entry_point_fails_fast_under_emscripten(self): + with mock.patch.object(iclabel_module, '_IS_EMSCRIPTEN', True): + with self.assertRaisesRegex(RuntimeError, r'use await iclabel_async\(\.\.\.\)'): + iclabel_module.iclabel(None) + + def test_sync_onnx_backend_fails_fast_under_emscripten(self): + with mock.patch.object(iclabel_onnx_module, '_IS_EMSCRIPTEN', True): + with self.assertRaisesRegex(RuntimeError, r'use await run_iclabel_net_async\(\.\.\.\)'): + iclabel_onnx_module.run_iclabel_net(None, None, None) + + +def _float32_classifications(eeg): + image, psdmed, autocorr = _reference_network_inputs(eeg) + import onnxruntime as ort + + float32_path = os.path.join( + os.path.dirname(__file__), '..', 'tools', 'iclabel', 'artifacts', 'iclabel_float32.onnx' + ) + (output,) = ort.InferenceSession(float32_path, providers=['CPUExecutionProvider']).run( + ['output'], {'image': image, 'psdmed': psdmed, 'autocorr': autocorr} + ) + return _reference_postprocess(output) + + +def _reference_network_inputs(eeg): + """Build ICLabel inputs independently from the production helpers.""" + features = ICL_feature_extractor(eeg, True) + features[0] = np.single( + np.concatenate([features[0], -features[0], features[0][:, ::-1, :, :], -features[0][:, ::-1, :, :]], axis=3) + ) + features[1] = np.single(np.tile(features[1], (1, 1, 1, 4))) + features[2] = np.single(np.tile(features[2], (1, 1, 1, 4))) + return tuple(np.transpose(feature, (3, 2, 0, 1)) for feature in features) + + +def _reference_postprocess(output): + output = output.T + output = np.reshape(output, (-1, 4), order='F') + output = np.mean(output, axis=1) + output = np.reshape(output, (7, -1), order='F') + return output.T + + +def test_postprocess_matches_independent_reference(): + output = np.arange(56, dtype=np.float32).reshape(4, 7, 2, 1) + + np.testing.assert_array_equal(iclabel_module._postprocess_network_output(output), _reference_postprocess(output)) + + @unittest.skipIf(os.getenv('EEGPREP_SKIP_MATLAB') == '1', "MATLAB not available") class TestICLabelEngines(unittest.TestCase): def setUp(self): self.EEG = pop_loadset(os.path.join(local_url, 'eeglab_data_with_ica_tmp.set')) def test_basic(self): - if not has_optional_dependency('torch'): - self.skipTest("PyTorch is not installed; install eegprep[torch] to run ICLabel parity") + if not has_optional_dependency('onnxruntime'): + self.skipTest("onnxruntime is not installed; install eegprep[iclabel] to run ICLabel parity") features_python = ICL_feature_extractor(self.EEG, True) print(f"\n{'=' * 60}") @@ -29,10 +121,11 @@ def test_basic(self): print(f"Python autocorr max: {np.max(features_python[2]):.6f}, min: {np.min(features_python[2]):.6f}") print(f"{'=' * 60}\n") - EEG_python = iclabel(self.EEG, algorithm='default', engine=None) + # Keep probability-level MATLAB parity against the preserved float32 reference; + # the package default int8 artifact is covered by the semantic gate tests. EEG_matlab = iclabel(self.EEG, algorithm='default', engine='matlab') - res1 = EEG_python['etc']['ic_classification']['ICLabel']['classifications'].flatten() + res1 = _float32_classifications(self.EEG).flatten() res2 = EEG_matlab['etc']['ic_classification']['ICLabel']['classifications'].flatten() # Diagnostic output @@ -67,5 +160,74 @@ def test_basic(self): self.assertTrue(np.allclose(res1, res2, rtol=1e-4, atol=1e-5), 'ICLabel results differ beyond tolerance') +class TestICLabelOnnxExport(unittest.TestCase): + """Cross-check the preserved float32 ONNX reference against its torch source.""" + + def setUp(self): + self.EEG = pop_loadset(os.path.join(local_url, 'eeglab_data_with_ica_tmp.set')) + + def test_onnx_matches_torch_on_sample_data(self): + if not has_optional_dependency('torch'): + self.skipTest("PyTorch is not installed; install eegprep[torch] to run the ONNX-vs-torch cross-check") + if not has_optional_dependency('onnxruntime'): + self.skipTest("onnxruntime is not installed; install eegprep[iclabel] to run ICLabel classification") + + import torch + + from eegprep.plugins.ICLabel.iclabel_net import ICLabelNet + + mat_path = os.path.join(os.path.dirname(__file__), '..', 'src', 'eegprep', 'plugins', 'ICLabel', 'netICL.mat') + model = ICLabelNet(mat_path) + model.eval() + + image, psdmed, autocorr = _reference_network_inputs(self.EEG) + + with torch.no_grad(): + torch_out = model(torch.from_numpy(image), torch.from_numpy(psdmed), torch.from_numpy(autocorr)).numpy() + + torch_final = _reference_postprocess(torch_out) + + onnx_final = _float32_classifications(self.EEG) + + diff = np.abs(torch_final - onnx_final) + print(f"\nONNX vs torch max abs diff: {diff.max():.2e}, mean abs diff: {diff.mean():.2e}") + + # Measured on this sample_data dataset while building the export + # (tools/iclabel/export_iclabel_onnx.py): max abs diff 1.43e-06, max + # rel diff 2.59e-05, from ordinary float32 op-ordering differences + # between eager torch execution and onnxruntime's fused CPU kernels. + # That is the same order of magnitude as the Python-vs-MATLAB gap + # tolerated above (~3.4e-06 abs / ~4.3e-05 rel), so this test reuses + # that already-established tolerance instead of introducing a new one. + self.assertTrue( + np.allclose(torch_final, onnx_final, rtol=1e-4, atol=1e-5), + 'ONNX and torch ICLabel outputs differ beyond tolerance', + ) + + +@unittest.skipUnless(has_optional_dependency('onnxruntime'), "install eegprep[iclabel] to run ICLabel runtime coverage") +class TestICLabelRuntime(unittest.TestCase): + def test_default_entry_point_uses_packaged_onnx_runtime(self): + eeg = pop_loadset(os.path.join(local_url, 'eeglab_data_with_ica_tmp.set')) + + output = iclabel(eeg) + classifications = output['etc']['ic_classification']['ICLabel']['classifications'] + reference = _float32_classifications(eeg) + + self.assertEqual(classifications.shape[1], 7) + self.assertTrue(np.isfinite(classifications).all()) + self.assertEqual(output['etc']['ic_classification']['ICLabel']['version'], 'default') + difference = np.abs(classifications - reference) + self.assertLess(float(difference.max()), 1e-2) + self.assertLess(float(difference.mean()), 1e-3) + np.testing.assert_array_equal(classifications.argmax(axis=1), reference.argmax(axis=1)) + + def rejection_flags(probabilities): + labelled = {'etc': {'ic_classification': {'ICLabel': {'classifications': probabilities}}}} + return eeg_icflag(labelled, DEFAULT_ICFLAG_THRESHOLDS)['reject']['gcompreject'] + + np.testing.assert_array_equal(rejection_flags(classifications), rejection_flags(reference)) + + if __name__ == '__main__': unittest.main() diff --git a/tests/test_iclabel_quantization.py b/tests/test_iclabel_quantization.py new file mode 100644 index 00000000..9828c82b --- /dev/null +++ b/tests/test_iclabel_quantization.py @@ -0,0 +1,326 @@ +from pathlib import Path + +import numpy as np +import pytest + +from eegprep.plugins.ICLabel.pop_icflag import DEFAULT_ICFLAG_THRESHOLDS +from tools.iclabel.quantize_iclabel_onnx import ( + DEFAULT_CALIBRATED_ARTIFACT, + DEFAULT_FLOAT32_ARTIFACT, + DEFAULT_FROZEN_MANIFEST, + DEFAULT_EVALUATION_FEATURES, + DEFAULT_WEIGHT_ONLY_ARTIFACT, + MAX_PROBABILITY_ABS_DIFF, + MIN_KEEP_REJECT_AGREEMENT, + MIN_TOP1_AGREEMENT, + compare_predictions, + evaluate_artifacts, + load_frozen_manifest, + load_verified_feature_archive, + network_inputs_from_features, + parity_gate_passes, + predict_features, + quantize_calibrated, + quantize_weight_only, + select_default_artifact, +) + + +def _reference_network_inputs(features): + topo, psdmed, autocorr = features + topo = np.single(np.concatenate([topo, -topo, topo[:, ::-1, :, :], -topo[:, ::-1, :, :]], axis=3)) + psdmed = np.single(np.tile(psdmed, (1, 1, 1, 4))) + autocorr = np.single(np.tile(autocorr, (1, 1, 1, 4))) + return { + "image": np.transpose(topo, (3, 2, 0, 1)), + "psdmed": np.transpose(psdmed, (3, 2, 0, 1)), + "autocorr": np.transpose(autocorr, (3, 2, 0, 1)), + } + + +def test_frozen_manifest_is_subject_balanced_and_calibration_disjoint(): + manifest = load_frozen_manifest(DEFAULT_FROZEN_MANIFEST) + + assert manifest["status"] == "frozen" + evaluation = manifest["evaluation"] + calibration = manifest["calibration"] + evaluation_subjects = [item["subject"] for item in evaluation["recordings"]] + calibration_subjects = [item["subject"] for item in calibration["recordings"]] + + assert evaluation_subjects == sorted(set(evaluation_subjects)) + assert calibration_subjects == sorted(set(calibration_subjects)) + assert set(evaluation_subjects).isdisjoint(calibration_subjects) + assert all(item["component_indices"] == list(range(31)) for item in evaluation["recordings"]) + assert all(item["component_indices"] == list(range(31)) for item in calibration["recordings"]) + assert {item["source_path"] for item in evaluation["recordings"]}.isdisjoint( + item["source_path"] for item in calibration["recordings"] + ) + assert manifest["selection_policy"]["confidence_based_filtering"] is False + assert manifest["selection_policy"]["subject_disjoint"] is True + assert manifest["selection_policy"]["recordings_per_subject"] == 1 + + +def test_compare_predictions_reports_overall_and_per_class_agreement(): + teacher = np.array( + [ + [0.02, 0.94, 0.01, 0.01, 0.01, 0.005, 0.005], + [0.02, 0.01, 0.94, 0.01, 0.01, 0.005, 0.005], + [0.80, 0.03, 0.03, 0.03, 0.03, 0.04, 0.04], + [0.05, 0.05, 0.05, 0.05, 0.05, 0.05, 0.70], + ] + ) + candidate = teacher.copy() + candidate[3] = [0.55, 0.05, 0.05, 0.05, 0.05, 0.05, 0.40] + + metrics = compare_predictions(teacher, candidate) + + assert metrics["top1_agreement"] == pytest.approx(0.75) + assert metrics["keep_reject_agreement"] == pytest.approx(1.0) + assert metrics["max_probability_abs_diff"] == pytest.approx(0.50) + assert metrics["mean_probability_abs_diff"] == pytest.approx(0.8 / 28) + assert metrics["per_class_agreement"]["Brain"]["count"] == 1 + assert metrics["per_class_agreement"]["Brain"]["agreement"] == pytest.approx(1.0) + assert metrics["per_class_agreement"]["Other"]["count"] == 1 + assert metrics["per_class_agreement"]["Other"]["agreement"] == pytest.approx(0.0) + + +def test_compare_predictions_exposes_rejection_threshold_boundaries(): + teacher = np.array( + [ + [0.9001, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0999], + [0.0, 0.9999, 0.0, 0.0, 0.0, 0.0, 0.0001], + [0.0, 0.0, 0.9, 0.0, 0.0, 0.0, 0.1], + ] + ) + candidate = np.array( + [ + [0.8999, 0.0, 0.0, 0.0, 0.0, 0.0, 0.1001], + [0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 0.0], + [0.0, 0.0, 0.9001, 0.0, 0.0, 0.0, 0.0999], + ] + ) + thresholds = np.array( + [ + [0.9, 1.0], + [0.9, 1.0], + [0.9, 1.0], + [np.nan, np.nan], + [np.nan, np.nan], + [np.nan, np.nan], + [np.nan, np.nan], + ] + ) + + metrics = compare_predictions(teacher, candidate, thresholds) + + assert metrics["top1_agreement"] == pytest.approx(1.0) + assert metrics["keep_reject_agreement"] == pytest.approx(0.0) + + +@pytest.mark.parametrize( + ("probability_drift", "expected"), + ((MAX_PROBABILITY_ABS_DIFF, True), (np.nextafter(MAX_PROBABILITY_ABS_DIFF, np.inf), False)), +) +def test_probability_gate_boundary_is_inclusive(probability_drift, expected): + metrics = { + "top1_agreement": MIN_TOP1_AGREEMENT, + "keep_reject_agreement": MIN_KEEP_REJECT_AGREEMENT, + "max_probability_abs_diff": probability_drift, + } + + assert parity_gate_passes(metrics) is expected + + +def test_frozen_feature_archive_hash_is_verified(tmp_path): + manifest = load_frozen_manifest(DEFAULT_FROZEN_MANIFEST) + with np.load(DEFAULT_EVALUATION_FEATURES, allow_pickle=False) as archive: + features = {name: np.asarray(archive[name]).copy() for name in ("topo", "psd", "autocorr")} + features["topo"][0, 0, 0, 0] += 1.0 + mutated = tmp_path / "evaluation_features.npz" + np.savez(mutated, **features) + + with pytest.raises(ValueError, match="SHA-256 mismatch"): + load_verified_feature_archive(mutated, manifest, "evaluation") + + +def test_quantization_inputs_match_independent_reference_transform(): + manifest = load_frozen_manifest(DEFAULT_FROZEN_MANIFEST) + features = load_verified_feature_archive(DEFAULT_EVALUATION_FEATURES, manifest, "evaluation") + actual = network_inputs_from_features(features) + expected = _reference_network_inputs(features) + + for name in expected: + np.testing.assert_array_equal(actual[name], expected[name]) + + +def test_artifact_selection_falls_back_to_float32_until_an_int8_gate_passes(): + report = { + "candidates": { + "weight_only": { + "artifact": "iclabel_int8_weight_only.onnx", + "top1_agreement": 0.99, + "keep_reject_agreement": 0.98, + "max_probability_abs_diff": 0.01, + }, + "calibrated": { + "artifact": "iclabel_int8_calibrated.onnx", + "top1_agreement": 0.99, + "keep_reject_agreement": 0.98, + "max_probability_abs_diff": 0.01, + }, + } + } + + assert select_default_artifact(report) == "iclabel.onnx" + + report["candidates"]["weight_only"]["keep_reject_agreement"] = 0.99 + assert select_default_artifact(report) == "iclabel_int8_weight_only.onnx" + + report["candidates"]["weight_only"]["keep_reject_agreement"] = 0.98 + report["candidates"]["calibrated"]["keep_reject_agreement"] = 0.99 + assert select_default_artifact(report) == "iclabel_int8_calibrated.onnx" + + +def test_thresholds_remain_the_existing_open_interval_defaults(): + expected = np.array( + [ + [np.nan, np.nan], + [0.9, 1.0], + [0.9, 1.0], + [np.nan, np.nan], + [np.nan, np.nan], + [np.nan, np.nan], + [np.nan, np.nan], + ] + ) + np.testing.assert_equal(DEFAULT_ICFLAG_THRESHOLDS, expected) + + +def test_packaged_artifact_matches_the_gate_selected_candidate(): + package_artifact = Path(__file__).parents[1] / "src" / "eegprep" / "plugins" / "ICLabel" / "iclabel.onnx" + assert package_artifact.read_bytes() == DEFAULT_WEIGHT_ONLY_ARTIFACT.read_bytes() + + +def test_packaged_default_artifact_passes_the_semantic_gate(): + pytest.importorskip("onnxruntime") + + package_artifact = Path(__file__).parents[1] / "src" / "eegprep" / "plugins" / "ICLabel" / "iclabel.onnx" + manifest = load_frozen_manifest(DEFAULT_FROZEN_MANIFEST) + features = load_verified_feature_archive(DEFAULT_EVALUATION_FEATURES, manifest, "evaluation") + teacher = predict_features(DEFAULT_FLOAT32_ARTIFACT, features) + candidate = predict_features(package_artifact, features) + + assert parity_gate_passes(compare_predictions(teacher, candidate)) + + +def test_committed_candidates_pass_evaluation_on_the_frozen_archive(): + pytest.importorskip("onnxruntime") + + report = evaluate_artifacts( + DEFAULT_FLOAT32_ARTIFACT, + DEFAULT_EVALUATION_FEATURES, + {"weight_only": DEFAULT_WEIGHT_ONLY_ARTIFACT, "calibrated": DEFAULT_CALIBRATED_ARTIFACT}, + ) + + assert report["float32_reference"]["sample_count"] == 217 + assert report["float32_reference"]["top1_agreement"] == pytest.approx(1.0) + assert report["float32_reference"]["keep_reject_agreement"] == pytest.approx(1.0) + assert report["feature_archives"]["evaluation"] == { + "path": "tools/iclabel/evaluation_features.npz", + "sha256": "dbb2c0cc063d29ab2e8e59fafb0f474df780a458ac17be9d4732abc0233c8b44", + "component_count": 217, + } + assert report["default_artifact"] == "iclabel_int8_weight_only.onnx" + assert report["gate"]["maximum_probability_abs_diff"] == MAX_PROBABILITY_ABS_DIFF + + expected_counts = { + "Brain": 18, + "Muscle": 1, + "Eye": 5, + "Heart": 0, + "Line Noise": 0, + "Channel Noise": 0, + "Other": 193, + } + # ONNX Runtime can make platform-specific boundary choices for calibrated + # int8 activations; the fixed promotion thresholds are the portable gate. + for candidate_name in ("weight_only", "calibrated"): + candidate = report["candidates"][candidate_name] + assert candidate["gate_pass"] is (candidate_name == "weight_only") + assert candidate["top1_agreement"] >= MIN_TOP1_AGREEMENT + assert candidate["keep_reject_agreement"] >= MIN_KEEP_REJECT_AGREEMENT + assert candidate["teacher_class_distribution"] == expected_counts + assert set(candidate["per_class_agreement"]) == set(expected_counts) + assert report["candidates"]["calibrated"]["max_probability_abs_diff"] > MAX_PROBABILITY_ABS_DIFF + + +def test_weight_only_matches_full_frozen_probabilities_and_threshold_boundaries(): + pytest.importorskip("onnxruntime") + + manifest = load_frozen_manifest(DEFAULT_FROZEN_MANIFEST) + features = load_verified_feature_archive(DEFAULT_EVALUATION_FEATURES, manifest, "evaluation") + teacher = predict_features(DEFAULT_FLOAT32_ARTIFACT, features) + candidate = predict_features(DEFAULT_WEIGHT_ONLY_ARTIFACT, features) + metrics = compare_predictions(teacher, candidate, DEFAULT_ICFLAG_THRESHOLDS) + + assert teacher.shape == (217, 7) + assert metrics["max_probability_abs_diff"] <= MAX_PROBABILITY_ABS_DIFF + assert metrics["mean_probability_abs_diff"] <= 0.001 + assert metrics["keep_reject_agreement"] == pytest.approx(1.0) + + +def test_weight_only_quantization_preserves_float_io_and_softmax(tmp_path): + onnx = pytest.importorskip("onnx") + ort = pytest.importorskip("onnxruntime") + + output_path = tmp_path / "weight-only.onnx" + quantize_weight_only(DEFAULT_FLOAT32_ARTIFACT, output_path) + model = onnx.load(output_path) + + input_types = {value.name: value.type.tensor_type.elem_type for value in model.graph.input} + output_types = {value.name: value.type.tensor_type.elem_type for value in model.graph.output} + assert set(input_types.values()) == {onnx.TensorProto.FLOAT} + assert set(output_types.values()) == {onnx.TensorProto.FLOAT} + assert any(node.op_type == "DequantizeLinear" for node in model.graph.node) + assert any(node.op_type == "Softmax" for node in model.graph.node) + + session = ort.InferenceSession(str(output_path), providers=["CPUExecutionProvider"]) + inputs = { + "image": np.zeros((2, 1, 32, 32), dtype=np.float32), + "psdmed": np.zeros((2, 1, 1, 100), dtype=np.float32), + "autocorr": np.zeros((2, 1, 1, 100), dtype=np.float32), + } + (output,) = session.run(["output"], inputs) + assert output.shape == (2, 7, 1, 1) + assert np.isfinite(output).all() + np.testing.assert_allclose(output.sum(axis=1), 1.0, rtol=1e-5, atol=1e-6) + + +def test_calibrated_quantization_consumes_float_feature_inputs(tmp_path): + onnx = pytest.importorskip("onnx") + ort = pytest.importorskip("onnxruntime") + + rng = np.random.default_rng(379) + calibration_inputs = { + "image": rng.standard_normal((4, 1, 32, 32)).astype(np.float32), + "psdmed": rng.standard_normal((4, 1, 1, 100)).astype(np.float32), + "autocorr": rng.standard_normal((4, 1, 1, 100)).astype(np.float32), + } + output_path = tmp_path / "calibrated.onnx" + + quantize_calibrated(DEFAULT_FLOAT32_ARTIFACT, output_path, calibration_inputs) + + assert output_path.exists() + assert output_path.stat().st_size < DEFAULT_FLOAT32_ARTIFACT.stat().st_size + model = onnx.load(output_path) + input_types = {value.name: value.type.tensor_type.elem_type for value in model.graph.input} + output_types = {value.name: value.type.tensor_type.elem_type for value in model.graph.output} + assert set(input_types.values()) == {onnx.TensorProto.FLOAT} + assert set(output_types.values()) == {onnx.TensorProto.FLOAT} + assert any(node.op_type == "Softmax" for node in model.graph.node) + (output,) = ort.InferenceSession(str(output_path), providers=["CPUExecutionProvider"]).run( + ["output"], calibration_inputs + ) + assert output.shape == (4, 7, 1, 1) + assert np.isfinite(output).all() + np.testing.assert_allclose(output.sum(axis=1), 1.0, rtol=1e-5, atol=1e-6) diff --git a/tests/test_pipeline.py b/tests/test_pipeline.py index 55d10f70..2824e376 100644 --- a/tests/test_pipeline.py +++ b/tests/test_pipeline.py @@ -139,8 +139,8 @@ def test_iclabel(self): """Test iclabel component classification.""" if not self.has_matlab_picard: self.skipTest("MATLAB EEGLAB Picard plugin is not installed") - if not has_optional_dependency('torch'): - self.skipTest("PyTorch is not installed; install eegprep[torch] to run ICLabel parity") + if not has_optional_dependency('onnxruntime'): + self.skipTest("onnxruntime is not installed; install eegprep[iclabel] to run ICLabel parity") # Prepare data: channel cleaning + burst cleaning + ICA EEG_py_ch, *_ = clean_artifacts(deepcopy(self.EEG), BurstCriterion='off', ChannelCriterion=0.8) @@ -174,8 +174,8 @@ def test_z_full_pipeline(self): """Test the complete pipeline end-to-end.""" if not self.has_matlab_picard: self.skipTest("MATLAB EEGLAB Picard plugin is not installed") - if not has_optional_dependency('torch'): - self.skipTest("PyTorch is not installed; install eegprep[torch] to run full pipeline parity") + if not has_optional_dependency('onnxruntime'): + self.skipTest("onnxruntime is not installed; install eegprep[iclabel] to run full pipeline parity") print("\n" + "=" * 80) print("Full Pipeline Test: clean_artifacts -> eeg_picard -> iclabel") diff --git a/tests/test_public_api_examples.py b/tests/test_public_api_examples.py index 2880309d..92deaf19 100644 --- a/tests/test_public_api_examples.py +++ b/tests/test_public_api_examples.py @@ -32,8 +32,8 @@ def test_public_api_and_plugins_example_runs() -> None: # One illustration script per user guide section. These run for real against # sample_data/, so a failure here means the documented workflow is broken. -# plot_reject_artifacts.py calls ICLabel, which needs the torch extra and raises -# ImportError without it rather than skipping; install eegprep[torch] to run. +# plot_reject_artifacts.py calls ICLabel, which needs the iclabel extra and raises +# ImportError without it rather than skipping; install eegprep[iclabel] to run. USER_GUIDE_EXAMPLES = ( "plot_quickstart_tour.py", "plot_data_structures.py", @@ -76,7 +76,7 @@ def test_package_resources_cover_public_workflows() -> None: assert package.joinpath("resources/skills/eegprep-cli.md").is_file() assert package.joinpath("resources/headplot/colin27headmesh.mat").is_file() assert package.joinpath("resources/montages/standard-10-5-342ch.locs").is_file() - assert package.joinpath("plugins/ICLabel/netICL.mat").is_file() + assert package.joinpath("plugins/ICLabel/iclabel.onnx").is_file() def test_browser_help_and_docs_do_not_describe_eegbrowser_as_excluded() -> None: @@ -147,7 +147,7 @@ def test_setuptools_package_data_covers_runtime_resources() -> None: "resources/headplot/mheadnew.transform", "resources/headplot/mheadnew.xyz", "resources/montages/standard-10-5-342ch.locs", - "plugins/ICLabel/netICL.mat", + "plugins/ICLabel/iclabel.onnx", } <= packaged assert "eeglab/**" in excluded assert not any(path.startswith("eeglab/") for path in packaged) diff --git a/tests/test_pyodide_harness.py b/tests/test_pyodide_harness.py new file mode 100644 index 00000000..eed8e063 --- /dev/null +++ b/tests/test_pyodide_harness.py @@ -0,0 +1,187 @@ +import shutil +import subprocess +import textwrap +from pathlib import Path + +import pytest + +from tools.check_pyodide_base_resolution import KNOWN_GAPS +from tools.pyodide.benchmark import MATMUL_CASES, MAX_ICA_ITERATIONS, benchmark_record +from tools.pyodide.compare_benchmarks import compare_reports +from tools.pyodide.compare_iclabel import compare_reports as compare_iclabel_reports +from tools.pyodide.prepare_docopt_wheel import verify_sha256 + + +def test_docopt_hash_check_rejects_modified_sdist(tmp_path): + archive = tmp_path / "docopt-0.6.2.tar.gz" + archive.write_bytes(b"locked sdist payload") + + with pytest.raises(ValueError, match="SHA-256 mismatch"): + verify_sha256(archive, "0" * 64) + + +def test_docopt_is_no_longer_a_static_known_gap(): + assert "docopt" not in KNOWN_GAPS + + +def test_benchmark_record_marks_max_iterations_as_not_converged(): + result = benchmark_record( + algorithm="runica", + platform="native", + shape=(64, 15_000), + seed=375, + seconds=1.25, + iterations=MAX_ICA_ITERATIONS, + converged=False, + backend="numpy", + ) + + assert result["iterations"] == MAX_ICA_ITERATIONS + assert result["converged"] is False + assert result["error"] is None + + +def test_benchmark_uses_runica_inner_loop_product_shapes(): + assert [(case["left_shape"], case["right_shape"]) for case in MATMUL_CASES] == [ + ((64, 64), (64, 49)), + ((64, 64), (64, 64)), + ] + + +def test_compare_reports_applies_convergence_and_blas_gates(): + def report(platform, runica_seconds, picard_seconds, picard_converged): + matmul = [] + for case in ("activation", "weight_update"): + for dtype, operation, blas_seconds in ( + ("float64", "dgemm", 0.5), + ("float32", "sgemm", 0.9 if case == "activation" else 0.5), + ): + common = { + "case": case, + "dtype": dtype, + "platform": "pyodide", + "seed": 375, + "left_shape": [64, 64], + "right_shape": [64, 49] if case == "activation" else [64, 64], + "median_seconds": 1.0, + } + matmul.append({**common, "operation": "numpy_matmul"}) + matmul.append({**common, "operation": operation, "median_seconds": blas_seconds}) + return { + "schema_version": 1, + "platform": platform, + "seed": 375, + "shape": [64, 15_000], + "runica_block": 49, + "thread_count": 1, + "ica": { + "runica": { + "median_seconds": runica_seconds, + "median_iterations": 100, + "all_converged": True, + }, + "picard": { + "median_seconds": picard_seconds, + "median_iterations": 50, + "all_converged": picard_converged, + }, + }, + "matmul": matmul, + } + + comparison = compare_reports(report("native", 10.0, 5.0, True), report("pyodide", 20.0, 4.0, True)) + + assert comparison["ica"]["runica"]["pyodide_speed_ratio"] == 2.0 + assert comparison["ica"]["runica"]["native_median_iterations"] == 100 + assert comparison["ica"]["runica"]["native_all_converged"] is True + assert comparison["matmul"][0]["native_blas_speedup_over_numpy"] == 2.0 + assert comparison["matmul"][0]["blas_speedup_over_numpy"] == 2.0 + assert comparison["decisions"]["picard_browser_default_retained"] is True + assert comparison["decisions"]["phase3_recommended"] is False + + +def test_compare_iclabel_reports_applies_established_numeric_tolerance(): + native = { + "schema_version": 1, + "platform": "native", + "dataset": "eeglab_data_with_ica_tmp.set", + "shape": [2, 7], + "classifications": [[0.1] * 7, [0.9] * 7], + } + pyodide = { + **native, + "platform": "pyodide", + "classifications": [[0.1 + 1e-6] * 7, [0.9 - 1e-6] * 7], + } + + comparison = compare_iclabel_reports(native, pyodide) + + assert comparison["allclose"] is True + assert comparison["max_absolute_difference"] <= 1e-5 + + +def test_compare_iclabel_reports_rejects_shape_mismatch(): + native = { + "schema_version": 1, + "platform": "native", + "dataset": "sample.set", + "shape": [1, 7], + "classifications": [[0.1] * 7], + } + pyodide = { + **native, + "platform": "pyodide", + "shape": [2, 7], + "classifications": [[0.1] * 7, [0.9] * 7], + } + + with pytest.raises(ValueError, match="shapes differ"): + compare_iclabel_reports(native, pyodide) + + +@pytest.mark.skipif(shutil.which("node") is None, reason="Node.js is required for the browser bridge test") +def test_iclabel_web_bridge_retries_failed_session_initialization(): + script = textwrap.dedent( + """ + import { createIcLabelWebBridge } from './tools/pyodide/iclabel_web_bridge.mjs'; + + let attempts = 0; + const ort = { + InferenceSession: { + create: async () => { + attempts += 1; + if (attempts === 1) throw new Error('transient model load failure'); + return { run: async () => ({ output: { data: [1] } }) }; + }, + }, + Tensor: class Tensor { + constructor(type, data, shape) { + this.type = type; + this.data = data; + this.shape = shape; + } + }, + }; + + const bridge = createIcLabelWebBridge(ort, new Uint8Array()); + const input = new Float32Array([0]); + const shape = [1]; + try { + await bridge.run(input, shape, input, shape, input, shape); + throw new Error('first run unexpectedly succeeded'); + } catch (error) { + if (error.message !== 'transient model load failure') throw error; + } + await bridge.run(input, shape, input, shape, input, shape); + if (attempts !== 2) throw new Error(`expected two session attempts, got ${attempts}`); + """ + ) + result = subprocess.run( + ["node", "--input-type=module"], + input=script, + text=True, + capture_output=True, + cwd=Path(__file__).parents[1], + check=False, + ) + assert result.returncode == 0, result.stderr diff --git a/tests/test_runica_matmul.py b/tests/test_runica_matmul.py new file mode 100644 index 00000000..1c2eaa8a --- /dev/null +++ b/tests/test_runica_matmul.py @@ -0,0 +1,91 @@ +import importlib +import subprocess +import sys + +import numpy as np + +from eegprep.functions.sigprocfunc import runica_matmul + + +runica_module = importlib.import_module("eegprep.functions.sigprocfunc.runica") + + +def test_runica_matmul_backends_agree(): + rng = np.random.default_rng(376) + left = rng.standard_normal((7, 11)) + right = rng.standard_normal((11, 5)) + + np.testing.assert_allclose( + runica_matmul._blas_matmul(left, right), + runica_matmul._numpy_matmul(left, right), + rtol=1e-12, + atol=1e-12, + ) + + +def test_runica_matmul_selects_numpy_on_native_platform(): + assert runica_matmul.BACKEND == "numpy.matmul" + assert runica_matmul.runica_matmul is runica_matmul._numpy_matmul + + +def test_runica_matmul_selects_dgemm_on_emscripten(monkeypatch): + original_platform = sys.platform + monkeypatch.setattr(sys, "platform", "emscripten") + try: + reloaded = importlib.reload(runica_matmul) + assert reloaded.BACKEND == "scipy.linalg.blas.dgemm" + assert reloaded.runica_matmul is reloaded._blas_matmul + finally: + monkeypatch.setattr(sys, "platform", original_platform) + importlib.reload(runica_matmul) + + +def test_runica_consumer_reads_the_selected_backend_module(monkeypatch): + original_platform = sys.platform + monkeypatch.setattr(sys, "platform", "emscripten") + try: + selected = importlib.reload(runica_matmul) + consumer = importlib.reload(runica_module) + assert consumer._runica_matmul is selected + assert consumer._runica_matmul.runica_matmul is selected._blas_matmul + finally: + monkeypatch.setattr(sys, "platform", original_platform) + importlib.reload(runica_matmul) + importlib.reload(runica_module) + + +def test_runica_executes_the_selected_backend(monkeypatch): + original_platform = sys.platform + monkeypatch.setattr(sys, "platform", "emscripten") + try: + selected = importlib.reload(runica_matmul) + consumer = importlib.reload(runica_module) + calls = [] + + def spy(left, right): + calls.append((left.shape, right.shape)) + return selected._blas_matmul(left, right) + + monkeypatch.setattr(selected, "runica_matmul", spy) + data = np.random.default_rng(385).standard_normal((3, 100)) + consumer.runica(data, maxsteps=1, verbose=False, rndreset="off") + assert calls + finally: + monkeypatch.setattr(sys, "platform", original_platform) + importlib.reload(runica_matmul) + importlib.reload(runica_module) + + +def test_runica_import_does_not_eagerly_load_mne(): + result = subprocess.run( + [ + sys.executable, + "-c", + ("import sys; from eegprep.functions.sigprocfunc.runica import runica; assert 'mne' not in sys.modules"), + ], + capture_output=True, + text=True, + check=False, + ) + + assert result.returncode == 0, result.stderr diff --git a/tests/test_session_contracts.py b/tests/test_session_contracts.py index 9cc65fc3..39561e45 100644 --- a/tests/test_session_contracts.py +++ b/tests/test_session_contracts.py @@ -6,6 +6,7 @@ import pytest from eegprep.functions.adminfunc.console import EEGPrepConsoleWorkspace +from eegprep.functions.adminfunc.storage import MemmapData from eegprep.functions.guifunc.session import EEGPrepSession, normalize_dataset_indices from tests.fixtures import ( assert_eeg_fields_close, @@ -106,6 +107,63 @@ def test_currentset_empty_single_and_multiple_console_values(): assert session.current_set_value() == [1, 2] +def test_dataset_state_token_changes_when_session_notifies(): + session = EEGPrepSession() + token = session.dataset_state_token() + + session.store_current(create_test_eeg(n_channels=2, n_samples=8), new=True) + + assert not session.dataset_state_unchanged(token) + + +def test_dataset_state_token_ignores_history_only_changes(): + session = EEGPrepSession() + session.store_current(create_test_eeg(n_channels=2, n_samples=8), new=True) + token = session.dataset_state_token() + + session.add_history("plot(EEG);") + + assert session.dataset_state_unchanged(token) + + +def test_dataset_state_token_detects_saved_dataset_mutations(): + session = EEGPrepSession() + session.store_current(create_test_eeg(n_channels=2, n_samples=8), new=True) + token = session.dataset_state_token() + + session.mark_current_saved() + + assert not session.dataset_state_unchanged(token) + + +def test_dataset_state_token_detects_selected_slot_replacement_without_notification(): + session = EEGPrepSession() + session.store_current(create_test_eeg(n_channels=2, n_samples=8), new=True) + token = session.dataset_state_token() + replacement = create_test_eeg(n_channels=2, n_samples=8) + session.ALLEEG[0] = replacement + session.EEG = replacement + + assert not session.dataset_state_unchanged(token) + + +def test_dataset_state_token_detects_memmap_data_mutation(tmp_path): + path = tmp_path / "data.fdt" + backing = np.memmap(path, dtype="float32", mode="w+", shape=(2, 2), order="F") + backing[:] = [[1.0, 2.0], [3.0, 4.0]] + backing.flush() + + eeg = create_test_eeg(n_channels=2, n_samples=2) + eeg["data"] = MemmapData(path, (2, 2), mode="r+") + session = EEGPrepSession() + session.store_current(eeg, new=True) + token = session.dataset_state_token() + + np.asarray(session.EEG["data"])[0, 0] = 9.0 + + assert not session.dataset_state_unchanged(token) + + def test_dataset_index_normalization_accepts_canonical_empty_values(): assert normalize_dataset_indices(0) == [] assert normalize_dataset_indices([0]) == [] diff --git a/tests/test_storage.py b/tests/test_storage.py index e991dc7f..08f43d37 100644 --- a/tests/test_storage.py +++ b/tests/test_storage.py @@ -68,6 +68,26 @@ def test_pop_saveset_twofiles_roundtrips_epoched_data_as_memmap(tmp_path: Path): assert reloaded["data"][1, 2, 1] == -123.0 +def test_memmap_common_mutation_views_advance_revision(tmp_path: Path): + path = tmp_path / "tracked.fdt" + backing = np.memmap(path, dtype="float32", mode="w+", shape=(2, 3), order="F") + backing[:] = np.arange(6, dtype=np.float32).reshape((2, 3)) + backing.flush() + data = MemmapData(path, (2, 3), mode="r+") + + data.flat[0] = 10.0 + assert data.mutation_revision == 1 + assert np.asarray(data).flat[0] == 10.0 + data.ravel(order="F")[1] = 11.0 + assert data.mutation_revision == 2 + data.sort() + assert data.mutation_revision == 3 + np.put(data, [2], [12.0]) + assert data.mutation_revision == 4 + + data.close() + + def test_pop_saveset_default_single_file_keeps_data_inline(tmp_path: Path): data = np.arange(6, dtype=np.float32).reshape((2, 3)) set_file = tmp_path / "inline.set" diff --git a/tests/test_utils_asr.py b/tests/test_utils_asr.py index d2bf5401..18bc62fd 100644 --- a/tests/test_utils_asr.py +++ b/tests/test_utils_asr.py @@ -1,3 +1,4 @@ +import sys import unittest import numpy as np from unittest.mock import patch @@ -535,20 +536,16 @@ def test_memory_management(self): # Create larger data that might trigger splitting large_data = np.random.randn(self.n_channels, 5000) * 0.5 - with patch('psutil.virtual_memory') as mock_vm: - # Mock low available memory to trigger splitting - mock_vm.return_value.free = 50 * 1024**2 # 50 MB - - with self.assertLogs('eegprep.plugins.clean_rawdata.asr_process', level='INFO') as log: - cleaned_data, new_state = asr_process( - large_data, - self.srate, - self.state, - max_mem=10, # Low memory limit - ) + with self.assertLogs('eegprep.plugins.clean_rawdata.asr_process', level='INFO') as log: + cleaned_data, new_state = asr_process( + large_data, + self.srate, + self.state, + max_mem=10, # Low memory limit, forces splitting into blocks + ) - # Check that splitting was logged - self.assertTrue(any('blocks' in msg for msg in log.output)) + # Check that splitting was logged + self.assertTrue(any('blocks' in msg for msg in log.output)) # Check output self.assertEqual(cleaned_data.shape, large_data.shape) @@ -556,14 +553,23 @@ def test_memory_management(self): def test_memory_error_handling(self): """Test error handling when memory is insufficient.""" - with patch('psutil.virtual_memory') as mock_vm: - # Mock extremely low memory - mock_vm.return_value.free = 1024 # 1 KB + with self.assertRaises(RuntimeError) as cm: + asr_process(self.test_data, self.srate, self.state, max_mem=0.001) - with self.assertRaises(RuntimeError) as cm: - asr_process(self.test_data, self.srate, self.state, max_mem=0.001) + self.assertIn('Not enough memory', str(cm.exception)) + + def test_default_max_mem_works_without_psutil(self): + """max_mem=None must resolve to the fixed default without importing psutil. - self.assertIn('Not enough memory', str(cm.exception)) + psutil has no WebAssembly build, so this path must not import it (e.g. under + Pyodide). Block the real import (rather than mocking asr_process's logic) to prove + no code path reaches for psutil. + """ + with patch.dict(sys.modules, {'psutil': None}): + cleaned_data, new_state = asr_process(self.test_data, self.srate, self.state, max_mem=None) + + self.assertEqual(cleaned_data.shape, self.test_data.shape) + self.assertTrue(np.all(np.isfinite(cleaned_data))) def test_rank_deficient_covariance_produces_sane_output(self): """Process genuinely rank-deficient data (singular covariance). diff --git a/tools/check_pyodide_base_resolution.py b/tools/check_pyodide_base_resolution.py new file mode 100644 index 00000000..26217bba --- /dev/null +++ b/tools/check_pyodide_base_resolution.py @@ -0,0 +1,184 @@ +"""Assert that eegprep's base install (no extras) resolves cleanly under Pyodide. + +Phase gate for issue #374 / epic #324: ``micropip.install("eegprep")`` must not need a +WebAssembly build for anything left in ``[project].dependencies``. Every package reachable +from eegprep's base dependencies in ``uv.lock`` must either ship prebuilt in the target +Pyodide distribution, or publish a pure-Python wheel (abi ``none``, platform ``any``) that +micropip can pull straight from PyPI. This only checks resolvability from lockfile and +package-index metadata; it does not run inside Pyodide (that is phase 2's live +``micropip.install`` harness). +""" + +from __future__ import annotations + +import json +import sys +import urllib.error +import urllib.request +from dataclasses import dataclass +from pathlib import Path + +import tomllib + +from packaging.markers import Marker +from packaging.utils import canonicalize_name + +REPO_ROOT = Path(__file__).resolve().parents[1] +UV_LOCK_PATH = REPO_ROOT / "uv.lock" + +# Pinned Pyodide release used to look up which packages ship prebuilt in the distribution. +# Bump when phase 2 wires up a live micropip.install harness against a newer release. +PYODIDE_VERSION = "0.29.5" +PYODIDE_LOCK_URL = f"https://cdn.jsdelivr.net/pyodide/v{PYODIDE_VERSION}/full/pyodide-lock.json" +PYPI_JSON_URL = "https://pypi.org/pypi/{name}/{version}/json" +HTTP_TIMEOUT_S = 30 + +# The live Phase 2 harness builds this exact sdist into a local universal wheel before asking +# micropip to resolve eegprep. Keep the package out of ``KNOWN_GAPS`` so the static check cannot +# accidentally turn a known local transport into a silent exception. +KNOWN_GAPS: dict[str, str] = {} +LOCAL_WHEEL_PACKAGES = {"docopt"} + + +def _fetch_json(url: str) -> dict: + request = urllib.request.Request(url, headers={"User-Agent": "eegprep-pyodide-resolution-check"}) + with urllib.request.urlopen(request, timeout=HTTP_TIMEOUT_S) as response: + return json.load(response) + + +def _pyodide_environment(pyodide_info: dict) -> dict[str, str]: + """Build a packaging.markers environment for the Pyodide/Emscripten target.""" + python_full_version = pyodide_info["python"] + python_version = ".".join(python_full_version.split(".")[:2]) + return { + "implementation_name": "cpython", + "implementation_version": python_full_version, + "os_name": "posix", + "platform_machine": "wasm32", + "platform_release": "", + "platform_system": "Emscripten", + "platform_version": "", + "python_full_version": python_full_version, + "platform_python_implementation": "CPython", + "python_version": python_version, + "sys_platform": "emscripten", + } + + +def _load_uv_lock_packages() -> list[dict]: + data = tomllib.loads(UV_LOCK_PATH.read_text()) + return data["package"] + + +def _resolve_dependency(edge: dict, packages_by_key: dict[tuple[str, str], dict], packages_by_name: dict) -> dict: + if "version" in edge: + return packages_by_key[(edge["name"], edge["version"])] + candidates = packages_by_name[edge["name"]] + if len(candidates) != 1: + raise RuntimeError(f"Ambiguous uv.lock dependency edge for {edge['name']!r} without a pinned version") + return candidates[0] + + +def _base_closure(environment: dict[str, str]) -> dict[str, str]: + """Return {package_name: version} reachable from eegprep's base dependencies under `environment`.""" + packages = _load_uv_lock_packages() + packages_by_key = {(pkg["name"], pkg["version"]): pkg for pkg in packages if "version" in pkg} + packages_by_name: dict[str, list[dict]] = {} + for pkg in packages: + packages_by_name.setdefault(pkg["name"], []).append(pkg) + + (eegprep,) = packages_by_name["eegprep"] + closure: dict[str, str] = {} + stack = list(eegprep.get("dependencies", [])) + while stack: + edge = stack.pop() + marker = edge.get("marker") + if marker and not Marker(marker).evaluate(environment): + continue + pkg = _resolve_dependency(edge, packages_by_key, packages_by_name) + if pkg["name"] in closure: + continue + closure[pkg["name"]] = pkg["version"] + stack.extend(pkg.get("dependencies", [])) + return closure + + +def _is_universal_wheel(filename: str) -> bool: + """A wheel usable under Pyodide/micropip regardless of Python minor version: abi=none, platform=any.""" + stem = filename[: -len(".whl")] + tags = stem.split("-") + if len(tags) < 3: + return False + abi_tag, platform_tag = tags[-2], tags[-1] + return abi_tag == "none" and platform_tag == "any" + + +def _find_universal_wheel(name: str, version: str) -> str | None: + try: + release = _fetch_json(PYPI_JSON_URL.format(name=name, version=version)) + except urllib.error.HTTPError as exc: + raise RuntimeError(f"PyPI lookup failed for {name}=={version}: {exc}") from exc + for url_info in release.get("urls", []): + if url_info.get("packagetype") == "bdist_wheel" and _is_universal_wheel(url_info["filename"]): + return url_info["filename"] + return None + + +@dataclass +class PackageCheck: + name: str + version: str + in_pyodide_lock: bool + universal_wheel: str | None + + @property + def ok(self) -> bool: + return self.in_pyodide_lock or self.universal_wheel is not None + + +def main() -> int: + pyodide_lock = _fetch_json(PYODIDE_LOCK_URL) + environment = _pyodide_environment(pyodide_lock["info"]) + pyodide_names = {canonicalize_name(name) for name in pyodide_lock["packages"]} + + closure = _base_closure(environment) + + results = [] + for name, version in sorted(closure.items()): + in_lock = canonicalize_name(name) in pyodide_names + wheel = None if in_lock else _find_universal_wheel(name, version) + results.append(PackageCheck(name, version, in_lock, wheel)) + + for result in results: + if result.in_pyodide_lock: + source = "pyodide-lock" + elif result.universal_wheel: + source = f"pypi wheel: {result.universal_wheel}" + elif result.name in LOCAL_WHEEL_PACKAGES: + source = "local pure-Python wheel built by the Phase 2 harness" + elif result.name in KNOWN_GAPS: + source = f"NOT FOUND, known gap: {KNOWN_GAPS[result.name]}" + else: + source = "NOT FOUND (no Pyodide build, no pure-Python wheel)" + status = ( + "ok" + if result.ok or result.name in LOCAL_WHEEL_PACKAGES + else ("known-gap" if result.name in KNOWN_GAPS else "FAIL") + ) + print(f"[{status}] {result.name}=={result.version}: {source}") + + failures = [ + result + for result in results + if not result.ok and result.name not in KNOWN_GAPS and result.name not in LOCAL_WHEEL_PACKAGES + ] + known = [result for result in results if not result.ok and result.name in KNOWN_GAPS] + print( + f"\nChecked {len(results)} base packages against Pyodide {PYODIDE_VERSION}: " + f"{len(failures)} failing, {len(known)} known gaps (not blocking)." + ) + return 1 if failures else 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/tools/iclabel/artifacts/iclabel_float32.onnx b/tools/iclabel/artifacts/iclabel_float32.onnx new file mode 100644 index 00000000..9c67b291 Binary files /dev/null and b/tools/iclabel/artifacts/iclabel_float32.onnx differ diff --git a/tools/iclabel/artifacts/iclabel_int8_calibrated.onnx b/tools/iclabel/artifacts/iclabel_int8_calibrated.onnx new file mode 100644 index 00000000..76384dc4 Binary files /dev/null and b/tools/iclabel/artifacts/iclabel_int8_calibrated.onnx differ diff --git a/tools/iclabel/artifacts/iclabel_int8_weight_only.onnx b/tools/iclabel/artifacts/iclabel_int8_weight_only.onnx new file mode 100644 index 00000000..0b136301 Binary files /dev/null and b/tools/iclabel/artifacts/iclabel_int8_weight_only.onnx differ diff --git a/tools/iclabel/calibration_features.npz b/tools/iclabel/calibration_features.npz new file mode 100644 index 00000000..ae70815c Binary files /dev/null and b/tools/iclabel/calibration_features.npz differ diff --git a/tools/iclabel/evaluation_features.npz b/tools/iclabel/evaluation_features.npz new file mode 100644 index 00000000..e7343af9 Binary files /dev/null and b/tools/iclabel/evaluation_features.npz differ diff --git a/tools/iclabel/evaluation_manifest.json b/tools/iclabel/evaluation_manifest.json new file mode 100644 index 00000000..253d7804 --- /dev/null +++ b/tools/iclabel/evaluation_manifest.json @@ -0,0 +1,208 @@ +{ + "status": "frozen", + "frozen_on": "2026-09-17", + "purpose": "Held-out subject-disjoint real independent components for ICLabel float32 versus int8 parity.", + "dataset": { + "nemar_dataset_id": "on002680", + "openneuro_source_id": "ds002680", + "nemar_version": "v1.0.0", + "nemar_doi": "10.82901/nemar.on002680.v1.0.0", + "openneuro_doi": "10.18112/openneuro.ds002680.v1.2.0", + "license": "CC0", + "catalog_url": "https://api.nemar.org/datasets/on002680", + "dataset_url": "https://nemar.org/dataset/on002680", + "mirror_repository": "https://github.com/nemarDatasets/on002680", + "mirror_commit": "5cadd7d16423313a103ba24655e31245ffba8a4e", + "nemar_zarr_source_commit": "a8ebb56c5a2c4ee91a30d0b0d29072b3271d4384", + "source_pointer_url_template": "https://raw.githubusercontent.com/nemarDatasets/on002680/5cadd7d16423313a103ba24655e31245ffba8a4e/{source_path}", + "s3_object_url_template": "https://nemar.s3.us-east-2.amazonaws.com/on002680/objects/{git_annex_key}", + "openneuro_s3_url_template": "https://s3.amazonaws.com/openneuro.org/ds002680/{source_path}" + }, + "selection_policy": { + "subjects": [ + "sub-002", + "sub-003", + "sub-004", + "sub-005", + "sub-006", + "sub-007", + "sub-008", + "sub-009", + "sub-010", + "sub-011", + "sub-012", + "sub-013", + "sub-014", + "sub-015" + ], + "evaluation_subjects": ["sub-002", "sub-003", "sub-004", "sub-005", "sub-006", "sub-007", "sub-008"], + "calibration_subjects": ["sub-009", "sub-010", "sub-011", "sub-012", "sub-013", "sub-014", "sub-015"], + "recordings_per_subject": 1, + "evaluation_session": "ses-01", + "calibration_session": "ses-01", + "run": "run-1", + "subject_disjoint": true, + "component_indices_are_zero_based": true, + "confidence_based_filtering": false, + "label_based_selection": false, + "selection_declared_before_quantized_evaluation": true, + "selection_bias_mitigation": "One complete recording is selected per subject in each split, seven subjects are equally represented per split, and every component emitted by the fixed ICA solve is retained. Evaluation and calibration subjects are disjoint. A pre-freeze audit found the session-01 and session-02 signal matrices identical for sampled subjects, so session identity is not treated as independence. No component or subject is selected or removed using float32 probabilities, labels, thresholds, or quantization results." + }, + "preprocessing": { + "input_format": "Public NEMAR EEGLAB .set recordings with embedded data.", + "loader": "eegprep.pop_loadset", + "ica": "eegprep.functions.popfunc.eeg_runica.eeg_runica(EEG, extended=1, seed=379, rndreset='off', maxsteps=512, verbose=False)", + "ica_component_order": "eeg_runica output order; no post-hoc label- or confidence-based reordering", + "thread_policy": "OPENBLAS_NUM_THREADS=1, OMP_NUM_THREADS=1, MKL_NUM_THREADS=1", + "feature_extraction": "eegprep.plugins.ICLabel.ICL_feature_extractor.ICL_feature_extractor(EEG, True)", + "additional_filtering_or_resampling": "none", + "network_inputs": "The saved float32 features are expanded and transposed with the same four-way augmentation and NCHW conversion used by iclabel.py. Quantization does not alter feature extraction, normalization, augmentation, or softmax." + }, + "generation": { + "script": "tools/iclabel/freeze_iclabel_evaluation.py", + "command": "MNE_DONTWRITE_HOME=true OPENBLAS_NUM_THREADS=1 OMP_NUM_THREADS=1 MKL_NUM_THREADS=1 uv run --no-sync python -m tools.iclabel.freeze_iclabel_evaluation --manifest tools/iclabel/evaluation_manifest.json --data-dir --output-dir tools/iclabel", + "data_directory_layout": "The script stores verified public .set objects below // and resumable per-recording feature checkpoints below /features//.npz.", + "resume_rule": "Existing recordings are accepted only when their manifest size and MD5 match; existing feature checkpoints are accepted only when their expected float32 shapes and finiteness checks pass." + }, + "evaluation": { + "component_count": 217, + "feature_archive": { + "path": "tools/iclabel/evaluation_features.npz", + "sha256": "dbb2c0cc063d29ab2e8e59fafb0f474df780a458ac17be9d4732abc0233c8b44", + "dtype": "float32", + "teacher_class_distribution": { + "Brain": 18, + "Muscle": 1, + "Eye": 5, + "Heart": 0, + "Line Noise": 0, + "Channel Noise": 0, + "Other": 193 + } + }, + "recordings": [ + { + "subject": "sub-002", + "source_path": "sub-002/ses-01/eeg/sub-002_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s27111528--f1afdab0abe67a78c043024e339847e1.set", + "size_bytes": 27111528, + "md5": "f1afdab0abe67a78c043024e339847e1" + }, + { + "subject": "sub-003", + "source_path": "sub-003/ses-01/eeg/sub-003_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s29796960--5fb386464fa9ba3e5721784fe4866e08.set", + "size_bytes": 29796960, + "md5": "5fb386464fa9ba3e5721784fe4866e08" + }, + { + "subject": "sub-004", + "source_path": "sub-004/ses-01/eeg/sub-004_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s27130536--4957f1826c4a31181a733a045ad458a2.set", + "size_bytes": 27130536, + "md5": "4957f1826c4a31181a733a045ad458a2" + }, + { + "subject": "sub-005", + "source_path": "sub-005/ses-01/eeg/sub-005_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s28284024--cabd9de4fa7aec31f848a4360d9c3865.set", + "size_bytes": 28284024, + "md5": "cabd9de4fa7aec31f848a4360d9c3865" + }, + { + "subject": "sub-006", + "source_path": "sub-006/ses-01/eeg/sub-006_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s27214856--ca68281bf7ec06dd1aa4d7591bb92db5.set", + "size_bytes": 27214856, + "md5": "ca68281bf7ec06dd1aa4d7591bb92db5" + }, + { + "subject": "sub-007", + "source_path": "sub-007/ses-01/eeg/sub-007_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s32855400--537ecaa9169e6ce82d980f0f1a2f3968.set", + "size_bytes": 32855400, + "md5": "537ecaa9169e6ce82d980f0f1a2f3968" + }, + { + "subject": "sub-008", + "source_path": "sub-008/ses-01/eeg/sub-008_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s27967984--fa74eef298dea1a89109dc816506cd49.set", + "size_bytes": 27967984, + "md5": "fa74eef298dea1a89109dc816506cd49" + } + ] + }, + "calibration": { + "component_count": 217, + "feature_archive": { + "path": "tools/iclabel/calibration_features.npz", + "sha256": "6d5ee10f8c2a62a336e3f6372e686e00adad763b3d25f92728477a025e7d9c77", + "dtype": "float32" + }, + "recordings": [ + { + "subject": "sub-009", + "source_path": "sub-009/ses-01/eeg/sub-009_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s29952960--fa3a60a1cab32edea706089d767ed6c2.set", + "size_bytes": 29952960, + "md5": "fa3a60a1cab32edea706089d767ed6c2" + }, + { + "subject": "sub-010", + "source_path": "sub-010/ses-01/eeg/sub-010_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s28014160--e4506f8a5f506fed93dd18d47ce3235c.set", + "size_bytes": 28014160, + "md5": "e4506f8a5f506fed93dd18d47ce3235c" + }, + { + "subject": "sub-011", + "source_path": "sub-011/ses-01/eeg/sub-011_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s31064192--99357f6d192139929b44b3373cc8e33e.set", + "size_bytes": 31064192, + "md5": "99357f6d192139929b44b3373cc8e33e" + }, + { + "subject": "sub-012", + "source_path": "sub-012/ses-01/eeg/sub-012_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s28173976--8e009a8be773b494bde2403155062f5b.set", + "size_bytes": 28173976, + "md5": "8e009a8be773b494bde2403155062f5b" + }, + { + "subject": "sub-013", + "source_path": "sub-013/ses-01/eeg/sub-013_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s27891160--e33fa885e4b28ff90348e0b9ecf6d746.set", + "size_bytes": 27891160, + "md5": "e33fa885e4b28ff90348e0b9ecf6d746" + }, + { + "subject": "sub-014", + "source_path": "sub-014/ses-01/eeg/sub-014_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s27532304--0f4427a4694ac72648a7fc428a6d8b0f.set", + "size_bytes": 27532304, + "md5": "0f4427a4694ac72648a7fc428a6d8b0f" + }, + { + "subject": "sub-015", + "source_path": "sub-015/ses-01/eeg/sub-015_ses-01_task-gonogo_run-1_eeg.set", + "component_indices": [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30], + "git_annex_key": "MD5E-s28665632--dcea224b0a2fbeed595b3ea1ea336d13.set", + "size_bytes": 28665632, + "md5": "dcea224b0a2fbeed595b3ea1ea336d13" + } + ] + } +} diff --git a/tools/iclabel/export_iclabel_onnx.py b/tools/iclabel/export_iclabel_onnx.py new file mode 100644 index 00000000..c82695fd --- /dev/null +++ b/tools/iclabel/export_iclabel_onnx.py @@ -0,0 +1,121 @@ +"""Export the reproducible float32 ICLabel default network to ONNX. + +Offline, developer-only tool: it produces the float32 reference artifact +``tools/iclabel/artifacts/iclabel_float32.onnx`` by default. The selected +artifact is copied to the package filename ``iclabel.onnx`` only after the +Phase 6 parity gate passes. This script is not installed with the package. + +Regenerate the float32 reference after changing ``netICL.mat`` or +``iclabel_net.py``: + + uv sync --group dev --extra torch --extra iclabel + uv run --no-sync python tools/iclabel/export_iclabel_onnx.py + +Provenance: the exported graph is ``ICLabelNet.forward`` (three convolutional +branches over the topography image, PSD, and autocorrelation features, +concatenated and passed through a final conv + softmax), pinned at +``ONNX_OPSET_VERSION`` below. Feature extraction, input normalization, the +4-way augmentation, and the post-network averaging (``iclabel.py``) are +unchanged by this export; they run in plain numpy before and after the +network call. + +torch is required only to run this script, never at ICLabel classification +time; the runtime backend is onnxruntime (``eegprep[iclabel]``). +""" + +import argparse +import logging +from pathlib import Path + +import numpy as np +import onnx +import onnxruntime as ort +import torch + +from eegprep.plugins.ICLabel.iclabel_net import ICLabelNet + +logger = logging.getLogger(__name__) + +REPO_ROOT = Path(__file__).resolve().parents[2] +ICLABEL_DIR = REPO_ROOT / 'src' / 'eegprep' / 'plugins' / 'ICLabel' +ICLABEL_ARTIFACT_DIR = REPO_ROOT / 'tools' / 'iclabel' / 'artifacts' +DEFAULT_MAT_PATH = ICLABEL_DIR / 'netICL.mat' +DEFAULT_ONNX_PATH = ICLABEL_ARTIFACT_DIR / 'iclabel_float32.onnx' + +# Pinned per issue #377. The network only uses Conv2d, LeakyReLU, Softmax, +# Concat, and Reshape, all supported since opset 7-9, so the choice is driven +# by runtime compatibility rather than op coverage: opset 17 has been +# supported by onnxruntime since 1.14 (2023), giving broad compatibility +# today while staying modern enough for the onnxruntime-web/wasm target +# planned for phase 5. +ONNX_OPSET_VERSION = 17 + +INPUT_NAMES = ('image', 'psdmed', 'autocorr') +OUTPUT_NAMES = ('output',) +_INPUT_SHAPES = ((1, 1, 32, 32), (1, 1, 1, 100), (1, 1, 1, 100)) + + +def export(mat_path: Path = DEFAULT_MAT_PATH, onnx_path: Path = DEFAULT_ONNX_PATH) -> Path: + """Export the ICLabel default network to ``onnx_path`` with a pinned opset.""" + model = ICLabelNet(str(mat_path)) + model.eval() + onnx_path.parent.mkdir(parents=True, exist_ok=True) + + dummy_inputs = tuple(torch.zeros(shape, dtype=torch.float32) for shape in _INPUT_SHAPES) + dynamic_axes = {name: {0: 'batch'} for name in (*INPUT_NAMES, *OUTPUT_NAMES)} + + torch.onnx.export( + model, + dummy_inputs, + str(onnx_path), + input_names=list(INPUT_NAMES), + output_names=list(OUTPUT_NAMES), + opset_version=ONNX_OPSET_VERSION, + dynamic_axes=dynamic_axes, + do_constant_folding=True, + dynamo=False, + ) + onnx.checker.check_model(str(onnx_path)) + return onnx_path + + +def verify(mat_path: Path, onnx_path: Path, batch_size: int = 16, seed: int = 0) -> float: + """Compare torch and onnxruntime outputs on random inputs; return max abs diff.""" + model = ICLabelNet(str(mat_path)) + model.eval() + + rng = np.random.default_rng(seed) + inputs = { + name: rng.standard_normal((batch_size, *shape[1:])).astype(np.float32) + for name, shape in zip(INPUT_NAMES, _INPUT_SHAPES) + } + + with torch.no_grad(): + torch_out = model(*(torch.from_numpy(inputs[name]) for name in INPUT_NAMES)).numpy() + + session = ort.InferenceSession(str(onnx_path), providers=['CPUExecutionProvider']) + (onnx_out,) = session.run(list(OUTPUT_NAMES), inputs) + + max_abs_diff = float(np.max(np.abs(torch_out - onnx_out))) + return max_abs_diff + + +def main() -> None: + logging.basicConfig(level=logging.INFO, force=True) + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument('--mat-path', type=Path, default=DEFAULT_MAT_PATH) + parser.add_argument('--onnx-path', type=Path, default=DEFAULT_ONNX_PATH) + args = parser.parse_args() + + onnx_path = export(args.mat_path, args.onnx_path) + size_mb = onnx_path.stat().st_size / (1024 * 1024) + logger.info("Exported %s (opset %d, %.2f MB)", onnx_path, ONNX_OPSET_VERSION, size_mb) + + max_abs_diff = verify(args.mat_path, onnx_path) + logger.info("torch vs onnxruntime max abs diff on random inputs: %.3e", max_abs_diff) + if max_abs_diff > 1e-4: + raise RuntimeError(f"ONNX export diverges from torch beyond tolerance: max abs diff {max_abs_diff:.3e}") + + +if __name__ == '__main__': + main() diff --git a/tools/iclabel/freeze_iclabel_evaluation.py b/tools/iclabel/freeze_iclabel_evaluation.py new file mode 100644 index 00000000..26a502bb --- /dev/null +++ b/tools/iclabel/freeze_iclabel_evaluation.py @@ -0,0 +1,205 @@ +"""Download the frozen ICLabel recordings and build float32 feature archives. + +The manifest is the authority for recording selection and checksums. This tool +only downloads the public NEMAR objects named by that manifest, runs the fixed +EEGPrep ICA and ICLabel feature path, and writes resumable derived archives. +It does not run a quantized model or select a shipped artifact. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +import os +import tempfile +from collections.abc import Mapping, Sequence +from pathlib import Path +from urllib.parse import quote +from urllib.error import HTTPError +from urllib.request import Request, urlopen + +import numpy as np + +from eegprep import ICL_feature_extractor, eeg_runica, pop_loadset + +from tools.iclabel.quantize_iclabel_onnx import load_frozen_manifest + + +COMPONENT_COUNT = 31 +DEFAULT_DATA_DIR = Path("/private/tmp/eegprep-phase6-iclabel-data") +DEFAULT_OUTPUT_DIR = Path(__file__).parent +_CHUNK_SIZE = 1024 * 1024 + + +def _verify_recording(path: Path, recording: Mapping[str, object]) -> None: + expected_size = int(recording["size_bytes"]) + expected_md5 = str(recording["md5"]) + if path.stat().st_size != expected_size: + raise ValueError(f"Size mismatch for {path}: expected {expected_size} bytes") + digest = hashlib.md5(usedforsecurity=False) + with path.open("rb") as handle: + for chunk in iter(lambda: handle.read(_CHUNK_SIZE), b""): + digest.update(chunk) + if digest.hexdigest() != expected_md5: + raise ValueError(f"MD5 mismatch for {path}: expected {expected_md5}") + + +def _download_recording(manifest: Mapping[str, object], recording: Mapping[str, object], path: Path) -> None: + if path.exists(): + _verify_recording(path, recording) + return + + dataset = manifest["dataset"] + urls = [ + str(dataset["s3_object_url_template"]).format(git_annex_key=quote(str(recording["git_annex_key"]), safe="")), + str(dataset["openneuro_s3_url_template"]).format(source_path=quote(str(recording["source_path"]), safe="/")), + ] + path.parent.mkdir(parents=True, exist_ok=True) + partial_path = path.with_name(path.name + ".part") + for url_index, url in enumerate(urls): + request = Request(url, headers={"User-Agent": "EEGPrep-ICLabel-evaluation/1"}) + for attempt in range(3): + print(f"DOWNLOAD {recording['source_path']} <- {url}", flush=True) + try: + with urlopen(request, timeout=180) as response, partial_path.open("wb") as handle: + while chunk := response.read(_CHUNK_SIZE): + handle.write(chunk) + except HTTPError as error: + partial_path.unlink(missing_ok=True) + if url_index == 0 and error.code in (403, 404): + print(f"FALLBACK public OpenNeuro S3 after NEMAR HTTP {error.code}", flush=True) + break + raise + try: + _verify_recording(partial_path, recording) + except ValueError: + partial_path.unlink(missing_ok=True) + if attempt < 2: + print("RETRY after checksum mismatch", flush=True) + continue + if url_index == 0: + print("FALLBACK public OpenNeuro S3 after NEMAR checksum mismatch", flush=True) + break + raise + os.replace(partial_path, path) + return + raise RuntimeError(f"Unable to download a checksum-valid object for {recording['source_path']}") + + +def _feature_archive_is_valid(path: Path) -> bool: + if not path.exists(): + return False + try: + with np.load(path, allow_pickle=False) as archive: + arrays = [np.asarray(archive[name], dtype=np.float32) for name in ("topo", "psd", "autocorr")] + except (KeyError, OSError, ValueError): + return False + return ( + arrays[0].shape == (32, 32, 1, COMPONENT_COUNT) + and arrays[1].shape == (1, 100, 1, COMPONENT_COUNT) + and arrays[2].shape == (1, 100, 1, COMPONENT_COUNT) + and all(np.isfinite(array).all() for array in arrays) + ) + + +def _extract_recording_features(raw_path: Path) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + eeg = pop_loadset(str(raw_path)) + print(f"ICA {raw_path}", flush=True) + eeg = eeg_runica( + eeg, + extended=1, + seed=379, + rndreset="off", + maxsteps=512, + verbose=False, + ) + print(f"FEATURES {raw_path}", flush=True) + features = tuple(np.asarray(feature, dtype=np.float32) for feature in ICL_feature_extractor(eeg, True)) + if len(features) != 3: + raise ValueError(f"ICLabel feature extraction returned {len(features)} arrays for {raw_path}") + if any(not np.isfinite(feature).all() for feature in features): + raise ValueError(f"ICLabel feature extraction returned non-finite values for {raw_path}") + expected_shapes = ((32, 32, 1, COMPONENT_COUNT), (1, 100, 1, COMPONENT_COUNT), (1, 100, 1, COMPONENT_COUNT)) + if tuple(feature.shape for feature in features) != expected_shapes: + raise ValueError( + f"Unexpected ICLabel feature shapes for {raw_path}: {tuple(feature.shape for feature in features)}" + ) + return features + + +def _recording_features( + manifest: Mapping[str, object], + split_name: str, + recording: Mapping[str, object], + data_dir: Path, +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + raw_path = data_dir / split_name / str(recording["source_path"]) + _download_recording(manifest, recording, raw_path) + checkpoint = data_dir / "features" / split_name / f"{recording['subject']}.npz" + if _feature_archive_is_valid(checkpoint): + print(f"CHECKPOINT {split_name}/{recording['subject']}", flush=True) + with np.load(checkpoint, allow_pickle=False) as archive: + return tuple(np.asarray(archive[name], dtype=np.float32) for name in ("topo", "psd", "autocorr")) + + features = _extract_recording_features(raw_path) + checkpoint.parent.mkdir(parents=True, exist_ok=True) + temporary = checkpoint.with_name(checkpoint.name + ".part") + with temporary.open("wb") as handle: + np.savez_compressed(handle, topo=features[0], psd=features[1], autocorr=features[2]) + os.replace(temporary, checkpoint) + return features + + +def _write_archive(path: Path, features: Sequence[np.ndarray]) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + with tempfile.NamedTemporaryFile(dir=path.parent, prefix=f".{path.stem}-", suffix=".npz", delete=False) as handle: + temporary = Path(handle.name) + np.savez_compressed(handle, topo=features[0], psd=features[1], autocorr=features[2]) + os.replace(temporary, path) + + +def build_split( + manifest: Mapping[str, object], + split_name: str, + data_dir: Path, + output_path: Path, + limit: int | None = None, +) -> None: + split = manifest[split_name] + recordings = split["recordings"] + if limit is not None: + recordings = recordings[:limit] + extracted = [] + for index, recording in enumerate(recordings, start=1): + print(f"RECORDING {split_name} {index}/{len(recordings)} {recording['subject']}", flush=True) + extracted.append(_recording_features(manifest, split_name, recording, data_dir)) + if limit is None and len(extracted) != int(split["component_count"]) // COMPONENT_COUNT: + raise ValueError(f"{split_name} archive did not process all manifest recordings") + features = tuple(np.concatenate([recording[index] for recording in extracted], axis=3) for index in range(3)) + expected_count = len(extracted) * COMPONENT_COUNT + if any(feature.shape[3] != expected_count for feature in features): + raise ValueError(f"{split_name} archive has the wrong component count") + _write_archive(output_path, features) + print(f"ARCHIVE {split_name} {output_path} components={expected_count}", flush=True) + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--manifest", type=Path, default=Path(__file__).with_name("evaluation_manifest.json")) + parser.add_argument("--data-dir", type=Path, required=True) + parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR) + parser.add_argument("--split", choices=("all", "evaluation", "calibration"), default="all") + parser.add_argument("--limit", type=int, default=None, help=argparse.SUPPRESS) + args = parser.parse_args() + + manifest = load_frozen_manifest(args.manifest) + splits = ("evaluation", "calibration") if args.split == "all" else (args.split,) + for split_name in splits: + output_name = f"{split_name}_features.npz" + build_split(manifest, split_name, args.data_dir, args.output_dir / output_name, args.limit) + print(json.dumps({"status": "complete", "splits": list(splits)}), flush=True) + + +if __name__ == "__main__": + main() diff --git a/tools/iclabel/iclabel_net_load_py_measures.py b/tools/iclabel/iclabel_net_load_py_measures.py index 9c678fae..8766e815 100644 --- a/tools/iclabel/iclabel_net_load_py_measures.py +++ b/tools/iclabel/iclabel_net_load_py_measures.py @@ -5,10 +5,14 @@ via ``system('... iclabel_net_load_py_measures.py')``. It reuses the single canonical :class:`~eegprep.plugins.ICLabel.iclabel_net.ICLabelNet` definition so the parity harness exercises the same network as production ICLabel. + +``netICL.mat`` is a repo dev asset, not a packaged runtime resource (the +wheel ships ``iclabel.onnx`` instead), so it is resolved from the source +tree rather than through ``importlib.resources``. """ import logging -from importlib.resources import files +from pathlib import Path import scipy.io import torch @@ -17,9 +21,11 @@ logger = logging.getLogger(__name__) +REPO_ROOT = Path(__file__).resolve().parents[2] + if __name__ == "__main__": - net_path = files("eegprep").joinpath("plugins").joinpath("ICLabel").joinpath("netICL.mat") + net_path = REPO_ROOT / "src" / "eegprep" / "plugins" / "ICLabel" / "netICL.mat" model = ICLabelNet(str(net_path)) data = scipy.io.loadmat('python_temp_reformated.mat') image_mat = data['grid'][0][0] diff --git a/tools/iclabel/parity_report.json b/tools/iclabel/parity_report.json new file mode 100644 index 00000000..29faa39d --- /dev/null +++ b/tools/iclabel/parity_report.json @@ -0,0 +1,245 @@ +{ + "candidates": { + "calibrated": { + "artifact": "iclabel_int8_calibrated.onnx", + "candidate_class_distribution": { + "Brain": 16, + "Channel Noise": 0, + "Eye": 5, + "Heart": 0, + "Line Noise": 0, + "Muscle": 2, + "Other": 194 + }, + "gate_pass": false, + "max_probability_abs_diff": 0.09521264, + "mean_probability_abs_diff": 0.0058981734, + "keep_reject_agreement": 1.0, + "per_class_agreement": { + "Brain": { + "agreement": 0.8333333333333334, + "count": 18 + }, + "Channel Noise": { + "agreement": null, + "count": 0 + }, + "Eye": { + "agreement": 1.0, + "count": 5 + }, + "Heart": { + "agreement": null, + "count": 0 + }, + "Line Noise": { + "agreement": null, + "count": 0 + }, + "Muscle": { + "agreement": 1.0, + "count": 1 + }, + "Other": { + "agreement": 0.9948186528497409, + "count": 193 + } + }, + "sample_count": 217, + "sha256": "a98fdb146678daf74c588898e477ff3af234a682449a1712e4e96d3af01974b1", + "size_bytes": 2970375, + "teacher_class_distribution": { + "Brain": 18, + "Channel Noise": 0, + "Eye": 5, + "Heart": 0, + "Line Noise": 0, + "Muscle": 1, + "Other": 193 + }, + "top1_agreement": 0.9815668202764977 + }, + "weight_only": { + "artifact": "iclabel_int8_weight_only.onnx", + "candidate_class_distribution": { + "Brain": 18, + "Channel Noise": 0, + "Eye": 5, + "Heart": 0, + "Line Noise": 0, + "Muscle": 1, + "Other": 193 + }, + "gate_pass": true, + "max_probability_abs_diff": 0.013459265, + "mean_probability_abs_diff": 0.00080229656, + "keep_reject_agreement": 1.0, + "per_class_agreement": { + "Brain": { + "agreement": 1.0, + "count": 18 + }, + "Channel Noise": { + "agreement": null, + "count": 0 + }, + "Eye": { + "agreement": 1.0, + "count": 5 + }, + "Heart": { + "agreement": null, + "count": 0 + }, + "Line Noise": { + "agreement": null, + "count": 0 + }, + "Muscle": { + "agreement": 1.0, + "count": 1 + }, + "Other": { + "agreement": 1.0, + "count": 193 + } + }, + "sample_count": 217, + "sha256": "8c0537d299e8776eef418d4a697547ad0b317902d656f69f64440169c564d8e3", + "size_bytes": 2932897, + "teacher_class_distribution": { + "Brain": 18, + "Channel Noise": 0, + "Eye": 5, + "Heart": 0, + "Line Noise": 0, + "Muscle": 1, + "Other": 193 + }, + "top1_agreement": 1.0 + } + }, + "class_names": [ + "Brain", + "Muscle", + "Eye", + "Heart", + "Line Noise", + "Channel Noise", + "Other" + ], + "default_artifact": "iclabel_int8_weight_only.onnx", + "feature_archives": { + "calibration": { + "component_count": 217, + "path": "tools/iclabel/calibration_features.npz", + "sha256": "6d5ee10f8c2a62a336e3f6372e686e00adad763b3d25f92728477a025e7d9c77" + }, + "evaluation": { + "component_count": 217, + "path": "tools/iclabel/evaluation_features.npz", + "sha256": "dbb2c0cc063d29ab2e8e59fafb0f474df780a458ac17be9d4732abc0233c8b44" + } + }, + "float32_reference": { + "artifact": "iclabel_float32.onnx", + "candidate_class_distribution": { + "Brain": 18, + "Channel Noise": 0, + "Eye": 5, + "Heart": 0, + "Line Noise": 0, + "Muscle": 1, + "Other": 193 + }, + "max_probability_abs_diff": 0.0, + "mean_probability_abs_diff": 0.0, + "keep_reject_agreement": 1.0, + "per_class_agreement": { + "Brain": { + "agreement": 1.0, + "count": 18 + }, + "Channel Noise": { + "agreement": null, + "count": 0 + }, + "Eye": { + "agreement": 1.0, + "count": 5 + }, + "Heart": { + "agreement": null, + "count": 0 + }, + "Line Noise": { + "agreement": null, + "count": 0 + }, + "Muscle": { + "agreement": 1.0, + "count": 1 + }, + "Other": { + "agreement": 1.0, + "count": 193 + } + }, + "sample_count": 217, + "sha256": "1f16c67ebdbd8e6a30759ab628b61b693a0c4011e35a20e538acbebc1d4ebce4", + "size_bytes": 11626631, + "teacher_class_distribution": { + "Brain": 18, + "Channel Noise": 0, + "Eye": 5, + "Heart": 0, + "Line Noise": 0, + "Muscle": 1, + "Other": 193 + }, + "top1_agreement": 1.0 + }, + "gate": { + "minimum_keep_reject_agreement": 0.99, + "minimum_top1_agreement": 0.95, + "maximum_probability_abs_diff": 0.015 + }, + "manifest": "evaluation_manifest.json", + "quantization": { + "calibrated": "ONNX Runtime static per-channel Conv QDQ with MinMax calibration and float model I/O", + "feature_dtype": "float32", + "input_normalization_and_augmentation": "unchanged from iclabel.py", + "output_softmax": "unchanged from the float32 graph", + "weight_only": "per-output-channel int8 Conv weights with float DequantizeLinear before Conv" + }, + "thresholds": [ + [ + null, + null + ], + [ + 0.9, + 1.0 + ], + [ + 0.9, + 1.0 + ], + [ + null, + null + ], + [ + null, + null + ], + [ + null, + null + ], + [ + null, + null + ] + ] +} diff --git a/tools/iclabel/quantize_iclabel_onnx.py b/tools/iclabel/quantize_iclabel_onnx.py new file mode 100644 index 00000000..52900c59 --- /dev/null +++ b/tools/iclabel/quantize_iclabel_onnx.py @@ -0,0 +1,464 @@ +"""Build and evaluate ICLabel int8 ONNX artifacts. + +This developer-only tool deliberately leaves ICLabel feature extraction, +normalization, augmentation, and softmax outside quantization. The weight-only +candidate stores Conv weights as int8 with float dequantization before each +Conv. The calibrated candidate uses ONNX Runtime static QDQ quantization for +Conv weights and activations, while preserving float model inputs and output +softmax. +""" + +from __future__ import annotations + +import argparse +import hashlib +import json +from collections.abc import Mapping, Sequence +from pathlib import Path +from typing import Any + +import numpy as np + +from eegprep.plugins.ICLabel.eeg_icflag import eeg_icflag +from eegprep.plugins.ICLabel.pop_icflag import DEFAULT_ICFLAG_THRESHOLDS + + +CLASS_NAMES = ("Brain", "Muscle", "Eye", "Heart", "Line Noise", "Channel Noise", "Other") +MIN_TOP1_AGREEMENT = 0.95 +MIN_KEEP_REJECT_AGREEMENT = 0.99 +MAX_PROBABILITY_ABS_DIFF = 0.015 +_INPUT_NAMES = ("image", "psdmed", "autocorr") +_OUTPUT_NAME = "output" +_REPO_ROOT = Path(__file__).resolve().parents[2] +DEFAULT_FLOAT32_ARTIFACT = _REPO_ROOT / "tools" / "iclabel" / "artifacts" / "iclabel_float32.onnx" +DEFAULT_FROZEN_MANIFEST = Path(__file__).with_name("evaluation_manifest.json") +DEFAULT_EVALUATION_FEATURES = Path(__file__).with_name("evaluation_features.npz") +DEFAULT_CALIBRATION_FEATURES = Path(__file__).with_name("calibration_features.npz") +DEFAULT_ARTIFACT_DIR = Path(__file__).with_name("artifacts") +DEFAULT_WEIGHT_ONLY_ARTIFACT = DEFAULT_ARTIFACT_DIR / "iclabel_int8_weight_only.onnx" +DEFAULT_CALIBRATED_ARTIFACT = DEFAULT_ARTIFACT_DIR / "iclabel_int8_calibrated.onnx" +DEFAULT_REPORT = Path(__file__).with_name("parity_report.json") +_HASH_CHUNK_SIZE = 1024 * 1024 + + +def _sha256(path: Path) -> str: + digest = hashlib.sha256() + with path.open("rb") as handle: + for chunk in iter(lambda: handle.read(_HASH_CHUNK_SIZE), b""): + digest.update(chunk) + return digest.hexdigest() + + +def load_frozen_manifest(path: Path = DEFAULT_FROZEN_MANIFEST) -> dict[str, Any]: + """Load and validate the committed, pre-quantization evaluation manifest.""" + with Path(path).open(encoding="utf-8") as handle: + manifest = json.load(handle) + _validate_frozen_manifest(manifest) + return manifest + + +def _validate_frozen_manifest(manifest: Mapping[str, Any]) -> None: + if manifest.get("status") != "frozen": + raise ValueError("ICLabel evaluation manifest must have status='frozen'") + policy = manifest.get("selection_policy") + if not isinstance(policy, Mapping): + raise ValueError("ICLabel evaluation manifest is missing selection_policy") + if policy.get("confidence_based_filtering") is not False: + raise ValueError("ICLabel evaluation selection cannot use confidence filtering") + if policy.get("label_based_selection") is not False: + raise ValueError("ICLabel evaluation selection cannot use label filtering") + if policy.get("selection_declared_before_quantized_evaluation") is not True: + raise ValueError("ICLabel evaluation selection must be declared before quantized evaluation") + if policy.get("subject_disjoint") is not True: + raise ValueError("ICLabel evaluation and calibration subjects must be disjoint") + + splits = {} + for split_name in ("evaluation", "calibration"): + split = manifest.get(split_name) + if not isinstance(split, Mapping) or not isinstance(split.get("recordings"), list): + raise ValueError(f"ICLabel manifest is missing {split_name} recordings") + splits[split_name] = split["recordings"] + feature_archive = split.get("feature_archive") + if not isinstance(feature_archive, Mapping): + raise ValueError(f"ICLabel manifest is missing {split_name} feature archive metadata") + if not isinstance(feature_archive.get("path"), str) or not isinstance(feature_archive.get("sha256"), str): + raise ValueError(f"ICLabel {split_name} feature archive metadata must include path and sha256") + if split.get("component_count") != len(split["recordings"]) * 31: + raise ValueError(f"ICLabel {split_name} component_count does not match its recordings") + for recording in split["recordings"]: + if not isinstance(recording, Mapping): + raise ValueError(f"ICLabel {split_name} recording must be an object") + if recording.get("component_indices") != list(range(31)): + raise ValueError(f"ICLabel {split_name} recordings must retain all 31 components") + for field in ("subject", "source_path", "git_annex_key", "size_bytes", "md5"): + if field not in recording: + raise ValueError(f"ICLabel {split_name} recording is missing {field}") + + evaluation_subjects = [recording["subject"] for recording in splits["evaluation"]] + calibration_subjects = [recording["subject"] for recording in splits["calibration"]] + if evaluation_subjects != sorted(set(evaluation_subjects)): + raise ValueError("ICLabel evaluation subjects must be unique and sorted") + if calibration_subjects != sorted(set(calibration_subjects)): + raise ValueError("ICLabel calibration subjects must be unique and sorted") + if set(evaluation_subjects).intersection(calibration_subjects): + raise ValueError("ICLabel evaluation and calibration subjects must be disjoint") + if policy.get("evaluation_subjects") != evaluation_subjects: + raise ValueError("ICLabel evaluation subject policy does not match its recordings") + if policy.get("calibration_subjects") != calibration_subjects: + raise ValueError("ICLabel calibration subject policy does not match its recordings") + evaluation_paths = {recording["source_path"] for recording in splits["evaluation"]} + calibration_paths = {recording["source_path"] for recording in splits["calibration"]} + if evaluation_paths.intersection(calibration_paths): + raise ValueError("ICLabel evaluation and calibration recordings must be disjoint") + + +def load_feature_archive(path: Path) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + """Load saved float32 topography, PSD, and autocorrelation features.""" + with np.load(path, allow_pickle=False) as archive: + features = tuple(np.asarray(archive[name], dtype=np.float32) for name in ("topo", "psd", "autocorr")) + _validate_feature_arrays(features) + return features + + +def load_verified_feature_archive( + path: Path, + manifest: Mapping[str, Any], + split_name: str, +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + """Load a feature archive only when its frozen hash and count match the manifest.""" + split = manifest.get(split_name) + if not isinstance(split, Mapping): + raise ValueError(f"ICLabel manifest is missing {split_name}") + feature_archive = split.get("feature_archive") + if not isinstance(feature_archive, Mapping): + raise ValueError(f"ICLabel manifest is missing {split_name} feature archive metadata") + expected_sha256 = feature_archive.get("sha256") + if not isinstance(expected_sha256, str): + raise ValueError(f"ICLabel {split_name} feature archive metadata is missing sha256") + path = Path(path) + actual_sha256 = _sha256(path) + # Fail closed: gate metrics are only meaningful for the exact frozen input. + if actual_sha256 != expected_sha256: + raise ValueError( + f"ICLabel {split_name} feature archive SHA-256 mismatch: expected {expected_sha256}, got {actual_sha256}" + ) + features = load_feature_archive(path) + expected_component_count = split.get("component_count") + actual_component_count = features[0].shape[3] + if actual_component_count != expected_component_count: + raise ValueError( + f"ICLabel {split_name} feature archive has {actual_component_count} components; " + f"expected {expected_component_count}" + ) + return features + + +def _validate_feature_arrays(features: Sequence[np.ndarray]) -> None: + if len(features) != 3: + raise ValueError("ICLabel feature archive must contain topo, psd, and autocorr arrays") + shapes = [array.shape for array in features] + if len(shapes[0]) != 4 or shapes[0][:3] != (32, 32, 1): + raise ValueError(f"ICLabel topography features must have shape (32, 32, 1, n), got {shapes[0]}") + if len(shapes[1]) != 4 or len(shapes[2]) != 4 or shapes[1][:3] != (1, 100, 1) or shapes[2][:3] != (1, 100, 1): + raise ValueError(f"ICLabel PSD/autocorrelation features must have shape (1, 100, 1, n), got {shapes[1:]}") + if shapes[0][3] == 0 or not (shapes[0][3] == shapes[1][3] == shapes[2][3]): + raise ValueError("ICLabel feature arrays must contain the same component count") + if not all(np.isfinite(array).all() for array in features): + raise ValueError("ICLabel feature arrays must be finite") + + +def network_inputs_from_features(features: Sequence[np.ndarray]) -> dict[str, np.ndarray]: + """Build the frozen quantization inputs independently of runtime ICLabel.""" + _validate_feature_arrays(features) + topo, psd, autocorr = features + topo = np.single(np.concatenate([topo, -topo, topo[:, ::-1, :, :], -topo[:, ::-1, :, :]], axis=3)) + psd = np.single(np.tile(psd, (1, 1, 1, 4))) + autocorr = np.single(np.tile(autocorr, (1, 1, 1, 4))) + return { + "image": np.transpose(topo, (3, 2, 0, 1)), + "psdmed": np.transpose(psd, (3, 2, 0, 1)), + "autocorr": np.transpose(autocorr, (3, 2, 0, 1)), + } + + +def _run_network(model_path: Path, inputs: Mapping[str, np.ndarray]) -> np.ndarray: + import onnxruntime as ort + + session = ort.InferenceSession(str(model_path), providers=["CPUExecutionProvider"]) + (output,) = session.run([_OUTPUT_NAME], dict(inputs)) + return np.asarray(output, dtype=np.float32) + + +def predict_features(model_path: Path, features: Sequence[np.ndarray]) -> np.ndarray: + """Run an ICLabel artifact and return one 7-class row per component.""" + output = _run_network(model_path, network_inputs_from_features(features)) + output = output.T + output = np.reshape(output, (-1, 4), order="F") + output = np.mean(output, axis=1) + return np.reshape(output, (7, -1), order="F").T + + +def _rejection_flags(classifications: np.ndarray, thresholds: np.ndarray) -> np.ndarray: + eeg = {"etc": {"ic_classification": {"ICLabel": {"classifications": classifications}}}} + return np.asarray(eeg_icflag(eeg, thresholds)["reject"]["gcompreject"], dtype=bool) + + +def compare_predictions( + teacher: np.ndarray, + candidate: np.ndarray, + thresholds: np.ndarray = DEFAULT_ICFLAG_THRESHOLDS, +) -> dict[str, Any]: + """Compare probabilities, class labels, and ICLabel keep/reject decisions.""" + teacher = np.asarray(teacher, dtype=float) + candidate = np.asarray(candidate, dtype=float) + if teacher.shape != candidate.shape or teacher.ndim != 2 or teacher.shape[1] != len(CLASS_NAMES): + raise ValueError("ICLabel predictions must be matching (components, 7) arrays") + if not np.isfinite(teacher).all() or not np.isfinite(candidate).all(): + raise ValueError("ICLabel predictions must be finite") + + teacher_labels = np.argmax(teacher, axis=1) + candidate_labels = np.argmax(candidate, axis=1) + teacher_reject = _rejection_flags(teacher, thresholds) + candidate_reject = _rejection_flags(candidate, thresholds) + per_class: dict[str, dict[str, Any]] = {} + for class_index, class_name in enumerate(CLASS_NAMES): + selected = teacher_labels == class_index + count = int(np.sum(selected)) + per_class[class_name] = { + "count": count, + "agreement": None if count == 0 else float(np.mean(candidate_labels[selected] == class_index)), + } + + return { + "sample_count": int(teacher.shape[0]), + "max_probability_abs_diff": float(np.max(np.abs(teacher - candidate))), + "mean_probability_abs_diff": float(np.mean(np.abs(teacher - candidate))), + "top1_agreement": float(np.mean(teacher_labels == candidate_labels)), + "keep_reject_agreement": float(np.mean(teacher_reject == candidate_reject)), + "teacher_class_distribution": { + name: int(np.sum(teacher_labels == index)) for index, name in enumerate(CLASS_NAMES) + }, + "candidate_class_distribution": { + name: int(np.sum(candidate_labels == index)) for index, name in enumerate(CLASS_NAMES) + }, + "per_class_agreement": per_class, + } + + +def parity_gate_passes(metrics: Mapping[str, Any]) -> bool: + """Return whether the frozen probability and semantic parity gates pass.""" + return ( + float(metrics["top1_agreement"]) >= MIN_TOP1_AGREEMENT + and float(metrics["keep_reject_agreement"]) >= MIN_KEEP_REJECT_AGREEMENT + and float(metrics["max_probability_abs_diff"]) <= MAX_PROBABILITY_ABS_DIFF + ) + + +def select_default_artifact(report: Mapping[str, Any]) -> str: + """Select the first int8 candidate that passes the fixed parity gate.""" + candidates = report.get("candidates", {}) + for name in ("weight_only", "calibrated"): + candidate = candidates.get(name, {}) + if parity_gate_passes(candidate): + return str(candidate["artifact"]) + return "iclabel.onnx" + + +def quantize_weight_only(model_input: Path, model_output: Path) -> Path: + """Write a Conv-weight-only int8 QDQ artifact with float inputs and outputs.""" + import onnx + from onnx import helper, numpy_helper + + model = onnx.load(str(model_input)) + initializers = {initializer.name: initializer for initializer in model.graph.initializer} + replacement_initializers = [] + quantized_weight_names = set() + dequant_nodes = {} + graph_nodes = list(model.graph.node) + for node in graph_nodes: + if node.op_type != "Conv" or len(node.input) < 2: + continue + weight_name = node.input[1] + initializer = initializers.get(weight_name) + if initializer is None: + raise ValueError(f"ICLabel Conv weight {weight_name!r} is not an initializer") + weights = numpy_helper.to_array(initializer).astype(np.float32, copy=False) + axes = tuple(range(1, weights.ndim)) + scale = np.max(np.abs(weights), axis=axes) / 127.0 + scale = np.where(scale == 0, 1.0, scale).astype(np.float32) + reshape = (scale.shape[0],) + (1,) * (weights.ndim - 1) + quantized = np.clip(np.rint(weights / scale.reshape(reshape)), -127, 127).astype(np.int8) + int8_name = f"{weight_name}_int8" + scale_name = f"{weight_name}_scale" + zero_point_name = f"{weight_name}_zero_point" + dequant_name = f"{weight_name}_dequantized" + quantized_weight_names.add(weight_name) + replacement_initializers.extend( + [ + numpy_helper.from_array(quantized, name=int8_name), + numpy_helper.from_array(scale, name=scale_name), + numpy_helper.from_array(np.zeros(scale.shape, dtype=np.int8), name=zero_point_name), + ] + ) + dequant_nodes[id(node)] = helper.make_node( + "DequantizeLinear", + [int8_name, scale_name, zero_point_name], + [dequant_name], + name=f"{weight_name}_dequantize", + axis=0, + ) + node.input[1] = dequant_name + + if not dequant_nodes: + raise ValueError("ICLabel model contains no Conv weights to quantize") + kept_initializers = [ + initializer for initializer in model.graph.initializer if initializer.name not in quantized_weight_names + ] + del model.graph.initializer[:] + model.graph.initializer.extend(kept_initializers + replacement_initializers) + nodes = [] + for node in graph_nodes: + nodes.append(dequant_nodes.get(id(node))) + nodes.append(node) + del model.graph.node[:] + model.graph.node.extend(node for node in nodes if node is not None) + onnx.checker.check_model(model) + model_output.parent.mkdir(parents=True, exist_ok=True) + onnx.save(model, str(model_output)) + return model_output + + +class _CalibrationDataReader: + def __init__(self, inputs: Mapping[str, np.ndarray]): + self._inputs = {name: np.asarray(inputs[name], dtype=np.float32) for name in _INPUT_NAMES} + batches = {array.shape[0] for array in self._inputs.values()} + if len(batches) != 1 or not batches: + raise ValueError("Calibration inputs must have matching non-empty batch dimensions") + self._read = False + + def get_next(self) -> dict[str, np.ndarray] | None: + if self._read: + return None + self._read = True + return self._inputs + + +def quantize_calibrated( + model_input: Path, + model_output: Path, + calibration_inputs: Mapping[str, np.ndarray], +) -> Path: + """Write a calibrated int8 Conv QDQ artifact from float feature inputs.""" + from onnxruntime.quantization import CalibrationMethod, QuantFormat, QuantType, quantize_static + + model_output.parent.mkdir(parents=True, exist_ok=True) + quantize_static( + model_input=str(model_input), + model_output=str(model_output), + calibration_data_reader=_CalibrationDataReader(calibration_inputs), + quant_format=QuantFormat.QDQ, + op_types_to_quantize=["Conv"], + per_channel=True, + activation_type=QuantType.QUInt8, + weight_type=QuantType.QInt8, + calibrate_method=CalibrationMethod.MinMax, + ) + return model_output + + +def evaluate_artifacts( + float32_artifact: Path, + evaluation_features: Path, + candidate_artifacts: Mapping[str, Path], + manifest: Path = DEFAULT_FROZEN_MANIFEST, +) -> dict[str, Any]: + """Evaluate candidates against the float32 teacher on the frozen archive.""" + frozen_manifest = load_frozen_manifest(manifest) + features = load_verified_feature_archive(evaluation_features, frozen_manifest, "evaluation") + evaluation_archive = frozen_manifest["evaluation"]["feature_archive"] + teacher = predict_features(float32_artifact, features) + thresholds = np.asarray(DEFAULT_ICFLAG_THRESHOLDS, dtype=float) + report: dict[str, Any] = { + "manifest": str(Path(manifest).name), + "feature_archives": { + "evaluation": { + "path": str(evaluation_archive["path"]), + "sha256": str(evaluation_archive["sha256"]), + "component_count": int(features[0].shape[3]), + } + }, + "class_names": list(CLASS_NAMES), + "thresholds": [[None if np.isnan(value) else float(value) for value in row] for row in thresholds], + "gate": { + "minimum_top1_agreement": MIN_TOP1_AGREEMENT, + "minimum_keep_reject_agreement": MIN_KEEP_REJECT_AGREEMENT, + "maximum_probability_abs_diff": MAX_PROBABILITY_ABS_DIFF, + }, + "quantization": { + "feature_dtype": "float32", + "input_normalization_and_augmentation": "unchanged from iclabel.py", + "output_softmax": "unchanged from the float32 graph", + "weight_only": "per-output-channel int8 Conv weights with float DequantizeLinear before Conv", + "calibrated": "ONNX Runtime static per-channel Conv QDQ with MinMax calibration and float model I/O", + }, + "float32_reference": { + "artifact": Path(float32_artifact).name, + "size_bytes": Path(float32_artifact).stat().st_size, + "sha256": _sha256(float32_artifact), + **compare_predictions(teacher, teacher, thresholds), + }, + "candidates": {}, + } + for name, artifact in candidate_artifacts.items(): + metrics = compare_predictions(teacher, predict_features(artifact, features), thresholds) + report["candidates"][name] = { + "artifact": Path(artifact).name, + "size_bytes": Path(artifact).stat().st_size, + "sha256": _sha256(artifact), + "gate_pass": parity_gate_passes(metrics), + **metrics, + } + report["default_artifact"] = select_default_artifact(report) + return report + + +def main() -> None: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--float32-artifact", type=Path, default=DEFAULT_FLOAT32_ARTIFACT) + parser.add_argument("--manifest", type=Path, default=DEFAULT_FROZEN_MANIFEST) + parser.add_argument("--evaluation-features", type=Path, default=DEFAULT_EVALUATION_FEATURES) + parser.add_argument("--calibration-features", type=Path, default=DEFAULT_CALIBRATION_FEATURES) + parser.add_argument("--output-dir", type=Path, default=DEFAULT_ARTIFACT_DIR) + parser.add_argument("--report", type=Path, default=DEFAULT_REPORT) + args = parser.parse_args() + + frozen_manifest = load_frozen_manifest(args.manifest) + calibration_features = load_verified_feature_archive(args.calibration_features, frozen_manifest, "calibration") + args.output_dir.mkdir(parents=True, exist_ok=True) + weight_only = quantize_weight_only(args.float32_artifact, args.output_dir / DEFAULT_WEIGHT_ONLY_ARTIFACT.name) + calibrated = quantize_calibrated( + args.float32_artifact, + args.output_dir / DEFAULT_CALIBRATED_ARTIFACT.name, + network_inputs_from_features(calibration_features), + ) + report = evaluate_artifacts( + args.float32_artifact, + args.evaluation_features, + {"weight_only": weight_only, "calibrated": calibrated}, + args.manifest, + ) + calibration_archive = frozen_manifest["calibration"]["feature_archive"] + report["feature_archives"]["calibration"] = { + "path": str(calibration_archive["path"]), + "sha256": str(calibration_archive["sha256"]), + "component_count": int(calibration_features[0].shape[3]), + } + args.report.parent.mkdir(parents=True, exist_ok=True) + with args.report.open("w", encoding="utf-8") as handle: + json.dump(report, handle, indent=2, sort_keys=True) + handle.write("\n") + print(json.dumps(report, indent=2, sort_keys=True)) + + +if __name__ == "__main__": + main() diff --git a/tools/pyodide/benchmark.py b/tools/pyodide/benchmark.py new file mode 100644 index 00000000..ac0755c1 --- /dev/null +++ b/tools/pyodide/benchmark.py @@ -0,0 +1,264 @@ +"""Run the deterministic Phase 2 ICA and matrix-multiplication benchmarks.""" + +from __future__ import annotations + +import argparse +import json +import statistics +import time +import warnings +from collections.abc import Callable +from typing import Any + +import numpy as np +from scipy.linalg import blas as scipy_blas +from threadpoolctl import threadpool_limits + +from eegprep.functions.sigprocfunc.runica import runica +from eegprep.functions.sigprocfunc.runica_matmul import BACKEND as RUNICA_MATMUL_BACKEND + + +SEED = 375 +CHANNELS = 64 +FRAMES = 15_000 +MAX_ICA_ITERATIONS = 512 +WARMUP_RUNS = 1 +MEASURED_RUNS = 3 +RUNICA_BLOCK = 49 +MATMUL_REPETITIONS = 3 +THREAD_COUNT = 1 + +# Pyodide 0.29.5 has no pthread support. Keep the native comparison on the +# same one-thread budget; browser concurrency belongs at the Web Worker/job +# level rather than inside one ICA call. + +# These are the two products called out in runica's training loop for 64 +# channels and the default block heuristic at 15,000 frames: the activation +# product and the final square weight-update product. +MATMUL_CASES = ( + { + "name": "activation", + "left_shape": (CHANNELS, CHANNELS), + "right_shape": (CHANNELS, RUNICA_BLOCK), + }, + { + "name": "weight_update", + "left_shape": (CHANNELS, CHANNELS), + "right_shape": (CHANNELS, CHANNELS), + }, +) + + +def benchmark_record( + *, + algorithm: str, + platform: str, + shape: tuple[int, int], + seed: int, + seconds: float, + iterations: int | None, + converged: bool | None, + backend: str, + error: str | None = None, +) -> dict[str, Any]: + """Return one JSON-serializable benchmark measurement.""" + return { + "algorithm": algorithm, + "platform": platform, + "shape": list(shape), + "seed": seed, + "seconds": float(seconds), + "iterations": None if iterations is None else int(iterations), + "converged": converged, + "backend": backend, + "error": error, + } + + +def _median(values: list[float]) -> float: + if not values: + raise ValueError("Cannot calculate a median from no measurements") + return float(statistics.median(values)) + + +def _runica_once(data: np.ndarray) -> tuple[float, int, bool]: + result: tuple[Any, ...] + start = time.perf_counter() + result = runica( + data.copy(), + extended=1, + maxsteps=MAX_ICA_ITERATIONS, + verbose=False, + rndreset="off", + ) + seconds = time.perf_counter() - start + iterations = len(result[5]) + return seconds, int(iterations), int(iterations) < MAX_ICA_ITERATIONS + + +def _run_picard_once(data: np.ndarray) -> tuple[float, int, bool]: + # Keep this call aligned with eeg_picard.py. return_n_iter is the underlying + # Picard API's telemetry and does not change production behavior. + from picard import picard + + start = time.perf_counter() + with warnings.catch_warnings(record=True) as caught_warnings: + warnings.simplefilter("always") + _weighting, _unmixing, _sources, iterations = picard( + data.copy(), + ortho=False, + fun="tanh", + verbose=False, + m=10, + max_iter=MAX_ICA_ITERATIONS, + tol=1e-7, + centering=True, + whiten=True, + w_init=np.eye(CHANNELS), + return_n_iter=True, + random_state=SEED, + ) + seconds = time.perf_counter() - start + did_not_converge = any("did not converge" in str(item.message).lower() for item in caught_warnings) + return seconds, int(iterations), not did_not_converge and int(iterations) < MAX_ICA_ITERATIONS + + +def _ica_summary( + *, platform: str, data: np.ndarray, name: str, operation: Callable[[np.ndarray], tuple[float, int, bool]] +) -> dict[str, Any]: + for _ in range(WARMUP_RUNS): + operation(data) + + records = [] + for _ in range(MEASURED_RUNS): + try: + seconds, iterations, converged = operation(data) + records.append( + benchmark_record( + algorithm=name, + platform=platform, + shape=(CHANNELS, FRAMES), + seed=SEED, + seconds=seconds, + iterations=iterations, + converged=converged, + backend=RUNICA_MATMUL_BACKEND if name == "runica" else "picard", + ) + ) + except Exception as exc: # Keep the raw failure visible before failing the gate. + records.append( + benchmark_record( + algorithm=name, + platform=platform, + shape=(CHANNELS, FRAMES), + seed=SEED, + seconds=0.0, + iterations=None, + converged=None, + backend=RUNICA_MATMUL_BACKEND if name == "runica" else "picard", + error=f"{type(exc).__name__}: {exc}", + ) + ) + + successful = [record for record in records if record["error"] is None] + return { + "algorithm": name, + "backend": RUNICA_MATMUL_BACKEND if name == "runica" else "picard", + "warmup_runs": WARMUP_RUNS, + "measured_runs": MEASURED_RUNS, + "runs": records, + "median_seconds": _median([float(record["seconds"]) for record in successful]) if successful else None, + "median_iterations": _median([float(record["iterations"]) for record in successful]) if successful else None, + "all_converged": bool(successful) and all(record["converged"] for record in successful), + } + + +def _matmul_operation(name: str, left: np.ndarray, right: np.ndarray) -> Callable[[], np.ndarray]: + if name == "numpy_matmul": + return lambda: left @ right + if name in ("dgemm", "sgemm"): + blas_operation = getattr(scipy_blas, name) + return lambda: blas_operation(alpha=1.0, a=left, b=right) + raise ValueError(f"Unknown matrix multiplication operation: {name}") + + +def _matmul_summary(*, platform: str, rng: np.random.RandomState) -> list[dict[str, Any]]: + summaries = [] + for case in MATMUL_CASES: + for dtype, blas_name in ((np.float64, "dgemm"), (np.float32, "sgemm")): + left = rng.standard_normal(case["left_shape"]).astype(dtype) + right = rng.standard_normal(case["right_shape"]).astype(dtype) + for operation_name in ("numpy_matmul", blas_name): + operation = _matmul_operation(operation_name, left, right) + measurements = [] + checksum = 0.0 + with np.errstate(divide="ignore", over="ignore", invalid="ignore"): + for _ in range(WARMUP_RUNS): + operation() + for _ in range(MEASURED_RUNS): + start = time.perf_counter() + for _ in range(MATMUL_REPETITIONS): + output = operation() + checksum += float(output.flat[0]) + measurements.append(time.perf_counter() - start) + summaries.append( + { + "case": case["name"], + "left_shape": list(case["left_shape"]), + "right_shape": list(case["right_shape"]), + "dtype": np.dtype(dtype).name, + "operation": operation_name, + "platform": platform, + "seed": SEED, + "warmup_runs": WARMUP_RUNS, + "measured_runs": MEASURED_RUNS, + "repetitions": MATMUL_REPETITIONS, + "seconds": measurements, + "median_seconds": _median(measurements), + "checksum": checksum, + } + ) + return summaries + + +def run_benchmark(platform: str) -> dict[str, Any]: + """Run all Phase 2 measurements on ``platform`` and return the report.""" + rng = np.random.RandomState(SEED) + data = rng.standard_normal((CHANNELS, FRAMES)).astype(np.float64) + with threadpool_limits(limits=THREAD_COUNT): + ica = { + "runica": _ica_summary(platform=platform, data=data, name="runica", operation=_runica_once), + "picard": _ica_summary(platform=platform, data=data, name="picard", operation=_run_picard_once), + } + matmul = _matmul_summary(platform=platform, rng=rng) + + report = { + "schema_version": 1, + "platform": platform, + "seed": SEED, + "shape": [CHANNELS, FRAMES], + "max_ica_iterations": MAX_ICA_ITERATIONS, + "thread_count": THREAD_COUNT, + "threading_mode": "single-threaded", + "runica_block": RUNICA_BLOCK, + "warmup_runs": WARMUP_RUNS, + "measured_runs": MEASURED_RUNS, + "ica": ica, + "matmul": matmul, + } + errors = [record["error"] for summary in ica.values() for record in summary["runs"] if record["error"] is not None] + if errors: + raise RuntimeError("Benchmark operation failed: " + "; ".join(errors)) + return report + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--platform", choices=("native", "pyodide"), required=True) + args = parser.parse_args() + print(json.dumps(run_benchmark(args.platform), sort_keys=True)) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tools/pyodide/compare_benchmarks.py b/tools/pyodide/compare_benchmarks.py new file mode 100644 index 00000000..7a6c3518 --- /dev/null +++ b/tools/pyodide/compare_benchmarks.py @@ -0,0 +1,175 @@ +"""Compare native and Pyodide Phase 2 benchmark reports.""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path +from typing import Any + + +BLAS_SPEEDUP_GATE = 1.2 +REQUIRED_ICA_ALGORITHMS = ("runica", "picard") +REQUIRED_MATMUL_CASES = ("activation", "weight_update") + + +def _load_report(path: Path) -> dict[str, Any]: + report = json.loads(path.read_text()) + if report.get("schema_version") != 1: + raise ValueError(f"Unsupported benchmark schema in {path}") + if tuple(report.get("shape", ())) != (64, 15_000): + raise ValueError(f"Unexpected benchmark shape in {path}: {report.get('shape')}") + return report + + +def _matmul_lookup(report: dict[str, Any], case: str, dtype: str, operation: str) -> dict[str, Any]: + matches = [ + item + for item in report["matmul"] + if item["case"] == case and item["dtype"] == dtype and item["operation"] == operation + ] + if len(matches) != 1: + raise ValueError(f"Expected one {operation} result for {case}/{dtype}, found {len(matches)}") + return matches[0] + + +def compare_reports(native: dict[str, Any], pyodide: dict[str, Any]) -> dict[str, Any]: + """Return the Phase 2 decisions and ratios for two valid reports.""" + if native.get("platform") != "native" or pyodide.get("platform") != "pyodide": + raise ValueError("Reports must be labeled native and pyodide") + if ( + native["seed"] != pyodide["seed"] + or native["runica_block"] != pyodide["runica_block"] + or native["thread_count"] != pyodide["thread_count"] + ): + raise ValueError("Native and Pyodide reports do not use the same benchmark configuration") + + ica = {} + for algorithm in REQUIRED_ICA_ALGORITHMS: + native_summary = native["ica"][algorithm] + pyodide_summary = pyodide["ica"][algorithm] + if native_summary["median_seconds"] is None or pyodide_summary["median_seconds"] is None: + raise ValueError(f"Missing successful {algorithm} timing") + ica[algorithm] = { + "native_median_seconds": native_summary["median_seconds"], + "native_median_iterations": native_summary["median_iterations"], + "native_all_converged": native_summary["all_converged"], + "pyodide_median_seconds": pyodide_summary["median_seconds"], + "pyodide_median_iterations": pyodide_summary["median_iterations"], + "pyodide_all_converged": pyodide_summary["all_converged"], + "pyodide_speed_ratio": pyodide_summary["median_seconds"] / native_summary["median_seconds"], + } + + picard = ica["picard"] + runica_result = ica["runica"] + picard_browser_default_retained = bool( + picard["pyodide_all_converged"] + and runica_result["pyodide_median_iterations"] is not None + and picard["pyodide_median_iterations"] is not None + and picard["pyodide_median_iterations"] < runica_result["pyodide_median_iterations"] + ) + + matmul = [] + phase3_gate_results = [] + for case in REQUIRED_MATMUL_CASES: + for dtype, blas_name in (("float64", "dgemm"), ("float32", "sgemm")): + native_numpy_result = _matmul_lookup(native, case, dtype, "numpy_matmul") + native_blas_result = _matmul_lookup(native, case, dtype, blas_name) + numpy_result = _matmul_lookup(pyodide, case, dtype, "numpy_matmul") + blas_result = _matmul_lookup(pyodide, case, dtype, blas_name) + native_speedup = native_numpy_result["median_seconds"] / native_blas_result["median_seconds"] + speedup = numpy_result["median_seconds"] / blas_result["median_seconds"] + phase3_gate_results.append(speedup >= BLAS_SPEEDUP_GATE) + matmul.append( + { + "case": case, + "dtype": dtype, + "native_numpy_median_seconds": native_numpy_result["median_seconds"], + "native_blas": blas_name, + "native_blas_median_seconds": native_blas_result["median_seconds"], + "native_blas_speedup_over_numpy": native_speedup, + "numpy_median_seconds": numpy_result["median_seconds"], + "blas": blas_name, + "blas_median_seconds": blas_result["median_seconds"], + "blas_speedup_over_numpy": speedup, + } + ) + + return { + "schema_version": 1, + "shape": pyodide["shape"], + "seed": pyodide["seed"], + "ica": ica, + "decisions": { + "picard_browser_default_retained": picard_browser_default_retained, + "phase3_recommended": all(phase3_gate_results), + "phase3_gate": f"scipy BLAS >= {BLAS_SPEEDUP_GATE:.1f}x faster than NumPy @ for both Pyodide runica shapes and dtypes", + }, + "matmul": matmul, + } + + +def _markdown(comparison: dict[str, Any]) -> str: + lines = [ + "# Phase 2 Pyodide benchmark comparison", + "", + f"Input: `{comparison['shape'][0]} x {comparison['shape'][1]}`, seed `{comparison['seed']}`.", + "", + "## ICA", + "", + "| Algorithm | Native median (s) | Native median iterations | Native converged | Pyodide median (s) | Pyodide median iterations | Pyodide converged |", + "| --- | ---: | ---: | :---: | ---: | ---: | :---: |", + ] + for algorithm, result in comparison["ica"].items(): + lines.append( + f"| {algorithm} | {result['native_median_seconds']:.6f} | {result['native_median_iterations']:.1f} | " + f"{result['native_all_converged']} | {result['pyodide_median_seconds']:.6f} | " + f"{result['pyodide_median_iterations']:.1f} | {result['pyodide_all_converged']} |" + ) + lines.extend( + [ + "", + "## Matrix multiplication", + "", + "| runica product | dtype | Native NumPy @ (s) | Native BLAS (s) | Native speedup | Pyodide NumPy @ (s) | Pyodide BLAS (s) | Pyodide speedup |", + "| --- | --- | ---: | ---: | ---: | ---: | ---: | ---: |", + ] + ) + for result in comparison["matmul"]: + lines.append( + f"| {result['case']} | {result['dtype']} | {result['native_numpy_median_seconds']:.6f} | " + f"{result['native_blas_median_seconds']:.6f} | {result['native_blas_speedup_over_numpy']:.2f}x | " + f"{result['numpy_median_seconds']:.6f} | {result['blas_median_seconds']:.6f} | " + f"{result['blas_speedup_over_numpy']:.2f}x |" + ) + decisions = comparison["decisions"] + lines.extend( + [ + "", + "## Decisions", + "", + f"- Retain Picard as the browser default: `{decisions['picard_browser_default_retained']}`.", + f"- Recommend Phase 3 BLAS work: `{decisions['phase3_recommended']}`.", + f"- Gate: {decisions['phase3_gate']}.", + "", + ] + ) + return "\n".join(lines) + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--native", type=Path, required=True) + parser.add_argument("--pyodide", type=Path, required=True) + parser.add_argument("--report", type=Path) + args = parser.parse_args() + comparison = compare_reports(_load_report(args.native), _load_report(args.pyodide)) + output = json.dumps(comparison, indent=2, sort_keys=True) + print(output) + if args.report: + args.report.write_text(_markdown(comparison) + "\n") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tools/pyodide/compare_iclabel.py b/tools/pyodide/compare_iclabel.py new file mode 100644 index 00000000..51bae037 --- /dev/null +++ b/tools/pyodide/compare_iclabel.py @@ -0,0 +1,90 @@ +"""Compare native and ONNX Runtime Web ICLabel classification reports.""" + +from __future__ import annotations + +import argparse +import json +from pathlib import Path +from typing import Any + +import numpy as np + + +RTOL = 1e-4 +ATOL = 1e-5 + + +def _load_report(path: Path) -> dict[str, Any]: + report = json.loads(path.read_text()) + if report.get("schema_version") != 1: + raise ValueError(f"Unsupported ICLabel parity schema in {path}") + if report.get("platform") not in {"native", "pyodide"}: + raise ValueError(f"Invalid ICLabel parity platform in {path}") + classifications = np.asarray(report.get("classifications"), dtype=np.float32) + if classifications.ndim != 2 or classifications.shape[1] != 7: + raise ValueError(f"Unexpected ICLabel classification shape in {path}: {classifications.shape}") + if list(classifications.shape) != report.get("shape"): + raise ValueError(f"ICLabel report shape does not match classifications in {path}") + if not np.isfinite(classifications).all(): + raise ValueError(f"ICLabel report contains non-finite values in {path}") + return report + + +def compare_reports(native: dict[str, Any], pyodide: dict[str, Any]) -> dict[str, Any]: + """Return numeric parity metrics and the Phase 5 pass/fail decision.""" + if native.get("platform") != "native" or pyodide.get("platform") != "pyodide": + raise ValueError("Reports must be labeled native and pyodide") + if native.get("dataset") != pyodide.get("dataset"): + raise ValueError("Native and Pyodide reports use different datasets") + + native_values = np.asarray(native["classifications"], dtype=np.float32) + pyodide_values = np.asarray(pyodide["classifications"], dtype=np.float32) + if native_values.shape != pyodide_values.shape: + raise ValueError(f"Classification shapes differ: {native_values.shape} != {pyodide_values.shape}") + difference = np.abs(native_values - pyodide_values) + scale = np.maximum(np.maximum(np.abs(native_values), np.abs(pyodide_values)), np.float32(1e-10)) + max_absolute = float(np.max(difference, initial=0.0)) + max_relative = float(np.max(difference / scale, initial=0.0)) + return { + "schema_version": 1, + "dataset": native["dataset"], + "shape": list(native_values.shape), + "rtol": RTOL, + "atol": ATOL, + "max_absolute_difference": max_absolute, + "max_relative_difference": max_relative, + "allclose": bool(np.allclose(native_values, pyodide_values, rtol=RTOL, atol=ATOL)), + } + + +def _markdown(comparison: dict[str, Any]) -> str: + return "\n".join( + ( + "# Phase 5 ICLabel browser parity", + "", + f"Dataset: `{comparison['dataset']}`; shape: `{comparison['shape']}`.", + "", + f"- `allclose`: `{comparison['allclose']}`", + f"- maximum absolute difference: `{comparison['max_absolute_difference']:.3e}`", + f"- maximum relative difference: `{comparison['max_relative_difference']:.3e}`", + f"- tolerance: `rtol={comparison['rtol']}`, `atol={comparison['atol']}`", + "", + ) + ) + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--native", type=Path, required=True) + parser.add_argument("--pyodide", type=Path, required=True) + parser.add_argument("--report", type=Path) + args = parser.parse_args() + comparison = compare_reports(_load_report(args.native), _load_report(args.pyodide)) + print(json.dumps(comparison, indent=2, sort_keys=True)) + if args.report: + args.report.write_text(_markdown(comparison)) + return 0 if comparison["allclose"] else 1 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tools/pyodide/iclabel_parity.py b/tools/pyodide/iclabel_parity.py new file mode 100644 index 00000000..820d105e --- /dev/null +++ b/tools/pyodide/iclabel_parity.py @@ -0,0 +1,87 @@ +"""Run ICLabel on the checked-in ICA sample data for native/browser parity.""" + +from __future__ import annotations + +import argparse +import asyncio +import json +import os +from pathlib import Path +import sys + +import numpy as np + +from eegprep.functions.popfunc.pop_loadset import pop_loadset + + +DATASET_NAME = "eeglab_data_with_ica_tmp.set" +CLASS_COUNT = 7 + + +async def _classifications(eeg: dict, platform: str) -> np.ndarray: + if platform == "native": + from eegprep.plugins.ICLabel.iclabel import iclabel + + classified = iclabel(eeg) + else: + from eegprep.plugins.ICLabel.iclabel import iclabel_async + + classified = await iclabel_async(eeg) + classifications = np.asarray( + classified["etc"]["ic_classification"]["ICLabel"]["classifications"], + dtype=np.float32, + ) + if classifications.ndim != 2 or classifications.shape[1] != CLASS_COUNT: + raise ValueError(f"Unexpected ICLabel classification shape: {classifications.shape}") + if not np.isfinite(classifications).all(): + raise ValueError("ICLabel classifications contain non-finite values") + return classifications + + +async def _run_parity(platform: str, sample_data_dir: Path) -> dict: + """Classify the same sample dataset on one platform and return JSON data.""" + dataset_path = sample_data_dir / DATASET_NAME + eeg = pop_loadset(dataset_path) + classifications = await _classifications(eeg, platform) + return { + "schema_version": 1, + "platform": platform, + "dataset": DATASET_NAME, + "shape": list(classifications.shape), + "classifications": classifications.tolist(), + } + + +def run_parity(platform: str, sample_data_dir: Path) -> dict: + """Classify the same sample dataset on one platform and return JSON data.""" + return asyncio.run(_run_parity(platform, sample_data_dir)) + + +def _parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--platform", choices=("native", "pyodide"), required=True) + parser.add_argument( + "--sample-data-dir", + type=Path, + default=Path(os.environ.get("EEGPREP_SAMPLE_DATA", "sample_data")), + ) + return parser.parse_args() + + +async def _async_main() -> int: + args = _parse_args() + print(json.dumps(await _run_parity(args.platform, args.sample_data_dir), sort_keys=True)) + return 0 + + +def main() -> int: + args = _parse_args() + print(json.dumps(run_parity(args.platform, args.sample_data_dir), sort_keys=True)) + return 0 + + +if __name__ == "__main__": + if sys.platform == "emscripten": + __eegprep_async_result__ = _async_main() + else: + raise SystemExit(main()) diff --git a/tools/pyodide/iclabel_web_bridge.mjs b/tools/pyodide/iclabel_web_bridge.mjs new file mode 100644 index 00000000..7f5fba1d --- /dev/null +++ b/tools/pyodide/iclabel_web_bridge.mjs @@ -0,0 +1,35 @@ +/** + * Build the ONNX Runtime Web bridge consumed by the Pyodide ICLabel adapter. + * + * The Python side passes flat Float32Array values and explicit tensor shapes. + * Returning the output typed array from the Promise keeps the Pyodide boundary + * asynchronous without exposing ONNX Runtime objects to Python. + */ +export function createIcLabelWebBridge(ort, modelBytes, options = {}) { + const sessionOptions = { + executionProviders: ["wasm"], + numThreads: 1, + ...options, + }; + let sessionPromise; + + const getSession = () => { + if (sessionPromise === undefined) { + sessionPromise = ort.InferenceSession.create(modelBytes, sessionOptions).catch((error) => { + sessionPromise = undefined; + throw error; + }); + } + return sessionPromise; + }; + + return { + run: (imageData, imageShape, psdmedData, psdmedShape, autocorrData, autocorrShape) => + getSession().then((session) => + session.run({ + image: new ort.Tensor("float32", imageData, imageShape), + psdmed: new ort.Tensor("float32", psdmedData, psdmedShape), + autocorr: new ort.Tensor("float32", autocorrData, autocorrShape), + }).then((outputs) => outputs.output.data)), + }; +} diff --git a/tools/pyodide/package-lock.json b/tools/pyodide/package-lock.json new file mode 100644 index 00000000..e8d6fe3c --- /dev/null +++ b/tools/pyodide/package-lock.json @@ -0,0 +1,193 @@ +{ + "name": "eegprep-pyodide-runner", + "lockfileVersion": 3, + "requires": true, + "packages": { + "": { + "name": "eegprep-pyodide-runner", + "dependencies": { + "onnxruntime-web": "1.30.0", + "pyodide": "0.29.5" + } + }, + "node_modules/@protobufjs/aspromise": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/aspromise/-/aspromise-1.1.2.tgz", + "integrity": "sha512-j+gKExEuLmKwvz3OgROXtrJ2UG2x8Ch2YZUxahh+s1F2HZ+wAceUNLkvy6zKCPVRkU++ZWQrdxsUeQXmcg4uoQ==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/base64": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/base64/-/base64-1.1.2.tgz", + "integrity": "sha512-AZkcAA5vnN/v4PDqKyMR5lx7hZttPDgClv83E//FMNhR2TMcLUhfRUBHCmSl0oi9zMgDDqRUJkSxO3wm85+XLg==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/codegen": { + "version": "2.0.5", + "resolved": "https://registry.npmjs.org/@protobufjs/codegen/-/codegen-2.0.5.tgz", + "integrity": "sha512-zgXFLzW3Ap33e6d0Wlj4MGIm6Ce8O89n/apUaGNB/jx+hw+ruWEp7EwGUshdLKVRCxZW12fp9r40E1mQrf/34g==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/eventemitter": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/@protobufjs/eventemitter/-/eventemitter-1.1.1.tgz", + "integrity": "sha512-vW1GmwMZNnL+gMRaovlh9yZX74kc+TTU3FObkkurpMaRtBfLP3ldjS9KQWlwZgraRE0+dheEEoAxdzcJQ8eXZg==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/fetch": { + "version": "1.1.1", + "resolved": "https://registry.npmjs.org/@protobufjs/fetch/-/fetch-1.1.1.tgz", + "integrity": "sha512-GpptLrs57adMSuHi3VNj0mAF8dwh36LMaYF6XyJ6JMWlVsc+t42tm1HSEDmOs3A8fC9yyeisgLhsTVQokOZ0zw==", + "license": "BSD-3-Clause", + "dependencies": { + "@protobufjs/aspromise": "^1.1.1" + } + }, + "node_modules/@protobufjs/float": { + "version": "1.0.2", + "resolved": "https://registry.npmjs.org/@protobufjs/float/-/float-1.0.2.tgz", + "integrity": "sha512-Ddb+kVXlXst9d+R9PfTIxh1EdNkgoRe5tOX6t01f1lYWOvJnSPDBlG241QLzcyPdoNTsblLUdujGSE4RzrTZGQ==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/path": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/path/-/path-1.1.2.tgz", + "integrity": "sha512-6JOcJ5Tm08dOHAbdR3GrvP+yUUfkjG5ePsHYczMFLq3ZmMkAD98cDgcT2iA1lJ9NVwFd4tH/iSSoe44YWkltEA==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/pool": { + "version": "1.1.0", + "resolved": "https://registry.npmjs.org/@protobufjs/pool/-/pool-1.1.0.tgz", + "integrity": "sha512-0kELaGSIDBKvcgS4zkjz1PeddatrjYcmMWOlAuAPwAeccUrPHdUqo/J6LiymHHEiJT5NrF1UVwxY14f+fy4WQw==", + "license": "BSD-3-Clause" + }, + "node_modules/@protobufjs/utf8": { + "version": "1.1.2", + "resolved": "https://registry.npmjs.org/@protobufjs/utf8/-/utf8-1.1.2.tgz", + "integrity": "sha512-b1UQwcEZ4yCnMCD8DAL1VlbvBJE9/IX4FTIp7BG1xYpf29SLazLSrqUkj4w7Y5y7cCVP6E5tcqqcI0xemPkHug==", + "license": "BSD-3-Clause" + }, + "node_modules/@types/emscripten": { + "version": "1.41.6", + "resolved": "https://registry.npmjs.org/@types/emscripten/-/emscripten-1.41.6.tgz", + "integrity": "sha512-uN+9i8bFT5CUcZfyIEYDrSueACEyKGbUs5kC/72DGlZZoinh84sJfVV0i8UOJD1asdzkvLPBRrKs41kZ8MdEXg==", + "license": "MIT" + }, + "node_modules/@types/node": { + "version": "26.6.1", + "resolved": "https://registry.npmjs.org/@types/node/-/node-26.6.1.tgz", + "integrity": "sha512-VqGJBMCtdhqkBUCcBLvywI0NJ+KLuVzgNnlBUNFOQjqVxzo2lxLUNg1DSey8+u2u6ktswSAxg+s68QLzWHNOuA==", + "license": "MIT", + "dependencies": { + "undici-types": "~8.9.0" + } + }, + "node_modules/flatbuffers": { + "version": "25.9.23", + "resolved": "https://registry.npmjs.org/flatbuffers/-/flatbuffers-25.9.23.tgz", + "integrity": "sha512-MI1qs7Lo4Syw0EOzUl0xjs2lsoeqFku44KpngfIduHBYvzm8h2+7K8YMQh1JtVVVrUvhLpNwqVi4DERegUJhPQ==", + "license": "Apache-2.0" + }, + "node_modules/guid-typescript": { + "version": "1.0.9", + "resolved": "https://registry.npmjs.org/guid-typescript/-/guid-typescript-1.0.9.tgz", + "integrity": "sha512-Y8T4vYhEfwJOTbouREvG+3XDsjr8E3kIr7uf+JZ0BYloFsttiHU0WfvANVsR7TxNUJa/WpCnw/Ino/p+DeBhBQ==", + "license": "ISC" + }, + "node_modules/long": { + "version": "5.3.2", + "resolved": "https://registry.npmjs.org/long/-/long-5.3.2.tgz", + "integrity": "sha512-mNAgZ1GmyNhD7AuqnTG3/VQ26o760+ZYBPKjPvugO8+nLbYfX6TVpJPseBvopbdY+qpZ/lKUnmEc1LeZYS3QAA==", + "license": "Apache-2.0" + }, + "node_modules/onnxruntime-common": { + "version": "1.30.0", + "resolved": "https://registry.npmjs.org/onnxruntime-common/-/onnxruntime-common-1.30.0.tgz", + "integrity": "sha512-7fdVWjAID1dVhH/G8qK3APARunV4VkBFoCQAP7qp4Wkab0mrorvmc+sqiT+mKXOzDqdjN5j+/Z9nb4gzNPWcyA==", + "license": "MIT" + }, + "node_modules/onnxruntime-web": { + "version": "1.30.0", + "resolved": "https://registry.npmjs.org/onnxruntime-web/-/onnxruntime-web-1.30.0.tgz", + "integrity": "sha512-q0y+JrrtukXSzsBWEMccVfqX25LRmosXHF+CaRJmg8pZClzcV7svNc4rKY3jL02Vb7QmRMDs1SigqR4CXAfKYQ==", + "license": "MIT", + "dependencies": { + "flatbuffers": "^25.1.24", + "guid-typescript": "^1.0.9", + "long": "^5.2.3", + "onnxruntime-common": "1.30.0", + "platform": "^1.3.6", + "protobufjs": "^7.2.4" + } + }, + "node_modules/platform": { + "version": "1.3.6", + "resolved": "https://registry.npmjs.org/platform/-/platform-1.3.6.tgz", + "integrity": "sha512-fnWVljUchTro6RiCFvCXBbNhJc2NijN7oIQxbwsyL0buWJPG85v81ehlHI9fXrJsMNgTofEoWIQeClKpgxFLrg==", + "license": "MIT" + }, + "node_modules/protobufjs": { + "version": "7.6.6", + "resolved": "https://registry.npmjs.org/protobufjs/-/protobufjs-7.6.6.tgz", + "integrity": "sha512-dYDWdjSl5RNb7SgPxGQcRU+GtvP7s2fpkrY0r432PcOIaZ0/rBcxEZnQN67iJhFuQiVw754JDoPruPCNdGsbjg==", + "hasInstallScript": true, + "license": "BSD-3-Clause", + "dependencies": { + "@protobufjs/aspromise": "^1.1.2", + "@protobufjs/base64": "^1.1.2", + "@protobufjs/codegen": "^2.0.5", + "@protobufjs/eventemitter": "^1.1.1", + "@protobufjs/fetch": "^1.1.1", + "@protobufjs/float": "^1.0.2", + "@protobufjs/path": "^1.1.2", + "@protobufjs/pool": "^1.1.0", + "@protobufjs/utf8": "^1.1.1", + "@types/node": ">=13.7.0", + "long": "^5.3.2" + }, + "engines": { + "node": ">=12.0.0" + } + }, + "node_modules/pyodide": { + "version": "0.29.5", + "resolved": "https://registry.npmjs.org/pyodide/-/pyodide-0.29.5.tgz", + "integrity": "sha512-TkYrUv9m8QmfImADKRmpwZulL48E0uUaBSO+dvRQKyAJc0qh9gu7jn8btyHPh3NZm1r2IH2YrJhIr7Q2B9urAw==", + "license": "MPL-2.0", + "dependencies": { + "@types/emscripten": "^1.41.4", + "ws": "^8.5.0" + }, + "engines": { + "node": ">=18.0.0" + } + }, + "node_modules/undici-types": { + "version": "8.9.0", + "resolved": "https://registry.npmjs.org/undici-types/-/undici-types-8.9.0.tgz", + "integrity": "sha512-KTDyRTYX8sWmKXAikPHHSyc63CRPETMctyjKFupcC6OBLXT3xsN0e9aF7m+mIXutFWpUXuedtowG7iLOzp0kQg==", + "license": "MIT" + }, + "node_modules/ws": { + "version": "8.21.3", + "resolved": "https://registry.npmjs.org/ws/-/ws-8.21.3.tgz", + "integrity": "sha512-201TZ/kPWxoPr/OKWjquZR1SWKXcvxdH+e1xrx89b3YbmzLMFCLfnaG1HFIgWzJOEWZ7MvpK++odZufgYR50Rw==", + "license": "MIT", + "engines": { + "node": ">=10.0.0" + }, + "peerDependencies": { + "bufferutil": "^4.0.1", + "utf-8-validate": ">=5.0.2" + }, + "peerDependenciesMeta": { + "bufferutil": { + "optional": true + }, + "utf-8-validate": { + "optional": true + } + } + } + } +} diff --git a/tools/pyodide/package.json b/tools/pyodide/package.json new file mode 100644 index 00000000..a2fc4729 --- /dev/null +++ b/tools/pyodide/package.json @@ -0,0 +1,8 @@ +{ + "name": "eegprep-pyodide-runner", + "private": true, + "dependencies": { + "onnxruntime-web": "1.30.0", + "pyodide": "0.29.5" + } +} diff --git a/tools/pyodide/prepare_docopt_wheel.py b/tools/pyodide/prepare_docopt_wheel.py new file mode 100644 index 00000000..c9187e8c --- /dev/null +++ b/tools/pyodide/prepare_docopt_wheel.py @@ -0,0 +1,108 @@ +"""Build the pure-Python docopt wheel required by the Pyodide harness.""" + +from __future__ import annotations + +import argparse +import hashlib +import shutil +import subprocess +import tarfile +import tempfile +import urllib.request +import zipfile +from pathlib import Path + + +DOCOPT_VERSION = "0.6.2" +DOCOPT_SDIST_URL = ( + "https://files.pythonhosted.org/packages/a2/55/8f8cab2afd404cf578136ef2cc5dfb50baa1761b68c9da1fb1e4eed343c9/" + "docopt-0.6.2.tar.gz" +) +DOCOPT_SDIST_SHA256 = "49b3a825280bd66b3aa83585ef59c4a8c82f2c8a522dbe754a8bc8d08c85c491" +HTTP_TIMEOUT_S = 60 + + +def verify_sha256(path: Path, expected: str) -> None: + """Raise when ``path`` does not match the expected SHA-256 digest.""" + digest = hashlib.sha256() + with path.open("rb") as source: + for chunk in iter(lambda: source.read(1024 * 1024), b""): + digest.update(chunk) + actual = digest.hexdigest() + if actual != expected: + raise ValueError(f"SHA-256 mismatch for {path.name}: expected {expected}, got {actual}") + + +def _download_sdist(destination: Path) -> None: + request = urllib.request.Request(DOCOPT_SDIST_URL, headers={"User-Agent": "eegprep-pyodide-harness"}) + with urllib.request.urlopen(request, timeout=HTTP_TIMEOUT_S) as response, destination.open("wb") as output: + shutil.copyfileobj(response, output) + + +def _extract_sdist(archive_path: Path, destination: Path) -> Path: + destination.mkdir(parents=True, exist_ok=True) + with tarfile.open(archive_path, mode="r:gz") as archive: + for member in archive.getmembers(): + target = (destination / member.name).resolve() + if destination.resolve() not in target.parents: + raise ValueError(f"Unsafe path in docopt archive: {member.name}") + if member.issym() or member.islnk(): + raise ValueError(f"Links are not allowed in docopt archive: {member.name}") + archive.extractall(destination) + + roots = [path for path in destination.iterdir() if path.is_dir() and (path / "setup.py").is_file()] + if len(roots) != 1: + raise RuntimeError(f"Expected one docopt source directory, found {len(roots)}") + return roots[0] + + +def _validate_wheel(wheel_path: Path) -> None: + expected_prefix = f"docopt-{DOCOPT_VERSION}-" + if not wheel_path.name.startswith(expected_prefix) or not wheel_path.name.endswith("-none-any.whl"): + raise RuntimeError(f"Expected a universal docopt wheel, got {wheel_path.name}") + with zipfile.ZipFile(wheel_path) as wheel: + wheel_metadata = [name for name in wheel.namelist() if name.endswith("/WHEEL")] + if len(wheel_metadata) != 1: + raise RuntimeError("Generated docopt wheel has no unique WHEEL metadata file") + metadata = wheel.read(wheel_metadata[0]).decode("utf-8") + if "Root-Is-Purelib: true" not in metadata: + raise RuntimeError("Generated docopt wheel is not marked as pure Python") + + +def build_docopt_wheel(output_dir: Path) -> Path: + """Build and validate ``docopt==0.6.2`` into ``output_dir``.""" + output_dir = output_dir.resolve() + output_dir.mkdir(parents=True, exist_ok=True) + if any(output_dir.iterdir()): + raise ValueError(f"Docopt output directory must be empty: {output_dir}") + + with tempfile.TemporaryDirectory(prefix="eegprep-docopt-") as temporary: + temporary_path = Path(temporary) + archive_path = temporary_path / "docopt-0.6.2.tar.gz" + _download_sdist(archive_path) + verify_sha256(archive_path, DOCOPT_SDIST_SHA256) + source_root = _extract_sdist(archive_path, temporary_path / "source") + subprocess.run( + ["uv", "build", "--wheel", "--out-dir", str(output_dir)], + cwd=source_root, + check=True, + ) + + wheels = sorted(output_dir.glob(f"docopt-{DOCOPT_VERSION}-*.whl")) + if len(wheels) != 1: + raise RuntimeError(f"Expected one generated docopt wheel, found {len(wheels)}") + _validate_wheel(wheels[0]) + return wheels[0] + + +def main() -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--output-dir", type=Path, required=True) + args = parser.parse_args() + wheel = build_docopt_wheel(args.output_dir) + print(f"Built {wheel} from {DOCOPT_SDIST_URL}") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tools/pyodide/run_pyodide.mjs b/tools/pyodide/run_pyodide.mjs new file mode 100755 index 00000000..acdf369d --- /dev/null +++ b/tools/pyodide/run_pyodide.mjs @@ -0,0 +1,208 @@ +#!/usr/bin/env node + +import { readFileSync, writeFileSync } from "node:fs"; +import { basename, dirname } from "node:path"; +import { pathToFileURL } from "node:url"; + +const USAGE = `Usage: run_pyodide.mjs --pyodide-module PATH --wheel PATH --docopt-wheel PATH --script PATH [--sample-data-dir PATH] [--output PATH] [--iclabel-model PATH --onnxruntime-web-module PATH] -- [script args...]`; +const PYODIDE_PACKAGES = [ + "certifi", + "charset-normalizer", + "click", + "contourpy", + "cycler", + "decorator", + "fonttools", + "fsspec", + "h5py", + "idna", + "jinja2", + "joblib", + "kiwisolver", + "lazy-loader", + "markupsafe", + "matplotlib", + "mpmath", + "narwhals", + "numpy", + "packaging", + "pandas", + "pillow", + "platformdirs", + "pyparsing", + "python-dateutil", + "pytz", + "pyyaml", + "requests", + "scikit-learn", + "scipy", + "six", + "sympy", + "threadpoolctl", + "tqdm", + "typing-extensions", + "tzdata", + "urllib3", + "wrapt", +]; +const MICROPIP_CONSTRAINTS = ["mne==1.10.0", "sympy==1.14.0", "threadpoolctl==3.6.0"]; + +function fail(message) { + console.error(`${message}\n${USAGE}`); + process.exitCode = 2; +} + +function parseArguments(argv) { + const separator = argv.indexOf("--"); + const options = separator === -1 ? argv : argv.slice(0, separator); + const scriptArgs = separator === -1 ? [] : argv.slice(separator + 1); + const result = { scriptArgs }; + for (let index = 0; index < options.length; index += 1) { + const option = options[index]; + if ( + ![ + "--pyodide-module", + "--wheel", + "--docopt-wheel", + "--script", + "--sample-data-dir", + "--output", + "--iclabel-model", + "--onnxruntime-web-module", + ].includes(option) + ) { + throw new Error(`Unknown option: ${option}`); + } + const value = options[index + 1]; + if (!value || value.startsWith("--")) { + throw new Error(`Missing value for ${option}`); + } + result[option.slice(2).replaceAll("-", "_")] = value; + index += 1; + } + for (const required of ["pyodide_module", "wheel", "docopt_wheel", "script"]) { + if (!result[required]) { + throw new Error(`Missing required option --${required.replaceAll("_", "-")}`); + } + } + if (Boolean(result.iclabel_model) !== Boolean(result.onnxruntime_web_module)) { + throw new Error("--iclabel-model and --onnxruntime-web-module must be provided together"); + } + return result; +} + +function copyToPyodide(pyodide, hostPath, guestPath) { + pyodide.FS.writeFile(guestPath, readFileSync(hostPath)); +} + +async function main() { + let args; + try { + args = parseArguments(process.argv.slice(2)); + } catch (error) { + fail(error.message); + return; + } + + const pyodideModule = await import(pathToFileURL(args.pyodide_module).href); + const pyodide = await pyodideModule.loadPyodide(); + const stdout = []; + const stderr = []; + pyodide.setStdout({ + batched: (message) => { + stdout.push(message); + process.stdout.write(message); + }, + }); + pyodide.setStderr({ + batched: (message) => { + stderr.push(message); + process.stderr.write(message); + }, + }); + + await pyodide.loadPackage(["micropip", ...PYODIDE_PACKAGES]); + + if (args.iclabel_model) { + const ortModule = await import(pathToFileURL(args.onnxruntime_web_module).href); + const ort = ortModule.default ?? ortModule; + ort.env.wasm.wasmPaths = `${dirname(args.onnxruntime_web_module)}/`; + const bridgeModule = await import(new URL("./iclabel_web_bridge.mjs", import.meta.url).href); + const bridge = bridgeModule.createIcLabelWebBridge(ort, new Uint8Array(readFileSync(args.iclabel_model))); + globalThis.eegprep_iclabel_web = bridge; + pyodide.globals.set("eegprep_iclabel_web", bridge); + } + + const scriptPath = "/tmp/eegprep-script.py"; + copyToPyodide(pyodide, args.script, scriptPath); + + if (args.sample_data_dir) { + const sampleData = [ + "eeglab_data.set", + "eeglab_data.fdt", + "eeglab_data_with_ica_tmp.set", + "eeglab_data_with_ica_tmp.fdt", + ]; + pyodide.FS.mkdirTree("/tmp/eegprep-sample-data"); + for (const name of sampleData) { + copyToPyodide(pyodide, `${args.sample_data_dir}/${name}`, `/tmp/eegprep-sample-data/${name}`); + } + } + + // Keep the original PEP 427 names: micropip uses the filename to identify + // a local wheel. ``emfs:`` makes the transport work through Pyodide's + // virtual filesystem in both the Node harness and a browser worker. + const docoptWheelPath = `/tmp/${basename(args.docopt_wheel)}`; + const eegprepWheelPath = `/tmp/${basename(args.wheel)}`; + copyToPyodide(pyodide, args.docopt_wheel, docoptWheelPath); + copyToPyodide(pyodide, args.wheel, eegprepWheelPath); + + const installCode = ` +import micropip +await micropip.install("threadpoolctl==3.6.0", reinstall=True) +await micropip.install("sympy==1.14.0", reinstall=True) +await micropip.install(${JSON.stringify(`emfs:${docoptWheelPath}`)}) +await micropip.install( + ${JSON.stringify(`emfs:${eegprepWheelPath}`)}, + constraints=${JSON.stringify(MICROPIP_CONSTRAINTS)}, +) +`; + await pyodide.runPythonAsync(installCode); + + // Package loading and installation are streamed for diagnostics, but a + // requested output file must contain only the target script's JSON/text. + stdout.length = 0; + stderr.length = 0; + + const scriptCode = ` +import os +import runpy +import inspect +import sys +${args.sample_data_dir ? 'os.environ["EEGPREP_SAMPLE_DATA"] = "/tmp/eegprep-sample-data"' : ""} +sys.argv = ${JSON.stringify([scriptPath, ...args.scriptArgs])} +try: + script_globals = runpy.run_path(${JSON.stringify(scriptPath)}, run_name="__main__") +except SystemExit as exc: + if exc.code not in (None, 0): + raise +else: + async_result = script_globals.get("__eegprep_async_result__") + if inspect.isawaitable(async_result): + result = await async_result + if result not in (None, 0): + raise SystemExit(result) +`; + try { + await pyodide.runPythonAsync(scriptCode); + } finally { + if (args.output) { + writeFileSync(args.output, stdout.join("")); + } + } +} + +main().catch((error) => { + console.error(error?.stack ?? error); + process.exitCode = 1; +}); diff --git a/tools/pyodide/run_pyodide.sh b/tools/pyodide/run_pyodide.sh new file mode 100755 index 00000000..d528c4f1 --- /dev/null +++ b/tools/pyodide/run_pyodide.sh @@ -0,0 +1,118 @@ +#!/usr/bin/env bash +set -euo pipefail + +SCRIPT_DIR=$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd) +REPO_ROOT=$(cd -- "$SCRIPT_DIR/../.." && pwd) +WORKTREE_TMP=$(mktemp -d "${TMPDIR:-/tmp}/eegprep-pyodide.XXXXXX") +trap 'rm -rf "$WORKTREE_TMP"' EXIT + +usage() { + echo "Usage: run_pyodide.sh --wheel PATH --script PATH [--docopt-wheel PATH] [--sample-data-dir PATH] [--output PATH] [--iclabel-web] -- [script args...]" >&2 +} + +WHEEL_PATH="" +DOCOPT_WHEEL_PATH="" +SCRIPT_PATH="" +SAMPLE_DATA_DIR="" +OUTPUT_PATH="" +ICLABEL_WEB=false + +while (($# > 0)); do + case "$1" in + --wheel) + WHEEL_PATH=${2:-} + shift 2 + ;; + --docopt-wheel) + DOCOPT_WHEEL_PATH=${2:-} + shift 2 + ;; + --script) + SCRIPT_PATH=${2:-} + shift 2 + ;; + --sample-data-dir) + SAMPLE_DATA_DIR=${2:-} + shift 2 + ;; + --output) + OUTPUT_PATH=${2:-} + shift 2 + ;; + --iclabel-web) + ICLABEL_WEB=true + shift + ;; + --) + shift + break + ;; + *) + usage + exit 2 + ;; + esac +done + +if [[ -z "$WHEEL_PATH" || -z "$SCRIPT_PATH" ]]; then + usage + exit 2 +fi +if [[ ! -f "$WHEEL_PATH" || ! -f "$SCRIPT_PATH" ]]; then + echo "Wheel and script paths must name existing files" >&2 + exit 2 +fi + +if [[ -z "$DOCOPT_WHEEL_PATH" ]]; then + DOCOPT_DIR="$WORKTREE_TMP/docopt" + uv run --no-sync python "$REPO_ROOT/tools/pyodide/prepare_docopt_wheel.py" --output-dir "$DOCOPT_DIR" + DOCOPT_WHEEL_PATH=$(find "$DOCOPT_DIR" -maxdepth 1 -type f -name 'docopt-0.6.2-*.whl' -print -quit) +fi +if [[ ! -f "$DOCOPT_WHEEL_PATH" ]]; then + echo "docopt wheel path must name an existing file" >&2 + exit 2 +fi + +NPM_PREFIX="$WORKTREE_TMP/npm" +mkdir -p "$NPM_PREFIX" +cp "$REPO_ROOT/tools/pyodide/package.json" "$NPM_PREFIX/package.json" +cp "$REPO_ROOT/tools/pyodide/package-lock.json" "$NPM_PREFIX/package-lock.json" +npm ci --ignore-scripts --no-audit --no-fund --prefix "$NPM_PREFIX" +PYODIDE_MODULE="$NPM_PREFIX/node_modules/pyodide/pyodide.mjs" + +NODE_ARGS=( + "$SCRIPT_DIR/run_pyodide.mjs" + --pyodide-module "$PYODIDE_MODULE" + --wheel "$WHEEL_PATH" + --docopt-wheel "$DOCOPT_WHEEL_PATH" + --script "$SCRIPT_PATH" +) +if [[ "$ICLABEL_WEB" == true ]]; then + ICLABEL_MODEL_PATH="$WORKTREE_TMP/iclabel.onnx" + uv run --no-sync python - "$WHEEL_PATH" "$ICLABEL_MODEL_PATH" <<'PY' +import sys +import zipfile + +wheel_path, output_path = sys.argv[1:] +model_name = "eegprep/plugins/ICLabel/iclabel.onnx" +with zipfile.ZipFile(wheel_path) as wheel: + try: + model = wheel.read(model_name) + except KeyError as exc: + raise SystemExit(f"Wheel is missing {model_name}") from exc +with open(output_path, "wb") as handle: + handle.write(model) +PY + NODE_ARGS+=( + --iclabel-model "$ICLABEL_MODEL_PATH" + --onnxruntime-web-module "$NPM_PREFIX/node_modules/onnxruntime-web/dist/ort.bundle.min.mjs" + ) +fi +if [[ -n "$SAMPLE_DATA_DIR" ]]; then + NODE_ARGS+=(--sample-data-dir "$SAMPLE_DATA_DIR") +fi +if [[ -n "$OUTPUT_PATH" ]]; then + NODE_ARGS+=(--output "$OUTPUT_PATH") +fi + +node "${NODE_ARGS[@]}" -- "$@" diff --git a/tools/pyodide/smoke.py b/tools/pyodide/smoke.py new file mode 100644 index 00000000..aa0be25d --- /dev/null +++ b/tools/pyodide/smoke.py @@ -0,0 +1,43 @@ +"""Run a small continuous EEGPrep pipeline against the checked-in sample data.""" + +from __future__ import annotations + +import json +import os +from pathlib import Path + +import numpy as np + +from eegprep.functions.adminfunc.eeg_checkset import eeg_checkset +from eegprep.functions.popfunc.pop_loadset import pop_loadset +from eegprep.functions.popfunc.pop_reref import pop_reref + + +def main() -> int: + sample_dir = Path(os.environ["EEGPREP_SAMPLE_DATA"]) + dataset_path = sample_dir / "eeglab_data.set" + eeg = pop_loadset(dataset_path) + eeg = eeg_checkset(eeg) + rereferenced, command = pop_reref(eeg, [], return_com=True) + data = np.asarray(rereferenced["data"]) + if data.shape != (int(eeg["nbchan"]), int(eeg["pnts"])): + raise AssertionError(f"Unexpected continuous data shape after rereferencing: {data.shape}") + if not np.isfinite(data).all(): + raise AssertionError("Rereferenced sample data contains non-finite values") + print( + json.dumps( + { + "dataset": dataset_path.name, + "nbchan": int(rereferenced["nbchan"]), + "pnts": int(rereferenced["pnts"]), + "trials": int(rereferenced["trials"]), + "command": command, + }, + sort_keys=True, + ) + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/uv.lock b/uv.lock index 56004772..b11115f9 100644 --- a/uv.lock +++ b/uv.lock @@ -2,8 +2,10 @@ version = 1 revision = 3 requires-python = ">=3.12" resolution-markers = [ - "python_full_version >= '3.13' and sys_platform != 'darwin'", - "python_full_version >= '3.13' and sys_platform == 'darwin'", + "python_full_version >= '3.14' and sys_platform != 'darwin'", + "python_full_version == '3.13.*' and sys_platform != 'darwin'", + "python_full_version >= '3.14' and sys_platform == 'darwin'", + "python_full_version == '3.13.*' and sys_platform == 'darwin'", "python_full_version < '3.13' and sys_platform != 'darwin'", "python_full_version < '3.13' and sys_platform == 'darwin'", ] @@ -525,11 +527,8 @@ dependencies = [ { name = "mne" }, { name = "neo" }, { name = "numpy" }, - { name = "oct2py" }, { name = "packaging" }, - { name = "psutil" }, { name = "pybids" }, - { name = "pyedflib" }, { name = "python-picard" }, { name = "pyyaml" }, { name = "scipy", version = "1.15.3", source = { registry = "https://pypi.org/simple" }, marker = "sys_platform != 'darwin'" }, @@ -542,6 +541,10 @@ all = [ { name = "ipython" }, { name = "myst-parser" }, { name = "numpydoc" }, + { name = "oct2py" }, + { name = "onnx" }, + { name = "onnxruntime" }, + { name = "psutil" }, { name = "pydata-sphinx-theme" }, { name = "pyqtgraph" }, { name = "pyside6" }, @@ -569,17 +572,28 @@ docs = [ { name = "sphinx-togglebutton" }, { name = "sphinxcontrib-spelling" }, ] +eeglab = [ + { name = "oct2py" }, +] gui = [ { name = "pyqtgraph" }, { name = "pyside6" }, ] +iclabel = [ + { name = "onnxruntime" }, +] +sys = [ + { name = "psutil" }, +] torch = [ + { name = "onnx" }, { name = "torch" }, ] [package.dev-dependencies] dev = [ { name = "ipython" }, + { name = "pyedflib" }, { name = "pyqtgraph" }, { name = "pyside6" }, { name = "pytest" }, @@ -597,8 +611,11 @@ requires-dist = [ { name = "eeglabio", specifier = ">=0.1.2" }, { name = "eegprep", extras = ["console"], marker = "extra == 'all'" }, { name = "eegprep", extras = ["docs"], marker = "extra == 'all'" }, + { name = "eegprep", extras = ["eeglab"], marker = "extra == 'all'" }, { name = "eegprep", extras = ["gui"], marker = "extra == 'all'" }, { name = "eegprep", extras = ["gui"], marker = "extra == 'console'" }, + { name = "eegprep", extras = ["iclabel"], marker = "extra == 'all'" }, + { name = "eegprep", extras = ["sys"], marker = "extra == 'all'" }, { name = "eegprep", extras = ["torch"], marker = "extra == 'all'" }, { name = "h5py", specifier = ">=3.12.1" }, { name = "ipython", marker = "extra == 'console'", specifier = ">=8.0" }, @@ -608,12 +625,13 @@ requires-dist = [ { name = "neo", specifier = ">=0.14.2" }, { name = "numpy", specifier = ">=2.1.0" }, { name = "numpydoc", marker = "extra == 'docs'", specifier = ">=1.6.0" }, - { name = "oct2py", specifier = ">=5.5.0" }, + { name = "oct2py", marker = "extra == 'eeglab'", specifier = ">=5.5.0" }, + { name = "onnx", marker = "extra == 'torch'", specifier = ">=1.14" }, + { name = "onnxruntime", marker = "extra == 'iclabel'", specifier = ">=1.18" }, { name = "packaging", specifier = ">=23.0" }, - { name = "psutil", specifier = ">=7.0.0" }, + { name = "psutil", marker = "extra == 'sys'", specifier = ">=7.0.0" }, { name = "pybids", specifier = ">=0.4" }, { name = "pydata-sphinx-theme", marker = "extra == 'docs'", specifier = ">=0.14.0" }, - { name = "pyedflib", specifier = ">=0.1.42" }, { name = "pyqtgraph", marker = "extra == 'gui'", specifier = ">=0.13.7" }, { name = "pyside6", marker = "extra == 'gui'", specifier = ">=6.6" }, { name = "python-picard", specifier = ">=0.8,<0.9" }, @@ -629,11 +647,12 @@ requires-dist = [ { name = "threadpoolctl", specifier = ">=3.5.0" }, { name = "torch", marker = "extra == 'torch'", specifier = ">=2.0" }, ] -provides-extras = ["torch", "gui", "console", "docs", "all"] +provides-extras = ["torch", "iclabel", "gui", "console", "docs", "eeglab", "sys", "all"] [package.metadata.requires-dev] dev = [ { name = "ipython", specifier = ">=8.0" }, + { name = "pyedflib", specifier = ">=0.1.42" }, { name = "pyqtgraph", specifier = ">=0.13.7" }, { name = "pyside6", specifier = ">=6.6" }, { name = "pytest", specifier = ">=8.0" }, @@ -663,6 +682,14 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/81/47/dd9a212ef6e343a6857485ffe25bba537304f1913bdbed446a23f7f592e1/filelock-3.29.0-py3-none-any.whl", hash = "sha256:96f5f6344709aa1572bbf631c640e4ebeeb519e08da902c39a001882f30ac258", size = 39812, upload-time = "2026-04-19T15:39:08.752Z" }, ] +[[package]] +name = "flatbuffers" +version = "25.12.19" +source = { registry = "https://pypi.org/simple" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/e8/2d/d2a548598be01649e2d46231d151a6c56d10b964d94043a335ae56ea2d92/flatbuffers-25.12.19-py2.py3-none-any.whl", hash = "sha256:7634f50c427838bb021c2d66a3d1168e9d199b0607e6329399f04846d42e20b4", size = 26661, upload-time = "2025-12-19T23:16:13.622Z" }, +] + [[package]] name = "fonttools" version = "4.62.1" @@ -1292,6 +1319,93 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/a5/2b/72e0cc706c2884549323976f1d6221b74e0700f22e827e60b8413204eebb/metakernel-1.0.0-py3-none-any.whl", hash = "sha256:9f90660ce583baffe0235316f3e1a0abf75bbf60039e811ccbedd68ef612fa44", size = 204586, upload-time = "2026-03-24T01:24:11.774Z" }, ] +[[package]] +name = "ml-dtypes" +version = "0.5.4" +source = { registry = "https://pypi.org/simple" } +resolution-markers = [ + "python_full_version >= '3.14' and sys_platform != 'darwin'", + "python_full_version >= '3.14' and sys_platform == 'darwin'", +] +dependencies = [ + { name = "numpy" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/0e/4a/c27b42ed9b1c7d13d9ba8b6905dece787d6259152f2309338aed29b2447b/ml_dtypes-0.5.4.tar.gz", hash = "sha256:8ab06a50fb9bf9666dd0fe5dfb4676fa2b0ac0f31ecff72a6c3af8e22c063453", size = 692314, upload-time = "2025-11-17T22:32:31.031Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/a8/b8/3c70881695e056f8a32f8b941126cf78775d9a4d7feba8abcb52cb7b04f2/ml_dtypes-0.5.4-cp312-cp312-macosx_10_13_universal2.whl", hash = "sha256:a174837a64f5b16cab6f368171a1a03a27936b31699d167684073ff1c4237dac", size = 676927, upload-time = "2025-11-17T22:31:48.182Z" }, + { url = "https://files.pythonhosted.org/packages/54/0f/428ef6881782e5ebb7eca459689448c0394fa0a80bea3aa9262cba5445ea/ml_dtypes-0.5.4-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:a7f7c643e8b1320fd958bf098aa7ecf70623a42ec5154e3be3be673f4c34d900", size = 5028464, upload-time = "2025-11-17T22:31:50.135Z" }, + { url = "https://files.pythonhosted.org/packages/3a/cb/28ce52eb94390dda42599c98ea0204d74799e4d8047a0eb559b6fd648056/ml_dtypes-0.5.4-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:9ad459e99793fa6e13bd5b7e6792c8f9190b4e5a1b45c63aba14a4d0a7f1d5ff", size = 5009002, upload-time = "2025-11-17T22:31:52.001Z" }, + { url = "https://files.pythonhosted.org/packages/f5/f0/0cfadd537c5470378b1b32bd859cf2824972174b51b873c9d95cfd7475a5/ml_dtypes-0.5.4-cp312-cp312-win_amd64.whl", hash = "sha256:c1a953995cccb9e25a4ae19e34316671e4e2edaebe4cf538229b1fc7109087b7", size = 212222, upload-time = "2025-11-17T22:31:53.742Z" }, + { url = "https://files.pythonhosted.org/packages/16/2e/9acc86985bfad8f2c2d30291b27cd2bb4c74cea08695bd540906ed744249/ml_dtypes-0.5.4-cp312-cp312-win_arm64.whl", hash = "sha256:9bad06436568442575beb2d03389aa7456c690a5b05892c471215bfd8cf39460", size = 160793, upload-time = "2025-11-17T22:31:55.358Z" }, + { url = "https://files.pythonhosted.org/packages/d9/a1/4008f14bbc616cfb1ac5b39ea485f9c63031c4634ab3f4cf72e7541f816a/ml_dtypes-0.5.4-cp313-cp313-macosx_10_13_universal2.whl", hash = "sha256:8c760d85a2f82e2bed75867079188c9d18dae2ee77c25a54d60e9cc79be1bc48", size = 676888, upload-time = "2025-11-17T22:31:56.907Z" }, + { url = "https://files.pythonhosted.org/packages/d3/b7/dff378afc2b0d5a7d6cd9d3209b60474d9819d1189d347521e1688a60a53/ml_dtypes-0.5.4-cp313-cp313-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:ce756d3a10d0c4067172804c9cc276ba9cc0ff47af9078ad439b075d1abdc29b", size = 5036993, upload-time = "2025-11-17T22:31:58.497Z" }, + { url = "https://files.pythonhosted.org/packages/eb/33/40cd74219417e78b97c47802037cf2d87b91973e18bb968a7da48a96ea44/ml_dtypes-0.5.4-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:533ce891ba774eabf607172254f2e7260ba5f57bdd64030c9a4fcfbd99815d0d", size = 5010956, upload-time = "2025-11-17T22:31:59.931Z" }, + { url = "https://files.pythonhosted.org/packages/e1/8b/200088c6859d8221454825959df35b5244fa9bdf263fd0249ac5fb75e281/ml_dtypes-0.5.4-cp313-cp313-win_amd64.whl", hash = "sha256:f21c9219ef48ca5ee78402d5cc831bd58ea27ce89beda894428bc67a52da5328", size = 212224, upload-time = "2025-11-17T22:32:01.349Z" }, + { url = "https://files.pythonhosted.org/packages/8f/75/dfc3775cb36367816e678f69a7843f6f03bd4e2bcd79941e01ea960a068e/ml_dtypes-0.5.4-cp313-cp313-win_arm64.whl", hash = "sha256:35f29491a3e478407f7047b8a4834e4640a77d2737e0b294d049746507af5175", size = 160798, upload-time = "2025-11-17T22:32:02.864Z" }, + { url = "https://files.pythonhosted.org/packages/4f/74/e9ddb35fd1dd43b1106c20ced3f53c2e8e7fc7598c15638e9f80677f81d4/ml_dtypes-0.5.4-cp313-cp313t-macosx_10_13_universal2.whl", hash = "sha256:304ad47faa395415b9ccbcc06a0350800bc50eda70f0e45326796e27c62f18b6", size = 702083, upload-time = "2025-11-17T22:32:04.08Z" }, + { url = "https://files.pythonhosted.org/packages/74/f5/667060b0aed1aa63166b22897fdf16dca9eb704e6b4bbf86848d5a181aa7/ml_dtypes-0.5.4-cp313-cp313t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:6a0df4223b514d799b8a1629c65ddc351b3efa833ccf7f8ea0cf654a61d1e35d", size = 5354111, upload-time = "2025-11-17T22:32:05.546Z" }, + { url = "https://files.pythonhosted.org/packages/40/49/0f8c498a28c0efa5f5c95a9e374c83ec1385ca41d0e85e7cf40e5d519a21/ml_dtypes-0.5.4-cp313-cp313t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:531eff30e4d368cb6255bc2328d070e35836aa4f282a0fb5f3a0cd7260257298", size = 5366453, upload-time = "2025-11-17T22:32:07.115Z" }, + { url = "https://files.pythonhosted.org/packages/8c/27/12607423d0a9c6bbbcc780ad19f1f6baa2b68b18ce4bddcdc122c4c68dc9/ml_dtypes-0.5.4-cp313-cp313t-win_amd64.whl", hash = "sha256:cb73dccfc991691c444acc8c0012bee8f2470da826a92e3a20bb333b1a7894e6", size = 225612, upload-time = "2025-11-17T22:32:08.615Z" }, + { url = "https://files.pythonhosted.org/packages/e5/80/5a5929e92c72936d5b19872c5fb8fc09327c1da67b3b68c6a13139e77e20/ml_dtypes-0.5.4-cp313-cp313t-win_arm64.whl", hash = "sha256:3bbbe120b915090d9dd1375e4684dd17a20a2491ef25d640a908281da85e73f1", size = 164145, upload-time = "2025-11-17T22:32:09.782Z" }, + { url = "https://files.pythonhosted.org/packages/72/4e/1339dc6e2557a344f5ba5590872e80346f76f6cb2ac3dd16e4666e88818c/ml_dtypes-0.5.4-cp314-cp314-macosx_10_13_universal2.whl", hash = "sha256:2b857d3af6ac0d39db1de7c706e69c7f9791627209c3d6dedbfca8c7e5faec22", size = 673781, upload-time = "2025-11-17T22:32:11.364Z" }, + { url = "https://files.pythonhosted.org/packages/04/f9/067b84365c7e83bda15bba2b06c6ca250ce27b20630b1128c435fb7a09aa/ml_dtypes-0.5.4-cp314-cp314-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:805cef3a38f4eafae3a5bf9ebdcdb741d0bcfd9e1bd90eb54abd24f928cd2465", size = 5036145, upload-time = "2025-11-17T22:32:12.783Z" }, + { url = "https://files.pythonhosted.org/packages/c6/bb/82c7dcf38070b46172a517e2334e665c5bf374a262f99a283ea454bece7c/ml_dtypes-0.5.4-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:14a4fd3228af936461db66faccef6e4f41c1d82fcc30e9f8d58a08916b1d811f", size = 5010230, upload-time = "2025-11-17T22:32:14.38Z" }, + { url = "https://files.pythonhosted.org/packages/e9/93/2bfed22d2498c468f6bcd0d9f56b033eaa19f33320389314c19ef6766413/ml_dtypes-0.5.4-cp314-cp314-win_amd64.whl", hash = "sha256:8c6a2dcebd6f3903e05d51960a8058d6e131fe69f952a5397e5dbabc841b6d56", size = 221032, upload-time = "2025-11-17T22:32:15.763Z" }, + { url = "https://files.pythonhosted.org/packages/76/a3/9c912fe6ea747bb10fe2f8f54d027eb265db05dfb0c6335e3e063e74e6e8/ml_dtypes-0.5.4-cp314-cp314-win_arm64.whl", hash = "sha256:5a0f68ca8fd8d16583dfa7793973feb86f2fbb56ce3966daf9c9f748f52a2049", size = 163353, upload-time = "2025-11-17T22:32:16.932Z" }, + { url = "https://files.pythonhosted.org/packages/cd/02/48aa7d84cc30ab4ee37624a2fd98c56c02326785750cd212bc0826c2f15b/ml_dtypes-0.5.4-cp314-cp314t-macosx_10_13_universal2.whl", hash = "sha256:bfc534409c5d4b0bf945af29e5d0ab075eae9eecbb549ff8a29280db822f34f9", size = 702085, upload-time = "2025-11-17T22:32:18.175Z" }, + { url = "https://files.pythonhosted.org/packages/5a/e7/85cb99fe80a7a5513253ec7faa88a65306be071163485e9a626fce1b6e84/ml_dtypes-0.5.4-cp314-cp314t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:2314892cdc3fcf05e373d76d72aaa15fda9fb98625effa73c1d646f331fcecb7", size = 5355358, upload-time = "2025-11-17T22:32:19.7Z" }, + { url = "https://files.pythonhosted.org/packages/79/2b/a826ba18d2179a56e144aef69e57fb2ab7c464ef0b2111940ee8a3a223a2/ml_dtypes-0.5.4-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:0d2ffd05a2575b1519dc928c0b93c06339eb67173ff53acb00724502cda231cf", size = 5366332, upload-time = "2025-11-17T22:32:21.193Z" }, + { url = "https://files.pythonhosted.org/packages/84/44/f4d18446eacb20ea11e82f133ea8f86e2bf2891785b67d9da8d0ab0ef525/ml_dtypes-0.5.4-cp314-cp314t-win_amd64.whl", hash = "sha256:4381fe2f2452a2d7589689693d3162e876b3ddb0a832cde7a414f8e1adf7eab1", size = 236612, upload-time = "2025-11-17T22:32:22.579Z" }, + { url = "https://files.pythonhosted.org/packages/ad/3f/3d42e9a78fe5edf792a83c074b13b9b770092a4fbf3462872f4303135f09/ml_dtypes-0.5.4-cp314-cp314t-win_arm64.whl", hash = "sha256:11942cbf2cf92157db91e5022633c0d9474d4dfd813a909383bd23ce828a4b7d", size = 168825, upload-time = "2025-11-17T22:32:23.766Z" }, +] + +[[package]] +name = "ml-dtypes" +version = "0.6.0" +source = { registry = "https://pypi.org/simple" } +resolution-markers = [ + "python_full_version == '3.13.*' and sys_platform != 'darwin'", + "python_full_version == '3.13.*' and sys_platform == 'darwin'", + "python_full_version < '3.13' and sys_platform != 'darwin'", + "python_full_version < '3.13' and sys_platform == 'darwin'", +] +dependencies = [ + { name = "numpy" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/12/72/307d7c4bd0600601c7133fba5cb78af7db968152951c1cd473abb1cda782/ml_dtypes-0.6.0.tar.gz", hash = "sha256:5e60251d32ced5598972e4d5e06a2f044341f9291402551a3f6f0ec44f9299b0", size = 3032327, upload-time = "2026-08-13T14:14:40.215Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/84/6a/441eb053b078954f7fea284dfb288701884d0a1404d39babb858e1649023/ml_dtypes-0.6.0-cp312-cp312-macosx_10_13_universal2.whl", hash = "sha256:5359c588cc62de6f78d7430f06b65853d884955494d86d6ad90b6dd64a3f3a08", size = 565447, upload-time = "2026-08-13T14:14:01.737Z" }, + { url = "https://files.pythonhosted.org/packages/ed/cf/87e8a6c57eed63a91782a0d229856ddf73e138ce004dd71e2799a9dcdb33/ml_dtypes-0.6.0-cp312-cp312-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:37da32aa97749251025666d62372775019594577b9c9e9cfda83bed48d778fdb", size = 360227, upload-time = "2026-08-13T14:14:02.938Z" }, + { url = "https://files.pythonhosted.org/packages/c7/f9/7d76c1eae866f5d4636401b31b6d6dd90e4b4ced1fa7cfdfcca9c60e4bd3/ml_dtypes-0.6.0-cp312-cp312-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:3b4a480aa8fd54a1805b8ac10f3f91763926a74f73c0c364c10f9231854f4170", size = 409890, upload-time = "2026-08-13T14:14:04.248Z" }, + { url = "https://files.pythonhosted.org/packages/ba/db/9c61ec2760b5cbfb1c6558d5c991a6d8fd3271053c32db20506a9a90272b/ml_dtypes-0.6.0-cp312-cp312-win_amd64.whl", hash = "sha256:2a3e9d53925597fbffafd2a37048dadeddd0bdaba58058f6ae0869ed709a184d", size = 439333, upload-time = "2026-08-13T14:14:05.501Z" }, + { url = "https://files.pythonhosted.org/packages/6a/57/780ca3e5ab135b9fbdd8e5441abf5f801b30398371b691291e05ab9834c0/ml_dtypes-0.6.0-cp312-cp312-win_arm64.whl", hash = "sha256:6eaed129a4afe90694b8685e2f9b6294849f5eda4af9a15be83a4326eeebd775", size = 552268, upload-time = "2026-08-13T14:14:06.866Z" }, + { url = "https://files.pythonhosted.org/packages/50/51/fd1582b8f5ed8a9e7be0e161a6ea0dff70cb280479a12178df0b3a72700e/ml_dtypes-0.6.0-cp313-cp313-macosx_10_13_universal2.whl", hash = "sha256:084dfe51a7ad58b171f05115f8226ed4233a454a1611371947e806e76f0c638d", size = 565468, upload-time = "2026-08-13T14:14:08.5Z" }, + { url = "https://files.pythonhosted.org/packages/d2/22/20fd70ca6ed12446cb92d5b2a7745bd185f9d8b8cdeeadad976574398e6b/ml_dtypes-0.6.0-cp313-cp313-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:28d676428b104bb9717b0928bc5c5129f2d6b51b6727587cc4289e7bf8713cb5", size = 360232, upload-time = "2026-08-13T14:14:09.873Z" }, + { url = "https://files.pythonhosted.org/packages/89/a5/da8ae6c6f1babe4b68e3e55d43d39b529e29774f10e0910671a6b8c86eb8/ml_dtypes-0.6.0-cp313-cp313-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:26b1f1fa4f0435a2946859823f6e2bf06796f1e9f10f5a05b08a5e3c8f46ff69", size = 410169, upload-time = "2026-08-13T14:14:11.036Z" }, + { url = "https://files.pythonhosted.org/packages/e2/55/4561acefa00fa4bcbfb82ca6a48578b41f372cd7dd7cdd6eb4720abc2e5f/ml_dtypes-0.6.0-cp313-cp313-win_amd64.whl", hash = "sha256:fb87f46b4f7ad7b5d3ad8f4b452b024bd4229d44c8ff934798c1fe656210387a", size = 439357, upload-time = "2026-08-13T14:14:12.172Z" }, + { url = "https://files.pythonhosted.org/packages/b1/5d/6a01538e507ef0ed5e879985b13a92467bf8960696fb1131f8b8cadc60ff/ml_dtypes-0.6.0-cp313-cp313-win_arm64.whl", hash = "sha256:57ed0d6b4ac5e7868361303a9c57fbcf63b768236ee14456f585dfcf260d0292", size = 552278, upload-time = "2026-08-13T14:14:13.539Z" }, + { url = "https://files.pythonhosted.org/packages/d9/7a/97dc35667b7c9db33c5344c673cd27f87e34771875ea7100138726132ac9/ml_dtypes-0.6.0-cp314-cp314-macosx_10_15_universal2.whl", hash = "sha256:84fa136b8602c8c39e3b6cb24918960cd6f36cade7a70376f56770729cd56510", size = 562551, upload-time = "2026-08-13T14:14:14.774Z" }, + { url = "https://files.pythonhosted.org/packages/db/48/77f0ede10558d0d935da2e3276ed7e9c8cc2bad3463b9a0b66b03fc60be2/ml_dtypes-0.6.0-cp314-cp314-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:317be9967fb84b0ce4e80e6b1bf71213d21971621cf6f1e501a63602a95297bf", size = 360334, upload-time = "2026-08-13T14:14:16.079Z" }, + { url = "https://files.pythonhosted.org/packages/1c/b1/1831dd8c9b06c013085d31a2ac4f03392d43bd36bfc6ff591a08bcedc1cf/ml_dtypes-0.6.0-cp314-cp314-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:8f490c003369ce60e514a0c3b12374f05274c101fee1bead6740ec8a564032b0", size = 409966, upload-time = "2026-08-13T14:14:17.477Z" }, + { url = "https://files.pythonhosted.org/packages/ff/ad/9c32c53f823dda3742df19a79c10bc198365937873ea125ba65747440c23/ml_dtypes-0.6.0-cp314-cp314-win_amd64.whl", hash = "sha256:d574c2b28921dc72e869df248f1a278f6eee176a1f237c8642e1a71eb15f3977", size = 457224, upload-time = "2026-08-13T14:14:18.608Z" }, + { url = "https://files.pythonhosted.org/packages/41/3d/dd98205418a13353d41c52bf5326d8cbec515aace46174e23c6ea01c2978/ml_dtypes-0.6.0-cp314-cp314-win_arm64.whl", hash = "sha256:f4adb4af61516510d786cf8c01851a66f6d3ddfa79e1144deaa5b40d8507231e", size = 568378, upload-time = "2026-08-13T14:14:19.843Z" }, + { url = "https://files.pythonhosted.org/packages/65/36/32e7beef3281fed74883451477ad976364323206dbfaa95e948ba788dac7/ml_dtypes-0.6.0-cp314-cp314t-macosx_10_15_universal2.whl", hash = "sha256:3e169214e0d80ff1c038e1b3017e33c23e43bdf948d42d31de8283111c7e2fa3", size = 590177, upload-time = "2026-08-13T14:14:20.971Z" }, + { url = "https://files.pythonhosted.org/packages/d7/a2/99b3d9b3c984b3bd1e81d8244f1fa2f812e44060d853205b2df6271aa17c/ml_dtypes-0.6.0-cp314-cp314t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:573b11f3c327e17ef3826d266e676cf1149a1f3016f822a05f2306c55d8246bf", size = 363142, upload-time = "2026-08-13T14:14:22.463Z" }, + { url = "https://files.pythonhosted.org/packages/0c/fb/8091c0aee7f2712de99c7fd4b1642382644dec6a4962effe4f5b9d16a973/ml_dtypes-0.6.0-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:b76fa1d3f92967d58289ac47ab7458ede66e6f3527fff3e59142aee57d9307cd", size = 430645, upload-time = "2026-08-13T14:14:23.737Z" }, + { url = "https://files.pythonhosted.org/packages/c4/6f/962d2c589513b5930d05b6eae5fbd22ad8bbcf26bb763449f3d8f912360f/ml_dtypes-0.6.0-cp314-cp314t-win_amd64.whl", hash = "sha256:3be9911d953f97cddded4b9961d7b650473b7e55806d20f6176f8356dfe7b38e", size = 465667, upload-time = "2026-08-13T14:14:25.04Z" }, + { url = "https://files.pythonhosted.org/packages/aa/ca/bcb25e246edd19af5fa1cf6267040bd9977a7afca846e6cfd4a52078b44f/ml_dtypes-0.6.0-cp314-cp314t-win_arm64.whl", hash = "sha256:e74266ca8e97874a937b7646378c178025650a236584f7474d10d8086a6edea3", size = 572706, upload-time = "2026-08-13T14:14:26.296Z" }, + { url = "https://files.pythonhosted.org/packages/12/42/46cb442648e3c774d8cb25f2e1e41d496cdcc91fbe9c2a6f75c0b8df7af6/ml_dtypes-0.6.0-cp315-cp315-macosx_10_15_universal2.whl", hash = "sha256:b1b503864fada3f74fabf8d9fee7b4c1cbe956301e6fdece975d5f77c2fce958", size = 562550, upload-time = "2026-08-13T14:14:27.542Z" }, + { url = "https://files.pythonhosted.org/packages/07/56/844eff5af7a2d1a09d75df12c70225c3a6b6a771f95876b2bf5f7d10ad44/ml_dtypes-0.6.0-cp315-cp315-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:9c6ad60af4102789a5c09824004beade2f7f28cd1cd581ee5c170d9dc2fbb00e", size = 360332, upload-time = "2026-08-13T14:14:28.767Z" }, + { url = "https://files.pythonhosted.org/packages/b6/29/b7165a3a76364a5baa6aa4ee82a0adf73a3c014b8cd126120b62cc087992/ml_dtypes-0.6.0-cp315-cp315-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:d4f1b9329a251e4affe3bb58f4d3e2db22a714396fd7ffb40d0b5db423c24d17", size = 409964, upload-time = "2026-08-13T14:14:30.023Z" }, + { url = "https://files.pythonhosted.org/packages/c8/2e/f61c54a0544b6a170ac1bb89bcf406af53fb2deffc5476b6d2d3df5ba13e/ml_dtypes-0.6.0-cp315-cp315-win_amd64.whl", hash = "sha256:488c99ab181a2f59d9ec3b12c5fa11ec904e92be2c4ba18cded54dd7501208fe", size = 457249, upload-time = "2026-08-13T14:14:31.213Z" }, + { url = "https://files.pythonhosted.org/packages/63/00/bee1bc9faa02a46e7a851019fd23f47ca1f906609edbec8b6ba5decc3cc3/ml_dtypes-0.6.0-cp315-cp315-win_arm64.whl", hash = "sha256:de9d14748dbf3968951436ef514a29c9d1fe438aa680d110134ee2f7a9f9df18", size = 568381, upload-time = "2026-08-13T14:14:32.548Z" }, + { url = "https://files.pythonhosted.org/packages/72/f7/9a5edede28f73185fd51d75030ef7f11d76997bab3a92427d986e54fe2eb/ml_dtypes-0.6.0-cp315-cp315t-macosx_10_15_universal2.whl", hash = "sha256:e25bb3b0ad1217b60626e4ed45b10ca170c41d99fbe44a12bebc1e07ec4aad55", size = 589877, upload-time = "2026-08-13T14:14:33.695Z" }, + { url = "https://files.pythonhosted.org/packages/fd/81/d5924a141b850b606eb027493c9c3ca3c665cca5163af3f5b6e5e3345503/ml_dtypes-0.6.0-cp315-cp315t-manylinux_2_27_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:31f1ce979d31a357e95aa81812f20412c8c954fa43c44ee3ead1e1c8a78575ef", size = 362788, upload-time = "2026-08-13T14:14:34.996Z" }, + { url = "https://files.pythonhosted.org/packages/59/8f/3298e3f334832bc28dd144af6b99cdc93502a8687e71922ea68b0a319929/ml_dtypes-0.6.0-cp315-cp315t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:e2d6149f3a57f405bcad5fb41e03218b8373936253f23e1ca84c0108abbc3392", size = 430823, upload-time = "2026-08-13T14:14:36.44Z" }, + { url = "https://files.pythonhosted.org/packages/93/d2/f2dbf118f42ce4c325a139c9236737f436b7f8e00cd18701c99ef2405e6f/ml_dtypes-0.6.0-cp315-cp315t-win_amd64.whl", hash = "sha256:ce7563e0b1a4482cbc1b4a6272145e54e4489e54fe7428f94908c3d87103abfa", size = 465119, upload-time = "2026-08-13T14:14:37.776Z" }, + { url = "https://files.pythonhosted.org/packages/5a/ff/bda40387b5c5c64254595f4d81a12351770856acc5de4e6d43606a31f161/ml_dtypes-0.6.0-cp315-cp315t-win_arm64.whl", hash = "sha256:f6cb525101b6b903779188c1e9e9490c343b455ab822883e02cf01e5547338d2", size = 572666, upload-time = "2026-08-13T14:14:38.993Z" }, +] + [[package]] name = "mne" version = "1.12.1" @@ -1677,6 +1791,65 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/b3/67/b7a132f9929828f95654701b738a8d8735bd99f9b1b512281c4411e31cf3/octave_kernel-1.0.3-py3-none-any.whl", hash = "sha256:de989830cbd2904b90d63807c2a7a4e0237bfcb4953e61101173d4009835d518", size = 40456, upload-time = "2026-04-08T18:53:32.446Z" }, ] +[[package]] +name = "onnx" +version = "1.23.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "ml-dtypes", version = "0.5.4", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version >= '3.14'" }, + { name = "ml-dtypes", version = "0.6.0", source = { registry = "https://pypi.org/simple" }, marker = "python_full_version < '3.14'" }, + { name = "numpy" }, + { name = "protobuf" }, + { name = "typing-extensions" }, +] +sdist = { url = "https://files.pythonhosted.org/packages/29/3a/b73bfb8f506a7b5be43eebf71dbdd351c4116f8100e1103a91bfcbcd3502/onnx-1.23.0.tar.gz", hash = "sha256:3f4b7323ea3397c63aa6ff5d43abe16013ffefc0a96cb8a0be6527ff1e1add40", size = 6022619, upload-time = "2026-09-18T19:30:37.511Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/63/f9/54d8047c2f409cfc314a0df520c6e5703afe78b1395e8c44287724c4b37e/onnx-1.23.0-cp312-abi3-macosx_13_0_universal2.whl", hash = "sha256:6cc93aa89502acbf453a16d79b465d47d9c454f3a3cb3fc5507a8cb59ab53208", size = 9725329, upload-time = "2026-09-18T19:30:07.486Z" }, + { url = "https://files.pythonhosted.org/packages/c6/ff/a7d3fe41debd5b8947182a632b4e001bd3a7a103212e58ef6ca0c3daab91/onnx-1.23.0-cp312-abi3-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:95b77466b4f6bc7c2c63d6004b8db7d963fb17bf24aff13007752e9d2ebf9390", size = 8640232, upload-time = "2026-09-18T19:30:10.058Z" }, + { url = "https://files.pythonhosted.org/packages/09/86/63cdd5740f566fd9de2cadc106529adee0566f2fcafe69a65767eedc1681/onnx-1.23.0-cp312-abi3-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:f336004196a22fbdc16c62e7f26f20635af1826db6147c80ff3b4b8d428fc7ef", size = 8881347, upload-time = "2026-09-18T19:30:12.792Z" }, + { url = "https://files.pythonhosted.org/packages/cd/24/7844929ef08aee3f9cb857d5ce34bb60d44496c9cc5c66cdfd8255917283/onnx-1.23.0-cp312-abi3-pyemscripten_2026_0_wasm32.whl", hash = "sha256:05ad734997f09fcc9633fc2e3715033239792fc9f12aec6ec0c48e32b214e47e", size = 7314556, upload-time = "2026-09-18T19:30:15.456Z" }, + { url = "https://files.pythonhosted.org/packages/52/68/ee1ad44118d0031e28a1cabb1cc02a4c9e30273bdc25571c5847aec5e4e8/onnx-1.23.0-cp312-abi3-win32.whl", hash = "sha256:221037cb9ffe67b701fb180ca871f1ee9089f431d944f3d0313ffae68c356a80", size = 7736116, upload-time = "2026-09-18T19:30:18.06Z" }, + { url = "https://files.pythonhosted.org/packages/81/5a/766abb1411b9f9b06dc4e8fe9ff195d3b715535677c362ab3d950fa0e93f/onnx-1.23.0-cp312-abi3-win_amd64.whl", hash = "sha256:70a2f930b221f9dbdff62704838ce8bf81442787b25048f8b9cdc9118168799e", size = 7872197, upload-time = "2026-09-18T19:30:20.619Z" }, + { url = "https://files.pythonhosted.org/packages/94/34/483fa56b4d47a5210d64dccb522a60b46e0953ca3f72d75b3c4759266f5f/onnx-1.23.0-cp312-abi3-win_arm64.whl", hash = "sha256:dcd4d269818a13c48444c96b7f17c7539d8cda7810ef802431393b6cf5df794b", size = 8037224, upload-time = "2026-09-18T19:30:22.682Z" }, + { url = "https://files.pythonhosted.org/packages/4b/7a/3384704bbc90dba09bb6924399db717b8031704ad9bc80191c3cfe5cf88b/onnx-1.23.0-cp314-cp314t-macosx_13_0_universal2.whl", hash = "sha256:8e00eb5aaccbeef2203439e24b0e8f2bf01537fae3575ffbaacf7d8bc9d96e3b", size = 9730894, upload-time = "2026-09-18T19:30:24.982Z" }, + { url = "https://files.pythonhosted.org/packages/64/23/40a365a39f7f3148196e98f496b31b5588414b7898fe65f8343b402a6d1a/onnx-1.23.0-cp314-cp314t-manylinux_2_26_aarch64.manylinux_2_28_aarch64.whl", hash = "sha256:5741b60e189b5c8977beb8ec870c47bf8d2d938add5d5f11467facf6904de4cd", size = 8647161, upload-time = "2026-09-18T19:30:27.619Z" }, + { url = "https://files.pythonhosted.org/packages/86/dd/03c1b211b935cf6bf3f96f0d18a0ea62a3ac522212af4277ac47fba02c47/onnx-1.23.0-cp314-cp314t-manylinux_2_27_x86_64.manylinux_2_28_x86_64.whl", hash = "sha256:d6951c2ed343eb9dbc2b644507f9ed8d34c23558bae839b3726b5b8109d3ff8b", size = 8886389, upload-time = "2026-09-18T19:30:30.226Z" }, + { url = "https://files.pythonhosted.org/packages/b3/18/f4fe44070ffe4bb3a08898dac6c7a2eb78691025bce40bdaadfeba8795b3/onnx-1.23.0-cp314-cp314t-win_amd64.whl", hash = "sha256:3e9c9be1baebbc2b257c45937e4fb18f9e69a021d4207a909dc1f4407193fa21", size = 7910378, upload-time = "2026-09-18T19:30:32.962Z" }, + { url = "https://files.pythonhosted.org/packages/c4/92/b7e6f3a85f5ef39f012839cb77123544ada6e9cd69dadae4a66df5403b30/onnx-1.23.0-cp314-cp314t-win_arm64.whl", hash = "sha256:ae166bf343b95ee9d03bf023faac05216fe2c9eff44d1c9a4c01920657135c45", size = 8078741, upload-time = "2026-09-18T19:30:35.245Z" }, +] + +[[package]] +name = "onnxruntime" +version = "1.30.0" +source = { registry = "https://pypi.org/simple" } +dependencies = [ + { name = "flatbuffers" }, + { name = "numpy" }, + { name = "packaging" }, + { name = "protobuf" }, +] +wheels = [ + { url = "https://files.pythonhosted.org/packages/31/6f/48169f2e62b405bff5053cbd1d73fb5ce41ef7ecd13bb3bfcc191e689b8a/onnxruntime-1.30.0-cp312-cp312-macosx_14_0_arm64.whl", hash = "sha256:001ed726c9bd5e2bc92faade7d37d889e9606a350b7d5529f0227df2e3bb57fd", size = 21544867, upload-time = "2026-09-10T16:31:19.876Z" }, + { url = "https://files.pythonhosted.org/packages/16/bd/cbc5b8f91963689fdd622f463508c01d0aa95d3f944747b1e0b1eb2160b8/onnxruntime-1.30.0-cp312-cp312-manylinux_2_28_aarch64.whl", hash = "sha256:6c32a000d5139a38ba9349030b0032e3331acb559d596b22738d9d2b343a2b83", size = 21345202, upload-time = "2026-09-10T16:31:23.361Z" }, + { url = "https://files.pythonhosted.org/packages/34/35/e7f862dbacbc99fadd9b14a614e49c99bf0f35fd9927a82f096e3de33531/onnxruntime-1.30.0-cp312-cp312-manylinux_2_28_x86_64.whl", hash = "sha256:fa688e7891a6aa206636fe7372e27ee75fd17713289f6b4fc7b190e0a7de9328", size = 23585654, upload-time = "2026-09-10T16:31:26.65Z" }, + { url = "https://files.pythonhosted.org/packages/a6/13/0f1699f6de549c9324bc9112a2a85b14c517904cd11b562a654643b755a1/onnxruntime-1.30.0-cp312-cp312-win_amd64.whl", hash = "sha256:f3501472571f1b1eee50e017851e7929f5ea37312d2d8c2494a19e8fc58b4a38", size = 14311470, upload-time = "2026-09-10T16:31:31.273Z" }, + { url = "https://files.pythonhosted.org/packages/a2/17/02b13e5f51461f0453b18ab854e2d0bc1b6ec353241b05a2ee6b79e35d87/onnxruntime-1.30.0-cp312-cp312-win_arm64.whl", hash = "sha256:dc4c706f1935ebb62356e6a095b047859badd854482c40560888e95c328ed262", size = 14175072, upload-time = "2026-09-10T16:31:34.541Z" }, + { url = "https://files.pythonhosted.org/packages/f0/75/508454c5d01f31641dabc597fe559594c931a520a2673031179319d0afd8/onnxruntime-1.30.0-cp313-cp313-macosx_14_0_arm64.whl", hash = "sha256:05e4fc41711d1f4abd19a9124b5be7a65a506cb2670a7f14b16162ef13c58134", size = 21544990, upload-time = "2026-09-10T16:31:37.685Z" }, + { url = "https://files.pythonhosted.org/packages/89/06/e603c71f43f4fe3fd156a053af79cbed6e27a2c649f0988a67d97fedd39f/onnxruntime-1.30.0-cp313-cp313-manylinux_2_28_aarch64.whl", hash = "sha256:5327cf6aa15a02bad805fac8bd6882a62571e8b72f6f2938a8f37e6bd1966ce9", size = 21344996, upload-time = "2026-09-10T16:31:40.915Z" }, + { url = "https://files.pythonhosted.org/packages/f1/a1/ede48ab5dc54907a2999362777f541e132639fb06628ded1932058aa8a36/onnxruntime-1.30.0-cp313-cp313-manylinux_2_28_x86_64.whl", hash = "sha256:86f940afc801ea9681a4da8af84fbe95e1d9ea7d80903952cc1bfad54faad38f", size = 23585560, upload-time = "2026-09-10T16:31:44.941Z" }, + { url = "https://files.pythonhosted.org/packages/3c/dd/c57c529dbc6dd55eca24b12cfbeab1b6a690de72083824eca689085f55b0/onnxruntime-1.30.0-cp313-cp313-win_amd64.whl", hash = "sha256:4b63041bd623a9a9ac5e353948436c6fa7f43edd12d6b4a4ebc340bca959ba93", size = 14311378, upload-time = "2026-09-10T16:31:48.034Z" }, + { url = "https://files.pythonhosted.org/packages/11/2f/ef00b45b911e7f2115273a24ed1b43e73c8c6c9ec17f9bfa13f8492581f8/onnxruntime-1.30.0-cp313-cp313-win_arm64.whl", hash = "sha256:c389b6887fc95e0fcb80e89b156bc2cb18e662c29df55e9326fe64140e7d7b4f", size = 14175134, upload-time = "2026-09-10T16:31:50.799Z" }, + { url = "https://files.pythonhosted.org/packages/92/0a/284fd6fe701c9a8aff39dbc119c37582ca40e9cff1ea05425a7cc8606a02/onnxruntime-1.30.0-cp313-cp313t-manylinux_2_28_aarch64.whl", hash = "sha256:e5224ba2b00284cb1c48b3edcd303c109245d66df9ec1b861858fd6de672a38e", size = 21349453, upload-time = "2026-09-10T16:31:53.898Z" }, + { url = "https://files.pythonhosted.org/packages/d6/72/4f4466f8fa1ec267a9ef5e2f3bd175c203d3da0facbdf4049dc93abbe91c/onnxruntime-1.30.0-cp313-cp313t-manylinux_2_28_x86_64.whl", hash = "sha256:3f9e002417f1e3bbb31ed43dafa4f22ca1b2b68832244fdee192fcaa1ae19bca", size = 23579837, upload-time = "2026-09-10T16:31:56.82Z" }, + { url = "https://files.pythonhosted.org/packages/6a/03/05c9a9234688757d2876ddecf80bba908561ee11debf125bc1a427ae6f48/onnxruntime-1.30.0-cp314-cp314-macosx_14_0_arm64.whl", hash = "sha256:8b6169c16a48429890d2f4a0c774ebf54dfe9066a998514aad0518a16d398547", size = 21545313, upload-time = "2026-09-10T16:31:59.654Z" }, + { url = "https://files.pythonhosted.org/packages/c6/bc/1069e58b24779ba9d2fd479db5ecb3a15a6f49b585107c898819c0789558/onnxruntime-1.30.0-cp314-cp314-manylinux_2_28_aarch64.whl", hash = "sha256:d2184fddb6798136e7c478244391ca82443f5c757f59f15bb9e5ad2da5e03175", size = 21346830, upload-time = "2026-09-10T16:32:02.275Z" }, + { url = "https://files.pythonhosted.org/packages/f1/38/8138eed225c5bc6ddfc05879ecac7dacc63c34b9b6f99be72839c1f6dc49/onnxruntime-1.30.0-cp314-cp314-manylinux_2_28_x86_64.whl", hash = "sha256:8b611d24db2954545ce6bd9acd4670183cb368e4642450de7a9ab6474eb374ec", size = 23586333, upload-time = "2026-09-10T16:32:04.996Z" }, + { url = "https://files.pythonhosted.org/packages/4d/0a/748b86000fbf9518b5dfc5bdc1b924eeaca976203740fd9d30649e568ee6/onnxruntime-1.30.0-cp314-cp314-win_amd64.whl", hash = "sha256:3bdd1752a8502ac1a7ccc6e878d16db6943df574d54ed7f01058a6806ae05be4", size = 14675319, upload-time = "2026-09-10T16:32:07.5Z" }, + { url = "https://files.pythonhosted.org/packages/2f/db/db8e4c0cf6f70f1311060560630fd21763648e01dfbba5195a2592c052ef/onnxruntime-1.30.0-cp314-cp314-win_arm64.whl", hash = "sha256:83d543843cbd352cfa9996a6c7b92f8a480f114a18f95ff2a1acaf6921e85b6d", size = 14569360, upload-time = "2026-09-10T16:32:09.896Z" }, + { url = "https://files.pythonhosted.org/packages/24/0a/ec0a9d656e39b43887c378a5388b20c3b1e1ee43bded84580b0936466694/onnxruntime-1.30.0-cp314-cp314t-manylinux_2_28_aarch64.whl", hash = "sha256:265de607ba6f9814264e1d5d413fa7069d70f48e3049b5e460b3b04bdfcef294", size = 21348433, upload-time = "2026-09-10T16:32:12.772Z" }, + { url = "https://files.pythonhosted.org/packages/91/f0/40f74b7c00077e1e25627067ed98a70df1fef5c0e21b82849190d312554e/onnxruntime-1.30.0-cp314-cp314t-manylinux_2_28_x86_64.whl", hash = "sha256:67ad7f03433b6462c627d0f555dece80e6a26bc71e8542ced35cebd32142d1b7", size = 23579340, upload-time = "2026-09-10T16:32:15.532Z" }, +] + [[package]] name = "packaging" version = "26.2" @@ -1863,6 +2036,21 @@ wheels = [ { url = "https://files.pythonhosted.org/packages/84/03/0d3ce49e2505ae70cf43bc5bb3033955d2fc9f932163e84dc0779cc47f48/prompt_toolkit-3.0.52-py3-none-any.whl", hash = "sha256:9aac639a3bbd33284347de5ad8d68ecc044b91a762dc39b7c21095fcd6a19955", size = 391431, upload-time = "2025-08-27T15:23:59.498Z" }, ] +[[package]] +name = "protobuf" +version = "7.36.2" +source = { registry = "https://pypi.org/simple" } +sdist = { url = "https://files.pythonhosted.org/packages/d9/89/5b8517baa72f84a67b8a307ba953c91057af618bf40bf676f3c03551f8f0/protobuf-7.36.2.tar.gz", hash = "sha256:497d0463ff3316681da6c0b9e8d06cb465d61abce00b613ab42226175644d1bb", size = 512737, upload-time = "2026-09-17T20:07:59.326Z" } +wheels = [ + { url = "https://files.pythonhosted.org/packages/32/72/98342feb672507c8f3a69e34b4fa8961f608edba5c1a48a6f47156d92cb5/protobuf-7.36.2-cp310-abi3-macosx_10_9_universal2.whl", hash = "sha256:cbc70b17ee27e28894c7fee8bb04be1abead49e936bc70eb60052531eee2079e", size = 456039, upload-time = "2026-09-17T20:07:51.542Z" }, + { url = "https://files.pythonhosted.org/packages/b6/ea/91fdf7c2b8bbd49cde056f00a9df6773532987e1c00fe2830b895af95c7e/protobuf-7.36.2-cp310-abi3-manylinux2014_aarch64.whl", hash = "sha256:e11e1f0180583a2af89db6a2ecd9e8dc40aa6d2988ca175bfd0e6d12ea72d74e", size = 344219, upload-time = "2026-09-17T20:07:52.914Z" }, + { url = "https://files.pythonhosted.org/packages/17/ab/5fd5f8ece73fad885c5a09aa849b32d70472f954ba3a92d3bb5974ea953b/protobuf-7.36.2-cp310-abi3-manylinux2014_s390x.whl", hash = "sha256:f4fee11ec330d238b34a05c9b675f693c20415d1c5bd7d5320cc2f8a798eb9cf", size = 357223, upload-time = "2026-09-17T20:07:53.985Z" }, + { url = "https://files.pythonhosted.org/packages/db/f3/3996583dd2906297a637af12114deddf7658af6e683fedb83be061983fb5/protobuf-7.36.2-cp310-abi3-manylinux2014_x86_64.whl", hash = "sha256:89f23aa53c24553a2416fd4fd1ec06f74fa42b14b546d8883128813f775bbfd2", size = 343223, upload-time = "2026-09-17T20:07:54.931Z" }, + { url = "https://files.pythonhosted.org/packages/fc/1b/dcc64f358fcb51811b58ae40b3d28f820725f116d86487cc20bd4b130701/protobuf-7.36.2-cp310-abi3-win32.whl", hash = "sha256:912c1221170e16c08d1f086762f563dd61ff83c18b5fa6652952dfaded66f728", size = 442998, upload-time = "2026-09-17T20:07:55.826Z" }, + { url = "https://files.pythonhosted.org/packages/8a/55/b77bda4e5e5f5971fb51b07663694690e9afdb9402136c16a522bd621cad/protobuf-7.36.2-cp310-abi3-win_amd64.whl", hash = "sha256:a300819d441e078a5608c0d3c709796bb548136058fda017ae51d425b44fd353", size = 456514, upload-time = "2026-09-17T20:07:57.188Z" }, + { url = "https://files.pythonhosted.org/packages/e4/04/d52c7016b04b6c5108f26691f9d33ec82a9b65d041f1a9c771137693d618/protobuf-7.36.2-py3-none-any.whl", hash = "sha256:bdb3a345d48db958e6ce1f18e508beb0cc981d64f24088427549c866cd039f1e", size = 179806, upload-time = "2026-09-17T20:07:58.211Z" }, +] + [[package]] name = "psutil" version = "7.2.2" @@ -2374,7 +2562,8 @@ name = "scipy" version = "1.15.3" source = { registry = "https://pypi.org/simple" } resolution-markers = [ - "python_full_version >= '3.13' and sys_platform != 'darwin'", + "python_full_version >= '3.14' and sys_platform != 'darwin'", + "python_full_version == '3.13.*' and sys_platform != 'darwin'", "python_full_version < '3.13' and sys_platform != 'darwin'", ] dependencies = [ @@ -2404,7 +2593,8 @@ name = "scipy" version = "1.16.3" source = { registry = "https://pypi.org/simple" } resolution-markers = [ - "python_full_version >= '3.13' and sys_platform == 'darwin'", + "python_full_version >= '3.14' and sys_platform == 'darwin'", + "python_full_version == '3.13.*' and sys_platform == 'darwin'", "python_full_version < '3.13' and sys_platform == 'darwin'", ] dependencies = [