From fc996ad99763721e1a78ce70aa189bc998a4642d Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Mon, 28 Sep 2026 12:08:11 -0400 Subject: [PATCH 01/16] Prepare ModDotPlot v1.0.0 --- .github/workflows/black.yml | 19 - .github/workflows/ci.yml | 133 ++ .github/workflows/publish-to-pypi.yml | 183 +- LICENSE.ntHash | 21 + README.md | 44 +- benchmarks/benchmark_containment_matrix.py | 102 ++ benchmarks/benchmark_hashing.py | 256 +++ benchmarks/benchmark_sketch_reuse.py | 153 ++ pyproject.toml | 15 +- setup.cfg | 7 + setup.py | 24 + src/moddotplot/_nthash.cpp | 342 ++++ src/moddotplot/const.py | 4 +- src/moddotplot/estimate_identity.py | 755 ++++++-- src/moddotplot/interactive.py | 78 +- src/moddotplot/moddotplot.py | 807 ++++++--- src/moddotplot/native_render.py | 451 +++++ src/moddotplot/parse_fasta.py | 563 +++++- src/moddotplot/static_plots.py | 1897 ++++++++++---------- tests/test_algorithms.py | 269 +++ tests/test_annotation_track.py | 256 +++ tests/test_cli_integration.py | 191 ++ tests/test_cli_runtime.py | 234 +++ tests/test_direction_cli.py | 119 ++ tests/test_direction_plot.py | 61 + tests/test_entrypoints.py | 54 + tests/test_fasta_parser.py | 256 +++ tests/test_grid.py | 480 +++++ tests/test_hash_benchmark.py | 64 + tests/test_interactive_export.py | 35 + tests/test_interactive_parser.py | 21 + tests/test_issue53_memory_safe_plotting.py | 105 ++ tests/test_modimizer_vectorization.py | 129 ++ tests/test_native_render.py | 204 +++ tests/test_nthash.py | 283 +++ tests/test_packaging_metadata.py | 110 ++ tests/test_sketch_cache.py | 170 ++ tests/test_sparse_containment.py | 148 ++ tests/test_static_customization.py | 210 +++ 39 files changed, 7646 insertions(+), 1607 deletions(-) delete mode 100644 .github/workflows/black.yml create mode 100644 .github/workflows/ci.yml create mode 100644 LICENSE.ntHash create mode 100644 benchmarks/benchmark_containment_matrix.py create mode 100644 benchmarks/benchmark_hashing.py create mode 100644 benchmarks/benchmark_sketch_reuse.py create mode 100644 setup.cfg create mode 100644 setup.py create mode 100644 src/moddotplot/_nthash.cpp create mode 100644 src/moddotplot/native_render.py create mode 100644 tests/test_algorithms.py create mode 100644 tests/test_annotation_track.py create mode 100644 tests/test_cli_integration.py create mode 100644 tests/test_cli_runtime.py create mode 100644 tests/test_direction_cli.py create mode 100644 tests/test_direction_plot.py create mode 100644 tests/test_entrypoints.py create mode 100644 tests/test_fasta_parser.py create mode 100644 tests/test_grid.py create mode 100644 tests/test_hash_benchmark.py create mode 100644 tests/test_interactive_export.py create mode 100644 tests/test_interactive_parser.py create mode 100644 tests/test_issue53_memory_safe_plotting.py create mode 100644 tests/test_modimizer_vectorization.py create mode 100644 tests/test_native_render.py create mode 100644 tests/test_nthash.py create mode 100644 tests/test_packaging_metadata.py create mode 100644 tests/test_sketch_cache.py create mode 100644 tests/test_sparse_containment.py create mode 100644 tests/test_static_customization.py diff --git a/.github/workflows/black.yml b/.github/workflows/black.yml deleted file mode 100644 index b1287d3..0000000 --- a/.github/workflows/black.yml +++ /dev/null @@ -1,19 +0,0 @@ -name: black - -# Controls when the action will run. -on: - push: - branches: [main, develop] - pull_request: - branches: [main, develop] - - workflow_dispatch: - -jobs: - black: - runs-on: ubuntu-latest - steps: - - uses: actions/checkout@v2 - - uses: psf/black@stable - with: - options: ". --check --verbose" \ No newline at end of file diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..53dff37 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,133 @@ +name: CI + +on: + push: + branches: [main, develop] + pull_request: + branches: [main, develop] + workflow_dispatch: + +permissions: + contents: read + +concurrency: + group: ci-${{ github.workflow }}-${{ github.ref }} + cancel-in-progress: true + +jobs: + tests: + name: Tests (Python ${{ matrix.python-version }}) + runs-on: ubuntu-latest + timeout-minutes: 45 + strategy: + fail-fast: false + matrix: + python-version: + - "3.8" + - "3.9" + - "3.10" + - "3.11" + - "3.12" + env: + MPLBACKEND: Agg + + steps: + - name: Check out source + uses: actions/checkout@v7 + with: + persist-credentials: false + + - name: Set up Python ${{ matrix.python-version }} + uses: actions/setup-python@v7 + with: + python-version: ${{ matrix.python-version }} + cache: pip + cache-dependency-path: pyproject.toml + + - name: Install package and test dependencies + run: | + python -m pip install --upgrade pip + python -m pip install --editable ".[test]" + + - name: Run unit and integration tests + run: >- + python -m pytest + --cov=moddotplot + --cov-report=term-missing + --cov-report=xml + --junitxml=junit.xml + + - name: Store test reports + if: always() + uses: actions/upload-artifact@v7 + with: + name: test-reports-python-${{ matrix.python-version }} + path: | + coverage.xml + junit.xml + if-no-files-found: warn + + style: + name: Formatting + runs-on: ubuntu-latest + timeout-minutes: 10 + + steps: + - name: Check out source + uses: actions/checkout@v7 + with: + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@v7 + with: + python-version: "3.12" + cache: pip + cache-dependency-path: pyproject.toml + + - name: Install Black + run: python -m pip install black + + - name: Check formatting + run: python -m black --check src tests benchmarks setup.py + + package: + name: Build and validate package + runs-on: ubuntu-latest + timeout-minutes: 15 + + steps: + - name: Check out source + uses: actions/checkout@v7 + with: + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@v7 + with: + python-version: "3.12" + cache: pip + cache-dependency-path: pyproject.toml + + - name: Install packaging tools + run: python -m pip install build twine + + - name: Build wheel and source distribution + run: python -m build + + - name: Validate distribution metadata + run: python -m twine check --strict dist/* + + - name: Install and inspect the wheel + run: | + python -c "import pathlib; wheel = next(pathlib.Path('dist').glob('*.whl')); assert 'cp38-abi3' in wheel.name, wheel.name" + python -m pip install --no-deps --force-reinstall dist/*.whl + python -c "import importlib.metadata as m, pathlib, tomllib; expected = tomllib.loads(pathlib.Path('pyproject.toml').read_text())['project']; actual = m.metadata('ModDotPlot'); assert m.version('ModDotPlot') == expected['version']; assert actual['Requires-Python'] == expected['requires-python']" + python -c "from moddotplot import _nthash; hashes, mask = _nthash.hash_kmers('ACGTACGT', 3, True); assert len(hashes) == 48 and mask == b''" + + - name: Store distributions + uses: actions/upload-artifact@v7 + with: + name: python-package-distributions + path: dist/ + if-no-files-found: error diff --git a/.github/workflows/publish-to-pypi.yml b/.github/workflows/publish-to-pypi.yml index 44f8944..fd41b58 100644 --- a/.github/workflows/publish-to-pypi.yml +++ b/.github/workflows/publish-to-pypi.yml @@ -1,53 +1,156 @@ -name: Publish Python 🐍 distribution 📦 to PyPI and TestPyPI +name: Publish release to PyPI -on: push +on: + push: + tags: + - "v*" + +permissions: + contents: read jobs: - build: - name: Build distribution 📦 + verify: + name: Verify release and run tests runs-on: ubuntu-latest + timeout-minutes: 45 + env: + MPLBACKEND: Agg + + steps: + - name: Check out tagged source + uses: actions/checkout@v7 + with: + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@v7 + with: + python-version: "3.12" + cache: pip + cache-dependency-path: pyproject.toml + + - name: Verify tag matches the package version + env: + RELEASE_TAG: ${{ github.ref_name }} + run: >- + python -c + "import os, pathlib, tomllib; + version = tomllib.loads(pathlib.Path('pyproject.toml').read_text())['project']['version']; + tag = os.environ['RELEASE_TAG']; + assert tag == f'v{version}', f'tag {tag!r} does not match package version v{version}'" + + - name: Install package and tests + run: | + python -m pip install --upgrade pip + python -m pip install ".[test]" + + - name: Run the release test suite + run: python -m pytest + + sdist: + name: Build source distribution + needs: verify + runs-on: ubuntu-latest + timeout-minutes: 15 + + steps: + - name: Check out tagged source + uses: actions/checkout@v7 + with: + persist-credentials: false + + - name: Set up Python + uses: actions/setup-python@v7 + with: + python-version: "3.12" + + - name: Install packaging tools + run: python -m pip install build twine + + - name: Build source distribution + run: python -m build --sdist --outdir dist + + - name: Validate distribution metadata + run: python -m twine check --strict dist/* + + - name: Store source distribution for publication + uses: actions/upload-artifact@v7 + with: + name: source-distribution + path: dist/ + if-no-files-found: error + retention-days: 1 + + wheels: + name: Build wheel (${{ matrix.os }}) + needs: verify + runs-on: ${{ matrix.os }} + timeout-minutes: 45 + strategy: + fail-fast: false + matrix: + os: + - ubuntu-latest + - macos-15-intel + - windows-latest steps: - - uses: actions/checkout@v4 - with: - persist-credentials: false - - name: Set up Python - uses: actions/setup-python@v5 - with: - python-version: "3.x" - - - name: Install pypa/build - run: >- - python3 -m - pip install - build - --user - - name: Build a binary wheel and a source tarball - run: python3 -m build - - name: Store the distribution packages - uses: actions/upload-artifact@v4 - with: - name: python-package-distributions - path: dist/ - - publish-to-pypi: - name: >- - Publish Python 🐍 distribution 📦 to PyPI - if: startsWith(github.ref, 'refs/tags/') # only publish to PyPI on tag pushes + - name: Check out tagged source + uses: actions/checkout@v7 + with: + persist-credentials: false + + # The extension uses CPython's stable ABI. Building once with CPython 3.9 + # produces a cp38-abi3 wheel that supports every Python version declared + # by the package without compiling the same binary repeatedly. + - name: Build stable-ABI wheel + uses: pypa/cibuildwheel@v4.2.0 + env: + CIBW_BUILD: "cp39-*" + CIBW_SKIP: "*-musllinux_*" + CIBW_ARCHS_MACOS: universal2 + CIBW_TEST_COMMAND: >- + python -c "from moddotplot import _nthash; + hashes, mask = _nthash.hash_kmers('ACGTACGT', 3, True); + assert len(hashes) == 48 and mask == b''" + with: + output-dir: wheelhouse + + - name: Store wheel for publication + uses: actions/upload-artifact@v7 + with: + name: wheel-${{ matrix.os }} + path: wheelhouse/ + if-no-files-found: error + retention-days: 1 + + publish: + name: Publish distributions to PyPI needs: - - build + - sdist + - wheels runs-on: ubuntu-latest environment: name: pypi - url: https://pypi.org/p/ModDotPlot + url: https://pypi.org/project/ModDotPlot permissions: - id-token: write # IMPORTANT: mandatory for trusted publishing + id-token: write steps: - - name: Download all the dists - uses: actions/download-artifact@v4 - with: - name: python-package-distributions - path: dist/ - - name: Publish distribution 📦 to PyPI - uses: pypa/gh-action-pypi-publish@release/v1 + - name: Download validated distributions + uses: actions/download-artifact@v8 + with: + path: dist/ + pattern: "*" + merge-multiple: true + + - name: Check distribution metadata + run: | + python -m pip install twine + python -m twine check --strict dist/* + + - name: Publish distributions with Trusted Publishing + uses: pypa/gh-action-pypi-publish@release/v1 + with: + attestations: true + print-hash: true diff --git a/LICENSE.ntHash b/LICENSE.ntHash new file mode 100644 index 0000000..6c587d2 --- /dev/null +++ b/LICENSE.ntHash @@ -0,0 +1,21 @@ +MIT License + +Copyright (c) 2018 Hamid Mohamadi + +Permission is hereby granted, free of charge, to any person obtaining a copy +of this software and associated documentation files (the "Software"), to deal +in the Software without restriction, including without limitation the rights +to use, copy, modify, merge, publish, distribute, sublicense, and/or sell +copies of the Software, and to permit persons to whom the Software is +furnished to do so, subject to the following conditions: + +The above copyright notice and this permission notice shall be included in all +copies or substantial portions of the Software. + +THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR +IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY, +FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE +AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER +LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM, +OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE +SOFTWARE. diff --git a/README.md b/README.md index e540b53..fe52380 100644 --- a/README.md +++ b/README.md @@ -1,7 +1,7 @@ ![](images/logo.png) --- [![PyPI](https://img.shields.io/pypi/v/ModDotPlot?color=blue&label=PyPI)](https://pypi.org/project/ModDotPlot/) -[![CI](https://github.com/marbl/ModDotPlot/actions/workflows/black.yml/badge.svg)](https://github.com/marbl/ModDotPlot/actions/workflows/black.yml) +[![CI](https://github.com/marbl/ModDotPlot/actions/workflows/ci.yml/badge.svg)](https://github.com/marbl/ModDotPlot/actions/workflows/ci.yml) - [](#) - [Cite](#cite) @@ -37,6 +37,12 @@ If you use ModDotPlot for your research, please cite our software! _ModDotPlot_ is a dot plot visualization tool designed to be used at scale, both for smaller sequences and whole genomes. _ModDotPlot_ is the spiritual successor to [StainedGlass](https://mrvollger.github.io/StainedGlass/). The core algorithm breaks an input sequence down into intervals of sketched *k*-mers called **mod**imizers. This enables the rapid approximation of the Average Nucleotide Identity between combinations of intervals! +Version 1.0.0 uses a bundled [ntHash2](https://github.com/BirolLab/ntHash) implementation for k-mer hashing, replacing the previous `mmh3` runtime dependency. Hash values and exact sketches therefore differ from pre-1.0 releases; regenerate data instead of mixing sketches produced by the two algorithms. Previously saved interactive matrices remain loadable because they contain completed matrices rather than raw hashes. + +FASTA parsing and static BED annotation rendering are also built into ModDotPlot in version 1.0.0, replacing the previous `pysam` and `pyGenomeTracks` runtime dependencies. Plain FASTA, gzip-compressed FASTA, and BGZF-compressed FASTA inputs remain supported, and static annotations produce both PNG and the selected SVG, PDF, or PostScript vector format. + +Static triangle plots, annotation layouts, and multi-sequence grids are now composed directly with Matplotlib. This replaces the previous `CairoSVG`, `svgutils`, and `patchworklib` image-conversion and SVG-composition dependencies while retaining raster and vector output formats. + ![](images/demo.gif) If you're interested in learning more about _ModDotPlot_ and how to visualize tandem repeats, we have an in-depth [YouTube video tutorial](https://www.youtube.com/watch?v=_7sQaljB_ys&t=2321s&pp=ygUXYWxleCBzd2VldGVuIG1vZGRvdHBsb3Q%3D) hosted by the [BioDiversity Genomics Academy](https://thebgacademy.org). @@ -45,7 +51,7 @@ If you're interested in learning more about _ModDotPlot_ and how to visualize ta ## Installation -_ModDotPlot_ can be installed by running `pip install moddotplot`. It requires Python 3.7+ to run. Alternatively, you can download the current release from GitHub by using: +_ModDotPlot_ can be installed by running `pip install moddotplot`. Version 1.0.0 supports Python 3.8 through 3.12; Python 3.13 and newer are not supported by the pinned plotting stack. Alternatively, you can download the current release from GitHub by using: ``` git clone https://github.com/marbl/ModDotPlot.git @@ -74,7 +80,7 @@ Finally, confirm that the installation was installed correctly and that your ver | | | | (_) | (_| | | |__| | (_) | |_ | | | | (_) | |_ |_| |_|\___/ \__,_| |_____/ \___/ \__| |_| |_|\___/ \__| - v0.9.8 + v1.0.0 usage: moddotplot [-h] {interactive,static} ... @@ -111,7 +117,7 @@ Running _ModDotPlot_ in static mode quickly create plots under the specified out ![](images/moddotplot_output.png) -All plots and histograms are output in a vectorized (default: `.svg`) and rasterized `.png` image. [Plotnine](https://plotnine.readthedocs.io/en/v0.12.4/) is the Python plotting library used, with [CairoSVG](https://cairosvg.org) used for converting between image formats. +Plots and histograms are output as both rasterized `.png` images and vector graphics (default: `.svg`). [Plotnine](https://plotnine.readthedocs.io/en/v0.12.4/) provides the primary plotting interface, while Matplotlib directly renders triangle plots, annotation layouts, multi-sequence grids, and each requested output format. _ModDotPlot_ supports highly customizable plotting features in static mode. See [static mode commands](#static-mode-commands) for a complete list of features. @@ -136,7 +142,7 @@ Fasta files to input. Multifasta files are accepted. Interactive mode will only `-b / --bed <.bed file>` -Input bedfile used for dotplot annotation (note: this is not the same as the paired-end bed file produced by ModDotPlot). If selected, this will produce an annotated bedtrack image `_ANNOTATION_TRACK.svg` in static mode, and open an IGV js track in the interactive mode Dash application. The name in the bedfile must match the name of the fasta sequence header in order to produce a correct bed track. +Input bedfile used for dotplot annotation (note: this is not the same as the paired-end bed file produced by ModDotPlot). If selected, this will produce an annotated bedtrack image `_ANNOTATION_TRACK` as PNG and in the selected vector format in static mode, and open an IGV js track in the interactive mode Dash application. The name in the bedfile must match the name of the fasta sequence header in order to produce a correct bed track. `-k / --kmer ` @@ -152,7 +158,7 @@ Minimum sequence identity cutoff threshold. Default is 86. While it is possible `--delta ` -Each partition takes into account a fraction of its neighboring partitions k-mers. This is to avoid sub-optimal identity scores when partitons don't overlap identically. Default is 0.5, and the accepted range is between 0 and 1. Anything greater than 0.5 is not recommended. +Each partition includes a fraction of the adjacent windows' k-mers when estimating identity. This recovers repetitive matches that straddle different window boundaries. The default is 0.5, and the accepted range is between 0 and 1; values greater than 0.5 are not recommended. Set this to 0 only when strictly core-local comparisons are desired. `-m / --modimizer ` @@ -176,7 +182,7 @@ If set when 2 or more sequences are input into ModDotPlot, this will show an A v `--ambiguous ` -By default, k-mers that are homopolymers of ambiguous IUPAC codes (eg. NNNNNNNNNNN’s) are excluded from identity estimation. This results in gaps along the central diagonal for these regions. If desired, these can be kept by setting the `—-ambiguous` flag in both interactive and static mode. +By default, every k-mer window containing a non-ACGTU character is excluded from identity estimation without changing its genomic position. This produces gaps through regions containing ambiguous IUPAC bases. To include deterministic hashes for those windows, set the `--ambiguous` flag in either interactive or static mode. --- @@ -208,9 +214,9 @@ Skip output of histogram legend. Save .bedpe to file, but skip rendering of plots. -`--width ` +`--width ` -Adjust width of self dot plots. Default is 9 inches. +Adjust the output figure width. For a grid this is the width of the complete grid, not each cell. Default is 9 inches. `--dpi ` @@ -232,7 +238,7 @@ Window size. Unlike interactive mode, only one matrix will be created, so this r `--region ` -Plot only a particular range for a given sequence. Syntax is UCSC style (chr:start-end). +Plot only the requested 1-based, inclusive range for each named sequence. Syntax is `FASTA_ID:start-end`; the identifier must exactly match the FASTA header's first whitespace-delimited token. Supply one value per sequence when every grid row and column should be cropped, for example `--region sample.hap1:1-4000000 sample.hap2:1-4000000`. Region limits apply to self plots, pairwise plots, BEDPE coordinates, and grid axes. `--palette ` @@ -246,10 +252,22 @@ Add custom identity threshold breakpoints. Note that the number of breakpoints m Flip sequential order of color palette. Set to `-` by default for divergent palettes. -`--color ` +`--colors ` (legacy alias: `--color`) List of custom colors in hexcode format can be entered sequentially, mapped from low to high identity. +`--plot-direction ` + +With FASTA input, additionally create strand-direction plots. Canonical matches that are also present in a forward-only sketch are blue; canonical-only matches, representing reverse orientation, are pink. This option reads each input in both canonical and forward-only modes and cannot be reconstructed from a loaded BEDPE file. + +`--grid ` + +Create a square grid containing every self comparison on the diagonal and every pairwise comparison off the diagonal. The grid is rendered as one Matplotlib figure and supports three or more input sequences, although large grids become visually dense. + +`--grid-only ` + +Create the comparison grid without writing the individual dotplots. + `-t / --axes-ticks ` Custom tickmarks for x and y axis. Values outside of the `--axes-limits` will not be shown. @@ -331,9 +349,9 @@ Using `samtools faidx` will result in a genomic range being added to a fasta fil #### Adding custom bed file annotations -If providing a custom annotation file using `--bed/b`, _ModDotPlot_ will output additional files: +If providing a custom BED3-BED9 annotation file using `--bed/-b`, _ModDotPlot_ will output additional files: -- An annotation track `_ANNOTATION_TRACK`, containing . Colors for ranges are set using the 9th column of the bedfile. +- A collapsed annotation track `_ANNOTATION_TRACK` in PNG and the selected SVG, PDF, or PostScript vector format. Interval colors use the BED `itemRgb` value in column 9 when present, with a default color for BED3-BED8 records or invalid RGB values. - The annotation track overlayed with a self-identity dotplot `_ANNOTATED` for each sequence present in the annotation track. ``` diff --git a/benchmarks/benchmark_containment_matrix.py b/benchmarks/benchmark_containment_matrix.py new file mode 100644 index 0000000..8e3df71 --- /dev/null +++ b/benchmarks/benchmark_containment_matrix.py @@ -0,0 +1,102 @@ +#!/usr/bin/env python3 +"""Benchmark the exact sparse containment-matrix implementation. + +The generated sketches model the default ``delta=0.5`` layout: each expanded +window contains its core plus half of each adjacent core. This exercises the +same dimensions and sketch cardinalities as a roughly 100 Mb sequence at the +default 1,000-bin resolution without allocating a genome-sized hash array. + +Example:: + + python benchmarks/benchmark_containment_matrix.py + python benchmarks/benchmark_containment_matrix.py --resolution 2000 +""" + +from __future__ import annotations + +import argparse +import resource +import sys +import time +from pathlib import Path + +import numpy as np + +try: + from moddotplot.estimate_identity import ( + pairwiseContainmentMatrix, + selfContainmentMatrix, + ) +except ModuleNotFoundError: # Permit running from an uninstalled source tree. + sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) + from moddotplot.estimate_identity import ( + pairwiseContainmentMatrix, + selfContainmentMatrix, + ) + + +def build_sketches(resolution: int, sketch_size: int, seed: int): + rng = np.random.default_rng(seed) + core = [ + np.unique( + rng.integers(0, np.iinfo(np.uint64).max, sketch_size, dtype=np.uint64) + ) + for _ in range(resolution) + ] + expanded = [] + halfway = sketch_size // 2 + for index, sketch in enumerate(core): + pieces = [sketch] + if index: + pieces.append(core[index - 1][halfway:]) + if index + 1 < resolution: + pieces.append(core[index + 1][:halfway]) + expanded.append(np.unique(np.concatenate(pieces))) + return core, expanded + + +def main(argv=None): + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--resolution", type=int, default=1000) + parser.add_argument("--sketch-size", type=int, default=1612) + parser.add_argument("--identity", type=float, default=86) + parser.add_argument("--kmer", type=int, default=21) + parser.add_argument("--seed", type=int, default=519) + args = parser.parse_args(argv) + + core, expanded = build_sketches(args.resolution, args.sketch_size, args.seed) + compact_bytes = sum(sketch.nbytes for sketch in core + expanded) + + started = time.perf_counter() + self_matrix = selfContainmentMatrix( + core, expanded, args.kmer, args.identity, ambiguous=False + ) + self_seconds = time.perf_counter() - started + + started = time.perf_counter() + pair_matrix = pairwiseContainmentMatrix( + core, + core, + expanded, + expanded, + args.identity, + args.kmer, + supress_progress=True, + ) + pair_seconds = time.perf_counter() - started + + print(f"resolution: {args.resolution}") + print(f"core hashes/window: {args.sketch_size}") + print(f"expanded hashes/window: about {args.sketch_size * 2}") + print(f"prepared sketch memory: {compact_bytes / 2**20:.2f} MiB") + print(f"self matrix: {self_seconds:.3f} s ({self_matrix.shape})") + print(f"pair matrix: {pair_seconds:.3f} s ({pair_matrix.shape})") + print("two-self-plus-pair estimate: " f"{2 * self_seconds + pair_seconds:.3f} s") + peak_rss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss + # macOS reports bytes; Linux and other supported Unix platforms report KiB. + peak_rss_mib = peak_rss / (2**20 if sys.platform == "darwin" else 2**10) + print(f"peak process RSS: {peak_rss_mib:.2f} MiB") + + +if __name__ == "__main__": + main() diff --git a/benchmarks/benchmark_hashing.py b/benchmarks/benchmark_hashing.py new file mode 100644 index 0000000..0c30371 --- /dev/null +++ b/benchmarks/benchmark_hashing.py @@ -0,0 +1,256 @@ +#!/usr/bin/env python3 +"""Compare ModDotPlot's ntHash2 path with the removed mmh3 implementation. + +``mmh3`` is deliberately optional. When it is installed, this utility +recreates the legacy ModDotPlot loop for an apples-to-apples migration +benchmark. Otherwise it reports ntHash2 throughput on its own. + +Examples:: + + python benchmarks/benchmark_hashing.py --length 1000000 --repeats 7 + python benchmarks/benchmark_hashing.py --fasta sequence.fa --kmer 21 + python benchmarks/benchmark_hashing.py --json results.json +""" + +from __future__ import annotations + +import argparse +import gc +import gzip +import json +import random +import statistics +import sys +import time +from pathlib import Path +from typing import Callable, Dict, Iterable, List, Optional, Sequence, Tuple + + +try: + from moddotplot.parse_fasta import _hash_sequence +except ModuleNotFoundError: # Permit running from an uninstalled source tree. + sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) + from moddotplot.parse_fasta import _hash_sequence + + +DNA_ALPHABET = "ACGT" +REVERSE_COMPLEMENT = str.maketrans("ACGT", "TGCA") + + +def _load_mmh3(): + """Return the optional legacy module without making it a dependency.""" + try: + import mmh3 # type: ignore[import-not-found] + except ImportError: + return None + return mmh3 + + +def _legacy_mmh3_hashes( + sequence: str, k: int, canonical: bool, mmh3_module +) -> List[int]: + """Reproduce ModDotPlot's pre-ntHash2 per-k-mer implementation.""" + result = [] + for start in range(max(len(sequence) - k + 1, 0)): + kmer = sequence[start : start + k].upper() + forward = mmh3_module.hash(kmer) + if canonical: + reverse = mmh3_module.hash(kmer[::-1].translate(REVERSE_COMPLEMENT)) + result.append(min(forward, reverse)) + else: + result.append(forward) + return result + + +def _moddotplot_nthash2_hashes(sequence: str, k: int, canonical: bool): + """Exercise the compact batch hashing path used by ModDotPlot's CLI.""" + return _hash_sequence( + sequence, + k, + fw_only=not canonical, + ambiguous=False, + ) + + +def _read_first_fasta(path: Path) -> Tuple[str, str]: + opener = gzip.open if path.suffix == ".gz" else open + name: Optional[str] = None + chunks: List[str] = [] + with opener(path, "rt") as handle: + for line in handle: + line = line.strip() + if line.startswith(">"): + if name is not None: + break + name = line[1:].split()[0] or path.name + elif name is not None: + chunks.append(line) + if name is None: + raise ValueError(f"No FASTA record found in {path}") + return name, "".join(chunks) + + +def _time_call(function: Callable[[], Sequence[int]]) -> Tuple[float, int, int]: + gc.collect() + gc.disable() + try: + started = time.perf_counter_ns() + result = function() + elapsed = (time.perf_counter_ns() - started) / 1e9 + finally: + gc.enable() + count = len(result) + # Touch the result without adding a full O(n) checksum to the timing. + raw_result = getattr(result, "data", result) + checksum = 0 if count == 0 else int(raw_result[0]) ^ int(raw_result[-1]) + return elapsed, count, checksum + + +def _summary(samples: Iterable[float], count: int) -> Dict[str, float]: + values = list(samples) + median = statistics.median(values) + return { + "minimum_seconds": min(values), + "median_seconds": median, + "mean_seconds": statistics.mean(values), + "stdev_seconds": statistics.stdev(values) if len(values) > 1 else 0.0, + "maximum_seconds": max(values), + "median_million_hashes_per_second": count / median / 1_000_000, + } + + +def benchmark(sequence: str, k: int, repeats: int, seed: int) -> Dict[str, object]: + mmh3_module = _load_mmh3() + rng = random.Random(seed) + report: Dict[str, object] = { + "sequence_length": len(sequence), + "kmer_length": k, + "repeats": repeats, + "legacy_mmh3_available": mmh3_module is not None, + "modes": {}, + } + + # Warm native code, imports, and allocators before recording samples. + warm_sequence = sequence[: max(k, min(len(sequence), 10_000))] + _moddotplot_nthash2_hashes(warm_sequence, k, canonical=True) + if mmh3_module is not None: + _legacy_mmh3_hashes(warm_sequence, k, True, mmh3_module) + + for mode, canonical in (("forward", False), ("canonical", True)): + implementations: Dict[str, Callable[[], Sequence[int]]] = { + "nthash2": lambda canonical=canonical: _moddotplot_nthash2_hashes( + sequence, k, canonical + ) + } + if mmh3_module is not None: + implementations[ + "legacy_mmh3" + ] = lambda canonical=canonical: _legacy_mmh3_hashes( + sequence, k, canonical, mmh3_module + ) + + raw: Dict[str, List[float]] = {name: [] for name in implementations} + counts: Dict[str, int] = {} + checksums: Dict[str, int] = {} + for _ in range(repeats): + order = list(implementations) + rng.shuffle(order) + for name in order: + elapsed, count, checksum = _time_call(implementations[name]) + raw[name].append(elapsed) + counts[name] = count + checksums[name] = checksum + + expected_count = max(len(sequence) - k + 1, 0) + if any(count != expected_count for count in counts.values()): + raise RuntimeError( + f"Unexpected k-mer cardinality: expected {expected_count}, got {counts}" + ) + + mode_report: Dict[str, object] = { + "hash_count": expected_count, + "implementations": { + name: { + **_summary(samples, counts[name]), + "raw_seconds": samples, + "checksum": checksums[name], + } + for name, samples in raw.items() + }, + } + if mmh3_module is not None: + summaries = mode_report["implementations"] + mode_report["speedup_over_legacy_mmh3"] = ( + summaries["legacy_mmh3"]["median_seconds"] + / summaries["nthash2"]["median_seconds"] + ) + report["modes"][mode] = mode_report + + return report + + +def _print_report(label: str, report: Dict[str, object]) -> None: + print( + f"Input: {label}; {report['sequence_length']:,} bases; " + f"k={report['kmer_length']}; {report['repeats']} repeats" + ) + print("ntHash2 result: compact NumPy uint64 array; legacy result: Python int list") + if not report["legacy_mmh3_available"]: + print("Legacy comparison: skipped (optional mmh3 is not installed)") + print( + f"{'mode':<10} {'implementation':<14} {'median (s)':>12} " + f"{'Mhash/s':>10} {'speedup':>10}" + ) + for mode, mode_report in report["modes"].items(): + speedup = mode_report.get("speedup_over_legacy_mmh3") + for implementation, stats in mode_report["implementations"].items(): + shown_speedup = ( + f"{speedup:.2f}x" + if implementation == "nthash2" and speedup is not None + else "-" + ) + print( + f"{mode:<10} {implementation:<14} " + f"{stats['median_seconds']:>12.6f} " + f"{stats['median_million_hashes_per_second']:>10.2f} " + f"{shown_speedup:>10}" + ) + + +def main(argv: Optional[Sequence[str]] = None) -> int: + parser = argparse.ArgumentParser(description=__doc__) + source = parser.add_mutually_exclusive_group() + source.add_argument("--fasta", type=Path, help="benchmark the first FASTA record") + source.add_argument( + "--length", type=int, default=1_000_000, help="synthetic sequence length" + ) + parser.add_argument("--kmer", type=int, default=21) + parser.add_argument("--repeats", type=int, default=7) + parser.add_argument("--seed", type=int, default=20260927) + parser.add_argument("--json", type=Path, help="also write raw results as JSON") + args = parser.parse_args(argv) + + if args.kmer <= 0: + parser.error("--kmer must be positive") + if args.length < 0: + parser.error("--length cannot be negative") + if args.repeats <= 0: + parser.error("--repeats must be positive") + + if args.fasta: + label, sequence = _read_first_fasta(args.fasta) + else: + label = f"synthetic(seed={args.seed})" + sequence = "".join( + random.Random(args.seed).choices(DNA_ALPHABET, k=args.length) + ) + + report = benchmark(sequence, args.kmer, args.repeats, args.seed) + _print_report(label, report) + if args.json: + args.json.write_text(json.dumps(report, indent=2) + "\n") + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmarks/benchmark_sketch_reuse.py b/benchmarks/benchmark_sketch_reuse.py new file mode 100644 index 0000000..484bc21 --- /dev/null +++ b/benchmarks/benchmark_sketch_reuse.py @@ -0,0 +1,153 @@ +#!/usr/bin/env python3 +"""Benchmark prepared-sketch reuse for a two-sequence static grid. + +The benchmark isolates the work before matrix comparison. With a fixed window, +the uncached workflow prepares both sequences for their self matrices and then +prepares both again for the pairwise matrix (four preparations). The cache +performs the same access pattern with two preparations and two exact hits. + +Example:: + + python benchmarks/benchmark_sketch_reuse.py --length 1000000 --repeats 5 +""" + +from __future__ import annotations + +import argparse +import gc +import math +import statistics +import sys +import time +from pathlib import Path +from typing import Callable, Dict, List, Sequence + +import numpy as np + +try: + from moddotplot.estimate_identity import ( + ModimizerSketchCache, + prepare_modimizer_sketches, + ) +except ModuleNotFoundError: # Permit running from an uninstalled source tree. + sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) + from moddotplot.estimate_identity import ( + ModimizerSketchCache, + prepare_modimizer_sketches, + ) + + +def _time(function: Callable[[], None]) -> float: + gc.collect() + started = time.perf_counter() + function() + return time.perf_counter() - started + + +def _configuration(length: int, resolution: int, modimizer: int): + window = math.ceil(length / resolution) + effective_modimizer = min(window, modimizer) + raw_sparsity = round(window / effective_modimizer) + if raw_sparsity <= effective_modimizer: + sparsity = 2 ** int(math.log2(raw_sparsity)) + else: + sparsity = 2 ** (int(math.log2(raw_sparsity - 1)) + 1) + return window, sparsity, round(window / sparsity) + + +def benchmark( + length: int, + resolution: int, + modimizer: int, + delta: float, + kmer: int, + repeats: int, +) -> None: + rng = np.random.default_rng(56) + sequences = [ + rng.integers(0, np.iinfo(np.uint64).max, length, dtype=np.uint64) + for _ in range(2) + ] + window, sparsity, expectation = _configuration(length, resolution, modimizer) + + def prepare(sequence): + return prepare_modimizer_sketches( + length, + sequence, + window, + sparsity, + delta, + kmer, + False, + expectation, + ) + + def uncached(): + for index in (0, 1, 1, 0): + prepare(sequences[index]) + + def cached(): + cache = ModimizerSketchCache(max_entries=2) + for index in (0, 1, 1, 0): + cache.get_or_prepare( + index, + length, + sequences[index], + window, + sparsity, + delta, + kmer, + False, + expectation, + ) + cache.clear() + + # Warm NumPy dispatch and allocators before recording samples. + prepare(sequences[0][: min(length, window)]) + samples: Dict[str, List[float]] = {"uncached": [], "cached": []} + for repeat in range(repeats): + order: Sequence[str] = ( + ("uncached", "cached") if repeat % 2 == 0 else ("cached", "uncached") + ) + for name in order: + samples[name].append(_time(uncached if name == "uncached" else cached)) + + uncached_median = statistics.median(samples["uncached"]) + cached_median = statistics.median(samples["cached"]) + print( + f"{length:,} hashes/sequence; window={window:,}; delta={delta}; " + f"resolution={resolution}; repeats={repeats}" + ) + print(f"uncached (4 preparations): {uncached_median:.6f} s") + print(f"cached (2 preparations): {cached_median:.6f} s") + print(f"preparation speedup: {uncached_median / cached_median:.2f}x") + + +def main(argv=None) -> int: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--length", type=int, default=100_000) + parser.add_argument("--resolution", type=int, default=100) + parser.add_argument("--modimizer", type=int, default=100) + parser.add_argument("--delta", type=float, default=0.5) + parser.add_argument("--kmer", type=int, default=21) + parser.add_argument("--repeats", type=int, default=5) + args = parser.parse_args(argv) + if min(args.length, args.resolution, args.modimizer, args.kmer, args.repeats) <= 0: + parser.error( + "length, resolution, modimizer, kmer, and repeats must be positive" + ) + if args.delta < 0: + parser.error("delta must be non-negative") + benchmark( + args.length, + args.resolution, + args.modimizer, + args.delta, + args.kmer, + args.repeats, + ) + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/pyproject.toml b/pyproject.toml index 87556b2..2217c0f 100755 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,23 +4,19 @@ build-backend = "setuptools.build_meta" [project] name = "ModDotPlot" -version = "0.9.9" -requires-python = ">= 3.7" +version = "1.0.0" +requires-python = ">=3.8,<3.13" dependencies = [ - "pysam", "pandas", + "matplotlib", "plotly", "dash", "plotnine==0.12.4", "palettable", - "mmh3", "setproctitle", "numpy", + "scipy", "pillow", - "patchworklib==0.6.3", - "cairosvg", - "pygenometracks", - "svgutils", "cooler", ] authors = [ @@ -42,5 +38,6 @@ moddotplot = "moddotplot.__main__:main" # development dependency groups test = [ "pytest", - "pytest-cov" + "pytest-cov", + "tomli>=1.1; python_version < '3.11'", ] diff --git a/setup.cfg b/setup.cfg new file mode 100644 index 0000000..a57e107 --- /dev/null +++ b/setup.cfg @@ -0,0 +1,7 @@ +[metadata] +license_files = + LICENSE + LICENSE.ntHash + +[bdist_wheel] +py_limited_api = cp38 diff --git a/setup.py b/setup.py new file mode 100644 index 0000000..4123808 --- /dev/null +++ b/setup.py @@ -0,0 +1,24 @@ +import sys + +from setuptools import Extension, setup + + +if sys.platform == "win32": + compile_args = ["/O2", "/std:c++17"] +else: + compile_args = ["-O3", "-std=c++17"] + + +setup( + exclude_package_data={"moddotplot": ["*.cpp"]}, + ext_modules=[ + Extension( + "moddotplot._nthash", + sources=["src/moddotplot/_nthash.cpp"], + define_macros=[("Py_LIMITED_API", "0x03080000")], + extra_compile_args=compile_args, + language="c++", + py_limited_api=True, + ) + ], +) diff --git a/src/moddotplot/_nthash.cpp b/src/moddotplot/_nthash.cpp new file mode 100644 index 0000000..7a2a521 --- /dev/null +++ b/src/moddotplot/_nthash.cpp @@ -0,0 +1,342 @@ +#ifndef Py_LIMITED_API +#define Py_LIMITED_API 0x03080000 +#endif +#define PY_SSIZE_T_CLEAN +#include + +#include +#include +#include + +// The rolling hash kernel below is derived from ntHash v2.4.0, commit +// c26bd4572a19de81e30d55042dbd33c1fd21d4b6. ntHash is distributed under +// the MIT license; see LICENSE.ntHash in the source distribution. + +namespace { + +constexpr uint64_t MASK_33 = (uint64_t{1} << 33) - 1; +constexpr uint64_t MASK_31 = (uint64_t{1} << 31) - 1; +constexpr uint64_t SEED_A = 0x3c8bfbb395c60474ULL; +constexpr uint64_t SEED_C = 0x3193c18562a02b4cULL; +constexpr uint64_t SEED_G = 0x20323ed082572324ULL; +constexpr uint64_t SEED_T = 0x295549f54be24456ULL; +constexpr uint64_t FNV_OFFSET_BASIS = 14695981039346656037ULL; +constexpr uint64_t FNV_PRIME = 1099511628211ULL; + +inline char +upper_ascii(char value) +{ + return value >= 'a' && value <= 'z' + ? static_cast(value - ('a' - 'A')) + : value; +} + +inline uint64_t +seed(char value) +{ + switch (upper_ascii(value)) { + case 'A': return SEED_A; + case 'C': return SEED_C; + case 'G': return SEED_G; + case 'T': + case 'U': return SEED_T; + default: return 0; + } +} + +inline uint64_t +complement_seed(char value) +{ + switch (upper_ascii(value)) { + case 'A': return SEED_T; + case 'C': return SEED_G; + case 'G': return SEED_C; + case 'T': + case 'U': return SEED_A; + default: return 0; + } +} + +inline bool +valid(char value) +{ + return seed(value) != 0; +} + +inline uint64_t +rotate_left_width(uint64_t value, unsigned amount, unsigned width) +{ + amount %= width; + if (amount == 0) { + return value; + } + const uint64_t mask = width == 33 ? MASK_33 : MASK_31; + return ((value << amount) | (value >> (width - amount))) & mask; +} + +inline uint64_t +rotate_right_width(uint64_t value, unsigned amount, unsigned width) +{ + amount %= width; + if (amount == 0) { + return value; + } + const uint64_t mask = width == 33 ? MASK_33 : MASK_31; + return ((value >> amount) | (value << (width - amount))) & mask; +} + +inline uint64_t +srol(uint64_t value, unsigned amount = 1) +{ + const uint64_t low = rotate_left_width(value & MASK_33, amount, 33); + const uint64_t high = rotate_left_width(value >> 33, amount, 31); + return low | (high << 33); +} + +inline uint64_t +sror(uint64_t value) +{ + const uint64_t low = rotate_right_width(value & MASK_33, 1, 33); + const uint64_t high = rotate_right_width(value >> 33, 1, 31); + return low | (high << 33); +} + +inline void +base_hashes(const char* sequence, + unsigned k, + uint64_t& forward, + uint64_t& reverse) +{ + forward = 0; + reverse = 0; + for (unsigned i = 0; i < k; ++i) { + forward = srol(forward) ^ seed(sequence[i]); + reverse = srol(reverse) ^ complement_seed(sequence[k - i - 1]); + } +} + +inline uint64_t +canonical(uint64_t forward, uint64_t reverse) +{ + // Unsigned overflow supplies the modulo-2^64 operation specified by + // ntHash2's canonical hash. + return forward + reverse; +} + +inline unsigned char +normalized_forward_char(char value) +{ + const char normalized = upper_ascii(value); + return static_cast(normalized == 'U' ? 'T' : normalized); +} + +inline unsigned char +normalized_reverse_complement_char(char value) +{ + switch (upper_ascii(value)) { + case 'A': return 'T'; + case 'C': return 'G'; + case 'G': return 'C'; + case 'T': + case 'U': return 'A'; + case 'R': return 'Y'; + case 'Y': return 'R'; + case 'M': return 'K'; + case 'K': return 'M'; + case 'S': return 'S'; + case 'W': return 'W'; + case 'H': return 'D'; + case 'B': return 'V'; + case 'V': return 'B'; + case 'D': return 'H'; + case 'N': return 'N'; + default: return static_cast(upper_ascii(value)); + } +} + +inline uint64_t +fnv1a(const char* sequence, unsigned k, bool reverse_complement) +{ + uint64_t hash = FNV_OFFSET_BASIS; + for (unsigned i = 0; i < k; ++i) { + const unsigned char value = reverse_complement + ? normalized_reverse_complement_char( + sequence[k - i - 1]) + : normalized_forward_char(sequence[i]); + hash ^= value; + hash *= FNV_PRIME; + } + return hash; +} + +PyObject* +hash_kmers(PyObject*, PyObject* args) +{ + const char* sequence = nullptr; + Py_ssize_t sequence_length = 0; + Py_ssize_t requested_k = 0; + int use_canonical = 0; + if (!PyArg_ParseTuple(args, + "s#np:hash_kmers", + &sequence, + &sequence_length, + &requested_k, + &use_canonical)) { + return nullptr; + } + + if (requested_k < 1 || + requested_k > std::numeric_limits::max()) { + PyErr_SetString(PyExc_ValueError, "k must be between 1 and 65535"); + return nullptr; + } + + const unsigned k = static_cast(requested_k); + const Py_ssize_t count = + sequence_length >= requested_k ? sequence_length - requested_k + 1 : 0; + if (count > std::numeric_limits::max() / + static_cast(sizeof(uint64_t))) { + PyErr_SetString(PyExc_OverflowError, "hash output is too large"); + return nullptr; + } + + bool has_ambiguous_base = false; + if (count > 0) { + Py_BEGIN_ALLOW_THREADS + for (Py_ssize_t i = 0; i < sequence_length; ++i) { + if (!valid(sequence[i])) { + has_ambiguous_base = true; + break; + } + } + Py_END_ALLOW_THREADS + } + + PyObject* hash_bytes = PyBytes_FromStringAndSize( + nullptr, count * static_cast(sizeof(uint64_t))); + if (hash_bytes == nullptr) { + return nullptr; + } + PyObject* mask_bytes = PyBytes_FromStringAndSize( + nullptr, has_ambiguous_base ? count : 0); + if (mask_bytes == nullptr) { + Py_DECREF(hash_bytes); + return nullptr; + } + + char* hash_output = PyBytes_AsString(hash_bytes); + char* mask_output = + has_ambiguous_base ? PyBytes_AsString(mask_bytes) : nullptr; + if (hash_output == nullptr || (has_ambiguous_base && mask_output == nullptr)) { + Py_DECREF(hash_bytes); + Py_DECREF(mask_bytes); + return nullptr; + } + + if (count > 0) { + Py_BEGIN_ALLOW_THREADS + unsigned invalid_count = 0; + for (unsigned i = 0; i < k; ++i) { + invalid_count += !valid(sequence[i]); + } + + uint64_t forward = 0; + uint64_t reverse = 0; + bool rolling = false; + + for (Py_ssize_t position = 0; position < count; ++position) { + if (position > 0) { + const bool previous_window_was_valid = invalid_count == 0; + invalid_count -= !valid(sequence[position - 1]); + invalid_count += !valid(sequence[position + requested_k - 1]); + + if (invalid_count == 0 && previous_window_was_valid) { + const char outgoing = sequence[position - 1]; + const char incoming = sequence[position + requested_k - 1]; + forward = + srol(forward) ^ seed(incoming) ^ srol(seed(outgoing), k); + reverse ^= srol(complement_seed(incoming), k); + reverse ^= complement_seed(outgoing); + reverse = sror(reverse); + rolling = true; + } else if (invalid_count != 0) { + rolling = false; + } + } + + const bool ambiguous = invalid_count != 0; + uint64_t value = 0; + if (!ambiguous) { + if (!rolling) { + base_hashes(sequence + position, k, forward, reverse); + rolling = true; + } + value = use_canonical ? canonical(forward, reverse) : forward; + } else { + const uint64_t fallback_forward = + fnv1a(sequence + position, k, false); + value = use_canonical + ? canonical(fallback_forward, + fnv1a(sequence + position, k, true)) + : fallback_forward; + } + + std::memcpy(hash_output + position * sizeof(uint64_t), + &value, + sizeof(value)); + if (mask_output != nullptr) { + mask_output[position] = static_cast(ambiguous); + } + } + Py_END_ALLOW_THREADS + } + + PyObject* result = PyTuple_Pack(2, hash_bytes, mask_bytes); + Py_DECREF(hash_bytes); + Py_DECREF(mask_bytes); + return result; +} + +PyMethodDef methods[] = { + { "hash_kmers", + hash_kmers, + METH_VARARGS, + "hash_kmers(sequence, k, canonical, /)\n--\n\n" + "Hash every positional k-mer in one native call.\n\n" + "Return (hashes, ambiguity_mask). hashes contains packed native-endian " + "uint64 values. ambiguity_mask is empty when every window is valid; " + "otherwise it contains one byte per window, with 1 marking a window " + "that contains a non-ACGTU character." }, + { nullptr, nullptr, 0, nullptr }, +}; + +PyModuleDef module = { + PyModuleDef_HEAD_INIT, + "_nthash", + "Stable-ABI batch bindings for the vendored ntHash2 rolling hash.", + -1, + methods, +}; + +} // namespace + +PyMODINIT_FUNC +PyInit__nthash() +{ + PyObject* created_module = PyModule_Create(&module); + if (created_module == nullptr) { + return nullptr; + } + if (PyModule_AddStringConstant( + created_module, "ALGORITHM", "ntHash_v2") < 0 || + PyModule_AddStringConstant( + created_module, "UPSTREAM_VERSION", "2.4.0") < 0 || + PyModule_AddStringConstant( + created_module, + "UPSTREAM_COMMIT", + "c26bd4572a19de81e30d55042dbd33c1fd21d4b6") < 0) { + Py_DECREF(created_module); + return nullptr; + } + return created_module; +} diff --git a/src/moddotplot/const.py b/src/moddotplot/const.py index 692ee54..d30b1e1 100755 --- a/src/moddotplot/const.py +++ b/src/moddotplot/const.py @@ -1,4 +1,4 @@ -VERSION = "0.9.9" +VERSION = "1.0.0" COLS = [ "#query_name", "query_start", @@ -9,7 +9,7 @@ "perID_by_events", ] -ASCII_ART = """ +ASCII_ART = r""" __ __ _ _____ _ _____ _ _ | \/ | | | | __ \ | | | __ \| | | | | \ / | ___ __| | | | | | ___ | |_ | |__) | | ___ | |_ diff --git a/src/moddotplot/estimate_identity.py b/src/moddotplot/estimate_identity.py index 7b76b71..01602f8 100644 --- a/src/moddotplot/estimate_identity.py +++ b/src/moddotplot/estimate_identity.py @@ -1,5 +1,7 @@ #!/usr/bin/env python3 import math +from collections import OrderedDict +from dataclasses import dataclass import numpy as np from moddotplot.const import ( SEQUENTIAL_PALETTES, @@ -7,52 +9,174 @@ QUALITATIVE_PALETTES, ) from palettable import colorbrewer -from typing import List, Set, Dict, Tuple -import mmh3 +from typing import Collection, Hashable, List, Set, Dict, Tuple import pandas as pd import cooler +from scipy.sparse import csr_matrix from moddotplot.parse_fasta import printProgressBar -def removeAmbiguousBases(mod_list, k): - # Ambiguous IUPAC codes - bases_to_remove = ["R", "Y", "M", "K", "S", "W", "H", "B", "V", "D", "N"] - kmers_to_remove = set() - for i in range(len(bases_to_remove)): - result_string = str(bases_to_remove[i]) * k - kmers_to_remove.add(mmh3.hash(result_string)) - mod_set = set(mod_list) - # Remove homopolymers of ambiguous nucleotides - mod_set.difference_update(kmers_to_remove) - return mod_set +@dataclass(frozen=True) +class PreparedModimizerSketches: + """Core and expanded window sketches prepared for a matrix calculation. + The compact sorted NumPy arrays are deliberately retained by reference. A + prepared value can therefore be shared by self and pairwise calculations + without copying sketch data. Public conversion helpers continue to return + sets for backward compatibility. + """ -def createSelfMatrix( + core: List[Collection[int]] + neighbors: List[Collection[int]] + + +def prepare_modimizer_sketches( sequence_length, sequence, window_size, sparsity, delta, k, - identity, ambiguous, - sketch_size, + expectation, ): - no_neighbors = partitionOverlaps(sequence, window_size, 0, sequence_length, k) + """Partition one sequence and build the sketches used by matrix routines.""" + + core_partitions = partitionOverlaps(sequence, window_size, 0, sequence_length, k) if delta > 0: - neighbors = partitionOverlaps(sequence, window_size, delta, sequence_length, k) + neighbor_partitions = partitionOverlaps( + sequence, window_size, delta, sequence_length, k + ) else: - neighbors = no_neighbors + neighbor_partitions = core_partitions + + # Prepared values stay as compact sorted ndarrays. The public conversion + # helpers still return sets for API compatibility, but keeping millions of + # selected hashes as Python ints in Python hash tables costs roughly ten + # times more memory on chromosome-sized inputs. + core = _convert_to_modimizer_arrays( + core_partitions, sparsity, ambiguous, k, expectation + ) + if neighbor_partitions is core_partitions: + neighbors = core + else: + neighbors = _convert_to_modimizer_arrays( + neighbor_partitions, sparsity, ambiguous, k, expectation + ) + return PreparedModimizerSketches(core=core, neighbors=neighbors) + - neighbors_mods = convertToModimizers(neighbors, sparsity, ambiguous, k, sketch_size) - no_neighbors_mods = convertToModimizers( - no_neighbors, sparsity, ambiguous, k, sketch_size +def create_self_matrix_from_sketches(prepared, k, identity, ambiguous): + """Build a self matrix from an already prepared sequence sketch.""" + + return selfContainmentMatrix( + prepared.core, prepared.neighbors, k, identity, ambiguous ) - matrix = selfContainmentMatrix( - no_neighbors_mods, neighbors_mods, k, identity, ambiguous + + +def create_pairwise_matrix_from_sketches( + prepared_x, prepared_y, identity, k, supress_progress=False +): + """Build a pairwise matrix from two already prepared sequence sketches.""" + + return pairwiseContainmentMatrix( + prepared_x.core, + prepared_y.core, + prepared_x.neighbors, + prepared_y.neighbors, + identity, + k, + supress_progress, ) - return matrix + + +class ModimizerSketchCache: + """Small LRU cache for prepared sequence sketches. + + ``source_key`` identifies the underlying sequence slice. Calculation + parameters are incorporated automatically, preventing reuse when a + pairwise plot uses a different window size from the corresponding self + plot. The default capacity matches a two-sequence comparison and bounds + retained memory for larger grids. + """ + + def __init__(self, max_entries=2): + if max_entries < 1: + raise ValueError("max_entries must be at least one") + self.max_entries = max_entries + self._cache = OrderedDict() + + def get_or_prepare( + self, + source_key: Hashable, + sequence_length, + sequence, + window_size, + sparsity, + delta, + k, + ambiguous, + expectation, + ): + cache_key = ( + source_key, + sequence_length, + window_size, + sparsity, + delta, + k, + ambiguous, + expectation, + ) + try: + prepared = self._cache.pop(cache_key) + except KeyError: + # Evict before constructing the replacement so a miss cannot + # briefly exceed the configured memory bound by one full sketch. + if len(self._cache) >= self.max_entries: + self._cache.popitem(last=False) + prepared = prepare_modimizer_sketches( + sequence_length, + sequence, + window_size, + sparsity, + delta, + k, + ambiguous, + expectation, + ) + self._cache[cache_key] = prepared + return prepared + + def clear(self): + """Release references to all retained sketches.""" + + self._cache.clear() + + +def createSelfMatrix( + sequence_length, + sequence, + window_size, + sparsity, + delta, + k, + identity, + ambiguous, + sketch_size, +): + prepared = prepare_modimizer_sketches( + sequence_length, + sequence, + window_size, + sparsity, + delta, + k, + ambiguous, + sketch_size, + ) + return create_self_matrix_from_sketches(prepared, k, identity, ambiguous) def createPairwiseMatrix( @@ -68,110 +192,207 @@ def createPairwiseMatrix( ambiguous, expectation, ): - no_neighbors_large = partitionOverlaps(larger_seq, window_size, 0, larger_length, k) - no_neighbors_small = partitionOverlaps( - smaller_seq, window_size, 0, smaller_length, k - ) - if delta > 0: - neighbors_large = partitionOverlaps( - larger_seq, window_size, delta, larger_length, k - ) - neighbors_small = partitionOverlaps( - smaller_seq, window_size, delta, smaller_length, k - ) - else: - neighbors_large = no_neighbors_large - neighbors_small = no_neighbors_small - - neighbors_mods_large = convertToModimizers( - neighbors_large, sparsity, ambiguous, k, expectation - ) - no_neighbors_mods_large = convertToModimizers( - no_neighbors_large, sparsity, ambiguous, k, expectation - ) - neighbors_mods_small = convertToModimizers( - neighbors_small, sparsity, ambiguous, k, expectation - ) - no_neighbors_mods_small = convertToModimizers( - no_neighbors_small, sparsity, ambiguous, k, expectation + prepared_large = prepare_modimizer_sketches( + larger_length, + larger_seq, + window_size, + sparsity, + delta, + k, + ambiguous, + expectation, ) - matrix = pairwiseContainmentMatrix( - no_neighbors_mods_large, - no_neighbors_mods_small, - neighbors_mods_large, - neighbors_mods_small, - identity, + prepared_small = prepare_modimizer_sketches( + smaller_length, + smaller_seq, + window_size, + sparsity, + delta, k, - False, + ambiguous, + expectation, + ) + return create_pairwise_matrix_from_sketches( + prepared_large, prepared_small, identity, k ) - return matrix def partitionOverlaps( lst: List[int], win: int, delta: float, seq_len: int, k: int ) -> List[List[int]]: - kmer_list = [] - kmer_to_genomic_coordinate_offset = win - k + 1 + if win <= 0: + raise ValueError("window size must be greater than zero") + if k <= 0: + raise ValueError("k-mer size must be greater than zero") + + kmer_count = min(len(lst), max(seq_len, 0)) + if kmer_count == 0: + return [] + + # A sequence with n bases has n - k + 1 k-mers. Reconstruct the genomic + # length so that every partition starts on a multiple of ``win``. The old + # counter started the second partition at ``win - k + 2`` and then advanced + # by ``win``, making all but the first window begin k - 2 bases too early. + sequence_length = kmer_count + k - 1 delta_offset = win * delta + kmer_list = [] + + # A trailing genomic fragment shorter than k has no k-mer and therefore no + # matrix cell, so iterate over valid k-mer starts rather than base length. + for window_start in range(0, kmer_count, win): + window_end = min(window_start + win, sequence_length) + expanded_start = max(0, int(round(window_start - delta_offset))) + expanded_end = min(sequence_length, int(round(window_end + delta_offset))) + + # K-mers are indexed by their genomic start. Subtracting k - 1 from + # the right boundary excludes k-mers that cross out of the interval. + start_index = min(expanded_start, kmer_count) + end_index = min(max(expanded_end - k + 1, start_index), kmer_count) + kmer_list.append(lst[start_index:end_index]) - # Set the first window to contain win - k + 1 kmers. - starting_end_index = int(round(kmer_to_genomic_coordinate_offset + delta_offset)) - kmer_list.append(lst[0:starting_end_index]) - counter = win - k + 1 - - # Set normal windows - while counter <= (seq_len - win): - start_index = counter + 1 - end_index = win + counter + 1 - delta_start_index = int(round(start_index - delta_offset)) - delta_end_index = int(round(end_index + delta_offset)) - if delta_end_index > seq_len: - delta_end_index = seq_len - try: - kmer_list.append(lst[delta_start_index:delta_end_index]) - except Exception as e: - print("Error in appending list of kmers...\n") - print(e) - kmer_list.append(lst[delta_start_index:seq_len]) - counter += win - - # Set the last window to get the remainder - if counter <= seq_len - 2: - final_start_index = int(round(counter + 1 - delta_offset)) - kmer_list.append(lst[final_start_index:seq_len]) - - # Test that last value was added on correctly - try: - assert kmer_list[-1][-1] == lst[-1] - except (AssertionError, IndexError) as e: - print(f"Error: Last k-mer does not match original sequence: {e}\n") return kmer_list +def _valid_hashes(partition): + """Return the unmasked hashes in *partition* as a flat NumPy array. + + FASTA input uses a ``uint64`` ``MaskedArray``. Keeping that representation + here is important: iterating over a masked array produces a Python object + for every genomic k-mer and was the dominant cost of sketch construction. + The object-array fallback retains compatibility with legacy callers that + pass lists containing ``None`` or masked scalar values. + """ + + if isinstance(partition, np.ma.MaskedArray): + values = np.asarray(partition.data).reshape(-1) + raw_mask = partition.mask + # ``getmaskarray`` materializes a full all-False array for ``nomask``. + # Avoid allocating and scanning that temporary for every 100 kb window. + if raw_mask is np.ma.nomask or (np.ndim(raw_mask) == 0 and not bool(raw_mask)): + return values + mask = np.asarray(raw_mask, dtype=bool).reshape(-1) + if np.any(mask): + values = values[~mask] + return values + + values = np.asarray(partition) + # NumPy may infer ``float64`` for a Python list containing uint64-range + # integers on some supported versions, irreversibly rounding hash values. + # Send non-integral legacy sequences through the exact Python-int fallback + # below; native numeric arrays can still be converted in bulk. + legacy_requires_exact_conversion = ( + not isinstance(partition, np.ndarray) and values.dtype.kind not in "biu" + ) + if values.dtype.kind != "O" and not legacy_requires_exact_conversion: + values = values.reshape(-1) + # ``populateModimizers`` historically called int() on every value. + # Hash arrays are already integral, but retain that behavior for older + # callers that provide a floating-point sequence. + if values.dtype.kind not in "biu": + values = values.astype(np.int64) + return values + + # Mixed Python sequences cannot be converted to a numeric array until + # ``None`` and masked sentinels have been removed. This path is for API + # compatibility; normal FASTA processing always takes the vectorized path + # above. + object_values = np.asarray(partition, dtype=object).reshape(-1) + valid_values = [ + int(kmer) + for kmer in object_values + if kmer is not None and not np.ma.is_masked(kmer) + ] + if not valid_values: + return np.empty(0, dtype=np.uint64) + + try: + return np.asarray(valid_values, dtype=np.uint64) + except (OverflowError, ValueError): + # Preserve arbitrary-size and negative Python integers for legacy + # callers. NumPy still performs the modulo and uniqueness operations + # in bulk, albeit with object arithmetic. + return np.asarray(valid_values, dtype=object) + + +def _divisible_hashes(values, sparsity): + """Select hashes divisible by an integer sparsity without Python loops.""" + + if values.size == 0: + return values + + # All production sparsities are powers of two. A bit mask avoids creating + # the full-size temporary remainder array and is valid for signed and + # unsigned integer hashes alike. + if values.dtype.kind in "biu" and sparsity & (sparsity - 1) == 0: + return values[np.bitwise_and(values, sparsity - 1) == 0] + return values[np.remainder(values, sparsity) == 0] + + +def _populate_modimizer_array(partition, sparsity, ambiguous, expectation, k): + """Build one adaptive modimizer sketch using vectorized NumPy operations. + + ``ambiguous`` and ``k`` remain in the public signature for compatibility; + ambiguity is represented by the mask on ``partition``. If a sketch is + smaller than half its expectation, sparsity is repeatedly halved just as + before. The sparsity now stays integer throughout that fallback rather + than becoming a float after the first recursive call. + """ + + del ambiguous, k + + current_sparsity = int(sparsity) + if current_sparsity < 1: + raise ValueError("sparsity must be a positive integer") + + values = _valid_hashes(partition) + minimum_size = round(expectation / 2) + + while True: + selected = _divisible_hashes(values, current_sparsity) + unique = np.unique(selected) + if unique.size >= minimum_size or current_sparsity == 1: + return unique + current_sparsity = max(1, current_sparsity // 2) + + def populateModimizers(partition, sparsity, ambiguous, expectation, k): - mod_set = set() - for kmer in partition: - if kmer % sparsity == 0: - mod_set.add(kmer) - if not ambiguous: - mod_set = removeAmbiguousBases(mod_set, k) - if (len(mod_set) < round(expectation / 2)) and (sparsity > 1): - populateModimizers(partition, sparsity / 2, ambiguous, expectation, k) - return mod_set + """Build one adaptive sketch and return its historical ``set`` type.""" + + unique = _populate_modimizer_array(partition, sparsity, ambiguous, expectation, k) + # ``tolist`` converts NumPy integer scalars to Python ints, preserving the + # public return type while internal prepared sketches retain compact arrays. + return set(unique.tolist()) + + +def _convert_to_modimizer_arrays( + kmer_list, sparsity: int, ambiguous: bool, k: int, expectation: int +): + return [ + _populate_modimizer_array(partition, sparsity, ambiguous, expectation, k) + for partition in kmer_list + ] def convertToModimizers( kmer_list: List[List[int]], sparsity: int, ambiguous: bool, k: int, expectation: int ) -> List[Set[int]]: - mod_total = [] - for partition in kmer_list: - mod_set = populateModimizers(partition, sparsity, ambiguous, expectation, k) - mod_total.append(mod_set) - return mod_total + return [ + populateModimizers(partition, sparsity, ambiguous, expectation, k) + for partition in kmer_list + ] def convertMatrixToBed( - matrix, window_size, id_threshold, x_name, y_name, self_identity, x_offset, y_offset + matrix, + window_size, + id_threshold, + x_name, + y_name, + self_identity, + x_offset, + y_offset, + x_end=None, + y_end=None, ): bed = [ ( @@ -196,6 +417,15 @@ def convertMatrixToBed( start_y = y * window_size + y_offset end_y = start_y + window_size - 1 + # The final matrix window can be shorter than ``window_size``. + # Keep exported coordinates inside the exact sequence or + # requested region instead of allowing the last tile to + # overhang it. + if x_end is not None: + end_x = min(end_x, x_end) + if y_end is not None: + end_y = min(end_y, y_end) + bed.append( ( x_name, @@ -296,41 +526,174 @@ def containment_neighbors( k: int, ) -> float: """ - Calculate the containment neighbors based on four sets and an identity threshold. + Calculate symmetric containment using the opposite expanded window. + + ``set1`` and ``set2`` are sketches of the two core windows, while ``set3`` + and ``set4`` are the corresponding sketches expanded by ``delta``. A core + is compared with the *other* window's expanded sketch in both directions. + This is what lets a repeat that straddles a partition boundary match a core + window instead of being missed solely because the partitions are offset. + + The maximum directional containment is thresholded after both directions + have been evaluated. Applying the cutoff to only the first direction made + the result depend on argument order in the original implementation. Args: set1 (Set[int]): The first set. set2 (Set[int]): The second set. - set3 (Set[int]): The third set. - set4 (Set[int]): The fourth set. + set3 (Set[int]): Expanded sketch corresponding to ``set1``. + set4 (Set[int]): Expanded sketch corresponding to ``set2``. identity (int): The identity threshold. k (int): Kmer value. Returns: float: The containment neighbors value. """ - len_a = len(set1) - len_b = len(set2) - - intersection_a_b_prime = len(set1 & set4) - if len_a != 0: - containment_a_b_prime = intersection_a_b_prime / len_a - else: - # If len_a is zero, handle it by setting containment_a_b_prime to a default value - containment_a_b_prime = 0 + containment_a_b_expanded = len(set1 & set4) / len(set1) if set1 else 0.0 + containment_b_a_expanded = len(set2 & set3) / len(set2) if set2 else 0.0 + symmetric_containment = max(containment_a_b_expanded, containment_b_a_expanded) - if binomial_distance(containment_a_b_prime, k) < identity / 100: + if binomial_distance(symmetric_containment, k) < identity / 100: return 0.0 + return symmetric_containment + +def _sketch_intersection_counts(sketches_a, sketches_b): + """Return every pairwise sketch intersection count as a dense array. + + The previous matrix implementation performed two Python set intersections + for every output cell. At the default resolution that means roughly two + million intersections per matrix, each scanning about 1,600 hashes. Here + hashes are coordinate-compressed once and the sketches are represented as + sparse incidence matrices. Sparse matrix multiplication then calculates + the exact same intersection counts in compiled code. + + Only the small ``len(sketches_a) x len(sketches_b)`` result is dense. The + incidence matrices retain one entry per selected hash, so chromosome size + does not create a dense hash universe. Production ntHash values take the + fast unsigned-64-bit path; the mapping fallback preserves exact behavior + for legacy callers using negative, oversized, or other hashable values. + """ + + rows = len(sketches_a) + cols = len(sketches_b) + lengths_a = np.fromiter( + (len(sketch) for sketch in sketches_a), dtype=np.int64, count=rows + ) + lengths_b = np.fromiter( + (len(sketch) for sketch in sketches_b), dtype=np.int64, count=cols + ) + total_a = int(lengths_a.sum()) + total_b = int(lengths_b.sum()) + total = total_a + total_b + + count_dtype = np.int32 + if total == 0: + return np.zeros((rows, cols), dtype=count_dtype) + + collections = (sketches_a, sketches_b) + compact_uint64_arrays = all( + isinstance(sketch, np.ndarray) + and sketch.ndim == 1 + and sketch.dtype == np.dtype(np.uint64) + for sketches in collections + for sketch in sketches + ) + + if compact_uint64_arrays: + # This is the normal prepared-sketch path. Concatenating the arrays in + # C avoids boxing millions of hashes back into Python integers. + hashes = np.concatenate( + tuple(sketch for sketches in collections for sketch in sketches) + ) + unique_hashes, inverse = np.unique(hashes, return_inverse=True) + hash_count = len(unique_hashes) + del hashes, unique_hashes else: - intersection_a_prime_b = len(set2 & set3) - if len_b != 0: - containment_a_prime_b = intersection_a_prime_b / len_b - else: - # If len_a is zero, handle it by setting containment_a_b_prime to a default value - containment_a_prime_b = 0 + uint64_max = np.iinfo(np.uint64).max + uint64_compatible = all( + isinstance(value, (int, np.integer)) and 0 <= int(value) <= uint64_max + for sketches in collections + for sketch in sketches + for value in sketch + ) + + if not compact_uint64_arrays and uint64_compatible: + hashes = np.fromiter( + ( + int(value) + for sketches in collections + for sketch in sketches + for value in sketch + ), + dtype=np.uint64, + count=total, + ) + unique_hashes, inverse = np.unique(hashes, return_inverse=True) + hash_count = len(unique_hashes) + del hashes, unique_hashes + elif not compact_uint64_arrays: + # This compatibility path is not used by FASTA processing, but keeps + # the public matrix functions exact for arbitrary hashable set values. + hash_columns = {} + inverse = np.empty(total, dtype=np.int64) + position = 0 + for sketches in collections: + for sketch in sketches: + for value in sketch: + try: + column = hash_columns[value] + except KeyError: + column = len(hash_columns) + hash_columns[value] = column + inverse[position] = column + position += 1 + hash_count = len(hash_columns) + + # SciPy uses 32-bit sparse indices whenever dimensions and nonzero counts + # fit. Keeping that representation saves tens of megabytes at r=1000. + index_dtype = ( + np.int32 + if hash_count <= np.iinfo(np.int32).max and total <= np.iinfo(np.int32).max + else np.int64 + ) + indices_a = inverse[:total_a].astype(index_dtype, copy=False) + indices_b = inverse[total_a:].astype(index_dtype, copy=False) + del inverse + + indptr_a = np.empty(rows + 1, dtype=index_dtype) + indptr_b = np.empty(cols + 1, dtype=index_dtype) + indptr_a[0] = 0 + indptr_b[0] = 0 + np.cumsum(lengths_a, out=indptr_a[1:]) + np.cumsum(lengths_b, out=indptr_b[1:]) + + incidence_a = csr_matrix( + ( + np.ones(total_a, dtype=count_dtype), + indices_a, + indptr_a, + ), + shape=(rows, hash_count), + ) + incidence_b = csr_matrix( + ( + np.ones(total_b, dtype=count_dtype), + indices_b, + indptr_b, + ), + shape=(cols, hash_count), + ) + return (incidence_a @ incidence_b.T).toarray() + - return max(containment_a_b_prime, containment_a_prime_b) +def _identity_matrix_from_containment(containment_matrix, identity, k): + """Apply ModDotPlot's k-mer identity transform and cutoff in place.""" + + identities = np.power(containment_matrix, 1.0 / k) + identities[identities < identity / 100] = 0.0 + identities *= 100.0 + return identities def selfContainmentMatrix( @@ -352,35 +715,29 @@ def selfContainmentMatrix( np.ndarray: A NumPy array representing the self-containment matrix. """ n = len(mod_set) - progress_thresholds = round(n / 77) - - if progress_thresholds == 0: - progress_thresholds = 1 - + if len(mod_set_neighbors) != n: + raise IndexError("core and expanded self sketches must have equal lengths") printProgressBar(0, n, prefix="Progress:", suffix="Complete", length=40) - containment_matrix = np.empty((n, n)) - - for w in range(n): - if w % progress_thresholds == 0: - printProgressBar(w, n, prefix="Progress:", suffix="Complete", length=40) - containment_matrix[w, w] = 100.0 - if len(mod_set[w]) == 0 and not ambiguous: - containment_matrix[w, w] = 0 - - for r in range(w + 1, n): - c_hat = binomial_distance( - containment_neighbors( - mod_set[w], - mod_set[r], - mod_set_neighbors[w], - mod_set_neighbors[r], - identity, - k, - ), - k, - ) - containment_matrix[r, w] = c_hat * 100.0 - containment_matrix[w, r] = c_hat * 100.0 + intersection_counts = _sketch_intersection_counts(mod_set, mod_set_neighbors) + core_sizes = np.fromiter((len(sketch) for sketch in mod_set), dtype=float, count=n) + directional_containment = np.zeros((n, n), dtype=float) + np.divide( + intersection_counts, + core_sizes[:, np.newaxis], + out=directional_containment, + where=core_sizes[:, np.newaxis] != 0, + ) + symmetric_containment = np.maximum( + directional_containment, directional_containment.T + ) + containment_matrix = _identity_matrix_from_containment( + symmetric_containment, identity, k + ) + + diagonal = np.full(n, 100.0) + if not ambiguous: + diagonal[core_sizes == 0] = 0.0 + np.fill_diagonal(containment_matrix, diagonal) printProgressBar( n, n, prefix="Progress:", suffix="Completed", length=40 @@ -390,10 +747,10 @@ def selfContainmentMatrix( def pairwiseContainmentMatrix( - mod_set_x: List[int], - mod_set_y: List[int], - mod_set_x_neighbors: List[List[int]], - mod_set_y_neighbors: List[List[int]], + mod_set_x: List[Set[int]], + mod_set_y: List[Set[int]], + mod_set_x_neighbors: List[Set[int]], + mod_set_y_neighbors: List[Set[int]], identity: int, k: int, supress_progress: bool, @@ -402,53 +759,61 @@ def pairwiseContainmentMatrix( Calculate an updated identity matrix using specified parameters. Args: - mod_set_x (List[int]): List of values for the x-axis. - mod_set_y (List[int]): List of values for the y-axis. - mod_set_x_neighbors (List[List[int]]): List of lists representing neighbors of mod_set_x values. - mod_set_y_neighbors (List[List[int]]): List of lists representing neighbors of mod_set_y values. + mod_set_x (List[Set[int]]): Modimizer sets for columns on the x-axis. + mod_set_y (List[Set[int]]): Modimizer sets for rows on the y-axis. + mod_set_x_neighbors (List[Set[int]]): Neighbor sets for x-axis windows. + mod_set_y_neighbors (List[Set[int]]): Neighbor sets for y-axis windows. identity (int): Resolution parameter. k (int): Value for the k parameter in the binomial_distance function. supress_progress (bool): if true supresses the progress bar Returns: - np.ndarray: An identity matrix containing containment values. + np.ndarray: A ``(len(mod_set_y), len(mod_set_x))`` identity matrix. """ - n = max(len(mod_set_y), len(mod_set_x)) - progress_thresholds = round(n / 77) - if progress_thresholds == 0: - progress_thresholds = 1 + rows = len(mod_set_y) + cols = len(mod_set_x) + if len(mod_set_x_neighbors) != cols or len(mod_set_y_neighbors) != rows: + raise IndexError("core and expanded pairwise sketches must have equal lengths") + if not supress_progress: + printProgressBar(0, rows, prefix="Progress:", suffix="Complete", length=40) + x_core_sizes = np.fromiter( + (len(sketch) for sketch in mod_set_x), dtype=float, count=cols + ) + y_core_sizes = np.fromiter( + (len(sketch) for sketch in mod_set_y), dtype=float, count=rows + ) + # core X against expanded Y, transposed into the public (Y, X) layout. + x_to_y_counts = _sketch_intersection_counts(mod_set_x, mod_set_y_neighbors).T + x_to_y = np.zeros((rows, cols), dtype=float) + np.divide( + x_to_y_counts, + x_core_sizes[np.newaxis, :], + out=x_to_y, + where=x_core_sizes[np.newaxis, :] != 0, + ) if not supress_progress: - printProgressBar(0, n, prefix="Progress:", suffix="Complete", length=40) - containment_matrix = np.zeros((n, n), dtype=float) - - for w in range(len(mod_set_y)): - if not supress_progress: - if w % progress_thresholds == 0: - printProgressBar(w, n, prefix="Progress:", suffix="Complete", length=40) - for q in range(n): - try: - containment_matrix[w, q] = ( - binomial_distance( - containment_neighbors( - mod_set_x[q], - mod_set_y[w], - mod_set_x_neighbors[q], - mod_set_y_neighbors[w], - identity, - k, - ), - k, - ) - * 100.0 - ) - # Bandaid solution for too sequences that are too small. - except IndexError as e: - pass + printProgressBar( + rows // 2, rows, prefix="Progress:", suffix="Complete", length=40 + ) + + # core Y against expanded X already has the public (Y, X) orientation. + y_to_x_counts = _sketch_intersection_counts(mod_set_y, mod_set_x_neighbors) + y_to_x = np.zeros((rows, cols), dtype=float) + np.divide( + y_to_x_counts, + y_core_sizes[:, np.newaxis], + out=y_to_x, + where=y_core_sizes[:, np.newaxis] != 0, + ) + symmetric_containment = np.maximum(x_to_y, y_to_x) + containment_matrix = _identity_matrix_from_containment( + symmetric_containment, identity, k + ) if not supress_progress: printProgressBar( - n, n, prefix="Progress:", suffix="Completed", length=40 + rows, rows, prefix="Progress:", suffix="Completed", length=40 ) # show completed progress bar print("\n") return containment_matrix diff --git a/src/moddotplot/interactive.py b/src/moddotplot/interactive.py index 5f8dfde..4e207b5 100644 --- a/src/moddotplot/interactive.py +++ b/src/moddotplot/interactive.py @@ -21,6 +21,55 @@ log.setLevel(logging.ERROR) +def figure_to_bed(figure, default_identity=86.0): + """Convert the heatmap in a Dash figure into BEDPE rows. + + The interactive callback receives a JSON-compatible Plotly figure rather + than the original matrix metadata, so derive the window and coordinate + offsets from the trace axes. Keeping this logic outside the callback also + makes the export path independently testable. + """ + trace = figure["data"][0] + matrix = np.asarray(trace["z"], dtype=float) + + positive_values = matrix[np.isfinite(matrix) & (matrix > 0)] + identity = ( + float(np.min(positive_values)) + if positive_values.size + else float(default_identity) + ) + + x_values = np.asarray(trace.get("x", []), dtype=float) + y_values = np.asarray(trace.get("y", []), dtype=float) + window_size = float(trace.get("dx", 0)) + if x_values.size > 1: + window_size = float(x_values[1] - x_values[0]) + elif y_values.size > 1: + window_size = float(y_values[1] - y_values[0]) + if window_size <= 0: + raise ValueError("Unable to determine a positive window size from the plot") + + x_offset = float(x_values[0]) if x_values.size else float(trace.get("x0", 0)) + y_offset = float(y_values[0]) if y_values.size else float(trace.get("y0", 0)) + layout = figure.get("layout", {}) + x_name = layout.get("xaxis", {}).get("title", {}).get("text", "x") + y_name = layout.get("yaxis", {}).get("title", {}).get("text", "y") + self_identity = x_name == y_name + filename = f"{x_name}.bedpe" if self_identity else f"{x_name}-{y_name}.bedpe" + + rows = convertMatrixToBed( + matrix, + window_size, + identity, + x_name, + y_name, + self_identity, + x_offset, + y_offset, + ) + return rows, filename + + def find_closest_elements(value, sorted_list): # Initialize variables to store the indices of the closest elements closest_index1 = None @@ -290,7 +339,7 @@ def halving_sequence(size, start): ), html.Div( html.Button( - "Save Matrix to Bed File", + "Save Matrix to BEDPE File", id="save-bed", n_clicks=0, disabled=False, @@ -1345,29 +1394,10 @@ def update_button_state(text_value): def save_bed(n_clicks, figure): global clicked_values if n_clicks > 0: - window_size = figure["data"][0]["x"][1] - figure["data"][0]["x"][0] try: - identity = round( - min([val for val in figure["data"][0]["z"][0] if val > 0]) - ) - x_axis_name = figure["layout"]["xaxis"]["title"]["text"] - y_axis_name = figure["layout"]["yaxis"]["title"]["text"] - if x_axis_name == y_axis_name: - selfy = True - title_hi = x_axis_name + ".bed" - else: - selfy = False - title_hi = x_axis_name + "-" + y_axis_name + ".bed" - except: - identity = 86 - x_axis_name = "x" - y_axis_name = "y" - selfy = False - title_hi = "x-y.bed" - pls = np.array(figure["data"][0]["z"]) - tr = convertMatrixToBed( - pls, window_size, identity, x_axis_name, y_axis_name, selfy - ) + tr, title_hi = figure_to_bed(figure) + except (KeyError, TypeError, ValueError) as error: + return f"Unable to save BEDPE file: {error}", 0 if not output_dir: bedfile_output = os.path.join("./", title_hi) else: @@ -1377,7 +1407,7 @@ def save_bed(n_clicks, figure): with open(bedfile_output, "w") as bedfile: for row in tr: bedfile.write("\t".join(map(str, row)) + "\n") - msg = f"Saved bed file to {bedfile_output}\n" + msg = f"Saved BEDPE file to {bedfile_output}\n" return msg, 0 # Make sure to return a tuple of values for the outputs else: return ( diff --git a/src/moddotplot/moddotplot.py b/src/moddotplot/moddotplot.py index 9667ce7..626522d 100755 --- a/src/moddotplot/moddotplot.py +++ b/src/moddotplot/moddotplot.py @@ -1,6 +1,7 @@ #!/usr/bin/env python3 import sys from moddotplot.parse_fasta import ( + HASH_ALGORITHM, readKmersFromFile, getInputHeaders, isValidFasta, @@ -16,6 +17,9 @@ convertMatrixToCool, createSelfMatrix, createPairwiseMatrix, + create_self_matrix_from_sketches, + create_pairwise_matrix_from_sketches, + ModimizerSketchCache, partitionOverlaps, ) from moddotplot.interactive import run_dash @@ -23,13 +27,36 @@ import argparse import math -from moddotplot.static_plots import read_df_from_file, create_plots, create_grid import json import numpy as np import pickle import os +# Static plotting pulls in the Plotnine and Matplotlib stacks. Keep those +# imports behind the static command boundary so ``--help`` and interactive +# mode do not pay their startup cost. +read_df_from_file = None +create_plots = None +create_grid = None +create_direction_plot = None + + +def _load_static_plotting(): + global read_df_from_file, create_plots, create_grid, create_direction_plot + + from moddotplot import static_plots + + if read_df_from_file is None: + read_df_from_file = static_plots.read_df_from_file + if create_plots is None: + create_plots = static_plots.create_plots + if create_grid is None: + create_grid = static_plots.create_grid + if create_direction_plot is None: + create_direction_plot = static_plots.create_direction_plot + + def get_parser(): """ Argument parsing for stand-alone runs. @@ -40,7 +67,7 @@ def get_parser(): description="ModDotPlot: Visualization of Tandem Repeats", ) subparsers = parser.add_subparsers( - dest="command", help="Choose mode: interactive or static" + dest="command", required=True, help="Choose mode: interactive or static" ) interactive_parser = subparsers.add_parser( "interactive", help="Interactive mode commands" @@ -112,7 +139,7 @@ def get_parser(): "--delta", default=0.5, type=float, - help="Fraction of neighboring partition to include in identity estimation. Must be between 0 and 1, use > 0.5 is not recommended.", + help="Fraction of each neighboring window included when estimating identity. Default: 0.5.", ) interactive_parser.add_argument( @@ -151,7 +178,13 @@ def get_parser(): interactive_parser.add_argument( "--ambiguous", action="store_true", - help="Preserve diagonal when handling strings of ambiguous homopolymers (eg. long runs of N's).", + help="Include k-mer windows containing non-ACGTU IUPAC bases instead of masking them.", + ) + + interactive_parser.add_argument( + "--forward", + action="store_true", + help="Enforce forward only k-mers instead of canonical k-mers. Warning: only use if you want strand-specific output!", ) interactive_parser.add_argument( @@ -249,7 +282,7 @@ def get_parser(): "--delta", default=0.5, type=float, - help="Fraction of neighboring partition to include in identity estimation. Must be between 0 and 1, use > 0.5 is not recommended.", + help="Fraction of each neighboring window included when estimating identity. Default: 0.5.", ) static_parser.add_argument( @@ -330,6 +363,8 @@ def get_parser(): static_parser.add_argument( "--colors", + "--color", + dest="colors", default=None, nargs="+", help="Use a custom color palette, entered in either hexcode or rgb format.", @@ -374,7 +409,7 @@ def get_parser(): static_parser.add_argument( "--ambiguous", action="store_true", - help="Preserve diagonal when handling strings of ambiguous homopolymers (eg. long runs of N's).", + help="Include k-mer windows containing non-ACGTU IUPAC bases instead of masking them.", ) static_parser.add_argument( @@ -405,10 +440,122 @@ def get_parser(): return parser +def _apply_static_config(args, config): + """Apply static-mode JSON configuration values to parsed arguments.""" + # TODO: Remove args that are interactive only + args.fasta = config.get("fasta") + args.load = config.get("load") + args.bed = config.get("bed") + + # Distance matrix commands + args.kmer = config.get("kmer", args.kmer) + args.modimizer = config.get("modimizer", args.modimizer) + args.resolution = config.get("resolution", args.resolution) + args.window = config.get("window", args.window) + args.region = config.get("region", args.region) + args.identity = config.get("identity", args.identity) + args.delta = config.get("delta", args.delta) + args.output_dir = config.get("output_dir", args.output_dir) + args.compare = config.get("compare", args.compare) + args.compare_only = config.get("compare_only", args.compare_only) + args.compare_order = config.get("compare_order", args.compare_order) + + args.cooler = config.get("cooler", args.cooler) + args.no_bedpe = config.get("no_bedpe", args.no_bedpe) + args.no_plot = config.get("no_plot", args.no_plot) + args.no_hist = config.get("no_hist", args.no_hist) + args.width = config.get("width", args.width) + args.axes_limits = config.get("axes_limits", args.axes_limits) + args.dpi = config.get("dpi", args.dpi) + args.palette = config.get("palette", args.palette) + args.palette_orientation = config.get( + "palette_orientation", args.palette_orientation + ) + args.colors = config.get("colors", config.get("color", args.colors)) + args.axes_ticks = config.get("axes_ticks", args.axes_ticks) + args.axes_number = config.get("axes_number", args.axes_number) + args.breakpoints = config.get("breakpoints", args.breakpoints) + args.bin_freq = config.get("bin_freq", args.bin_freq) + args.forward = config.get("forward", args.forward) + args.plot_direction = config.get("plot_direction", args.plot_direction) + args.ambiguous = config.get("ambiguous", args.ambiguous) + args.grid = config.get("grid", args.grid) + args.grid_only = config.get("grid_only", args.grid_only) + args.vector = config.get("vector", args.vector) + args.deraster = config.get("deraster", args.deraster) + + return args + + +def _parse_region_arguments(region_arguments, sequence_names): + """Validate CLI regions and index them by their exact FASTA identifier.""" + + if not region_arguments: + return {} + if isinstance(region_arguments, str): + region_arguments = [region_arguments] + + available_names = { + parsed[0] if (parsed := extractRegion(name)) else name + for name in sequence_names + } + regions = {} + for value in region_arguments: + parsed = extractRegion(value) + if not parsed: + raise ValueError(f"invalid region {value!r}; expected FASTA_ID:start-end") + sequence_name, start, end = parsed + if start < 1 or end < start: + raise ValueError( + f"invalid region {value!r}; coordinates must satisfy 1 <= start <= end" + ) + if sequence_name not in available_names: + raise ValueError(f"region {value!r} does not match any FASTA identifier") + if sequence_name in regions: + raise ValueError( + f"multiple regions were provided for FASTA identifier {sequence_name!r}" + ) + regions[sequence_name] = parsed + return regions + + +def _slice_kmers_for_region(kmers, region, kmer_size, sequence_start=1): + """Return the exact k-mer slice for a 1-based inclusive base interval.""" + + _sequence_name, start, end = region + sequence_base_length = len(kmers) + kmer_size - 1 + sequence_end = sequence_start + sequence_base_length - 1 + if start < sequence_start or end > sequence_end: + raise ValueError( + f"region {start}-{end} is outside the available interval " + f"{sequence_start}-{sequence_end}" + ) + region_base_length = end - start + 1 + if region_base_length < kmer_size: + raise ValueError( + f"region length {region_base_length} is shorter than k-mer size {kmer_size}" + ) + + # A base interval [start, end] contains k-mers beginning at genomic + # positions start through end - k + 1, inclusive. Translate those positions + # into the available sequence's zero-based coordinates. + local_start = start - sequence_start + local_stop = end - sequence_start - kmer_size + 2 + selected = kmers[local_start:local_stop] + expected_count = region_base_length - kmer_size + 1 + if len(selected) != expected_count: + raise ValueError( + f"region produced {len(selected)} k-mers; expected {expected_count}" + ) + return selected + + def main(): print(ASCII_ART) print(f"v{VERSION} \n") args = get_parser().parse_args() + if args.command == "static": + _load_static_plotting() # -----------MUTUALLY EXCLUSIVE: INTERACTIVE OR STATIC MODE----------- if args.command == "interactive": print(f"Running ModDotPlot in interactive mode\n") @@ -454,42 +601,16 @@ def main(): if args.config: with open(args.config, "r") as f: config = json.load(f) - # TODO: Remove args that are interactive only - args.fasta = config.get("fasta") - args.load = config.get("load") - args.bed = config.get("bed") - - # Distance matrix commands - args.kmer = config.get("kmer", args.kmer) - args.modimizer = config.get("modimizer", args.modimizer) - args.resolution = config.get("resolution", args.resolution) - args.window = config.get("window", args.window) - args.identity = config.get("identity", args.identity) - args.delta = config.get("delta", args.delta) - args.output_dir = config.get("output_dir", args.output_dir) - args.compare = config.get("compare", args.compare) - args.compare_only = config.get("compare_only", args.compare_only) - - args.no_bedpe = config.get("no_bedpe", args.no_bedpe) - args.no_plot = config.get("no_plot", args.no_plot) - args.no_hist = config.get("no_hist", args.no_hist) - args.width = config.get("width", args.width) - args.axes_limits = config.get("axes_limits", args.axes_limits) - args.dpi = config.get("dpi", args.dpi) - args.palette = config.get("palette", args.palette) - args.palette_orientation = config.get( - "palette_orientation", args.palette_orientation - ) - args.colors = config.get("color", args.colors) - args.axes_ticks = config.get("axes_ticks", args.axes_ticks) - args.breakpoints = config.get("breakpoints", args.breakpoints) - args.bin_freq = config.get("bin_freq", args.bin_freq) - args.axes_limits = config.get("axes_limits", args.axes_limits) - args.axes_ticks = config.get("axes_ticks", args.axes_ticks) - args.vector = config.get("vector", args.vector) - args.deraster = config.get("deraster", args.deraster) + _apply_static_config(args, config) # -----------INPUT COMMAND VALIDATION----------- + if args.plot_direction and getattr(args, "load", None): + print( + "Error: --plot-direction requires FASTA input because strand " + "orientation cannot be recovered from a BEDPE file.\n" + ) + sys.exit(2) + # TODO: More tests! if args.breakpoints: # Check that start value for breakpoints = identity threshold value @@ -582,7 +703,7 @@ def main(): is_freq=args.bin_freq, xlim=args.axes_limits, custom_colors=args.colors, - custom_breakpoints=args.axes_ticks, + custom_breakpoints=args.breakpoints, from_file=df, is_pairwise=True, axes_labels=args.axes_ticks, @@ -615,23 +736,25 @@ def main(): is_freq=args.bin_freq, xlim=xlim_val_grid, custom_colors=args.colors, - custom_breakpoints=args.axes_ticks, + custom_breakpoints=args.breakpoints, axes_label=args.axes_ticks, is_bed=True, width=args.width, breaks=args.axes_ticks, deraster=args.deraster, vector_format=args.vector, + dpi=args.dpi, ) sys.exit(0) # -----------INPUT SEQUENCE VALIDATION----------- seq_list = [] fasta_list = args.fasta.copy() + fasta_headers = {} for i in args.fasta: try: - isValidFasta(i) headers = getInputHeaders(i) + fasta_headers[i] = headers if len(headers) > 1: print(f"File {i} contains multiple fasta entries.\n") @@ -644,14 +767,65 @@ def main(): ) fasta_list.remove(i) + try: + region_by_name = _parse_region_arguments( + getattr(args, "region", None), seq_list + ) + except ValueError as error: + print(f"Error: {error}.\n") + sys.exit(2) + # -----------LOAD SEQUENCES INTO MEMORY----------- kmer_list = [] for i in fasta_list: if args.forward: - kmer_list.append(readKmersFromFile(i, args.kmer, False, True)) + kmer_list.append( + readKmersFromFile( + i, + args.kmer, + False, + True, + args.ambiguous, + region_by_name, + fasta_headers[i], + ) + ) else: - kmer_list.append(readKmersFromFile(i, args.kmer, False, False)) + kmer_list.append( + readKmersFromFile( + i, + args.kmer, + False, + False, + args.ambiguous, + region_by_name, + fasta_headers[i], + ) + ) k_list = [item for sublist in kmer_list for item in sublist] + + # Direction plots need both canonical and forward-only hashes. Load the + # opposite representation only when requested so normal runs retain their + # existing memory footprint. + direction_k_list = None + if args.command == "static" and args.plot_direction: + alternate_kmer_list = [ + readKmersFromFile( + path, + args.kmer, + False, + not args.forward, + args.ambiguous, + region_by_name, + fasta_headers[path], + ) + for path in fasta_list + ] + direction_k_list = [item for sublist in alternate_kmer_list for item in sublist] + if len(direction_k_list) != len(k_list): + raise ValueError( + "Canonical and forward-only FASTA parsing produced different sequence counts" + ) # Throw error if compare only selected with one sequence. if len(k_list) < 2 and args.compare_only: print( @@ -661,8 +835,8 @@ def main(): # -----------LAUNCH INTERACTIVE MODE----------- if args.command == "interactive": - # Single sequence, can set window length immediately. - hgi = len(max(k_list)) + # Use the longest sequence to size the shared interactive image pyramid. + hgi = max(len(kmers) for kmers in k_list) hgi = hgi + args.kmer - 1 min_window_size = 0 window_lengths = [] @@ -784,6 +958,8 @@ def main(): "max_window_size": window_lengths[-1], "resolution": args.resolution, "kmer_length": args.kmer, + "hash_algorithm": HASH_ALGORITHM, + "format_version": 2, "title": f"{seq_list[j]}", "sparsities": sparsities, } @@ -878,6 +1054,8 @@ def main(): "max_window_size": window_lengths[-1], "resolution": args.resolution, "kmer_length": args.kmer, + "hash_algorithm": HASH_ALGORITHM, + "format_version": 2, "title": f"{larger_name}-{smaller_name}", "sparsities": sparsities, } @@ -940,13 +1118,43 @@ def main(): if args.grid or args.grid_only: grid_val_singles = [] grid_val_single_names = [] - new_sequences = list(zip(seq_list, k_list)) + if direction_k_list is None: + new_sequences = list(zip(seq_list, k_list)) + else: + new_sequences = list(zip(seq_list, k_list, direction_k_list)) if args.compare_order == "size": sequences = sorted(new_sequences, key=lambda seq: len(seq[1]), reverse=True) else: sequences = new_sequences + + # Record exact base-coordinate bounds independently of sparse hits. + # Renderers must not infer these bounds from BEDPE rows: thresholding + # can remove edge windows, and a partial final window can extend past + # the selected interval. + selected_intervals = [] + for sequence in sequences: + sequence_name = sequence[0] + header_range = extractRegion(sequence_name) + base_name = header_range[0] if header_range else sequence_name + selected_range = region_by_name.get(base_name) + if selected_range: + interval_start, interval_end = selected_range[1:] + else: + interval_start = int(header_range[1]) if header_range else 1 + interval_end = interval_start + len(sequence[1]) + args.kmer - 2 + selected_intervals.append((interval_start, interval_end)) + grid_axis_bounds = ( + min(start for start, _end in selected_intervals), + max(end for _start, end in selected_intervals), + ) + sketch_cache = ( + ModimizerSketchCache(max_entries=2) if args.grid or args.grid_only else None + ) if len(sequences) > 6 and (args.grid or args.grid_only): - print("Too many sequences to create a grid. Skipping. \n") + print( + f"Creating a large {len(sequences)}x{len(sequences)} grid; " + "rendering may take additional time and memory.\n" + ) # Create output directory, if doesn't exist: if (args.output_dir) and not os.path.exists(args.output_dir): @@ -954,55 +1162,49 @@ def main(): # -----------COMPUTE SELF-IDENTITY PLOTS----------- if not args.compare_only: for i in range(len(sequences)): - seq_length = len(sequences[i][1]) - seq_name = sequences[i][0] - seq_range = extractRegion(seq_name) + sequence_name = sequences[i][0] + header_range = extractRegion(sequence_name) + base_name = header_range[0] if header_range else sequence_name + sequence_start = int(header_range[1]) if header_range else 1 + seq_range = region_by_name.get(base_name) + matrix_sequence = sequences[i][1] + alternate_sequence = sequences[i][2] if args.plot_direction else None + if seq_range: - seq_name = seq_range[0] - # If region, then I only want to use the subsequence. - try: - if args.region: - subseq_start_pos = None - subseq_end_pos = None - for region in args.region: - chrom, lower_bound, upper_bound = extractRegion(region) - if chrom == seq_name: - subseq_start_pos = lower_bound - subseq_end_pos = upper_bound - seq_start_pos = lower_bound - # Validate bounds - if subseq_start_pos < 1 or subseq_end_pos > seq_length: - print( - f"Error: region {region} is out of bounds for {seq_name}. Will use entire sequence.\n" - ) - subseq_start_pos = 1 - subseq_end_pos = seq_length - seq_name = sequences[i][0] - break - print( - f"Using region {seq_name}:{subseq_start_pos}-{subseq_end_pos}\n" - ) - # Change sequence length, and use a subsequence instead. - seq_length = ( - subseq_end_pos - subseq_start_pos + 1 - args.kmer - ) - seq_range = seq_name, subseq_start_pos, subseq_end_pos - seq_name = ( - f"{seq_name}:{subseq_start_pos}-{subseq_end_pos}" - ) - if not subseq_end_pos or not subseq_start_pos: - print( - f"Error: region {args.region} not found in {seq_name}. Will use entire sequence.\n" + selected_sequence_start = seq_range[1] + try: + matrix_sequence = _slice_kmers_for_region( + matrix_sequence, + seq_range, + args.kmer, + sequence_start=selected_sequence_start, + ) + if args.plot_direction: + alternate_sequence = _slice_kmers_for_region( + alternate_sequence, + seq_range, + args.kmer, + sequence_start=selected_sequence_start, ) - seq_range = None - except Exception as e: - print( - f"Error obtaining region for {seq_name}. Will use entire sequence: {e}\n" - ) - if not seq_range: - seq_start_pos = 1 + except ValueError as error: + print(f"Error: invalid region for {base_name}: {error}.\n") + sys.exit(2) + _, seq_start_pos, subseq_end_pos = seq_range + seq_name = f"{base_name}:{seq_start_pos}-{subseq_end_pos}" + print(f"Using region {seq_name}\n") else: - seq_start_pos = int(seq_range[1]) + seq_start_pos = sequence_start + subseq_end_pos = ( + seq_start_pos + len(matrix_sequence) + args.kmer - 2 + ) + seq_name = sequence_name + + plot_axis_bounds = args.axes_limits or ( + seq_start_pos, + subseq_end_pos, + ) + + seq_length = len(matrix_sequence) win = args.window res = args.resolution if args.window: @@ -1034,13 +1236,11 @@ def main(): print(f"\tWindow size w: {win}\n") print(f"\tModimizer sketch size: {expectation}\n") print(f"\tPlot Resolution r: {res}\n") - if args.region and seq_range: - subseq = sequences[i][1][ - subseq_start_pos : (subseq_end_pos - args.kmer + 1) - ] + + if sketch_cache is None: self_mat = createSelfMatrix( seq_length, - subseq, + matrix_sequence, win, seq_sparsity, args.delta, @@ -1050,9 +1250,30 @@ def main(): expectation, ) else: - self_mat = createSelfMatrix( + source_region = (seq_range[1], seq_range[2]) if seq_range else None + prepared_self = sketch_cache.get_or_prepare( + (i, source_region), + seq_length, + matrix_sequence, + win, + seq_sparsity, + args.delta, + args.kmer, + args.ambiguous, + expectation, + ) + self_mat = create_self_matrix_from_sketches( + prepared_self, args.kmer, args.identity, args.ambiguous + ) + # The cache owns the reusable reference. Keeping this loop + # local alive can pin an evicted sketch until all grid + # calculations finish. + del prepared_self + direction_self_mat = None + if args.plot_direction and not args.no_plot and not args.grid_only: + direction_self_mat = createSelfMatrix( seq_length, - sequences[i][1], + alternate_sequence, win, seq_sparsity, args.delta, @@ -1070,6 +1291,8 @@ def main(): True, seq_start_pos, seq_start_pos, + subseq_end_pos, + subseq_end_pos, ) if args.grid or args.grid_only: grid_val_singles.append(bed) @@ -1102,17 +1325,13 @@ def main(): except Exception as e: print(f"Error creating cooler file: {e}") + bedpe_path = os.path.join(args.output_dir or ".", seq_name) + if (not args.no_bedpe) or ((not args.no_plot) and (not args.grid_only)): + os.makedirs(bedpe_path, exist_ok=True) + if not args.no_bedpe: # Log saving bed file - bedpe_path = "." - if not args.output_dir: - bedpe_path = os.path.join(bedpe_path, seq_name) - os.makedirs(bedpe_path, exist_ok=True) - bedfile_output = os.path.join(seq_name, seq_name + ".bedpe") - else: - bedpe_path = os.path.join(args.output_dir, seq_name) - os.makedirs(bedpe_path, exist_ok=True) - bedfile_output = os.path.join(bedpe_path, seq_name + ".bedpe") + bedfile_output = os.path.join(bedpe_path, seq_name + ".bedpe") with open(bedfile_output, "w") as bedfile: for row in bed: @@ -1133,7 +1352,7 @@ def main(): width=args.width, dpi=args.dpi, is_freq=args.bin_freq, - xlim=args.axes_limits, + xlim=plot_axis_bounds, custom_colors=args.colors, custom_breakpoints=args.breakpoints, from_file=None, @@ -1144,6 +1363,30 @@ def main(): deraster=args.deraster, annotation=args.bed, ) + if args.plot_direction: + if args.forward: + canonical_matrix = direction_self_mat + forward_matrix = self_mat + else: + canonical_matrix = self_mat + forward_matrix = direction_self_mat + create_direction_plot( + canonical_matrix=canonical_matrix, + forward_matrix=forward_matrix, + window_size=win, + directory=bedpe_path, + name_x=seq_name, + name_y=seq_name, + self_identity=True, + width=args.width, + dpi=args.dpi, + vector_format=args.vector, + deraster=args.deraster, + xlim=plot_axis_bounds, + axes_labels=args.axes_ticks, + x_offset=seq_start_pos, + y_offset=seq_start_pos, + ) # -----------COMPUTE COMPARATIVE PLOTS----------- # TODO: Optimize computations so that largest sequence doesn't need to be redone all the time @@ -1155,111 +1398,106 @@ def main(): if args.grid or args.grid_only: grid_val_doubles = [] grid_val_double_names = [] - xlim_val_grid = 0 + xlim_val_grid = args.axes_limits or grid_axis_bounds for i in range(len(sequences)): for j in range(i + 1, len(sequences)): # Larger = x, smaller = y. This is pre-sorted earlier. larger_seq = sequences[i][1] smaller_seq = sequences[j][1] - larger_length = len(larger_seq) - smaller_length = len(smaller_seq) - larger_seq_name = sequences[i][0] - smaller_seq_name = sequences[j][0] - larger_seq_range = extractRegion(larger_seq_name) - if not larger_seq_range: - larger_seq_start_pos = 1 - else: - larger_seq_start_pos = int(larger_seq_range[1]) - larger_seq_name = larger_seq_range[0] - smaller_seq_range = extractRegion(smaller_seq_name) - if not smaller_seq_range: - smaller_seq_start_pos = 1 - else: - smaller_seq_start_pos = int(smaller_seq_range[1]) - smaller_seq_name = smaller_seq_range[0] + larger_direction_seq = ( + sequences[i][2] if args.plot_direction else None + ) + smaller_direction_seq = ( + sequences[j][2] if args.plot_direction else None + ) + larger_sequence_name = sequences[i][0] + smaller_sequence_name = sequences[j][0] + larger_header_range = extractRegion(larger_sequence_name) + smaller_header_range = extractRegion(smaller_sequence_name) + larger_base_name = ( + larger_header_range[0] + if larger_header_range + else larger_sequence_name + ) + smaller_base_name = ( + smaller_header_range[0] + if smaller_header_range + else smaller_sequence_name + ) + larger_sequence_start = ( + int(larger_header_range[1]) if larger_header_range else 1 + ) + smaller_sequence_start = ( + int(smaller_header_range[1]) if smaller_header_range else 1 + ) + larger_seq_range = region_by_name.get(larger_base_name) + smaller_seq_range = region_by_name.get(smaller_base_name) + larger_subseq = larger_seq + smaller_subseq = smaller_seq + larger_direction_subseq = larger_direction_seq + smaller_direction_subseq = smaller_direction_seq + larger_seq_start_pos = larger_sequence_start + smaller_seq_start_pos = smaller_sequence_start + larger_seq_end_pos = ( + larger_seq_start_pos + len(larger_subseq) + args.kmer - 2 + ) + smaller_seq_end_pos = ( + smaller_seq_start_pos + len(smaller_subseq) + args.kmer - 2 + ) + larger_seq_name = larger_sequence_name + smaller_seq_name = smaller_sequence_name try: - if args.region: - subseq_start_pos = None - subseq_end_pos = None - for region in args.region: - chrom, lower_bound, upper_bound = extractRegion(region) - if chrom == larger_seq_name: - larger_subseq_start_pos = lower_bound - larger_subseq_end_pos = upper_bound - larger_seq_start_pos = lower_bound - # Validate bounds - if ( - larger_seq_start_pos < 1 - or larger_subseq_end_pos > larger_length - ): - print( - f"Error: region {region} is out of bounds for {larger_seq_name}. Will use entire sequence.\n" - ) - larger_subseq_start_pos = 1 - larger_subseq_end_pos = larger_length - larger_seq_name = sequences[i][0] - break - print( - f"Using region {larger_seq_name}:{larger_subseq_start_pos}-{larger_subseq_end_pos}\n" - ) - # Change sequence length, and use a subsequence instead. - larger_length = ( - larger_subseq_end_pos - - larger_subseq_start_pos - + 1 - - args.kmer - ) - larger_seq_range = ( - larger_seq_name, - larger_subseq_start_pos, - larger_subseq_end_pos, - ) - larger_seq_name = f"{larger_seq_name}:{larger_subseq_start_pos}-{larger_subseq_end_pos}" - - if chrom == smaller_seq_name: - smaller_subseq_start_pos = lower_bound - smaller_subseq_end_pos = upper_bound - smaller_seq_start_pos = lower_bound - # Validate bounds - if ( - smaller_seq_start_pos < 1 - or smaller_subseq_end_pos > smaller_length - ): - print( - f"Error: region {region} is out of bounds for {smaller_seq_name}. Will use entire sequence.\n" - ) - smaller_subseq_start_pos = 1 - smaller_subseq_end_pos = smaller_length - smaller_seq_name = sequences[j][0] - break - print( - f"Using region {smaller_seq_name}:{smaller_subseq_start_pos}-{smaller_subseq_end_pos}\n" - ) - # Change sequence length, and use a subsequence instead. - smaller_length = ( - smaller_subseq_end_pos - - smaller_subseq_start_pos - + 1 - - args.kmer - ) - smaller_seq_range = ( - smaller_seq_name, - smaller_subseq_start_pos, - smaller_subseq_end_pos, - ) - smaller_seq_name = f"{smaller_seq_name}:{smaller_subseq_start_pos}-{smaller_subseq_end_pos}" - # This is wrong. Might be fine to leave alone - if not larger_subseq_end_pos or not larger_subseq_start_pos: - print( - f"Error: region {args.region} not found in {seq_name}. Will use entire sequence.\n" + if larger_seq_range: + selected_larger_start = larger_seq_range[1] + larger_subseq = _slice_kmers_for_region( + larger_seq, + larger_seq_range, + args.kmer, + sequence_start=selected_larger_start, + ) + if args.plot_direction: + larger_direction_subseq = _slice_kmers_for_region( + larger_direction_seq, + larger_seq_range, + args.kmer, + sequence_start=selected_larger_start, ) - seq_range = None - except Exception as e: - print( - f"Error obtaining region for {seq_name}. Will use entire sequence: {e}\n" - ) + _, larger_seq_start_pos, larger_end = larger_seq_range + larger_seq_end_pos = larger_end + larger_seq_name = f"{larger_base_name}:{larger_seq_start_pos}-{larger_end}" + print(f"Using region {larger_seq_name}\n") + + if smaller_seq_range: + selected_smaller_start = smaller_seq_range[1] + smaller_subseq = _slice_kmers_for_region( + smaller_seq, + smaller_seq_range, + args.kmer, + sequence_start=selected_smaller_start, + ) + if args.plot_direction: + smaller_direction_subseq = _slice_kmers_for_region( + smaller_direction_seq, + smaller_seq_range, + args.kmer, + sequence_start=selected_smaller_start, + ) + _, smaller_seq_start_pos, smaller_end = smaller_seq_range + smaller_seq_end_pos = smaller_end + smaller_seq_name = f"{smaller_base_name}:{smaller_seq_start_pos}-{smaller_end}" + print(f"Using region {smaller_seq_name}\n") + except ValueError as error: + print(f"Error: invalid comparison region: {error}.\n") + sys.exit(2) + + larger_length = len(larger_subseq) + smaller_length = len(smaller_subseq) + pair_axis_bounds = args.axes_limits or ( + min(larger_seq_start_pos, smaller_seq_start_pos), + max(larger_seq_end_pos, smaller_seq_end_pos), + ) win = args.window res = args.resolution @@ -1290,25 +1528,7 @@ def main(): print(f"\tModimizer sketch size: {expectation}\n") print(f"\tPlot Resolution r: {res}\n") - if args.region and (larger_seq_range or smaller_seq_range): - if larger_seq_range: - larger_subseq = larger_seq[ - larger_subseq_start_pos : ( - larger_subseq_end_pos - args.kmer + 1 - ) - ] - else: - larger_subseq = larger_seq - - if smaller_seq_range: - smaller_subseq = smaller_seq[ - smaller_subseq_start_pos : ( - smaller_subseq_end_pos - args.kmer + 1 - ) - ] - else: - smaller_subseq = smaller_seq - + if sketch_cache is None: pair_mat = createPairwiseMatrix( smaller_length, larger_length, @@ -1323,11 +1543,54 @@ def main(): expectation, ) else: - pair_mat = createPairwiseMatrix( + smaller_source_region = ( + (smaller_seq_range[1], smaller_seq_range[2]) + if smaller_seq_range + else None + ) + larger_source_region = ( + (larger_seq_range[1], larger_seq_range[2]) + if larger_seq_range + else None + ) + prepared_smaller = sketch_cache.get_or_prepare( + (j, smaller_source_region), smaller_length, + smaller_subseq, + win, + seq_sparsity, + args.delta, + args.kmer, + args.ambiguous, + expectation, + ) + prepared_larger = sketch_cache.get_or_prepare( + (i, larger_source_region), larger_length, - smaller_seq, - larger_seq, + larger_subseq, + win, + seq_sparsity, + args.delta, + args.kmer, + args.ambiguous, + expectation, + ) + pair_mat = create_pairwise_matrix_from_sketches( + prepared_smaller, + prepared_larger, + args.identity, + args.kmer, + ) + # Avoid retaining entries after the bounded cache + # evicts or clears them. + del prepared_smaller, prepared_larger + direction_pair_mat = None + if args.plot_direction and not args.no_plot and not args.grid_only: + direction_pair_mat = createPairwiseMatrix( + smaller_length, + larger_length, + smaller_direction_subseq, + larger_direction_subseq, win, seq_sparsity, args.delta, @@ -1336,8 +1599,15 @@ def main(): args.ambiguous, expectation, ) + canonical_pair_mat = ( + direction_pair_mat if args.forward else pair_mat + ) + else: + canonical_pair_mat = pair_mat # Throw error if the matrix is empty - if np.all(pair_mat == 0) and (not (args.grid or args.grid_only)): + if np.all(canonical_pair_mat == 0) and not ( + args.grid or args.grid_only + ): print( f"The pairwise identity matrix for {sequences[i][0]} and {sequences[j][0]} is empty. Skipping.\n" ) @@ -1387,41 +1657,28 @@ def main(): False, larger_seq_start_pos, smaller_seq_start_pos, + larger_seq_end_pos, + smaller_seq_end_pos, ) if args.grid or args.grid_only: grid_val_doubles.append(bed) grid_val_double_names.append( [larger_seq_name, smaller_seq_name] ) - xlim_val_grid = max(larger_length, xlim_val_grid) + bedfile_prefix = larger_seq_name + "_" + smaller_seq_name + bedpe_path = os.path.join( + args.output_dir or ".", bedfile_prefix + ) + if (not args.no_bedpe) or ( + (not args.no_plot) and (not args.grid_only) + ): + os.makedirs(bedpe_path, exist_ok=True) + if not args.no_bedpe: # Log saving bed file - bedpe_path = "." - if not args.output_dir: - bedfile_prefix = ( - larger_seq_name + "_" + smaller_seq_name - ) - bedpe_path = os.path.join(bedpe_path, bedfile_prefix) - os.makedirs(bedpe_path, exist_ok=True) - bedfile_output = os.path.join( - bedpe_path, - bedfile_prefix + "_COMPARE.bedpe", - ) - else: - bedfile_prefix = ( - larger_seq_name + "_" + smaller_seq_name - ) - bedpe_path = os.path.join( - args.output_dir, bedfile_prefix - ) - os.makedirs(bedpe_path, exist_ok=True) - bedfile_output = os.path.join( - bedpe_path, - larger_seq_name - + "_" - + smaller_seq_name - + "_COMPARE.bedpe", - ) + bedfile_output = os.path.join( + bedpe_path, bedfile_prefix + "_COMPARE.bedpe" + ) with open(bedfile_output, "w") as bedfile: for row in bed: bedfile.write("\t".join(map(str, row)) + "\n") @@ -1441,7 +1698,7 @@ def main(): width=args.width, dpi=args.dpi, is_freq=args.bin_freq, - xlim=args.axes_limits, + xlim=pair_axis_bounds, custom_colors=args.colors, custom_breakpoints=args.breakpoints, from_file=None, @@ -1452,8 +1709,37 @@ def main(): deraster=args.deraster, annotation=args.bed, ) + if args.plot_direction: + if args.forward: + canonical_matrix = direction_pair_mat + forward_matrix = pair_mat + else: + canonical_matrix = pair_mat + forward_matrix = direction_pair_mat + create_direction_plot( + canonical_matrix=canonical_matrix, + forward_matrix=forward_matrix, + window_size=win, + directory=bedpe_path, + name_x=larger_seq_name, + name_y=smaller_seq_name, + self_identity=False, + width=args.width, + dpi=args.dpi, + vector_format=args.vector, + deraster=args.deraster, + xlim=pair_axis_bounds, + axes_labels=args.axes_ticks, + x_offset=larger_seq_start_pos, + y_offset=smaller_seq_start_pos, + ) + + if sketch_cache is not None: + sketch_cache.clear() if args.grid or args.grid_only: + if args.axes_limits: + xlim_val_grid = args.axes_limits print(f"Creating a {len(sequences)}x{len(sequences)} grid.\n") create_grid( singles=grid_val_singles, @@ -1466,13 +1752,14 @@ def main(): is_freq=args.bin_freq, xlim=xlim_val_grid, custom_colors=args.colors, - custom_breakpoints=args.axes_ticks, + custom_breakpoints=args.breakpoints, axes_label=args.axes_ticks, is_bed=False, width=args.width, breaks=args.axes_ticks, deraster=args.deraster, vector_format=args.vector, + dpi=args.dpi, ) diff --git a/src/moddotplot/native_render.py b/src/moddotplot/native_render.py new file mode 100644 index 0000000..50be2d0 --- /dev/null +++ b/src/moddotplot/native_render.py @@ -0,0 +1,451 @@ +"""Sparse, native Matplotlib rendering primitives for ModDotPlot. + +The functions in this module deliberately know nothing about named palettes or +how percent-identity bins are calculated. Callers provide a sequence or +mapping of colors after applying their chosen palette. Keeping those concerns +separate makes the geometry usable by individual dotplots, triangle plots, and +multi-panel grids without creating intermediate SVG files. +""" + +from dataclasses import dataclass +from pathlib import Path +from typing import Callable, Mapping, Optional, Sequence, Tuple, Union + +import matplotlib.pyplot as plt +from matplotlib.axes import Axes +from matplotlib.collections import PolyCollection +from matplotlib.figure import Figure +from matplotlib.ticker import FuncFormatter +import numpy as np +import pandas as pd + + +ColorSource = Union[Sequence[str], Mapping[object, str]] +TickFormatter = Callable[[float, int], str] + + +@dataclass(frozen=True) +class TriangleLayout: + """Axes belonging to a native triangle figure. + + ``annotation_axis`` is ``None`` for an unannotated layout. When present it + shares its x-axis with ``triangle_axis``, so BED features and transformed + triangle tiles use the same genomic coordinates. + """ + + figure: Figure + triangle_axis: Axes + annotation_axis: Optional[Axes] + + +def tile_width(dataframe: pd.DataFrame) -> float: + """Return the legacy tile width, ``max(q_en - q_st)``. + + ModDotPlot historically passes ``q_st`` and ``r_st`` to ``geom_tile`` as + tile centers. The native renderer intentionally preserves that behavior + for the first release of the new plotting pipeline. + """ + + _require_columns(dataframe, ("q_st", "q_en")) + if dataframe.empty: + return 0.0 + widths = _numeric_values(dataframe, "q_en") - _numeric_values(dataframe, "q_st") + if not np.all(np.isfinite(widths)) or np.any(widths <= 0): + raise ValueError("Tile widths must be finite and greater than zero") + return float(np.max(widths)) + + +def rectangular_tile_vertices( + dataframe: pd.DataFrame, *, transpose: bool = False +) -> Sequence[np.ndarray]: + """Build one square polygon per sparse input row. + + The squares are centered on the start columns and all use the maximum query + interval width, matching the existing plotnine tile semantics. Transpose + swaps query and reference positions without copying or expanding the sparse + input into a dense genomic matrix. + """ + + _require_columns(dataframe, ("q_st", "q_en", "r_st")) + if dataframe.empty: + return [] + + window = tile_width(dataframe) + half_window = window / 2.0 + query = _numeric_values(dataframe, "q_st") + reference = _numeric_values(dataframe, "r_st") + if not np.all(np.isfinite(query)) or not np.all(np.isfinite(reference)): + raise ValueError("Tile centers must contain only finite numbers") + if transpose: + query, reference = reference, query + + return [ + np.asarray( + [ + (q - half_window, r - half_window), + (q + half_window, r - half_window), + (q + half_window, r + half_window), + (q - half_window, r + half_window), + ], + dtype=float, + ) + for q, r in zip(query, reference) + ] + + +def transform_triangle_points(points: np.ndarray) -> np.ndarray: + """Rotate dotplot coordinates into triangle coordinates. + + For each ``(query, reference)`` point, the result is + ``((query + reference) / 2, (reference - query) / 2)``. + """ + + values = np.asarray(points, dtype=float) + if values.ndim != 2 or values.shape[1] != 2: + raise ValueError("Triangle points must be an N-by-2 array") + if not np.all(np.isfinite(values)): + raise ValueError("Triangle points must contain only finite numbers") + query = values[:, 0] + reference = values[:, 1] + return np.column_stack(((query + reference) / 2.0, (reference - query) / 2.0)) + + +def triangle_tile_vertices( + dataframe: pd.DataFrame, *, baseline: float = 0.0 +) -> Sequence[np.ndarray]: + """Transform sparse square tiles and clip them at the triangle baseline.""" + + if not np.isfinite(baseline): + raise ValueError("Triangle baseline must be finite") + polygons = [] + for rectangle in rectangular_tile_vertices(dataframe): + transformed = transform_triangle_points(rectangle) + clipped = _clip_polygon_above_baseline(transformed, float(baseline)) + if len(clipped) >= 3: + polygons.append(clipped) + return polygons + + +def draw_rectangular_tiles( + axis: Axes, + dataframe: pd.DataFrame, + colors: ColorSource, + *, + color_column: str = "discrete", + transpose: bool = False, + rasterized: bool = True, + edgecolors: str = "none", + linewidth: float = 0.0, + alpha: float = 1.0, +) -> PolyCollection: + """Draw sparse rectangular tiles and return their ``PolyCollection``.""" + + vertices = rectangular_tile_vertices(dataframe, transpose=transpose) + facecolors = _resolve_facecolors(dataframe, colors, color_column) + collection = PolyCollection( + vertices, + facecolors=facecolors, + edgecolors=edgecolors, + linewidths=linewidth, + rasterized=rasterized, + alpha=alpha, + ) + axis.add_collection(collection) + return collection + + +def draw_triangle_tiles( + axis: Axes, + dataframe: pd.DataFrame, + colors: ColorSource, + *, + color_column: str = "discrete", + baseline: float = 0.0, + rasterized: bool = True, + edgecolors: str = "none", + linewidth: float = 0.0, + alpha: float = 1.0, +) -> PolyCollection: + """Draw transformed sparse tiles, clipped to ``y >= baseline``.""" + + vertices = triangle_tile_vertices(dataframe, baseline=baseline) + # A tile can be wholly below a non-default baseline, so resolve only the + # corresponding leading colors. The normal ModDotPlot baseline is zero + # and its upper-triangle input retains every row. + facecolors = _resolve_facecolors(dataframe, colors, color_column) + if len(vertices) != len(facecolors): + facecolors = _triangle_visible_facecolors( + dataframe, colors, color_column, baseline + ) + collection = PolyCollection( + vertices, + facecolors=facecolors, + edgecolors=edgecolors, + linewidths=linewidth, + rasterized=rasterized, + alpha=alpha, + ) + axis.add_collection(collection) + return collection + + +def genomic_scale(maximum: float) -> Tuple[float, str]: + """Return the divisor and unit used for genomic tick labels.""" + + if not np.isfinite(maximum): + raise ValueError("Genomic axis maximum must be finite") + maximum = abs(float(maximum)) + if maximum < 200_000: + return 1_000.0, "Kbp" + if maximum > 200_000_000: + return 1_000_000_000.0, "Gbp" + return 1_000_000.0, "Mbp" + + +def genomic_tick_formatter(maximum: float) -> FuncFormatter: + """Create a Matplotlib formatter matching ModDotPlot's genomic scaling.""" + + divisor, _ = genomic_scale(maximum) + + def format_tick(value: float, _position: int) -> str: + scaled = value / divisor + return f"{scaled:g}" + + return FuncFormatter(format_tick) + + +def configure_dotplot_axis( + axis: Axes, + region_start: float, + region_end: float, + *, + breaks: Optional[Sequence[float]] = None, + formatter: Optional[TickFormatter] = None, + show_x: bool = True, + show_y: bool = True, + equal_aspect: bool = True, +) -> Axes: + """Configure limits and scaled ticks for a rectangular dotplot axis.""" + + _validate_region(region_start, region_end) + axis.set_xlim(region_start, region_end) + axis.set_ylim(region_start, region_end) + if breaks is not None: + axis.set_xticks(breaks) + axis.set_yticks(breaks) + # Matplotlib deliberately expands view limits to make every fixed tick + # visible. Re-apply the genomic interval so custom or automatically + # generated ticks can never add whitespace outside the data bounds. + axis.set_xlim(region_start, region_end) + axis.set_ylim(region_start, region_end) + tick_formatter = formatter or genomic_tick_formatter(region_end) + axis.xaxis.set_major_formatter( + FuncFormatter(tick_formatter) if formatter else tick_formatter + ) + axis.yaxis.set_major_formatter( + FuncFormatter(tick_formatter) if formatter else tick_formatter + ) + axis.tick_params(axis="x", labelbottom=show_x, bottom=show_x) + axis.tick_params(axis="y", labelleft=show_y, left=show_y) + if equal_aspect: + axis.set_aspect("equal", adjustable="box") + return axis + + +def configure_triangle_axis( + axis: Axes, + region_start: float, + region_end: float, + *, + breaks: Optional[Sequence[float]] = None, + formatter: Optional[TickFormatter] = None, + baseline: float = 0.0, + label: bool = True, +) -> Axes: + """Configure a transformed triangle axis with a genomic x-axis.""" + + _validate_region(region_start, region_end) + if not np.isfinite(baseline): + raise ValueError("Triangle baseline must be finite") + axis.set_xlim(region_start, region_end) + axis.set_ylim(baseline, baseline + (region_end - region_start) / 2.0) + if breaks is not None: + axis.set_xticks(breaks) + axis.set_xlim(region_start, region_end) + tick_formatter = formatter or genomic_tick_formatter(region_end) + axis.xaxis.set_major_formatter( + FuncFormatter(tick_formatter) if formatter else tick_formatter + ) + axis.set_yticks([]) + axis.set_aspect("equal", adjustable="box") + axis.spines["left"].set_visible(False) + axis.spines["right"].set_visible(False) + axis.spines["top"].set_visible(False) + if label: + _, unit = genomic_scale(region_end) + axis.set_xlabel(f"Genomic Position ({unit})") + return axis + + +def create_triangle_layout( + width: float, + *, + with_annotation: bool = False, + annotation_height: float = 0.8, + hspace: float = 0.05, +) -> TriangleLayout: + """Create triangle axes, optionally with a shared BED annotation axis.""" + + width = float(width) + annotation_height = float(annotation_height) + if not np.isfinite(width) or width <= 0: + raise ValueError("Figure width must be finite and greater than zero") + if not np.isfinite(annotation_height) or annotation_height <= 0: + raise ValueError("Annotation height must be finite and greater than zero") + + triangle_height = width / 2.0 + extra_height = annotation_height if with_annotation else 0.0 + figure = plt.figure(figsize=(width, triangle_height + extra_height)) + if not with_annotation: + triangle_axis = figure.add_subplot(1, 1, 1) + return TriangleLayout(figure, triangle_axis, None) + + grid = figure.add_gridspec( + 2, + 1, + height_ratios=(triangle_height, annotation_height), + hspace=hspace, + ) + triangle_axis = figure.add_subplot(grid[0, 0]) + annotation_axis = figure.add_subplot(grid[1, 0], sharex=triangle_axis) + triangle_axis.tick_params(axis="x", labelbottom=False) + return TriangleLayout(figure, triangle_axis, annotation_axis) + + +def save_figure_pair( + figure: Figure, + output_prefix: Union[str, Path], + vector_format: str, + dpi: int, + *, + transparent: bool = False, + bbox_inches: Optional[str] = "tight", +) -> Tuple[Path, Path]: + """Save one figure directly as PNG and SVG, PDF, or PostScript.""" + + vector_format = str(vector_format).lower().lstrip(".") + if vector_format not in {"svg", "pdf", "ps"}: + raise ValueError("Vector format must be one of: svg, pdf, ps") + if int(dpi) <= 0: + raise ValueError("DPI must be greater than zero") + + prefix = Path(output_prefix) + prefix.parent.mkdir(parents=True, exist_ok=True) + # Sequence identifiers frequently contain dots (for example + # ``PAN010.chr14``), so ``Path.with_suffix`` would discard part of a valid + # output prefix. Append extensions exactly as the legacy CLI does. + png_path = Path(f"{prefix}.png") + vector_path = Path(f"{prefix}.{vector_format}") + save_options = { + "dpi": int(dpi), + "transparent": transparent, + "bbox_inches": bbox_inches, + } + figure.savefig(png_path, format="png", **save_options) + figure.savefig(vector_path, format=vector_format, **save_options) + return png_path, vector_path + + +def _require_columns(dataframe: pd.DataFrame, columns: Sequence[str]) -> None: + missing = [column for column in columns if column not in dataframe.columns] + if missing: + raise ValueError(f"Missing required tile columns: {', '.join(missing)}") + + +def _numeric_values(dataframe: pd.DataFrame, column: str) -> np.ndarray: + try: + return pd.to_numeric(dataframe[column], errors="raise").to_numpy(dtype=float) + except (TypeError, ValueError) as error: + raise ValueError(f"Tile column '{column}' must contain only numbers") from error + + +def _resolve_facecolors( + dataframe: pd.DataFrame, colors: ColorSource, color_column: str +) -> Sequence[str]: + if dataframe.empty: + return [] + _require_columns(dataframe, (color_column,)) + values = dataframe[color_column] + + if isinstance(colors, Mapping): + try: + return [colors[value] for value in values] + except KeyError as error: + raise ValueError( + f"No color was provided for tile category {error.args[0]!r}" + ) from error + + if isinstance(colors, str): + palette = [colors] + else: + palette = list(colors) + if not palette: + raise ValueError("At least one tile color is required") + + if isinstance(values.dtype, pd.CategoricalDtype): + categories = list(values.cat.categories) + else: + categories = list(pd.unique(values)) + if len(palette) < len(categories): + raise ValueError( + "The color sequence must contain at least one color per tile category" + ) + color_by_category = dict(zip(categories, palette)) + return [color_by_category[value] for value in values] + + +def _triangle_visible_facecolors( + dataframe: pd.DataFrame, + colors: ColorSource, + color_column: str, + baseline: float, +) -> Sequence[str]: + all_colors = _resolve_facecolors(dataframe, colors, color_column) + rectangles = rectangular_tile_vertices(dataframe) + return [ + color + for rectangle, color in zip(rectangles, all_colors) + if len( + _clip_polygon_above_baseline( + transform_triangle_points(rectangle), float(baseline) + ) + ) + >= 3 + ] + + +def _clip_polygon_above_baseline(polygon: np.ndarray, baseline: float) -> np.ndarray: + """Clip a convex polygon against the half-plane ``y >= baseline``.""" + + if len(polygon) == 0: + return np.empty((0, 2), dtype=float) + output = [] + previous = polygon[-1] + previous_inside = previous[1] >= baseline + for current in polygon: + current_inside = current[1] >= baseline + if current_inside != previous_inside: + fraction = (baseline - previous[1]) / (current[1] - previous[1]) + output.append(previous + fraction * (current - previous)) + if current_inside: + output.append(current) + previous = current + previous_inside = current_inside + return np.asarray(output, dtype=float).reshape((-1, 2)) + + +def _validate_region(region_start: float, region_end: float) -> None: + if not np.isfinite(region_start) or not np.isfinite(region_end): + raise ValueError("Genomic axis limits must be finite") + if region_end <= region_start: + raise ValueError("Genomic region end must be greater than its start") diff --git a/src/moddotplot/parse_fasta.py b/src/moddotplot/parse_fasta.py index 35a23af..b0a75f3 100644 --- a/src/moddotplot/parse_fasta.py +++ b/src/moddotplot/parse_fasta.py @@ -1,15 +1,311 @@ -from enum import unique -from typing import Iterable, List, Sequence -import pysam +from typing import ( + Iterable, + Iterator, + List, + NamedTuple, + Optional, + Sequence, + TextIO, + Tuple, +) import sys -import mmh3 import os import pickle import re import numpy as np import gzip -tab_b = bytes.maketrans(b"ACTG", b"TGAC") +from moddotplot import _nthash + +HASH_ALGORITHM = _nthash.ALGORITHM + + +class FastaIndexEntry(NamedTuple): + name: str + length: int + offset: int + line_bases: int + line_width: int + + +def _is_gzip(filename: str) -> bool: + with open(filename, "rb") as probe: + return probe.read(2) == b"\x1f\x8b" + + +def _open_fasta_text(filename: str) -> TextIO: + """Open a plain, gzip, or BGZF FASTA file as text. + + Compression is detected from the gzip magic bytes rather than the filename + extension. BGZF is a blocked form of gzip and is decoded by Python's gzip + reader as a concatenated gzip stream. + """ + if _is_gzip(filename): + return gzip.open(filename, "rt", encoding="ascii", newline=None) + return open(filename, "rt", encoding="ascii", newline=None) + + +def _read_fasta_index(filename: str) -> Optional[List[FastaIndexEntry]]: + """Read a fresh samtools-style ``.fai`` index when one is available.""" + + if _is_gzip(filename): + return None + index_path = f"{os.fspath(filename)}.fai" + if not os.path.isfile(index_path): + return None + if os.path.getmtime(index_path) < os.path.getmtime(filename): + return None + + entries = [] + seen_names = set() + try: + with open(index_path, "rt", encoding="utf-8") as index: + for line_number, raw_line in enumerate(index, start=1): + fields = raw_line.rstrip("\r\n").split("\t") + if len(fields) < 5: + return None + name = fields[0] + if not name or name in seen_names: + return None + seen_names.add(name) + entry = FastaIndexEntry( + name, + int(fields[1]), + int(fields[2]), + int(fields[3]), + int(fields[4]), + ) + if entry.length < 0 or entry.offset < 0: + return None + if entry.length and ( + entry.line_bases <= 0 or entry.line_width < entry.line_bases + ): + return None + entries.append(entry) + except (OSError, UnicodeError, ValueError): + return None + return entries or None + + +def _iter_fasta_headers(filename: str) -> Iterator[str]: + """Yield FASTA identifiers without assembling or validating sequences.""" + + seen_ids = set() + found_header = False + with _open_fasta_text(filename) as fasta: + for line_number, raw_line in enumerate(fasta, start=1): + if raw_line.startswith(">"): + description = raw_line[1:].strip() + if not description: + raise ValueError( + f"Invalid FASTA {filename!s}: empty header at line {line_number}" + ) + sequence_id = description.split(maxsplit=1)[0] + if sequence_id in seen_ids: + raise ValueError( + f"Invalid FASTA {filename!s}: duplicate sequence identifier " + f"{sequence_id!r} at line {line_number}" + ) + seen_ids.add(sequence_id) + found_header = True + yield sequence_id + elif not found_header and raw_line.strip(): + raise ValueError( + f"Invalid FASTA {filename!s}: sequence data before the first " + f"header at line {line_number}" + ) + + if not found_header: + raise ValueError(f"Invalid FASTA {filename!s}: no FASTA records found") + + +def _fetch_indexed_region( + filename: str, entry: FastaIndexEntry, start: int, end: int +) -> str: + """Fetch one 1-based inclusive interval directly from an indexed FASTA.""" + + if entry.length == 0 and start == 1 and end == 0: + return "" + if start < 1 or end < start or end > entry.length: + raise ValueError( + f"region {entry.name}:{start}-{end} is outside sequence length " + f"{entry.length}" + ) + start_index = start - 1 + end_index = end - 1 + start_byte = ( + entry.offset + + (start_index // entry.line_bases) * entry.line_width + + start_index % entry.line_bases + ) + end_byte = ( + entry.offset + + (end_index // entry.line_bases) * entry.line_width + + end_index % entry.line_bases + ) + with open(filename, "rb") as fasta: + fasta.seek(start_byte) + raw_sequence = fasta.read(end_byte - start_byte + 1) + + sequence_bytes = raw_sequence.replace(b"\n", b"").replace(b"\r", b"") + expected_length = end - start + 1 + if len(sequence_bytes) != expected_length: + raise ValueError( + f"FASTA index for {entry.name!r} returned {len(sequence_bytes)} bases; " + f"expected {expected_length}" + ) + if re.search(rb"[\t\v\f ]", sequence_bytes): + raise ValueError(f"Invalid FASTA {filename!s}: whitespace within sequence data") + return sequence_bytes.decode("ascii") + + +def _iter_selected_fasta_records( + filename: str, regions, single_record: bool +) -> Iterator[Tuple[str, str]]: + """Stream selected intervals, stopping early for a one-record FASTA.""" + + sequence_id = None + sequence_parts = [] + sequence_position = 0 + selected_region = None + selection_complete = False + seen_ids = set() + + def selected_sequence(): + sequence = "".join(sequence_parts) + if selected_region: + _name, start, end = selected_region + expected_length = end - start + 1 + if len(sequence) != expected_length: + raise ValueError( + f"region {sequence_id}:{start}-{end} is outside sequence length " + f"{sequence_position}" + ) + return sequence + + with _open_fasta_text(filename) as fasta: + for line_number, raw_line in enumerate(fasta, start=1): + if raw_line.startswith(">"): + if sequence_id is not None: + yield sequence_id, selected_sequence() + + description = raw_line[1:].strip() + if not description: + raise ValueError( + f"Invalid FASTA {filename!s}: empty header at line {line_number}" + ) + sequence_id = description.split(maxsplit=1)[0] + if sequence_id in seen_ids: + raise ValueError( + f"Invalid FASTA {filename!s}: duplicate sequence identifier " + f"{sequence_id!r} at line {line_number}" + ) + seen_ids.add(sequence_id) + sequence_parts = [] + sequence_position = 0 + selected_region = regions.get(sequence_id) if regions else None + selection_complete = False + continue + + if sequence_id is None: + if not raw_line.strip(): + continue + raise ValueError( + f"Invalid FASTA {filename!s}: sequence data before the first " + f"header at line {line_number}" + ) + if selection_complete: + continue + + line = raw_line.strip() + if not line: + continue + if any(character.isspace() for character in line): + raise ValueError( + f"Invalid FASTA {filename!s}: whitespace within sequence data " + f"at line {line_number}" + ) + + line_start = sequence_position + 1 + line_end = sequence_position + len(line) + if selected_region: + _name, start, end = selected_region + overlap_start = max(start, line_start) + overlap_end = min(end, line_end) + if overlap_start <= overlap_end: + local_start = overlap_start - line_start + local_end = overlap_end - line_start + 1 + sequence_parts.append(line[local_start:local_end]) + sequence_position = line_end + if sequence_position >= end: + selection_complete = True + if single_record: + yield sequence_id, selected_sequence() + return + else: + sequence_parts.append(line) + sequence_position = line_end + + if sequence_id is None: + raise ValueError(f"Invalid FASTA {filename!s}: no FASTA records found") + yield sequence_id, selected_sequence() + + +def _iter_fasta_records(filename: str) -> Iterator[Tuple[str, str]]: + """Yield ``(identifier, sequence)`` records from a FASTA file. + + Identifiers follow the convention used by ``pysam.FastaFile``: only the + first whitespace-delimited token after ``>`` is retained. Empty sequence + records and blank lines are accepted, while malformed content, empty + identifiers, and duplicate identifiers raise ``ValueError`` with the + offending line number. + """ + sequence_id = None + sequence_parts = [] + seen_ids = set() + + with _open_fasta_text(filename) as fasta: + for line_number, raw_line in enumerate(fasta, start=1): + line = raw_line.strip() + if not line: + continue + + if line.startswith(">"): + if sequence_id is not None: + yield sequence_id, "".join(sequence_parts) + + description = line[1:].strip() + if not description: + raise ValueError( + f"Invalid FASTA {filename!s}: empty header at line {line_number}" + ) + + sequence_id = description.split(maxsplit=1)[0] + if sequence_id in seen_ids: + raise ValueError( + f"Invalid FASTA {filename!s}: duplicate sequence identifier " + f"{sequence_id!r} at line {line_number}" + ) + seen_ids.add(sequence_id) + sequence_parts = [] + continue + + if sequence_id is None: + raise ValueError( + f"Invalid FASTA {filename!s}: sequence data before the first " + f"header at line {line_number}" + ) + if any(character.isspace() for character in line): + raise ValueError( + f"Invalid FASTA {filename!s}: whitespace within sequence data " + f"at line {line_number}" + ) + sequence_parts.append(line) + + if sequence_id is None: + raise ValueError(f"Invalid FASTA {filename!s}: no FASTA records found") + + yield sequence_id, "".join(sequence_parts) def extractRegion(seq_name): @@ -20,10 +316,14 @@ def extractRegion(seq_name): - Extended format: "HG002_chr13_MATERNAL:1-4000000:1000000-3000000" (keeps the last range). """ - # Match chromosome + one or more ranges separated by colons - region_pattern = r"^([a-zA-Z0-9_]+)(?::(\d+-\d+))+" + # FASTA identifiers commonly contain periods, hyphens, and other + # punctuation. Treat everything before the first trailing coordinate + # range as the identifier instead of restricting it to ``\w`` characters. + # Repeated ranges are retained for compatibility with saved names such as + # ``sample:1-4000000:1000000-3000000``; the final range wins. + region_pattern = r"^(.+?)(?::\d+-\d+)+$" - match = re.match(region_pattern, seq_name) + match = re.fullmatch(region_pattern, seq_name) if match: chrom = match.group(1) # Get all ranges (everything after the chrom) @@ -37,57 +337,107 @@ def extractRegion(seq_name): return None -def generateKmersFromFasta( - seq: Sequence[str], k: int, quiet: bool, fw_only: bool -) -> Iterable[int]: +def _hash_sequence( + seq: Sequence[str], + k: int, + fw_only: bool, + ambiguous: bool = False, +) -> np.ma.MaskedArray: + """Bulk-hash every genomic k-mer start with ntHash2. + + ntHash2 ordinarily skips windows containing non-ACGTU characters. The + native wrapper instead returns one hash for every start together with an + ambiguity mask, allowing downstream window slicing to remain aligned with + genomic coordinates. Unless ``ambiguous`` is requested, those fallback + hashes stay masked and therefore cannot enter a modimizer sketch. + """ + if k <= 0: + raise ValueError("k-mer size must be greater than zero") + n = len(seq) + total_kmers = max(n - k + 1, 0) + + if isinstance(seq, str): + sequence = seq.upper().encode("ascii") + else: + sequence = bytes(seq).upper() + + hash_buffer, ambiguity_buffer = _nthash.hash_kmers(sequence, k, not fw_only) + hashes = np.frombuffer(hash_buffer, dtype=np.uint64) + if len(hashes) != total_kmers: + raise RuntimeError( + "ntHash2 returned an unexpected number of hashes: " + f"expected {total_kmers}, received {len(hashes)}" + ) + + mask = np.ma.nomask + if ambiguity_buffer: + ambiguous_windows = np.frombuffer(ambiguity_buffer, dtype=np.uint8) + if len(ambiguous_windows) != total_kmers: + raise RuntimeError( + "ntHash2 returned an ambiguity mask with an unexpected length: " + f"expected {total_kmers}, received {len(ambiguous_windows)}" + ) + if not ambiguous: + mask = ambiguous_windows.astype(bool, copy=False) + + return np.ma.MaskedArray(hashes, mask=mask, copy=False) + + +def generateKmersFromFasta( + seq: Sequence[str], + k: int, + quiet: bool, + fw_only: bool, + ambiguous: bool = False, +) -> Iterable[Optional[int]]: + """Yield position-preserving ntHash2 values for a sequence. + + The public iterator remains compatible with existing callers. Ambiguous + windows are yielded as ``None`` unless ``ambiguous`` is enabled, while the + FASTA reader below keeps the compact bulk NumPy representation in memory. + """ + total_kmers = max(len(seq) - k + 1, 0) if not quiet: - progress_thresholds = round(n / 77) - printProgressBar(0, n - k + 1, prefix="Progress:", suffix="Complete", length=40) + printProgressBar( + 0, total_kmers, prefix="Progress:", suffix="Complete", length=40 + ) - for i in range(n - k + 1): - if not quiet: - if i % progress_thresholds == 0: - printProgressBar( - i, n - k + 1, prefix="Progress:", suffix="Complete", length=40 - ) - if i == n - k: - printProgressBar( - n - k + 1, - n - k + 1, - prefix="Progress:", - suffix="Completed", - length=40, - ) - # Remove case sensitivity - kmer = seq[i : i + k].upper() - fh = mmh3.hash(kmer) - if fw_only: - yield fh - else: - # Calculate reverse complement hash directly without the need for translation - rc = mmh3.hash(kmer[::-1].translate(tab_b)) - yield fh if fh < rc else rc + hashes = _hash_sequence(seq, k, fw_only, ambiguous) + progress_threshold = max(round(total_kmers / 77), 1) + for index, kmer_hash in enumerate(hashes): + if not quiet and index % progress_threshold == 0: + printProgressBar( + index, + total_kmers, + prefix="Progress:", + suffix="Complete", + length=40, + ) + + yield None if np.ma.is_masked(kmer_hash) else int(kmer_hash) + + if not quiet and index == total_kmers - 1: + printProgressBar( + total_kmers, + total_kmers, + prefix="Progress:", + suffix="Completed", + length=40, + ) def isValidFasta(file_path): try: - open_func = gzip.open if file_path.endswith(".gz") else open - with open_func(file_path, "rt") as file: - in_sequence = False - for line in file: - line = line.strip() - if line.startswith(">"): - in_sequence = True - elif in_sequence and not line: - # Empty line encountered after the sequence header - in_sequence = False - elif in_sequence: - pass - return in_sequence + for _sequence_id, _sequence in _iter_fasta_records(file_path): + pass + return True except FileNotFoundError: print(f"Unable to find fasta {file_path}. Check filename and/or directory!\n") sys.exit(5) + except ValueError as error: + print(f"An error occurred: {error}") + return False except Exception as e: print(f"An error occurred: {str(e)}") sys.exit(6) @@ -145,8 +495,12 @@ def printProgressBar( fill="█", printEnd="\r", ): - percent = f"{100 * (iteration / total):.{decimals}f}" - filledLength = int(length * iteration // total) + if total <= 0: + percent = f"{100:.{decimals}f}" + filledLength = length + else: + percent = f"{100 * (iteration / total):.{decimals}f}" + filledLength = int(length * iteration // total) bar = [fill] * filledLength + ["-"] * (length - filledLength) bar_str = "".join(bar) print(f"\r{prefix} |{bar_str}| {percent}% {suffix}", end=printEnd) @@ -155,39 +509,100 @@ def printProgressBar( def readKmersFromFile( - filename: str, ksize: int, quiet: bool, fw_only: bool -) -> List[List[int]]: + filename: str, + ksize: int, + quiet: bool, + fw_only: bool, + ambiguous: bool = False, + regions=None, + record_ids=None, +) -> List[np.ma.MaskedArray]: """ Given a filename and an integer k, returns a list of all k-mers found in the sequences in the file. """ all_kmers = [] - seq = pysam.FastaFile(filename) - - for seq_id in seq.references: - print(f"Retrieving k-mers from {seq_id}.... \n") - kmers_for_seq = [] - for kmer_hash in generateKmersFromFasta( - seq.fetch(seq_id), ksize, quiet, fw_only - ): - kmers_for_seq.append(kmer_hash) + record_ids = ( + list(record_ids) if record_ids is not None else getInputHeaders(filename) + ) + fasta_index = _read_fasta_index(filename) + indexed_entries = ( + {entry.name: entry for entry in fasta_index} + if fasta_index and [entry.name for entry in fasta_index] == record_ids + else None + ) + + if indexed_entries is not None: + + def indexed_records(): + for seq_id in record_ids: + entry = indexed_entries[seq_id] + if regions and seq_id in regions: + _name, start, end = regions[seq_id] + else: + start, end = 1, entry.length + sequence_label = ( + f"{seq_id}:{start}-{end}" + if regions and seq_id in regions + else seq_id + ) + print(f"Retrieving k-mers from {sequence_label}.... \n") + sequence = _fetch_indexed_region(filename, entry, start, end) + yield seq_id, sequence, sequence_label + + selected_records = indexed_records() + else: + selected_records = ( + ( + seq_id, + sequence, + ( + f"{seq_id}:{regions[seq_id][1]}-{regions[seq_id][2]}" + if regions and seq_id in regions + else seq_id + ), + ) + for seq_id, sequence in _iter_selected_fasta_records( + filename, regions, single_record=len(record_ids) == 1 + ) + ) + + for seq_id, sequence, sequence_label in selected_records: + if indexed_entries is None: + print(f"Retrieving k-mers from {sequence_label}.... \n") + if len(sequence) < ksize: + if regions and seq_id in regions: + _name, start, end = regions[seq_id] + raise ValueError( + f"region {seq_id}:{start}-{end} is shorter than k-mer size " + f"{ksize}" + ) + + total_kmers = max(len(sequence) - ksize + 1, 0) + if not quiet: + printProgressBar( + 0, total_kmers, prefix="Progress:", suffix="Complete", length=40 + ) + kmers_for_seq = _hash_sequence(sequence, ksize, fw_only, ambiguous) + if not quiet: + printProgressBar( + total_kmers, + total_kmers, + prefix="Progress:", + suffix="Completed", + length=40, + ) all_kmers.append(kmers_for_seq) - print(f"\n{seq_id} k-mers retrieved! \n") + print(f"\n{sequence_label} k-mers retrieved! \n") return all_kmers def getInputHeaders(filename: str) -> List[str]: - header_list = [] - try: - seq = pysam.FastaFile(filename) - except OSError: - seq = None - for seq_id in seq.references: - header_list.append(seq_id) - - return header_list + fasta_index = _read_fasta_index(filename) + if fasta_index is not None: + return [entry.name for entry in fasta_index] + return list(_iter_fasta_headers(filename)) def getInputSeqLength(filename: str) -> List[int]: - seq = pysam.FastaFile(filename) - return seq.lengths + return [len(sequence) for _sequence_id, sequence in _iter_fasta_records(filename)] diff --git a/src/moddotplot/static_plots.py b/src/moddotplot/static_plots.py index 40d3d5a..e542cc0 100755 --- a/src/moddotplot/static_plots.py +++ b/src/moddotplot/static_plots.py @@ -21,111 +21,104 @@ element_text, theme_light, geom_blank, - annotate, - element_rect, - coord_flip, theme_minimal, - geom_raster, ) -import svgutils.transform as sg -import cairosvg import pandas as pd import numpy as np -import glob -from PIL import Image -import patchworklib as pw import math import os -import xml.etree.ElementTree as ET -import sys import re -from moddotplot.parse_fasta import printProgressBar -from lxml import etree -from pygenometracks.utilities import get_region import matplotlib.pyplot as plt +from matplotlib.patches import Rectangle +from matplotlib.ticker import ScalarFormatter +from moddotplot.native_render import ( + configure_dotplot_axis, + configure_triangle_axis, + create_triangle_layout, + draw_rectangular_tiles, + draw_triangle_tiles, + genomic_scale, + genomic_tick_formatter, + save_figure_pair, +) from moddotplot.const import ( DIVERGING_PALETTES, QUALITATIVE_PALETTES, SEQUENTIAL_PALETTES, ) -from typing import List from palettable.colorbrewer import qualitative, sequential, diverging -import logging -# Set log level BEFORE importing pygenometracks -for name in logging.root.manager.loggerDict: - if name.startswith("pygenometracks"): - logging.getLogger(name).setLevel(logging.CRITICAL) - logging.getLogger(name).propagate = False # Don't pass to root logger -# Also make sure the root logger isn’t outputting debug messages -logging.basicConfig(level=logging.CRITICAL) +DEFAULT_ANNOTATION_COLOR = "#4C72B0" +REGION_SUFFIX_PATTERN = re.compile(r"(?::\d+-\d+)+$") -from pygenometracks.tracksClass import PlotTracks +def display_sequence_name(name): + """Return a sequence name without appended region coordinates.""" -def is_plot_empty(p): - # Check if the plot has data or any layers - return len(p.layers) == 0 and p.data.empty + return REGION_SUFFIX_PATTERN.sub("", str(name)) -def check_pascal(single_val, double_val): - try: - if len(single_val) == 2: - assert len(double_val) == 1 - elif len(single_val) == 3: - assert len(double_val) == 3 - elif len(single_val) == 4: - assert len(double_val) == 6 - elif len(single_val) == 5: - assert len(double_val) == 10 - elif len(single_val) == 6: - assert len(double_val) == 15 - elif len(single_val) == 0: - assert len(double_val) == (1 or 3 or 6 or 10 or 15) - except AssertionError as e: - print( - f"Missing bed files required to create grid. Please verify all bed files are included." - ) - sys.exit(8) +def _fit_grid_sequence_labels(figure, axes, minimum_size=2.0): + """Shrink grid sequence headings until each fits inside its own panel.""" + figure.canvas.draw() + renderer = figure.canvas.get_renderer() + ratios = [] -def generate_ini_file( - bedfile, ininame, chrom, color_value="bed_rgb", x_axis=True, display="collapsed" -): - try: - thing = chrom.split(":")[0] - sections = [ - "[spacer]", - "# height of space in cm (optional)", - "height = 0.5", - "", - f"[{thing}]", - f"file = {bedfile}", - f"Title=", - "height = 1", - f"display = {display}", - f"color = {color_value}", - "labels = false", - "fontsize = 10", - "file_type = bed", - ] - - if x_axis: - sections.insert(0, "[x-axis]") - - ini_content = "\n".join(sections) - with open(f"{ininame}.ini", "w") as file: - file.write(ini_content) - print(f"Successfully generated {ininame}.ini\n") - return f"{ininame}.ini" - except Exception as err: - print(f"Error producing ini file: {err}\n") - return None + for axis in axes[0, :]: + title = axis.title + if title.get_text(): + title_box = title.get_window_extent(renderer=renderer) + axis_box = axis.get_window_extent(renderer=renderer) + if title_box.width: + ratios.append(axis_box.width * 0.9 / title_box.width) + + for axis in axes[:, 0]: + label = axis.yaxis.label + if label.get_text(): + label_box = label.get_window_extent(renderer=renderer) + axis_box = axis.get_window_extent(renderer=renderer) + if label_box.height: + ratios.append(axis_box.height * 0.9 / label_box.height) + + if not ratios: + return + scale = min(1.0, min(ratios)) + if scale >= 1.0: + return + + artists = [axis.title for axis in axes[0, :]] + [ + axis.yaxis.label for axis in axes[:, 0] + ] + for artist in artists: + artist.set_fontsize(max(minimum_size, artist.get_fontsize() * scale)) + + +def _resolve_native_colors(palette, palette_orientation, custom_colors=None): + """Resolve plot colors with the same orientation rules as plotnine paths.""" + if palette in DIVERGING_PALETTES: + palette_colors = getattr(diverging, palette).hex_colors + palette_orientation = "-" if palette_orientation == "+" else "+" + elif palette in QUALITATIVE_PALETTES: + palette_colors = getattr(qualitative, palette).hex_colors + elif palette in SEQUENTIAL_PALETTES: + palette_colors = getattr(sequential, palette).hex_colors + else: + palette_colors = diverging.Spectral_11.hex_colors + palette_orientation = "-" + + colors = palette_colors[::-1] if palette_orientation == "-" else palette_colors + return list(custom_colors) if custom_colors else list(colors) + + +def is_plot_empty(p): + # Check if the plot has data or any layers + return len(p.layers) == 0 and p.data.empty def read_annotation_bed(filepath): - """Reads a BED file into a Pandas DataFrame and ensures correct formatting.""" + """Read the BED3-BED9 subset used by ModDotPlot annotations.""" col_names = [ "chrom", "start", @@ -136,250 +129,164 @@ def read_annotation_bed(filepath): "thickStart", "thickEnd", "itemRgb", - ] # Include additional fields + ] - df = pd.read_csv(filepath, sep="\t", comment="#", header=None) + try: + df = pd.read_csv(filepath, sep="\t", comment="#", header=None, dtype=str) + except pd.errors.EmptyDataError: + return pd.DataFrame(columns=col_names[:3]) - # Ensure at least three required columns exist - if df.shape[1] < 3: + if not 3 <= df.shape[1] <= len(col_names): raise ValueError( - "Invalid BED file: must have at least 3 columns (chrom, start, end)." + "Invalid BED file: expected between 3 and 9 tab-separated columns." ) - # Rename only the expected columns df.columns = col_names[: df.shape[1]] + df["chrom"] = df["chrom"].astype(str) + + for column in ("start", "end"): + try: + values = pd.to_numeric(df[column], errors="raise") + except (TypeError, ValueError) as error: + raise ValueError( + f"Invalid BED file: '{column}' must contain only integers." + ) from error + if values.isna().any() or not np.all(np.isfinite(values)): + raise ValueError( + f"Invalid BED file: '{column}' must contain only finite integers." + ) + if np.any(values % 1 != 0): + raise ValueError( + f"Invalid BED file: '{column}' must contain only integers." + ) + df[column] = values.astype(np.int64) - # Ensure start and end columns contain valid integers - if df["start"].isna().any() or df["end"].isna().any(): + if (df["start"] < 0).any(): + raise ValueError("Invalid BED file: 'start' must be non-negative.") + if (df["end"] <= df["start"]).any(): raise ValueError( - "Invalid BED file: 'start' and 'end' columns must be integers and contain no missing values." + "Invalid BED file: 'end' must be greater than 'start' for every interval." ) return df -def make_svg_background_transparent(svg_path, output_path=None): - """ - Makes the background of an SVG file transparent by removing/modifying background fills. - - Args: - svg_path: Path to input SVG file - output_path: Path to output SVG file (if None, overwrites input) - """ - import xml.etree.ElementTree as ET - import re - - if output_path is None: - output_path = svg_path - - # Parse the SVG - tree = ET.parse(svg_path) - root = tree.getroot() - - # Define SVG namespace - ns = {"svg": "http://www.w3.org/2000/svg"} - - # Remove background rectangles/paths that cover the entire canvas - # Get SVG dimensions for comparison - width = root.get("width", "0") - height = root.get("height", "0") - - # Extract numeric values - width_num = float(re.sub(r"[a-zA-Z%]+", "", width)) if width != "0" else 0 - height_num = float(re.sub(r"[a-zA-Z%]+", "", height)) if height != "0" else 0 - - # Find and modify background elements - elements_to_modify = [] - - # Check all paths, rectangles, and other elements - for elem in root.iter(): - if ( - elem.tag.endswith("path") - or elem.tag.endswith("rect") - or elem.tag.endswith("polygon") - ): - # Check if this element has a background-like fill - style = elem.get("style", "") - fill = elem.get("fill", "") - - # Look for background colors (light colors, white, etc.) - background_colors = [ - "#ffffff", - "#f0ffff", - "white", - "lightblue", - "lightgray", - "lightgrey", - ] - - is_background = False - current_fill = None - - if "fill:" in style: - # Extract fill from style - fill_match = re.search(r"fill:\s*([^;]+)", style) - if fill_match: - current_fill = fill_match.group(1).strip() - elif fill: - current_fill = fill - - if current_fill and any( - bg_color in current_fill.lower() for bg_color in background_colors - ): - is_background = True - - # For paths, check if it covers a large area (likely background) - if elem.tag.endswith("path"): - d = elem.get("d", "") - # Simple heuristic: if path starts at 0,0 and covers large area, it's likely background - if "M 0" in d and current_fill: - is_background = True - - # For rectangles, check if it covers the full canvas - if elem.tag.endswith("rect"): - x = float(elem.get("x", 0)) - y = float(elem.get("y", 0)) - w = float(elem.get("width", 0)) - h = float(elem.get("height", 0)) - - # If rectangle covers most/all of the canvas, it's likely background - if x <= 1 and y <= 1 and w >= width_num * 0.9 and h >= height_num * 0.9: - is_background = True - - if is_background: - elements_to_modify.append(elem) - - # Modify the background elements - for elem in elements_to_modify: - style = elem.get("style", "") - - if "fill:" in style: - # Replace fill in style - new_style = re.sub(r"fill:\s*[^;]+", "fill: transparent", style) - elem.set("style", new_style) - elif elem.get("fill"): - # Replace fill attribute - elem.set("fill", "transparent") - - # Save the modified SVG - tree.write(output_path, encoding="unicode", xml_declaration=True) - - -def make_all_svg_backgrounds_transparent(directory): - """ - Makes all SVG files in a directory have transparent backgrounds. - """ - - svg_files = glob.glob(os.path.join(directory, "*.svg")) - - for svg_file in svg_files: - make_svg_background_transparent(svg_file) - - print(f"Processed {len(svg_files)} SVG files") +def _annotation_color(value, fallback=DEFAULT_ANNOTATION_COLOR): + """Return a Matplotlib color for a BED ``itemRgb`` value.""" + if value is None or pd.isna(value): + return fallback - -def run_pygenometracks(inifile, region, output_file, width): - output_dir = os.path.dirname(output_file) - if output_dir and not os.path.exists(output_dir): - os.makedirs(output_dir) # Create directory if it doesn't exist + fields = [field.strip() for field in str(value).split(",")] + if len(fields) != 3: + return fallback try: - trp = PlotTracks( - inifile, - width, - fig_height=None, - fontsize=10, - dpi=300, - track_label_width=0.05, - plot_regions=region, - plot_width=width, - ) - - # Extract the chromosome, start, and end from the region - chrom, start, end = region[0] - - # Call the plot method directly to generate the image - fig = trp.plot(output_file, chrom, start, end) - - return trp - - except Exception as e: - if "No valid intervals were found" in str(e): - print(f"No valid intervals found in BED file for region {region[0]}") - print( - "This is expected when the BED file doesn't overlap with the query region." - ) - return None - else: - print(f"Error in run_pygenometracks: {e}") - raise e - + channels = tuple(int(field) for field in fields) + except ValueError: + return fallback + if any(channel < 0 or channel > 255 for channel in channels): + return fallback + return tuple(channel / 255 for channel in channels) -def test_pygenometracks_direct(inifile, chrom, start, end, output_file, width=40): - """ - Direct test function to call PlotTracks.plot() without any coordinate validation. - This bypasses the region checking that happens during PlotTracks initialization. - - Args: - inifile: Path to the tracks configuration file - chrom: Chromosome name - start: Start coordinate - end: End coordinate - output_file: Output image file path - width: Figure width in cm - """ - output_dir = os.path.dirname(output_file) - if output_dir and not os.path.exists(output_dir): - os.makedirs(output_dir) - - # Create a dummy region for initialization (this gets overridden in plot()) - dummy_region = [(chrom, 1, 1000)] - - trp = PlotTracks( - inifile, - width, - fig_height=None, - fontsize=10, - dpi=300, - track_label_width=0.05, - plot_regions=dummy_region, - plot_width=width, +def _visible_annotation_intervals( + bed_df, chrom, region_start, region_end, fallback=DEFAULT_ANNOTATION_COLOR +): + """Select and clip BED intervals to a plotted genomic region.""" + if region_end <= region_start: + raise ValueError("Annotation region end must be greater than its start.") + if bed_df.empty: + return [] + + intervals = [] + matching = bed_df[bed_df["chrom"] == str(chrom)] + has_item_rgb = "itemRgb" in matching.columns + for row in matching.itertuples(index=False): + interval_start = int(row.start) + interval_end = int(row.end) + if interval_end <= interval_start: + continue + clipped_start = max(interval_start, region_start) + clipped_end = min(interval_end, region_end) + if clipped_end <= clipped_start: + continue + rgb = getattr(row, "itemRgb", None) if has_item_rgb else None + intervals.append((clipped_start, clipped_end, _annotation_color(rgb, fallback))) + return intervals + + +def draw_annotation_track( + axis, + bed_df, + chrom, + region_start, + region_end, + fallback=DEFAULT_ANNOTATION_COLOR, +): + """Draw a label-free, collapsed BED track and return its interval count.""" + intervals = _visible_annotation_intervals( + bed_df, chrom, region_start, region_end, fallback ) + for interval_start, interval_end, color in intervals: + axis.add_patch( + Rectangle( + (interval_start, 0.2), + interval_end - interval_start, + 0.6, + facecolor=color, + edgecolor=color, + linewidth=0.25, + ) + ) - # Call plot() directly with your desired coordinates - fig = trp.plot(output_file, chrom, start, end) - - return fig - - -def reverse_pascal(double_vals): - if len(double_vals) == 1: - return 2 - elif len(double_vals) == 3: - return 3 - elif len(double_vals) == 6: - return 4 - elif len(double_vals) == 10: - return 5 - elif len(double_vals) == 15: - return 6 - else: - sys.exit(9) + axis.set_xlim(region_start, region_end) + axis.set_ylim(0, 1) + axis.set_yticks([]) + formatter = ScalarFormatter(useOffset=False) + formatter.set_scientific(False) + axis.xaxis.set_major_formatter(formatter) + axis.tick_params(axis="x", labelsize=8, length=3) + axis.spines["left"].set_visible(False) + axis.spines["right"].set_visible(False) + axis.spines["top"].set_visible(False) + axis.set_facecolor("none") + return len(intervals) + + +def render_annotation_track( + bed_df, + chrom, + region_start, + region_end, + output_prefix, + width_cm, + dpi, + vector_format="svg", +): + """Render a collapsed BED track as PNG and the selected vector format.""" + if vector_format not in {"svg", "pdf", "ps"}: + raise ValueError(f"Unsupported vector format: {vector_format}") + figure, axis = plt.subplots(figsize=(max(float(width_cm), 2.54) / 2.54, 2.0 / 2.54)) + try: + interval_count = draw_annotation_track( + axis, bed_df, chrom, region_start, region_end + ) + if interval_count == 0: + return False -# Hardcoding for now, I have the formula.... I'm just lazy -def transpose_order(double_vals): - if len(double_vals) == 1: - return [0] - elif len(double_vals) == 3: - return [0, 1, 2] - elif len(double_vals) == 6: - return [0, 1, 3, 2, 4, 5] - elif len(double_vals) == 10: - return [0, 1, 4, 2, 5, 7, 3, 6, 8, 9] - elif len(double_vals) == 15: - return [0, 1, 5, 2, 6, 9, 3, 7, 10, 12, 4, 8, 11, 13, 14] + figure.subplots_adjust(left=0.02, right=0.995, bottom=0.34, top=0.96) + figure.savefig( + f"{output_prefix}.{vector_format}", + format=vector_format, + dpi=dpi, + transparent=vector_format != "ps", + facecolor="white" if vector_format == "ps" else "none", + ) + figure.savefig(f"{output_prefix}.png", format="png", dpi=dpi, facecolor="white") + return True + finally: + plt.close(figure) def check_st_en_equality(df): @@ -415,8 +322,24 @@ def make_scale(vals: list) -> list: return make_m(scaled) +def _dotplot_tiles(mapping, deraster=False, **kwargs): + """Create tiles without materializing a genomic-coordinate-sized image. + + ``plotnine.geom_raster`` expands sparse coordinates into an RGBA array whose + dimensions are derived from the smallest coordinate spacing. A 496 Mb + sequence plotted in 2 kb windows can therefore request roughly + 248,000-by-248,000 pixels even when only a small fraction of those cells + contain matches. ``geom_tile`` draws only the cells present in the input + dataframe. Setting ``raster=True`` keeps the default compact, rasterized + appearance in vector output, while ``--deraster`` leaves the tiles as + vectors. + """ + return geom_tile(mapping, raster=not deraster, **kwargs) + + def get_colors(sdf, ncolors, is_freq, custom_breakpoints): - assert ncolors > 2 and ncolors < 12 + if ncolors < 1: + raise ValueError("At least one color is required") try: bot = math.floor(min(sdf["perID_by_events"])) except ValueError: @@ -431,11 +354,24 @@ def get_colors(sdf, ncolors, is_freq, custom_breakpoints): else: breaks = [bot + i * interval for i in range(ncolors + 1)] if custom_breakpoints: - np.asarray(custom_breakpoints, dtype=np.float64) + breaks = np.asarray(custom_breakpoints, dtype=np.float64) + if len(breaks) != ncolors + 1: + raise ValueError( + "The number of breakpoints must equal the number of colors plus one" + ) + if not np.all(np.isfinite(breaks)): + raise ValueError("Breakpoints must contain only finite numbers") + if np.any(np.diff(breaks) <= 0): + raise ValueError("Breakpoints must be strictly increasing") labels = np.arange(len(breaks) - 1) - # corner case of only one %id value - if len(breaks) == 1: - return pd.factorize([1] * len(sdf["perID_by_events"]))[0] + # A dataset containing only 100% identity creates repeated default bin + # edges; frequency bins likewise collapse to one edge when every value is + # equal. Both cases represent a single category and must not reach + # ``pandas.cut``, which requires unique edges. + if len(np.unique(breaks)) < 2: + return pd.Categorical( + np.zeros(len(sdf["perID_by_events"]), dtype=int), categories=[0] + ) else: tmp = pd.cut( sdf["perID_by_events"], bins=breaks, labels=labels, include_lowest=True @@ -541,13 +477,46 @@ def generate_breaks(min_number, max_number, min_breaks=5, max_breaks=9): # Round down min_number to the nearest multiple of magnitude min_aligned = int(min_number // magnitude * magnitude) - # Generate breakpoints + # Generate only breakpoints that fall inside the requested interval. + # Matplotlib expands an axis when ``set_ticks`` includes an out-of-range + # value, which previously turned a ~103 Mb grid into a 125 Mb grid. upper_bound = int(min_aligned + (threshold + 1) * magnitude) - breaks = list(range(min_aligned, upper_bound, int(magnitude))) + breaks = [ + value + for value in range(min_aligned, upper_bound, int(magnitude)) + if min_number <= value <= max_number + ] return breaks +def _requested_axis_bounds(requested_limit): + """Return an exact ``(start, end)`` pair when one was supplied.""" + + if isinstance(requested_limit, (tuple, list)): + if len(requested_limit) != 2: + raise ValueError("Axis bounds must contain exactly two values") + start, end = map(float, requested_limit) + if end <= start: + raise ValueError("Axis end must be greater than axis start") + return start, end + return None + + +def _data_axis_limits(sdf, requested_limit=None): + """Resolve plot limits, preserving exact region bounds when provided.""" + + requested_bounds = _requested_axis_bounds(requested_limit) + if requested_bounds is not None: + return requested_bounds + + minimum = float(min(sdf["q_st"].min(), sdf["r_st"].min())) + maximum = float(max(sdf["q_en"].max(), sdf["r_en"].max())) + if requested_limit: + maximum = max(maximum, float(requested_limit)) + return minimum, maximum + + def make_dot( sdf, name_x, @@ -562,10 +531,12 @@ def make_dot( width, is_pairwise, ): + display_x = display_sequence_name(name_x) + display_y = display_sequence_name(name_y) if is_pairwise: - title_name = f"Comparative Plot: {name_x} vs {name_y}" + title_name = f"Comparative Plot: {display_x} vs {display_y}" else: - title_name = f"Self-Identity Plot: {name_x}" + title_name = f"Self-Identity Plot: {display_x}" title_length = 2 * width if len(title_name) > 50: title_length = 1.5 * width @@ -591,24 +562,27 @@ def make_dot( new_hexcodes = hexcodes[::-1] if palette_orientation == "-" else hexcodes if colors: new_hexcodes = colors # Override colors if provided - if not xlim: - xlim = 0 - # Determine maximum genomic position for scaling - min_val = min(sdf["q_st"].min(), sdf["r_st"].min()) - max_val = max(sdf["q_en"].max(), sdf["r_en"].max(), xlim) + # Determine the exact genomic interval. A two-value limit is supplied by + # FASTA mode so blank edge windows do not shrink or extend the plot. + min_val, max_val = _data_axis_limits(sdf, xlim) # If user provides breaks, convert to ints if not breaks: breaks = generate_breaks(int(min_val), int(max_val)) else: - [int(x) for x in breaks] - xlim = xlim or 0 + breaks = [int(x) for x in breaks] # Compute window size (handling exceptions) try: window = max(sdf["q_en"] - sdf["q_st"]) except ValueError: # Empty dataframe case return ggplot(aes(x=[], y=[])) + theme_minimal() + # Region-qualified names remain in BEDPE data and filenames, but plot + # headings should show only the underlying FASTA identifier. + sdf = sdf.copy() + sdf["q"] = sdf["q"].map(display_sequence_name) + sdf["r"] = sdf["r"].map(display_sequence_name) + # Determine axis label scale based on genomic position size if max_val < 200_000: x_label = "Genomic Position (Kbp)" @@ -625,14 +599,14 @@ def make_dot( plot_background=element_blank(), panel_background=element_blank(), axis_line=element_line(color="black"), - axis_text=element_text(family=["DejaVu Sans"], size=width), + axis_text=element_text(family=["DejaVu Sans"], size=width * 2), axis_ticks_major=element_line( size=(width), color="black" ), # Increased tick length title=element_text( family=["DejaVu Sans"], size=title_length, hjust=0.5 ), # Center title - axis_title_x=element_text(size=(width * 1.4), family=["DejaVu Sans"]), + axis_title_x=element_text(size=(width * 2.8), family=["DejaVu Sans"]), strip_background=element_blank(), # Remove facet strip background strip_text=element_text( size=(width * 1.2), family=["DejaVu Sans"] @@ -656,9 +630,9 @@ def make_dot( + labs(x=x_label, y="", title=title_name) ) - # Select either geom_raster or geom_tile depending on deraster flag - p = ggplot_args + (geom_tile if deraster else geom_raster)( - aes(x="q_st", y="r_st", fill="discrete", height=window, width=window) + p = ggplot_args + _dotplot_tiles( + aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), + deraster, ) return p @@ -676,6 +650,7 @@ def make_dot_grid( deraster, width, ): + title_name = display_sequence_name(title_name) # Select the color palette if hasattr(diverging, palette): function_name = getattr(diverging, palette) @@ -706,7 +681,7 @@ def make_dot_grid( if not breaks: breaks = generate_breaks(int(min_val), int(max_val)) else: - [int(x) for x in breaks] + breaks = [int(x) for x in breaks] xlim = xlim or 0 # Compute window size (handling exceptions) try: @@ -754,14 +729,189 @@ def make_dot_grid( + labs(x="", y="", title="") ) - # Select either geom_raster or geom_tile depending on deraster flag - p = ggplot_args + (geom_tile if deraster else geom_raster)( - aes(x="q_st", y="r_st", fill="discrete", height=window, width=window) + p = ggplot_args + _dotplot_tiles( + aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), + deraster, ) return p +def direction_dataframe( + canonical_matrix, + forward_matrix, + window_size, + name_x, + name_y, + self_identity, + x_offset=0, + y_offset=0, +): + """Build plotting records that distinguish forward and reverse matches. + + Canonical k-mers match in either orientation, while forward-only k-mers + match only same-strand sequence. A canonical hit missing from the + forward-only matrix therefore represents a reverse-orientation match. + """ + canonical_matrix = np.asarray(canonical_matrix, dtype=float) + forward_matrix = np.asarray(forward_matrix, dtype=float) + if canonical_matrix.shape != forward_matrix.shape: + raise ValueError("Canonical and forward matrices must have matching shapes") + + records = [] + for x_index, y_index in np.argwhere(canonical_matrix > 0): + if self_identity and x_index > y_index: + continue + query_start = x_index * window_size + x_offset + reference_start = y_index * window_size + y_offset + records.append( + { + "q": name_x, + "q_st": query_start, + "q_en": query_start + window_size - 1, + "r": name_y, + "r_st": reference_start, + "r_en": reference_start + window_size - 1, + "direction": ( + "Forward" if forward_matrix[x_index, y_index] > 0 else "Reverse" + ), + } + ) + + dataframe = pd.DataFrame.from_records( + records, + columns=["q", "q_st", "q_en", "r", "r_st", "r_en", "direction"], + ) + dataframe["direction"] = pd.Categorical( + dataframe["direction"], categories=["Forward", "Reverse"], ordered=True + ) + return dataframe + + +def create_direction_plot( + canonical_matrix, + forward_matrix, + window_size, + directory, + name_x, + name_y, + self_identity, + width, + dpi, + vector_format, + deraster=False, + xlim=None, + axes_labels=None, + x_offset=0, + y_offset=0, +): + """Save a blue/pink plot showing match orientation.""" + dataframe = direction_dataframe( + canonical_matrix, + forward_matrix, + window_size, + name_x, + name_y, + self_identity, + x_offset, + y_offset, + ) + if dataframe.empty: + print(f"No directional matches found for {name_x} and {name_y}. Skipping.\n") + return None + + dataframe = dataframe.assign( + q_position=dataframe["q_st"] + window_size / 2, + r_position=dataframe["r_st"] + window_size / 2, + ) + requested_bounds = _requested_axis_bounds(xlim) + if requested_bounds is not None: + min_val, max_val = requested_bounds + else: + min_val = min(dataframe["q_st"].min(), dataframe["r_st"].min()) + max_val = max( + dataframe["q_en"].max() + 1, + dataframe["r_en"].max() + 1, + xlim or 0, + ) + breaks = ( + [int(value) for value in axes_labels] + if axes_labels + else generate_breaks(int(min_val), int(max_val)) + ) + title = ( + f"Direction Plot: {display_sequence_name(name_x)}" + if self_identity + else f"Direction Plot: {display_sequence_name(name_x)} vs " + f"{display_sequence_name(name_y)}" + ) + plot = ( + ggplot(dataframe) + + _dotplot_tiles( + aes( + x="q_position", + y="r_position", + fill="direction", + height=window_size, + width=window_size, + ), + deraster, + ) + + scale_fill_manual( + values={"Forward": "#2166AC", "Reverse": "#D01C8B"}, + name="Direction", + ) + + scale_x_continuous( + labels=make_scale, limits=[min_val, max_val], breaks=breaks + ) + + scale_y_continuous( + labels=make_scale, limits=[min_val, max_val], breaks=breaks + ) + + coord_fixed(ratio=1) + + labs( + x="Genomic Position", + y="", + title=title, + caption="Blue: forward Pink: reverse", + ) + + theme_light() + + theme( + legend_position="none", + panel_grid_major=element_blank(), + panel_grid_minor=element_blank(), + axis_text=element_text(family=["DejaVu Sans"], size=width), + title=element_text(size=width * 1.4, hjust=0.5), + axis_title_x=element_text(size=width * 1.2, family=["DejaVu Sans"]), + ) + ) + + os.makedirs(directory, exist_ok=True) + filename = ( + f"{name_x}_DIRECTION" if self_identity else f"{name_x}_{name_y}_DIRECTION" + ) + prefix = os.path.join(directory, filename) + ggsave( + plot, + width=width, + height=width, + dpi=dpi, + format=vector_format, + filename=f"{prefix}.{vector_format}", + verbose=False, + ) + ggsave( + plot, + width=width, + height=width, + dpi=dpi, + format="png", + filename=f"{prefix}.png", + verbose=False, + ) + print(f"Direction plots saved to {prefix}.png and {prefix}.{vector_format}.\n") + return plot + + def make_dot_final( sdf, width, @@ -773,6 +923,32 @@ def make_dot_final( transpose=False, deraster=False, ): + if sdf.empty: + max_val = xlim or 1 + if not breaks: + breaks = generate_breaks(0, int(max_val)) + else: + breaks = [int(x) for x in breaks] + return ( + ggplot(sdf) + + geom_blank() + + scale_x_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) + + scale_y_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) + + coord_fixed(ratio=1) + + labs(x=None, y=None, title=None) + + theme( + panel_grid_major=element_blank(), + panel_grid_minor=element_blank(), + plot_background=element_blank(), + panel_background=element_blank(), + axis_line=element_line(color="black"), + axis_text=element_text(family=["DejaVu Sans"], size=width), + axis_ticks_major=element_line(), + axis_title_x=element_blank(), + axis_title_y=element_blank(), + ) + ) + if hasattr(diverging, palette): function_name = getattr(diverging, palette) elif hasattr(qualitative, palette): @@ -802,7 +978,7 @@ def make_dot_final( if not breaks: breaks = generate_breaks(int(min_val), int(max_val)) else: - [int(x) for x in breaks] + breaks = [int(x) for x in breaks] xlim = xlim or 0 max_val = max(sdf["q_en"].max(), sdf["r_en"].max(), xlim) @@ -824,8 +1000,9 @@ def make_dot_final( if deraster: p = ( ggplot(sdf) - + geom_tile( - aes(x=x_col, y=y_col, fill="discrete", height=window, width=window) + + _dotplot_tiles( + aes(x=x_col, y=y_col, fill="discrete", height=window, width=window), + deraster, ) + scale_color_discrete(guide=False) + scale_fill_manual(values=new_hexcodes, guide=False) @@ -848,8 +1025,9 @@ def make_dot_final( else: p = ( ggplot(sdf) - + geom_raster( - aes(x=x_col, y=y_col, fill="discrete", height=window, width=window) + + _dotplot_tiles( + aes(x=x_col, y=y_col, fill="discrete", height=window, width=window), + deraster, ) + scale_color_discrete(guide=False) + scale_fill_manual(values=new_hexcodes, guide=False) @@ -887,6 +1065,7 @@ def make_tri( deraster, width, ): + title_name = display_sequence_name(title_name) # Select the color palette if hasattr(diverging, palette): function_name = getattr(diverging, palette) @@ -917,7 +1096,7 @@ def make_tri( if not breaks: breaks = generate_breaks(int(min_val), int(max_val)) else: - [int(x) for x in breaks] + breaks = [int(x) for x in breaks] xlim = xlim or 0 # Compute window size (handling exceptions) try: @@ -936,8 +1115,9 @@ def make_tri( if not deraster: tri = ( ggplot(sdf) - + geom_raster( + + _dotplot_tiles( aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), + deraster, alpha=1.0, ) # Ensure full opacity + scale_fill_manual(values=new_hexcodes, guide=False) @@ -1001,8 +1181,9 @@ def make_tri( else: tri = ( ggplot(sdf) - + geom_tile( + + _dotplot_tiles( aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), + deraster, alpha=1.0, ) # Ensure full opacity + scale_fill_manual(values=new_hexcodes, guide=False) @@ -1071,194 +1252,8 @@ def make_tri( return tri, axis -def rotate_vectorized_tri(svg_path, scale_x, scale_y): - # Define SVG namespace - ns = {"svg": "http://www.w3.org/2000/svg"} - - # Parse the SVG file - tree = ET.parse(svg_path) - root = tree.getroot() - - # Find all elements with id="PolyCollection_1" - g_elements = root.find(".//svg:g[@id='PolyCollection_1']", namespaces=ns) - - if g_elements is not None: - # Apply the rotation transform to the group - scale_factor = 1 / math.sqrt(2) - transform = f"rotate(45 0 0) translate({scale_x}, {scale_y}) scale({scale_factor}, {scale_factor})" - g_elements.set("transform", transform) - - # Save the modified SVG - - viewBox = root.get("viewBox") - - if viewBox: - min_x, min_y, width, height = map(float, viewBox.split()) - new_min_y = min_y + height / 2 # Move down by half the height - new_height = height / 2 # Reduce height by half - root.set("viewBox", f"{min_x} {new_min_y} {width} {new_height}") - else: - print("No viewBox found. Consider adding one manually.") - - # Hacky, but it works to halve the height - height_svg = root.get("height") - if height_svg: - current_height = re.match(r"(\d*\.?\d+)([a-zA-Z%]*)", height_svg) - if current_height: - numeric_height, unit = current_height.groups() - numeric_height = float(numeric_height) - - root.set("height", f"{numeric_height / 1.8}pt") - tree.write(svg_path) - - -def rotate_rasterized_tri(svg_path, shift_x, shift_y): - # Load the SVG file - tree = ET.parse(svg_path) - root = tree.getroot() - - # Namespace handling - ns = {"svg": "http://www.w3.org/2000/svg"} - # Find all image elements with base64 embedded data - for image in root.findall(".//svg:image", ns): - href = image.get("{http://www.w3.org/1999/xlink}href", "") - if href.startswith("data:image/png;base64,"): - # Get the current width and height of the image - width = float(image.get("width", 0)) / math.sqrt(2) - height = float(image.get("height", 0)) / math.sqrt(2) - # Set the new width and height - image.set("width", str(width)) - image.set("height", str(height)) - # Apply a 270-degree rotation (about the top-left corner of the image) - transform = image.get("transform", "") - new_transform = ( - f"rotate(45, 0, 0) translate({shift_x}, {shift_y}) {transform}" - if transform - else f"rotate(45, 0, {height}) translate({shift_x}, {shift_y})" - ) - image.set("transform", new_transform) - # Update viewbox - viewBox = root.get("viewBox") - - if viewBox: - min_x, min_y, width, height = map(float, viewBox.split()) - new_min_y = min_y + height / 2 # Move down by half the height - new_height = height / 2 # Reduce height by half - root.set("viewBox", f"{min_x} {new_min_y} {width} {new_height}") - else: - print("No viewBox found. Consider adding one manually.") - - # Hacky, but it works to halve the height - height_svg = root.get("height") - if height_svg: - current_height = re.match(r"(\d*\.?\d+)([a-zA-Z%]*)", height_svg) - if current_height: - numeric_height, unit = current_height.groups() - numeric_height = float(numeric_height) - - root.set("height", f"{numeric_height / 1.8}pt") - else: - print("Warning: Could not parse height attribute.") - - # Save the modified SVG back to the same file - tree.write(svg_path) - - -from lxml import etree - - -def append_svg(svg1_path, svg2_path, output_path): - """Appends SVG2 to the bottom of SVG1 without modifying its width.""" - # Load SVG1 and SVG2 - tree1 = etree.parse(svg1_path) - root1 = tree1.getroot() - - tree2 = etree.parse(svg2_path) - root2 = tree2.getroot() - - # Extract width and height of SVG1 - width1 = float(root1.get("width", "0").replace("pt", "")) - height1 = float(root1.get("height", "0").replace("pt", "")) - - # Extract width and height of SVG2 - width2 = float(root2.get("width", "0").replace("pt", "")) - height2 = float(root2.get("height", "0").replace("pt", "")) - - # Update SVG1 height to accommodate SVG2 - new_height = height1 + height2 - root1.set("height", f"{new_height}pt") - - # Create a translation group for SVG2 and shift it down - group = etree.Element("g", attrib={"transform": f"translate(0,{height1 + 400})"}) - for child in root2: - group.append(child) - - # Append translated SVG2 to SVG1 - root1.append(group) - - # Save the new merged SVG - tree1.write(output_path, pretty_print=True, xml_declaration=True, encoding="utf-8") - - -def get_svg_size(svg_path): - """Helper to extract width and height of an SVG in pt units.""" - import xml.etree.ElementTree as ET - - tree = ET.parse(svg_path) - root = tree.getroot() - width = root.get("width") - height = root.get("height") - return width, height - - -def parse_size(size_str): - """Convert '648pt' or '800px' -> float(648).""" - return float(re.sub(r"[a-zA-Z]+", "", size_str)) - - -def merge_annotation_tri(svg1_path, svg2_path, output_path, deraster, width): - """Merges two SVG files into a single SVG file with proper size.""" - - w1, h1 = get_svg_size(svg1_path) - w2, h2 = get_svg_size(svg2_path) - - # Ensure they are floats - w1, h1 = parse_size(w1), parse_size(h1) - w2, h2 = parse_size(w2), parse_size(h2) - # Determine total size - total_width = max(w1, w2) - total_height = h1 + h2 - # Create figure - fig = sg.SVGFigure(f"{total_width}px", f"{total_height}px") - - # Load SVGs - svg1 = sg.fromfile(svg1_path).getroot() - svg2 = sg.fromfile(svg2_path).getroot() - make_svg_background_transparent(svg2_path) - - # Position - if deraster: - # Its not perfect for width > 18, but good enough - adjust_svg1 = (0, -5 * (h1 / 6)) - adjust_svg2 = ((9 - (width / 2)), 5 * (h1 / 6)) - svg1.moveto(adjust_svg1[0], adjust_svg1[1]) - svg2.scale(1.077 + (width / 1000)) - svg2.moveto(adjust_svg2[0], adjust_svg2[1]) - else: - adjust_svg1 = (0, -5 * (h1 / 6)) - adjust_svg2 = (10 + (width / 4.5), 5 * (h1 / 6)) - svg1.moveto(adjust_svg1[0], adjust_svg1[1]) - svg2.moveto(adjust_svg2[0], adjust_svg2[1]) - scaling_factor = width / 2 - svg2.scale(1.034 + (scaling_factor / 1000)) - - # Append and save - fig.append([svg1, svg2]) - fig.set_size((f"{total_width}px", f"{total_height}px")) - fig.save(output_path) - - def make_tri_axis(sdf, title_name, palette, palette_orientation, colors, breaks, xlim): + title_name = display_sequence_name(title_name) if not breaks: breaks = True else: @@ -1394,10 +1389,156 @@ def make_hist(sdf, palette, palette_orientation, custom_colors, custom_breakpoin return p -def create_grid( +def _triangle_limits(sdf, xlim): + if sdf.empty: + raise ValueError("Cannot render a triangle plot without identity tiles") + requested_bounds = _requested_axis_bounds(xlim) + if requested_bounds is not None: + region_start, region_end = requested_bounds + else: + region_start = max(float(sdf["q_st"].min()), float(sdf["r_st"].min())) + region_end = max( + float(sdf["q_en"].max()), + float(sdf["r_en"].max()), + float(xlim or 0), + ) + if region_end <= region_start: + raise ValueError("Triangle plot end must be greater than its start") + return region_start, region_end + + +def _build_triangle_figure( + sdf, + title, + palette, + palette_orientation, + custom_colors, + axes_labels, + xlim, + deraster, + width, + annotation_df=None, + annotation_chrom=None, +): + """Build a triangle, optionally aligned with a BED annotation track.""" + region_start, region_end = _triangle_limits(sdf, xlim) + breaks = ( + [float(value) for value in axes_labels] + if axes_labels + else generate_breaks(int(region_start), int(region_end)) + ) + colors = _resolve_native_colors(palette, palette_orientation, custom_colors) + with_annotation = annotation_df is not None + layout = create_triangle_layout(width, with_annotation=with_annotation) + try: + draw_triangle_tiles( + layout.triangle_axis, + sdf, + colors, + rasterized=not deraster, + ) + configure_triangle_axis( + layout.triangle_axis, + region_start, + region_end, + breaks=breaks, + label=not with_annotation, + ) + layout.triangle_axis.set_title( + display_sequence_name(title), fontsize=max(10, width * 1.4) + ) + + if with_annotation: + if not annotation_chrom: + raise ValueError("An annotation chromosome is required") + draw_annotation_track( + layout.annotation_axis, + annotation_df, + annotation_chrom, + region_start, + region_end, + ) + layout.annotation_axis.set_xticks(breaks) + layout.annotation_axis.xaxis.set_major_formatter( + genomic_tick_formatter(region_end) + ) + _, unit = genomic_scale(region_end) + layout.annotation_axis.set_xlabel(f"Genomic Position ({unit})") + + layout.figure.subplots_adjust( + left=0.08, + right=0.98, + bottom=0.14 if not with_annotation else 0.10, + top=0.90, + hspace=0.05, + ) + except Exception: + plt.close(layout.figure) + raise + return layout.figure + + +def _ordered_grid_names(single_names, double_names): + names = [] + for name in single_names: + if name in names: + raise ValueError(f"Duplicate self comparison for sequence {name!r}") + names.append(name) + for pair in double_names: + if len(pair) != 2: + raise ValueError("Each grid comparison must name exactly two sequences") + for name in pair: + if name not in names: + names.append(name) + if not names: + raise ValueError("No sequence names were provided for the grid") + return names + + +def _read_grid_dataframe( + matrix, + is_bed, + palette, + palette_orientation, + is_freq, + custom_colors, + custom_breakpoints, +): + return read_df( + None if is_bed else [matrix], + palette, + palette_orientation, + is_freq, + custom_colors, + custom_breakpoints, + matrix if is_bed else None, + ) + + +def _grid_axis_limits(dataframes, requested_limit): + requested_bounds = _requested_axis_bounds(requested_limit) + if requested_bounds is not None: + return requested_bounds + if requested_limit: + return 0.0, float(requested_limit) + minima = [ + float(min(dataframe["q_st"].min(), dataframe["r_st"].min())) + for dataframe in dataframes + if not dataframe.empty + ] + maxima = [ + float(max(dataframe["q_en"].max(), dataframe["r_en"].max())) + for dataframe in dataframes + if not dataframe.empty + ] + if not maxima: + return 0.0, 1.0 + return min(minima), max(maxima) + + +def _build_grid_figure( singles, doubles, - directory, palette, palette_orientation, single_names, @@ -1411,283 +1552,201 @@ def create_grid( width, breaks, deraster, - vector_format, ): - new_index = [] - transpose_index = [] - check_pascal(singles, doubles) - # Singles can be empty if not selected - for i in range(len(single_names)): - for j in range(i + 1, len(single_names)): - try: - index = double_names.index([single_names[i], single_names[j]]) - transpose_index.append(0) - except: - index = double_names.index([single_names[j], single_names[i]]) - transpose_index.append(1) - - new_index.append(index) - - single_list = [] - double_list = [] - single_heatmap_list = [] - normal_heatmap_list = [] - transpose_heatmap_list = [] - for matrix in singles: - if is_bed: - df = read_df( - None, - palette, - palette_orientation, - is_freq, - custom_colors, - custom_breakpoints, - matrix, - ) - else: - df = read_df( - [matrix], - palette, - palette_orientation, - is_freq, - custom_colors, - custom_breakpoints, - None, - ) - single_list.append(df) - for matrix in doubles: - if is_bed: - df = read_df( - None, - palette, - palette_orientation, - is_freq, - custom_colors, - custom_breakpoints, - matrix, - ) - else: - df = read_df( - [matrix], - palette, - palette_orientation, - is_freq, - custom_colors, - custom_breakpoints, - None, - ) - double_list.append(df) - # This is the diagonals - for plot in single_list: - heatmap = make_dot_final( - sdf=plot, - width=width, - palette=palette, - palette_orientation=palette_orientation, - colors=custom_colors, - breaks=axes_label, - xlim=xlim, - transpose=False, - deraster=deraster, - ) - single_heatmap_list.append(heatmap) - for indie in new_index: - xd = new_index.index(indie) - # These are the non-diagonal transposes and normals! - if transpose_index[xd] == 0: - heatmap = make_dot_final( - sdf=double_list[indie], - width=width, - palette=palette, - palette_orientation=palette_orientation, - colors=custom_colors, - breaks=axes_label, - xlim=xlim, - transpose=True, - deraster=deraster, - ) - normal_heatmap_list.append(heatmap) - heatmap_t = make_dot_final( - double_list[indie], - width, - palette, - palette_orientation, - custom_colors, - axes_label, - xlim, - transpose=False, - deraster=deraster, - ) - transpose_heatmap_list.append(heatmap_t) - else: - heatmap = make_dot_final( - double_list[indie], - width, - palette, - palette_orientation, - custom_colors, - axes_label, - xlim, - transpose=True, - deraster=deraster, - ) - normal_heatmap_list.append(heatmap) - heatmap_t = make_dot_final( - double_list[indie], - width, - palette, - palette_orientation, - custom_colors, - axes_label, - xlim, - transpose=False, - deraster=deraster, - ) - transpose_heatmap_list.append(heatmap_t) - - assert len(transpose_heatmap_list) == len(normal_heatmap_list) - single_length = len(single_heatmap_list) - if single_length == 0: - single_length = reverse_pascal(len(normal_heatmap_list)) - - normal_counter = 0 - trans_counter = 0 - trans_to_use = transpose_order(normal_heatmap_list) - start_grid = pw.Brick(figsize=(9, 9)) - n = single_length * single_length - - if n > 9: - print(f"This might take a while\n...\n") - - printProgressBar(0, n, prefix="Progress:", suffix="Complete", length=40) - tots = 0 - col_names = pw.Brick(figsize=(width / 4.5, width)) - row_names = pw.Brick(figsize=(width, width / 4.5)) - for i in range(single_length): - row_grid = pw.Brick(figsize=(width, width)) - for j in range(single_length): - if i == j: - if len(single_heatmap_list) == 0: - g1 = pw.Brick(figsize=(width, width)) - else: - g1 = pw.load_ggplot(single_heatmap_list[i], figsize=(width, width)) + """Build a native Matplotlib comparison grid and return its axes. - elif i < j: - g1 = pw.load_ggplot( - normal_heatmap_list[normal_counter], figsize=(width, width) - ) - normal_counter += 1 - elif i > j: - g1 = pw.load_ggplot( - transpose_heatmap_list[trans_to_use[trans_counter]], - figsize=(width, width), - ) - trans_counter += 1 - if j == 0: - row_grid = g1 - else: - row_grid = row_grid | g1 - tots += 1 - printProgressBar(tots, n, prefix="Progress:", suffix="Complete", length=40) - if i == 0: - start_grid = row_grid - else: - start_grid = row_grid / start_grid - for w in range(single_length): - p1 = ( - ggplot() - + geom_blank() - + annotate( # Use geom_blank to create a plot with no data - "text", - x=0, - y=0, - label=single_names[w], - size=width * 3, - angle=90, - ha="center", - va="center", - ) - + theme( - # Center the plot area and make backgrounds transparent - axis_title_x=element_blank(), - axis_title_y=element_blank(), - axis_ticks=element_blank(), - axis_text=element_blank(), - plot_background=element_rect( - fill="none" - ), # Transparent plot background - panel_background=element_rect( - fill="none" - ), # Transparent panel background - panel_grid=element_blank(), - aspect_ratio=0.5, # Adjust the aspect ratio for the desired width/height - ) - + coord_flip() # Rotate the plot 90 degrees counterclockwise + ``width`` controls the complete square figure, not each panel, so memory + use does not grow quadratically in pixels as sequences are added. + """ + if len(singles) != len(single_names): + raise ValueError("Self-comparison matrices and names must have equal lengths") + if len(doubles) != len(double_names): + raise ValueError("Pairwise matrices and names must have equal lengths") + + names = _ordered_grid_names(single_names, double_names) + # Keep columns in input order and reverse rows so self-comparisons occupy + # the anti-diagonal: bottom-left to top-right. + row_names = list(reversed(names)) + single_frames = {} + for name, matrix in zip(single_names, singles): + single_frames[name] = _read_grid_dataframe( + matrix, + is_bed, + palette, + palette_orientation, + is_freq, + custom_colors, + custom_breakpoints, ) - p2 = ( - ggplot() - + geom_blank() - + annotate( # Use geom_blank to create a plot with no data - "text", - x=0, - y=0, - label=single_names[w], - size=width * 3, - ha="center", - va="center", - ) - + theme( - # Center the plot area and make backgrounds transparent - axis_title_x=element_blank(), - axis_title_y=element_blank(), - axis_ticks=element_blank(), - axis_text=element_blank(), - plot_background=element_rect( - fill="none" - ), # Transparent plot background - panel_background=element_rect( - fill="none" - ), # Transparent panel background - panel_grid=element_blank(), - aspect_ratio=0.5, # Adjust the aspect ratio for the desired width/height + + pair_frames = {} + pair_orientations = {} + for pair, matrix in zip(double_names, doubles): + query_name, reference_name = pair + if query_name == reference_name: + raise ValueError("Pairwise grid comparisons must use two distinct names") + key = frozenset((query_name, reference_name)) + if key in pair_frames: + raise ValueError( + f"Duplicate pairwise comparison for {query_name!r} and {reference_name!r}" ) + pair_frames[key] = _read_grid_dataframe( + matrix, + is_bed, + palette, + palette_orientation, + is_freq, + custom_colors, + custom_breakpoints, ) - g1 = pw.load_ggplot(p1, figsize=(width / 4.5, width)) - g2 = pw.load_ggplot(p2, figsize=(width, width / 4.5)) + pair_orientations[key] = (query_name, reference_name) + + missing_pairs = [ + (names[row], names[column]) + for row in range(len(names)) + for column in range(row + 1, len(names)) + if frozenset((names[row], names[column])) not in pair_frames + ] + if missing_pairs: + formatted = ", ".join(f"{left}/{right}" for left, right in missing_pairs) + raise ValueError(f"Missing pairwise grid comparisons: {formatted}") + + all_frames = list(single_frames.values()) + list(pair_frames.values()) + axis_start, axis_end = _grid_axis_limits(all_frames, xlim) + axis_breaks = axes_label or breaks + if not axis_breaks: + axis_breaks = generate_breaks(int(axis_start), int(axis_end)) + axis_breaks = [float(value) for value in axis_breaks] + colors = _resolve_native_colors(palette, palette_orientation, custom_colors) + + grid_size = len(names) + figure_width = max(float(width), 2.0) + heading_size = max(6.0, min(12.0, figure_width * 1.2)) + # Numeric genomic labels are intentionally twice the previous size. The + # inverse grid-size factor keeps larger grids proportionate. + tick_size = max(8.0, min(18.0, figure_width * 3.0 / grid_size)) + figure, axes = plt.subplots( + grid_size, + grid_size, + figsize=(figure_width, figure_width), + sharex=True, + sharey=True, + squeeze=False, + ) + try: + for row, row_name in enumerate(row_names): + for column, column_name in enumerate(names): + axis = axes[row, column] + dataframe = None + transpose = False + if row_name == column_name: + dataframe = single_frames.get(row_name) + else: + key = frozenset((row_name, column_name)) + dataframe = pair_frames[key] + query_name, reference_name = pair_orientations[key] + if (query_name, reference_name) == (column_name, row_name): + transpose = False + elif (query_name, reference_name) == (row_name, column_name): + transpose = True + else: + raise ValueError( + f"Grid comparison names do not match {row_name!r}/{column_name!r}" + ) + + if dataframe is not None and not dataframe.empty: + draw_rectangular_tiles( + axis, + dataframe, + colors, + transpose=transpose, + rasterized=not deraster, + ) + configure_dotplot_axis( + axis, + axis_start, + axis_end, + breaks=axis_breaks, + show_x=row == grid_size - 1, + show_y=column == 0, + ) + axis.tick_params(axis="both", labelsize=tick_size) + axis.grid(False) + if row == 0: + axis.set_title( + display_sequence_name(column_name), fontsize=heading_size + ) + if column == 0: + axis.set_ylabel( + display_sequence_name(row_name), fontsize=heading_size + ) - if w == 0: - col_names = g1 - row_names = g2 - else: - col_names = g1 / col_names - row_names = row_names | g2 - # Create a ghost 2x2 - pghost = ( - ggplot() - + geom_blank() - + annotate( # Use geom_blank to create a plot with no data - "text", x=0, y=0, label="", size=32, ha="center", va="center" - ) - + theme( - # Center the plot area and make backgrounds transparent - axis_title_x=element_blank(), - axis_title_y=element_blank(), - axis_ticks=element_blank(), - axis_text=element_blank(), - plot_background=element_rect(fill="none"), # Transparent plot background - panel_background=element_rect(fill="none"), # Transparent panel background - panel_grid=element_blank(), - aspect_ratio=0.5, # Adjust the aspect ratio for the desired width/height + figure.subplots_adjust( + left=0.10, + right=0.98, + bottom=0.08, + top=0.92, + wspace=0.08, + hspace=0.08, ) + _fit_grid_sequence_labels(figure, axes) + except Exception: + plt.close(figure) + raise + return figure, axes + + +def create_grid( + singles, + doubles, + directory, + palette, + palette_orientation, + single_names, + double_names, + is_freq, + xlim, + custom_colors, + custom_breakpoints, + axes_label, + is_bed, + width, + breaks, + deraster, + vector_format, + dpi=300, +): + figure, axes = _build_grid_figure( + singles=singles, + doubles=doubles, + palette=palette, + palette_orientation=palette_orientation, + single_names=single_names, + double_names=double_names, + is_freq=is_freq, + xlim=xlim, + custom_colors=custom_colors, + custom_breakpoints=custom_breakpoints, + axes_label=axes_label, + is_bed=is_bed, + width=width, + breaks=breaks, + deraster=deraster, ) - ghosty = pw.load_ggplot(pghost, figsize=(2, 2)) - col_names = ghosty / col_names - start_grid = col_names | (row_names / start_grid) - gridname = f"{single_length}x{single_length}_GRID" - print(f"\nGrid complete! Saving to {directory}/{gridname}...\n") - start_grid.savefig(f"{directory}/{gridname}.png") - start_grid.savefig(f"{directory}/{gridname}.{vector_format}", format=vector_format) - print(f"Grid saved successfully!\n") + grid_size = axes.shape[0] + grid_prefix = os.path.join(directory, f"{grid_size}x{grid_size}_GRID") + print(f"\nGrid complete! Saving to {grid_prefix}...\n") + try: + save_figure_pair( + figure, + grid_prefix, + vector_format, + dpi, + bbox_inches="tight", + ) + finally: + plt.close(figure) + print("Grid saved successfully!\n") def create_plots( @@ -1732,62 +1791,46 @@ def create_plots( sdf, palette, palette_orientation, custom_colors, custom_breakpoints ) + annotation_track_created = False + annotation_bed_df = None + annotation_chrom = None # Just doing triangle plots for now. if annotation: - print(f"Generating ini file for annotation track:\n") + print("Generating annotation track:\n") iniprefix = plot_filename - inifile = generate_ini_file( - bedfile=annotation, - ininame=iniprefix, - chrom=name_x, - ) - min_val = max(sdf["q_st"].min(), sdf["r_st"].min()) - max_val = max(sdf["q_en"].max(), sdf["r_en"].max()) - if not xlim: - xlim = max_val - region = [(name_x.split(":")[0], min_val, xlim)] - if inifile: - try: - # Check if the BED file has valid intervals for this region first - bed_df = read_annotation_bed(annotation) - chrom_name = name_x.split(":")[0] - - # Filter for the chromosome and region of interest - valid_intervals = bed_df[ - (bed_df["chrom"] == chrom_name) - & (bed_df["end"] >= min_val) - & (bed_df["start"] <= xlim) - ] - - if valid_intervals.empty: - print( - f"No valid intervals found in {annotation} for region {chrom_name}:{min_val}-{xlim}.\n" - ) - print("Skipping annotation track generation.\n") - else: - bed_track = run_pygenometracks( - inifile=inifile, - region=region, - output_file=f"{iniprefix}_ANNOTATION_TRACK.svg", - width=width * 2.05, - ) - bed_track.plot( - f"{iniprefix}_ANNOTATION_TRACK.svg", - name_x.split(":")[0], - min_val, - xlim, - ) - bed_track.plot( - f"{iniprefix}_ANNOTATION_TRACK.png", - name_x.split(":")[0], - min_val, - xlim, - ) - print(f"\nAnnotation track saved to {iniprefix}_ANNOTATION_TRACK\n") - - except Exception as e: - print(f"Error processing annotation file {annotation}: {e}\n") + requested_bounds = _requested_axis_bounds(xlim) + if requested_bounds is not None: + min_val, annotation_end = requested_bounds + else: + min_val = max(sdf["q_st"].min(), sdf["r_st"].min()) + max_val = max(sdf["q_en"].max(), sdf["r_en"].max()) + annotation_end = xlim or max_val + chrom_name = name_x.split(":")[0] + try: + bed_df = read_annotation_bed(annotation) + annotation_track_created = render_annotation_track( + bed_df=bed_df, + chrom=chrom_name, + region_start=min_val, + region_end=annotation_end, + output_prefix=f"{iniprefix}_ANNOTATION_TRACK", + width_cm=width * 2.05, + dpi=dpi, + vector_format=vector_format, + ) + if annotation_track_created: + annotation_bed_df = bed_df + annotation_chrom = chrom_name + print(f"\nAnnotation track saved to {iniprefix}_ANNOTATION_TRACK\n") + else: + print( + f"No valid intervals found in {annotation} for region " + f"{chrom_name}:{min_val}-{annotation_end}.\n" + ) print("Skipping annotation track generation.\n") + except Exception as e: + print(f"Error processing annotation file {annotation}: {e}\n") + print("Skipping annotation track generation.\n") if is_pairwise: heatmap = make_dot( @@ -1863,18 +1906,6 @@ def create_plots( print( f"Producing dotplots with derasterization turned off. This may take a while...\n" ) - tri_plot = make_tri( - sdf, - plot_filename, - palette, - palette_orientation, - custom_colors, - axes_labels, - xlim, - axes_tick_number, - deraster, - width, - ) full_plot = make_dot( check_st_en_equality(sdf), name_x, @@ -1908,128 +1939,52 @@ def create_plots( verbose=False, ) tri_prefix = f"{plot_filename}_TRI" - ggsave( - tri_plot[0], + triangle_figure = _build_triangle_figure( + sdf=sdf, + title=name_x, + palette=palette, + palette_orientation=palette_orientation, + custom_colors=custom_colors, + axes_labels=axes_labels, + xlim=xlim, + deraster=deraster, width=width, - height=width, - dpi=dpi, - format="svg", - filename=f"{tri_prefix}.svg", - verbose=False, ) - if annotation: - anno_prefix = f"{plot_filename}_PRE_ANNOTATED" - annotated_tri = tri_plot[0] + theme( - axis_title_x=element_blank(), - axis_line_x=element_blank(), - axis_text_x=element_blank(), - axis_ticks_minor_x=element_blank(), - axis_ticks=element_blank(), + try: + save_figure_pair( + triangle_figure, + tri_prefix, + vector_format, + dpi, + bbox_inches="tight", ) - ggsave( - annotated_tri, + finally: + plt.close(triangle_figure) + + if annotation_track_created: + annotated_figure = _build_triangle_figure( + sdf=sdf, + title=name_x, + palette=palette, + palette_orientation=palette_orientation, + custom_colors=custom_colors, + axes_labels=axes_labels, + xlim=xlim, + deraster=deraster, width=width, - height=width, - dpi=dpi, - format="svg", - filename=f"{anno_prefix}.svg", - verbose=False, - ) - # These scaling values were determined thorugh much trial and error. Please don't delete :) - if deraster: - scaling_values = (46.62 * width, -3.75 * width) - rotate_vectorized_tri( - f"{tri_prefix}.svg", scaling_values[0], scaling_values[1] + annotation_df=annotation_bed_df, + annotation_chrom=annotation_chrom, ) - if annotation: - rotate_vectorized_tri( - f"{anno_prefix}.svg", scaling_values[0], scaling_values[1] - ) - try: - cairosvg.svg2png( - url=f"{tri_prefix}.svg", write_to=f"{tri_prefix}.png", dpi=dpi - ) - except: - print(f"Error installing cairosvg. Unable to convert svg file. \n") - else: - scaling_values = (44.6 * width, -23 * width) - rotate_rasterized_tri( - f"{tri_prefix}.svg", scaling_values[0], scaling_values[1] - ) - if annotation: - if annotation: - rotate_rasterized_tri( - f"{anno_prefix}.svg", scaling_values[0], scaling_values[1] - ) try: - cairosvg.svg2png( - url=f"{tri_prefix}.svg", write_to=f"{tri_prefix}.png", dpi=dpi + save_figure_pair( + annotated_figure, + f"{tri_prefix}_ANNOTATED", + vector_format, + dpi, + bbox_inches="tight", ) - except: - print(f"Error installing cairosvg. Unable to convert svg file. \n") - if annotation: - # Only merge if annotation was successfully created - if os.path.exists(f"{iniprefix}_ANNOTATION_TRACK.svg"): - make_svg_background_transparent(f"{iniprefix}_ANNOTATION_TRACK.svg") - merge_annotation_tri( - f"{anno_prefix}.svg", - f"{iniprefix}_ANNOTATION_TRACK.svg", - f"{tri_prefix}_ANNOTATED.svg", - deraster, - width, - ) - if os.path.exists(f"{anno_prefix}.svg"): - os.remove(f"{anno_prefix}.svg") - cairosvg.svg2png( - url=f"{tri_prefix}_ANNOTATED.svg", - write_to=f"{tri_prefix}_ANNOTATED.png", - dpi=dpi, - ) - try: - if vector_format != "svg": - if vector_format == "pdf": - cairosvg.svg2pdf( - url=f"{tri_prefix}_ANNOTATED.svg", - write_to=f"{tri_prefix}_ANNOTATED.pdf", - ) - cairosvg.svg2pdf( - url=f"{iniprefix}_ANNOTATION_TRACK.svg", - write_to=f"{iniprefix}_ANNOTATION_TRACK.pdf", - ) - elif vector_format == "ps": - cairosvg.svg2ps( - url=f"{tri_prefix}_ANNOTATED.svg", - write_to=f"{tri_prefix}_ANNOTATED.ps", - ) - cairosvg.svg2ps( - url=f"{tri_prefix}_ANNOTATION_TRACK.svg", - write_to=f"{tri_prefix}_ANNOTATION_TRACK.ps", - ) - if os.path.exists(f"{iniprefix}_ANNOTATION_TRACK.svg"): - os.remove(f"{iniprefix}_ANNOTATION_TRACK.svg") - if os.path.exists(f"{tri_prefix}_ANNOTATED.svg"): - os.remove(f"{tri_prefix}_ANNOTATED.svg") - except Exception as e: - print(f"Error converting annotated SVG: {e}") - else: - print("Annotation file not created, skipping merge step.") - # Convert from svg to selected vector format. Ignore error if user has issues with cairosvg. - try: - if vector_format != "svg": - if vector_format == "pdf": - cairosvg.svg2pdf( - url=f"{tri_prefix}.svg", write_to=f"{tri_prefix}.pdf" - ) - if os.path.exists(f"{tri_prefix}.svg"): - os.remove(f"{tri_prefix}.svg") - elif vector_format == "ps": - cairosvg.svg2pdf( - url=f"{tri_prefix}.svg", write_to=f"{tri_prefix}.ps" - ) - if os.path.exists(f"{tri_prefix}.svg"): - os.remove(f"{tri_prefix}.svg") - except: - pass + finally: + plt.close(annotated_figure) if no_hist: print( diff --git a/tests/test_algorithms.py b/tests/test_algorithms.py new file mode 100644 index 0000000..4483ac9 --- /dev/null +++ b/tests/test_algorithms.py @@ -0,0 +1,269 @@ +import numpy as np +import pytest + +from moddotplot.estimate_identity import ( + containment_neighbors, + convertMatrixToBed, + createSelfMatrix, + pairwiseContainmentMatrix, + partitionOverlaps, + populateModimizers, +) +from moddotplot.parse_fasta import generateKmersFromFasta, printProgressBar + + +def test_bed_conversion_clamps_partial_windows_to_exact_region_end(): + bed = convertMatrixToBed( + np.ones((2, 2)), + window_size=100, + id_threshold=80, + x_name="query", + y_name="reference", + self_identity=False, + x_offset=101, + y_offset=201, + x_end=250, + y_end=350, + ) + + assert max(row[2] for row in bed[1:]) == 250 + assert max(row[5] for row in bed[1:]) == 350 + + +def test_populate_modimizers_returns_denser_recursive_fallback(): + result = populateModimizers( + partition=[1, 2, 3, 4], + sparsity=4, + ambiguous=True, + expectation=4, + k=1, + ) + + assert result == {2, 4} + + +def test_populate_modimizers_can_fall_back_all_the_way_to_sparsity_one(): + result = populateModimizers( + partition=[1, 3, 5], + sparsity=8, + ambiguous=True, + expectation=10, + k=1, + ) + + assert result == {1, 3, 5} + + +def test_populate_modimizers_empty_partition_terminates_at_sparsity_one(): + assert populateModimizers([], 8, True, 10, 1) == set() + + +def test_partition_overlaps_uses_consistent_genomic_window_boundaries(): + # Twenty-six 5-mers represent a 30-base sequence. Each 10-base window + # contains six internal 5-mers; boundary-spanning k-mers are excluded. + kmers = list(range(26)) + + assert partitionOverlaps(kmers, win=10, delta=0, seq_len=26, k=5) == [ + list(range(0, 6)), + list(range(10, 16)), + list(range(20, 26)), + ] + + +def test_partition_overlaps_expands_each_window_in_genomic_coordinates(): + kmers = list(range(26)) + + assert partitionOverlaps(kmers, win=10, delta=0.5, seq_len=26, k=5) == [ + list(range(0, 11)), + list(range(5, 21)), + list(range(15, 26)), + ] + + +def test_partition_overlaps_omits_trailing_fragment_without_a_kmer(): + # Eleven bases produce seven 5-mers. With a 10-base window, the final + # one-base fragment cannot produce another partition. + assert partitionOverlaps(list(range(7)), win=10, delta=0, seq_len=7, k=5) == [ + list(range(6)) + ] + + +@pytest.mark.parametrize(("win", "k"), [(0, 5), (10, 0)]) +def test_partition_overlaps_rejects_nonpositive_sizes(win, k): + with pytest.raises(ValueError): + partitionOverlaps([1, 2, 3], win=win, delta=0, seq_len=3, k=k) + + +def test_containment_uses_matching_flanks_outside_core_windows(): + core_a = {1, 2, 3, 4} + core_b = {5, 6, 7, 8} + + result = containment_neighbors( + core_a, + core_b, + core_a | core_b, + core_a | core_b, + identity=0, + k=1, + ) + + # Neighbor expansion deliberately permits this match: it is the mechanism + # used to recover repeats that straddle different partition boundaries. + # Consequently, callers that need strictly core-local identity must use + # delta=0 rather than silently expecting the expanded sketches to be + # ignored. + assert result == 1.0 + + +def test_delta_half_recovers_repeat_shifted_across_window_boundary(): + # The {1, 2, 3, 4} repeat fills window 0, but its second occurrence starts + # halfway through window 1 and ends halfway through window 2. Core-only + # comparison sees just half of it in window 2 and falls below the cutoff; + # delta=0.5 expands window 2 far enough to recover the full repeat. + hashes = [1, 2, 3, 4, 90, 91, 1, 2, 3, 4, 92, 93] + + without_neighbors = createSelfMatrix(len(hashes), hashes, 4, 1, 0, 1, 75, True, 4) + with_neighbors = createSelfMatrix(len(hashes), hashes, 4, 1, 0.5, 1, 75, True, 4) + + assert without_neighbors[0, 2] == 0.0 + assert with_neighbors[0, 2] == 100.0 + + +def test_neighbor_containment_applies_cutoff_after_both_directions(): + # A -> expanded B is 1/2, while B -> expanded A is 3/4. The stronger + # direction must be considered before applying the threshold. + core_a = {1, 2} + core_b = {3, 4, 5, 6} + expanded_a = {1, 2, 3, 4, 5} + expanded_b = {1} + + assert ( + containment_neighbors(core_a, core_b, expanded_a, expanded_b, identity=75, k=1) + == 0.75 + ) + assert ( + containment_neighbors(core_a, core_b, expanded_a, expanded_b, identity=76, k=1) + == 0.0 + ) + + +def test_containment_cutoff_is_independent_of_argument_order(): + larger_sketch = {1, 2, 3, 4} + contained_sketch = {1, 2} + + forward = containment_neighbors( + larger_sketch, + contained_sketch, + larger_sketch, + contained_sketch, + identity=75, + k=1, + ) + reverse = containment_neighbors( + contained_sketch, + larger_sketch, + contained_sketch, + larger_sketch, + identity=75, + k=1, + ) + + assert forward == reverse == 1.0 + + +def test_pairwise_containment_matrix_is_rectangular_and_keeps_axis_orientation(): + matrix = pairwiseContainmentMatrix( + mod_set_x=[{1}, {2}, {3}], + mod_set_y=[{1}, {3}], + mod_set_x_neighbors=[{1}, {2}, {3}], + mod_set_y_neighbors=[{1}, {3}], + identity=0, + k=1, + supress_progress=True, + ) + + np.testing.assert_array_equal( + matrix, + np.array( + [ + [100.0, 0.0, 0.0], + [0.0, 0.0, 100.0], + ] + ), + ) + assert matrix.shape == (2, 3) + + +def test_pairwise_containment_matrix_supports_more_rows_than_columns(): + matrix = pairwiseContainmentMatrix( + mod_set_x=[{2}], + mod_set_y=[{1}, {2}, {3}], + mod_set_x_neighbors=[{2}], + mod_set_y_neighbors=[{1}, {2}, {3}], + identity=0, + k=1, + supress_progress=True, + ) + + np.testing.assert_array_equal(matrix, np.array([[0.0], [100.0], [0.0]])) + assert matrix.shape == (3, 1) + + +@pytest.mark.parametrize( + ("mod_set_x", "mod_set_y", "expected_shape"), + [ + ([], [{1}, {2}], (2, 0)), + ([{1}, {2}], [], (0, 2)), + ([], [], (0, 0)), + ], +) +def test_pairwise_containment_matrix_preserves_empty_axis_dimensions( + mod_set_x, mod_set_y, expected_shape +): + matrix = pairwiseContainmentMatrix( + mod_set_x=mod_set_x, + mod_set_y=mod_set_y, + mod_set_x_neighbors=list(mod_set_x), + mod_set_y_neighbors=list(mod_set_y), + identity=0, + k=1, + supress_progress=True, + ) + + assert matrix.shape == expected_shape + + +def test_pairwise_containment_matrix_does_not_hide_misaligned_neighbor_data(): + with pytest.raises(IndexError): + pairwiseContainmentMatrix( + mod_set_x=[{1}, {2}], + mod_set_y=[{1}], + mod_set_x_neighbors=[{1}], + mod_set_y_neighbors=[{1}], + identity=0, + k=1, + supress_progress=True, + ) + + +@pytest.mark.parametrize("sequence", ["", "A", "AC"]) +def test_generate_kmers_shorter_than_k_with_progress_returns_empty(sequence, capsys): + assert list(generateKmersFromFasta(sequence, 3, quiet=False, fw_only=True)) == [] + assert "100.0%" in capsys.readouterr().out + + +def test_generate_one_kmer_with_progress_does_not_use_zero_modulus(capsys): + result = list(generateKmersFromFasta("ACG", 3, quiet=False, fw_only=True)) + + assert result == [np.uint64(0xB13A5310100F646E)] + output = capsys.readouterr().out + assert "100.0%" in output + assert "Completed" in output + + +def test_print_progress_bar_accepts_zero_total(capsys): + printProgressBar(0, 0, prefix="Progress:", suffix="Completed", length=4) + + output = capsys.readouterr().out + assert "|████|" in output + assert "100.0%" in output diff --git a/tests/test_annotation_track.py b/tests/test_annotation_track.py new file mode 100644 index 0000000..a58ad36 --- /dev/null +++ b/tests/test_annotation_track.py @@ -0,0 +1,256 @@ +import xml.etree.ElementTree as ET +from pathlib import Path + +import matplotlib.pyplot as plt +from matplotlib.colors import to_rgba +import pandas as pd +import pytest + +import moddotplot.static_plots as static_plots +from moddotplot.static_plots import ( + DEFAULT_ANNOTATION_COLOR, + draw_annotation_track, + read_annotation_bed, + render_annotation_track, +) + + +@pytest.mark.parametrize( + ("record", "expected_columns"), + [ + ("chr1\t10\t20\n", ["chrom", "start", "end"]), + ( + "chr1\t10\t20\tfeature\t42\t+\t11\t19\t12,34,56\n", + [ + "chrom", + "start", + "end", + "name", + "score", + "strand", + "thickStart", + "thickEnd", + "itemRgb", + ], + ), + ], +) +def test_read_annotation_bed_accepts_bed3_through_bed9( + tmp_path, record, expected_columns +): + bed_path = tmp_path / "annotations.bed" + bed_path.write_text(record) + + dataframe = read_annotation_bed(bed_path) + + assert list(dataframe.columns) == expected_columns + assert dataframe.loc[0, "chrom"] == "chr1" + assert dataframe.loc[0, "start"] == 10 + assert dataframe.loc[0, "end"] == 20 + + +@pytest.mark.parametrize( + "record", + [ + "chr1\tstart\t20\n", + "chr1\t10.5\t20\n", + "chr1\t-1\t20\n", + "chr1\t20\t20\n", + "chr1\t21\t20\n", + "chr1\t1\t2\ta\t0\t+\t1\t2\t0,0,0\textra\n", + ], +) +def test_read_annotation_bed_rejects_malformed_records(tmp_path, record): + bed_path = tmp_path / "invalid.bed" + bed_path.write_text(record) + + with pytest.raises(ValueError, match="Invalid BED file"): + read_annotation_bed(bed_path) + + +def test_draw_annotation_track_filters_clips_and_uses_rgb(tmp_path): + bed_path = tmp_path / "annotations.bed" + bed_path.write_text( + "chr1\t0\t12\tleft\t0\t+\t0\t12\t255,0,0\n" + "chr1\t18\t30\tright\t0\t+\t18\t30\tnot-a-color\n" + "chr1\t20\t25\toutside\t0\t+\t20\t25\t0,255,0\n" + "chr2\t10\t20\twrong-chrom\t0\t+\t10\t20\t0,0,255\n" + ) + dataframe = read_annotation_bed(bed_path) + figure, axis = plt.subplots() + + interval_count = draw_annotation_track(axis, dataframe, "chr1", 10, 20) + + assert interval_count == 2 + assert axis.get_xlim() == pytest.approx((10, 20)) + assert len(axis.patches) == 2 + assert axis.patches[0].get_x() == 10 + assert axis.patches[0].get_width() == 2 + assert axis.patches[0].get_facecolor() == pytest.approx(to_rgba((1, 0, 0))) + assert axis.patches[1].get_x() == 18 + assert axis.patches[1].get_width() == 2 + assert axis.patches[1].get_facecolor() == pytest.approx( + to_rgba(DEFAULT_ANNOTATION_COLOR) + ) + assert {patch.get_y() for patch in axis.patches} == {0.2} + plt.close(figure) + + +def test_render_annotation_track_writes_expected_svg_and_png(tmp_path): + bed_path = tmp_path / "annotations.bed" + bed_path.write_text("chr1\t10\t20\n") + dataframe = read_annotation_bed(bed_path) + output_prefix = tmp_path / "sample_ANNOTATION_TRACK" + + created = render_annotation_track( + dataframe, + "chr1", + 0, + 100, + output_prefix, + width_cm=20, + dpi=72, + ) + + svg_path = tmp_path / "sample_ANNOTATION_TRACK.svg" + png_path = tmp_path / "sample_ANNOTATION_TRACK.png" + assert created is True + assert svg_path.stat().st_size > 0 + assert png_path.read_bytes().startswith(b"\x89PNG\r\n\x1a\n") + ET.parse(svg_path) + + +def test_render_annotation_track_skips_empty_overlap(tmp_path): + bed_path = tmp_path / "annotations.bed" + bed_path.write_text("chr2\t10\t20\n") + dataframe = read_annotation_bed(bed_path) + output_prefix = tmp_path / "sample_ANNOTATION_TRACK" + + created = render_annotation_track( + dataframe, + "chr1", + 0, + 100, + output_prefix, + width_cm=20, + dpi=72, + ) + + assert created is False + assert not (tmp_path / "sample_ANNOTATION_TRACK.svg").exists() + assert not (tmp_path / "sample_ANNOTATION_TRACK.png").exists() + + +class _FakePlot: + def __add__(self, _other): + return self + + +def _stub_create_plots_dependencies(monkeypatch): + dataframe = pd.DataFrame( + [ + { + "q": "chr1", + "q_st": 0, + "q_en": 100, + "r": "chr1", + "r_st": 0, + "r_en": 100, + "perID_by_events": 100.0, + "discrete": 0, + } + ] + ) + monkeypatch.setattr(static_plots, "read_df", lambda *_args, **_kwargs: dataframe) + monkeypatch.setattr( + static_plots, "make_hist", lambda *_args, **_kwargs: _FakePlot() + ) + monkeypatch.setattr(static_plots, "make_dot", lambda *_args, **_kwargs: _FakePlot()) + + def fake_ggsave(*_args, **kwargs): + output = Path(kwargs["filename"]) + if output.suffix == ".png": + output.write_bytes(b"png") + else: + output.write_text('') + + monkeypatch.setattr(static_plots, "ggsave", fake_ggsave) + + +def _run_create_plots(output_dir, annotation, vector_format="svg"): + static_plots.create_plots( + sdf=None, + directory=str(output_dir), + name_x="chr1", + name_y="chr1", + palette="Spectral_11", + palette_orientation="+", + no_hist=True, + width=4, + dpi=72, + is_freq=False, + xlim=100, + custom_colors=None, + custom_breakpoints=None, + from_file=None, + is_pairwise=False, + axes_labels=None, + axes_tick_number=7, + vector_format=vector_format, + deraster=True, + annotation=str(annotation), + ) + + +def test_create_plots_creates_and_skips_annotation_artifacts(tmp_path, monkeypatch): + _stub_create_plots_dependencies(monkeypatch) + + matching_dir = tmp_path / "matching" + matching_dir.mkdir() + matching_bed = tmp_path / "matching.bed" + matching_bed.write_text("chr1\t10\t20\n") + _run_create_plots(matching_dir, matching_bed) + + assert (matching_dir / "chr1_ANNOTATION_TRACK.svg").exists() + assert (matching_dir / "chr1_ANNOTATION_TRACK.png").exists() + assert (matching_dir / "chr1_TRI_ANNOTATED.svg").exists() + assert (matching_dir / "chr1_TRI_ANNOTATED.png").exists() + assert not (matching_dir / "chr1_PRE_ANNOTATED.svg").exists() + + empty_dir = tmp_path / "empty" + empty_dir.mkdir() + nonmatching_bed = tmp_path / "nonmatching.bed" + nonmatching_bed.write_text("chr2\t10\t20\n") + _run_create_plots(empty_dir, nonmatching_bed) + + assert not (empty_dir / "chr1_ANNOTATION_TRACK.svg").exists() + assert not (empty_dir / "chr1_ANNOTATION_TRACK.png").exists() + assert not (empty_dir / "chr1_PRE_ANNOTATED.svg").exists() + assert not (empty_dir / "chr1_TRI_ANNOTATED.svg").exists() + assert not (empty_dir / "chr1_TRI_ANNOTATED.png").exists() + + +@pytest.mark.parametrize( + ("vector_format", "magic"), + [("svg", b"alpha\n" + + "ACGT" * 300 + + "\n>beta\n" + + "ACGT" * 275 + + "\n>gamma\n" + + "ACGT" * 250 + + "\n" + ) + + +def test_static_cli_computes_all_self_and_pairwise_outputs(tmp_path): + fasta = tmp_path / "three.fa" + output = tmp_path / "static" + _write_multifasta(fasta) + + result = _run_cli( + "static", + "--fasta", + fasta, + "--window", + 100, + "--modimizer", + 10, + "--identity", + 80, + "--compare", + "--no-plot", + "--output-dir", + output, + ) + + assert result.returncode == 0, result.stderr + result.stdout + bedpe_files = sorted(path.relative_to(output) for path in output.rglob("*.bedpe")) + assert bedpe_files == [ + Path("alpha/alpha.bedpe"), + Path("alpha_beta/alpha_beta_COMPARE.bedpe"), + Path("alpha_gamma/alpha_gamma_COMPARE.bedpe"), + Path("beta/beta.bedpe"), + Path("beta_gamma/beta_gamma_COMPARE.bedpe"), + Path("gamma/gamma.bedpe"), + ] + assert all(path.stat().st_size > 0 for path in output.rglob("*.bedpe")) + + +def test_static_grid_regions_with_dotted_headers_crop_every_output(tmp_path): + names = [ + "PAN010.chr14.haplotype1.paternal", + "PAN010.chr14.haplotype2.maternal", + "PAN027.chr14.paternal", + ] + fasta = tmp_path / "three-dotted.fa" + fasta.write_text( + "".join(f">{name}\n{'ACGT' * 300}\n" for name in names), + encoding="ascii", + ) + output = tmp_path / "regions" + + result = _run_cli( + "static", + "--grid", + "--fasta", + fasta, + "--region", + *(f"{name}:1-400" for name in names), + "--window", + 50, + "--modimizer", + 10, + "--identity", + 80, + "--no-plot", + "--output-dir", + output, + ) + + assert result.returncode == 0, result.stderr + result.stdout + assert result.stdout.count("Sequence length n: 400") == 3 + assert "Sequence length n: 1200" not in result.stdout + bedpe_files = list(output.rglob("*.bedpe")) + assert len(bedpe_files) == 6 + for bedpe in bedpe_files: + rows = [line.split("\t") for line in bedpe.read_text().splitlines()[1:]] + assert rows + coordinates = [int(row[index]) for row in rows for index in (1, 2, 4, 5)] + assert min(coordinates) >= 1 + assert max(coordinates) <= 400 + + assert (output / "3x3_GRID.png").is_file() + assert (output / "3x3_GRID.svg").is_file() + + +def test_static_cli_reads_gzip_and_renders_bed_annotations(tmp_path): + fasta = tmp_path / "one.fa.gz" + with gzip.open(fasta, "wt", encoding="ascii") as stream: + stream.write(">alpha description\n" + "ACGT" * 300 + "\n") + + annotation = tmp_path / "annotations.bed" + annotation.write_text("alpha\t100\t300\tfeature\t0\t+\t100\t300\t12,34,56\n") + output = tmp_path / "annotated" + + result = _run_cli( + "static", + "--fasta", + fasta, + "--bed", + annotation, + "--window", + 100, + "--modimizer", + 10, + "--identity", + 80, + "--no-hist", + "--output-dir", + output, + ) + + assert result.returncode == 0, result.stderr + result.stdout + sequence_output = output / "alpha" + expected = [ + sequence_output / "alpha_ANNOTATION_TRACK.svg", + sequence_output / "alpha_ANNOTATION_TRACK.png", + sequence_output / "alpha_TRI_ANNOTATED.svg", + sequence_output / "alpha_TRI_ANNOTATED.png", + ] + assert all(path.stat().st_size > 0 for path in expected) + ET.parse(sequence_output / "alpha_ANNOTATION_TRACK.svg") + ET.parse(sequence_output / "alpha_TRI_ANNOTATED.svg") + assert not list(sequence_output.glob("*.ini")) + assert not (tmp_path / "one.fa.gz.fai").exists() + + +def test_interactive_cli_forward_mode_saves_matrix_without_launching_server(tmp_path): + fasta = tmp_path / "one.fa" + fasta.write_text(">alpha\n" + "ACGT" * 300 + "\n") + output = tmp_path / "interactive" + + result = _run_cli( + "interactive", + "--fasta", + fasta, + "--window", + 100, + "--resolution", + 10, + "--modimizer", + 10, + "--quick", + "--forward", + "--save", + "--no-plot", + "--output-dir", + output, + ) + + assert result.returncode == 0, result.stderr + result.stdout + saved = output / "interactive_matrices" + assert (saved / "alpha_0.npz").is_file() + assert (saved / "metadata.pkl").is_file() + assert "Saved matrices" in result.stdout diff --git a/tests/test_cli_runtime.py b/tests/test_cli_runtime.py new file mode 100644 index 0000000..b70f787 --- /dev/null +++ b/tests/test_cli_runtime.py @@ -0,0 +1,234 @@ +import sys + +import numpy as np +import pytest + +import moddotplot.moddotplot as cli + + +def _patch_fasta_input(monkeypatch, names, kmers): + monkeypatch.setattr(cli, "isValidFasta", lambda _path: True) + monkeypatch.setattr(cli, "getInputHeaders", lambda _path: names) + monkeypatch.setattr(cli, "readKmersFromFile", lambda *_args: kmers) + + +def _patch_static_calculation(monkeypatch, plot_calls, pair_calls=None): + monkeypatch.setattr(cli, "createSelfMatrix", lambda *_args: np.full((1, 1), 100.0)) + + def create_pairwise(*args): + if pair_calls is not None: + pair_calls.append(args) + return np.full((1, 1), 95.0) + + monkeypatch.setattr(cli, "createPairwiseMatrix", create_pairwise) + monkeypatch.setattr( + cli, + "convertMatrixToBed", + lambda *_args, **_kwargs: [["header"], ["value"]], + ) + monkeypatch.setattr(cli, "create_plots", lambda **kwargs: plot_calls.append(kwargs)) + + +def test_main_without_subcommand_reports_parser_error(monkeypatch, capsys): + monkeypatch.setattr(sys, "argv", ["moddotplot"]) + + with pytest.raises(SystemExit) as exc_info: + cli.main() + + captured = capsys.readouterr() + assert exc_info.value.code == 2 + assert "the following arguments are required: command" in captured.err + assert "{interactive,static}" in captured.err + + +@pytest.mark.parametrize("option", ["--colors", "--color"]) +def test_static_parser_accepts_color_option_aliases(option): + args = cli.get_parser().parse_args( + ["static", "--fasta", "sequence.fa", option, "#010203", "#abcdef"] + ) + + assert args.colors == ["#010203", "#abcdef"] + + +def test_static_config_prefers_canonical_colors_key(): + args = cli.get_parser().parse_args(["static", "--fasta", "sequence.fa"]) + + cli._apply_static_config( + args, + { + "fasta": ["sequence.fa"], + "colors": ["#canonical"], + "color": ["#legacy"], + }, + ) + + assert args.colors == ["#canonical"] + + +def test_static_config_supports_legacy_color_key(): + args = cli.get_parser().parse_args(["static", "--fasta", "sequence.fa"]) + + cli._apply_static_config(args, {"fasta": ["sequence.fa"], "color": ["#legacy"]}) + + assert args.colors == ["#legacy"] + + +def test_static_delta_defaults_to_half_window(): + args = cli.get_parser().parse_args(["static", "--fasta", "sequence.fa"]) + + assert args.delta == 0.5 + + +def test_region_kmer_slice_uses_one_based_inclusive_base_coordinates(): + kmers = np.arange(980) # 980 21-mers represent a 1,000-base sequence. + + selected = cli._slice_kmers_for_region(kmers, ("chrA", 101, 400), 21) + + assert len(selected) == 280 + np.testing.assert_array_equal(selected, np.arange(100, 380)) + + +@pytest.mark.parametrize( + "region", + [("chrA", 0, 100), ("chrA", 900, 1_001), ("chrA", 10, 20)], +) +def test_region_kmer_slice_rejects_invalid_bounds_or_short_intervals(region): + with pytest.raises(ValueError): + cli._slice_kmers_for_region(np.arange(980), region, 21) + + +def test_interactive_window_uses_longest_sequence_by_length(monkeypatch): + # The shorter k-mer list is lexicographically greater. This reproduces the + # old len(max(k_list)) bug while keeping the test computation tiny. + short_kmers = [9] * 100 + long_kmers = [1] * 1000 + _patch_fasta_input(monkeypatch, ["short", "long"], [short_kmers, long_kmers]) + monkeypatch.setattr(cli, "partitionOverlaps", lambda *_args: []) + monkeypatch.setattr(cli, "convertToModimizers", lambda *_args: []) + monkeypatch.setattr(cli, "selfContainmentMatrix", lambda *_args: np.zeros((1, 1))) + dash_calls = [] + monkeypatch.setattr(cli, "run_dash", lambda *args: dash_calls.append(args)) + monkeypatch.setattr( + sys, + "argv", + [ + "moddotplot", + "interactive", + "--fasta", + "sequence.fa", + "--resolution", + "10", + "--quick", + ], + ) + + cli.main() + + metadata = dash_calls[0][1] + assert {entry["min_window_size"] for entry in metadata} == {102} + assert {entry["max_window_size"] for entry in metadata} == {102} + + +def test_no_bedpe_self_plot_uses_sequence_output_directory(monkeypatch, tmp_path): + _patch_fasta_input(monkeypatch, ["chrA"], [[1] * 1000]) + plot_calls = [] + _patch_static_calculation(monkeypatch, plot_calls) + monkeypatch.setattr( + sys, + "argv", + [ + "moddotplot", + "static", + "--fasta", + "sequence.fa", + "--resolution", + "10", + "--no-bedpe", + "--output-dir", + str(tmp_path), + ], + ) + + cli.main() + + expected_directory = tmp_path / "chrA" + assert expected_directory.is_dir() + assert plot_calls[0]["directory"] == str(expected_directory) + assert not list(tmp_path.rglob("*.bedpe")) + + +def test_unmatched_region_fails_instead_of_silently_using_full_sequences( + monkeypatch, tmp_path, capsys +): + larger = [1] * 1000 + smaller = [2] * 800 + _patch_fasta_input(monkeypatch, ["chrA", "chrB"], [larger, smaller]) + plot_calls = [] + pair_calls = [] + _patch_static_calculation(monkeypatch, plot_calls, pair_calls) + monkeypatch.setattr( + sys, + "argv", + [ + "moddotplot", + "static", + "--fasta", + "sequences.fa", + "--compare-only", + "--region", + "missing:1-100", + "--resolution", + "10", + "--no-bedpe", + "--output-dir", + str(tmp_path), + ], + ) + + with pytest.raises(SystemExit) as error: + cli.main() + + captured = capsys.readouterr() + assert error.value.code == 2 + assert "does not match any FASTA identifier" in captured.out + assert not pair_calls + assert not plot_calls + assert not list(tmp_path.rglob("*.bedpe")) + + +def test_compare_only_region_can_target_just_one_sequence(monkeypatch, tmp_path): + larger = [1] * 1000 + smaller = [2] * 800 + _patch_fasta_input(monkeypatch, ["chrA", "chrB"], [larger, smaller]) + plot_calls = [] + pair_calls = [] + _patch_static_calculation(monkeypatch, plot_calls, pair_calls) + monkeypatch.setattr( + sys, + "argv", + [ + "moddotplot", + "static", + "--fasta", + "sequences.fa", + "--compare-only", + "--region", + "chrB:101-400", + "--resolution", + "10", + "--no-bedpe", + "--output-dir", + str(tmp_path), + ], + ) + + cli.main() + + pair_args = pair_calls[0] + assert pair_args[0] == 280 + assert len(pair_args[2]) == 280 + assert pair_args[3] is larger + assert plot_calls[0]["name_x"] == "chrA" + assert plot_calls[0]["name_y"] == "chrB:101-400" + assert plot_calls[0]["xlim"] == (1, 1020) + assert not list(tmp_path.rglob("*.bedpe")) diff --git a/tests/test_direction_cli.py b/tests/test_direction_cli.py new file mode 100644 index 0000000..451eaea --- /dev/null +++ b/tests/test_direction_cli.py @@ -0,0 +1,119 @@ +import sys + +import numpy as np + +import moddotplot.moddotplot as cli + + +def _patch_common_static_io(monkeypatch, names, hash_sets, tmp_path): + monkeypatch.setattr(cli, "isValidFasta", lambda _path: True) + monkeypatch.setattr(cli, "getInputHeaders", lambda _path: names) + + def read_hashes( + _path, + _kmer, + _quiet, + forward_only, + ambiguous=False, + regions=None, + record_ids=None, + ): + assert ambiguous is False + assert not regions + assert record_ids == names + return hash_sets[forward_only] + + monkeypatch.setattr(cli, "readKmersFromFile", read_hashes) + monkeypatch.setattr( + cli, + "convertMatrixToBed", + lambda *_args, **_kwargs: [["header"], ["value"]], + ) + monkeypatch.setattr(cli, "create_plots", lambda **_kwargs: None) + monkeypatch.setattr(cli, "create_grid", lambda **_kwargs: None) + monkeypatch.setattr(cli, "read_df_from_file", lambda _path: None) + monkeypatch.setattr( + sys, + "argv", + [ + "moddotplot", + "static", + "--fasta", + "sequence.fa", + "--resolution", + "10", + "--no-bedpe", + "--plot-direction", + "--output-dir", + str(tmp_path), + ], + ) + + +def test_static_direction_plot_receives_canonical_and_forward_self_matrices( + monkeypatch, tmp_path +): + canonical_hashes = [[1] * 1000] + forward_hashes = [[2] * 1000] + _patch_common_static_io( + monkeypatch, + ["chrA"], + {False: canonical_hashes, True: forward_hashes}, + tmp_path, + ) + canonical_matrix = np.array([[100.0, 90.0], [90.0, 100.0]]) + forward_matrix = np.array([[100.0, 0.0], [0.0, 100.0]]) + + def create_self(_length, sequence, *_args): + return canonical_matrix if sequence is canonical_hashes[0] else forward_matrix + + monkeypatch.setattr(cli, "createSelfMatrix", create_self) + direction_calls = [] + monkeypatch.setattr( + cli, + "create_direction_plot", + lambda **kwargs: direction_calls.append(kwargs), + ) + + cli.main() + + assert len(direction_calls) == 1 + call = direction_calls[0] + assert call["canonical_matrix"] is canonical_matrix + assert call["forward_matrix"] is forward_matrix + assert call["self_identity"] is True + assert call["name_x"] == call["name_y"] == "chrA" + + +def test_static_direction_plot_receives_pairwise_matrices(monkeypatch, tmp_path): + canonical_hashes = [[1] * 1000, [2] * 800] + forward_hashes = [[3] * 1000, [4] * 800] + _patch_common_static_io( + monkeypatch, + ["chrA", "chrB"], + {False: canonical_hashes, True: forward_hashes}, + tmp_path, + ) + sys.argv.extend(["--compare-only"]) + canonical_matrix = np.array([[90.0]]) + forward_matrix = np.array([[0.0]]) + + def create_pair(_y_length, _x_length, y_sequence, _x_sequence, *_args): + return canonical_matrix if y_sequence is canonical_hashes[1] else forward_matrix + + monkeypatch.setattr(cli, "createPairwiseMatrix", create_pair) + direction_calls = [] + monkeypatch.setattr( + cli, + "create_direction_plot", + lambda **kwargs: direction_calls.append(kwargs), + ) + + cli.main() + + assert len(direction_calls) == 1 + call = direction_calls[0] + assert call["canonical_matrix"] is canonical_matrix + assert call["forward_matrix"] is forward_matrix + assert call["self_identity"] is False + assert (call["name_x"], call["name_y"]) == ("chrA", "chrB") diff --git a/tests/test_direction_plot.py b/tests/test_direction_plot.py new file mode 100644 index 0000000..60f3d3c --- /dev/null +++ b/tests/test_direction_plot.py @@ -0,0 +1,61 @@ +import numpy as np +import pytest + +from moddotplot.static_plots import create_direction_plot, direction_dataframe + + +def test_direction_dataframe_classifies_canonical_only_hits_as_reverse(): + canonical = np.array([[100.0, 95.0], [95.0, 100.0]]) + forward = np.array([[100.0, 0.0], [0.0, 100.0]]) + + result = direction_dataframe( + canonical, + forward, + window_size=10, + name_x="chr1", + name_y="chr1", + self_identity=True, + x_offset=100, + y_offset=100, + ) + + assert result["direction"].tolist() == ["Forward", "Reverse", "Forward"] + assert result[["q_st", "r_st"]].values.tolist() == [ + [100, 100], + [100, 110], + [110, 110], + ] + + +def test_direction_dataframe_requires_equal_shapes(): + with pytest.raises(ValueError, match="matching shapes"): + direction_dataframe( + np.ones((2, 2)), + np.ones((2, 3)), + 10, + "x", + "y", + False, + ) + + +def test_create_direction_plot_writes_raster_and_vector_outputs(tmp_path): + canonical = np.array([[100.0, 95.0], [95.0, 100.0]]) + forward = np.array([[100.0, 0.0], [0.0, 100.0]]) + + plot = create_direction_plot( + canonical, + forward, + window_size=10, + directory=tmp_path, + name_x="chr1", + name_y="chr1", + self_identity=True, + width=1, + dpi=72, + vector_format="svg", + ) + + assert plot is not None + assert (tmp_path / "chr1_DIRECTION.png").stat().st_size > 0 + assert (tmp_path / "chr1_DIRECTION.svg").stat().st_size > 0 diff --git a/tests/test_entrypoints.py b/tests/test_entrypoints.py new file mode 100644 index 0000000..b764478 --- /dev/null +++ b/tests/test_entrypoints.py @@ -0,0 +1,54 @@ +import os +from pathlib import Path +import subprocess +import sys + + +PROJECT_ROOT = Path(__file__).resolve().parents[1] + + +def _run_module(*arguments): + environment = os.environ.copy() + environment["PYTHONPATH"] = str(PROJECT_ROOT / "src") + return subprocess.run( + [sys.executable, "-m", "moddotplot", *arguments], + cwd=PROJECT_ROOT, + env=environment, + capture_output=True, + text=True, + check=False, + ) + + +def test_help_does_not_import_static_rendering_stack(): + result = subprocess.run( + [ + sys.executable, + "-c", + ( + "import sys; import moddotplot.moddotplot; " + "assert 'moddotplot.static_plots' not in sys.modules" + ), + ], + cwd=PROJECT_ROOT, + env={**os.environ, "PYTHONPATH": str(PROJECT_ROOT / "src")}, + capture_output=True, + text=True, + check=False, + ) + + assert result.returncode == 0, result.stderr + + +def test_module_help_is_available(): + result = _run_module("--help") + + assert result.returncode == 0 + assert "{interactive,static}" in result.stdout + + +def test_module_without_subcommand_has_clean_usage_error(): + result = _run_module() + + assert result.returncode == 2 + assert "the following arguments are required: command" in result.stderr diff --git a/tests/test_fasta_parser.py b/tests/test_fasta_parser.py new file mode 100644 index 0000000..06bb5e9 --- /dev/null +++ b/tests/test_fasta_parser.py @@ -0,0 +1,256 @@ +import gzip +import struct +import zlib + +import numpy as np +import pytest +import moddotplot.parse_fasta as fasta_parser + +from moddotplot.parse_fasta import ( + _iter_fasta_records, + extractRegion, + generateKmersFromFasta, + getInputHeaders, + getInputSeqLength, + isValidFasta, + readKmersFromFile, +) + + +@pytest.mark.parametrize( + ("value", "expected"), + [ + ("chrA:1-400", ("chrA", 1, 400)), + ( + "PAN010.chr14.haplotype1.paternal:1-4000000", + ("PAN010.chr14.haplotype1.paternal", 1, 4_000_000), + ), + ("sample-name:10-20", ("sample-name", 10, 20)), + ("sample:1-400:101-300", ("sample", 101, 300)), + ], +) +def test_extract_region_accepts_realistic_fasta_identifiers(value, expected): + assert extractRegion(value) == expected + + +@pytest.mark.parametrize("value", ["sample", "sample:1", "sample:one-two"]) +def test_extract_region_rejects_missing_or_malformed_coordinates(value): + assert extractRegion(value) is None + + +def _bgzf_block(data): + compressor = zlib.compressobj(level=6, wbits=-15) + compressed = compressor.compress(data) + compressor.flush() + block_size = 18 + len(compressed) + 8 + if block_size > 65_536: + raise ValueError("test BGZF block is too large") + + header = struct.pack("alpha descriptive header\r\n" + b"ac\r\n" + b"\r\n" + b"gT\r\n" + b">beta another description\r\n" + b"NN\r\n" + b">empty" + ) + + assert list(_iter_fasta_records(fasta)) == [ + ("alpha", "acgT"), + ("beta", "NN"), + ("empty", ""), + ] + assert getInputHeaders(fasta) == ["alpha", "beta", "empty"] + assert getInputSeqLength(fasta) == [4, 2, 0] + assert isValidFasta(fasta) is True + assert not (tmp_path / "records.fa.fai").exists() + + +def test_gzip_is_detected_by_content_instead_of_extension(tmp_path): + fasta = tmp_path / "compressed.data" + with gzip.open(fasta, "wt", encoding="ascii") as output: + output.write(">alpha description\nAC\nGT\n>beta\ntt") + + assert list(_iter_fasta_records(fasta)) == [ + ("alpha", "ACGT"), + ("beta", "tt"), + ] + assert getInputHeaders(fasta) == ["alpha", "beta"] + assert getInputSeqLength(fasta) == [4, 2] + + +def test_bgzf_concatenated_blocks_are_read_as_one_fasta_stream(tmp_path): + fasta = tmp_path / "records.fa.bgz" + fasta.write_bytes( + _bgzf_block(b">alpha description\nAC") + + _bgzf_block(b"GT\n>beta\ntt\n") + + _bgzf_block(b"") + ) + + assert list(_iter_fasta_records(fasta)) == [ + ("alpha", "ACGT"), + ("beta", "tt"), + ] + + +def test_read_kmers_preserves_record_order_and_public_return_shape(tmp_path): + fasta = tmp_path / "records.fa" + fasta.write_text(">alpha description\nACGT\n>beta\nTTAA\n") + + result = readKmersFromFile( + str(fasta), ksize=3, quiet=True, fw_only=True, ambiguous=False + ) + + assert isinstance(result, list) + assert len(result) == 2 + assert all(isinstance(record, np.ma.MaskedArray) for record in result) + assert result[0].tolist() == list( + generateKmersFromFasta("ACGT", 3, quiet=True, fw_only=True) + ) + assert result[1].tolist() == list( + generateKmersFromFasta("TTAA", 3, quiet=True, fw_only=True) + ) + + +def test_read_kmers_hashes_only_the_requested_region(tmp_path): + sequence = "ACGT" * 250 + fasta = tmp_path / "region.fa" + fasta.write_text(f">sample.with.dots\n{sequence}\n") + + result = readKmersFromFile( + str(fasta), + ksize=21, + quiet=True, + fw_only=True, + ambiguous=False, + regions={"sample.with.dots": ("sample.with.dots", 101, 400)}, + ) + + assert len(result[0]) == 280 + assert result[0].tolist() == list( + generateKmersFromFasta(sequence[100:400], 21, quiet=True, fw_only=True) + ) + + +def test_indexed_region_uses_random_access_instead_of_record_parser( + tmp_path, monkeypatch +): + sequence = "ACGT" * 250 + fasta = tmp_path / "indexed.fa" + header = b">sample.with.dots\n" + fasta.write_bytes(header + sequence.encode("ascii") + b"\n") + (tmp_path / "indexed.fa.fai").write_text( + f"sample.with.dots\t{len(sequence)}\t{len(header)}\t{len(sequence)}\t{len(sequence) + 1}\n" + ) + + def fail_streaming(*_args, **_kwargs): + raise AssertionError("indexed FASTA should not use the streaming fallback") + + monkeypatch.setattr(fasta_parser, "_iter_selected_fasta_records", fail_streaming) + result = readKmersFromFile( + str(fasta), + ksize=21, + quiet=True, + fw_only=True, + regions={"sample.with.dots": ("sample.with.dots", 101, 400)}, + record_ids=["sample.with.dots"], + ) + + assert len(result[0]) == 280 + assert result[0].tolist() == list( + generateKmersFromFasta(sequence[100:400], 21, quiet=True, fw_only=True) + ) + + +def test_single_record_stream_stops_after_requested_region(tmp_path): + fasta = tmp_path / "streamed.fa" + fasta.write_text(">sample\nACGT\nACGT\nBRO KEN\n") + + result = readKmersFromFile( + str(fasta), + ksize=3, + quiet=True, + fw_only=True, + regions={"sample": ("sample", 1, 8)}, + record_ids=["sample"], + ) + + assert result[0].tolist() == list( + generateKmersFromFasta("ACGTACGT", 3, quiet=True, fw_only=True) + ) + + +def test_header_reader_does_not_assemble_sequence_records(tmp_path, monkeypatch): + fasta = tmp_path / "headers.fa" + fasta.write_text(">alpha\nACGT\n>beta\nTTAA\n") + + monkeypatch.setattr( + fasta_parser, + "_iter_fasta_records", + lambda *_args: (_ for _ in ()).throw(AssertionError("full parser used")), + ) + + assert getInputHeaders(fasta) == ["alpha", "beta"] + + +@pytest.mark.parametrize( + ("contents", "message"), + [ + (b"ACGT\n", "sequence data before the first header"), + (b"> \nACGT\n", "empty header"), + (b">alpha\nAC GT\n", "whitespace within sequence data"), + (b">alpha\nAC\n>alpha description\nGT\n", "duplicate sequence identifier"), + (b"\n\n", "no FASTA records found"), + ], +) +def test_malformed_fasta_has_a_clear_error(tmp_path, contents, message): + fasta = tmp_path / "malformed.fa" + fasta.write_bytes(contents) + + with pytest.raises(ValueError, match=message): + list(_iter_fasta_records(fasta)) + + +def test_public_header_reader_rejects_duplicate_first_token_ids(tmp_path): + fasta = tmp_path / "duplicate.fa" + fasta.write_text(">alpha first\nAC\n>alpha second\nGT\n") + + with pytest.raises(ValueError, match="duplicate sequence identifier 'alpha'"): + getInputHeaders(fasta) + + +def test_validation_returns_false_and_reports_malformed_fasta(tmp_path, capsys): + fasta = tmp_path / "malformed.fa" + fasta.write_text(">\nACGT\n") + + assert isValidFasta(fasta) is False + assert "empty header at line 1" in capsys.readouterr().out + + +def test_validation_preserves_missing_file_exit_code(tmp_path, capsys): + missing = tmp_path / "missing.fa" + + with pytest.raises(SystemExit) as error: + isValidFasta(missing) + + assert error.value.code == 5 + assert "Unable to find fasta" in capsys.readouterr().out + + +def test_validation_preserves_unreadable_compression_exit_code(tmp_path, capsys): + fasta = tmp_path / "corrupt.fa.gz" + fasta.write_bytes(b"\x1f\x8bnot-a-valid-gzip-stream") + + with pytest.raises(SystemExit) as error: + isValidFasta(fasta) + + assert error.value.code == 6 + assert "An error occurred" in capsys.readouterr().out diff --git a/tests/test_grid.py b/tests/test_grid.py new file mode 100644 index 0000000..95a1bfa --- /dev/null +++ b/tests/test_grid.py @@ -0,0 +1,480 @@ +import matplotlib.pyplot as plt +from matplotlib.colors import to_rgba +import numpy as np +import pytest + +from moddotplot.static_plots import _build_grid_figure, create_grid + + +BED_HEADER = ( + "#query_name", + "query_start", + "query_end", + "reference_name", + "reference_start", + "reference_end", + "perID_by_events", +) + + +def _bed(matrix, query_name, reference_name, *, self_identity): + """Build the in-memory BEDPE representation accepted by ``create_grid``.""" + rows = [BED_HEADER] + for query_index, matrix_row in enumerate(matrix): + for reference_index, identity in enumerate(matrix_row): + if self_identity and query_index > reference_index: + continue + if identity < 86: + continue + query_start = query_index * 10 + reference_start = reference_index * 10 + rows.append( + ( + query_name, + query_start, + query_start + 9, + reference_name, + reference_start, + reference_start + 9, + float(identity), + ) + ) + return rows + + +def _records(query_name, reference_name, cells): + """Build BEDPE rows from ``(q_start, r_start, identity)`` cells.""" + return [BED_HEADER] + [ + ( + query_name, + query_start, + query_start + 10, + reference_name, + reference_start, + reference_start + 10, + float(identity), + ) + for query_start, reference_start, identity in cells + ] + + +def _grid_kwargs( + *, + singles, + doubles, + single_names, + double_names, + custom_colors=None, + custom_breakpoints=None, + deraster=False, +): + return { + "singles": singles, + "doubles": doubles, + "palette": "Spectral_11", + "palette_orientation": "+", + "single_names": single_names, + "double_names": double_names, + "is_freq": False, + "xlim": 100, + "custom_colors": custom_colors, + "custom_breakpoints": custom_breakpoints, + "axes_label": [0, 50, 100], + "is_bed": False, + "width": 1, + "breaks": [0, 50, 100], + "deraster": deraster, + } + + +def _basic_two_sequence_grid(**overrides): + names = ["sequence_a", "sequence_b"] + kwargs = _grid_kwargs( + singles=[ + _records(names[0], names[0], [(10, 20, 91)]), + _records(names[1], names[1], [(30, 40, 92)]), + ], + doubles=[_records(names[0], names[1], [(20, 70, 95)])], + single_names=names, + double_names=[[names[0], names[1]]], + ) + kwargs.update(overrides) + return kwargs + + +def _create_grid(tmp_path, *, vector_format="svg", **kwargs): + create_grid( + directory=tmp_path, + vector_format=vector_format, + dpi=72, + **kwargs, + ) + + +def _collection_center(axis): + vertices = [ + path.vertices + for collection in axis.collections + for path in collection.get_paths() + if path.vertices.size + ] + assert vertices, "expected the grid cell to contain a plotted collection" + points = np.concatenate(vertices) + return ( + (points[:, 0].min() + points[:, 0].max()) / 2, + (points[:, 1].min() + points[:, 1].max()) / 2, + ) + + +def _all_artist_colors(axes): + colors = set() + for axis in axes.flat: + for collection in axis.collections: + colors.update(tuple(color) for color in collection.get_facecolors()) + colors.update(to_rgba(patch.get_facecolor()) for patch in axis.patches) + return colors + + +def test_create_grid_handles_empty_pairwise_comparison(tmp_path): + names = ["sequence_a", "sequence_b", "sequence_c"] + self_matrix = [[100, 92], [92, 100]] + pair_matrix = [[91, 94], [96, 99]] + empty_pair_matrix = [[0, 0], [0, 0]] + + singles = [_bed(self_matrix, name, name, self_identity=True) for name in names] + doubles = [ + _bed(pair_matrix, "sequence_a", "sequence_b", self_identity=False), + _bed(pair_matrix, "sequence_a", "sequence_c", self_identity=False), + _bed( + empty_pair_matrix, + "sequence_b", + "sequence_c", + self_identity=False, + ), + ] + + # An all-zero comparison produces a BED table containing only its header. + assert len(doubles[-1]) == 1 + + _create_grid( + tmp_path, + **_grid_kwargs( + singles=singles, + doubles=doubles, + single_names=names, + double_names=[ + ["sequence_a", "sequence_b"], + ["sequence_a", "sequence_c"], + ["sequence_b", "sequence_c"], + ], + ), + ) + + for suffix in ("png", "svg"): + output = tmp_path / f"3x3_GRID.{suffix}" + assert output.is_file() + assert output.stat().st_size > 0 + + +@pytest.mark.parametrize( + ("vector_format", "signature"), + [("svg", b" 0 + + +def test_benchmark_cli_clearly_reports_skipped_legacy_comparison(monkeypatch, capsys): + benchmark = _load_benchmark_module() + monkeypatch.setattr(benchmark, "_load_mmh3", lambda: None) + + assert benchmark.main(["--length", "100", "--kmer", "5", "--repeats", "1"]) == 0 + + output = capsys.readouterr().out + assert "Legacy comparison: skipped (optional mmh3 is not installed)" in output + assert "forward" in output + assert "canonical" in output + assert "nthash2" in output diff --git a/tests/test_interactive_export.py b/tests/test_interactive_export.py new file mode 100644 index 0000000..140e03c --- /dev/null +++ b/tests/test_interactive_export.py @@ -0,0 +1,35 @@ +import pytest + +from moddotplot.interactive import figure_to_bed + + +def _figure(z, *, x=(100, 110), y=(200, 210), x_name="chr1", y_name="chr2"): + return { + "data": [{"z": z, "x": list(x), "y": list(y)}], + "layout": { + "xaxis": {"title": {"text": x_name}}, + "yaxis": {"title": {"text": y_name}}, + }, + } + + +def test_figure_to_bed_includes_coordinate_offsets(): + rows, filename = figure_to_bed(_figure([[98.0, 0.0], [91.0, 99.0]])) + + assert filename == "chr1-chr2.bedpe" + assert rows[1][:6] == ("chr1", 100, 109, "chr2", 200, 209) + assert rows[-1][:6] == ("chr1", 110, 119, "chr2", 210, 219) + + +def test_figure_to_bed_names_self_identity_export(): + rows, filename = figure_to_bed( + _figure([[100.0, 95.0], [95.0, 100.0]], x_name="chr1", y_name="chr1") + ) + + assert filename == "chr1.bedpe" + assert len(rows) == 4 # header plus the upper triangle + + +def test_figure_to_bed_rejects_missing_window_size(): + with pytest.raises(ValueError, match="positive window size"): + figure_to_bed(_figure([[100.0]], x=(), y=())) diff --git a/tests/test_interactive_parser.py b/tests/test_interactive_parser.py new file mode 100644 index 0000000..264c69b --- /dev/null +++ b/tests/test_interactive_parser.py @@ -0,0 +1,21 @@ +from moddotplot.moddotplot import get_parser + + +def test_interactive_forward_defaults_to_false(): + args = get_parser().parse_args(["interactive", "--fasta", "sequence.fa"]) + + assert args.forward is False + + +def test_interactive_delta_defaults_to_half_window(): + args = get_parser().parse_args(["interactive", "--fasta", "sequence.fa"]) + + assert args.delta == 0.5 + + +def test_interactive_accepts_forward_flag(): + args = get_parser().parse_args( + ["interactive", "--fasta", "sequence.fa", "--forward"] + ) + + assert args.forward is True diff --git a/tests/test_issue53_memory_safe_plotting.py b/tests/test_issue53_memory_safe_plotting.py new file mode 100644 index 0000000..b12b8d5 --- /dev/null +++ b/tests/test_issue53_memory_safe_plotting.py @@ -0,0 +1,105 @@ +import matplotlib.pyplot as plt +import pandas as pd +from plotnine.geoms.geom_tile import geom_tile + +from moddotplot.static_plots import make_dot, make_dot_final, make_dot_grid, make_tri + + +GENOME_SIZE = 496_000_000 +WINDOW_SIZE = 2_000 + + +def _large_sparse_dotplot_data(): + # The adjacent first two coordinates establish a 2 kb resolution while the + # final coordinate establishes the ~496 Mb extent from issue #53. A + # geom_raster layer would try to allocate about 248,000**2 RGBA pixels. + starts = [0, WINDOW_SIZE, GENOME_SIZE - WINDOW_SIZE] + return pd.DataFrame( + { + "q": ["chr1"] * len(starts), + "q_st": starts, + "q_en": [start + WINDOW_SIZE - 1 for start in starts], + "r": ["chr1"] * len(starts), + "r_st": starts, + "r_en": [start + WINDOW_SIZE - 1 for start in starts], + "discrete": pd.Categorical([0, 1, 2]), + } + ) + + +def _make_full_plot(data, deraster=False): + return make_dot( + sdf=data, + name_x="chr1", + name_y="chr1", + palette="Spectral_11", + palette_orientation="+", + colors=None, + breaks=[0, GENOME_SIZE // 2, GENOME_SIZE], + num_ticks=3, + xlim=GENOME_SIZE, + deraster=deraster, + width=2, + is_pairwise=False, + ) + + +def _assert_tile_layer(plot, rasterized): + layer = plot.layers[0] + assert isinstance(layer.geom, geom_tile) + assert layer.geom._kwargs["raster"] is rasterized + + +def test_large_sparse_plot_does_not_build_coordinate_sized_raster(): + plot = _make_full_plot(_large_sparse_dotplot_data()) + + _assert_tile_layer(plot, rasterized=True) + figure = plot.draw(show=False) + try: + axis = figure.axes[0] + assert not axis.images + assert len(axis.collections) == 1 + assert axis.collections[0].get_rasterized() is True + assert len(axis.collections[0].get_paths()) == 3 + finally: + plt.close(figure) + + +def test_deraster_keeps_memory_safe_tiles_as_vectors(): + plot = _make_full_plot(_large_sparse_dotplot_data(), deraster=True) + + _assert_tile_layer(plot, rasterized=False) + figure = plot.draw(show=False) + try: + assert figure.axes[0].collections[0].get_rasterized() is False + finally: + plt.close(figure) + + +def test_grid_and_triangle_paths_use_the_same_memory_safe_geometry(): + data = _large_sparse_dotplot_data() + common = { + "sdf": data, + "palette": "Spectral_11", + "palette_orientation": "+", + "colors": None, + "breaks": [0, GENOME_SIZE // 2, GENOME_SIZE], + "xlim": GENOME_SIZE, + "deraster": False, + "width": 2, + } + + grid_cell = make_dot_final(**common) + grid_plot = make_dot_grid( + **common, + title_name="grid", + on_diagonal=True, + ) + triangle, _ = make_tri( + **common, + title_name="triangle", + num_ticks=3, + ) + + for plot in (grid_cell, grid_plot, triangle): + _assert_tile_layer(plot, rasterized=True) diff --git a/tests/test_modimizer_vectorization.py b/tests/test_modimizer_vectorization.py new file mode 100644 index 0000000..45370d7 --- /dev/null +++ b/tests/test_modimizer_vectorization.py @@ -0,0 +1,129 @@ +import numpy as np +import pytest + +import moddotplot.estimate_identity as estimate_identity +from moddotplot.estimate_identity import convertToModimizers, populateModimizers + + +def _scalar_reference(partition, sparsity, expectation): + """Straightforward reference for the adaptive integer-sparsity algorithm.""" + + current_sparsity = int(sparsity) + while True: + modimizers = { + int(kmer) + for kmer in partition + if kmer is not None + and not np.ma.is_masked(kmer) + and int(kmer) % current_sparsity == 0 + } + if len(modimizers) >= round(expectation / 2) or current_sparsity == 1: + return modimizers + current_sparsity = max(1, current_sparsity // 2) + + +@pytest.mark.parametrize("sparsity", [1, 2, 8, 64]) +@pytest.mark.parametrize("expectation", [0, 1, 12, 200]) +def test_vectorized_modimizers_match_scalar_reference(sparsity, expectation): + rng = np.random.default_rng(441) + values = rng.integers(0, 2**63, size=4096, dtype=np.uint64) + values[100:160] = values[:60] + mask = rng.random(values.size) < 0.13 + partition = np.ma.MaskedArray(values, mask=mask) + + assert populateModimizers( + partition, + sparsity=sparsity, + ambiguous=False, + expectation=expectation, + k=21, + ) == _scalar_reference(partition, sparsity, expectation) + + +def test_legacy_sequence_skips_none_and_masked_values(): + partition = [1, None, np.ma.masked, 2, np.uint64(4), 4, 7] + + assert populateModimizers(partition, 4, False, 4, 21) == {2, 4} + + +def test_adaptive_sparsity_remains_integer(monkeypatch): + observed_sparsities = [] + original = estimate_identity._divisible_hashes + + def record_sparsity(values, sparsity): + observed_sparsities.append(sparsity) + return original(values, sparsity) + + monkeypatch.setattr(estimate_identity, "_divisible_hashes", record_sparsity) + + assert populateModimizers([], 8, False, 100, 21) == set() + assert observed_sparsities == [8, 4, 2, 1] + assert all(type(value) is int for value in observed_sparsities) + + +def test_convert_to_modimizers_preserves_partition_order(): + hashes = np.ma.MaskedArray( + np.arange(24, dtype=np.uint64), + mask=[False] * 8 + [True] * 4 + [False] * 12, + ) + + sketches = convertToModimizers( + [hashes[:8], hashes[8:16], hashes[16:]], + sparsity=4, + ambiguous=False, + k=5, + expectation=1, + ) + + assert sketches == [{0, 4}, {12}, {16, 20}] + assert all(type(value) is int for sketch in sketches for value in sketch) + + +def test_numpy_fast_path_does_not_iterate_python_scalars(): + class NonIterableArray(np.ndarray): + def __iter__(self): + raise AssertionError("the NumPy fast path must not iterate in Python") + + values = np.arange(100_000, dtype=np.uint64).view(NonIterableArray) + + result = populateModimizers(values, 64, False, 1, 21) + + assert result == set(range(0, 100_000, 64)) + + +def test_nomask_fast_path_does_not_materialize_a_mask(monkeypatch): + partition = np.ma.MaskedArray( + np.arange(100_000, dtype=np.uint64), mask=np.ma.nomask + ) + + def fail_if_called(*_args, **_kwargs): + raise AssertionError("nomask should not be expanded into a boolean array") + + monkeypatch.setattr(estimate_identity.np.ma, "getmaskarray", fail_if_called) + + assert populateModimizers(partition, 64, False, 1, 21) == set(range(0, 100_000, 64)) + + +def test_uniqueness_work_scales_with_selected_candidates(monkeypatch): + # Filtering must happen before uniqueness. At the production-like sparsity + # below, only 1/1024 of the input should reach the more expensive sort. + values = np.arange(2**20, dtype=np.uint64) + unique_input_sizes = [] + original = estimate_identity.np.unique + + def record_unique_input(selected): + unique_input_sizes.append(selected.size) + return original(selected) + + monkeypatch.setattr(estimate_identity.np, "unique", record_unique_input) + + result = populateModimizers(values, 1024, False, 1, 21) + + assert len(result) == 1024 + assert unique_input_sizes == [1024] + + +@pytest.mark.parametrize("sparsity", [0, -1]) +def test_modimizers_reject_nonpositive_sparsity(sparsity): + with pytest.raises(ValueError, match="positive integer"): + populateModimizers([1, 2, 3], sparsity, False, 1, 21) diff --git a/tests/test_native_render.py b/tests/test_native_render.py new file mode 100644 index 0000000..7b9db57 --- /dev/null +++ b/tests/test_native_render.py @@ -0,0 +1,204 @@ +from pathlib import Path + +import matplotlib.pyplot as plt +from matplotlib.colors import to_rgba +import numpy as np +import pandas as pd +import pytest + +from moddotplot.native_render import ( + configure_dotplot_axis, + configure_triangle_axis, + create_triangle_layout, + draw_rectangular_tiles, + draw_triangle_tiles, + genomic_scale, + rectangular_tile_vertices, + save_figure_pair, + tile_width, + transform_triangle_points, + triangle_tile_vertices, +) + + +def _tiles(): + return pd.DataFrame( + { + "q_st": [0, 10], + "q_en": [4, 14], + "r_st": [10, 20], + "r_en": [14, 24], + "discrete": pd.Categorical([0, 1], categories=[0, 1]), + } + ) + + +def test_rectangular_tiles_preserve_center_and_max_query_width_semantics(): + data = _tiles() + + assert tile_width(data) == 4 + vertices = rectangular_tile_vertices(data) + transposed = rectangular_tile_vertices(data, transpose=True) + + np.testing.assert_allclose(vertices[0], [[-2, 8], [2, 8], [2, 12], [-2, 12]]) + np.testing.assert_allclose(transposed[0], [[8, -2], [12, -2], [12, 2], [8, 2]]) + + +def test_rectangular_collection_is_sparse_colored_and_optionally_rasterized(): + figure, axis = plt.subplots() + try: + collection = draw_rectangular_tiles( + axis, + _tiles(), + {0: "#010203", 1: "#abcdef"}, + rasterized=False, + ) + + assert len(collection.get_paths()) == 2 + assert collection.get_rasterized() is False + np.testing.assert_allclose( + collection.get_facecolors(), + [to_rgba("#010203"), to_rgba("#abcdef")], + ) + assert not axis.images + finally: + plt.close(figure) + + +def test_triangle_transform_and_baseline_clipping(): + points = np.asarray([[2, 6], [4, 4], [8, 2]]) + np.testing.assert_allclose( + transform_triangle_points(points), [[4, 2], [4, 0], [5, -3]] + ) + + diagonal = pd.DataFrame( + { + "q_st": [10], + "q_en": [14], + "r_st": [10], + "r_en": [14], + "discrete": [0], + } + ) + polygon = triangle_tile_vertices(diagonal)[0] + assert np.min(polygon[:, 1]) == pytest.approx(0) + assert np.max(polygon[:, 1]) == pytest.approx(2) + assert np.min(polygon[:, 0]) == pytest.approx(8) + assert np.max(polygon[:, 0]) == pytest.approx(12) + + +def test_triangle_collection_omits_tiles_wholly_below_baseline(): + data = pd.DataFrame( + { + "q_st": [0, 20], + "q_en": [4, 24], + "r_st": [10, 0], + "r_en": [14, 4], + "discrete": ["visible", "hidden"], + } + ) + figure, axis = plt.subplots() + try: + collection = draw_triangle_tiles( + axis, + data, + {"visible": "red", "hidden": "blue"}, + rasterized=True, + ) + assert len(collection.get_paths()) == 1 + assert collection.get_rasterized() is True + np.testing.assert_allclose(collection.get_facecolors(), [to_rgba("red")]) + assert np.min(collection.get_paths()[0].vertices[:, 1]) >= 0 + finally: + plt.close(figure) + + +def test_axis_configuration_supports_scaled_and_custom_tick_formatters(): + figure, (dot_axis, triangle_axis) = plt.subplots(1, 2) + try: + configure_dotplot_axis( + dot_axis, + 0, + 1_000_000, + breaks=[0, 500_000, 1_000_000, 1_250_000], + formatter=lambda value, _position: f"{value / 1000:.0f}k", + ) + configure_triangle_axis( + triangle_axis, + 0, + 1_000_000, + breaks=[0, 1_000_000], + ) + + assert dot_axis.get_xlim() == pytest.approx((0, 1_000_000)) + assert dot_axis.xaxis.get_major_formatter()(500_000, 0) == "500k" + assert triangle_axis.get_ylim() == pytest.approx((0, 500_000)) + assert triangle_axis.xaxis.get_major_formatter()(1_000_000, 0) == "1" + assert triangle_axis.get_xlabel() == "Genomic Position (Mbp)" + assert genomic_scale(100_000) == (1_000.0, "Kbp") + assert genomic_scale(500_000_000) == (1_000_000_000.0, "Gbp") + finally: + plt.close(figure) + + +def test_annotated_triangle_layout_shares_genomic_x_axis(): + layout = create_triangle_layout(6, with_annotation=True) + try: + assert layout.annotation_axis is not None + assert layout.triangle_axis.get_shared_x_axes().joined( + layout.triangle_axis, layout.annotation_axis + ) + assert tuple(layout.figure.get_size_inches()) == pytest.approx((6, 3.8)) + finally: + plt.close(layout.figure) + + +@pytest.mark.parametrize( + ("vector_format", "magic"), + [("svg", b" np.iinfo(np.uint32).max for value in hashes) + + +@pytest.mark.parametrize("as_bytes", [False, True]) +def test_native_canonical_hashes_are_reverse_complement_invariant(as_bytes): + sequence = "AACCGTACACTGGACTGAGTCT" + reverse_complement = _reverse_complement(sequence) + if as_bytes: + sequence = sequence.encode("ascii") + reverse_complement = reverse_complement.encode("ascii") + + forward_raw, forward_mask = _nthash.hash_kmers(sequence, 7, True) + reverse_raw, reverse_mask = _nthash.hash_kmers(reverse_complement, 7, True) + + assert forward_mask == reverse_mask == b"" + np.testing.assert_array_equal( + _unpack_hashes(forward_raw), _unpack_hashes(reverse_raw)[::-1] + ) + + +def test_forward_hashes_remain_strand_specific(): + sequence = "AACCGTACACTGGACTGAGTCT" + reverse_complement = _reverse_complement(sequence) + + forward = _hash_sequence(sequence, 7, fw_only=True).data + reverse = _hash_sequence(reverse_complement, 7, fw_only=True).data[::-1] + + assert np.any(forward != reverse) + + +def test_native_mask_marks_every_and_only_ambiguous_window(): + sequence = "ACGTNRYACGT" + k = 3 + raw_hashes, raw_mask = _nthash.hash_kmers(sequence, k, True) + + assert len(raw_hashes) == (len(sequence) - k + 1) * np.dtype(np.uint64).itemsize + assert list(raw_mask) == [0, 0, 1, 1, 1, 1, 1, 0, 0] + assert len(_unpack_hashes(raw_hashes)) == len(raw_mask) == 9 + + +def test_public_ambiguity_flag_masks_or_retains_position_preserving_fallbacks(): + sequence = "ACGTNRYACGT" + expected_mask = np.array([0, 0, 1, 1, 1, 1, 1, 0, 0], dtype=bool) + + excluded = _hash_sequence(sequence, 3, fw_only=False, ambiguous=False) + retained = _hash_sequence(sequence, 3, fw_only=False, ambiguous=True) + + assert len(excluded) == len(retained) == len(sequence) - 3 + 1 + np.testing.assert_array_equal(np.ma.getmaskarray(excluded), expected_mask) + assert not np.ma.getmaskarray(retained).any() + np.testing.assert_array_equal( + excluded.data[~expected_mask], retained.data[~expected_mask] + ) + np.testing.assert_array_equal( + excluded.data[expected_mask], retained.data[expected_mask] + ) + + +def test_generator_keeps_ambiguous_window_positions_as_none_by_default(): + sequence = "ACGTNRYACGT" + + excluded = list( + generateKmersFromFasta(sequence, 3, quiet=True, fw_only=False, ambiguous=False) + ) + retained = list( + generateKmersFromFasta(sequence, 3, quiet=True, fw_only=False, ambiguous=True) + ) + + assert len(excluded) == len(retained) == len(sequence) - 3 + 1 + assert [value is None for value in excluded] == [ + False, + False, + True, + True, + True, + True, + True, + False, + False, + ] + assert all(isinstance(value, int) for value in retained) + + +def test_ambiguous_canonical_fallback_is_reverse_complement_invariant(): + sequence = "AURYNACGTRYSWKMBDHVNACGTU" + reverse_complement = _reverse_complement(sequence) + + forward = _hash_sequence(sequence, 5, fw_only=False, ambiguous=True) + reverse = _hash_sequence(reverse_complement, 5, fw_only=False, ambiguous=True) + + np.testing.assert_array_equal(forward.data, reverse.data[::-1]) + + +@pytest.mark.parametrize("sequence", ["", "A", "AC"]) +def test_short_sequences_return_empty_uint64_results(sequence): + raw_hashes, raw_mask = _nthash.hash_kmers(sequence, 3, True) + hashes = _hash_sequence(sequence, 3, fw_only=False) + + assert raw_hashes == raw_mask == b"" + assert hashes.shape == (0,) + assert hashes.dtype == np.dtype(np.uint64) + assert list(generateKmersFromFasta(sequence, 3, quiet=True, fw_only=False)) == [] + + +@pytest.mark.parametrize("k", [0, -1, 65536]) +def test_native_rejects_unsupported_kmer_sizes(k): + with pytest.raises(ValueError, match="k must be between 1 and 65535"): + _nthash.hash_kmers("ACGT", k, True) + + +def test_hashing_is_ascii_case_insensitive_and_treats_u_as_t(): + dna = _hash_sequence("ACGTTACG", 4, fw_only=False) + lower = _hash_sequence("acgttacg", 4, fw_only=False) + rna = _hash_sequence("ACGUUACG", 4, fw_only=False) + + np.testing.assert_array_equal(dna.data, lower.data) + np.testing.assert_array_equal(dna.data, rna.data) + + +def test_mask_survives_window_partitioning_without_coordinate_collapse(): + hashes = _hash_sequence("AAAAANAAAAA", 3, fw_only=False, ambiguous=False) + + partitions = partitionOverlaps(hashes, win=6, delta=0, seq_len=len(hashes), k=3) + + assert [len(partition) for partition in partitions] == [4, 3] + assert [np.ma.getmaskarray(partition).tolist() for partition in partitions] == [ + [False, False, False, True], + [False, False, False], + ] + + +def test_modimizer_population_skips_masked_hashes_and_returns_python_ints(): + partition = np.ma.MaskedArray( + np.array([2, 100, 4, 200, 6], dtype=np.uint64), + mask=[False, True, False, True, False], + ) + + result = populateModimizers( + partition, sparsity=2, ambiguous=False, expectation=1, k=3 + ) + + assert result == {2, 4, 6} + assert all(type(value) is int for value in result) + + +def test_ambiguity_flag_controls_whether_fallbacks_enter_modimizer_sketches(): + excluded = _hash_sequence("NNNNN", 3, fw_only=False, ambiguous=False) + retained = _hash_sequence("NNNNN", 3, fw_only=False, ambiguous=True) + + excluded_modimizers = populateModimizers( + excluded, sparsity=1, ambiguous=False, expectation=1, k=3 + ) + retained_modimizers = populateModimizers( + retained, sparsity=1, ambiguous=True, expectation=1, k=3 + ) + + assert excluded_modimizers == set() + assert len(retained_modimizers) == 1 diff --git a/tests/test_packaging_metadata.py b/tests/test_packaging_metadata.py new file mode 100644 index 0000000..fa7a0d7 --- /dev/null +++ b/tests/test_packaging_metadata.py @@ -0,0 +1,110 @@ +from pathlib import Path + +try: + import tomllib +except ModuleNotFoundError: # pragma: no cover + # Exercised by the Python 3.8-3.10 CI jobs. + import tomli as tomllib + +from moddotplot.const import VERSION + + +PROJECT_ROOT = Path(__file__).resolve().parents[1] + + +def project_metadata(): + with (PROJECT_ROOT / "pyproject.toml").open("rb") as pyproject: + return tomllib.load(pyproject)["project"] + + +def test_runtime_and_distribution_versions_match(): + assert project_metadata()["version"] == VERSION + + +def test_declared_python_floor_matches_documentation(): + assert project_metadata()["requires-python"] == ">=3.8,<3.13" + readme = (PROJECT_ROOT / "README.md").read_text() + assert "supports Python 3.8 through 3.12" in readme + assert "Python 3.13 and newer are not supported" in readme + + +def test_plotnine_supports_declared_python_floor(): + dependencies = project_metadata()["dependencies"] + assert "plotnine==0.12.4" in dependencies + + +def test_mmh3_is_not_a_runtime_dependency(): + dependencies = project_metadata()["dependencies"] + normalized_names = { + dependency.split(";", 1)[0] + .split("[", 1)[0] + .split("=", 1)[0] + .split("<", 1)[0] + .split(">", 1)[0] + .strip() + .lower() + for dependency in dependencies + } + + assert "mmh3" not in normalized_names + + +def test_replaced_compiled_and_genome_track_dependencies_are_not_runtime_dependencies(): + dependencies = project_metadata()["dependencies"] + normalized_names = { + dependency.split(";", 1)[0] + .split("[", 1)[0] + .split("=", 1)[0] + .split("<", 1)[0] + .split(">", 1)[0] + .strip() + .lower() + for dependency in dependencies + } + + assert "pysam" not in normalized_names + assert "pygenometracks" not in normalized_names + assert "matplotlib" in normalized_names + + +def test_svg_composition_dependencies_are_not_runtime_dependencies(): + dependencies = project_metadata()["dependencies"] + normalized_names = { + dependency.split(";", 1)[0] + .split("[", 1)[0] + .split("=", 1)[0] + .split("<", 1)[0] + .split(">", 1)[0] + .strip() + .lower() + for dependency in dependencies + } + + assert "cairosvg" not in normalized_names + assert "svgutils" not in normalized_names + assert "patchworklib" not in normalized_names + assert "matplotlib" in normalized_names + + +def test_ci_covers_every_supported_python_minor(): + workflow = (PROJECT_ROOT / ".github/workflows/ci.yml").read_text() + for minor in range(8, 13): + assert f' - "3.{minor}"' in workflow + assert ' - "3.13"' not in workflow + + +def test_release_workflow_is_tag_gated_and_uses_trusted_publishing(): + workflow = (PROJECT_ROOT / ".github/workflows/publish-to-pypi.yml").read_text() + setup_config = (PROJECT_ROOT / "setup.cfg").read_text() + + assert ' - "v*"' in workflow + assert "Verify tag matches the package version" in workflow + assert "id-token: write" in workflow + assert "pypa/gh-action-pypi-publish@release/v1" in workflow + assert "pypa/cibuildwheel@v4.2.0" in workflow + assert 'CIBW_BUILD: "cp39-*"' in workflow + assert "py_limited_api = cp38" in setup_config + for runner in ("ubuntu-latest", "macos-15-intel", "windows-latest"): + assert f" - {runner}" in workflow + assert " - sdist" in workflow + assert " - wheels" in workflow diff --git a/tests/test_sketch_cache.py b/tests/test_sketch_cache.py new file mode 100644 index 0000000..4b8f9a4 --- /dev/null +++ b/tests/test_sketch_cache.py @@ -0,0 +1,170 @@ +import gc +import sys +import weakref + +import numpy as np +import pytest + +import moddotplot.estimate_identity as estimate_identity +import moddotplot.moddotplot as cli +from moddotplot.estimate_identity import ( + ModimizerSketchCache, + PreparedModimizerSketches, + create_pairwise_matrix_from_sketches, + create_self_matrix_from_sketches, + createPairwiseMatrix, + createSelfMatrix, + prepare_modimizer_sketches, +) + + +def test_prepared_sketch_matrix_results_match_compatibility_apis(): + first = np.arange(1, 57, dtype=np.uint64) + second = np.arange(29, 85, dtype=np.uint64) + parameters = { + "window_size": 20, + "sparsity": 2, + "delta": 0.5, + "k": 5, + "ambiguous": True, + "expectation": 8, + } + + prepared_first = prepare_modimizer_sketches(len(first), first, **parameters) + prepared_second = prepare_modimizer_sketches(len(second), second, **parameters) + + expected_self = createSelfMatrix( + len(first), + first, + parameters["window_size"], + parameters["sparsity"], + parameters["delta"], + parameters["k"], + 0, + parameters["ambiguous"], + parameters["expectation"], + ) + actual_self = create_self_matrix_from_sketches( + prepared_first, parameters["k"], 0, parameters["ambiguous"] + ) + np.testing.assert_array_equal(actual_self, expected_self) + + expected_pair = createPairwiseMatrix( + len(first), + len(second), + first, + second, + parameters["window_size"], + parameters["sparsity"], + parameters["delta"], + parameters["k"], + 0, + parameters["ambiguous"], + parameters["expectation"], + ) + actual_pair = create_pairwise_matrix_from_sketches( + prepared_first, prepared_second, 0, parameters["k"], True + ) + np.testing.assert_array_equal(actual_pair, expected_pair) + + +def test_sketch_cache_reuses_exact_configuration_without_copying(monkeypatch): + calls = [] + cache_sizes_during_prepare = [] + prepared = PreparedModimizerSketches(core=[{1}], neighbors=[{1, 2}]) + + def fake_prepare(*args): + calls.append(args) + cache_sizes_during_prepare.append(len(cache._cache)) + return prepared + + monkeypatch.setattr(estimate_identity, "prepare_modimizer_sketches", fake_prepare) + cache = ModimizerSketchCache(max_entries=2) + arguments = ("sequence-a", 100, [1, 2, 3], 10, 2, 0.5, 5, False, 5) + + first = cache.get_or_prepare(*arguments) + second = cache.get_or_prepare(*arguments) + + assert first is prepared + assert second is first + assert len(calls) == 1 + + cache.get_or_prepare("sequence-a", 100, [1, 2, 3], 10, 2, 0.25, 5, False, 5) + assert len(calls) == 2 + + cache.get_or_prepare("sequence-b", 100, [1, 2, 3], 10, 2, 0.5, 5, False, 5) + assert len(calls) == 3 + assert cache_sizes_during_prepare == [0, 1, 1] + assert len(cache._cache) == 2 + + +@pytest.mark.parametrize( + ("lengths", "sizing_arguments", "expected_preparations"), + [ + ((1000, 1000), ("--window", "100"), 2), + # With default-style resolution sizing, the shorter self sketch is an + # exact pairwise hit while the longer sequence needs the shorter + # pairwise window: three preparations instead of four. + ((1000, 900), ("--resolution", "10"), 3), + ], +) +def test_two_sequence_grid_reuses_and_releases_prepared_sketches( + monkeypatch, tmp_path, lengths, sizing_arguments, expected_preparations +): + sequences = [[1] * lengths[0], [2] * lengths[1]] + monkeypatch.setattr(cli, "isValidFasta", lambda _path: True) + monkeypatch.setattr(cli, "getInputHeaders", lambda _path: ["chrA", "chrB"]) + monkeypatch.setattr(cli, "readKmersFromFile", lambda *_args: sequences) + + prepare_calls = [] + prepared_references = [] + + def fake_prepare(*args): + prepare_calls.append(args) + prepared = PreparedModimizerSketches(core=[{1}], neighbors=[{1, 2}]) + prepared_references.append(weakref.ref(prepared)) + return prepared + + monkeypatch.setattr(estimate_identity, "prepare_modimizer_sketches", fake_prepare) + monkeypatch.setattr( + cli, + "create_self_matrix_from_sketches", + lambda *_args: np.full((1, 1), 100.0), + ) + monkeypatch.setattr( + cli, + "create_pairwise_matrix_from_sketches", + lambda *_args: np.full((1, 1), 95.0), + ) + monkeypatch.setattr( + cli, + "convertMatrixToBed", + lambda *_args, **_kwargs: [["header"], ["value"]], + ) + + def assert_sketches_released_before_render(**_kwargs): + gc.collect() + assert all(reference() is None for reference in prepared_references) + + monkeypatch.setattr(cli, "create_grid", assert_sketches_released_before_render) + monkeypatch.setattr(cli, "create_plots", lambda **_kwargs: None) + monkeypatch.setattr( + sys, + "argv", + [ + "moddotplot", + "static", + "--fasta", + "sequences.fa", + "--grid-only", + *sizing_arguments, + "--no-bedpe", + "--output-dir", + str(tmp_path), + ], + ) + + cli.main() + + assert len(prepare_calls) == expected_preparations + assert {call[1][0] for call in prepare_calls} == {1, 2} diff --git a/tests/test_sparse_containment.py b/tests/test_sparse_containment.py new file mode 100644 index 0000000..afbe808 --- /dev/null +++ b/tests/test_sparse_containment.py @@ -0,0 +1,148 @@ +import numpy as np +import pytest + +import moddotplot.estimate_identity as estimate_identity +from moddotplot.estimate_identity import ( + _sketch_intersection_counts, + pairwiseContainmentMatrix, + prepare_modimizer_sketches, + selfContainmentMatrix, +) + + +def _scalar_identity(core_a, core_b, expanded_a, expanded_b, identity, k): + core_a = set(core_a) + core_b = set(core_b) + expanded_a = set(expanded_a) + expanded_b = set(expanded_b) + a_to_b = len(core_a & expanded_b) / len(core_a) if core_a else 0.0 + b_to_a = len(core_b & expanded_a) / len(core_b) if core_b else 0.0 + estimated_identity = max(a_to_b, b_to_a) ** (1.0 / k) + return estimated_identity * 100 if estimated_identity >= identity / 100 else 0.0 + + +def _scalar_pairwise(core_x, core_y, expanded_x, expanded_y, identity, k): + return np.asarray( + [ + [ + _scalar_identity( + core_x[x], core_y[y], expanded_x[x], expanded_y[y], identity, k + ) + for x in range(len(core_x)) + ] + for y in range(len(core_y)) + ], + dtype=float, + ).reshape(len(core_y), len(core_x)) + + +def _random_sketches(rng, count, universe, maximum_size): + sketches = [] + expanded = [] + for _ in range(count): + size = int(rng.integers(0, maximum_size + 1)) + core = set(rng.choice(universe, size=size, replace=False).tolist()) + additions = set( + rng.choice( + universe, + size=int(rng.integers(0, maximum_size + 1)), + replace=False, + ).tolist() + ) + sketches.append(core) + expanded.append(core | additions) + return sketches, expanded + + +@pytest.mark.parametrize(("identity", "k"), [(0, 1), (75, 5), (86, 21), (99, 31)]) +def test_sparse_pairwise_matches_scalar_reference(identity, k): + rng = np.random.default_rng(911) + core_x, expanded_x = _random_sketches(rng, 7, 200, 30) + core_y, expanded_y = _random_sketches(rng, 5, 200, 30) + + expected = _scalar_pairwise(core_x, core_y, expanded_x, expanded_y, identity, k) + actual = pairwiseContainmentMatrix( + core_x, + core_y, + expanded_x, + expanded_y, + identity, + k, + supress_progress=True, + ) + + np.testing.assert_allclose(actual, expected) + assert actual.shape == (5, 7) + + +@pytest.mark.parametrize("ambiguous", [False, True]) +def test_sparse_self_matches_scalar_reference_including_empty_diagonal(ambiguous): + rng = np.random.default_rng(77) + core, expanded = _random_sketches(rng, 8, 150, 25) + core[3] = set() + expanded[3] = set() + expected = _scalar_pairwise(core, core, expanded, expanded, 86, 21) + np.fill_diagonal(expected, 100.0) + if not ambiguous: + expected[3, 3] = 0.0 + + actual = selfContainmentMatrix(core, expanded, 21, 86, ambiguous) + + np.testing.assert_allclose(actual, expected) + + +def test_sparse_hash_compression_does_not_alias_out_of_range_python_ints(): + # Casting these values to uint64 would alias -1 and 2**64 - 1. The generic + # compatibility path must continue to treat them as distinct hashes. + sketches_a = [{-1}, {2**64 - 1}, {2**64 + 1}] + sketches_b = [{-1}, {2**64 - 1}, {1}] + + np.testing.assert_array_equal( + _sketch_intersection_counts(sketches_a, sketches_b), + np.diag([1, 1, 0]).astype(np.int32), + ) + + +def test_sparse_counts_use_wide_accumulator(): + # uint8 sparse multiplication silently wraps 300 to 44. + sketch = set(range(300)) + counts = _sketch_intersection_counts([sketch], [sketch]) + + assert counts.dtype == np.int32 + assert counts[0, 0] == 300 + + +def test_matrix_path_does_not_fall_back_to_per_cell_set_intersections(monkeypatch): + def fail_if_called(*_args, **_kwargs): + raise AssertionError("per-cell Python containment was used") + + monkeypatch.setattr(estimate_identity, "containment_neighbors", fail_if_called) + core = [{index, index + 1} for index in range(250)] + expanded = [sketch | {index + 2} for index, sketch in enumerate(core)] + + matrix = pairwiseContainmentMatrix( + core, core, expanded, expanded, 0, 21, supress_progress=True + ) + + assert matrix.shape == (250, 250) + + +def test_prepared_sketches_use_compact_uint64_arrays(): + hashes = np.arange(20_000, dtype=np.uint64) + prepared = prepare_modimizer_sketches( + len(hashes), + hashes, + window_size=1_000, + sparsity=8, + delta=0.5, + k=21, + ambiguous=False, + expectation=125, + ) + + sketches = prepared.core + prepared.neighbors + assert all(isinstance(sketch, np.ndarray) for sketch in sketches) + assert all(sketch.dtype == np.uint64 for sketch in sketches) + assert sum(sketch.nbytes for sketch in sketches) == 8 * sum( + len(sketch) for sketch in sketches + ) diff --git a/tests/test_static_customization.py b/tests/test_static_customization.py new file mode 100644 index 0000000..2ab7424 --- /dev/null +++ b/tests/test_static_customization.py @@ -0,0 +1,210 @@ +import pandas as pd +import pytest + +from moddotplot.static_plots import ( + display_sequence_name, + generate_breaks, + get_colors, + make_dot, +) + + +def test_get_colors_uses_string_custom_breakpoints(): + identity_scores = pd.DataFrame( + { + "perID_by_events": [ + 86.0, + 89.0, + 90.0, + 97.6, + 98.1, + 98.9, + 99.6, + 100.0, + ] + } + ) + custom_breakpoints = [ + "86", + "90", + "97.5", + "97.75", + "98.0", + "98.25", + "98.5", + "98.75", + "99.0", + "99.25", + "99.5", + "100.0", + ] + + bins = get_colors( + identity_scores, + ncolors=11, + is_freq=False, + custom_breakpoints=custom_breakpoints, + ) + + assert bins.astype(int).tolist() == [0, 0, 0, 2, 4, 7, 10, 10] + + +def test_make_dot_uses_custom_color_scale(): + custom_colors = ["#010203", "#456789", "#abcdef"] + plot_data = pd.DataFrame( + { + "q": ["query"] * 3, + "q_st": [0, 10, 20], + "q_en": [10, 20, 30], + "r": ["reference"] * 3, + "r_st": [0, 10, 20], + "r_en": [10, 20, 30], + "discrete": pd.Categorical([0, 1, 2], categories=[0, 1, 2]), + } + ) + + plot = make_dot( + sdf=plot_data, + name_x="query", + name_y="reference", + palette="Spectral_11", + palette_orientation="+", + colors=custom_colors, + breaks=[0, 10, 20, 30], + num_ticks=4, + xlim=30, + deraster=False, + width=4, + is_pairwise=True, + ) + + fill_scale = plot.scales.get_scales("fill") + assert fill_scale.palette(len(custom_colors)) == custom_colors + + +def test_make_dot_honors_exact_region_bounds(): + plot_data = pd.DataFrame( + { + "q": ["query"], + "q_st": [200], + "q_en": [300], + "r": ["reference"], + "r_st": [200], + "r_en": [300], + "discrete": pd.Categorical([0]), + } + ) + + plot = make_dot( + sdf=plot_data, + name_x="query", + name_y="reference", + palette="Spectral_11", + palette_orientation="+", + colors=None, + breaks=None, + num_ticks=4, + xlim=(101, 400), + deraster=False, + width=4, + is_pairwise=True, + ) + + assert plot.scales.get_scales("x").limits == (101.0, 400.0) + assert plot.scales.get_scales("y").limits == (101.0, 400.0) + + +def test_display_names_omit_region_and_full_axes_are_twice_as_large(): + name = "PAN010.chr14.haplotype1.paternal:1-4000000" + plot_data = pd.DataFrame( + { + "q": [name], + "q_st": [1], + "q_en": [100], + "r": [name], + "r_st": [1], + "r_en": [100], + "discrete": pd.Categorical([0]), + } + ) + + plot = make_dot( + sdf=plot_data, + name_x=name, + name_y=name, + palette="Spectral_11", + palette_orientation="+", + colors=None, + breaks=[1, 50, 100], + num_ticks=3, + xlim=(1, 100), + deraster=False, + width=4, + is_pairwise=False, + ) + + assert display_sequence_name(name) == "PAN010.chr14.haplotype1.paternal" + assert ":1-4000000" not in plot.labels.title + assert plot.data["q"].unique().tolist() == ["PAN010.chr14.haplotype1.paternal"] + + figure = plot.draw(show=False) + try: + axis = figure.axes[0] + assert axis.get_xticklabels()[0].get_fontsize() == pytest.approx(8) + # Plotnine rounds text sizes to whole points: 2 * (width * 1.4) = 11.2. + assert axis.xaxis.label.get_fontsize() == pytest.approx(11) + finally: + import matplotlib.pyplot as plt + + plt.close(figure) + + +def test_generated_breaks_never_extend_past_sequence_length(): + breaks = generate_breaks(1, 103_156_783) + + assert breaks + assert all(1 <= value <= 103_156_783 for value in breaks) + assert breaks[-1] == 100_000_000 + + +def test_get_colors_allows_large_custom_palettes(): + scores = pd.DataFrame({"perID_by_events": [0.0, 50.0, 100.0]}) + breaks = list(range(0, 105, 5)) + + bins = get_colors(scores, ncolors=20, is_freq=False, custom_breakpoints=breaks) + + assert bins.astype(int).tolist() == [0, 9, 19] + + +@pytest.mark.parametrize("is_freq", [False, True]) +def test_get_colors_handles_a_single_perfect_identity_bin(is_freq): + scores = pd.DataFrame({"perID_by_events": [100.0, 100.0, 100.0]}) + + bins = get_colors( + scores, + ncolors=11, + is_freq=is_freq, + custom_breakpoints=None, + ) + + assert bins.tolist() == [0, 0, 0] + + +@pytest.mark.parametrize( + ("breakpoints", "message"), + [ + ([0, 50], "number of breakpoints"), + ([0, 75, 50, 100], "strictly increasing"), + ([0, float("nan"), 50, 100], "finite"), + ], +) +def test_get_colors_rejects_invalid_custom_breakpoints(breakpoints, message): + scores = pd.DataFrame({"perID_by_events": [0.0, 50.0, 100.0]}) + + with pytest.raises(ValueError, match=message): + get_colors( + scores, + ncolors=3, + is_freq=False, + custom_breakpoints=breakpoints, + ) From a825b3045199b5f683539142c71584e333acf976 Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Mon, 28 Sep 2026 17:28:15 -0400 Subject: [PATCH 02/16] Complete v1.0.0 plotting and compatibility updates --- .github/workflows/ci.yml | 9 +- .github/workflows/publish-to-pypi.yml | 4 +- CITATION.cff | 49 ++- README.md | 48 ++- pyproject.toml | 7 +- src/moddotplot/annotations.py | 114 ++++++ src/moddotplot/const.py | 1 + src/moddotplot/interactive.py | 280 ++++++++++--- src/moddotplot/moddotplot.py | 470 ++++++++++++++++++--- src/moddotplot/native_render.py | 59 ++- src/moddotplot/plot_summary.py | 113 ++++++ src/moddotplot/static_plots.py | 562 +++++++++++++++++--------- tests/test_annotation_track.py | 163 +++++++- tests/test_cli_integration.py | 25 ++ tests/test_cli_runtime.py | 148 ++++++- tests/test_direction_cli.py | 187 +++++++-- tests/test_entrypoints.py | 14 +- tests/test_grid.py | 83 +++- tests/test_interactive_annotations.py | 192 +++++++++ tests/test_interactive_parser.py | 15 + tests/test_packaging_metadata.py | 16 +- tests/test_plot_fonts.py | 73 ++++ tests/test_plot_summary.py | 61 +++ tests/test_static_customization.py | 48 +++ 24 files changed, 2351 insertions(+), 390 deletions(-) create mode 100644 src/moddotplot/annotations.py create mode 100644 src/moddotplot/plot_summary.py create mode 100644 tests/test_interactive_annotations.py create mode 100644 tests/test_plot_fonts.py create mode 100644 tests/test_plot_summary.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 53dff37..e00998a 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -23,11 +23,10 @@ jobs: fail-fast: false matrix: python-version: - - "3.8" - - "3.9" - - "3.10" - "3.11" - "3.12" + - "3.13" + - "3.14" env: MPLBACKEND: Agg @@ -81,7 +80,7 @@ jobs: - name: Set up Python uses: actions/setup-python@v7 with: - python-version: "3.12" + python-version: "3.14" cache: pip cache-dependency-path: pyproject.toml @@ -105,7 +104,7 @@ jobs: - name: Set up Python uses: actions/setup-python@v7 with: - python-version: "3.12" + python-version: "3.14" cache: pip cache-dependency-path: pyproject.toml diff --git a/.github/workflows/publish-to-pypi.yml b/.github/workflows/publish-to-pypi.yml index fd41b58..903deed 100644 --- a/.github/workflows/publish-to-pypi.yml +++ b/.github/workflows/publish-to-pypi.yml @@ -100,13 +100,13 @@ jobs: with: persist-credentials: false - # The extension uses CPython's stable ABI. Building once with CPython 3.9 + # The extension uses CPython's stable ABI. Building once with CPython 3.11 # produces a cp38-abi3 wheel that supports every Python version declared # by the package without compiling the same binary repeatedly. - name: Build stable-ABI wheel uses: pypa/cibuildwheel@v4.2.0 env: - CIBW_BUILD: "cp39-*" + CIBW_BUILD: "cp311-*" CIBW_SKIP: "*-musllinux_*" CIBW_ARCHS_MACOS: universal2 CIBW_TEST_COMMAND: >- diff --git a/CITATION.cff b/CITATION.cff index 786cba5..655e674 100644 --- a/CITATION.cff +++ b/CITATION.cff @@ -1,15 +1,34 @@ -@article{10.1093/bioinformatics/btae493, - author = {Sweeten, Alexander P and Schatz, Michael C and Phillippy, Adam M}, - title = "{ModDotPlot—rapid and interactive visualization of tandem repeats}", - journal = {Bioinformatics}, - volume = {40}, - number = {8}, - pages = {btae493}, - year = {2024}, - month = {08}, - abstract = "{A common method for analyzing genomic repeats is to produce a sequence similarity matrix visualized via a dot plot. Innovative approaches such as StainedGlass have improved upon this classic visualization by rendering dot plots as a heatmap of sequence identity, enabling researchers to better visualize multi-megabase tandem repeat arrays within centromeres and other heterochromatic regions of the genome. However, computing the similarity estimates for heatmaps requires high computational overhead and can suffer from decreasing accuracy.In this work, we introduce ModDotPlot, an interactive and alignment-free dot plot viewer. By approximating average nucleotide identity via a k-mer-based containment index, ModDotPlot produces accurate plots orders of magnitude faster than StainedGlass. We accomplish this through the use of a hierarchical modimizer scheme that can visualize the full 128 Mb genome of Arabidopsis thaliana in under 5 min on a laptop. ModDotPlot is bundled with a graphical user interface supporting real-time interactive navigation of entire chromosomes.ModDotPlot is available at https://github.com/marbl/ModDotPlot.}", - issn = {1367-4811}, - doi = {10.1093/bioinformatics/btae493}, - url = {https://doi.org/10.1093/bioinformatics/btae493}, - eprint = {https://academic.oup.com/bioinformatics/article-pdf/40/8/btae493/58809824/btae493.pdf}, -} +cff-version: 1.2.0 +message: "If you use ModDotPlot in your research, please cite the article below." +type: software +title: "ModDotPlot" +version: "1.0.0" +authors: + - family-names: "Sweeten" + given-names: "Alexander P." + - family-names: "Schatz" + given-names: "Michael C." + - family-names: "Phillippy" + given-names: "Adam M." +repository-code: "https://github.com/marbl/ModDotPlot" +url: "https://github.com/marbl/ModDotPlot" +license: MIT +preferred-citation: + type: article + title: "ModDotPlot—rapid and interactive visualization of tandem repeats" + authors: + - family-names: "Sweeten" + given-names: "Alexander P." + - family-names: "Schatz" + given-names: "Michael C." + - family-names: "Phillippy" + given-names: "Adam M." + journal: "Bioinformatics" + year: 2024 + month: 8 + volume: 40 + issue: 8 + start: "btae493" + issn: "1367-4811" + doi: "10.1093/bioinformatics/btae493" + url: "https://doi.org/10.1093/bioinformatics/btae493" diff --git a/README.md b/README.md index fe52380..76fc40b 100644 --- a/README.md +++ b/README.md @@ -51,7 +51,7 @@ If you're interested in learning more about _ModDotPlot_ and how to visualize ta ## Installation -_ModDotPlot_ can be installed by running `pip install moddotplot`. Version 1.0.0 supports Python 3.8 through 3.12; Python 3.13 and newer are not supported by the pinned plotting stack. Alternatively, you can download the current release from GitHub by using: +_ModDotPlot_ can be installed by running `pip install moddotplot`. Version 1.0.0 supports Python 3.11 through 3.14 and uses the current Matplotlib 3.11 and Plotnine 0.15 release lines. Alternatively, you can download the current release from GitHub by using: ``` git clone https://github.com/marbl/ModDotPlot.git @@ -82,14 +82,14 @@ Finally, confirm that the installation was installed correctly and that your ver v1.0.0 -usage: moddotplot [-h] {interactive,static} ... +usage: moddotplot [-h] [{static,interactive}] ... ModDotPlot: Visualization of Tandem Repeats positional arguments: - {interactive,static} Choose mode: interactive or static - interactive Interactive mode commands - static Static mode commands + {static,interactive} Choose mode; static is used when omitted + static Static mode commands (default) + interactive Interactive mode commands (deprecated; explicit use only) options: -h, --help show this help message and exit @@ -101,7 +101,13 @@ Note that running `moddotplot -h` might take a while at first! This is because t ## Usage -_ModDotPlot_ must be run either in `static` mode, or `interactive` mode: +_ModDotPlot_ runs in `static` mode by default. The `static` subcommand remains +available for compatibility and clarity, so these forms are equivalent: + +``` +moddotplot -f sequence.fa +moddotplot static -f sequence.fa +``` ### Static Mode @@ -117,7 +123,9 @@ Running _ModDotPlot_ in static mode quickly create plots under the specified out ![](images/moddotplot_output.png) -Plots and histograms are output as both rasterized `.png` images and vector graphics (default: `.svg`). [Plotnine](https://plotnine.readthedocs.io/en/v0.12.4/) provides the primary plotting interface, while Matplotlib directly renders triangle plots, annotation layouts, multi-sequence grids, and each requested output format. +Plots and histograms are output as both rasterized `.png` images and vector graphics (default: `.svg`). [Plotnine](https://plotnine.org/) provides the primary plotting interface, while Matplotlib directly renders triangle plots, annotation layouts, multi-sequence grids, and each requested output format. Grid axes state their genomic unit (Kbp, Mbp, or Gbp). Plot text uses Helvetica by default with an automatic DejaVu Sans fallback if Helvetica cannot render a glyph. + +Every directory containing generated static plots also receives a `plot_summary.txt` reproducibility record. It lists the creation time, absolute plot and input paths, window size, any selected region or annotation BED file, and the exact command used for the run. _ModDotPlot_ supports highly customizable plotting features in static mode. See [static mode commands](#static-mode-commands) for a complete list of features. @@ -128,6 +136,10 @@ _ModDotPlot_ supports highly customizable plotting features in static mode. See moddotplot interactive ``` +Interactive mode is deprecated and maintenance-only. It remains available, but +will not receive new features. It runs only when the `interactive` subcommand is +explicitly provided. + Running _ModDotPlot_ in interactive mode will launch a [Dash application](https://plotly.com/dash/) on your machine's localhost. Open any web browser and go to `http://127.0.0.1:` to view the interactive plot (this should happen automatically, but depending on your environment you might need to copy and paste this URL into your web browser). Running `Ctrl+C` on the command line will exit the Dash application. The default port number used by Dash is `8050`, but this can be customized using the `--port` command (see [interactive mode commands](#interactive-mode-commands) for further info, and [Sample run - Port Forwarding](#sample-run---port-forwarding) for tips on running interactive mode on an HPC environment). --- @@ -140,9 +152,9 @@ The following arguments are the same in both interactive and static mode: Fasta files to input. Multifasta files are accepted. Interactive mode will only support a maximum of two sequences at a time. -`-b / --bed <.bed file>` +`-b / --bed <.bed file(s)>` -Input bedfile used for dotplot annotation (note: this is not the same as the paired-end bed file produced by ModDotPlot). If selected, this will produce an annotated bedtrack image `_ANNOTATION_TRACK` as PNG and in the selected vector format in static mode, and open an IGV js track in the interactive mode Dash application. The name in the bedfile must match the name of the fasta sequence header in order to produce a correct bed track. +Input BED3-BED9 annotation file used for dotplot annotation (this is not the paired-end BEDPE file produced by ModDotPlot). The BED chromosome field must match a FASTA header, excluding any trailing `:start-end` region suffix. Static mode accepts one BED file and produces an annotation track plus annotated triangle output. Interactive mode accepts one or more BED files, combines their matching intervals, and displays a collapsed track beneath the x axis. Comparative interactive plots also display a track beside the y axis when that sequence has matching annotations. BED `itemRgb` colors are used when present. `-k / --kmer ` @@ -258,7 +270,7 @@ List of custom colors in hexcode format can be entered sequentially, mapped from `--plot-direction ` -With FASTA input, additionally create strand-direction plots. Canonical matches that are also present in a forward-only sketch are blue; canonical-only matches, representing reverse orientation, are pink. This option reads each input in both canonical and forward-only modes and cannot be reconstructed from a loaded BEDPE file. +With FASTA input, retain the standard ANI-colored plots and additionally create a `directionality` subfolder. Direction plots use blue for same-orientation matches and pink for reverse-orientation matches, with darker shades representing stronger ANI. Self-comparisons are named `_DIRECTION_FULL`, `_DIRECTION_TRI`, and `_DIRECTION_HIST`; a requested grid is named `_DIRECTION_GRID`. This option reads each input in both canonical and forward-only modes and cannot be reconstructed from a loaded BEDPE file. `--grid ` @@ -391,6 +403,20 @@ moddotplot static -f sequences/*_MATERNAL*.fa --compare-only ### Interactive Mode Commands +`-b / --bed <.bed file> [<.bed file> ...]` + +Add one or more BED3-BED9 annotation files. A self-identity plot shows the +matching track beneath its x axis. A comparative plot shows independent x- and +y-axis tracks when BED chromosome names match both FASTA headers. A FASTA +header such as `chr14_MATERNAL:1-4000000` matches BED chromosome +`chr14_MATERNAL`, and the interactive axes retain those genomic coordinates. +For example: + +``` +moddotplot interactive -f sample1.fa sample2.fa --compare \ + --bed sample1.bed sample2.bed +``` + `--port ` Port to display ModDotPlot on. Default is 8050, this can be changed to any accepted port. @@ -491,6 +517,4 @@ For bug reports or general usage questions, please raise a GitHub issue, or emai - Mac users might encounter the following unexpected command line output: `/bin/sh: lscpu: command not found`. This is a known issue with Plotnine, the Python plotting library used by ModDotPlot. This can be safely ignored. -- If you encounter an error with the following traceback: `rv = reductor(4) TypeError: cannot pickle 'generator' object`, ths means that you have a newer version of Plotnine that is incompatible with ModDotPlot. Please uninstall plotnine and reinstall version 0.12.4 `pip install plotnine==0.12.4`. - - The error ` UserWarning: h5py is running against HDF5 1.xx.x when it was built against 1.xx.x, this may cause problems` is due to the h5py library used by cooler having conflicting versions in the dependency tree. This can also be safely ignored, but if you want to remove this message run `pip uninstall -y h5py` `pip install --no-binary=h5py h5py` diff --git a/pyproject.toml b/pyproject.toml index 2217c0f..6bceec0 100755 --- a/pyproject.toml +++ b/pyproject.toml @@ -5,13 +5,13 @@ build-backend = "setuptools.build_meta" [project] name = "ModDotPlot" version = "1.0.0" -requires-python = ">=3.8,<3.13" +requires-python = ">=3.11,<3.15" dependencies = [ "pandas", - "matplotlib", + "matplotlib>=3.11.2", "plotly", "dash", - "plotnine==0.12.4", + "plotnine>=0.15.8,<0.16", "palettable", "setproctitle", "numpy", @@ -39,5 +39,4 @@ moddotplot = "moddotplot.__main__:main" test = [ "pytest", "pytest-cov", - "tomli>=1.1; python_version < '3.11'", ] diff --git a/src/moddotplot/annotations.py b/src/moddotplot/annotations.py new file mode 100644 index 0000000..33af709 --- /dev/null +++ b/src/moddotplot/annotations.py @@ -0,0 +1,114 @@ +"""Shared BED annotation parsing and interval-selection helpers.""" + +import numpy as np +import pandas as pd + + +DEFAULT_ANNOTATION_COLOR = "#4C72B0" +BED_COLUMNS = [ + "chrom", + "start", + "end", + "name", + "score", + "strand", + "thickStart", + "thickEnd", + "itemRgb", +] + + +def read_annotation_bed(filepath): + """Read the BED3-BED9 subset used by ModDotPlot annotations.""" + + try: + dataframe = pd.read_csv(filepath, sep="\t", comment="#", header=None, dtype=str) + except pd.errors.EmptyDataError: + return pd.DataFrame(columns=BED_COLUMNS[:3]) + + if not 3 <= dataframe.shape[1] <= len(BED_COLUMNS): + raise ValueError( + "Invalid BED file: expected between 3 and 9 tab-separated columns." + ) + + dataframe.columns = BED_COLUMNS[: dataframe.shape[1]] + dataframe["chrom"] = dataframe["chrom"].astype(str) + + for column in ("start", "end"): + try: + values = pd.to_numeric(dataframe[column], errors="raise") + except (TypeError, ValueError) as error: + raise ValueError( + f"Invalid BED file: '{column}' must contain only integers." + ) from error + if values.isna().any() or not np.all(np.isfinite(values)): + raise ValueError( + f"Invalid BED file: '{column}' must contain only finite integers." + ) + if np.any(values % 1 != 0): + raise ValueError( + f"Invalid BED file: '{column}' must contain only integers." + ) + dataframe[column] = values.astype(np.int64) + + if (dataframe["start"] < 0).any(): + raise ValueError("Invalid BED file: 'start' must be non-negative.") + if (dataframe["end"] <= dataframe["start"]).any(): + raise ValueError( + "Invalid BED file: 'end' must be greater than 'start' for every interval." + ) + + return dataframe + + +def read_annotation_beds(filepaths): + """Read and combine one or more annotation BED files.""" + + frames = [read_annotation_bed(filepath) for filepath in filepaths] + if not frames: + return pd.DataFrame(columns=BED_COLUMNS[:3]) + return pd.concat(frames, ignore_index=True, sort=False) + + +def annotation_color(value, fallback=DEFAULT_ANNOTATION_COLOR): + """Return an RGB tuple for a BED ``itemRgb`` value.""" + + if value is None or pd.isna(value): + return fallback + + fields = [field.strip() for field in str(value).split(",")] + if len(fields) != 3: + return fallback + + try: + channels = tuple(int(field) for field in fields) + except ValueError: + return fallback + if any(channel < 0 or channel > 255 for channel in channels): + return fallback + return tuple(channel / 255 for channel in channels) + + +def visible_annotation_intervals( + bed_df, chrom, region_start, region_end, fallback=DEFAULT_ANNOTATION_COLOR +): + """Select and clip BED intervals to a plotted genomic region.""" + + if region_end <= region_start: + raise ValueError("Annotation region end must be greater than its start.") + if bed_df.empty: + return [] + + intervals = [] + matching = bed_df[bed_df["chrom"] == str(chrom)] + has_item_rgb = "itemRgb" in matching.columns + for row in matching.itertuples(index=False): + interval_start = int(row.start) + interval_end = int(row.end) + clipped_start = max(interval_start, region_start) + clipped_end = min(interval_end, region_end) + if clipped_end <= clipped_start: + continue + rgb = getattr(row, "itemRgb", None) if has_item_rgb else None + intervals.append((clipped_start, clipped_end, annotation_color(rgb, fallback))) + return intervals diff --git a/src/moddotplot/const.py b/src/moddotplot/const.py index d30b1e1..facb179 100755 --- a/src/moddotplot/const.py +++ b/src/moddotplot/const.py @@ -1,4 +1,5 @@ VERSION = "1.0.0" +DIRECTION_COLORS = {"Forward": "#2166AC", "Reverse": "#D01C8B"} COLS = [ "#query_name", "query_start", diff --git a/src/moddotplot/interactive.py b/src/moddotplot/interactive.py index 4e207b5..6173a8d 100644 --- a/src/moddotplot/interactive.py +++ b/src/moddotplot/interactive.py @@ -15,12 +15,164 @@ import plotly.graph_objs as go import logging import os +from moddotplot.annotations import visible_annotation_intervals +from moddotplot.parse_fasta import extractRegion + +INTERACTIVE_FONT_FAMILY = "Helvetica, 'DejaVu Sans', sans-serif" # Prevent HTTP protocol requests from showing up in terminal log = logging.getLogger("werkzeug") log.setLevel(logging.ERROR) +def _plotly_annotation_color(color): + """Convert a shared annotation color into a Plotly-compatible value.""" + + if isinstance(color, str): + return color + channels = [round(float(channel) * 255) for channel in color] + return f"rgb({channels[0]},{channels[1]},{channels[2]})" + + +def interactive_axis_bounds(plot_metadata, axis): + """Return genomic bounds for an interactive matrix axis. + + Region-suffixed FASTA identifiers carry the genomic offset needed to align + BED annotations. Older saved metadata has no explicit start/end fields, so + derive them from the axis name while retaining the historical zero-based + bounds for ordinary identifiers. + """ + + size = float(plot_metadata[f"{axis}_size"]) + explicit_start = plot_metadata.get(f"{axis}_start") + explicit_end = plot_metadata.get(f"{axis}_end") + if explicit_start is not None and explicit_end is not None: + return float(explicit_start), float(explicit_end) + + name = plot_metadata[f"{axis}_name"] + parsed_region = extractRegion(name) + if not parsed_region: + return 0.0, size + + _chrom, start, declared_end = parsed_region + expected_size = declared_end - start + 1 + end = declared_end if math.isclose(expected_size, size) else start + size - 1 + return float(start), float(end) + + +def _annotation_chromosome(name): + parsed_region = extractRegion(name) + return parsed_region[0] if parsed_region else name + + +def add_annotation_tracks(figure, plot_metadata, annotations): + """Add collapsed BED tracks aligned to the matrix's visible axes. + + Self-identity plots receive an x-axis track. Comparative plots receive an + x-axis track and, when the y sequence also has annotations, a y-axis track. + Shapes use genomic axis references, so they remain aligned during Plotly + zooming and panning without adding callback state. + """ + + if annotations is None or annotations.empty: + return figure + + x_start, x_end = interactive_axis_bounds(plot_metadata, "x") + x_intervals = visible_annotation_intervals( + annotations, + _annotation_chromosome(plot_metadata["x_name"]), + x_start, + x_end, + ) + y_intervals = [] + if not plot_metadata.get("self", False): + y_start, y_end = interactive_axis_bounds(plot_metadata, "y") + y_intervals = visible_annotation_intervals( + annotations, + _annotation_chromosome(plot_metadata["y_name"]), + y_start, + y_end, + ) + + if x_intervals: + figure.add_shape( + type="rect", + x0=x_start, + x1=x_end, + y0=-0.14, + y1=-0.09, + xref="x", + yref="paper", + fillcolor="#F0F0F0", + line=dict(color="#B0B0B0", width=0.5), + ) + for start, end, color in x_intervals: + figure.add_shape( + type="rect", + x0=start, + x1=end, + y0=-0.14, + y1=-0.09, + xref="x", + yref="paper", + fillcolor=_plotly_annotation_color(color), + line=dict(color=_plotly_annotation_color(color), width=0.5), + ) + bottom_margin = figure.layout.margin.b or 0 + figure.update_layout(margin=dict(b=max(bottom_margin, 140))) + figure.update_xaxes(title_standoff=70) + + if y_intervals: + figure.add_shape( + type="rect", + x0=-0.14, + x1=-0.09, + y0=y_start, + y1=y_end, + xref="paper", + yref="y", + fillcolor="#F0F0F0", + line=dict(color="#B0B0B0", width=0.5), + ) + for start, end, color in y_intervals: + figure.add_shape( + type="rect", + x0=-0.14, + x1=-0.09, + y0=start, + y1=end, + xref="paper", + yref="y", + fillcolor=_plotly_annotation_color(color), + line=dict(color=_plotly_annotation_color(color), width=0.5), + ) + left_margin = figure.layout.margin.l or 0 + figure.update_layout(margin=dict(l=max(left_margin, 140))) + figure.update_yaxes(title_standoff=70) + + return figure + + +def preserve_zoom_ranges(figure, x_range, y_range): + """Keep a callback-generated figure inside the user's requested viewport. + + Plotly includes layout shapes when autoranging. BED tracks intentionally + span the complete genomic interval, so a rebuilt image-pyramid figure must + restore the zoom ranges after adding those shapes. + """ + + x_start, x_end = map(float, x_range) + y_start, y_end = map(float, y_range) + if not all(np.isfinite((x_start, x_end, y_start, y_end))): + raise ValueError("Zoom ranges must contain only finite values") + if x_end <= x_start or y_end <= y_start: + raise ValueError("Zoom range ends must be greater than their starts") + + figure.update_xaxes(range=[x_start, x_end], autorange=False) + figure.update_yaxes(range=[y_start, y_end], autorange=False) + return figure + + def figure_to_bed(figure, default_identity=86.0): """Convert the heatmap in a Dash figure into BEDPE rows. @@ -101,7 +253,16 @@ def find_closest_elements(value, sorted_list): return closest_index1, closest_index2 -def run_dash(matrices, metadata, axes, sparsity, identity, port_number, output_dir): +def run_dash( + matrices, + metadata, + axes, + sparsity, + identity, + port_number, + output_dir, + annotations=None, +): # Run Dash app app = dash.Dash(__name__, prevent_initial_callbacks="initial_duplicate") app.title = "ModDotPlot" @@ -182,10 +343,10 @@ def halving_sequence(size, start): hoverinfo="all", hovertemplate=hover_template_text, name="", - x0=0, + x0=main_x_axis[0], dx=current_metadata["max_window_size"], xtype="scaled", - y0=0, + y0=main_y_axis[0], dy=current_metadata["max_window_size"], ytype="scaled", ) @@ -232,13 +393,17 @@ def halving_sequence(size, start): fig.update_layout( height=800, width=800, - hoverlabel=dict(bgcolor="white", font_size=16, font_family="Helvetica"), + font=dict(family=INTERACTIVE_FONT_FAMILY), + hoverlabel=dict( + bgcolor="white", font_size=16, font_family=INTERACTIVE_FONT_FAMILY + ), yaxis_scaleanchor="x", title=fig_title, - title_font=dict(size=title_size, family="Helvetica, Arial, sans-serif"), + title_font=dict(size=title_size, family=INTERACTIVE_FONT_FAMILY), title_x=0.5, title_y=0.95, ) + add_annotation_tracks(fig, current_metadata, annotations) colorscales = px.colors.named_colorscales() colornames = px.colors.named_colorscales() @@ -294,7 +459,7 @@ def halving_sequence(size, start): "paddingBottom": "40px", "paddingLeft": "40px", "width": "100px", - "fontFamily": "Helvetica, Arial, sans-serif", + "fontFamily": INTERACTIVE_FONT_FAMILY, }, # Added padding to separate the content ), html.Div( @@ -330,7 +495,7 @@ def halving_sequence(size, start): "display": "none" if len(titles) < 2 else "block", "width": "fit-content" * 2, # Set width to fit content - "fontFamily": "Helvetica, Arial, sans-serif", + "fontFamily": INTERACTIVE_FONT_FAMILY, "paddingTop": "10px", "paddingLeft": "25px", "paddingBottom": "30px", @@ -354,7 +519,7 @@ def halving_sequence(size, start): ], id="window-div", style={ - "fontFamily": "Helvetica, Arial, sans-serif", + "fontFamily": INTERACTIVE_FONT_FAMILY, "paddingTop": "10px", "paddingLeft": "45px", "paddingBot": "30px", @@ -365,7 +530,7 @@ def halving_sequence(size, start): html.Div( f"Minimum Window Size: {current_metadata['min_window_size']}", style={ - "fontFamily": "Helvetica, Arial, sans-serif", + "fontFamily": INTERACTIVE_FONT_FAMILY, "paddingTop": "10px", "paddingLeft": "45px", "paddingBot": "30px", @@ -388,7 +553,7 @@ def halving_sequence(size, start): "paddingBottom": "30px", "width": "260px", "textAlign": "left", - "fontFamily": "Helvetica, Arial, sans-serif", + "fontFamily": INTERACTIVE_FONT_FAMILY, }, ), html.Div(id="content"), @@ -396,7 +561,7 @@ def halving_sequence(size, start): "Color Palette", id="color-palette-text", style={ - "fontFamily": "Helvetica, Arial, sans-serif", + "fontFamily": INTERACTIVE_FONT_FAMILY, "paddingLeft": "45px", # Added padding to separate the content }, ), @@ -1279,7 +1444,7 @@ def halving_sequence(size, start): "paddingTop": "10px", "paddingLeft": "20px", "width": "250px", - "fontFamily": "Helvetica, Arial, sans-serif", + "fontFamily": INTERACTIVE_FONT_FAMILY, }, # Added padding to separate the content ) ], @@ -1290,7 +1455,7 @@ def halving_sequence(size, start): "Coordinate Log:", id="coordinate-text", style={ - "fontFamily": "Helvetica, Arial, sans-serif", + "fontFamily": INTERACTIVE_FONT_FAMILY, "paddingLeft": "45px", # Added padding to separate the content "paddingBottom": "10px", }, @@ -1307,7 +1472,7 @@ def halving_sequence(size, start): "height": "40px", "width": "62%", "margin-left": "22px", - "fontFamily": "Helvetica, Arial, sans-serif", + "fontFamily": INTERACTIVE_FONT_FAMILY, }, ), html.Button( @@ -1322,7 +1487,7 @@ def halving_sequence(size, start): ), ], style={ - "fontFamily": "Helvetica, Arial, sans-serif", + "fontFamily": INTERACTIVE_FONT_FAMILY, }, ), ], @@ -1477,10 +1642,20 @@ def update_dotplot( if len(updated_info["x_name"]) + len(updated_info["y_name"]) > 22: title_size = 16 new_main_level = image_pyramid[0] - fig = go.Figure(data=[heatmap]) current_color = getInteractiveColor(getMatchingColors(color), "+") - new_heatmap = heatmap - new_heatmap.update(dict(colorscale=current_color, z=new_main_level)) + new_heatmap = go.Heatmap(heatmap.to_plotly_json()) + new_heatmap.update( + dict( + colorscale=current_color, + z=new_main_level, + x=image_axes[0], + y=image_axes[1], + x0=image_axes[0][0], + y0=image_axes[1][0], + dx=updated_info["max_window_size"], + dy=updated_info["max_window_size"], + ) + ) masked_data = np.where( (new_heatmap["z"] < threshold_range[0]) @@ -1525,16 +1700,22 @@ def update_dotplot( fig.update_yaxes(title_text=updated_info["y_name"], title_font=dict(size=18)) fig.update_layout( - hoverlabel=dict(bgcolor="white", font_size=16, font_family="Helvetica"), + font=dict(family=INTERACTIVE_FONT_FAMILY), + hoverlabel=dict( + bgcolor="white", + font_size=16, + font_family=INTERACTIVE_FONT_FAMILY, + ), title=updated_title, # Add your title here title_font=dict( - size=title_size, family="Helvetica" + size=title_size, family=INTERACTIVE_FONT_FAMILY ), # Adjust the title font size if needed title_x=0.5, title_y=0.95, ) fig.update_layout(yaxis_scaleanchor="x") + add_annotation_tracks(fig, updated_info, annotations) if relayoutData is not None: # TODO: Pan mode should stay in pan mode @@ -1543,32 +1724,30 @@ def update_dotplot( y_start_range = relayoutData.get("yaxis.range[0]") x_end_range = relayoutData.get("xaxis.range[1]") y_end_range = relayoutData.get("yaxis.range[1]") + x_axis_start, x_axis_end = interactive_axis_bounds(updated_info, "x") + y_axis_start, y_axis_end = interactive_axis_bounds(updated_info, "y") # Check that selected range is in bounds, snap to boundary otherwise # TODO: still broken here if x_start_range is not None: - if x_start_range < 0: - print("x axis out of bounds! Shifting to x=0") - relayoutData["xaxis.range[0]"] = 0 - x_start_range = 0 - if y_start_range < 0: - print("y axis out of bounds! Shifting to y=0") - relayoutData["yaxis.range[0]"] = 0 - y_start_range = 0 - if x_end_range > updated_info["x_size"]: - print( - f"x axis out of bounds! Shifting to x={updated_info['x_size']}" - ) - relayoutData["xaxis.range[1]"] = updated_info["x_size"] - x_end_range = updated_info["x_size"] - if y_end_range > updated_info["y_size"]: - print( - f"y axis out of bounds! Shifting to y={updated_info['y_size']}" - ) - relayoutData["yaxis.range[1]"] = updated_info["y_size"] - y_end_range = updated_info["y_size"] - - if x_start_range >= 0 and y_start_range >= 0: + if x_start_range < x_axis_start: + print(f"x axis out of bounds! Shifting to x={x_axis_start:g}") + relayoutData["xaxis.range[0]"] = x_axis_start + x_start_range = x_axis_start + if y_start_range < y_axis_start: + print(f"y axis out of bounds! Shifting to y={y_axis_start:g}") + relayoutData["yaxis.range[0]"] = y_axis_start + y_start_range = y_axis_start + if x_end_range > x_axis_end: + print(f"x axis out of bounds! Shifting to x={x_axis_end:g}") + relayoutData["xaxis.range[1]"] = x_axis_end + x_end_range = x_axis_end + if y_end_range > y_axis_end: + print(f"y axis out of bounds! Shifting to y={y_axis_end:g}") + relayoutData["yaxis.range[1]"] = y_axis_end + y_end_range = y_axis_end + + if x_start_range >= x_axis_start and y_start_range >= y_axis_start: x_begin = round(relayoutData["xaxis.range[0]"]) x_end = round(relayoutData["xaxis.range[1]"]) @@ -1681,15 +1860,24 @@ def update_dotplot( title_text=updated_info["y_name"], title_font=dict(size=18) ) current_fig.update_layout( + font=dict(family=INTERACTIVE_FONT_FAMILY), hoverlabel=dict( - bgcolor="white", font_size=16, font_family="Helvetica" + bgcolor="white", + font_size=16, + font_family=INTERACTIVE_FONT_FAMILY, ), - title=fig_title, # Add your title here - title_font=dict(size=28, family="Helvetica"), + title=updated_title, # Add your title here + title_font=dict(size=28, family=INTERACTIVE_FONT_FAMILY), title_x=0.5, # Center horizontally title_y=0.95, # Center vertically ) current_fig.update_layout(yaxis_scaleanchor="x") + add_annotation_tracks(current_fig, updated_info, annotations) + preserve_zoom_ranges( + current_fig, + (x_start_range, x_end_range), + (y_start_range, y_end_range), + ) return ( current_fig, "", @@ -1764,7 +1952,7 @@ def update_dotplot( .custom-slider .rc-slider-handle { border: 2px solid black; /* Change the handle border color to gray */ - font-family: Helvetica, Arial, sans-serif; + font-family: Helvetica, "DejaVu Sans", sans-serif; } #main_color { justify-content:center; diff --git a/src/moddotplot/moddotplot.py b/src/moddotplot/moddotplot.py index 626522d..08b243f 100755 --- a/src/moddotplot/moddotplot.py +++ b/src/moddotplot/moddotplot.py @@ -22,7 +22,8 @@ ModimizerSketchCache, partitionOverlaps, ) -from moddotplot.interactive import run_dash +from moddotplot.interactive import interactive_axis_bounds, run_dash +from moddotplot.annotations import read_annotation_beds from moddotplot.const import ASCII_ART, VERSION import argparse @@ -31,6 +32,9 @@ import numpy as np import pickle import os +import shlex + +from moddotplot.plot_summary import PlotSummaryWriter # Static plotting pulls in the Plotnine and Matplotlib stacks. Keep those @@ -39,11 +43,16 @@ read_df_from_file = None create_plots = None create_grid = None -create_direction_plot = None + +COMMANDS = frozenset({"interactive", "static"}) +INTERACTIVE_DEPRECATION_MESSAGE = ( + "Warning: interactive mode is deprecated and maintenance-only. " + "It remains available, but will not receive new features." +) def _load_static_plotting(): - global read_df_from_file, create_plots, create_grid, create_direction_plot + global read_df_from_file, create_plots, create_grid from moddotplot import static_plots @@ -53,8 +62,6 @@ def _load_static_plotting(): create_plots = static_plots.create_plots if create_grid is None: create_grid = static_plots.create_grid - if create_direction_plot is None: - create_direction_plot = static_plots.create_direction_plot def get_parser(): @@ -67,12 +74,21 @@ def get_parser(): description="ModDotPlot: Visualization of Tandem Repeats", ) subparsers = parser.add_subparsers( - dest="command", required=True, help="Choose mode: interactive or static" + dest="command", + required=False, + help="Choose mode; static is used when omitted", + ) + static_parser = subparsers.add_parser( + "static", help="Static mode commands (default)" ) interactive_parser = subparsers.add_parser( - "interactive", help="Interactive mode commands" + "interactive", + help="Interactive mode commands (deprecated; explicit use only)", + description=( + "Deprecated interactive mode. This mode remains available but is " + "maintenance-only and will not receive new features." + ), ) - static_parser = subparsers.add_parser("static", help="Static mode commands") # -----------INTERACTIVE MODE SUBCOMMANDS----------- interactive_input_group = interactive_parser.add_mutually_exclusive_group( @@ -149,6 +165,17 @@ def get_parser(): help="Directory name for saving matrices and coordinate logs. Defaults to working directory.", ) + interactive_parser.add_argument( + "-b", + "--bed", + default=None, + nargs="+", + help=( + "BED3-BED9 annotation file(s). Tracks are shown for BED chromosome " + "names matching the FASTA headers on each matrix axis." + ), + ) + compare_group.add_argument( "--compare", action="store_true", @@ -358,7 +385,10 @@ def get_parser(): static_parser.add_argument( "--plot-direction", action="store_true", - help="Create a plot containing the direction of each k-mer array (relative to the first array). Arrays with inversions will be highlighted in blue (forward) and pink (reverse).", + help=( + "Color matches in _FULL, _TRI, and grid plots by k-mer orientation: " + "blue for the same orientation and pink for reverse orientation." + ), ) static_parser.add_argument( @@ -440,6 +470,23 @@ def get_parser(): return parser +def _arguments_with_default_command(arguments=None): + """Return CLI arguments with static mode inserted when no mode is given.""" + + normalized = list(sys.argv[1:] if arguments is None else arguments) + if normalized[:1] and normalized[0] in ("-h", "--help"): + return normalized + if not normalized or normalized[0] not in COMMANDS: + normalized.insert(0, "static") + return normalized + + +def parse_args(arguments=None): + """Parse command-line arguments, defaulting omitted subcommands to static.""" + + return get_parser().parse_args(_arguments_with_default_command(arguments)) + + def _apply_static_config(args, config): """Apply static-mode JSON configuration values to parsed arguments.""" # TODO: Remove args that are interactive only @@ -550,14 +597,92 @@ def _slice_kmers_for_region(kmers, region, kmer_size, sequence_start=1): return selected +def _bedpe_window_sizes(dataframe): + """Extract the window sizes represented by a loaded BEDPE dataframe.""" + + for start, end in (("query_start", "query_end"), ("q_st", "q_en")): + if start in dataframe and end in dataframe: + sizes = dataframe[end] - dataframe[start] + return sorted({int(size) for size in sizes if size > 0}) + return [] + + +def _regions_from_names(names): + """Return normalized region strings embedded in sequence names.""" + + regions = [] + for name in names: + parsed = extractRegion(name) + if parsed: + region = f"{parsed[0]}:{parsed[1]}-{parsed[2]}" + if region not in regions: + regions.append(region) + return regions + + +def _annotate_bed_directions( + bed, + canonical_matrix, + forward_matrix, + window_size, + x_offset, + y_offset, +): + """Add forward/reverse orientation labels to canonical BEDPE rows.""" + + canonical_matrix = np.asarray(canonical_matrix) + forward_matrix = np.asarray(forward_matrix) + if canonical_matrix.shape != forward_matrix.shape: + raise ValueError("Canonical and forward matrices must have matching shapes") + if not bed: + return bed + + header = list(bed[0]) + try: + query_start_index = header.index("query_start") + reference_start_index = header.index("reference_start") + except ValueError as error: + raise ValueError( + "BEDPE data is missing direction coordinate columns" + ) from error + + annotated = [tuple([*header, "direction"])] + for row in bed[1:]: + query_index = round((int(row[query_start_index]) - x_offset) / window_size) + reference_index = round( + (int(row[reference_start_index]) - y_offset) / window_size + ) + try: + direction = ( + "Forward" + if forward_matrix[query_index, reference_index] > 0 + else "Reverse" + ) + except IndexError as error: + raise ValueError( + "BEDPE coordinates fall outside direction matrices" + ) from error + annotated.append(tuple([*row, direction])) + return annotated + + def main(): print(ASCII_ART) print(f"v{VERSION} \n") - args = get_parser().parse_args() + args = parse_args() + summary_writer = PlotSummaryWriter(shlex.join(sys.argv)) + annotation_df = None + if args.command == "interactive" and args.bed: + try: + annotation_df = read_annotation_beds(args.bed) + except (OSError, ValueError) as error: + print(f"Error reading annotation BED file(s): {error}", file=sys.stderr) + sys.exit(2) if args.command == "static": _load_static_plotting() # -----------MUTUALLY EXCLUSIVE: INTERACTIVE OR STATIC MODE----------- if args.command == "interactive": + print(INTERACTIVE_DEPRECATION_MESSAGE, file=sys.stderr) print(f"Running ModDotPlot in interactive mode\n") # -----------LOAD MATRICES FOR INTERACTIVE MODE----------- if hasattr(args, "load") and args.load: @@ -571,16 +696,16 @@ def main(): for i in range(len(matrices)): matrix_axes = [] for matrix in matrices[i]: + x_start, x_end = interactive_axis_bounds(metadata[i], "x") + y_start, y_end = interactive_axis_bounds(metadata[i], "y") x_axis = [ - j * round(metadata[i]["x_size"] / matrix.shape[0]) - for j in range(matrix.shape[0]) + value + for value in np.linspace(x_start, x_end, matrix.shape[0] + 1) ] y_axis = [ - j * round(metadata[i]["y_size"] / matrix.shape[1]) - for j in range(matrix.shape[1]) + value + for value in np.linspace(y_start, y_end, matrix.shape[1] + 1) ] - x_axis.append(metadata[i]["x_size"]) - y_axis.append(metadata[i]["y_size"]) matrix_axes.append(x_axis) matrix_axes.append(y_axis) axes.append(matrix_axes) @@ -592,6 +717,7 @@ def main(): args.identity, args.port, args.output_dir, + annotation_df, ) sys.exit(0) elif args.command == "static": @@ -633,12 +759,21 @@ def main(): single_val_name = [] double_val_name = [] xlim_val_grid = 0 + loaded_window_sizes = [] + loaded_regions = [] for bed in args.load: # If args.load is provided as input, run static mode directly from the paired-end bed file. Skip counting input k-mers. df = read_df_from_file(bed) unique_query_names = df["#query_name"].unique() unique_reference_names = df["reference_name"].unique() + bed_window_sizes = _bedpe_window_sizes(df) + bed_regions = _regions_from_names( + [*unique_query_names, *unique_reference_names] + ) + if args.grid or args.grid_only: + loaded_window_sizes.extend(bed_window_sizes) + loaded_regions.extend(bed_regions) assert len(unique_query_names) == len(unique_reference_names) assert len(unique_reference_names) == 1 self_id_scores = df[df["#query_name"] == df["reference_name"]] @@ -652,7 +787,7 @@ def main(): os.makedirs(args.output_dir) if len(self_id_scores) > 1: if not args.grid_only: - create_plots( + plot_files = create_plots( sdf=None, directory=args.output_dir if args.output_dir else ".", name_x=unique_query_names[0], @@ -674,6 +809,14 @@ def main(): deraster=args.deraster, annotation=args.bed, ) + summary_writer.add( + args.output_dir, + plot_files or [], + window_sizes=bed_window_sizes, + regions=bed_regions, + bed_file=args.bed, + bedpe_inputs=[bed], + ) if args.grid or args.grid_only: single_vals.append(df) single_val_name.append(unique_query_names[0]) @@ -690,7 +833,7 @@ def main(): if len(pairwise_id_scores) > 1: if not args.grid_only: # Potentially sort - create_plots( + plot_files = create_plots( sdf=None, directory=args.output_dir if args.output_dir else ".", name_x=unique_query_names[0], @@ -712,6 +855,14 @@ def main(): deraster=args.deraster, annotation=args.bed, ) + summary_writer.add( + args.output_dir, + plot_files or [], + window_sizes=bed_window_sizes, + regions=bed_regions, + bed_file=args.bed, + bedpe_inputs=[bed], + ) if args.grid or args.grid_only: double_vals.append(df) double_val_name.append( @@ -725,7 +876,7 @@ def main(): print( f"Creating a {len(single_val_name)}x{len(single_val_name)} grid.\n" ) - create_grid( + plot_files = create_grid( singles=single_vals, doubles=double_vals, directory=args.output_dir if args.output_dir else ".", @@ -745,6 +896,14 @@ def main(): vector_format=args.vector, dpi=args.dpi, ) + summary_writer.add( + args.output_dir, + plot_files or [], + window_sizes=loaded_window_sizes, + regions=loaded_regions, + bed_file=args.bed, + bedpe_inputs=args.load, + ) sys.exit(0) # -----------INPUT SEQUENCE VALIDATION----------- @@ -767,6 +926,13 @@ def main(): ) fasta_list.remove(i) + fasta_source_by_name = {} + for fasta_path, headers in fasta_headers.items(): + for header in headers: + parsed_header = extractRegion(header) + base_header = parsed_header[0] if parsed_header else header + fasta_source_by_name.setdefault(base_header, fasta_path) + try: region_by_name = _parse_region_arguments( getattr(args, "region", None), seq_list @@ -1096,9 +1262,11 @@ def main(): axes = [] for matrices_set, meta in zip(matrices, metadata): matrix_axes = [] + x_start, x_end = interactive_axis_bounds(meta, "x") + y_start, y_end = interactive_axis_bounds(meta, "y") for matrix in matrices_set: - x_axis = np.linspace(0, meta["x_size"], matrix.shape[0] + 1) - y_axis = np.linspace(0, meta["y_size"], matrix.shape[1] + 1) + x_axis = np.linspace(x_start, x_end, matrix.shape[0] + 1) + y_axis = np.linspace(y_start, y_end, matrix.shape[1] + 1) matrix_axes.append(x_axis) matrix_axes.append(y_axis) axes.append(matrix_axes) @@ -1110,14 +1278,22 @@ def main(): args.identity, args.port, args.output_dir, + annotation_df, ) # -----------SETUP STATIC MODE----------- elif args.command == "static": # -----------SET SPARSITY VALUE----------- + direction_rendering = args.plot_direction and ( + not args.no_plot or args.grid or args.grid_only + ) if args.grid or args.grid_only: grid_val_singles = [] grid_val_single_names = [] + grid_window_sizes = [] + if direction_rendering: + direction_grid_val_singles = [] + direction_grid_val_doubles = [] if direction_k_list is None: new_sequences = list(zip(seq_list, k_list)) else: @@ -1270,7 +1446,7 @@ def main(): # calculations finish. del prepared_self direction_self_mat = None - if args.plot_direction and not args.no_plot and not args.grid_only: + if direction_rendering: direction_self_mat = createSelfMatrix( seq_length, alternate_sequence, @@ -1294,9 +1470,41 @@ def main(): subseq_end_pos, subseq_end_pos, ) + plot_bed = bed + if direction_rendering: + if args.forward: + canonical_matrix = direction_self_mat + forward_matrix = self_mat + canonical_bed = convertMatrixToBed( + canonical_matrix, + win, + args.identity, + seq_name, + seq_name, + True, + seq_start_pos, + seq_start_pos, + subseq_end_pos, + subseq_end_pos, + ) + else: + canonical_matrix = self_mat + forward_matrix = direction_self_mat + canonical_bed = bed + plot_bed = _annotate_bed_directions( + canonical_bed, + canonical_matrix, + forward_matrix, + win, + seq_start_pos, + seq_start_pos, + ) if args.grid or args.grid_only: grid_val_singles.append(bed) grid_val_single_names.append(seq_name) + grid_window_sizes.append(win) + if direction_rendering: + direction_grid_val_singles.append(plot_bed) if args.cooler: try: @@ -1341,7 +1549,7 @@ def main(): ) if (not args.no_plot) and (not args.grid_only): - create_plots( + plot_files = create_plots( sdf=[bed], directory=bedpe_path, name_x=seq_name, @@ -1363,29 +1571,50 @@ def main(): deraster=args.deraster, annotation=args.bed, ) - if args.plot_direction: - if args.forward: - canonical_matrix = direction_self_mat - forward_matrix = self_mat - else: - canonical_matrix = self_mat - forward_matrix = direction_self_mat - create_direction_plot( - canonical_matrix=canonical_matrix, - forward_matrix=forward_matrix, - window_size=win, - directory=bedpe_path, + self_region = ( + [f"{base_name}:{seq_range[1]}-{seq_range[2]}"] + if seq_range + else _regions_from_names([sequence_name]) + ) + summary_writer.add( + bedpe_path, + plot_files or [], + fasta_files=[fasta_source_by_name[base_name]], + window_sizes=[win], + regions=self_region, + bed_file=args.bed, + ) + if direction_rendering: + direction_directory = os.path.join(bedpe_path, "directionality") + direction_files = create_plots( + sdf=[plot_bed], + directory=direction_directory, name_x=seq_name, name_y=seq_name, - self_identity=True, + palette=args.palette, + palette_orientation=args.palette_orientation, + no_hist=args.no_hist, width=args.width, dpi=args.dpi, - vector_format=args.vector, - deraster=args.deraster, + is_freq=args.bin_freq, xlim=plot_axis_bounds, + custom_colors=args.colors, + custom_breakpoints=args.breakpoints, + from_file=None, + is_pairwise=False, axes_labels=args.axes_ticks, - x_offset=seq_start_pos, - y_offset=seq_start_pos, + axes_tick_number=args.axes_number, + vector_format=args.vector, + deraster=args.deraster, + annotation=None, + ) + summary_writer.add( + direction_directory, + direction_files or [], + fasta_files=[fasta_source_by_name[base_name]], + window_sizes=[win], + regions=self_region, + bed_file=args.bed, ) # -----------COMPUTE COMPARATIVE PLOTS----------- @@ -1585,7 +1814,7 @@ def main(): # evicts or clears them. del prepared_smaller, prepared_larger direction_pair_mat = None - if args.plot_direction and not args.no_plot and not args.grid_only: + if direction_rendering: direction_pair_mat = createPairwiseMatrix( smaller_length, larger_length, @@ -1660,11 +1889,43 @@ def main(): larger_seq_end_pos, smaller_seq_end_pos, ) + plot_bed = bed + if direction_rendering: + if args.forward: + canonical_matrix = direction_pair_mat + forward_matrix = pair_mat + canonical_bed = convertMatrixToBed( + canonical_matrix, + win, + args.identity, + larger_seq_name, + smaller_seq_name, + False, + larger_seq_start_pos, + smaller_seq_start_pos, + larger_seq_end_pos, + smaller_seq_end_pos, + ) + else: + canonical_matrix = pair_mat + forward_matrix = direction_pair_mat + canonical_bed = bed + plot_bed = _annotate_bed_directions( + canonical_bed, + canonical_matrix, + forward_matrix, + win, + larger_seq_start_pos, + smaller_seq_start_pos, + ) if args.grid or args.grid_only: grid_val_doubles.append(bed) grid_val_double_names.append( [larger_seq_name, smaller_seq_name] ) + grid_window_sizes.append(win) + if direction_rendering: + direction_grid_val_doubles.append(plot_bed) bedfile_prefix = larger_seq_name + "_" + smaller_seq_name bedpe_path = os.path.join( args.output_dir or ".", bedfile_prefix @@ -1687,7 +1948,7 @@ def main(): ) if (not args.no_plot) and (not args.grid_only): - create_plots( + plot_files = create_plots( sdf=[bed], directory=bedpe_path, name_x=larger_seq_name, @@ -1709,29 +1970,70 @@ def main(): deraster=args.deraster, annotation=args.bed, ) - if args.plot_direction: - if args.forward: - canonical_matrix = direction_pair_mat - forward_matrix = pair_mat - else: - canonical_matrix = pair_mat - forward_matrix = direction_pair_mat - create_direction_plot( - canonical_matrix=canonical_matrix, - forward_matrix=forward_matrix, - window_size=win, - directory=bedpe_path, + pair_regions = [] + if larger_seq_range: + pair_regions.append( + f"{larger_base_name}:{larger_seq_range[1]}-" + f"{larger_seq_range[2]}" + ) + else: + pair_regions.extend( + _regions_from_names([larger_sequence_name]) + ) + if smaller_seq_range: + pair_regions.append( + f"{smaller_base_name}:{smaller_seq_range[1]}-" + f"{smaller_seq_range[2]}" + ) + else: + pair_regions.extend( + _regions_from_names([smaller_sequence_name]) + ) + pair_fasta_files = [ + fasta_source_by_name[name] + for name in (larger_base_name, smaller_base_name) + ] + summary_writer.add( + bedpe_path, + plot_files or [], + fasta_files=pair_fasta_files, + window_sizes=[win], + regions=pair_regions, + bed_file=args.bed, + ) + if direction_rendering: + direction_directory = os.path.join( + bedpe_path, "directionality" + ) + direction_files = create_plots( + sdf=[plot_bed], + directory=direction_directory, name_x=larger_seq_name, name_y=smaller_seq_name, - self_identity=False, + palette=args.palette, + palette_orientation=args.palette_orientation, + no_hist=args.no_hist, width=args.width, dpi=args.dpi, - vector_format=args.vector, - deraster=args.deraster, + is_freq=args.bin_freq, xlim=pair_axis_bounds, + custom_colors=args.colors, + custom_breakpoints=args.breakpoints, + from_file=None, + is_pairwise=True, axes_labels=args.axes_ticks, - x_offset=larger_seq_start_pos, - y_offset=smaller_seq_start_pos, + axes_tick_number=args.axes_number, + vector_format=args.vector, + deraster=args.deraster, + annotation=None, + ) + summary_writer.add( + direction_directory, + direction_files or [], + fasta_files=pair_fasta_files, + window_sizes=[win], + regions=pair_regions, + bed_file=args.bed, ) if sketch_cache is not None: @@ -1741,7 +2043,7 @@ def main(): if args.axes_limits: xlim_val_grid = args.axes_limits print(f"Creating a {len(sequences)}x{len(sequences)} grid.\n") - create_grid( + plot_files = create_grid( singles=grid_val_singles, doubles=grid_val_doubles, directory=args.output_dir if args.output_dir else ".", @@ -1761,6 +2063,54 @@ def main(): vector_format=args.vector, dpi=args.dpi, ) + grid_directory = args.output_dir if args.output_dir else "." + grid_regions = [ + f"{name}:{region[1]}-{region[2]}" + for name, region in region_by_name.items() + ] + for embedded_region in _regions_from_names( + [sequence[0] for sequence in sequences] + ): + if embedded_region not in grid_regions: + grid_regions.append(embedded_region) + summary_writer.add( + grid_directory, + plot_files or [], + fasta_files=fasta_list, + window_sizes=grid_window_sizes, + regions=grid_regions, + bed_file=args.bed, + ) + if direction_rendering: + direction_directory = os.path.join(grid_directory, "directionality") + direction_files = create_grid( + singles=direction_grid_val_singles, + doubles=direction_grid_val_doubles, + directory=direction_directory, + palette=args.palette, + palette_orientation=args.palette_orientation, + single_names=grid_val_single_names, + double_names=grid_val_double_names, + is_freq=args.bin_freq, + xlim=xlim_val_grid, + custom_colors=args.colors, + custom_breakpoints=args.breakpoints, + axes_label=args.axes_ticks, + is_bed=False, + width=args.width, + breaks=args.axes_ticks, + deraster=args.deraster, + vector_format=args.vector, + dpi=args.dpi, + ) + summary_writer.add( + direction_directory, + direction_files or [], + fasta_files=fasta_list, + window_sizes=grid_window_sizes, + regions=grid_regions, + bed_file=args.bed, + ) if __name__ == "__main__": diff --git a/src/moddotplot/native_render.py b/src/moddotplot/native_render.py index 50be2d0..a79f5f3 100644 --- a/src/moddotplot/native_render.py +++ b/src/moddotplot/native_render.py @@ -15,6 +15,7 @@ from matplotlib.axes import Axes from matplotlib.collections import PolyCollection from matplotlib.figure import Figure +from matplotlib.text import Text from matplotlib.ticker import FuncFormatter import numpy as np import pandas as pd @@ -23,6 +24,56 @@ ColorSource = Union[Sequence[str], Mapping[object, str]] TickFormatter = Callable[[float, int], str] +DEFAULT_FONT_FAMILY = "Helvetica" +FALLBACK_FONT_FAMILY = "DejaVu Sans" +MIN_TEXT_SIZE = 8.0 +MIN_TITLE_SIZE = 10.0 + + +def clamped_font_size( + width: float, + multiplier: float, + minimum: float = MIN_TEXT_SIZE, + maximum: Optional[float] = None, +) -> float: + """Scale a font with figure width without allowing unreadable sizes.""" + + size = float(width) * float(multiplier) + if not np.isfinite(size): + raise ValueError("Calculated font size must be finite") + size = max(float(minimum), size) + if maximum is not None: + size = min(float(maximum), size) + return size + + +def set_figure_font_family(figure: Figure, family: str) -> None: + """Set every existing text artist in a figure to one font family.""" + + for artist in figure.findobj(match=Text): + artist.set_fontfamily(family) + + +def is_glyph_loading_error(error: BaseException) -> bool: + """Return whether Matplotlib failed while loading a font glyph.""" + + return ( + isinstance(error, RuntimeError) and "failed to load glyph" in str(error).lower() + ) + + +def save_with_font_fallback(figure: Figure, save: Callable[[], None]) -> None: + """Save with Helvetica, retrying the complete operation with DejaVu Sans.""" + + set_figure_font_family(figure, DEFAULT_FONT_FAMILY) + try: + save() + except RuntimeError as error: + if not is_glyph_loading_error(error): + raise + set_figure_font_family(figure, FALLBACK_FONT_FAMILY) + save() + @dataclass(frozen=True) class TriangleLayout: @@ -351,8 +402,12 @@ def save_figure_pair( "transparent": transparent, "bbox_inches": bbox_inches, } - figure.savefig(png_path, format="png", **save_options) - figure.savefig(vector_path, format=vector_format, **save_options) + + def save_outputs() -> None: + figure.savefig(png_path, format="png", **save_options) + figure.savefig(vector_path, format=vector_format, **save_options) + + save_with_font_fallback(figure, save_outputs) return png_path, vector_path diff --git a/src/moddotplot/plot_summary.py b/src/moddotplot/plot_summary.py new file mode 100644 index 0000000..0c268df --- /dev/null +++ b/src/moddotplot/plot_summary.py @@ -0,0 +1,113 @@ +"""Reproducibility summaries for static plot output directories.""" + +from __future__ import annotations + +from datetime import datetime +from pathlib import Path +from typing import Iterable, Optional + + +def _absolute_paths(paths: Optional[Iterable[str]]) -> list[str]: + """Return stable, de-duplicated absolute paths in input order.""" + + resolved = [] + for path in paths or (): + absolute = str(Path(path).expanduser().resolve()) + if absolute not in resolved: + resolved.append(absolute) + return resolved + + +class PlotSummaryWriter: + """Collect and write one human-readable summary per plot directory.""" + + filename = "plot_summary.txt" + + def __init__(self, command: str, created_at: Optional[datetime] = None): + self.command = command + self.created_at = created_at + self._records: dict[str, list[dict[str, object]]] = {} + + def add( + self, + directory: str, + plot_files: Iterable[str], + *, + fasta_files: Optional[Iterable[str]] = None, + window_sizes: Optional[Iterable[int]] = None, + regions: Optional[Iterable[str]] = None, + bed_file: Optional[str] = None, + bedpe_inputs: Optional[Iterable[str]] = None, + ) -> Optional[Path]: + """Add a plot record and immediately refresh its directory summary. + + Empty plot records are ignored. This keeps mocked or skipped renderers + from claiming that files were produced when they were not. + """ + + plots = _absolute_paths(plot_files) + if not plots: + return None + + output_directory = str(Path(directory).expanduser().resolve()) + record = { + "created_at": (self.created_at or datetime.now().astimezone()).isoformat( + timespec="seconds" + ), + "plot_files": plots, + "fasta_files": _absolute_paths(fasta_files), + "window_sizes": list( + dict.fromkeys(int(size) for size in window_sizes or ()) + ), + "regions": list(dict.fromkeys(regions or ())), + "bed_file": _absolute_paths([bed_file])[0] if bed_file else None, + "bedpe_inputs": _absolute_paths(bedpe_inputs), + } + records = self._records.setdefault(output_directory, []) + if record not in records: + records.append(record) + return self._write(output_directory) + + def _write(self, directory: str) -> Path: + output_directory = Path(directory) + output_directory.mkdir(parents=True, exist_ok=True) + summary_path = output_directory / self.filename + lines = [ + "ModDotPlot Plot Summary", + "", + f"Command: {self.command}", + "", + f"Output directory: {output_directory}", + ] + + for record in self._records[directory]: + lines.extend(["", f"Created: {record['created_at']}"]) + self._append_list(lines, "Plot files", record["plot_files"]) + self._append_list( + lines, + "FASTA files", + record["fasta_files"], + empty="None (plot loaded from BEDPE)", + ) + window_sizes = [f"{size} bp" for size in record["window_sizes"]] + self._append_list(lines, "Window sizes", window_sizes, empty="Unknown") + self._append_list( + lines, + "Regions", + record["regions"], + empty="None (whole sequence)", + ) + lines.extend(["", f"BED annotation file: {record['bed_file'] or 'None'}"]) + if record["bedpe_inputs"]: + self._append_list(lines, "Input BEDPE files", record["bedpe_inputs"]) + + summary_path.write_text("\n".join(lines) + "\n", encoding="utf-8") + return summary_path + + @staticmethod + def _append_list(lines, label, values, *, empty=None): + lines.extend(["", f"{label}:"]) + if values: + lines.extend(f" - {value}" for value in values) + elif empty is not None: + lines.append(f" - {empty}") diff --git a/src/moddotplot/static_plots.py b/src/moddotplot/static_plots.py index e542cc0..6edee35 100755 --- a/src/moddotplot/static_plots.py +++ b/src/moddotplot/static_plots.py @@ -29,9 +29,15 @@ import os import re import matplotlib.pyplot as plt +from matplotlib.colors import to_hex, to_rgb from matplotlib.patches import Rectangle from matplotlib.ticker import ScalarFormatter from moddotplot.native_render import ( + DEFAULT_FONT_FAMILY, + FALLBACK_FONT_FAMILY, + MIN_TEXT_SIZE, + MIN_TITLE_SIZE, + clamped_font_size, configure_dotplot_axis, configure_triangle_axis, create_triangle_layout, @@ -39,28 +45,61 @@ draw_triangle_tiles, genomic_scale, genomic_tick_formatter, + is_glyph_loading_error, save_figure_pair, + save_with_font_fallback, + set_figure_font_family, ) from moddotplot.const import ( + DIRECTION_COLORS, DIVERGING_PALETTES, QUALITATIVE_PALETTES, SEQUENTIAL_PALETTES, ) from palettable.colorbrewer import qualitative, sequential, diverging +from moddotplot.annotations import ( + DEFAULT_ANNOTATION_COLOR, + annotation_color as _annotation_color, + read_annotation_bed, + visible_annotation_intervals as _visible_annotation_intervals, +) -DEFAULT_ANNOTATION_COLOR = "#4C72B0" REGION_SUFFIX_PATTERN = re.compile(r"(?::\d+-\d+)+$") +def _plot_font_theme(family=DEFAULT_FONT_FAMILY): + """Apply one family to every Plotnine text themeable.""" + + font = element_text(family=[family]) + return theme( + text=font, + title=element_text(family=[family]), + axis_text=element_text(family=[family]), + strip_text=element_text(family=[family]), + legend_text=element_text(family=[family]), + ) + + +def _save_plot(plot, **kwargs): + """Save a Plotnine plot in Helvetica, retrying on glyph-load failure.""" + + try: + ggsave(plot + _plot_font_theme(), **kwargs) + except RuntimeError as error: + if not is_glyph_loading_error(error): + raise + ggsave(plot + _plot_font_theme(FALLBACK_FONT_FAMILY), **kwargs) + + def display_sequence_name(name): """Return a sequence name without appended region coordinates.""" return REGION_SUFFIX_PATTERN.sub("", str(name)) -def _fit_grid_sequence_labels(figure, axes, minimum_size=2.0): - """Shrink grid sequence headings until each fits inside its own panel.""" +def _fit_grid_sequence_labels(figure, axes, minimum_size=MIN_TEXT_SIZE): + """Fit grid headings without shrinking them below a readable size.""" figure.canvas.draw() renderer = figure.canvas.get_renderer() @@ -82,15 +121,24 @@ def _fit_grid_sequence_labels(figure, axes, minimum_size=2.0): if label_box.height: ratios.append(axis_box.height * 0.9 / label_box.height) - if not ratios: + artists = [axis.title for axis in axes[0, :]] + [ + axis.yaxis.label for axis in axes[:, 0] + ] + artists = [artist for artist in artists if artist.get_text()] + if not ratios or not artists: return scale = min(1.0, min(ratios)) if scale >= 1.0: return - artists = [axis.title for axis in axes[0, :]] + [ - axis.yaxis.label for axis in axes[:, 0] - ] + sizes = [artist.get_fontsize() for artist in artists] + smallest_scaled_size = min(size * scale for size in sizes) + if smallest_scaled_size < minimum_size: + enlargement = minimum_size / smallest_scaled_size + width, height = figure.get_size_inches() + figure.set_size_inches(width * enlargement, height * enlargement, forward=True) + scale = min(1.0, scale * enlargement) + for artist in artists: artist.set_fontsize(max(minimum_size, artist.get_fontsize() * scale)) @@ -112,107 +160,52 @@ def _resolve_native_colors(palette, palette_orientation, custom_colors=None): return list(custom_colors) if custom_colors else list(colors) -def is_plot_empty(p): - # Check if the plot has data or any layers - return len(p.layers) == 0 and p.data.empty - - -def read_annotation_bed(filepath): - """Read the BED3-BED9 subset used by ModDotPlot annotations.""" - col_names = [ - "chrom", - "start", - "end", - "name", - "score", - "strand", - "thickStart", - "thickEnd", - "itemRgb", - ] +DIRECTION_ANI_COLUMN = "direction_ani" - try: - df = pd.read_csv(filepath, sep="\t", comment="#", header=None, dtype=str) - except pd.errors.EmptyDataError: - return pd.DataFrame(columns=col_names[:3]) - if not 3 <= df.shape[1] <= len(col_names): - raise ValueError( - "Invalid BED file: expected between 3 and 9 tab-separated columns." - ) +def _direction_ani_style(dataframe): + """Return data and colors for direction hue plus ANI intensity. - df.columns = col_names[: df.shape[1]] - df["chrom"] = df["chrom"].astype(str) - - for column in ("start", "end"): - try: - values = pd.to_numeric(df[column], errors="raise") - except (TypeError, ValueError) as error: - raise ValueError( - f"Invalid BED file: '{column}' must contain only integers." - ) from error - if values.isna().any() or not np.all(np.isfinite(values)): - raise ValueError( - f"Invalid BED file: '{column}' must contain only finite integers." - ) - if np.any(values % 1 != 0): - raise ValueError( - f"Invalid BED file: '{column}' must contain only integers." - ) - df[column] = values.astype(np.int64) - - if (df["start"] < 0).any(): - raise ValueError("Invalid BED file: 'start' must be non-negative.") - if (df["end"] <= df["start"]).any(): - raise ValueError( - "Invalid BED file: 'end' must be greater than 'start' for every interval." - ) - - return df - - -def _annotation_color(value, fallback=DEFAULT_ANNOTATION_COLOR): - """Return a Matplotlib color for a BED ``itemRgb`` value.""" - if value is None or pd.isna(value): - return fallback + Direction selects the blue or pink hue. The existing ordered ANI bins + control saturation: weak matches are pale and the strongest bin reaches + the base direction color. Data without ANI bins retains the solid legacy + direction colors, which keeps the low-level rendering API usable. + """ - fields = [field.strip() for field in str(value).split(",")] - if len(fields) != 3: - return fallback + if "direction" not in dataframe.columns: + return dataframe, None, None + if "discrete" not in dataframe.columns: + return dataframe, DIRECTION_COLORS, "direction" - try: - channels = tuple(int(field) for field in fields) - except ValueError: - return fallback - if any(channel < 0 or channel > 255 for channel in channels): - return fallback - return tuple(channel / 255 for channel in channels) + categories = ( + list(dataframe["discrete"].cat.categories) + if isinstance(dataframe["discrete"].dtype, pd.CategoricalDtype) + else list(pd.unique(dataframe["discrete"].dropna())) + ) + if not categories: + return dataframe, DIRECTION_COLORS, "direction" + + colors = {} + category_count = len(categories) + for direction, base_color in DIRECTION_COLORS.items(): + base_rgb = np.asarray(to_rgb(base_color)) + for index, category in enumerate(categories): + fraction = 1.0 if category_count == 1 else index / (category_count - 1) + strength = 0.25 + (0.75 * fraction) + rgb = np.ones(3) + ((base_rgb - np.ones(3)) * strength) + colors[f"{direction}:{category}"] = to_hex(rgb) + + styled = dataframe.copy() + styled[DIRECTION_ANI_COLUMN] = [ + f"{direction}:{category}" + for direction, category in zip(styled["direction"], styled["discrete"]) + ] + return styled, colors, DIRECTION_ANI_COLUMN -def _visible_annotation_intervals( - bed_df, chrom, region_start, region_end, fallback=DEFAULT_ANNOTATION_COLOR -): - """Select and clip BED intervals to a plotted genomic region.""" - if region_end <= region_start: - raise ValueError("Annotation region end must be greater than its start.") - if bed_df.empty: - return [] - - intervals = [] - matching = bed_df[bed_df["chrom"] == str(chrom)] - has_item_rgb = "itemRgb" in matching.columns - for row in matching.itertuples(index=False): - interval_start = int(row.start) - interval_end = int(row.end) - if interval_end <= interval_start: - continue - clipped_start = max(interval_start, region_start) - clipped_end = min(interval_end, region_end) - if clipped_end <= clipped_start: - continue - rgb = getattr(row, "itemRgb", None) if has_item_rgb else None - intervals.append((clipped_start, clipped_end, _annotation_color(rgb, fallback))) - return intervals +def is_plot_empty(p): + # Check if the plot has data or any layers + return len(p.layers) == 0 and p.data.empty def draw_annotation_track( @@ -276,14 +269,20 @@ def render_annotation_track( return False figure.subplots_adjust(left=0.02, right=0.995, bottom=0.34, top=0.96) - figure.savefig( - f"{output_prefix}.{vector_format}", - format=vector_format, - dpi=dpi, - transparent=vector_format != "ps", - facecolor="white" if vector_format == "ps" else "none", - ) - figure.savefig(f"{output_prefix}.png", format="png", dpi=dpi, facecolor="white") + + def save_outputs(): + figure.savefig( + f"{output_prefix}.{vector_format}", + format=vector_format, + dpi=dpi, + transparent=vector_format != "ps", + facecolor="white" if vector_format == "ps" else "none", + ) + figure.savefig( + f"{output_prefix}.png", format="png", dpi=dpi, facecolor="white" + ) + + save_with_font_fallback(figure, save_outputs) return True finally: plt.close(figure) @@ -542,6 +541,8 @@ def make_dot( title_length = 1.5 * width elif len(title_name) > 80: title_length = width + sdf, direction_colors, direction_column = _direction_ani_style(sdf) + direction_coloring = direction_colors is not None # Select the color palette if hasattr(diverging, palette): function_name = getattr(diverging, palette) @@ -562,6 +563,8 @@ def make_dot( new_hexcodes = hexcodes[::-1] if palette_orientation == "-" else hexcodes if colors: new_hexcodes = colors # Override colors if provided + fill_column = direction_column if direction_coloring else "discrete" + fill_colors = direction_colors if direction_coloring else new_hexcodes # Determine the exact genomic interval. A two-value limit is supplied by # FASTA mode so blank edge windows do not shrink or extend the plot. min_val, max_val = _data_axis_limits(sdf, xlim) @@ -599,25 +602,33 @@ def make_dot( plot_background=element_blank(), panel_background=element_blank(), axis_line=element_line(color="black"), - axis_text=element_text(family=["DejaVu Sans"], size=width * 2), + axis_text=element_text( + family=[DEFAULT_FONT_FAMILY], + size=clamped_font_size(width, 2.0), + ), axis_ticks_major=element_line( size=(width), color="black" ), # Increased tick length title=element_text( - family=["DejaVu Sans"], size=title_length, hjust=0.5 + family=[DEFAULT_FONT_FAMILY], + size=max(MIN_TITLE_SIZE, title_length), + hjust=0.5, ), # Center title - axis_title_x=element_text(size=(width * 2.8), family=["DejaVu Sans"]), + axis_title_x=element_text( + size=clamped_font_size(width, 2.8), + family=[DEFAULT_FONT_FAMILY], + ), strip_background=element_blank(), # Remove facet strip background strip_text=element_text( - size=(width * 1.2), family=["DejaVu Sans"] + size=clamped_font_size(width, 1.2), family=[DEFAULT_FONT_FAMILY] ), # Customize facet label text size (optional) ) # Construct the plot arguments ggplot_args = ( ggplot(sdf) - + scale_color_discrete(guide=False) - + scale_fill_manual(values=new_hexcodes, guide=False) + + scale_color_discrete(guide=None) + + scale_fill_manual(values=fill_colors, guide=None) + common_theme + scale_x_continuous( labels=make_scale, limits=[min_val, max_val], breaks=breaks @@ -631,7 +642,7 @@ def make_dot( ) p = ggplot_args + _dotplot_tiles( - aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), + aes(x="q_st", y="r_st", fill=fill_column, height=window, width=window), deraster, ) @@ -705,23 +716,32 @@ def make_dot_grid( plot_background=element_blank(), panel_background=element_blank(), axis_line=element_line(color="black"), - axis_text=element_text(family=["DejaVu Sans"], size=width), + axis_text=element_text( + family=[DEFAULT_FONT_FAMILY], size=clamped_font_size(width, 1.0) + ), axis_ticks_major=element_line( size=(width), color="black" ), # Increased tick length - title=element_text(size=(width * 1.2), alpha=0), - axis_title_x=element_text(size=(width * 1.2), family=["DejaVu Sans"]), + title=element_text( + size=clamped_font_size(width, 1.2, MIN_TITLE_SIZE), + family=[DEFAULT_FONT_FAMILY], + alpha=0, + ), + axis_title_x=element_text( + size=clamped_font_size(width, 1.2), + family=[DEFAULT_FONT_FAMILY], + ), strip_background=element_blank(), # Remove facet strip background strip_text=element_text( - size=(width * 1.2), family=["DejaVu Sans"] + size=clamped_font_size(width, 1.2), family=[DEFAULT_FONT_FAMILY] ), # Customize facet label text size (optional) ) # Construct the plot arguments ggplot_args = ( ggplot(sdf) - + scale_color_discrete(guide=False) - + scale_fill_manual(values=new_hexcodes, guide=False) + + scale_color_discrete(guide=None) + + scale_fill_manual(values=new_hexcodes, guide=None) + common_theme + scale_x_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) + scale_y_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) @@ -879,9 +899,19 @@ def create_direction_plot( legend_position="none", panel_grid_major=element_blank(), panel_grid_minor=element_blank(), - axis_text=element_text(family=["DejaVu Sans"], size=width), - title=element_text(size=width * 1.4, hjust=0.5), - axis_title_x=element_text(size=width * 1.2, family=["DejaVu Sans"]), + axis_text=element_text( + family=[DEFAULT_FONT_FAMILY], + size=clamped_font_size(width, 1.0), + ), + title=element_text( + family=[DEFAULT_FONT_FAMILY], + size=clamped_font_size(width, 1.4, MIN_TITLE_SIZE), + hjust=0.5, + ), + axis_title_x=element_text( + size=clamped_font_size(width, 1.2), + family=[DEFAULT_FONT_FAMILY], + ), ) ) @@ -890,7 +920,7 @@ def create_direction_plot( f"{name_x}_DIRECTION" if self_identity else f"{name_x}_{name_y}_DIRECTION" ) prefix = os.path.join(directory, filename) - ggsave( + _save_plot( plot, width=width, height=width, @@ -899,7 +929,7 @@ def create_direction_plot( filename=f"{prefix}.{vector_format}", verbose=False, ) - ggsave( + _save_plot( plot, width=width, height=width, @@ -942,7 +972,10 @@ def make_dot_final( plot_background=element_blank(), panel_background=element_blank(), axis_line=element_line(color="black"), - axis_text=element_text(family=["DejaVu Sans"], size=width), + axis_text=element_text( + family=[DEFAULT_FONT_FAMILY], + size=clamped_font_size(width, 1.0), + ), axis_ticks_major=element_line(), axis_title_x=element_blank(), axis_title_y=element_blank(), @@ -1004,8 +1037,8 @@ def make_dot_final( aes(x=x_col, y=y_col, fill="discrete", height=window, width=window), deraster, ) - + scale_color_discrete(guide=False) - + scale_fill_manual(values=new_hexcodes, guide=False) + + scale_color_discrete(guide=None) + + scale_fill_manual(values=new_hexcodes, guide=None) + theme( legend_position="none", panel_grid_major=element_blank(), @@ -1013,9 +1046,12 @@ def make_dot_final( plot_background=element_blank(), panel_background=element_blank(), axis_line=element_line(color="black"), - axis_text=element_text(family=["DejaVu Sans"], size=width), + axis_text=element_text( + family=[DEFAULT_FONT_FAMILY], + size=clamped_font_size(width, 1.0), + ), axis_ticks_major=element_line(), - title=element_text(family=["Dejavu Sans"]), + title=element_text(family=[DEFAULT_FONT_FAMILY]), ) + scale_x_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) + scale_y_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) @@ -1029,8 +1065,8 @@ def make_dot_final( aes(x=x_col, y=y_col, fill="discrete", height=window, width=window), deraster, ) - + scale_color_discrete(guide=False) - + scale_fill_manual(values=new_hexcodes, guide=False) + + scale_color_discrete(guide=None) + + scale_fill_manual(values=new_hexcodes, guide=None) + theme( legend_position="none", panel_grid_major=element_blank(), @@ -1038,9 +1074,12 @@ def make_dot_final( plot_background=element_blank(), panel_background=element_blank(), axis_line=element_line(color="black"), - axis_text=element_text(family=["DejaVu Sans"], size=width), + axis_text=element_text( + family=[DEFAULT_FONT_FAMILY], + size=clamped_font_size(width, 1.0), + ), axis_ticks_major=element_line(), - title=element_text(family=["Dejavu Sans"]), + title=element_text(family=[DEFAULT_FONT_FAMILY]), ) + scale_x_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) + scale_y_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) @@ -1120,8 +1159,8 @@ def make_tri( deraster, alpha=1.0, ) # Ensure full opacity - + scale_fill_manual(values=new_hexcodes, guide=False) - + scale_color_discrete(guide=False) + + scale_fill_manual(values=new_hexcodes, guide=None) + + scale_color_discrete(guide=None) + scale_x_continuous( labels=make_scale, limits=[min_val, max_val], breaks=breaks ) @@ -1136,14 +1175,24 @@ def make_tri( panel_grid_minor=element_blank(), plot_background=element_blank(), panel_background=element_blank(), - axis_text=element_text(family=["DejaVu Sans"], size=width), + axis_text=element_text( + family=[DEFAULT_FONT_FAMILY], + size=clamped_font_size(width, 1.0), + ), axis_line_x=element_line(), axis_line_y=element_blank(), axis_ticks_major_x=element_line(), axis_ticks_major_y=element_blank(), axis_ticks_major=element_line(size=(width)), - title=element_text(size=(width * 1.4), hjust=0.5), - axis_title_x=element_text(size=(width * 1.4), family=["DejaVu Sans"]), + title=element_text( + family=[DEFAULT_FONT_FAMILY], + size=clamped_font_size(width, 1.4, MIN_TITLE_SIZE), + hjust=0.5, + ), + axis_title_x=element_text( + size=clamped_font_size(width, 1.4), + family=[DEFAULT_FONT_FAMILY], + ), axis_text_y=element_blank(), ) ) @@ -1153,8 +1202,8 @@ def make_tri( aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), alpha=0, ) - + scale_color_discrete(guide=False) - + scale_fill_manual(values=new_hexcodes, guide=False) + + scale_color_discrete(guide=None) + + scale_fill_manual(values=new_hexcodes, guide=None) + scale_x_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) + scale_y_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) + coord_fixed(ratio=1) @@ -1166,7 +1215,10 @@ def make_tri( plot_background=element_blank(), panel_background=element_blank(), axis_line=element_line(color="black"), - axis_text=element_text(family=["DejaVu Sans"], size=width), + axis_text=element_text( + family=[DEFAULT_FONT_FAMILY], + size=clamped_font_size(width, 1.0), + ), axis_ticks_major=element_line(), axis_line_x=element_line(), axis_line_y=element_blank(), @@ -1175,7 +1227,10 @@ def make_tri( axis_text_x=element_line(), axis_text_y=element_blank(), plot_title=element_blank(), - axis_title_x=element_text(size=(width * 1.2), family=["DejaVu Sans"]), + axis_title_x=element_text( + size=clamped_font_size(width, 1.2), + family=[DEFAULT_FONT_FAMILY], + ), ) ) else: @@ -1186,8 +1241,8 @@ def make_tri( deraster, alpha=1.0, ) # Ensure full opacity - + scale_fill_manual(values=new_hexcodes, guide=False) - + scale_color_discrete(guide=False) + + scale_fill_manual(values=new_hexcodes, guide=None) + + scale_color_discrete(guide=None) + scale_x_continuous( labels=make_scale, limits=[min_val, max_val], breaks=breaks ) @@ -1202,7 +1257,10 @@ def make_tri( panel_grid_minor=element_blank(), plot_background=element_blank(), panel_background=element_blank(), - axis_text=element_text(family=["DejaVu Sans"], size=width), + axis_text=element_text( + family=[DEFAULT_FONT_FAMILY], + size=clamped_font_size(width, 1.0), + ), axis_line_x=element_line(), axis_line_y=element_blank(), axis_ticks_major_x=element_line(), @@ -1210,7 +1268,10 @@ def make_tri( axis_ticks_major=element_line(), axis_text_y=element_blank(), title=element_blank(), - axis_title_x=element_text(size=(width * 1.2), family=["DejaVu Sans"]), + axis_title_x=element_text( + size=clamped_font_size(width, 1.2), + family=[DEFAULT_FONT_FAMILY], + ), ) ) axis = ( @@ -1219,8 +1280,8 @@ def make_tri( aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), alpha=0, ) - + scale_color_discrete(guide=False) - + scale_fill_manual(values=new_hexcodes, guide=False) + + scale_color_discrete(guide=None) + + scale_fill_manual(values=new_hexcodes, guide=None) + scale_x_continuous( labels=make_scale, limits=[min_val, max_val], breaks=breaks ) @@ -1236,7 +1297,7 @@ def make_tri( plot_background=element_blank(), panel_background=element_blank(), axis_line=element_line(color="black"), - axis_text=element_text(family=["DejaVu Sans"]), + axis_text=element_text(family=[DEFAULT_FONT_FAMILY]), axis_ticks_major=element_line(), axis_line_x=element_line(), axis_line_y=element_blank(), @@ -1245,7 +1306,10 @@ def make_tri( axis_text_x=element_line(), axis_text_y=element_blank(), plot_title=element_blank(), - axis_title_x=element_text(size=(width * 1.2), family=["DejaVu Sans"]), + axis_title_x=element_text( + size=clamped_font_size(width, 1.2), + family=[DEFAULT_FONT_FAMILY], + ), ) ) @@ -1300,10 +1364,10 @@ def make_tri_axis(sdf, title_name, palette, palette_orientation, colors, breaks, aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), alpha=0, ) - + scale_color_discrete(guide=False) + + scale_color_discrete(guide=None) + scale_fill_manual( values=new_hexcodes, - guide=False, + guide=None, ) + theme( legend_position="none", @@ -1313,7 +1377,7 @@ def make_tri_axis(sdf, title_name, palette, palette_orientation, colors, breaks, panel_background=element_blank(), axis_line=element_line(color="black"), # Adjust axis line size axis_text=element_text( - family=["DejaVu Sans"] + family=[DEFAULT_FONT_FAMILY] ), # Change axis text font and size axis_ticks_major=element_line(), axis_line_x=element_line(), # Keep the x-axis line @@ -1374,13 +1438,16 @@ def make_hist(sdf, palette, palette_orientation, custom_colors, custom_breakpoin if count > 1e6: extra = "\n(thousands)" + sdf, direction_colors, direction_column = _direction_ani_style(sdf) + fill_column = direction_column or "discrete" + fill_colors = direction_colors or new_hexcodes p = ( - ggplot(data=sdf, mapping=aes(x="perID_by_events", fill="discrete")) + ggplot(data=sdf, mapping=aes(x="perID_by_events", fill=fill_column)) + geom_histogram(bins=300) + scale_color_cmap(cmap_name="plasma") - + scale_fill_manual(new_hexcodes) + + scale_fill_manual(fill_colors) + theme_light() - + theme(text=element_text(family=["DejaVu Sans"])) + + _plot_font_theme() + theme(legend_position="none") + coord_cartesian(xlim=(bot, 100)) + xlab("% Identity Estimate") @@ -1427,7 +1494,11 @@ def _build_triangle_figure( if axes_labels else generate_breaks(int(region_start), int(region_end)) ) - colors = _resolve_native_colors(palette, palette_orientation, custom_colors) + sdf, direction_colors, direction_column = _direction_ani_style(sdf) + colors = direction_colors or _resolve_native_colors( + palette, palette_orientation, custom_colors + ) + color_column = direction_column or "discrete" with_annotation = annotation_df is not None layout = create_triangle_layout(width, with_annotation=with_annotation) try: @@ -1435,6 +1506,7 @@ def _build_triangle_figure( layout.triangle_axis, sdf, colors, + color_column=color_column, rasterized=not deraster, ) configure_triangle_axis( @@ -1444,8 +1516,15 @@ def _build_triangle_figure( breaks=breaks, label=not with_annotation, ) + if with_annotation: + # The annotation axis is the sole genomic axis in the combined + # figure. Remove the triangle baseline and its tick marks. + layout.triangle_axis.spines["bottom"].set_visible(False) + layout.triangle_axis.tick_params(axis="x", bottom=False, labelbottom=False) layout.triangle_axis.set_title( - display_sequence_name(title), fontsize=max(10, width * 1.4) + display_sequence_name(title), + fontsize=clamped_font_size(width, 1.4, MIN_TITLE_SIZE), + fontfamily=DEFAULT_FONT_FAMILY, ) if with_annotation: @@ -1472,6 +1551,23 @@ def _build_triangle_figure( top=0.90, hspace=0.05, ) + if with_annotation: + # ``set_aspect('equal', adjustable='box')`` narrows the triangle + # axis inside its GridSpec cell to retain 45-degree diagonals. A + # normal annotation axis keeps the full cell width, so shared data + # limits alone do not produce physical pixel alignment. Match the + # BED axis to the triangle's final horizontal bounds. + triangle_position = layout.triangle_axis.get_position() + annotation_position = layout.annotation_axis.get_position() + layout.annotation_axis.set_position( + [ + triangle_position.x0, + annotation_position.y0, + triangle_position.width, + annotation_position.height, + ] + ) + set_figure_font_family(layout.figure, DEFAULT_FONT_FAMILY) except Exception: plt.close(layout.figure) raise @@ -1613,6 +1709,7 @@ def _build_grid_figure( all_frames = list(single_frames.values()) + list(pair_frames.values()) axis_start, axis_end = _grid_axis_limits(all_frames, xlim) + _, axis_unit = genomic_scale(axis_end) axis_breaks = axes_label or breaks if not axis_breaks: axis_breaks = generate_breaks(int(axis_start), int(axis_end)) @@ -1621,10 +1718,21 @@ def _build_grid_figure( grid_size = len(names) figure_width = max(float(width), 2.0) - heading_size = max(6.0, min(12.0, figure_width * 1.2)) + heading_size = clamped_font_size(figure_width, 1.2, 8.0, 12.0) # Numeric genomic labels are intentionally twice the previous size. The # inverse grid-size factor keeps larger grids proportionate. - tick_size = max(8.0, min(18.0, figure_width * 3.0 / grid_size)) + tick_size = clamped_font_size( + figure_width, + 3.0 / grid_size, + MIN_TEXT_SIZE, + 18.0, + ) + axis_title_size = clamped_font_size( + figure_width, + 1.6 / grid_size, + MIN_TEXT_SIZE, + 14.0, + ) figure, axes = plt.subplots( grid_size, grid_size, @@ -1655,10 +1763,16 @@ def _build_grid_figure( ) if dataframe is not None and not dataframe.empty: + ( + dataframe, + direction_colors, + direction_column, + ) = _direction_ani_style(dataframe) draw_rectangular_tiles( axis, dataframe, - colors, + direction_colors or colors, + color_column=direction_column or "discrete", transpose=transpose, rasterized=not deraster, ) @@ -1674,22 +1788,37 @@ def _build_grid_figure( axis.grid(False) if row == 0: axis.set_title( - display_sequence_name(column_name), fontsize=heading_size + display_sequence_name(column_name), + fontsize=heading_size, + fontfamily=DEFAULT_FONT_FAMILY, ) if column == 0: axis.set_ylabel( - display_sequence_name(row_name), fontsize=heading_size + display_sequence_name(row_name), + fontsize=heading_size, + fontfamily=DEFAULT_FONT_FAMILY, ) + figure.supxlabel( + f"Genomic Position ({axis_unit})", + fontsize=axis_title_size, + fontfamily=DEFAULT_FONT_FAMILY, + ) + figure.supylabel( + f"Genomic Position ({axis_unit})", + fontsize=axis_title_size, + fontfamily=DEFAULT_FONT_FAMILY, + ) figure.subplots_adjust( - left=0.10, + left=0.14, right=0.98, - bottom=0.08, + bottom=0.12, top=0.92, wspace=0.08, hspace=0.08, ) _fit_grid_sequence_labels(figure, axes) + set_figure_font_family(figure, DEFAULT_FONT_FAMILY) except Exception: plt.close(figure) raise @@ -1734,7 +1863,14 @@ def create_grid( deraster=deraster, ) grid_size = axes.shape[0] - grid_prefix = os.path.join(directory, f"{grid_size}x{grid_size}_GRID") + directional = any( + "direction" in matrix.columns + if isinstance(matrix, pd.DataFrame) + else bool(matrix) and "direction" in matrix[0] + for matrix in [*singles, *doubles] + ) + grid_label = "DIRECTION_GRID" if directional else "GRID" + grid_prefix = os.path.join(directory, f"{grid_size}x{grid_size}_{grid_label}") print(f"\nGrid complete! Saving to {grid_prefix}...\n") try: save_figure_pair( @@ -1747,6 +1883,7 @@ def create_grid( finally: plt.close(figure) print("Grid saved successfully!\n") + return [f"{grid_prefix}.{vector_format}", f"{grid_prefix}.png"] def create_plots( @@ -1771,6 +1908,8 @@ def create_plots( deraster, annotation, ): + os.makedirs(directory, exist_ok=True) + created_files = [] df = read_df( sdf, palette, @@ -1781,6 +1920,7 @@ def create_plots( from_file, ) sdf = df + directional = "direction" in sdf.columns plot_filename = os.path.join(directory, name_x) @@ -1819,6 +1959,12 @@ def create_plots( vector_format=vector_format, ) if annotation_track_created: + created_files.extend( + [ + f"{iniprefix}_ANNOTATION_TRACK.{vector_format}", + f"{iniprefix}_ANNOTATION_TRACK.png", + ] + ) annotation_bed_df = bed_df annotation_chrom = chrom_name print(f"\nAnnotation track saved to {iniprefix}_ANNOTATION_TRACK\n") @@ -1848,41 +1994,55 @@ def create_plots( True, ) print(f"Creating plots and saving to {plot_filename}...\n") - ggsave( + full_suffix = "_DIRECTION_FULL" if directional else "_COMPARE" + hist_suffix = "_DIRECTION_HIST" if directional else "_COMPARE_HIST" + _save_plot( heatmap, width=width, height=width, dpi=dpi, format=vector_format, - filename=f"{plot_filename}_COMPARE.{vector_format}", + filename=f"{plot_filename}{full_suffix}.{vector_format}", verbose=False, ) - ggsave( + _save_plot( heatmap, width=width, height=width, dpi=dpi, format="png", - filename=f"{plot_filename}_COMPARE.png", + filename=f"{plot_filename}{full_suffix}.png", verbose=False, ) + created_files.extend( + [ + f"{plot_filename}{full_suffix}.{vector_format}", + f"{plot_filename}{full_suffix}.png", + ] + ) if not no_hist: - ggsave( + _save_plot( histy, width=3, height=3, dpi=dpi, format=vector_format, - filename=f"{plot_filename}_COMPARE_HIST.{vector_format}", + filename=f"{plot_filename}{hist_suffix}.{vector_format}", verbose=False, ) - ggsave( + created_files.extend( + [ + f"{plot_filename}{hist_suffix}.{vector_format}", + f"{plot_filename}{hist_suffix}.png", + ] + ) + _save_plot( histy, width=3, height=3, dpi=dpi, format="png", - filename=f"{plot_filename}_COMPARE_HIST.png", + filename=f"{plot_filename}{hist_suffix}.png", verbose=False, ) try: @@ -1890,15 +2050,16 @@ def create_plots( print( f"{plot_filename} comparative plots and histogram saved sucessfully. \n" ) - return 0 + return created_files except ValueError: print( f"{plot_filename} comparative plots and histogram saved sucessfully. \n" ) - return 0 + return created_files if no_hist: print( - f"{plot_filename}_COMPARE.{vector_format} and {plot_filename}_COMPARE.png saved sucessfully. \n" + f"{plot_filename}{full_suffix}.{vector_format} and " + f"{plot_filename}{full_suffix}.png saved sucessfully. \n" ) # Self-identity plots: Output _TRI, _FULL, and _HIST else: @@ -1920,25 +2081,34 @@ def create_plots( width, False, ) - ggsave( + full_suffix = "_DIRECTION_FULL" if directional else "_FULL" + tri_suffix = "_DIRECTION_TRI" if directional else "_TRI" + hist_suffix = "_DIRECTION_HIST" if directional else "_HIST" + _save_plot( full_plot, width=width, height=width, dpi=dpi, format=vector_format, - filename=f"{plot_filename}_FULL.{vector_format}", + filename=f"{plot_filename}{full_suffix}.{vector_format}", verbose=False, ) - ggsave( + created_files.extend( + [ + f"{plot_filename}{full_suffix}.{vector_format}", + f"{plot_filename}{full_suffix}.png", + ] + ) + _save_plot( full_plot, width=width, height=width, dpi=dpi, format="png", - filename=f"{plot_filename}_FULL.png", + filename=f"{plot_filename}{full_suffix}.png", verbose=False, ) - tri_prefix = f"{plot_filename}_TRI" + tri_prefix = f"{plot_filename}{tri_suffix}" triangle_figure = _build_triangle_figure( sdf=sdf, title=name_x, @@ -1960,6 +2130,7 @@ def create_plots( ) finally: plt.close(triangle_figure) + created_files.extend([f"{tri_prefix}.{vector_format}", f"{tri_prefix}.png"]) if annotation_track_created: annotated_figure = _build_triangle_figure( @@ -1985,30 +2156,43 @@ def create_plots( ) finally: plt.close(annotated_figure) + created_files.extend( + [ + f"{tri_prefix}_ANNOTATED.{vector_format}", + f"{tri_prefix}_ANNOTATED.png", + ] + ) if no_hist: print( f"Triangle plots and full plots for {plot_filename} saved sucessfully. \n" ) else: - ggsave( + _save_plot( histy, width=3, height=3, dpi=dpi, format=vector_format, - filename=plot_filename + f"_HIST.{vector_format}", + filename=plot_filename + f"{hist_suffix}.{vector_format}", verbose=False, ) - ggsave( + _save_plot( histy, width=3, height=3, dpi=dpi, format="png", - filename=plot_filename + "_HIST.png", + filename=plot_filename + f"{hist_suffix}.png", verbose=False, ) + created_files.extend( + [ + plot_filename + f"{hist_suffix}.{vector_format}", + plot_filename + f"{hist_suffix}.png", + ] + ) print( f"Triangle plots, full plots, and histogram for {plot_filename} saved sucessfully. \n" ) + return created_files diff --git a/tests/test_annotation_track.py b/tests/test_annotation_track.py index a58ad36..ee42f01 100644 --- a/tests/test_annotation_track.py +++ b/tests/test_annotation_track.py @@ -7,8 +7,10 @@ import pytest import moddotplot.static_plots as static_plots +from moddotplot.const import DIRECTION_COLORS from moddotplot.static_plots import ( DEFAULT_ANNOTATION_COLOR, + _build_triangle_figure, draw_annotation_track, read_annotation_bed, render_annotation_track, @@ -120,6 +122,44 @@ def test_render_annotation_track_writes_expected_svg_and_png(tmp_path): ET.parse(svg_path) +def test_triangle_uses_blue_and_pink_direction_colors(): + dataframe = pd.DataFrame( + { + "q": ["chr1"] * 4, + "q_st": [10, 20, 30, 40], + "q_en": [20, 30, 40, 50], + "r": ["chr1"] * 4, + "r_st": [10, 40, 60, 80], + "r_en": [20, 50, 70, 90], + "discrete": pd.Categorical([0, 1, 0, 1], categories=[0, 1]), + "direction": ["Forward", "Forward", "Reverse", "Reverse"], + } + ) + + figure = _build_triangle_figure( + sdf=dataframe, + title="chr1", + palette="Spectral_11", + palette_orientation="+", + custom_colors=["#000000", "#ffffff"], + axes_labels=[0, 50, 100], + xlim=(0, 100), + deraster=True, + width=4, + ) + try: + rendered = { + tuple(color) + for collection in figure.axes[0].collections + for color in collection.get_facecolors() + } + assert to_rgba(DIRECTION_COLORS["Forward"]) in rendered + assert to_rgba(DIRECTION_COLORS["Reverse"]) in rendered + assert len(rendered) == 4 + finally: + plt.close(figure) + + def test_render_annotation_track_skips_empty_overlap(tmp_path): bed_path = tmp_path / "annotations.bed" bed_path.write_text("chr2\t10\t20\n") @@ -141,12 +181,87 @@ def test_render_annotation_track_skips_empty_overlap(tmp_path): assert not (tmp_path / "sample_ANNOTATION_TRACK.png").exists() +def test_annotated_triangle_physically_aligns_axes_and_hides_heatmap_baseline( + tmp_path, +): + bed_path = tmp_path / "annotations.bed" + bed_path.write_text("chr1\t0\t100\n") + dataframe = pd.DataFrame( + { + "q_st": [0], + "q_en": [100], + "r_st": [0], + "r_en": [100], + "discrete": [0], + } + ) + + figure = _build_triangle_figure( + sdf=dataframe, + title="chr1:1-100", + palette="Spectral_11", + palette_orientation="+", + custom_colors=None, + axes_labels=None, + xlim=(0, 100), + deraster=False, + width=6, + annotation_df=read_annotation_bed(bed_path), + annotation_chrom="chr1", + ) + try: + triangle_axis, annotation_axis = figure.axes + triangle_position = triangle_axis.get_position() + annotation_position = annotation_axis.get_position() + + assert annotation_position.x0 == pytest.approx(triangle_position.x0) + assert annotation_position.width == pytest.approx(triangle_position.width) + assert not triangle_axis.spines["bottom"].get_visible() + assert not any( + tick.tick1line.get_visible() for tick in triangle_axis.xaxis.majorTicks + ) + assert annotation_axis.spines["bottom"].get_visible() + assert annotation_axis.get_xlabel() == "Genomic Position (Kbp)" + finally: + plt.close(figure) + + +def test_unannotated_triangle_keeps_its_genomic_axis(): + dataframe = pd.DataFrame( + { + "q_st": [0], + "q_en": [100], + "r_st": [0], + "r_en": [100], + "discrete": [0], + } + ) + + figure = _build_triangle_figure( + sdf=dataframe, + title="chr1", + palette="Spectral_11", + palette_orientation="+", + custom_colors=None, + axes_labels=None, + xlim=(0, 100), + deraster=False, + width=6, + ) + try: + triangle_axis = figure.axes[0] + assert triangle_axis.spines["bottom"].get_visible() + assert triangle_axis.get_xlabel() == "Genomic Position (Kbp)" + finally: + plt.close(figure) + + class _FakePlot: def __add__(self, _other): return self -def _stub_create_plots_dependencies(monkeypatch): +def _stub_create_plots_dependencies(monkeypatch, *, directional=False): dataframe = pd.DataFrame( [ { @@ -161,6 +276,8 @@ def _stub_create_plots_dependencies(monkeypatch): } ] ) + if directional: + dataframe["direction"] = ["Forward"] monkeypatch.setattr(static_plots, "read_df", lambda *_args, **_kwargs: dataframe) monkeypatch.setattr( static_plots, "make_hist", lambda *_args, **_kwargs: _FakePlot() @@ -178,7 +295,7 @@ def fake_ggsave(*_args, **kwargs): def _run_create_plots(output_dir, annotation, vector_format="svg"): - static_plots.create_plots( + return static_plots.create_plots( sdf=None, directory=str(output_dir), name_x="chr1", @@ -209,25 +326,63 @@ def test_create_plots_creates_and_skips_annotation_artifacts(tmp_path, monkeypat matching_dir.mkdir() matching_bed = tmp_path / "matching.bed" matching_bed.write_text("chr1\t10\t20\n") - _run_create_plots(matching_dir, matching_bed) + matching_outputs = _run_create_plots(matching_dir, matching_bed) assert (matching_dir / "chr1_ANNOTATION_TRACK.svg").exists() assert (matching_dir / "chr1_ANNOTATION_TRACK.png").exists() assert (matching_dir / "chr1_TRI_ANNOTATED.svg").exists() assert (matching_dir / "chr1_TRI_ANNOTATED.png").exists() assert not (matching_dir / "chr1_PRE_ANNOTATED.svg").exists() + assert str(matching_dir / "chr1_TRI_ANNOTATED.svg") in matching_outputs + assert str(matching_dir / "chr1_ANNOTATION_TRACK.svg") in matching_outputs empty_dir = tmp_path / "empty" empty_dir.mkdir() nonmatching_bed = tmp_path / "nonmatching.bed" nonmatching_bed.write_text("chr2\t10\t20\n") - _run_create_plots(empty_dir, nonmatching_bed) + empty_outputs = _run_create_plots(empty_dir, nonmatching_bed) assert not (empty_dir / "chr1_ANNOTATION_TRACK.svg").exists() assert not (empty_dir / "chr1_ANNOTATION_TRACK.png").exists() assert not (empty_dir / "chr1_PRE_ANNOTATED.svg").exists() assert not (empty_dir / "chr1_TRI_ANNOTATED.svg").exists() assert not (empty_dir / "chr1_TRI_ANNOTATED.png").exists() + assert not any("ANNOTATION" in output for output in empty_outputs) + + +def test_create_plots_uses_direction_output_names(tmp_path, monkeypatch): + _stub_create_plots_dependencies(monkeypatch, directional=True) + + outputs = static_plots.create_plots( + sdf=None, + directory=str(tmp_path / "directionality"), + name_x="chr1", + name_y="chr1", + palette="Spectral_11", + palette_orientation="+", + no_hist=False, + width=4, + dpi=72, + is_freq=False, + xlim=100, + custom_colors=None, + custom_breakpoints=None, + from_file=None, + is_pairwise=False, + axes_labels=None, + axes_tick_number=7, + vector_format="svg", + deraster=True, + annotation=None, + ) + + expected_stems = { + "chr1_DIRECTION_FULL", + "chr1_DIRECTION_TRI", + "chr1_DIRECTION_HIST", + } + assert {Path(output).stem for output in outputs} == expected_stems + assert all(Path(output).parent.name == "directionality" for output in outputs) @pytest.mark.parametrize( diff --git a/tests/test_cli_integration.py b/tests/test_cli_integration.py index d23add6..f8578a7 100644 --- a/tests/test_cli_integration.py +++ b/tests/test_cli_integration.py @@ -73,6 +73,30 @@ def test_static_cli_computes_all_self_and_pairwise_outputs(tmp_path): assert all(path.stat().st_size > 0 for path in output.rglob("*.bedpe")) +def test_omitted_subcommand_runs_static_mode(tmp_path): + fasta = tmp_path / "one.fa" + fasta.write_text(">alpha\n" + "ACGT" * 300 + "\n") + output = tmp_path / "implicit-static" + + result = _run_cli( + "--fasta", + fasta, + "--window", + 100, + "--modimizer", + 10, + "--identity", + 80, + "--no-plot", + "--output-dir", + output, + ) + + assert result.returncode == 0, result.stderr + result.stdout + assert "Running ModDotPlot in static mode" in result.stdout + assert (output / "alpha" / "alpha.bedpe").is_file() + + def test_static_grid_regions_with_dotted_headers_crop_every_output(tmp_path): names = [ "PAN010.chr14.haplotype1.paternal", @@ -189,3 +213,4 @@ def test_interactive_cli_forward_mode_saves_matrix_without_launching_server(tmp_ assert (saved / "alpha_0.npz").is_file() assert (saved / "metadata.pkl").is_file() assert "Saved matrices" in result.stdout + assert "interactive mode is deprecated and maintenance-only" in result.stderr diff --git a/tests/test_cli_runtime.py b/tests/test_cli_runtime.py index b70f787..7eeade2 100644 --- a/tests/test_cli_runtime.py +++ b/tests/test_cli_runtime.py @@ -1,4 +1,5 @@ import sys +import shlex import numpy as np import pytest @@ -29,7 +30,7 @@ def create_pairwise(*args): monkeypatch.setattr(cli, "create_plots", lambda **kwargs: plot_calls.append(kwargs)) -def test_main_without_subcommand_reports_parser_error(monkeypatch, capsys): +def test_main_without_arguments_defaults_to_static_parser(monkeypatch, capsys): monkeypatch.setattr(sys, "argv", ["moddotplot"]) with pytest.raises(SystemExit) as exc_info: @@ -37,8 +38,25 @@ def test_main_without_subcommand_reports_parser_error(monkeypatch, capsys): captured = capsys.readouterr() assert exc_info.value.code == 2 - assert "the following arguments are required: command" in captured.err - assert "{interactive,static}" in captured.err + assert "moddotplot static" in captured.err + assert "one of the arguments -c/--config -l/--load -f/--fasta is required" in ( + captured.err + ) + assert "the following arguments are required: command" not in captured.err + + +def test_parser_defaults_omitted_subcommand_to_static(): + args = cli.parse_args(["--fasta", "sequence.fa", "--no-plot"]) + + assert args.command == "static" + assert args.fasta == ["sequence.fa"] + assert args.no_plot + + +def test_parser_preserves_explicit_interactive_subcommand(): + args = cli.parse_args(["interactive", "--fasta", "sequence.fa"]) + + assert args.command == "interactive" @pytest.mark.parametrize("option", ["--colors", "--color"]) @@ -129,6 +147,78 @@ def test_interactive_window_uses_longest_sequence_by_length(monkeypatch): assert {entry["max_window_size"] for entry in metadata} == {102} +def test_interactive_load_passes_combined_beds_to_dash(monkeypatch, tmp_path): + first_bed = tmp_path / "first.bed" + second_bed = tmp_path / "second.bed" + first_bed.write_text("chrA\t1010\t1020\n") + second_bed.write_text("chrB\t30\t40\n") + metadata = [ + { + "x_name": "chrA:1001-1100", + "y_name": "chrB", + "x_size": 100, + "y_size": 100, + "self": False, + "max_window_size": 50, + "resolution": 2, + } + ] + monkeypatch.setattr( + cli, "extractFiles", lambda _path: ([[np.ones((2, 2))]], metadata) + ) + dash_calls = [] + monkeypatch.setattr(cli, "run_dash", lambda *args: dash_calls.append(args)) + monkeypatch.setattr( + sys, + "argv", + [ + "moddotplot", + "interactive", + "--load", + "saved-matrices", + "--bed", + str(first_bed), + str(second_bed), + ], + ) + + with pytest.raises(SystemExit) as exc_info: + cli.main() + + assert exc_info.value.code == 0 + x_axis, y_axis = dash_calls[0][2][0] + assert (x_axis[0], x_axis[-1]) == (1001, 1100) + assert (y_axis[0], y_axis[-1]) == (0, 100) + annotations = dash_calls[0][7] + assert annotations[["chrom", "start", "end"]].to_dict("records") == [ + {"chrom": "chrA", "start": 1010, "end": 1020}, + {"chrom": "chrB", "start": 30, "end": 40}, + ] + + +def test_interactive_rejects_invalid_annotation_bed(monkeypatch, tmp_path, capsys): + bed = tmp_path / "invalid.bed" + bed.write_text("chrA\tnot-a-coordinate\t20\n") + monkeypatch.setattr( + sys, + "argv", + [ + "moddotplot", + "interactive", + "--load", + "saved-matrices", + "--bed", + str(bed), + ], + ) + + with pytest.raises(SystemExit) as exc_info: + cli.main() + + assert exc_info.value.code == 2 + assert "Error reading annotation BED file(s)" in capsys.readouterr().err + + def test_no_bedpe_self_plot_uses_sequence_output_directory(monkeypatch, tmp_path): _patch_fasta_input(monkeypatch, ["chrA"], [[1] * 1000]) plot_calls = [] @@ -157,6 +247,58 @@ def test_no_bedpe_self_plot_uses_sequence_output_directory(monkeypatch, tmp_path assert not list(tmp_path.rglob("*.bedpe")) +def test_static_plot_directory_gets_reproducibility_summary(monkeypatch, tmp_path): + fasta = tmp_path / "source genome.fa" + fasta.touch() + _patch_fasta_input(monkeypatch, ["chrA"], [[1] * 1000]) + monkeypatch.setattr(cli, "createSelfMatrix", lambda *_args: np.full((1, 1), 100.0)) + monkeypatch.setattr( + cli, + "convertMatrixToBed", + lambda *_args, **_kwargs: [["header"], ["value"]], + ) + + def create_plot_files(**kwargs): + prefix = tmp_path / "chrA:101-400" / "chrA:101-400" + created = [] + for suffix in ("_FULL.svg", "_FULL.png", "_TRI.svg", "_TRI.png"): + path = prefix.parent / f"{prefix.name}{suffix}" + path.touch() + created.append(str(path)) + return created + + monkeypatch.setattr(cli, "create_plots", create_plot_files) + monkeypatch.setattr( + sys, + "argv", + [ + "moddotplot", + "-f", + str(fasta), + "--region", + "chrA:101-400", + "--resolution", + "10", + "--no-bedpe", + "--no-hist", + "--output-dir", + str(tmp_path), + ], + ) + + cli.main() + + summary = (tmp_path / "chrA:101-400" / "plot_summary.txt").read_text() + assert f"Command: {shlex.join(sys.argv)}" in summary + assert str(fasta.resolve()) in summary + assert "Window sizes:\n - 28 bp" in summary + assert "Regions:\n - chrA:101-400" in summary + assert "BED annotation file: None" in summary + assert ( + str((tmp_path / "chrA:101-400" / "chrA:101-400_TRI.svg").resolve()) in summary + ) + + def test_unmatched_region_fails_instead_of_silently_using_full_sequences( monkeypatch, tmp_path, capsys ): diff --git a/tests/test_direction_cli.py b/tests/test_direction_cli.py index 451eaea..93d254e 100644 --- a/tests/test_direction_cli.py +++ b/tests/test_direction_cli.py @@ -1,4 +1,5 @@ import sys +from pathlib import Path import numpy as np @@ -24,13 +25,61 @@ def read_hashes( return hash_sets[forward_only] monkeypatch.setattr(cli, "readKmersFromFile", read_hashes) - monkeypatch.setattr( - cli, - "convertMatrixToBed", - lambda *_args, **_kwargs: [["header"], ["value"]], - ) - monkeypatch.setattr(cli, "create_plots", lambda **_kwargs: None) - monkeypatch.setattr(cli, "create_grid", lambda **_kwargs: None) + + def convert_matrix_to_bed( + matrix, + window_size, + _identity, + name_x, + name_y, + self_identity, + x_offset, + y_offset, + *_args, + ): + rows = [ + ( + "#query_name", + "query_start", + "query_end", + "reference_name", + "reference_start", + "reference_end", + "perID_by_events", + ) + ] + for query_index, reference_index in np.argwhere(matrix > 0): + if self_identity and query_index > reference_index: + continue + query_start = query_index * window_size + x_offset + reference_start = reference_index * window_size + y_offset + rows.append( + ( + name_x, + query_start, + query_start + window_size - 1, + name_y, + reference_start, + reference_start + window_size - 1, + matrix[query_index, reference_index], + ) + ) + return rows + + monkeypatch.setattr(cli, "convertMatrixToBed", convert_matrix_to_bed) + plot_calls = [] + grid_calls = [] + + def capture_plot(**kwargs): + plot_calls.append(kwargs) + return [str(Path(kwargs["directory"]) / "mock_plot.png")] + + def capture_grid(**kwargs): + grid_calls.append(kwargs) + return [str(Path(kwargs["directory"]) / "mock_grid.png")] + + monkeypatch.setattr(cli, "create_plots", capture_plot) + monkeypatch.setattr(cli, "create_grid", capture_grid) monkeypatch.setattr(cli, "read_df_from_file", lambda _path: None) monkeypatch.setattr( sys, @@ -48,14 +97,13 @@ def read_hashes( str(tmp_path), ], ) + return plot_calls, grid_calls -def test_static_direction_plot_receives_canonical_and_forward_self_matrices( - monkeypatch, tmp_path -): +def test_static_direction_colors_standard_self_plot_rows(monkeypatch, tmp_path): canonical_hashes = [[1] * 1000] forward_hashes = [[2] * 1000] - _patch_common_static_io( + plot_calls, _grid_calls = _patch_common_static_io( monkeypatch, ["chrA"], {False: canonical_hashes, True: forward_hashes}, @@ -68,27 +116,58 @@ def create_self(_length, sequence, *_args): return canonical_matrix if sequence is canonical_hashes[0] else forward_matrix monkeypatch.setattr(cli, "createSelfMatrix", create_self) - direction_calls = [] - monkeypatch.setattr( - cli, - "create_direction_plot", - lambda **kwargs: direction_calls.append(kwargs), + cli.main() + + assert len(plot_calls) == 2 + assert "direction" not in plot_calls[0]["sdf"][0][0] + assert plot_calls[1]["directory"].endswith("directionality") + direction_bed = plot_calls[1]["sdf"][0] + assert direction_bed[0][-1] == "direction" + assert [row[-1] for row in direction_bed[1:]] == [ + "Forward", + "Reverse", + "Forward", + ] + assert plot_calls[1]["name_x"] == plot_calls[1]["name_y"] == "chrA" + direction_summary = tmp_path / "chrA" / "directionality" / "plot_summary.txt" + assert direction_summary.is_file() + assert "Plot group" not in direction_summary.read_text() + + +def test_direction_coloring_with_forward_option_keeps_reverse_hits( + monkeypatch, tmp_path +): + canonical_hashes = [[1] * 1000] + forward_hashes = [[2] * 1000] + plot_calls, _grid_calls = _patch_common_static_io( + monkeypatch, + ["chrA"], + {False: canonical_hashes, True: forward_hashes}, + tmp_path, ) + sys.argv.append("--forward") + canonical_matrix = np.array([[100.0, 90.0], [90.0, 100.0]]) + forward_matrix = np.array([[100.0, 0.0], [0.0, 100.0]]) + + def create_self(_length, sequence, *_args): + return canonical_matrix if sequence is canonical_hashes[0] else forward_matrix + monkeypatch.setattr(cli, "createSelfMatrix", create_self) cli.main() - assert len(direction_calls) == 1 - call = direction_calls[0] - assert call["canonical_matrix"] is canonical_matrix - assert call["forward_matrix"] is forward_matrix - assert call["self_identity"] is True - assert call["name_x"] == call["name_y"] == "chrA" + assert len(plot_calls) == 2 + direction_bed = plot_calls[1]["sdf"][0] + assert [row[-1] for row in direction_bed[1:]] == [ + "Forward", + "Reverse", + "Forward", + ] -def test_static_direction_plot_receives_pairwise_matrices(monkeypatch, tmp_path): +def test_static_direction_colors_standard_pairwise_plot_rows(monkeypatch, tmp_path): canonical_hashes = [[1] * 1000, [2] * 800] forward_hashes = [[3] * 1000, [4] * 800] - _patch_common_static_io( + plot_calls, _grid_calls = _patch_common_static_io( monkeypatch, ["chrA", "chrB"], {False: canonical_hashes, True: forward_hashes}, @@ -102,18 +181,56 @@ def create_pair(_y_length, _x_length, y_sequence, _x_sequence, *_args): return canonical_matrix if y_sequence is canonical_hashes[1] else forward_matrix monkeypatch.setattr(cli, "createPairwiseMatrix", create_pair) - direction_calls = [] - monkeypatch.setattr( - cli, - "create_direction_plot", - lambda **kwargs: direction_calls.append(kwargs), + cli.main() + + assert len(plot_calls) == 2 + assert "direction" not in plot_calls[0]["sdf"][0][0] + assert plot_calls[1]["directory"].endswith("directionality") + direction_bed = plot_calls[1]["sdf"][0] + assert direction_bed[0][-1] == "direction" + assert [row[-1] for row in direction_bed[1:]] == ["Reverse"] + assert (plot_calls[1]["name_x"], plot_calls[1]["name_y"]) == ("chrA", "chrB") + + +def test_grid_only_receives_direction_colored_self_and_pairwise_rows( + monkeypatch, tmp_path +): + canonical_hashes = [[1] * 1000, [2] * 800] + forward_hashes = [[3] * 1000, [4] * 800] + _plot_calls, grid_calls = _patch_common_static_io( + monkeypatch, + ["chrA", "chrB"], + {False: canonical_hashes, True: forward_hashes}, + tmp_path, ) + sys.argv.extend(["--grid-only"]) + monkeypatch.setattr(cli, "ModimizerSketchCache", lambda **_kwargs: None) + + def create_self(_length, sequence, *_args): + if sequence in canonical_hashes: + return np.array([[100.0, 95.0], [95.0, 100.0]]) + return np.array([[100.0, 0.0], [0.0, 100.0]]) + + def create_pair(_y_length, _x_length, y_sequence, _x_sequence, *_args): + if y_sequence in canonical_hashes: + return np.array([[95.0]]) + return np.array([[0.0]]) + + monkeypatch.setattr(cli, "createSelfMatrix", create_self) + monkeypatch.setattr(cli, "createPairwiseMatrix", create_pair) cli.main() - assert len(direction_calls) == 1 - call = direction_calls[0] - assert call["canonical_matrix"] is canonical_matrix - assert call["forward_matrix"] is forward_matrix - assert call["self_identity"] is False - assert (call["name_x"], call["name_y"]) == ("chrA", "chrB") + assert len(grid_calls) == 2 + assert all("direction" not in matrix[0] for matrix in grid_calls[0]["singles"]) + assert all("direction" not in matrix[0] for matrix in grid_calls[0]["doubles"]) + assert grid_calls[1]["directory"].endswith("directionality") + grid_call = grid_calls[1] + assert all(matrix[0][-1] == "direction" for matrix in grid_call["singles"]) + assert all(matrix[0][-1] == "direction" for matrix in grid_call["doubles"]) + assert any( + row[-1] == "Reverse" + for matrix in [*grid_call["singles"], *grid_call["doubles"]] + for row in matrix[1:] + ) + assert (tmp_path / "directionality" / "plot_summary.txt").is_file() diff --git a/tests/test_entrypoints.py b/tests/test_entrypoints.py index b764478..8d12f06 100644 --- a/tests/test_entrypoints.py +++ b/tests/test_entrypoints.py @@ -44,11 +44,19 @@ def test_module_help_is_available(): result = _run_module("--help") assert result.returncode == 0 - assert "{interactive,static}" in result.stdout + assert "{static,interactive}" in result.stdout + assert "static is used when omitted" in result.stdout + assert "Static mode commands (default)" in result.stdout + assert "Interactive mode commands (deprecated; explicit use" in result.stdout + assert "only)" in result.stdout -def test_module_without_subcommand_has_clean_usage_error(): +def test_module_without_arguments_defaults_to_static_parser(): result = _run_module() assert result.returncode == 2 - assert "the following arguments are required: command" in result.stderr + assert "static:" in result.stderr + assert "one of the arguments -c/--config -l/--load -f/--fasta is required" in ( + result.stderr + ) + assert "the following arguments are required: command" not in result.stderr diff --git a/tests/test_grid.py b/tests/test_grid.py index 95a1bfa..b692cf0 100644 --- a/tests/test_grid.py +++ b/tests/test_grid.py @@ -1,8 +1,10 @@ import matplotlib.pyplot as plt from matplotlib.colors import to_rgba import numpy as np +from pathlib import Path import pytest +from moddotplot.const import DIRECTION_COLORS from moddotplot.static_plots import _build_grid_figure, create_grid @@ -103,7 +105,7 @@ def _basic_two_sequence_grid(**overrides): def _create_grid(tmp_path, *, vector_format="svg", **kwargs): - create_grid( + return create_grid( directory=tmp_path, vector_format=vector_format, dpi=72, @@ -183,7 +185,7 @@ def test_create_grid_handles_empty_pairwise_comparison(tmp_path): def test_create_grid_writes_png_and_valid_selected_vector_format( tmp_path, vector_format, signature ): - _create_grid( + output_files = _create_grid( tmp_path, vector_format=vector_format, **_basic_two_sequence_grid(), @@ -191,6 +193,7 @@ def test_create_grid_writes_png_and_valid_selected_vector_format( png = tmp_path / "2x2_GRID.png" vector = tmp_path / f"2x2_GRID.{vector_format}" + assert output_files == [str(vector), str(png)] assert png.read_bytes().startswith(b"\x89PNG\r\n\x1a\n") if vector_format == "svg": assert signature in vector.read_bytes()[:1024] @@ -198,6 +201,26 @@ def test_create_grid_writes_png_and_valid_selected_vector_format( assert vector.read_bytes().startswith(signature) +def test_create_grid_uses_direction_output_name(tmp_path): + kwargs = _basic_two_sequence_grid() + direction_header = (*BED_HEADER, "direction") + kwargs["singles"] = [ + [direction_header] + [(*row, "Forward") for row in records[1:]] + for records in kwargs["singles"] + ] + kwargs["doubles"] = [ + [direction_header] + [(*row, "Reverse") for row in records[1:]] + for records in kwargs["doubles"] + ] + + outputs = _create_grid(tmp_path / "directionality", **kwargs) + + assert {Path(output).name for output in outputs} == { + "2x2_DIRECTION_GRID.svg", + "2x2_DIRECTION_GRID.png", + } + + def test_reversed_pair_metadata_orients_asymmetric_coordinates_by_grid_cell(): # The BED record is B (query) vs A (reference), while the requested grid is # ordered A, B across columns and B, A down rows. Pairwise comparisons sit @@ -341,6 +364,26 @@ def test_grid_exact_bounds_do_not_expand_to_next_nice_tick(): plt.close(figure) +@pytest.mark.parametrize( + ("axis_end", "unit"), + [(100_000, "Kbp"), (103_156_783, "Mbp"), (500_000_000, "Gbp")], +) +def test_grid_labels_both_axes_with_genomic_units(axis_end, unit): + kwargs = _basic_two_sequence_grid( + xlim=(1, axis_end), + axes_label=None, + breaks=None, + ) + + figure, _axes = _build_grid_figure(**kwargs) + try: + expected = f"Genomic Position ({unit})" + assert figure._supxlabel.get_text() == expected + assert figure._supylabel.get_text() == expected + finally: + plt.close(figure) + + def test_compare_only_grid_derives_sequence_names_from_pair_metadata(tmp_path): double_names = [ ["gamma", "alpha"], @@ -460,6 +503,42 @@ def test_grid_honors_custom_breakpoints_and_colors(): plt.close(figure) +def test_grid_uses_blue_and_pink_direction_colors_in_every_cell(): + kwargs = _basic_two_sequence_grid() + direction_header = (*BED_HEADER, "direction") + + def add_directions(records): + return [direction_header] + [ + (*row, "Forward" if index % 2 == 0 else "Reverse") + for index, row in enumerate(records[1:]) + ] + + kwargs["singles"] = [add_directions(records) for records in kwargs["singles"]] + pair = _records( + "sequence_a", + "sequence_b", + [(10, 10, 86), (30, 40, 100), (50, 60, 86), (70, 80, 100)], + ) + kwargs["doubles"] = [ + [direction_header] + + [ + (*row, direction) + for row, direction in zip( + pair[1:], ["Forward", "Forward", "Reverse", "Reverse"] + ) + ] + ] + + figure, axes = _build_grid_figure(**kwargs) + try: + rendered_colors = _all_artist_colors(axes) + assert to_rgba(DIRECTION_COLORS["Forward"]) in rendered_colors + assert to_rgba(DIRECTION_COLORS["Reverse"]) in rendered_colors + assert len(rendered_colors) >= 4 + finally: + plt.close(figure) + + @pytest.mark.parametrize( ("deraster", "expected_rasterized"), [(False, True), (True, False)] ) diff --git a/tests/test_interactive_annotations.py b/tests/test_interactive_annotations.py new file mode 100644 index 0000000..e646804 --- /dev/null +++ b/tests/test_interactive_annotations.py @@ -0,0 +1,192 @@ +import numpy as np +import plotly.graph_objs as go +import dash + +from moddotplot.annotations import read_annotation_beds +from moddotplot.interactive import ( + add_annotation_tracks, + interactive_axis_bounds, + preserve_zoom_ranges, + run_dash, +) + + +def _figure(): + figure = go.Figure(data=[go.Heatmap(z=[[100.0]])]) + figure.update_xaxes(title_text="chrX") + figure.update_yaxes(title_text="chrY") + return figure + + +def test_multiple_beds_add_tracks_to_both_comparative_axes(tmp_path): + x_bed = tmp_path / "x.bed" + y_bed = tmp_path / "y.bed" + x_bed.write_text("chrX\t10\t30\tx-feature\t0\t+\t10\t30\t255,0,0\n") + y_bed.write_text("chrY\t40\t80\ty-feature\t0\t+\t40\t80\t0,0,255\n") + annotations = read_annotation_beds([x_bed, y_bed]) + metadata = { + "x_name": "chrX", + "y_name": "chrY", + "x_size": 100, + "y_size": 120, + "self": False, + } + + figure = add_annotation_tracks(_figure(), metadata, annotations) + + x_shapes = [ + shape + for shape in figure.layout.shapes + if shape.xref == "x" and shape.yref == "paper" + ] + y_shapes = [ + shape + for shape in figure.layout.shapes + if shape.xref == "paper" and shape.yref == "y" + ] + assert len(x_shapes) == 2 # track background plus one interval + assert len(y_shapes) == 2 + assert (x_shapes[1].x0, x_shapes[1].x1) == (10, 30) + assert x_shapes[1].fillcolor == "rgb(255,0,0)" + assert (y_shapes[1].y0, y_shapes[1].y1) == (40, 80) + assert y_shapes[1].fillcolor == "rgb(0,0,255)" + assert figure.layout.margin.b >= 140 + assert figure.layout.margin.l >= 140 + + +def test_self_plot_adds_only_x_track_and_clips_to_sequence(tmp_path): + bed = tmp_path / "annotations.bed" + bed.write_text("chrX\t90\t130\nchrY\t10\t20\n") + annotations = read_annotation_beds([bed]) + metadata = { + "x_name": "chrX", + "y_name": "chrX", + "x_size": 100, + "y_size": 100, + "self": True, + } + + figure = add_annotation_tracks(_figure(), metadata, annotations) + + shapes = list(figure.layout.shapes) + assert len(shapes) == 2 + assert all(shape.xref == "x" and shape.yref == "paper" for shape in shapes) + assert (shapes[1].x0, shapes[1].x1) == (90, 100) + + +def test_region_suffixed_header_matches_base_bed_chromosome(tmp_path): + bed = tmp_path / "annotations.bed" + bed.write_text("chr14_MATERNAL\t14002000\t14004000\n") + annotations = read_annotation_beds([bed]) + metadata = { + "x_name": "chr14_MATERNAL:14000001-18000000", + "y_name": "chr14_MATERNAL:14000001-18000000", + "x_size": 4_000_000, + "y_size": 4_000_000, + "self": True, + } + + figure = add_annotation_tracks(_figure(), metadata, annotations) + + shapes = list(figure.layout.shapes) + assert interactive_axis_bounds(metadata, "x") == (14_000_001, 18_000_000) + assert len(shapes) == 2 + assert (shapes[0].x0, shapes[0].x1) == (14_000_001, 18_000_000) + assert (shapes[1].x0, shapes[1].x1) == (14_002_000, 14_004_000) + + +def test_nonmatching_bed_does_not_add_tracks(tmp_path): + bed = tmp_path / "annotations.bed" + bed.write_text("other\t10\t20\n") + annotations = read_annotation_beds([bed]) + metadata = { + "x_name": "chrX", + "y_name": "chrY", + "x_size": 100, + "y_size": 100, + "self": False, + } + + figure = add_annotation_tracks(_figure(), metadata, annotations) + + assert not figure.layout.shapes + + +def test_bed_shapes_do_not_override_explicit_zoom_range(tmp_path): + bed = tmp_path / "annotations.bed" + bed.write_text("chrX\t0\t1000000\n") + annotations = read_annotation_beds([bed]) + metadata = { + "x_name": "chrX", + "y_name": "chrX", + "x_size": 1_000_000, + "y_size": 1_000_000, + "self": True, + } + figure = add_annotation_tracks(_figure(), metadata, annotations) + + preserve_zoom_ranges( + figure, + x_range=(200_000, 300_000), + y_range=(400_000, 500_000), + ) + + assert tuple(figure.layout.xaxis.range) == (200_000, 300_000) + assert tuple(figure.layout.yaxis.range) == (400_000, 500_000) + assert figure.layout.xaxis.autorange is False + assert figure.layout.yaxis.autorange is False + assert figure.layout.shapes[0].x0 == 0 + assert figure.layout.shapes[0].x1 == 1_000_000 + + +def _find_component(component, component_id): + if getattr(component, "id", None) == component_id: + return component + children = getattr(component, "children", None) + if children is None: + return None + if not isinstance(children, (list, tuple)): + children = [children] + for child in children: + match = _find_component(child, component_id) + if match is not None: + return match + return None + + +def test_dash_initial_comparative_figure_contains_both_tracks(monkeypatch, tmp_path): + bed = tmp_path / "annotations.bed" + bed.write_text("chrX\t10\t30\nchrY\t40\t80\n") + annotations = read_annotation_beds([bed]) + metadata = [ + { + "x_name": "chrX", + "y_name": "chrY", + "x_size": 100, + "y_size": 120, + "self": False, + "min_window_size": 50, + "max_window_size": 50, + "resolution": 2, + "title": "chrX-chrY", + "sparsities": [1], + } + ] + apps = [] + monkeypatch.setattr(dash.Dash, "run", lambda app, **_kwargs: apps.append(app)) + + run_dash( + [[np.array([[100.0, 90.0], [90.0, 100.0]])]], + metadata, + [[[0, 50, 100], [0, 60, 120]]], + sparsity=1, + identity=86.0, + port_number=8050, + output_dir=None, + annotations=annotations, + ) + + graph = _find_component(apps[0].layout, "dotplot") + shapes = list(graph.figure.layout.shapes) + assert any(shape.xref == "x" and shape.yref == "paper" for shape in shapes) + assert any(shape.xref == "paper" and shape.yref == "y" for shape in shapes) diff --git a/tests/test_interactive_parser.py b/tests/test_interactive_parser.py index 264c69b..9bc6133 100644 --- a/tests/test_interactive_parser.py +++ b/tests/test_interactive_parser.py @@ -19,3 +19,18 @@ def test_interactive_accepts_forward_flag(): ) assert args.forward is True + + +def test_interactive_accepts_multiple_annotation_beds(): + args = get_parser().parse_args( + [ + "interactive", + "--fasta", + "sequence.fa", + "--bed", + "first.bed", + "second.bed", + ] + ) + + assert args.bed == ["first.bed", "second.bed"] diff --git a/tests/test_packaging_metadata.py b/tests/test_packaging_metadata.py index fa7a0d7..7698f0d 100644 --- a/tests/test_packaging_metadata.py +++ b/tests/test_packaging_metadata.py @@ -3,7 +3,7 @@ try: import tomllib except ModuleNotFoundError: # pragma: no cover - # Exercised by the Python 3.8-3.10 CI jobs. + # Retained for source-tree tooling that may run outside supported Python. import tomli as tomllib from moddotplot.const import VERSION @@ -22,15 +22,15 @@ def test_runtime_and_distribution_versions_match(): def test_declared_python_floor_matches_documentation(): - assert project_metadata()["requires-python"] == ">=3.8,<3.13" + assert project_metadata()["requires-python"] == ">=3.11,<3.15" readme = (PROJECT_ROOT / "README.md").read_text() - assert "supports Python 3.8 through 3.12" in readme - assert "Python 3.13 and newer are not supported" in readme + assert "supports Python 3.11 through 3.14" in readme def test_plotnine_supports_declared_python_floor(): dependencies = project_metadata()["dependencies"] - assert "plotnine==0.12.4" in dependencies + assert "plotnine>=0.15.8,<0.16" in dependencies + assert "matplotlib>=3.11.2" in dependencies def test_mmh3_is_not_a_runtime_dependency(): @@ -88,9 +88,9 @@ def test_svg_composition_dependencies_are_not_runtime_dependencies(): def test_ci_covers_every_supported_python_minor(): workflow = (PROJECT_ROOT / ".github/workflows/ci.yml").read_text() - for minor in range(8, 13): + for minor in range(11, 15): assert f' - "3.{minor}"' in workflow - assert ' - "3.13"' not in workflow + assert ' - "3.10"' not in workflow def test_release_workflow_is_tag_gated_and_uses_trusted_publishing(): @@ -102,7 +102,7 @@ def test_release_workflow_is_tag_gated_and_uses_trusted_publishing(): assert "id-token: write" in workflow assert "pypa/gh-action-pypi-publish@release/v1" in workflow assert "pypa/cibuildwheel@v4.2.0" in workflow - assert 'CIBW_BUILD: "cp39-*"' in workflow + assert 'CIBW_BUILD: "cp311-*"' in workflow assert "py_limited_api = cp38" in setup_config for runner in ("ubuntu-latest", "macos-15-intel", "windows-latest"): assert f" - {runner}" in workflow diff --git a/tests/test_plot_fonts.py b/tests/test_plot_fonts.py new file mode 100644 index 0000000..e2ca36e --- /dev/null +++ b/tests/test_plot_fonts.py @@ -0,0 +1,73 @@ +import matplotlib.pyplot as plt +from plotnine import ggplot + +import moddotplot.static_plots as static_plots +from moddotplot.native_render import ( + DEFAULT_FONT_FAMILY, + FALLBACK_FONT_FAMILY, + MIN_TEXT_SIZE, + MIN_TITLE_SIZE, + clamped_font_size, + save_figure_pair, +) + + +def test_width_scaled_fonts_have_readable_minimums(): + assert clamped_font_size(1, 1.0) == MIN_TEXT_SIZE + assert clamped_font_size(1, 1.4, MIN_TITLE_SIZE) == MIN_TITLE_SIZE + assert clamped_font_size(9, 1.4, MIN_TITLE_SIZE) == 12.6 + + +def test_native_outputs_default_to_helvetica(tmp_path): + figure, axis = plt.subplots() + title = axis.set_title("Helvetica title") + try: + save_figure_pair(figure, tmp_path / "helvetica", "svg", 72) + assert title.get_fontfamily() == [DEFAULT_FONT_FAMILY] + assert (tmp_path / "helvetica.png").stat().st_size > 0 + assert (tmp_path / "helvetica.svg").stat().st_size > 0 + finally: + plt.close(figure) + + +def test_native_outputs_retry_with_dejavu_on_glyph_failure(tmp_path, monkeypatch): + figure, axis = plt.subplots() + title = axis.set_title("Fallback title") + original_savefig = figure.savefig + attempted_families = [] + + def fail_for_helvetica(*args, **kwargs): + family = title.get_fontfamily()[0] + attempted_families.append(family) + if family == DEFAULT_FONT_FAMILY: + raise RuntimeError("failed to load glyph") + return original_savefig(*args, **kwargs) + + monkeypatch.setattr(figure, "savefig", fail_for_helvetica) + try: + save_figure_pair(figure, tmp_path / "fallback", "svg", 72) + assert attempted_families == [ + DEFAULT_FONT_FAMILY, + FALLBACK_FONT_FAMILY, + FALLBACK_FONT_FAMILY, + ] + assert title.get_fontfamily() == [FALLBACK_FONT_FAMILY] + assert (tmp_path / "fallback.png").stat().st_size > 0 + assert (tmp_path / "fallback.svg").stat().st_size > 0 + finally: + plt.close(figure) + + +def test_plotnine_outputs_retry_with_dejavu_on_glyph_failure(monkeypatch): + attempted_families = [] + + def fail_for_helvetica(plot, **_kwargs): + family = plot.theme.getp(("text", "family"))[0] + attempted_families.append(family) + if family == DEFAULT_FONT_FAMILY: + raise RuntimeError("failed to load glyph") + + monkeypatch.setattr(static_plots, "ggsave", fail_for_helvetica) + static_plots._save_plot(ggplot(), filename="unused.png") + + assert attempted_families == [DEFAULT_FONT_FAMILY, FALLBACK_FONT_FAMILY] diff --git a/tests/test_plot_summary.py b/tests/test_plot_summary.py new file mode 100644 index 0000000..fb1503b --- /dev/null +++ b/tests/test_plot_summary.py @@ -0,0 +1,61 @@ +from datetime import datetime, timezone + +from moddotplot.plot_summary import PlotSummaryWriter + + +def test_plot_summary_records_reproducibility_metadata_and_merges_groups(tmp_path): + first_plot = tmp_path / "chrA_FULL.png" + second_plot = tmp_path / "chrA_TRI.svg" + fasta = tmp_path / "source genome.fa" + annotation = tmp_path / "features.bed" + bedpe = tmp_path / "matrix.bedpe" + for path in (first_plot, second_plot, fasta, annotation, bedpe): + path.touch() + + writer = PlotSummaryWriter( + "moddotplot -f 'source genome.fa' --region chrA:1-4000000", + created_at=datetime(2026, 9, 28, 12, 30, tzinfo=timezone.utc), + ) + writer.add( + tmp_path, + [first_plot], + fasta_files=[fasta], + window_sizes=[4000], + regions=["chrA:1-4000000"], + bed_file=annotation, + ) + writer.add( + tmp_path, + [second_plot], + window_sizes=[4000], + bedpe_inputs=[bedpe], + ) + + summary = (tmp_path / "plot_summary.txt").read_text(encoding="utf-8") + assert "Created: 2026-09-28T12:30:00+00:00" in summary + assert ( + "Command: moddotplot -f 'source genome.fa' --region chrA:1-4000000" in summary + ) + assert str(first_plot.resolve()) in summary + assert str(second_plot.resolve()) in summary + assert str(fasta.resolve()) in summary + assert "4000 bp" in summary + assert "chrA:1-4000000" in summary + assert f"BED annotation file: {annotation.resolve()}" in summary + assert str(bedpe.resolve()) in summary + assert "Plot group" not in summary + assert "Created: 2026-09-28T12:30:00+00:00\n\nPlot files:" in summary + assert "Plot files:" in summary + assert "\n\nFASTA files:" in summary + assert "\n\nWindow sizes:" in summary + assert "\n\nRegions:" in summary + assert "\n\nBED annotation file:" in summary + + +def test_plot_summary_ignores_empty_plot_groups(tmp_path): + writer = PlotSummaryWriter("moddotplot --help") + + result = writer.add(tmp_path, [], fasta_files=["source.fa"]) + + assert result is None + assert not (tmp_path / "plot_summary.txt").exists() diff --git a/tests/test_static_customization.py b/tests/test_static_customization.py index 2ab7424..9916790 100644 --- a/tests/test_static_customization.py +++ b/tests/test_static_customization.py @@ -1,6 +1,8 @@ import pandas as pd import pytest +from matplotlib.colors import to_rgba +from moddotplot.const import DIRECTION_COLORS from moddotplot.static_plots import ( display_sequence_name, generate_breaks, @@ -82,6 +84,52 @@ def test_make_dot_uses_custom_color_scale(): assert fill_scale.palette(len(custom_colors)) == custom_colors +def test_make_dot_uses_direction_colors_when_orientation_is_present(): + plot_data = pd.DataFrame( + { + "q": ["query"] * 4, + "q_st": [0, 10, 20, 30], + "q_en": [10, 20, 30, 40], + "r": ["reference"] * 4, + "r_st": [0, 10, 20, 30], + "r_en": [10, 20, 30, 40], + "discrete": pd.Categorical([0, 1, 0, 1], categories=[0, 1]), + "direction": ["Forward", "Forward", "Reverse", "Reverse"], + } + ) + + plot = make_dot( + sdf=plot_data, + name_x="query", + name_y="reference", + palette="Spectral_11", + palette_orientation="+", + colors=["#000000", "#ffffff"], + breaks=[0, 10, 20, 30, 40], + num_ticks=3, + xlim=40, + deraster=True, + width=4, + is_pairwise=True, + ) + + figure = plot.draw(show=False) + try: + rendered = { + tuple(color) + for collection in figure.axes[0].collections + for color in collection.get_facecolors() + } + assert to_rgba(DIRECTION_COLORS["Forward"]) in rendered + assert to_rgba(DIRECTION_COLORS["Reverse"]) in rendered + assert len(rendered) == 4 + assert to_rgba("#000000") not in rendered + finally: + import matplotlib.pyplot as plt + + plt.close(figure) + + def test_make_dot_honors_exact_region_bounds(): plot_data = pd.DataFrame( { From 18f84234451814c730c565297491a15da4538c71 Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Tue, 29 Sep 2026 17:03:44 -0400 Subject: [PATCH 03/16] Optimize chromosome dotplot runtime and memory --- .gitignore | 6 + README.md | 33 +- src/moddotplot/__main__.py | 10 +- src/moddotplot/_nthash.cpp | 624 +++++++++++++++++++++ src/moddotplot/estimate_identity.py | 398 ++++++++++++-- src/moddotplot/moddotplot.py | 724 +++++++++++++++++++++++-- src/moddotplot/native_render.py | 29 +- src/moddotplot/parse_fasta.py | 347 ++++++++++-- src/moddotplot/static_plots.py | 481 +++++++++++----- tests/test_algorithms.py | 227 ++++++++ tests/test_annotation_track.py | 25 +- tests/test_cli_integration.py | 187 +++++++ tests/test_cli_runtime.py | 266 +++++++++ tests/test_fasta_parser.py | 226 ++++++++ tests/test_grid.py | 144 ++++- tests/test_native_full_plot.py | 224 ++++++++ tests/test_native_sequence_sketches.py | 153 ++++++ tests/test_plot_fonts.py | 45 +- tests/test_sparse_containment.py | 108 ++++ tests/test_static_customization.py | 12 + 20 files changed, 3956 insertions(+), 313 deletions(-) create mode 100644 .gitignore create mode 100644 tests/test_native_full_plot.py create mode 100644 tests/test_native_sequence_sketches.py diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..0d73943 --- /dev/null +++ b/.gitignore @@ -0,0 +1,6 @@ +__pycache__/ +*.py[cod] +*.so +build/ +*.egg-info/ +.pytest_cache/ diff --git a/README.md b/README.md index 76fc40b..fa9cbe8 100644 --- a/README.md +++ b/README.md @@ -150,7 +150,10 @@ The following arguments are the same in both interactive and static mode: `-f / --fasta ` -Fasta files to input. Multifasta files are accepted. Interactive mode will only support a maximum of two sequences at a time. +FASTA files to input. Multi-FASTA files are accepted. By default, every record +is analyzed; static mode can limit a run to named records with +`-s/--sequence`. Interactive mode will only support a maximum of two sequences +at a time. `-b / --bed <.bed file(s)>` @@ -210,6 +213,15 @@ Create a plot from a previously computed pairwise bed file. Skips Average Nucleo Run moddotplot static with a config file instead of command line args. Example syntax in `config/config.json`. Recommended when creating a really customized plot. Used instead of -f/--fasta. +`-s / --sequence [ ...]` + +In static mode, analyze only the requested records from the input FASTA +file(s). Each ID is matched to the first whitespace-delimited token in a FASTA +header. Exact matches are preferred, with an unambiguous case-insensitive +fallback (so `chr1` selects `Chr1`). Unknown, ambiguous, and duplicate requested +IDs are errors. When using a config file, provide the same list under the +`sequence` key, for example `"sequence": ["chr1", "chr2"]`. + `--cooler ` If set, will output a matrix as a cooler file for each input sequence, in addition to a bedpe file. @@ -226,6 +238,16 @@ Skip output of histogram legend. Save .bedpe to file, but skip rendering of plots. +`--processes <1-4>` + +Set the number of independent chromosome workers for self-only static runs. +When omitted, ModDotPlot uses two workers while rendering or up to four for a +`--no-plot` run on multi-record FASTA files that have random-access indexes +(`.fai`, plus `.gzi` for BGZF). This keeps default aggregate memory bounded; +`--processes 4` opts into maximum plotting throughput. Ordinary gzip and +unindexed inputs remain single-pass and sequential so they are not scanned +once per worker. Use `--processes 1` for explicitly serial execution. + `--width ` Adjust the output figure width. For a grid this is the width of the complete grid, not each cell. Default is 9 inches. @@ -274,7 +296,14 @@ With FASTA input, retain the standard ANI-colored plots and additionally create `--grid ` -Create a square grid containing every self comparison on the diagonal and every pairwise comparison off the diagonal. The grid is rendered as one Matplotlib figure and supports three or more input sequences, although large grids become visually dense. +Create a square grid containing every self comparison on the bottom-left-to-top-right diagonal and every pairwise comparison off the diagonal. Self-comparison cells are rendered as full symmetric dotplots. The shared genomic-axis titles appear only around the bottom-left cell without enlarging the output canvas. The grid is rendered as one Matplotlib figure and supports three or more input sequences, although large grids become visually dense. + +For example, select only `chr1` and `chr2` from a multi-FASTA input and create +their grid: + +``` +moddotplot -f ../../moddotplot-interactive/Col-CEN_v1.2.fasta -s chr1 chr2 --grid +``` `--grid-only ` diff --git a/src/moddotplot/__main__.py b/src/moddotplot/__main__.py index 31b5895..80bf3d0 100644 --- a/src/moddotplot/__main__.py +++ b/src/moddotplot/__main__.py @@ -5,7 +5,9 @@ from moddotplot.parse_fasta import * import setproctitle -# Set the process title to a custom name -setproctitle.setproctitle("ModDotPlot") - -sys.exit(main()) +if __name__ == "__main__": + # Keeping execution behind the standard guard makes this module safe to + # import in spawn-based chromosome worker processes and as a console-script + # entry point. + setproctitle.setproctitle("ModDotPlot") + sys.exit(main()) diff --git a/src/moddotplot/_nthash.cpp b/src/moddotplot/_nthash.cpp index 7a2a521..e24f427 100644 --- a/src/moddotplot/_nthash.cpp +++ b/src/moddotplot/_nthash.cpp @@ -4,9 +4,15 @@ #define PY_SSIZE_T_CLEAN #include +#include #include #include +#include #include +#include +#include +#include +#include // The rolling hash kernel below is derived from ntHash v2.4.0, commit // c26bd4572a19de81e30d55042dbd33c1fd21d4b6. ntHash is distributed under @@ -169,6 +175,609 @@ fnv1a(const char* sequence, unsigned k, bool reverse_complement) return hash; } +void +add_interval_layer(const char* sequence, + Py_ssize_t start, + Py_ssize_t end, + unsigned k, + bool use_canonical, + bool include_ambiguous, + unsigned sparsity_power, + std::vector& selected) +{ + if (start >= end) { + return; + } + unsigned invalid_count = 0; + for (unsigned i = 0; i < k; ++i) { + invalid_count += !valid(sequence[start + i]); + } + uint64_t forward = 0; + uint64_t reverse = 0; + bool rolling = false; + const uint64_t mask = sparsity_power == 0 + ? 0 + : (uint64_t{ 1 } << sparsity_power) - 1; + + for (Py_ssize_t position = start; position < end; ++position) { + if (position > start) { + const bool previous_window_was_valid = invalid_count == 0; + invalid_count -= !valid(sequence[position - 1]); + invalid_count += !valid(sequence[position + k - 1]); + if (invalid_count == 0 && previous_window_was_valid) { + const char outgoing = sequence[position - 1]; + const char incoming = sequence[position + k - 1]; + forward = + srol(forward) ^ seed(incoming) ^ srol(seed(outgoing), k); + reverse ^= srol(complement_seed(incoming), k); + reverse ^= complement_seed(outgoing); + reverse = sror(reverse); + rolling = true; + } else if (invalid_count != 0) { + rolling = false; + } + } + + const bool ambiguous = invalid_count != 0; + if (ambiguous && !include_ambiguous) { + continue; + } + uint64_t value = 0; + if (!ambiguous) { + if (!rolling) { + base_hashes(sequence + position, k, forward, reverse); + rolling = true; + } + value = use_canonical ? canonical(forward, reverse) : forward; + } else { + const uint64_t fallback_forward = fnv1a(sequence + position, k, false); + value = use_canonical + ? canonical(fallback_forward, + fnv1a(sequence + position, k, true)) + : fallback_forward; + } + if ((value & mask) == 0) { + selected.push_back(value); + } + } +} + +// The ordinary sparsity layer is populated during the single chromosome pass. +// Only a genuinely deficient interval is re-hashed at successively denser +// layers. Human windows overwhelmingly take the one-pass path, while the rare +// fallback remains bit-for-bit equivalent and bounded to that interval. +struct WindowSketch +{ + Py_ssize_t start; + Py_ssize_t end; + std::size_t output_index; + unsigned sparsity_power; + std::size_t minimum_size; + std::vector top_layer; + + WindowSketch(Py_ssize_t interval_start, + Py_ssize_t interval_end, + std::size_t index, + unsigned power, + std::size_t minimum) + : start(interval_start) + , end(interval_end) + , output_index(index) + , sparsity_power(power) + , minimum_size(minimum) + { + } + + void add(uint64_t value) + { + const uint64_t top_mask = sparsity_power == 0 + ? 0 + : (uint64_t{ 1 } << sparsity_power) - 1; + if ((value & top_mask) == 0) { + top_layer.push_back(value); + } + } + + std::vector finish(const char* sequence, + unsigned k, + bool use_canonical, + bool include_ambiguous) + { + std::sort(top_layer.begin(), top_layer.end()); + top_layer.erase( + std::unique(top_layer.begin(), top_layer.end()), top_layer.end()); + unsigned selected_power = sparsity_power; + while (top_layer.size() < minimum_size && selected_power > 0) { + --selected_power; + add_interval_layer(sequence, + start, + end, + k, + use_canonical, + include_ambiguous, + selected_power, + top_layer); + std::sort(top_layer.begin(), top_layer.end()); + top_layer.erase( + std::unique(top_layer.begin(), top_layer.end()), top_layer.end()); + } + return std::move(top_layer); + } +}; + +inline bool +is_power_of_two(Py_ssize_t value) +{ + return value > 0 && + (static_cast(value) & + (static_cast(value) - 1)) == 0; +} + +inline unsigned +power_of_two_exponent(Py_ssize_t value) +{ + unsigned power = 0; + while (value > 1) { + value >>= 1; + ++power; + } + return power; +} + +PyObject* +sketch_kmers(PyObject*, PyObject* args) +{ + const char* sequence = nullptr; + Py_ssize_t sequence_length = 0; + Py_ssize_t requested_k = 0; + int use_canonical = 0; + PyObject* intervals_object = nullptr; + Py_ssize_t sparsity = 0; + Py_ssize_t requested_minimum_size = 0; + int include_ambiguous = 0; + if (!PyArg_ParseTuple(args, + "s#npOnnp:sketch_kmers", + &sequence, + &sequence_length, + &requested_k, + &use_canonical, + &intervals_object, + &sparsity, + &requested_minimum_size, + &include_ambiguous)) { + return nullptr; + } + if (requested_k < 1 || + requested_k > std::numeric_limits::max()) { + PyErr_SetString(PyExc_ValueError, "k must be between 1 and 65535"); + return nullptr; + } + if (!is_power_of_two(sparsity)) { + PyErr_SetString(PyExc_ValueError, "sparsity must be a positive power of two"); + return nullptr; + } + + const Py_ssize_t count = sequence_length >= requested_k + ? sequence_length - requested_k + 1 + : 0; + const Py_ssize_t interval_count = PySequence_Size(intervals_object); + if (interval_count < 0) { + return nullptr; + } + const unsigned sparsity_power = power_of_two_exponent(sparsity); + const std::size_t minimum_size = requested_minimum_size > 0 + ? static_cast( + requested_minimum_size) + : 0; + + std::vector windows; + std::vector> output; + std::vector schedule; + try { + windows.reserve(static_cast(interval_count)); + for (Py_ssize_t index = 0; index < interval_count; ++index) { + PyObject* interval = PySequence_GetItem(intervals_object, index); + if (interval == nullptr) { + return nullptr; + } + Py_ssize_t start = 0; + Py_ssize_t end = 0; + const int parsed = PyArg_ParseTuple(interval, "nn", &start, &end); + Py_DECREF(interval); + if (!parsed) { + return nullptr; + } + if (start < 0 || end < start || end > count) { + PyErr_SetString( + PyExc_ValueError, + "interval bounds must satisfy 0 <= start <= end <= k-mer count"); + return nullptr; + } + windows.emplace_back(start, + end, + static_cast(index), + sparsity_power, + minimum_size); + } + + output.resize(static_cast(interval_count)); + schedule.resize(static_cast(interval_count)); + std::iota(schedule.begin(), schedule.end(), std::size_t{ 0 }); + std::sort(schedule.begin(), schedule.end(), [&](std::size_t left, + std::size_t right) { + if (windows[left].start != windows[right].start) { + return windows[left].start < windows[right].start; + } + return windows[left].end < windows[right].end; + }); + } catch (const std::bad_alloc&) { + PyErr_NoMemory(); + return nullptr; + } catch (const std::exception& error) { + PyErr_SetString(PyExc_RuntimeError, error.what()); + return nullptr; + } catch (...) { + PyErr_SetString(PyExc_RuntimeError, "failed to initialize native sketches"); + return nullptr; + } + + std::exception_ptr computation_error; + Py_BEGIN_ALLOW_THREADS + try { + std::vector active; + active.reserve(8); + std::size_t next_window = 0; + const unsigned k = static_cast(requested_k); + unsigned invalid_count = 0; + for (unsigned i = 0; i < k && count > 0; ++i) { + invalid_count += !valid(sequence[i]); + } + uint64_t forward = 0; + uint64_t reverse = 0; + bool rolling = false; + + for (Py_ssize_t position = 0; position < count; ++position) { + while (next_window < schedule.size() && + windows[schedule[next_window]].start <= position) { + const std::size_t index = schedule[next_window++]; + auto& window = windows[index]; + if (window.start == window.end) { + output[window.output_index] = window.finish( + sequence, k, use_canonical != 0, include_ambiguous != 0); + } else { + active.push_back(index); + } + } + auto kept_end = std::remove_if( + active.begin(), active.end(), [&](std::size_t index) { + auto& window = windows[index]; + if (window.end > position) { + return false; + } + output[window.output_index] = window.finish( + sequence, k, use_canonical != 0, include_ambiguous != 0); + return true; + }); + active.erase(kept_end, active.end()); + + if (position > 0) { + const bool previous_window_was_valid = invalid_count == 0; + invalid_count -= !valid(sequence[position - 1]); + invalid_count += !valid(sequence[position + requested_k - 1]); + if (invalid_count == 0 && previous_window_was_valid) { + const char outgoing = sequence[position - 1]; + const char incoming = sequence[position + requested_k - 1]; + forward = + srol(forward) ^ seed(incoming) ^ srol(seed(outgoing), k); + reverse ^= srol(complement_seed(incoming), k); + reverse ^= complement_seed(outgoing); + reverse = sror(reverse); + rolling = true; + } else if (invalid_count != 0) { + rolling = false; + } + } + + const bool ambiguous = invalid_count != 0; + if (ambiguous && !include_ambiguous) { + continue; + } + uint64_t value = 0; + if (!ambiguous) { + if (!rolling) { + base_hashes(sequence + position, k, forward, reverse); + rolling = true; + } + value = use_canonical ? canonical(forward, reverse) : forward; + } else { + const uint64_t fallback_forward = + fnv1a(sequence + position, k, false); + value = use_canonical + ? canonical(fallback_forward, + fnv1a(sequence + position, k, true)) + : fallback_forward; + } + for (const std::size_t index : active) { + windows[index].add(value); + } + } + + for (const std::size_t index : active) { + auto& window = windows[index]; + output[window.output_index] = window.finish( + sequence, k, use_canonical != 0, include_ambiguous != 0); + } + while (next_window < schedule.size()) { + auto& window = windows[schedule[next_window++]]; + output[window.output_index] = window.finish( + sequence, k, use_canonical != 0, include_ambiguous != 0); + } + } catch (...) { + computation_error = std::current_exception(); + } + Py_END_ALLOW_THREADS + + if (computation_error != nullptr) { + try { + std::rethrow_exception(computation_error); + } catch (const std::bad_alloc&) { + PyErr_NoMemory(); + } catch (const std::exception& error) { + PyErr_SetString(PyExc_RuntimeError, error.what()); + } catch (...) { + PyErr_SetString(PyExc_RuntimeError, "native sketch construction failed"); + } + return nullptr; + } + + PyObject* result = PyList_New(interval_count); + if (result == nullptr) { + return nullptr; + } + for (Py_ssize_t index = 0; index < interval_count; ++index) { + const auto& sketch = output[static_cast(index)]; + if (sketch.size() > static_cast( + std::numeric_limits::max()) / + sizeof(uint64_t)) { + Py_DECREF(result); + PyErr_SetString(PyExc_OverflowError, "sketch output is too large"); + return nullptr; + } + PyObject* packed = PyBytes_FromStringAndSize( + reinterpret_cast(sketch.data()), + static_cast(sketch.size() * sizeof(uint64_t))); + if (packed == nullptr) { + Py_DECREF(result); + return nullptr; + } + PyList_SetItem(result, index, packed); + } + return result; +} + +struct SketchInput +{ + const char* values; + std::size_t length; + std::size_t row; + bool right; + PyObject* owner; +}; + +struct OwnedSketchInputs +{ + std::vector values; + + OwnedSketchInputs() = default; + OwnedSketchInputs(const OwnedSketchInputs&) = delete; + OwnedSketchInputs& operator=(const OwnedSketchInputs&) = delete; + + ~OwnedSketchInputs() + { + // Destruction happens only while the calling thread owns the GIL: native + // computation catches exceptions before Py_END_ALLOW_THREADS. Keeping one + // reference per buffer makes the raw pointers safe even when callers pass + // a custom sequence that creates bytes objects on demand, or mutate an + // input list from another Python thread while the merge is running. + for (const auto& input : values) { + Py_DECREF(input.owner); + } + } +}; + +inline uint64_t +sketch_value_at(const SketchInput& input, std::size_t offset) +{ + uint64_t value = 0; + std::memcpy(&value, input.values + offset * sizeof(uint64_t), sizeof(value)); + return value; +} + +struct SketchCursor +{ + uint64_t value; + std::size_t input; + std::size_t offset; +}; + +struct CursorGreater +{ + bool operator()(const SketchCursor& left, const SketchCursor& right) const + { + return left.value > right.value; + } +}; + +bool +append_sketch_inputs(PyObject* sketches, + Py_ssize_t count, + bool right, + OwnedSketchInputs& inputs) +{ + for (Py_ssize_t row = 0; row < count; ++row) { + PyObject* sketch = PySequence_GetItem(sketches, row); + if (sketch == nullptr) { + return false; + } + char* buffer = nullptr; + Py_ssize_t buffer_length = 0; + const int status = + PyBytes_AsStringAndSize(sketch, &buffer, &buffer_length); + if (status < 0) { + Py_DECREF(sketch); + return false; + } + if (buffer_length < 0 || + buffer_length % static_cast(sizeof(uint64_t)) != 0) { + Py_DECREF(sketch); + PyErr_SetString( + PyExc_ValueError, "sketch buffers must contain packed uint64 values"); + return false; + } + try { + inputs.values.push_back( + { buffer, + static_cast(buffer_length) / sizeof(uint64_t), + static_cast(row), + right, + sketch }); + } catch (const std::bad_alloc&) { + Py_DECREF(sketch); + PyErr_NoMemory(); + return false; + } catch (const std::exception& error) { + Py_DECREF(sketch); + PyErr_SetString(PyExc_RuntimeError, error.what()); + return false; + } catch (...) { + Py_DECREF(sketch); + PyErr_SetString(PyExc_RuntimeError, "failed to retain sketch input"); + return false; + } + } + return true; +} + +PyObject* +intersection_counts(PyObject*, PyObject* args) +{ + PyObject* left_object = nullptr; + PyObject* right_object = nullptr; + if (!PyArg_ParseTuple( + args, "OO:intersection_counts", &left_object, &right_object)) { + return nullptr; + } + const Py_ssize_t rows = PySequence_Size(left_object); + if (rows < 0) { + return nullptr; + } + const Py_ssize_t columns = PySequence_Size(right_object); + if (columns < 0) { + return nullptr; + } + if (columns != 0 && + rows > std::numeric_limits::max() / columns) { + PyErr_SetString(PyExc_OverflowError, "intersection matrix is too large"); + return nullptr; + } + const Py_ssize_t cell_count = rows * columns; + if (cell_count > std::numeric_limits::max() / + static_cast(sizeof(int32_t))) { + PyErr_SetString(PyExc_OverflowError, "intersection matrix is too large"); + return nullptr; + } + if (rows > std::numeric_limits::max() - columns) { + PyErr_SetString(PyExc_OverflowError, "too many sketch inputs"); + return nullptr; + } + + OwnedSketchInputs inputs; + try { + inputs.values.reserve(static_cast(rows + columns)); + } catch (const std::bad_alloc&) { + PyErr_NoMemory(); + return nullptr; + } catch (const std::exception& error) { + PyErr_SetString(PyExc_RuntimeError, error.what()); + return nullptr; + } catch (...) { + PyErr_SetString(PyExc_RuntimeError, "failed to allocate sketch inputs"); + return nullptr; + } + // Use the dimensions captured above rather than asking arbitrary Python + // sequences for their lengths again. A stateful __len__ must not be able to + // produce row indices outside the matrix allocated from the first result. + if (!append_sketch_inputs(left_object, rows, false, inputs) || + !append_sketch_inputs(right_object, columns, true, inputs)) { + return nullptr; + } + + std::vector counts; + std::exception_ptr computation_error; + Py_BEGIN_ALLOW_THREADS + try { + counts.assign(static_cast(cell_count), int32_t{ 0 }); + std::priority_queue, + CursorGreater> + queue; + for (std::size_t index = 0; index < inputs.values.size(); ++index) { + if (inputs.values[index].length != 0) { + queue.push({ sketch_value_at(inputs.values[index], 0), index, 0 }); + } + } + + std::vector left_rows; + std::vector right_rows; + left_rows.reserve(static_cast(rows)); + right_rows.reserve(static_cast(columns)); + while (!queue.empty()) { + const uint64_t current_hash = queue.top().value; + left_rows.clear(); + right_rows.clear(); + do { + const SketchCursor cursor = queue.top(); + queue.pop(); + const SketchInput& input = inputs.values[cursor.input]; + (input.right ? right_rows : left_rows).push_back(input.row); + const std::size_t next_offset = cursor.offset + 1; + if (next_offset < input.length) { + queue.push( + { sketch_value_at(input, next_offset), cursor.input, next_offset }); + } + } while (!queue.empty() && queue.top().value == current_hash); + + for (const std::size_t left_row : left_rows) { + const std::size_t row_offset = + left_row * static_cast(columns); + for (const std::size_t right_row : right_rows) { + ++counts[row_offset + right_row]; + } + } + } + } catch (...) { + computation_error = std::current_exception(); + } + Py_END_ALLOW_THREADS + + if (computation_error != nullptr) { + try { + std::rethrow_exception(computation_error); + } catch (const std::bad_alloc&) { + PyErr_NoMemory(); + } catch (const std::exception& error) { + PyErr_SetString(PyExc_RuntimeError, error.what()); + } catch (...) { + PyErr_SetString(PyExc_RuntimeError, + "native sketch intersection failed"); + } + return nullptr; + } + return PyBytes_FromStringAndSize( + reinterpret_cast(counts.data()), + cell_count * static_cast(sizeof(int32_t))); +} + PyObject* hash_kmers(PyObject*, PyObject* args) { @@ -298,6 +907,21 @@ hash_kmers(PyObject*, PyObject* args) } PyMethodDef methods[] = { + { "intersection_counts", + intersection_counts, + METH_VARARGS, + "intersection_counts(left, right, /)\n--\n\n" + "Count intersections between sorted unique uint64 sketch buffers using a " + "bounded-memory k-way merge. Return a packed row-major int32 matrix." }, + { "sketch_kmers", + sketch_kmers, + METH_VARARGS, + "sketch_kmers(sequence, k, canonical, intervals, sparsity, " + "minimum_size, include_ambiguous, /)\n--\n\n" + "Hash a sequence and construct sorted, unique window sketches without " + "materializing its positional hash array. intervals contains zero-based, " + "half-open k-mer-index bounds. Return one packed native-endian uint64 " + "buffer per interval." }, { "hash_kmers", hash_kmers, METH_VARARGS, diff --git a/src/moddotplot/estimate_identity.py b/src/moddotplot/estimate_identity.py index 01602f8..094af69 100644 --- a/src/moddotplot/estimate_identity.py +++ b/src/moddotplot/estimate_identity.py @@ -14,6 +14,7 @@ import cooler from scipy.sparse import csr_matrix +from moddotplot import _nthash from moddotplot.parse_fasta import printProgressBar @@ -31,6 +32,64 @@ class PreparedModimizerSketches: neighbors: List[Collection[int]] +def prepare_sequence_sketches( + sequence, + window_size, + sparsity, + delta, + k, + ambiguous, + expectation, + canonical=True, +): + """Hash *sequence* directly into exact adaptive window sketches. + + Unlike :func:`prepare_modimizer_sketches`, this path never materializes the + chromosome-wide array containing one ``uint64`` per genomic k-mer. The + native kernel streams positional hashes through the currently active core + and expanded windows and returns only the sorted, unique sketches retained + by the existing adaptive-sparsity algorithm. + + The returned arrays and all interval/ambiguity semantics are identical to + hashing with ``_hash_sequence`` and then calling + :func:`prepare_modimizer_sketches`. + """ + + if k <= 0: + raise ValueError("k-mer size must be greater than zero") + if window_size <= 0: + raise ValueError("window size must be greater than zero") + + # ``s#`` accepts str and immutable bytes. Other contiguous byte-oriented + # sequence representations are normalized once, without uppercasing: the + # native ntHash implementation already handles ASCII case and U/T. + native_sequence = ( + sequence if isinstance(sequence, (str, bytes)) else bytes(sequence) + ) + kmer_count = max(len(native_sequence) - k + 1, 0) + core_bounds = _partition_bounds(kmer_count, window_size, 0, k) + neighbor_bounds = ( + _partition_bounds(kmer_count, window_size, delta, k) if delta > 0 else None + ) + bounds = core_bounds if neighbor_bounds is None else core_bounds + neighbor_bounds + packed_sketches = _nthash.sketch_kmers( + native_sequence, + k, + bool(canonical), + bounds, + int(sparsity), + round(expectation / 2), + bool(ambiguous), + ) + arrays = [np.frombuffer(packed, dtype=np.uint64) for packed in packed_sketches] + core = arrays[: len(core_bounds)] + if neighbor_bounds is None: + neighbors = core + else: + neighbors = arrays[len(core_bounds) :] + return PreparedModimizerSketches(core=core, neighbors=neighbors) + + def prepare_modimizer_sketches( sequence_length, sequence, @@ -229,13 +288,28 @@ def partitionOverlaps( if kmer_count == 0: return [] + return [ + lst[start:end] for start, end in _partition_bounds(kmer_count, win, delta, k) + ] + + +def _partition_bounds(kmer_count: int, win: int, delta: float, k: int): + """Return the half-open k-mer-index bounds used by window partitioning.""" + + if win <= 0: + raise ValueError("window size must be greater than zero") + if k <= 0: + raise ValueError("k-mer size must be greater than zero") + if kmer_count <= 0: + return [] + # A sequence with n bases has n - k + 1 k-mers. Reconstruct the genomic # length so that every partition starts on a multiple of ``win``. The old # counter started the second partition at ``win - k + 2`` and then advanced # by ``win``, making all but the first window begin k - 2 bases too early. sequence_length = kmer_count + k - 1 delta_offset = win * delta - kmer_list = [] + bounds = [] # A trailing genomic fragment shorter than k has no k-mer and therefore no # matrix cell, so iterate over valid k-mer starts rather than base length. @@ -248,9 +322,9 @@ def partitionOverlaps( # the right boundary excludes k-mers that cross out of the interval. start_index = min(expanded_start, kmer_count) end_index = min(max(expanded_end - k + 1, start_index), kmer_count) - kmer_list.append(lst[start_index:end_index]) + bounds.append((start_index, end_index)) - return kmer_list + return bounds def _valid_hashes(partition): @@ -382,7 +456,122 @@ def convertToModimizers( ] -def convertMatrixToBed( +BEDPE_HEADER = ( + "#query_name", + "query_start", + "query_end", + "reference_name", + "reference_start", + "reference_end", + "perID_by_events", +) + +# Bound the largest temporary threshold mask/nonzero result created while +# converting a dense identity matrix. The returned compatibility list can of +# course still be large; callers that need bounded output memory can consume +# ``iterMatrixToBedChunks`` directly. +DEFAULT_BEDPE_CHUNK_CELLS = 262_144 + + +def _iter_matrix_blocks(values, max_chunk_cells): + """Yield C-order matrix blocks containing at most ``max_chunk_cells``. + + Normally a block is a band of complete matrix rows. If one row itself is + wider than the configured bound, that row is split into consecutive column + blocks. In both cases the blocks, and cells within each block, retain the + same order as a nested row-then-column loop. + """ + + if isinstance(max_chunk_cells, bool) or not isinstance( + max_chunk_cells, (int, np.integer) + ): + raise TypeError("max_chunk_cells must be a positive integer") + max_chunk_cells = int(max_chunk_cells) + if max_chunk_cells <= 0: + raise ValueError("max_chunk_cells must be a positive integer") + + rows, cols = values.shape + if rows == 0 or cols == 0: + return + + if cols <= max_chunk_cells: + rows_per_band = max(1, max_chunk_cells // cols) + for row_start in range(0, rows, rows_per_band): + row_stop = min(row_start + rows_per_band, rows) + yield values[row_start:row_stop, :], row_start, 0 + return + + # An individual row exceeds the cell budget. Splitting it by columns is + # the only way to maintain the strict bound without changing C-order. + for row_index in range(rows): + for column_start in range(0, cols, max_chunk_cells): + column_stop = min(column_start + max_chunk_cells, cols) + yield ( + values[row_index : row_index + 1, column_start:column_stop], + row_index, + column_start, + ) + + +def _iter_matrix_to_bed_columns( + matrix, + window_size, + id_threshold, + self_identity, + x_offset, + y_offset, + x_end, + y_end, + max_chunk_cells, +): + """Yield column arrays for retained BEDPE cells in exact legacy order.""" + + values = np.asarray(matrix) + if values.ndim != 2: + raise ValueError("identity matrix must be two-dimensional") + + cutoff = id_threshold / 100 + for block, row_offset, column_offset in _iter_matrix_blocks( + values, max_chunk_cells + ): + # Threshold one bounded block at a time instead of allocating an + # additional matrix-sized boolean array. ``nonzero`` emits C-order + # indices, matching the historical nested x-then-y iteration. + local_x, local_y = np.nonzero(block >= cutoff) + if local_x.size == 0: + continue + + x_indices = local_x + row_offset + y_indices = local_y + column_offset + if self_identity: + upper_triangle = x_indices <= y_indices + if not bool(np.all(upper_triangle)): + x_indices = x_indices[upper_triangle] + y_indices = y_indices[upper_triangle] + if x_indices.size == 0: + continue + + start_x = x_indices * window_size + x_offset + end_x = start_x + window_size - 1 + start_y = y_indices * window_size + y_offset + end_y = start_y + window_size - 1 + if x_end is not None: + end_x = np.minimum(end_x, x_end) + if y_end is not None: + end_y = np.minimum(end_y, y_end) + + # Coordinate calculations may be floating point for interactive + # exports. Integer conversion deliberately happens after clipping, + # preserving the legacy scalar ``int(...)`` truncation semantics. + start_x = np.asarray(start_x).astype(np.int64, copy=False) + end_x = np.asarray(end_x).astype(np.int64, copy=False) + start_y = np.asarray(start_y).astype(np.int64, copy=False) + end_y = np.asarray(end_y).astype(np.int64, copy=False) + selected_values = np.asarray(values[x_indices, y_indices], dtype=float) + yield start_x, end_x, start_y, end_y, selected_values + + +def iterMatrixToBedChunks( matrix, window_size, id_threshold, @@ -393,50 +582,134 @@ def convertMatrixToBed( y_offset, x_end=None, y_end=None, + max_chunk_cells=DEFAULT_BEDPE_CHUNK_CELLS, ): - bed = [ - ( - "#query_name", - "query_start", - "query_end", - "reference_name", - "reference_start", - "reference_end", - "perID_by_events", + """Yield retained BEDPE records as bounded, columnar DataFrames. + + Every yielded frame contains at most ``max_chunk_cells`` records and uses + :data:`BEDPE_HEADER` as its exact column order. Empty matrix blocks are not + yielded. Consuming these frames incrementally avoids both a matrix-sized + threshold mask and the compatibility API's list of Python tuples. + """ + + for start_x, end_x, start_y, end_y, selected_values in _iter_matrix_to_bed_columns( + matrix, + window_size, + id_threshold, + self_identity, + x_offset, + y_offset, + x_end, + y_end, + max_chunk_cells, + ): + record_count = selected_values.size + yield pd.DataFrame( + { + BEDPE_HEADER[0]: np.full(record_count, x_name, dtype=object), + BEDPE_HEADER[1]: start_x, + BEDPE_HEADER[2]: end_x, + BEDPE_HEADER[3]: np.full(record_count, y_name, dtype=object), + BEDPE_HEADER[4]: start_y, + BEDPE_HEADER[5]: end_y, + BEDPE_HEADER[6]: selected_values, + }, + columns=BEDPE_HEADER, ) - ] - rows, cols = matrix.shape - for x in range(rows): - for y in range(cols): - value = matrix[x, y] - if (not self_identity) or (self_identity and x <= y): - if value >= id_threshold / 100: - start_x = x * window_size + x_offset - end_x = start_x + window_size - 1 - start_y = y * window_size + y_offset - end_y = start_y + window_size - 1 - - # The final matrix window can be shorter than ``window_size``. - # Keep exported coordinates inside the exact sequence or - # requested region instead of allowing the last tile to - # overhang it. - if x_end is not None: - end_x = min(end_x, x_end) - if y_end is not None: - end_y = min(end_y, y_end) - - bed.append( - ( - x_name, - int(start_x), - int(end_x), - y_name, - int(start_y), - int(end_y), - float(value), - ) - ) + +def convertMatrixToBedDataFrame( + matrix, + window_size, + id_threshold, + x_name, + y_name, + self_identity, + x_offset, + y_offset, + x_end=None, + y_end=None, + max_chunk_cells=DEFAULT_BEDPE_CHUNK_CELLS, +): + """Return BEDPE records as one columnar DataFrame. + + For fully bounded consumption, prefer :func:`iterMatrixToBedChunks`. + This convenience helper avoids the substantially larger list of Python + row tuples expected by the historical :func:`convertMatrixToBed` API. + """ + + chunks = list( + iterMatrixToBedChunks( + matrix, + window_size, + id_threshold, + x_name, + y_name, + self_identity, + x_offset, + y_offset, + x_end, + y_end, + max_chunk_cells, + ) + ) + if chunks: + return pd.concat(chunks, ignore_index=True, copy=False) + return pd.DataFrame( + { + BEDPE_HEADER[0]: pd.Series(dtype=object), + BEDPE_HEADER[1]: pd.Series(dtype=np.int64), + BEDPE_HEADER[2]: pd.Series(dtype=np.int64), + BEDPE_HEADER[3]: pd.Series(dtype=object), + BEDPE_HEADER[4]: pd.Series(dtype=np.int64), + BEDPE_HEADER[5]: pd.Series(dtype=np.int64), + BEDPE_HEADER[6]: pd.Series(dtype=float), + }, + columns=BEDPE_HEADER, + ) + + +def convertMatrixToBed( + matrix, + window_size, + id_threshold, + x_name, + y_name, + self_identity, + x_offset, + y_offset, + x_end=None, + y_end=None, + max_chunk_cells=DEFAULT_BEDPE_CHUNK_CELLS, +): + """Return the historical header-plus-row-tuples BEDPE representation.""" + + bed = [BEDPE_HEADER] + for start_x, end_x, start_y, end_y, selected_values in _iter_matrix_to_bed_columns( + matrix, + window_size, + id_threshold, + self_identity, + x_offset, + y_offset, + x_end, + y_end, + max_chunk_cells, + ): + bed.extend( + ( + x_name, + int(current_start_x), + int(current_end_x), + y_name, + int(current_start_y), + int(current_end_y), + float(value), + ) + for current_start_x, current_end_x, current_start_y, current_end_y, value in zip( + start_x, end_x, start_y, end_y, selected_values + ) + ) return bed @@ -600,6 +873,33 @@ def _sketch_intersection_counts(sketches_a, sketches_b): for sketch in sketches ) + # Prepared production sketches are sorted and unique. Merge their streams + # natively, retaining only one cursor per sketch and the dense output. This + # avoids the chromosome-scale concatenated hash array, int64 inverse map, + # coordinate universe, and two CSR incidence matrices. Unsorted or + # duplicate-bearing arrays continue through the compatibility path below. + if compact_uint64_arrays and all( + sketch.size < 2 or bool(np.all(sketch[1:] > sketch[:-1])) + for sketches in collections + for sketch in sketches + ): + + def packed(sketch): + # Native sequence sketches are zero-copy views of immutable bytes, + # so their original packed buffer can be reused. Compatibility + # arrays require one compact copy but no inverse/CSR structures. + return ( + sketch.base + if isinstance(sketch.base, bytes) and sketch.nbytes == len(sketch.base) + else sketch.tobytes() + ) + + packed_counts = _nthash.intersection_counts( + [packed(sketch) for sketch in sketches_a], + [packed(sketch) for sketch in sketches_b], + ) + return np.frombuffer(packed_counts, dtype=np.int32).reshape(rows, cols).copy() + if compact_uint64_arrays: # This is the normal prepared-sketch path. Concatenating the arrays in # C avoids boxing millions of hashes back into Python integers. @@ -690,10 +990,10 @@ def _sketch_intersection_counts(sketches_a, sketches_b): def _identity_matrix_from_containment(containment_matrix, identity, k): """Apply ModDotPlot's k-mer identity transform and cutoff in place.""" - identities = np.power(containment_matrix, 1.0 / k) - identities[identities < identity / 100] = 0.0 - identities *= 100.0 - return identities + np.power(containment_matrix, 1.0 / k, out=containment_matrix) + containment_matrix[containment_matrix < identity / 100] = 0.0 + containment_matrix *= 100.0 + return containment_matrix def selfContainmentMatrix( diff --git a/src/moddotplot/moddotplot.py b/src/moddotplot/moddotplot.py index 08b243f..d7bd529 100755 --- a/src/moddotplot/moddotplot.py +++ b/src/moddotplot/moddotplot.py @@ -3,6 +3,8 @@ from moddotplot.parse_fasta import ( HASH_ALGORITHM, readKmersFromFile, + iter_fasta_records, + supports_indexed_fasta_access, getInputHeaders, isValidFasta, extractFiles, @@ -13,12 +15,16 @@ convertToModimizers, selfContainmentMatrix, pairwiseContainmentMatrix, + BEDPE_HEADER, convertMatrixToBed, + convertMatrixToBedDataFrame, + iterMatrixToBedChunks, convertMatrixToCool, createSelfMatrix, createPairwiseMatrix, create_self_matrix_from_sketches, create_pairwise_matrix_from_sketches, + prepare_sequence_sketches, ModimizerSketchCache, partitionOverlaps, ) @@ -27,11 +33,15 @@ from moddotplot.const import ASCII_ART, VERSION import argparse +from concurrent.futures import ProcessPoolExecutor, as_completed +from dataclasses import dataclass +from itertools import islice import math import json import numpy as np import pickle import os +import multiprocessing import shlex from moddotplot.plot_summary import PlotSummaryWriter @@ -253,6 +263,17 @@ def get_parser(): nargs="+", ) + static_parser.add_argument( + "-s", + "--sequence", + default=None, + nargs="+", + help=( + "Analyze only these FASTA sequence identifiers. Exact matches are " + "preferred; an unambiguous case-insensitive match is also accepted." + ), + ) + # Add a mutually exclusive group for compare and compare only. static_compare_group = static_parser.add_mutually_exclusive_group(required=False) static_window_size_group = static_parser.add_mutually_exclusive_group( @@ -354,6 +375,19 @@ def get_parser(): "--no-hist", action="store_true", help="Skip output of histogram color legend." ) + static_parser.add_argument( + "--processes", + default=None, + type=int, + choices=range(1, 5), + metavar="N", + help=( + "Number of independent chromosome workers (1-4). The default " + "automatically uses up to four workers for indexed multi-record " + "FASTA input." + ), + ) + static_parser.add_argument( "--width", default=9, type=float, help="Plot width (also height for _FULL)." ) @@ -499,6 +533,7 @@ def _apply_static_config(args, config): args.modimizer = config.get("modimizer", args.modimizer) args.resolution = config.get("resolution", args.resolution) args.window = config.get("window", args.window) + args.sequence = config.get("sequence", args.sequence) args.region = config.get("region", args.region) args.identity = config.get("identity", args.identity) args.delta = config.get("delta", args.delta) @@ -511,6 +546,7 @@ def _apply_static_config(args, config): args.no_bedpe = config.get("no_bedpe", args.no_bedpe) args.no_plot = config.get("no_plot", args.no_plot) args.no_hist = config.get("no_hist", args.no_hist) + args.processes = config.get("processes", args.processes) args.width = config.get("width", args.width) args.axes_limits = config.get("axes_limits", args.axes_limits) args.dpi = config.get("dpi", args.dpi) @@ -534,6 +570,68 @@ def _apply_static_config(args, config): return args +def _select_fasta_headers(fasta_headers, requested_sequences): + """Filter FASTA headers using exact-first, case-insensitive selectors. + + The returned names retain their spelling and source order from the FASTA + inputs so downstream plot labels and ``--compare-order sequential`` remain + stable. A case-insensitive fallback makes common selectors such as ``chr1`` + work with records named ``Chr1`` while still rejecting ambiguous matches. + """ + + if not requested_sequences: + return ( + {path: list(headers) for path, headers in fasta_headers.items()}, + [header for headers in fasta_headers.values() for header in headers], + ) + if isinstance(requested_sequences, str): + requested_sequences = [requested_sequences] + + exact_matches = {} + folded_matches = {} + for path, headers in fasta_headers.items(): + for header in headers: + record = (path, header) + exact_matches.setdefault(header, []).append(record) + folded_matches.setdefault(header.casefold(), []).append(record) + + selected_records = set() + for selector in requested_sequences: + if not isinstance(selector, str) or not selector: + raise ValueError("sequence identifiers must be non-empty strings") + matches = exact_matches.get(selector) + if matches is None: + matches = folded_matches.get(selector.casefold(), []) + if not matches: + raise ValueError( + f"sequence {selector!r} does not match any FASTA identifier" + ) + if len(matches) > 1: + formatted = ", ".join( + f"{header!r} in {os.fspath(path)!r}" for path, header in matches + ) + raise ValueError( + f"sequence {selector!r} is ambiguous; it matches {formatted}" + ) + + record = matches[0] + if record in selected_records: + raise ValueError( + f"sequence {selector!r} selects FASTA identifier {record[1]!r}, " + "which was already requested" + ) + selected_records.add(record) + + selected_headers = {} + selected_names = [] + for path, headers in fasta_headers.items(): + retained = [header for header in headers if (path, header) in selected_records] + if retained: + selected_headers[path] = retained + selected_names.extend(retained) + return selected_headers, selected_names + + def _parse_region_arguments(region_arguments, sequence_names): """Validate CLI regions and index them by their exact FASTA identifier.""" @@ -666,6 +764,496 @@ def _annotate_bed_directions( return annotated +@dataclass(frozen=True) +class MatrixConfig: + """Effective per-sequence parameters for one identity matrix.""" + + window_size: int + resolution: int + modimizer: int + sparsity: int + expectation: int + + +def _matrix_config_for_length(kmer_count, args): + """Resolve matrix parameters without mutating the parsed CLI arguments. + + Automatic windows are never shorter than a k-mer. A requested resolution + that would create smaller windows is reduced to the highest meaningful + resolution for that sequence instead. This matters for chrM at the default + resolution, where 17-base windows cannot contain a 21-mer. + """ + + kmer_count = int(kmer_count) + if kmer_count <= 0: + raise ValueError("sequence must contain at least one k-mer") + + if args.window is not None: + window_size = int(args.window) + if window_size <= 0: + raise ValueError("window size must be greater than zero") + resolution = math.ceil(kmer_count / window_size) + else: + requested_resolution = int(args.resolution) + if requested_resolution <= 0: + raise ValueError("resolution must be greater than zero") + window_size = max( + int(args.kmer), math.ceil(kmer_count / requested_resolution) + ) + resolution = math.ceil(kmer_count / window_size) + + if window_size < 10: + raise ValueError("window size must be at least 10 bases") + + requested_modimizer = int(args.modimizer) + if requested_modimizer <= 0: + raise ValueError("modimizer sketch size must be greater than zero") + effective_modimizer = min(requested_modimizer, window_size) + sparsity_ratio = max(1, round(window_size / effective_modimizer)) + if sparsity_ratio <= effective_modimizer: + sparsity = 2 ** int(math.log2(sparsity_ratio)) + else: + sparsity = 2 ** (int(math.log2(sparsity_ratio - 1)) + 1) + expectation = round(window_size / sparsity) + return MatrixConfig( + window_size=window_size, + resolution=resolution, + modimizer=effective_modimizer, + sparsity=sparsity, + expectation=expectation, + ) + + +def _write_bedpe(path, rows): + """Write generic BEDPE rows in bounded batches.""" + + row_iterator = iter(rows) + with open(path, "w") as bedfile: + while batch := list(islice(row_iterator, 8192)): + bedfile.writelines( + "\t".join(map(str, row)) + "\n" for row in batch + ) + + +def _write_matrix_bedpe(path, chunks): + """Stream columnar BEDPE chunks without building Python row tuples.""" + + with open(path, "w") as bedfile: + bedfile.write("\t".join(BEDPE_HEADER) + "\n") + for frame in chunks: + bedfile.writelines( + "\t".join(map(str, row)) + "\n" + for row in frame.itertuples(index=False, name=None) + ) + + +def _annotate_bed_direction_frame( + frame, + forward_matrix, + window_size, + x_offset, + y_offset, +): + """Add orientation labels to a columnar BEDPE dataframe.""" + + annotated = frame.copy() + if annotated.empty: + annotated["direction"] = np.asarray([], dtype=object) + return annotated + query_indices = np.rint( + (annotated["query_start"].to_numpy() - x_offset) / window_size + ).astype(np.intp) + reference_indices = np.rint( + (annotated["reference_start"].to_numpy() - y_offset) / window_size + ).astype(np.intp) + try: + values = np.asarray(forward_matrix)[query_indices, reference_indices] + except IndexError as error: + raise ValueError( + "BEDPE coordinates fall outside direction matrices" + ) from error + annotated["direction"] = np.where(values > 0, "Forward", "Reverse") + return annotated + + +def _streaming_process_count(args, record_count, indexed_access): + """Resolve a safe worker count for independent chromosome plots.""" + + requested = _validated_process_count(getattr(args, "processes", None)) + + if record_count < 2 or not indexed_access: + return 1 + available_cpus = max(1, os.cpu_count() or 1) + # Plotting temporarily retains Matplotlib figures and encoded raster + # buffers in addition to a chromosome's sketches. Two automatic rendering + # workers keep aggregate RSS below the former all-chromosome pipeline on + # typical hosts; compute-only runs can safely use all four. Users may + # explicitly request up to four when throughput is the priority. + automatic_limit = 4 if getattr(args, "no_plot", False) else 2 + automatic = min(automatic_limit, available_cpus, record_count) + return min(requested, record_count) if requested is not None else automatic + + +def _validated_process_count(value): + """Normalize a CLI/config process count and enforce the public bound.""" + + if value is None: + return None + if isinstance(value, bool): + raise ValueError("processes must be an integer from 1 through 4") + try: + value = int(value) + except (TypeError, ValueError) as error: + raise ValueError("processes must be an integer from 1 through 4") from error + if value < 1 or value > 4: + raise ValueError("processes must be an integer from 1 through 4") + return value + + +def _process_static_self_task(task): + """Reopen and process one indexed FASTA record in a spawned worker.""" + + args, fasta_path, sequence_id, selected_region, summary_command = task + if not args.no_plot: + _load_static_plotting() + regions = {sequence_id: selected_region} if selected_region is not None else None + records = iter_fasta_records( + fasta_path, + regions=regions, + record_ids=[sequence_id], + ) + try: + current_id, sequence, sequence_label = next(records) + except StopIteration as error: + raise ValueError( + f"sequence {sequence_id!r} was not found in {fasta_path!r}" + ) from error + + _process_static_self_record( + args=args, + sequence_id=current_id, + sequence=sequence, + sequence_label=sequence_label, + source_path=fasta_path, + summary_writer=PlotSummaryWriter(summary_command), + ) + return sequence_id + + +def _process_static_self_record( + *, + args, + sequence_id, + sequence, + sequence_label, + source_path, + summary_writer, +): + """Compute and emit one independent static self plot. + + The input sequence and its compact prepared sketches are deliberately kept + local so processing the next FASTA record cannot retain this chromosome. + """ + + sequence_start = 1 + parsed_label = extractRegion(sequence_label) + if parsed_label: + _base_name, sequence_start, sequence_end = parsed_label + else: + sequence_end = len(sequence) + sequence_name = sequence_label + plot_axis_bounds = args.axes_limits or (sequence_start, sequence_end) + kmer_count = len(sequence) - args.kmer + 1 + if kmer_count <= 0: + raise ValueError( + f"sequence {sequence_name!r} is shorter than k-mer size {args.kmer}" + ) + + config = _matrix_config_for_length(kmer_count, args) + print(f"Computing self identity matrix for {sequence_name}... \n") + print(f"\tSequence length n: {len(sequence)}\n") + print(f"\tWindow size w: {config.window_size}\n") + print(f"\tModimizer sketch size: {config.expectation}\n") + print(f"\tPlot Resolution r: {config.resolution}\n") + + direction_rendering = args.plot_direction and not args.no_plot + prepared = prepare_sequence_sketches( + sequence, + config.window_size, + config.sparsity, + args.delta, + args.kmer, + args.ambiguous, + config.expectation, + canonical=not args.forward, + ) + alternate_prepared = None + if direction_rendering: + alternate_prepared = prepare_sequence_sketches( + sequence, + config.window_size, + config.sparsity, + args.delta, + args.kmer, + args.ambiguous, + config.expectation, + canonical=bool(args.forward), + ) + # No chromosome-sized sequence or positional-hash array survives into the + # sparse intersection and rendering phases. + del sequence + + self_mat = create_self_matrix_from_sketches( + prepared, args.kmer, args.identity, args.ambiguous + ) + del prepared + direction_self_mat = None + if alternate_prepared is not None: + direction_self_mat = create_self_matrix_from_sketches( + alternate_prepared, args.kmer, args.identity, args.ambiguous + ) + del alternate_prepared + + output_directory = os.path.join(args.output_dir or ".", sequence_name) + if (not args.no_bedpe) or (not args.no_plot): + os.makedirs(output_directory, exist_ok=True) + + if args.cooler: + try: + os.makedirs(output_directory, exist_ok=True) + cooler_output = os.path.join(output_directory, sequence_name + ".cooler") + convertMatrixToCool( + matrix=self_mat, + window_size=config.window_size, + id_threshold=args.identity, + x_name=sequence_name, + y_name=sequence_name, + self_identity=True, + x_offset=sequence_start, + y_offset=sequence_start, + chromsizes=kmer_count, + output_cool=cooler_output, + ) + print( + f"Saved self-identity matrix as a cooler file to {cooler_output}\n" + ) + except Exception as error: + print(f"Error creating cooler file: {error}") + + bed_frame = None + direction_frame = None + if not args.no_plot: + bed_frame = convertMatrixToBedDataFrame( + self_mat, + config.window_size, + args.identity, + sequence_name, + sequence_name, + True, + sequence_start, + sequence_start, + sequence_end, + sequence_end, + ) + if direction_rendering: + if args.forward: + canonical_frame = convertMatrixToBedDataFrame( + direction_self_mat, + config.window_size, + args.identity, + sequence_name, + sequence_name, + True, + sequence_start, + sequence_start, + sequence_end, + sequence_end, + ) + forward_matrix = self_mat + else: + canonical_frame = bed_frame + forward_matrix = direction_self_mat + direction_frame = _annotate_bed_direction_frame( + canonical_frame, + forward_matrix, + config.window_size, + sequence_start, + sequence_start, + ) + + if not args.no_bedpe: + bedfile_output = os.path.join(output_directory, sequence_name + ".bedpe") + if bed_frame is not None: + _write_matrix_bedpe(bedfile_output, [bed_frame]) + else: + _write_matrix_bedpe( + bedfile_output, + iterMatrixToBedChunks( + self_mat, + config.window_size, + args.identity, + sequence_name, + sequence_name, + True, + sequence_start, + sequence_start, + sequence_end, + sequence_end, + ), + ) + print( + f"Saved self-identity matrix as a paired-end bed file to {bedfile_output}\n" + ) + + # Rendering consumes only compact BEDPE columns. Release dense matrices + # before Matplotlib/Plotnine allocate figures and raster buffers. + del self_mat + if direction_self_mat is not None: + del direction_self_mat + + if not args.no_plot: + plot_files = create_plots( + sdf=None, + directory=output_directory, + name_x=sequence_name, + name_y=sequence_name, + palette=args.palette, + palette_orientation=args.palette_orientation, + no_hist=args.no_hist, + width=args.width, + dpi=args.dpi, + is_freq=args.bin_freq, + xlim=plot_axis_bounds, + custom_colors=args.colors, + custom_breakpoints=args.breakpoints, + from_file=bed_frame, + is_pairwise=False, + axes_labels=args.axes_ticks, + axes_tick_number=args.axes_number, + vector_format=args.vector, + deraster=args.deraster, + annotation=args.bed, + ) + summary_writer.add( + output_directory, + plot_files or [], + fasta_files=[source_path], + window_sizes=[config.window_size], + regions=_regions_from_names([sequence_name]), + bed_file=args.bed, + ) + if direction_rendering: + direction_directory = os.path.join(output_directory, "directionality") + direction_files = create_plots( + sdf=None, + directory=direction_directory, + name_x=sequence_name, + name_y=sequence_name, + palette=args.palette, + palette_orientation=args.palette_orientation, + no_hist=args.no_hist, + width=args.width, + dpi=args.dpi, + is_freq=args.bin_freq, + xlim=plot_axis_bounds, + custom_colors=args.colors, + custom_breakpoints=args.breakpoints, + from_file=direction_frame, + is_pairwise=False, + axes_labels=args.axes_ticks, + axes_tick_number=args.axes_number, + vector_format=args.vector, + deraster=args.deraster, + annotation=None, + ) + summary_writer.add( + direction_directory, + direction_files or [], + fasta_files=[source_path], + window_sizes=[config.window_size], + regions=_regions_from_names([sequence_name]), + bed_file=args.bed, + ) + + +def _run_streaming_static_self( + args, fasta_list, fasta_headers, region_by_name, summary_writer +): + """Run independent self plots with bounded, record-local memory. + + Indexed FASTA records can be reopened directly by up to four spawned + workers. Unindexed plain/gzip inputs retain the one-pass sequential path; + this avoids making every worker rescan or decompress the entire file. + """ + + if args.output_dir: + os.makedirs(args.output_dir, exist_ok=True) + + tasks = [ + ( + args, + fasta_path, + sequence_id, + region_by_name.get(sequence_id), + getattr(summary_writer, "command", ""), + ) + for fasta_path in fasta_list + for sequence_id in fasta_headers[fasta_path] + ] + unique_record_ids = {task[2] for task in tasks} + indexed_access = ( + len(unique_record_ids) == len(tasks) + and all(supports_indexed_fasta_access(path) for path in fasta_list) + ) + process_count = _streaming_process_count(args, len(tasks), indexed_access) + + if process_count > 1: + print( + f"Processing {len(tasks)} sequences with {process_count} " + "chromosome workers.\n" + ) + context = multiprocessing.get_context("spawn") + try: + with ProcessPoolExecutor( + max_workers=process_count, + mp_context=context, + ) as executor: + future_records = { + executor.submit(_process_static_self_task, task): task[2] + for task in tasks + } + for future in as_completed(future_records): + sequence_id = future_records[future] + try: + future.result() + except Exception as error: + for pending in future_records: + pending.cancel() + raise ValueError( + f"failed while processing sequence {sequence_id!r}: {error}" + ) from error + except ValueError: + raise + except (OSError, RuntimeError) as error: + raise ValueError(f"unable to run chromosome workers: {error}") from error + return + + for fasta_path in fasta_list: + for sequence_id, sequence, sequence_label in iter_fasta_records( + fasta_path, + regions=region_by_name, + record_ids=fasta_headers[fasta_path], + ): + _process_static_self_record( + args=args, + sequence_id=sequence_id, + sequence=sequence, + sequence_label=sequence_label, + source_path=fasta_path, + summary_writer=summary_writer, + ) + + def main(): print(ASCII_ART) print(f"v{VERSION} \n") @@ -678,8 +1266,6 @@ def main(): except (OSError, ValueError) as error: print(f"Error reading annotation BED file(s): {error}", file=sys.stderr) sys.exit(2) - if args.command == "static": - _load_static_plotting() # -----------MUTUALLY EXCLUSIVE: INTERACTIVE OR STATIC MODE----------- if args.command == "interactive": print(INTERACTIVE_DEPRECATION_MESSAGE, file=sys.stderr) @@ -729,7 +1315,29 @@ def main(): config = json.load(f) _apply_static_config(args, config) + try: + args.processes = _validated_process_count(args.processes) + except ValueError as error: + print(f"Error: {error}.", file=sys.stderr) + sys.exit(2) + + # Plotting imports are comparatively expensive. Compute-only static + # runs should not import Plotnine or initialize Matplotlib at all. + if ( + (not args.no_plot) + or args.grid + or args.grid_only + or getattr(args, "load", None) + ): + _load_static_plotting() + # -----------INPUT COMMAND VALIDATION----------- + if args.sequence and getattr(args, "load", None): + print( + "Error: --sequence requires FASTA input; it cannot be used with --load.\n" + ) + sys.exit(2) + if args.plot_direction and getattr(args, "load", None): print( "Error: --plot-direction requires FASTA input because strand " @@ -907,10 +1515,12 @@ def main(): sys.exit(0) # -----------INPUT SEQUENCE VALIDATION----------- - seq_list = [] - fasta_list = args.fasta.copy() + # Repeating an input path cannot add a distinct sequence, and allowing it + # here would hash the same records twice while the path-keyed header map + # contains them only once. + fasta_list = list(dict.fromkeys(args.fasta)) fasta_headers = {} - for i in args.fasta: + for i in fasta_list.copy(): try: headers = getInputHeaders(i) fasta_headers[i] = headers @@ -918,14 +1528,25 @@ def main(): if len(headers) > 1: print(f"File {i} contains multiple fasta entries.\n") - seq_list.extend(headers) # Add all headers to seq_list - except Exception as e: print( f"\nUnable to open {i}. Please check it is correctly formatted or compressed...\n" ) fasta_list.remove(i) + if not fasta_list: + print("Error: no readable FASTA input files remain.", file=sys.stderr) + sys.exit(2) + + try: + fasta_headers, seq_list = _select_fasta_headers( + fasta_headers, getattr(args, "sequence", None) + ) + except ValueError as error: + print(f"Error: {error}.\n") + sys.exit(2) + fasta_list = [path for path in fasta_list if path in fasta_headers] + fasta_source_by_name = {} for fasta_path, headers in fasta_headers.items(): for header in headers: @@ -941,6 +1562,37 @@ def main(): print(f"Error: {error}.\n") sys.exit(2) + # Independent static self plots have no cross-record dependency. Stream + # them directly instead of retaining every positional hash in ``k_list``. + # The file-existence condition preserves unit tests and third-party callers + # that replace the legacy reader with an in-memory stub. + streaming_static_self = ( + args.command == "static" + and bool(fasta_list) + and not args.compare + and not args.compare_only + and not args.grid + and not args.grid_only + and args.compare_order == "sequential" + and all( + os.path.isfile(path) and os.path.getsize(path) > 0 + for path in fasta_list + ) + ) + if streaming_static_self: + try: + _run_streaming_static_self( + args, + fasta_list, + fasta_headers, + region_by_name, + summary_writer, + ) + except (OSError, UnicodeError, ValueError) as error: + print(f"Error processing FASTA input: {error}", file=sys.stderr) + sys.exit(2) + return + # -----------LOAD SEQUENCES INTO MEMORY----------- kmer_list = [] for i in fasta_list: @@ -1381,29 +2033,15 @@ def main(): ) seq_length = len(matrix_sequence) - win = args.window - res = args.resolution - if args.window: - # Change the resolution of each plot - res = math.ceil(seq_length / args.window) - else: - win = math.ceil(seq_length / args.resolution) - - if win < args.modimizer: - args.modimizer = win - if win < 10: - print(f"Error: sequence too small for analysis.\n") - print( - f"ModDotPlot requires a minimum window size of 10. Sequences less than 10Kbp will not work with ModDotPlot under normal resolution. We recommend rerunning ModDotPlot with --r {math.ceil(seq_length / 10)}.\n" - ) - sys.exit(0) - - seq_sparsity = round(win / args.modimizer) - if seq_sparsity <= args.modimizer: - seq_sparsity = 2 ** int(math.log2(seq_sparsity)) - else: - seq_sparsity = 2 ** (int(math.log2(seq_sparsity - 1)) + 1) - expectation = round(win / seq_sparsity) + try: + matrix_config = _matrix_config_for_length(seq_length, args) + except ValueError as error: + print(f"Error: {error}.\n") + sys.exit(2) + win = matrix_config.window_size + res = matrix_config.resolution + seq_sparsity = matrix_config.sparsity + expectation = matrix_config.expectation print(f"Computing self identity matrix for {seq_name}... \n") # TODO: Logging here @@ -1728,21 +2366,17 @@ def main(): max(larger_seq_end_pos, smaller_seq_end_pos), ) - win = args.window - res = args.resolution - if args.window: - res = math.ceil(smaller_length / args.window) - else: - win = math.ceil(smaller_length / args.resolution) - if win < args.modimizer: - args.modimizer = win - - seq_sparsity = round(win / args.modimizer) - if seq_sparsity <= args.modimizer: - seq_sparsity = 2 ** int(math.log2(seq_sparsity)) - else: - seq_sparsity = 2 ** (int(math.log2(seq_sparsity - 1)) + 1) - expectation = round(win / seq_sparsity) + try: + matrix_config = _matrix_config_for_length( + smaller_length, args + ) + except ValueError as error: + print(f"Error: {error}.\n") + sys.exit(2) + win = matrix_config.window_size + res = matrix_config.resolution + seq_sparsity = matrix_config.sparsity + expectation = matrix_config.expectation print( f"Computing pairwise identity matrix for {larger_seq_name} and {smaller_seq_name}... \n" ) diff --git a/src/moddotplot/native_render.py b/src/moddotplot/native_render.py index a79f5f3..a27d73d 100644 --- a/src/moddotplot/native_render.py +++ b/src/moddotplot/native_render.py @@ -17,6 +17,7 @@ from matplotlib.figure import Figure from matplotlib.text import Text from matplotlib.ticker import FuncFormatter +from matplotlib.transforms import Bbox import numpy as np import pandas as pd @@ -108,7 +109,7 @@ def tile_width(dataframe: pd.DataFrame) -> float: def rectangular_tile_vertices( dataframe: pd.DataFrame, *, transpose: bool = False -) -> Sequence[np.ndarray]: +) -> np.ndarray: """Build one square polygon per sparse input row. The squares are centered on the start columns and all use the maximum query @@ -119,7 +120,7 @@ def rectangular_tile_vertices( _require_columns(dataframe, ("q_st", "q_en", "r_st")) if dataframe.empty: - return [] + return np.empty((0, 4, 2), dtype=float) window = tile_width(dataframe) half_window = window / 2.0 @@ -130,18 +131,16 @@ def rectangular_tile_vertices( if transpose: query, reference = reference, query - return [ - np.asarray( - [ - (q - half_window, r - half_window), - (q + half_window, r - half_window), - (q + half_window, r + half_window), - (q - half_window, r + half_window), - ], - dtype=float, - ) - for q, r in zip(query, reference) - ] + vertices = np.empty((query.size, 4, 2), dtype=float) + vertices[:, 0, 0] = query - half_window + vertices[:, 0, 1] = reference - half_window + vertices[:, 1, 0] = query + half_window + vertices[:, 1, 1] = reference - half_window + vertices[:, 2, 0] = query + half_window + vertices[:, 2, 1] = reference + half_window + vertices[:, 3, 0] = query - half_window + vertices[:, 3, 1] = reference + half_window + return vertices def transform_triangle_points(points: np.ndarray) -> np.ndarray: @@ -380,7 +379,7 @@ def save_figure_pair( dpi: int, *, transparent: bool = False, - bbox_inches: Optional[str] = "tight", + bbox_inches: Optional[Union[str, Bbox]] = "tight", ) -> Tuple[Path, Path]: """Save one figure directly as PNG and SVG, PDF, or PostScript.""" diff --git a/src/moddotplot/parse_fasta.py b/src/moddotplot/parse_fasta.py index b0a75f3..9d531d0 100644 --- a/src/moddotplot/parse_fasta.py +++ b/src/moddotplot/parse_fasta.py @@ -8,10 +8,12 @@ TextIO, Tuple, ) +from bisect import bisect_right import sys import os import pickle import re +import struct import numpy as np import gzip @@ -28,11 +30,52 @@ class FastaIndexEntry(NamedTuple): line_width: int +class BgzfIndexEntry(NamedTuple): + compressed_offset: int + uncompressed_offset: int + + def _is_gzip(filename: str) -> bool: with open(filename, "rb") as probe: return probe.read(2) == b"\x1f\x8b" +def _is_bgzf(filename: str) -> bool: + """Return whether *filename* starts with a BGZF gzip member. + + BGZF is distinguished from ordinary gzip by the ``BC`` extra subfield in + each member header. Checking the header prevents an unrelated or stale + ``.gzi`` file from making a normal gzip stream look seekable. + """ + + try: + with open(filename, "rb") as compressed: + fixed_header = compressed.read(12) + if ( + len(fixed_header) != 12 + or fixed_header[:3] != b"\x1f\x8b\x08" + or not fixed_header[3] & 0x04 + ): + return False + extra_length = struct.unpack_from(" len(extra): + return False + if subfield_id == b"BC" and subfield_length == 2: + return True + offset = subfield_end + return False + + def _open_fasta_text(filename: str) -> TextIO: """Open a plain, gzip, or BGZF FASTA file as text. @@ -48,8 +91,6 @@ def _open_fasta_text(filename: str) -> TextIO: def _read_fasta_index(filename: str) -> Optional[List[FastaIndexEntry]]: """Read a fresh samtools-style ``.fai`` index when one is available.""" - if _is_gzip(filename): - return None index_path = f"{os.fspath(filename)}.fai" if not os.path.isfile(index_path): return None @@ -87,6 +128,70 @@ def _read_fasta_index(filename: str) -> Optional[List[FastaIndexEntry]]: return entries or None +def _read_bgzf_index(filename: str) -> Optional[List[BgzfIndexEntry]]: + """Read a fresh samtools-style ``.gzi`` index for a BGZF stream. + + ``.gzi`` stores compressed and uncompressed offsets for every BGZF block + after the first. The implicit origin is added here so callers can binary + search every uncompressed FASTA byte offset, including offsets in block 0. + Malformed, stale, or unrelated indexes are ignored and the FASTA reader can + transparently fall back to sequential gzip decompression. + """ + + if not _is_bgzf(filename): + return None + index_path = f"{os.fspath(filename)}.gzi" + if not os.path.isfile(index_path): + return None + try: + if os.path.getmtime(index_path) < os.path.getmtime(filename): + return None + index_size = os.path.getsize(index_path) + compressed_size = os.path.getsize(filename) + with open(index_path, "rb") as index: + count_bytes = index.read(8) + if len(count_bytes) != 8: + return None + entry_count = struct.unpack("= compressed_size + ): + return None + previous = entry + return entries + + +def supports_indexed_fasta_access(filename: str) -> bool: + """Return whether individual records can be fetched without a full scan. + + Plain FASTA needs a fresh ``.fai``. Compressed FASTA additionally needs + to be BGZF with a usable ``.gzi``. This predicate is intentionally + conservative because the chromosome process pool must never make every + worker decompress an ordinary gzip stream from the beginning. + """ + + if _read_fasta_index(filename) is None: + return False + return not _is_gzip(filename) or _read_bgzf_index(filename) is not None + + def _iter_fasta_headers(filename: str) -> Iterator[str]: """Yield FASTA identifiers without assembling or validating sequences.""" @@ -119,8 +224,38 @@ def _iter_fasta_headers(filename: str) -> Iterator[str]: raise ValueError(f"Invalid FASTA {filename!s}: no FASTA records found") +def _read_bgzf_range( + filename: str, + bgzf_index: Sequence[BgzfIndexEntry], + start: int, + size: int, +) -> bytes: + """Read an uncompressed byte range from a BGZF stream.""" + + if start < 0 or size < 0: + raise ValueError("BGZF byte ranges must be non-negative") + uncompressed_offsets = [entry.uncompressed_offset for entry in bgzf_index] + block_number = bisect_right(uncompressed_offsets, start) - 1 + if block_number < 0: + raise ValueError("BGZF index does not contain the start of the stream") + block = bgzf_index[block_number] + skip = start - block.uncompressed_offset + + with open(filename, "rb") as compressed: + compressed.seek(block.compressed_offset) + with gzip.GzipFile(fileobj=compressed, mode="rb") as uncompressed: + if len(uncompressed.read(skip)) != skip: + raise ValueError("BGZF index points past the end of the FASTA stream") + return uncompressed.read(size) + + def _fetch_indexed_region( - filename: str, entry: FastaIndexEntry, start: int, end: int + filename: str, + entry: FastaIndexEntry, + start: int, + end: int, + *, + bgzf_index: Optional[Sequence[BgzfIndexEntry]] = None, ) -> str: """Fetch one 1-based inclusive interval directly from an indexed FASTA.""" @@ -143,9 +278,21 @@ def _fetch_indexed_region( + (end_index // entry.line_bases) * entry.line_width + end_index % entry.line_bases ) - with open(filename, "rb") as fasta: - fasta.seek(start_byte) - raw_sequence = fasta.read(end_byte - start_byte + 1) + byte_count = end_byte - start_byte + 1 + if _is_gzip(filename): + if bgzf_index is None: + bgzf_index = _read_bgzf_index(filename) + if bgzf_index is None: + raise ValueError( + f"Compressed FASTA {filename!s} does not have a usable BGZF .gzi index" + ) + raw_sequence = _read_bgzf_range( + filename, bgzf_index, start_byte, byte_count + ) + else: + with open(filename, "rb") as fasta: + fasta.seek(start_byte) + raw_sequence = fasta.read(byte_count) sequence_bytes = raw_sequence.replace(b"\n", b"").replace(b"\r", b"") expected_length = end - start + 1 @@ -160,16 +307,25 @@ def _fetch_indexed_region( def _iter_selected_fasta_records( - filename: str, regions, single_record: bool + filename: str, regions, record_ids=None ) -> Iterator[Tuple[str, str]]: - """Stream selected intervals, stopping early for a one-record FASTA.""" + """Stream requested FASTA records and intervals in file order.""" + + requested_ids = None if record_ids is None else list(record_ids) + if requested_ids is not None and len(requested_ids) != len(set(requested_ids)): + raise ValueError("FASTA record identifiers must be unique") + requested_set = None if requested_ids is None else set(requested_ids) + if requested_set == set(): + return sequence_id = None sequence_parts = [] sequence_position = 0 selected_region = None selection_complete = False + collect_sequence = False seen_ids = set() + yielded_ids = set() def selected_sequence(): sequence = "".join(sequence_parts) @@ -186,8 +342,11 @@ def selected_sequence(): with _open_fasta_text(filename) as fasta: for line_number, raw_line in enumerate(fasta, start=1): if raw_line.startswith(">"): - if sequence_id is not None: + if sequence_id is not None and collect_sequence: yield sequence_id, selected_sequence() + yielded_ids.add(sequence_id) + if requested_set is not None and yielded_ids == requested_set: + return description = raw_line[1:].strip() if not description: @@ -203,7 +362,10 @@ def selected_sequence(): seen_ids.add(sequence_id) sequence_parts = [] sequence_position = 0 - selected_region = regions.get(sequence_id) if regions else None + collect_sequence = requested_set is None or sequence_id in requested_set + selected_region = ( + regions.get(sequence_id) if collect_sequence and regions else None + ) selection_complete = False continue @@ -225,6 +387,8 @@ def selected_sequence(): f"Invalid FASTA {filename!s}: whitespace within sequence data " f"at line {line_number}" ) + if not collect_sequence: + continue line_start = sequence_position + 1 line_end = sequence_position + len(line) @@ -239,7 +403,10 @@ def selected_sequence(): sequence_position = line_end if sequence_position >= end: selection_complete = True - if single_record: + if ( + requested_set is not None + and (yielded_ids | {sequence_id}) == requested_set + ): yield sequence_id, selected_sequence() return else: @@ -248,7 +415,108 @@ def selected_sequence(): if sequence_id is None: raise ValueError(f"Invalid FASTA {filename!s}: no FASTA records found") - yield sequence_id, selected_sequence() + if collect_sequence: + yield sequence_id, selected_sequence() + + +def _sequence_label(sequence_id: str, regions) -> str: + if regions and sequence_id in regions: + _name, start, end = regions[sequence_id] + return f"{sequence_id}:{start}-{end}" + return sequence_id + + +def iter_fasta_records( + filename: str, regions=None, record_ids=None +) -> Iterator[Tuple[str, str, str]]: + """Yield selected records as ``(identifier, sequence, display_label)``. + + A supplied ``record_ids`` sequence controls output order; with no explicit + selection, records retain index or FASTA file order. Fresh ``.fai`` indexes + provide record metadata for every FASTA. Plain FASTA and BGZF inputs with a + valid ``.gzi`` are fetched directly, while ordinary gzip and unusable BGZF + indexes retain the sequential decompression fallback. + """ + + requested_ids = None if record_ids is None else list(record_ids) + if requested_ids is not None and len(requested_ids) != len(set(requested_ids)): + raise ValueError("FASTA record identifiers must be unique") + + fasta_index = _read_fasta_index(filename) + indexed_entries = ( + {entry.name: entry for entry in fasta_index} if fasta_index else None + ) + if indexed_entries is not None: + selected_ids = ( + list(indexed_entries) if requested_ids is None else requested_ids + ) + missing_ids = [ + sequence_id + for sequence_id in selected_ids + if sequence_id not in indexed_entries + ] + if missing_ids: + formatted = ", ".join(repr(sequence_id) for sequence_id in missing_ids) + raise ValueError(f"FASTA record(s) not found: {formatted}") + + compressed = _is_gzip(filename) + bgzf_index = _read_bgzf_index(filename) if compressed else None + if not compressed or bgzf_index is not None: + for sequence_id in selected_ids: + entry = indexed_entries[sequence_id] + if regions and sequence_id in regions: + _name, start, end = regions[sequence_id] + else: + start, end = 1, entry.length + sequence = _fetch_indexed_region( + filename, + entry, + start, + end, + bgzf_index=bgzf_index, + ) + yield sequence_id, sequence, _sequence_label(sequence_id, regions) + return + else: + selected_ids = requested_ids + + streamed_records = _iter_selected_fasta_records( + filename, regions, record_ids=selected_ids + ) + if selected_ids is None: + for sequence_id, sequence in streamed_records: + yield sequence_id, sequence, _sequence_label(sequence_id, regions) + return + + # Streaming naturally discovers records in file order. Buffer only records + # that precede the next explicitly requested identifier so the public API + # can preserve caller order without requiring a separate header scan. + pending = {} + seen_selected = set() + selected_position = 0 + for sequence_id, sequence in streamed_records: + pending[sequence_id] = sequence + seen_selected.add(sequence_id) + while ( + selected_position < len(selected_ids) + and selected_ids[selected_position] in pending + ): + selected_id = selected_ids[selected_position] + yield ( + selected_id, + pending.pop(selected_id), + _sequence_label(selected_id, regions), + ) + selected_position += 1 + + if selected_position != len(selected_ids): + missing_ids = [ + sequence_id + for sequence_id in selected_ids + if sequence_id not in seen_selected + ] + formatted = ", ".join(repr(sequence_id) for sequence_id in missing_ids) + raise ValueError(f"FASTA record(s) not found: {formatted}") def _iter_fasta_records(filename: str) -> Iterator[Tuple[str, str]]: @@ -521,54 +789,10 @@ def readKmersFromFile( Given a filename and an integer k, returns a list of all k-mers found in the sequences in the file. """ all_kmers = [] - record_ids = ( - list(record_ids) if record_ids is not None else getInputHeaders(filename) - ) - fasta_index = _read_fasta_index(filename) - indexed_entries = ( - {entry.name: entry for entry in fasta_index} - if fasta_index and [entry.name for entry in fasta_index] == record_ids - else None - ) - - if indexed_entries is not None: - - def indexed_records(): - for seq_id in record_ids: - entry = indexed_entries[seq_id] - if regions and seq_id in regions: - _name, start, end = regions[seq_id] - else: - start, end = 1, entry.length - sequence_label = ( - f"{seq_id}:{start}-{end}" - if regions and seq_id in regions - else seq_id - ) - print(f"Retrieving k-mers from {sequence_label}.... \n") - sequence = _fetch_indexed_region(filename, entry, start, end) - yield seq_id, sequence, sequence_label - - selected_records = indexed_records() - else: - selected_records = ( - ( - seq_id, - sequence, - ( - f"{seq_id}:{regions[seq_id][1]}-{regions[seq_id][2]}" - if regions and seq_id in regions - else seq_id - ), - ) - for seq_id, sequence in _iter_selected_fasta_records( - filename, regions, single_record=len(record_ids) == 1 - ) - ) - - for seq_id, sequence, sequence_label in selected_records: - if indexed_entries is None: - print(f"Retrieving k-mers from {sequence_label}.... \n") + for seq_id, sequence, sequence_label in iter_fasta_records( + filename, regions=regions, record_ids=record_ids + ): + print(f"Retrieving k-mers from {sequence_label}.... \n") if len(sequence) < ksize: if regions and seq_id in regions: _name, start, end = regions[seq_id] @@ -605,4 +829,7 @@ def getInputHeaders(filename: str) -> List[str]: def getInputSeqLength(filename: str) -> List[int]: + fasta_index = _read_fasta_index(filename) + if fasta_index is not None: + return [entry.length for entry in fasta_index] return [len(sequence) for _sequence_id, sequence in _iter_fasta_records(filename)] diff --git a/src/moddotplot/static_plots.py b/src/moddotplot/static_plots.py index 6edee35..ae0b8ec 100755 --- a/src/moddotplot/static_plots.py +++ b/src/moddotplot/static_plots.py @@ -92,14 +92,60 @@ def _save_plot(plot, **kwargs): ggsave(plot + _plot_font_theme(FALLBACK_FONT_FAMILY), **kwargs) +def _draw_and_save_plot_pair( + plot, + output_prefix, + *, + width, + height, + dpi, + vector_format, +): + """Build one Plotnine figure and save both raster and vector outputs. + + ``ggsave`` redraws a plot for every requested format. Large tile plots and + histograms therefore paid their complete scale/layout/rasterization cost + twice. Drawing once also guarantees that both files contain the same axes, + labels, and tile realization. + """ + + def draw(family): + styled = ( + plot + + _plot_font_theme(family) + + theme(figure_size=(float(width), float(height)), dpi=int(dpi)) + ) + return styled.draw(show=False) + + try: + figure = draw(DEFAULT_FONT_FAMILY) + except RuntimeError as error: + if not is_glyph_loading_error(error): + raise + figure = draw(FALLBACK_FONT_FAMILY) + try: + return save_figure_pair( + figure, + output_prefix, + vector_format, + dpi, + # Match plotnine/ggsave's requested physical canvas exactly. A + # tight bounding box changes both the raster dimensions and plot + # framing (for example, 3 in at 96 dpi no longer yields 288 px). + bbox_inches=figure.bbox_inches, + ) + finally: + plt.close(figure) + + def display_sequence_name(name): """Return a sequence name without appended region coordinates.""" return REGION_SUFFIX_PATTERN.sub("", str(name)) -def _fit_grid_sequence_labels(figure, axes, minimum_size=MIN_TEXT_SIZE): - """Fit grid headings without shrinking them below a readable size.""" +def _fit_grid_sequence_labels(figure, axes): + """Fit grid headings inside the existing figure canvas.""" figure.canvas.draw() renderer = figure.canvas.get_renderer() @@ -131,16 +177,8 @@ def _fit_grid_sequence_labels(figure, axes, minimum_size=MIN_TEXT_SIZE): if scale >= 1.0: return - sizes = [artist.get_fontsize() for artist in artists] - smallest_scaled_size = min(size * scale for size in sizes) - if smallest_scaled_size < minimum_size: - enlargement = minimum_size / smallest_scaled_size - width, height = figure.get_size_inches() - figure.set_size_inches(width * enlargement, height * enlargement, forward=True) - scale = min(1.0, scale * enlargement) - for artist in artists: - artist.set_fontsize(max(minimum_size, artist.get_fontsize() * scale)) + artist.set_fontsize(artist.get_fontsize() * scale) def _resolve_native_colors(palette, palette_orientation, custom_colors=None): @@ -289,14 +327,25 @@ def save_outputs(): def check_st_en_equality(df): - unequal_rows = df[(df["q_st"] != df["r_st"]) | (df["q_en"] != df["r_en"])] - unequal_rows.loc[:, ["q_en", "r_en", "q_st", "r_st"]] = unequal_rows[ - ["r_en", "q_en", "r_st", "q_st"] - ].values + """Complete a self-comparison across its diagonal without duplicate tiles.""" - df = pd.concat([df, unequal_rows], ignore_index=True) + if df.empty: + return df.copy() - return df + coordinate_columns = ["q_st", "q_en", "r_st", "r_en"] + unequal_rows = df[(df["q_st"] != df["r_st"]) | (df["q_en"] != df["r_en"])].copy() + if unequal_rows.empty: + return df.copy() + + mirrored_rows = unequal_rows.copy() + mirrored_rows.loc[:, coordinate_columns] = unequal_rows[ + ["r_st", "r_en", "q_st", "q_en"] + ].to_numpy() + + existing_coordinates = pd.MultiIndex.from_frame(df[coordinate_columns]) + mirrored_coordinates = pd.MultiIndex.from_frame(mirrored_rows[coordinate_columns]) + mirrored_rows = mirrored_rows.loc[~mirrored_coordinates.isin(existing_coordinates)] + return pd.concat([df, mirrored_rows], ignore_index=True) def make_k(vals): @@ -362,6 +411,15 @@ def get_colors(sdf, ncolors, is_freq, custom_breakpoints): raise ValueError("Breakpoints must contain only finite numbers") if np.any(np.diff(breaks) <= 0): raise ValueError("Breakpoints must be strictly increasing") + values = np.asarray(sdf["perID_by_events"], dtype=np.float64) + if values.size and ( + not np.all(np.isfinite(values)) + or values.min() < breaks[0] + or values.max() > breaks[-1] + ): + raise ValueError( + "Breakpoints must cover all finite identity values in the plot" + ) labels = np.arange(len(breaks) - 1) # A dataset containing only 100% identity creates repeated default bin # edges; frequency bins likewise collapse to one edge when every value is @@ -1456,19 +1514,215 @@ def make_hist(sdf, palette, palette_orientation, custom_colors, custom_breakpoin return p +def _missing_symmetric_rows(dataframe): + """Return rows whose coordinate-transposed counterpart is absent. + + FASTA self matrices normally contain only one triangle. Drawing those + rows a second time with transposed coordinates completes the full plot + without allocating a second, mirrored dataframe. Loaded BEDPE files may + already contain both triangles, so the general path checks coordinate + membership before selecting rows to mirror. + """ + + if dataframe.empty: + return dataframe.iloc[0:0] + + q_start = dataframe["q_st"] + q_end = dataframe["q_en"] + r_start = dataframe["r_st"] + r_end = dataframe["r_en"] + off_diagonal = (q_start != r_start) | (q_end != r_end) + if not off_diagonal.any(): + return dataframe.iloc[0:0] + + # The mirror collection needs coordinates and its resolved color only; do + # not copy names, identity estimates, or plotting-helper columns from a + # potentially large dataframe. + render_columns = ["q_st", "q_en", "r_st", "r_en"] + render_columns.extend( + column + for column in ("discrete", "direction", DIRECTION_ANI_COLUMN) + if column in dataframe.columns + ) + + # The production FASTA path is strictly triangular. Avoid building two + # MultiIndexes for that common case; a strict start-coordinate ordering + # proves that no transposed off-diagonal row can already be present. + start_difference = ( + q_start.loc[off_diagonal].to_numpy() + - r_start.loc[off_diagonal].to_numpy() + ) + if np.all(start_difference < 0) or np.all(start_difference > 0): + return dataframe.loc[off_diagonal, render_columns] + + existing = pd.MultiIndex.from_arrays([q_start, q_end, r_start, r_end]) + mirrored = pd.MultiIndex.from_arrays( + [r_start, r_end, q_start, q_end] + ) + return dataframe.loc[ + off_diagonal & ~mirrored.isin(existing), render_columns + ] + + +def _full_plot_limits(dataframe, requested_limit): + """Resolve exact full-plot bounds, including an empty sparse matrix.""" + + requested_bounds = _requested_axis_bounds(requested_limit) + if requested_bounds is not None: + return requested_bounds + if not dataframe.empty: + return _data_axis_limits(dataframe, requested_limit) + if requested_limit: + return 0.0, float(requested_limit) + raise ValueError( + "Cannot infer full-plot bounds from an empty identity table; " + "provide explicit axis bounds" + ) + + +def _build_full_figure( + sdf, + name_x, + name_y, + palette, + palette_orientation, + custom_colors, + axes_labels, + xlim, + deraster, + width, + is_pairwise, +): + """Build a sparse full or comparative dotplot with native Matplotlib. + + The tile geometry intentionally retains the historical Plotnine contract: + BEDPE start coordinates are tile centers and ``q_en - q_st`` is the tile + width. The default rasterizes only the sparse tile collections in vector + output; ``--deraster`` leaves each tile as vector geometry. + """ + + display_x = display_sequence_name(name_x) + display_y = display_sequence_name(name_y) + title = ( + f"Comparative Plot: {display_x} vs {display_y}" + if is_pairwise + else f"Self-Identity Plot: {display_x}" + ) + region_start, region_end = _full_plot_limits(sdf, xlim) + breaks = ( + [float(value) for value in axes_labels] + if axes_labels + else generate_breaks(int(region_start), int(region_end)) + ) + + styled, direction_colors, direction_column = _direction_ani_style(sdf) + colors = direction_colors or _resolve_native_colors( + palette, palette_orientation, custom_colors + ) + color_column = direction_column or "discrete" + + figure, axis = plt.subplots(figsize=(float(width), float(width))) + try: + draw_rectangular_tiles( + axis, + styled, + colors, + color_column=color_column, + rasterized=not deraster, + ) + if not is_pairwise: + missing_rows = _missing_symmetric_rows(styled) + if not missing_rows.empty: + draw_rectangular_tiles( + axis, + missing_rows, + colors, + color_column=color_column, + transpose=True, + rasterized=not deraster, + ) + + configure_dotplot_axis( + axis, + region_start, + region_end, + breaks=breaks, + ) + _divisor, unit = genomic_scale(region_end) + axis.set_xlabel( + f"Genomic Position ({unit})", + fontsize=clamped_font_size(width, 2.8), + fontfamily=DEFAULT_FONT_FAMILY, + ) + axis.tick_params( + axis="both", + labelsize=clamped_font_size(width, 2.0), + length=max(3.5, float(width)), + colors="black", + ) + axis.grid(False) + axis.set_facecolor("none") + for spine in axis.spines.values(): + spine.set_color("black") + + # Plotnine's one-cell facet supplies a query label above the panel and + # a reference label at its right edge. Retain those identifiers while + # placing the descriptive title independently above them. + axis.set_title( + display_x, + fontsize=clamped_font_size(width, 1.2), + fontfamily=DEFAULT_FONT_FAMILY, + pad=5, + ) + axis.set_ylabel( + display_y, + fontsize=clamped_font_size(width, 1.2), + fontfamily=DEFAULT_FONT_FAMILY, + rotation=-90, + labelpad=16, + ) + axis.yaxis.set_label_position("right") + + title_size = 2.0 * float(width) + if len(title) > 80: + title_size = float(width) + elif len(title) > 50: + title_size = 1.5 * float(width) + figure.suptitle( + title, + fontsize=max(MIN_TITLE_SIZE, title_size), + fontfamily=DEFAULT_FONT_FAMILY, + y=0.975, + ) + figure.subplots_adjust( + left=0.14, + right=0.87, + bottom=0.14, + top=0.84, + ) + set_figure_font_family(figure, DEFAULT_FONT_FAMILY) + except Exception: + plt.close(figure) + raise + return figure + + def _triangle_limits(sdf, xlim): - if sdf.empty: - raise ValueError("Cannot render a triangle plot without identity tiles") requested_bounds = _requested_axis_bounds(xlim) if requested_bounds is not None: region_start, region_end = requested_bounds - else: + elif not sdf.empty: region_start = max(float(sdf["q_st"].min()), float(sdf["r_st"].min())) region_end = max( float(sdf["q_en"].max()), float(sdf["r_en"].max()), float(xlim or 0), ) + else: + raise ValueError( + "Cannot infer triangle bounds from an empty identity table; " + "provide explicit axis bounds" + ) if region_end <= region_start: raise ValueError("Triangle plot end must be greater than its start") return region_start, region_end @@ -1763,6 +2017,8 @@ def _build_grid_figure( ) if dataframe is not None and not dataframe.empty: + if row_name == column_name: + dataframe = check_st_en_equality(dataframe) ( dataframe, direction_colors, @@ -1797,27 +2053,47 @@ def _build_grid_figure( display_sequence_name(row_name), fontsize=heading_size, fontfamily=DEFAULT_FONT_FAMILY, + labelpad=2, ) - figure.supxlabel( - f"Genomic Position ({axis_unit})", - fontsize=axis_title_size, - fontfamily=DEFAULT_FONT_FAMILY, - ) - figure.supylabel( - f"Genomic Position ({axis_unit})", + # Keep both shared genomic-axis titles with the bottom-left cell, where + # both sets of numeric tick labels are visible. Figure-wide titles sit + # far from that cell and enlarge tightly cropped output canvases. + bottom_left_axis = axes[-1, 0] + axis_title = f"Genomic Position ({axis_unit})" + bottom_left_axis.set_xlabel( + axis_title, fontsize=axis_title_size, fontfamily=DEFAULT_FONT_FAMILY, + labelpad=2, ) + bottom_margin = max(0.12, 0.35 / figure_width) + left_margin_inches = 0.72 + max(0.0, tick_size - MIN_TEXT_SIZE) / 72.0 + left_margin = max(0.14, left_margin_inches / figure_width) figure.subplots_adjust( - left=0.14, + left=left_margin, right=0.98, - bottom=0.12, + bottom=bottom_margin, top=0.92, wspace=0.08, hspace=0.08, ) _fit_grid_sequence_labels(figure, axes) + vertical_title = bottom_left_axis.annotate( + axis_title, + xy=(0, 0.5), + xycoords=bottom_left_axis.yaxis.label, + xytext=(-4, 0), + textcoords="offset points", + ha="center", + va="center", + rotation=90, + rotation_mode="anchor", + fontsize=axis_title_size, + fontfamily=DEFAULT_FONT_FAMILY, + annotation_clip=False, + ) + vertical_title.set_gid("grid-y-axis-title") set_figure_font_family(figure, DEFAULT_FONT_FAMILY) except Exception: plt.close(figure) @@ -1878,7 +2154,7 @@ def create_grid( grid_prefix, vector_format, dpi, - bbox_inches="tight", + bbox_inches=figure.bbox_inches, ) finally: plt.close(figure) @@ -1979,41 +2255,32 @@ def create_plots( print("Skipping annotation track generation.\n") if is_pairwise: - heatmap = make_dot( - sdf, - name_x, - name_y, - palette, - palette_orientation, - custom_colors, - axes_labels, - axes_tick_number, - xlim, - deraster, - width, - True, - ) print(f"Creating plots and saving to {plot_filename}...\n") full_suffix = "_DIRECTION_FULL" if directional else "_COMPARE" hist_suffix = "_DIRECTION_HIST" if directional else "_COMPARE_HIST" - _save_plot( - heatmap, - width=width, - height=width, - dpi=dpi, - format=vector_format, - filename=f"{plot_filename}{full_suffix}.{vector_format}", - verbose=False, - ) - _save_plot( - heatmap, + full_figure = _build_full_figure( + sdf=sdf, + name_x=name_x, + name_y=name_y, + palette=palette, + palette_orientation=palette_orientation, + custom_colors=custom_colors, + axes_labels=axes_labels, + xlim=xlim, + deraster=deraster, width=width, - height=width, - dpi=dpi, - format="png", - filename=f"{plot_filename}{full_suffix}.png", - verbose=False, + is_pairwise=True, ) + try: + save_figure_pair( + full_figure, + f"{plot_filename}{full_suffix}", + vector_format, + dpi, + bbox_inches=full_figure.bbox_inches, + ) + finally: + plt.close(full_figure) created_files.extend( [ f"{plot_filename}{full_suffix}.{vector_format}", @@ -2021,14 +2288,13 @@ def create_plots( ] ) if not no_hist: - _save_plot( + _draw_and_save_plot_pair( histy, + f"{plot_filename}{hist_suffix}", width=3, height=3, dpi=dpi, - format=vector_format, - filename=f"{plot_filename}{hist_suffix}.{vector_format}", - verbose=False, + vector_format=vector_format, ) created_files.extend( [ @@ -2036,27 +2302,11 @@ def create_plots( f"{plot_filename}{hist_suffix}.png", ] ) - _save_plot( - histy, - width=3, - height=3, - dpi=dpi, - format="png", - filename=f"{plot_filename}{hist_suffix}.png", - verbose=False, - ) - try: - if not heatmap.data: - print( - f"{plot_filename} comparative plots and histogram saved sucessfully. \n" - ) - return created_files - except ValueError: + if not no_hist: print( f"{plot_filename} comparative plots and histogram saved sucessfully. \n" ) - return created_files - if no_hist: + else: print( f"{plot_filename}{full_suffix}.{vector_format} and " f"{plot_filename}{full_suffix}.png saved sucessfully. \n" @@ -2067,47 +2317,38 @@ def create_plots( print( f"Producing dotplots with derasterization turned off. This may take a while...\n" ) - full_plot = make_dot( - check_st_en_equality(sdf), - name_x, - name_y, - palette, - palette_orientation, - custom_colors, - axes_labels, - axes_tick_number, - xlim, - deraster, - width, - False, - ) full_suffix = "_DIRECTION_FULL" if directional else "_FULL" tri_suffix = "_DIRECTION_TRI" if directional else "_TRI" hist_suffix = "_DIRECTION_HIST" if directional else "_HIST" - _save_plot( - full_plot, + full_figure = _build_full_figure( + sdf=sdf, + name_x=name_x, + name_y=name_y, + palette=palette, + palette_orientation=palette_orientation, + custom_colors=custom_colors, + axes_labels=axes_labels, + xlim=xlim, + deraster=deraster, width=width, - height=width, - dpi=dpi, - format=vector_format, - filename=f"{plot_filename}{full_suffix}.{vector_format}", - verbose=False, + is_pairwise=False, ) + try: + save_figure_pair( + full_figure, + f"{plot_filename}{full_suffix}", + vector_format, + dpi, + bbox_inches=full_figure.bbox_inches, + ) + finally: + plt.close(full_figure) created_files.extend( [ f"{plot_filename}{full_suffix}.{vector_format}", f"{plot_filename}{full_suffix}.png", ] ) - _save_plot( - full_plot, - width=width, - height=width, - dpi=dpi, - format="png", - filename=f"{plot_filename}{full_suffix}.png", - verbose=False, - ) tri_prefix = f"{plot_filename}{tri_suffix}" triangle_figure = _build_triangle_figure( sdf=sdf, @@ -2168,23 +2409,13 @@ def create_plots( f"Triangle plots and full plots for {plot_filename} saved sucessfully. \n" ) else: - _save_plot( + _draw_and_save_plot_pair( histy, + f"{plot_filename}{hist_suffix}", width=3, height=3, dpi=dpi, - format=vector_format, - filename=plot_filename + f"{hist_suffix}.{vector_format}", - verbose=False, - ) - _save_plot( - histy, - width=3, - height=3, - dpi=dpi, - format="png", - filename=plot_filename + f"{hist_suffix}.png", - verbose=False, + vector_format=vector_format, ) created_files.extend( [ diff --git a/tests/test_algorithms.py b/tests/test_algorithms.py index 4483ac9..7d25283 100644 --- a/tests/test_algorithms.py +++ b/tests/test_algorithms.py @@ -2,9 +2,12 @@ import pytest from moddotplot.estimate_identity import ( + BEDPE_HEADER, containment_neighbors, convertMatrixToBed, + convertMatrixToBedDataFrame, createSelfMatrix, + iterMatrixToBedChunks, pairwiseContainmentMatrix, partitionOverlaps, populateModimizers, @@ -12,6 +15,50 @@ from moddotplot.parse_fasta import generateKmersFromFasta, printProgressBar +def _scalar_bed_reference( + matrix, + window_size, + id_threshold, + x_name, + y_name, + self_identity, + x_offset, + y_offset, + x_end=None, + y_end=None, +): + """Original scalar implementation used as a parity oracle.""" + + bed = [BEDPE_HEADER] + for x in range(matrix.shape[0]): + for y in range(matrix.shape[1]): + value = matrix[x, y] + if self_identity and x > y: + continue + if not value >= id_threshold / 100: + continue + start_x = x * window_size + x_offset + end_x = start_x + window_size - 1 + start_y = y * window_size + y_offset + end_y = start_y + window_size - 1 + if x_end is not None: + end_x = min(end_x, x_end) + if y_end is not None: + end_y = min(end_y, y_end) + bed.append( + ( + x_name, + int(start_x), + int(end_x), + y_name, + int(start_y), + int(end_y), + float(value), + ) + ) + return bed + + def test_bed_conversion_clamps_partial_windows_to_exact_region_end(): bed = convertMatrixToBed( np.ones((2, 2)), @@ -30,6 +77,186 @@ def test_bed_conversion_clamps_partial_windows_to_exact_region_end(): assert max(row[5] for row in bed[1:]) == 350 +@pytest.mark.parametrize("self_identity", [False, True]) +def test_vectorized_bed_conversion_preserves_row_order_and_values(self_identity): + matrix = np.array( + [ + [0.0, 86.5, 0.0, 92.25], + [87.0, 0.0, 99.0, 0.0], + [0.0, 91.0, 88.0, 0.0], + ] + ) + expected = [ + ( + "#query_name", + "query_start", + "query_end", + "reference_name", + "reference_start", + "reference_end", + "perID_by_events", + ) + ] + for x in range(matrix.shape[0]): + for y in range(matrix.shape[1]): + value = matrix[x, y] + if (not self_identity or x <= y) and value >= 86 / 100: + expected.append( + ( + "x", + x * 100 + 11, + min(x * 100 + 110, 250), + "y", + y * 100 + 21, + min(y * 100 + 120, 350), + float(value), + ) + ) + + assert ( + convertMatrixToBed( + matrix, + window_size=100, + id_threshold=86, + x_name="x", + y_name="y", + self_identity=self_identity, + x_offset=11, + y_offset=21, + x_end=250, + y_end=350, + ) + == expected + ) + + +def test_bed_conversion_returns_only_header_when_no_tiles_pass(): + bed = convertMatrixToBed( + np.zeros((1_000, 1_000)), + window_size=100, + id_threshold=86, + x_name="x", + y_name="y", + self_identity=True, + x_offset=0, + y_offset=0, + ) + + assert len(bed) == 1 + + +@pytest.mark.parametrize("self_identity", [False, True]) +@pytest.mark.parametrize("max_chunk_cells", [1, 7, 64, 10_000]) +def test_chunked_bed_conversion_matches_scalar_reference_randomized( + self_identity, max_chunk_cells +): + rng = np.random.default_rng(709) + # Exercise a non-contiguous view as well as values immediately around the + # legacy threshold. NaN must remain filtered by the comparison. + source = rng.uniform(0.0, 1.5, size=(14, 24)) + matrix = source[::2, 1::2] + matrix[0, :4] = [0.859999, 0.86, 0.860001, np.nan] + kwargs = dict( + window_size=37, + id_threshold=86, + x_name="query", + y_name="reference", + self_identity=self_identity, + x_offset=13, + y_offset=29, + x_end=251, + y_end=411, + ) + + expected = _scalar_bed_reference(matrix, **kwargs) + actual = convertMatrixToBed(matrix, max_chunk_cells=max_chunk_cells, **kwargs) + + assert actual == expected + + +def test_bed_chunks_are_strictly_bounded_across_row_and_column_boundaries(): + matrix = np.arange(30, dtype=float).reshape(3, 10) + chunks = list( + iterMatrixToBedChunks( + matrix, + window_size=10, + id_threshold=0, + x_name="x", + y_name="y", + self_identity=False, + x_offset=0, + y_offset=0, + max_chunk_cells=4, + ) + ) + + # A ten-column row must be split because it is wider than the cap. The + # concatenated chunks still follow exact C order across every boundary. + assert [len(chunk) for chunk in chunks] == [4, 4, 2] * 3 + assert all(tuple(chunk.columns) == BEDPE_HEADER for chunk in chunks) + observed = [ + tuple(row) + for chunk in chunks + for row in chunk.itertuples(index=False, name=None) + ] + assert observed == _scalar_bed_reference(matrix, 10, 0, "x", "y", False, 0, 0)[1:] + + +def test_bed_dataframe_helper_preserves_columns_clipping_and_float_values(): + matrix = np.array([[0.85, 0.9, 1.25], [0.95, 0.1, 1.5]], dtype=np.float32) + kwargs = dict( + window_size=10.5, + id_threshold=86, + x_name="chrQ", + y_name="chrR", + self_identity=False, + x_offset=-3.25, + y_offset=101.75, + x_end=9.5, + y_end=119.25, + max_chunk_cells=2, + ) + + frame = convertMatrixToBedDataFrame(matrix, **kwargs) + expected = _scalar_bed_reference( + matrix, + **{key: value for key, value in kwargs.items() if key != "max_chunk_cells"}, + ) + + assert tuple(frame.columns) == BEDPE_HEADER + assert list(frame.itertuples(index=False, name=None)) == expected[1:] + assert frame["perID_by_events"].dtype == np.dtype(float) + + +@pytest.mark.parametrize("shape", [(0, 4), (4, 0), (0, 0)]) +def test_empty_bed_dataframe_has_exact_schema(shape): + frame = convertMatrixToBedDataFrame( + np.empty(shape), 100, 86, "x", "y", True, 0, 0, max_chunk_cells=1 + ) + + assert frame.empty + assert tuple(frame.columns) == BEDPE_HEADER + assert convertMatrixToBed( + np.empty(shape), 100, 86, "x", "y", True, 0, 0, max_chunk_cells=1 + ) == [BEDPE_HEADER] + + +@pytest.mark.parametrize("max_chunk_cells", [0, -1, 1.5, True]) +def test_bed_conversion_rejects_invalid_chunk_bound(max_chunk_cells): + with pytest.raises((TypeError, ValueError), match="positive integer"): + convertMatrixToBed( + np.ones((1, 1)), + 100, + 86, + "x", + "y", + False, + 0, + 0, + max_chunk_cells=max_chunk_cells, + ) + + def test_populate_modimizers_returns_denser_recursive_fallback(): result = populateModimizers( partition=[1, 2, 3, 4], diff --git a/tests/test_annotation_track.py b/tests/test_annotation_track.py index ee42f01..2e6adad 100644 --- a/tests/test_annotation_track.py +++ b/tests/test_annotation_track.py @@ -284,14 +284,23 @@ def _stub_create_plots_dependencies(monkeypatch, *, directional=False): ) monkeypatch.setattr(static_plots, "make_dot", lambda *_args, **_kwargs: _FakePlot()) - def fake_ggsave(*_args, **kwargs): - output = Path(kwargs["filename"]) - if output.suffix == ".png": - output.write_bytes(b"png") - else: - output.write_text('') - - monkeypatch.setattr(static_plots, "ggsave", fake_ggsave) + def fake_plot_pair( + _plot, + output_prefix, + *, + width, + height, + dpi, + vector_format, + ): + del width, height, dpi + png = Path(f"{output_prefix}.png") + vector = Path(f"{output_prefix}.{vector_format}") + png.write_bytes(b"png") + vector.write_text('') + return png, vector + + monkeypatch.setattr(static_plots, "_draw_and_save_plot_pair", fake_plot_pair) def _run_create_plots(output_dir, annotation, vector_format="svg"): diff --git a/tests/test_cli_integration.py b/tests/test_cli_integration.py index f8578a7..fd32770 100644 --- a/tests/test_cli_integration.py +++ b/tests/test_cli_integration.py @@ -39,6 +39,62 @@ def _write_multifasta(path): ) +def _write_indexed_multifasta(path): + records = [ + ("alpha", "ACGT" * 300), + ("beta", "ACGT" * 275), + ("gamma", "ACGT" * 250), + ] + fasta_parts = [] + index_lines = [] + byte_offset = 0 + for name, sequence in records: + header = f">{name}\n" + fasta_parts.extend((header, sequence, "\n")) + sequence_offset = byte_offset + len(header) + index_lines.append( + f"{name}\t{len(sequence)}\t{sequence_offset}\t" + f"{len(sequence)}\t{len(sequence) + 1}\n" + ) + byte_offset += len(header) + len(sequence) + 1 + path.write_text("".join(fasta_parts)) + path.with_name(path.name + ".fai").write_text("".join(index_lines)) + + +def test_indexed_self_plots_are_identical_with_spawned_chromosome_workers(tmp_path): + fasta = tmp_path / "indexed.fa" + serial_output = tmp_path / "serial" + parallel_output = tmp_path / "parallel" + _write_indexed_multifasta(fasta) + common = ( + "--fasta", + fasta, + "--window", + 100, + "--modimizer", + 10, + "--identity", + 80, + "--no-plot", + ) + + serial = _run_cli(*common, "--processes", 1, "--output-dir", serial_output) + parallel = _run_cli(*common, "--processes", 2, "--output-dir", parallel_output) + + assert serial.returncode == 0, serial.stderr + serial.stdout + assert parallel.returncode == 0, parallel.stderr + parallel.stdout + assert "with 2 chromosome workers" in parallel.stdout + serial_files = { + path.relative_to(serial_output): path.read_bytes() + for path in serial_output.rglob("*.bedpe") + } + parallel_files = { + path.relative_to(parallel_output): path.read_bytes() + for path in parallel_output.rglob("*.bedpe") + } + assert parallel_files == serial_files + + def test_static_cli_computes_all_self_and_pairwise_outputs(tmp_path): fasta = tmp_path / "three.fa" output = tmp_path / "static" @@ -97,6 +153,137 @@ def test_omitted_subcommand_runs_static_mode(tmp_path): assert (output / "alpha" / "alpha.bedpe").is_file() +def test_static_duplicate_fasta_path_is_processed_once(tmp_path): + fasta = tmp_path / "one.fa" + fasta.write_text(">alpha\n" + "ACGT" * 300 + "\n") + output = tmp_path / "deduplicated-input" + + result = _run_cli( + "--fasta", + fasta, + fasta, + "--window", + 100, + "--modimizer", + 10, + "--identity", + 80, + "--no-plot", + "--output-dir", + output, + ) + + assert result.returncode == 0, result.stderr + result.stdout + assert result.stdout.count("Computing self identity matrix for alpha") == 1 + assert [path.relative_to(output) for path in output.rglob("*.bedpe")] == [ + Path("alpha/alpha.bedpe") + ] + + +def test_static_sequence_selection_builds_grid_from_only_requested_records(tmp_path): + fasta = tmp_path / "chromosomes.fa" + fasta.write_text( + ">Chr1\n" + + "ACGT" * 300 + + "\n>Chr2\n" + + "ACGT" * 275 + + "\n>Chr3\n" + + "ACGT" * 250 + + "\n" + ) + output = tmp_path / "selected-grid" + + result = _run_cli( + "-f", + fasta, + "-s", + "chr1", + "chr2", + "--grid", + "--window", + 100, + "--modimizer", + 10, + "--identity", + 80, + "--no-plot", + "--output-dir", + output, + ) + + assert result.returncode == 0, result.stderr + result.stdout + assert "Creating a 2x2 grid" in result.stdout + assert "Chr1 k-mers retrieved" in result.stdout + assert "Chr2 k-mers retrieved" in result.stdout + assert "Chr3 k-mers retrieved" not in result.stdout + bedpe_files = sorted(path.relative_to(output) for path in output.rglob("*.bedpe")) + assert bedpe_files == [ + Path("Chr1/Chr1.bedpe"), + Path("Chr1_Chr2/Chr1_Chr2_COMPARE.bedpe"), + Path("Chr2/Chr2.bedpe"), + ] + for name in ("2x2_GRID.png", "2x2_GRID.svg"): + grid = output / name + assert grid.is_file() + assert grid.stat().st_size > 0 + assert not any( + "Chr3" in str(path.relative_to(output)) for path in output.rglob("*") + ) + + +def test_static_sequence_selection_streams_only_requested_self_record(tmp_path): + fasta = tmp_path / "chromosomes.fa" + _write_multifasta(fasta) + output = tmp_path / "selected-self" + + result = _run_cli( + "-f", + fasta, + "-s", + "beta", + "--window", + 100, + "--modimizer", + 10, + "--identity", + 80, + "--no-plot", + "--output-dir", + output, + ) + + assert result.returncode == 0, result.stderr + result.stdout + assert "Computing self identity matrix for beta" in result.stdout + assert "Computing self identity matrix for alpha" not in result.stdout + assert "Computing self identity matrix for gamma" not in result.stdout + assert [path.relative_to(output) for path in output.rglob("*.bedpe")] == [ + Path("beta/beta.bedpe") + ] + + +def test_static_sequence_selection_rejects_unknown_identifier(tmp_path): + fasta = tmp_path / "chromosomes.fa" + fasta.write_text(">Chr1\n" + "ACGT" * 300 + "\n>Chr2\n" + "ACGT" * 275 + "\n") + output = tmp_path / "unknown-selection" + + result = _run_cli( + "-f", + fasta, + "-s", + "chr1", + "missing", + "--grid", + "--output-dir", + output, + ) + + combined_output = result.stderr + result.stdout + assert result.returncode == 2, combined_output + assert "missing" in combined_output + assert "does not match any FASTA identifier" in combined_output + assert not output.exists() + + def test_static_grid_regions_with_dotted_headers_crop_every_output(tmp_path): names = [ "PAN010.chr14.haplotype1.paternal", diff --git a/tests/test_cli_runtime.py b/tests/test_cli_runtime.py index 7eeade2..0ee3de1 100644 --- a/tests/test_cli_runtime.py +++ b/tests/test_cli_runtime.py @@ -1,5 +1,6 @@ import sys import shlex +from types import SimpleNamespace import numpy as np import pytest @@ -7,6 +8,161 @@ import moddotplot.moddotplot as cli +def _matrix_args(**overrides): + values = { + "window": None, + "resolution": 1000, + "kmer": 21, + "modimizer": 1000, + } + values.update(overrides) + return SimpleNamespace(**values) + + +def test_matrix_config_caps_resolution_at_one_valid_kmer_per_window(): + args = _matrix_args() + + config = cli._matrix_config_for_length(16_549, args) + + assert config.window_size == 21 + assert config.resolution == 789 + assert config.modimizer == 21 + assert config.sparsity == 1 + assert config.expectation == 21 + assert args.modimizer == 1000 + + +def test_matrix_config_preserves_default_nuclear_chromosome_parameters(): + config = cli._matrix_config_for_length(248_387_308, _matrix_args()) + + assert config.window_size == 248_388 + assert config.resolution == 1000 + assert config.modimizer == 1000 + assert config.sparsity == 128 + assert config.expectation == 1941 + + +def test_streaming_self_runner_finishes_one_record_before_requesting_next( + monkeypatch, +): + events = [] + + def records(*_args, **_kwargs): + events.append("yield:first") + yield "first", "ACGT", "first" + events.append("yield:second") + yield "second", "TGCA", "second" + + def process(**kwargs): + events.append(f"process:{kwargs['sequence_id']}") + + monkeypatch.setattr(cli, "iter_fasta_records", records) + monkeypatch.setattr(cli, "_process_static_self_record", process) + args = SimpleNamespace(output_dir=None) + + cli._run_streaming_static_self( + args, + ["input.fa"], + {"input.fa": ["first", "second"]}, + {}, + object(), + ) + + assert events == [ + "yield:first", + "process:first", + "yield:second", + "process:second", + ] + + +def test_streaming_self_runner_submits_only_record_descriptors_to_bounded_pool( + monkeypatch, +): + submitted = [] + + class FinishedFuture: + def result(self): + return None + + def cancel(self): + return False + + class FakeExecutor: + def __init__(self, *, max_workers, mp_context): + assert max_workers == 2 + assert mp_context.get_start_method() == "spawn" + + def __enter__(self): + return self + + def __exit__(self, *_args): + return False + + def submit(self, function, task): + assert function is cli._process_static_self_task + submitted.append(task) + return FinishedFuture() + + monkeypatch.setattr(cli, "ProcessPoolExecutor", FakeExecutor) + monkeypatch.setattr(cli, "as_completed", lambda futures: list(futures)) + monkeypatch.setattr(cli, "supports_indexed_fasta_access", lambda _path: True) + monkeypatch.setattr( + cli, + "iter_fasta_records", + lambda *_args, **_kwargs: pytest.fail( + "the parent must not decode sequences for indexed worker tasks" + ), + ) + args = SimpleNamespace(output_dir=None, processes=2) + + cli._run_streaming_static_self( + args, + ["indexed.fa"], + {"indexed.fa": ["chr1", "chr2", "chr3"]}, + {"chr2": ("chr2", 10, 20)}, + SimpleNamespace(command="moddotplot -f indexed.fa"), + ) + + assert [task[2] for task in submitted] == ["chr1", "chr2", "chr3"] + assert submitted[1][3] == ("chr2", 10, 20) + assert all(len(task) == 5 for task in submitted) + + +@pytest.mark.parametrize( + ("requested", "records", "indexed", "cpus", "expected"), + [ + (None, 25, True, 10, 2), + (None, 3, True, 2, 2), + (4, 2, True, 10, 2), + (4, 25, False, 10, 1), + (2, 1, True, 10, 1), + ], +) +def test_streaming_process_count_is_bounded_and_index_aware( + monkeypatch, requested, records, indexed, cpus, expected +): + monkeypatch.setattr(cli.os, "cpu_count", lambda: cpus) + args = SimpleNamespace(processes=requested) + + assert cli._streaming_process_count(args, records, indexed) == expected + + +def test_compute_only_auto_process_count_can_use_four_workers(monkeypatch): + monkeypatch.setattr(cli.os, "cpu_count", lambda: 10) + args = SimpleNamespace(processes=None, no_plot=True) + + assert cli._streaming_process_count(args, 25, indexed_access=True) == 4 + + +@pytest.mark.parametrize("requested", [0, 5, "many"]) +def test_streaming_process_count_rejects_invalid_values(requested): + with pytest.raises(ValueError, match="1 through 4"): + cli._streaming_process_count( + SimpleNamespace(processes=requested), 2, indexed_access=True + ) + + def _patch_fasta_input(monkeypatch, names, kmers): monkeypatch.setattr(cli, "isValidFasta", lambda _path: True) monkeypatch.setattr(cli, "getInputHeaders", lambda _path: names) @@ -45,6 +201,24 @@ def test_main_without_arguments_defaults_to_static_parser(monkeypatch, capsys): assert "the following arguments are required: command" not in captured.err +def test_static_main_errors_when_all_fasta_inputs_are_unreadable( + monkeypatch, tmp_path, capsys +): + missing = tmp_path / "missing.fa" + monkeypatch.setattr( + sys, + "argv", + ["moddotplot", "--fasta", str(missing), "--no-plot"], + ) + + with pytest.raises(SystemExit) as exc_info: + cli.main() + + captured = capsys.readouterr() + assert exc_info.value.code == 2 + assert "no readable FASTA input files remain" in captured.err + + def test_parser_defaults_omitted_subcommand_to_static(): args = cli.parse_args(["--fasta", "sequence.fa", "--no-plot"]) @@ -59,6 +233,87 @@ def test_parser_preserves_explicit_interactive_subcommand(): assert args.command == "interactive" +@pytest.mark.parametrize("option", ["-s", "--sequence"]) +def test_static_parser_accepts_sequence_selection_aliases(option): + args = cli.parse_args(["--fasta", "sequence.fa", option, "chr1", "chr2", "--grid"]) + + assert args.command == "static" + assert args.sequence == ["chr1", "chr2"] + + +def test_static_parser_accepts_bounded_process_request(): + args = cli.parse_args(["--fasta", "sequence.fa", "--processes", "3"]) + + assert args.processes == 3 + + +@pytest.mark.parametrize("value", [0, 5, True, "many"]) +def test_process_count_validation_rejects_invalid_config_values(value): + with pytest.raises(ValueError, match="integer from 1 through 4"): + cli._validated_process_count(value) + + +def test_static_parser_rejects_process_requests_outside_public_bound(): + with pytest.raises(SystemExit) as exc_info: + cli.parse_args(["--fasta", "sequence.fa", "--processes", "5"]) + + assert exc_info.value.code == 2 + + +def test_columnar_bedpe_writer_matches_legacy_text(tmp_path): + matrix = np.array([[1.0, 0.91, 0.2], [0.91, 1.0, 0.88], [0.2, 0.88, 1.0]]) + kwargs = dict( + window_size=10, + id_threshold=86, + x_name="chr1", + y_name="chr1", + self_identity=True, + x_offset=1, + y_offset=1, + x_end=29, + y_end=29, + ) + expected_rows = cli.convertMatrixToBed(matrix, **kwargs) + output = tmp_path / "matrix.bedpe" + + cli._write_matrix_bedpe( + output, + cli.iterMatrixToBedChunks(matrix, max_chunk_cells=2, **kwargs), + ) + + expected = "".join("\t".join(map(str, row)) + "\n" for row in expected_rows) + assert output.read_text() == expected + + +def test_interactive_short_s_remains_save_flag(): + args = cli.get_parser().parse_args(["interactive", "--fasta", "sequence.fa", "-s"]) + + assert args.save is True + + +def test_sequence_selection_prefers_exact_names_and_accepts_casefold_fallback(): + selected, names = cli._select_fasta_headers( + {"genome.fa": ["Chr1", "chr1", "Chr2"]}, ["chr1", "CHR2"] + ) + + assert selected == {"genome.fa": ["chr1", "Chr2"]} + assert names == ["chr1", "Chr2"] + + +@pytest.mark.parametrize( + ("headers", "selectors", "message"), + [ + ({"genome.fa": ["Chr1"]}, ["chr1", "Chr1"], "already requested"), + ({"genome.fa": ["Chr1", "chr1"]}, ["CHR1"], "ambiguous"), + ], +) +def test_sequence_selection_rejects_duplicate_or_ambiguous_requests( + headers, selectors, message +): + with pytest.raises(ValueError, match=message): + cli._select_fasta_headers(headers, selectors) + + @pytest.mark.parametrize("option", ["--colors", "--color"]) def test_static_parser_accepts_color_option_aliases(option): args = cli.get_parser().parse_args( @@ -91,6 +346,17 @@ def test_static_config_supports_legacy_color_key(): assert args.colors == ["#legacy"] +def test_static_config_accepts_sequence_selection(): + args = cli.get_parser().parse_args(["static", "--fasta", "sequence.fa"]) + + cli._apply_static_config( + args, + {"fasta": ["sequence.fa"], "sequence": ["chr1", "chr2"]}, + ) + + assert args.sequence == ["chr1", "chr2"] + + def test_static_delta_defaults_to_half_window(): args = cli.get_parser().parse_args(["static", "--fasta", "sequence.fa"]) diff --git a/tests/test_fasta_parser.py b/tests/test_fasta_parser.py index 06bb5e9..9267416 100644 --- a/tests/test_fasta_parser.py +++ b/tests/test_fasta_parser.py @@ -1,4 +1,5 @@ import gzip +import os import struct import zlib @@ -12,8 +13,10 @@ generateKmersFromFasta, getInputHeaders, getInputSeqLength, + iter_fasta_records, isValidFasta, readKmersFromFile, + supports_indexed_fasta_access, ) @@ -51,6 +54,60 @@ def _bgzf_block(data): return header + extra + compressed + footer +def _indexed_fasta(records, line_bases=8): + contents = bytearray() + index_lines = [] + for name, sequence in records: + contents.extend(f">{name} description\n".encode("ascii")) + sequence_offset = len(contents) + for start in range(0, len(sequence), line_bases): + contents.extend(sequence[start : start + line_bases].encode("ascii")) + contents.extend(b"\n") + index_lines.append( + f"{name}\t{len(sequence)}\t{sequence_offset}\t" + f"{line_bases}\t{line_bases + 1}\n" + ) + return bytes(contents), "".join(index_lines) + + +def _write_bgzf(path, uncompressed, block_size=29): + blocks = [] + gzip_entries = [] + compressed_offset = 0 + uncompressed_offset = 0 + for start in range(0, len(uncompressed), block_size): + if blocks: + gzip_entries.append((compressed_offset, uncompressed_offset)) + data = uncompressed[start : start + block_size] + block = _bgzf_block(data) + blocks.append(block) + compressed_offset += len(block) + uncompressed_offset += len(data) + blocks.append(_bgzf_block(b"")) + path.write_bytes(b"".join(blocks)) + gzi = struct.pack("alpha\nACGT\n>beta\nTTAA\n") + + assert list(iter_fasta_records(fasta, record_ids=["beta", "alpha"])) == [ + ("beta", "TTAA", "beta"), + ("alpha", "ACGT", "alpha"), + ] + with pytest.raises(ValueError, match="FASTA record.*not found: 'missing'"): + list(iter_fasta_records(fasta, record_ids=["missing"])) + + def test_read_kmers_preserves_record_order_and_public_return_shape(tmp_path): fasta = tmp_path / "records.fa" fasta.write_text(">alpha description\nACGT\n>beta\nTTAA\n") @@ -120,6 +293,59 @@ def test_read_kmers_preserves_record_order_and_public_return_shape(tmp_path): ) +def test_read_kmers_hashes_only_selected_records_from_unindexed_fasta(tmp_path): + sequences = { + "Chr1": "ACGTAC", + "Chr2": "TTAACC", + "Chr3": "GGGGGG", + } + fasta = tmp_path / "selected.fa" + fasta.write_text( + "".join(f">{name}\n{sequence}\n" for name, sequence in sequences.items()) + ) + + result = readKmersFromFile( + str(fasta), + ksize=3, + quiet=True, + fw_only=True, + ambiguous=False, + record_ids=["Chr1", "Chr2"], + ) + + assert len(result) == 2 + for hashes, sequence in zip(result, (sequences["Chr1"], sequences["Chr2"])): + assert hashes.tolist() == list( + generateKmersFromFasta(sequence, 3, quiet=True, fw_only=True) + ) + + +def test_read_kmers_uses_index_for_selected_record_subset(tmp_path, monkeypatch): + fasta = tmp_path / "selected-indexed.fa" + fasta.write_bytes(b">alpha\nACGTAC\n>beta\nTTAACC\n>gamma\nGGGGGG\n") + (tmp_path / "selected-indexed.fa.fai").write_text( + "alpha\t6\t7\t6\t7\n" "beta\t6\t20\t6\t7\n" "gamma\t6\t34\t6\t7\n" + ) + + def fail_streaming(*_args, **_kwargs): + raise AssertionError("indexed subset should not use the streaming fallback") + + monkeypatch.setattr(fasta_parser, "_iter_selected_fasta_records", fail_streaming) + result = readKmersFromFile( + str(fasta), + ksize=3, + quiet=True, + fw_only=True, + ambiguous=False, + record_ids=["beta", "alpha"], + ) + + assert [hashes.tolist() for hashes in result] == [ + list(generateKmersFromFasta(sequence, 3, quiet=True, fw_only=True)) + for sequence in ("TTAACC", "ACGTAC") + ] + + def test_read_kmers_hashes_only_the_requested_region(tmp_path): sequence = "ACGT" * 250 fasta = tmp_path / "region.fa" diff --git a/tests/test_grid.py b/tests/test_grid.py index b692cf0..a744e6e 100644 --- a/tests/test_grid.py +++ b/tests/test_grid.py @@ -5,6 +5,7 @@ import pytest from moddotplot.const import DIRECTION_COLORS +from moddotplot.native_render import FALLBACK_FONT_FAMILY, set_figure_font_family from moddotplot.static_plots import _build_grid_figure, create_grid @@ -128,6 +129,18 @@ def _collection_center(axis): ) +def _collection_centers(axis): + return sorted( + ( + (path.vertices[:, 0].min() + path.vertices[:, 0].max()) / 2, + (path.vertices[:, 1].min() + path.vertices[:, 1].max()) / 2, + ) + for collection in axis.collections + for path in collection.get_paths() + if path.vertices.size + ) + + def _all_artist_colors(axes): colors = set() for axis in axes.flat: @@ -269,6 +282,56 @@ def test_self_comparisons_run_bottom_left_to_top_right(): plt.close(figure) +def test_self_comparison_grid_diagonal_uses_full_symmetric_dotplots(): + kwargs = _basic_two_sequence_grid( + singles=[ + _records( + "sequence_a", + "sequence_a", + [(10, 30, 91), (50, 50, 100)], + ), + _records( + "sequence_b", + "sequence_b", + [(20, 40, 92), (60, 60, 100)], + ), + ] + ) + + figure, axes = _build_grid_figure(**kwargs) + try: + assert _collection_centers(axes[1, 0]) == [ + (10, 30), + (30, 10), + (50, 50), + ] + assert _collection_centers(axes[0, 1]) == [ + (20, 40), + (40, 20), + (60, 60), + ] + assert len(_collection_centers(axes[0, 0])) == 1 + assert len(_collection_centers(axes[1, 1])) == 1 + finally: + plt.close(figure) + + +def test_full_self_comparison_input_is_not_mirrored_twice(): + kwargs = _basic_two_sequence_grid( + singles=[ + _records("sequence_a", "sequence_a", [(10, 30, 91), (30, 10, 91)]), + _records("sequence_b", "sequence_b", [(20, 40, 92), (40, 20, 92)]), + ] + ) + + figure, axes = _build_grid_figure(**kwargs) + try: + assert _collection_centers(axes[1, 0]) == [(10, 30), (30, 10)] + assert _collection_centers(axes[0, 1]) == [(20, 40), (40, 20)] + finally: + plt.close(figure) + + def test_grid_region_names_fit_panels_and_numeric_labels_are_doubled(): names = [ "PAN010.chr14.haplotype1.paternal:1-4000000", @@ -368,18 +431,91 @@ def test_grid_exact_bounds_do_not_expand_to_next_nice_tick(): ("axis_end", "unit"), [(100_000, "Kbp"), (103_156_783, "Mbp"), (500_000_000, "Gbp")], ) -def test_grid_labels_both_axes_with_genomic_units(axis_end, unit): +def test_grid_labels_genomic_units_only_on_bottom_left_cell(axis_end, unit): kwargs = _basic_two_sequence_grid( xlim=(1, axis_end), axes_label=None, breaks=None, ) - figure, _axes = _build_grid_figure(**kwargs) + figure, axes = _build_grid_figure(**kwargs) try: expected = f"Genomic Position ({unit})" - assert figure._supxlabel.get_text() == expected - assert figure._supylabel.get_text() == expected + bottom_left_axis = axes[-1, 0] + assert bottom_left_axis.get_xlabel() == expected + assert sum(axis.get_xlabel() == expected for axis in axes.flat) == 1 + assert getattr(figure, "_supxlabel", None) is None + assert getattr(figure, "_supylabel", None) is None + vertical_titles = [ + text + for axis in axes.flat + for text in axis.texts + if text.get_gid() == "grid-y-axis-title" + ] + assert len(vertical_titles) == 1 + assert vertical_titles[0].get_text() == expected + assert list(axis.get_ylabel() for axis in axes[:, 0]) == [ + "sequence_b", + "sequence_a", + ] + + figure.canvas.draw() + renderer = figure.canvas.get_renderer() + figure_bounds = figure.bbox + cell_bounds = bottom_left_axis.get_window_extent(renderer) + horizontal_bounds = bottom_left_axis.xaxis.label.get_window_extent(renderer) + vertical_bounds = vertical_titles[0].get_window_extent(renderer) + row_label_bounds = bottom_left_axis.yaxis.label.get_window_extent(renderer) + + assert cell_bounds.x0 <= horizontal_bounds.x0 + horizontal_bounds.width / 2 + assert horizontal_bounds.x0 + horizontal_bounds.width / 2 <= cell_bounds.x1 + assert cell_bounds.y0 <= vertical_bounds.y0 + vertical_bounds.height / 2 + assert vertical_bounds.y0 + vertical_bounds.height / 2 <= cell_bounds.y1 + assert vertical_bounds.x1 <= row_label_bounds.x0 + for bounds in (horizontal_bounds, vertical_bounds): + assert figure_bounds.contains(bounds.x0, bounds.y0) + assert figure_bounds.contains(bounds.x1, bounds.y1) + finally: + plt.close(figure) + + +@pytest.mark.parametrize(("width", "expected_width"), [(1, 2), (4, 4)]) +def test_grid_axis_labels_do_not_expand_saved_canvas(tmp_path, width, expected_width): + kwargs = _basic_two_sequence_grid(width=width) + + with plt.rc_context({"savefig.bbox": "tight"}): + _create_grid(tmp_path, **kwargs) + + image = plt.imread(tmp_path / "2x2_GRID.png") + assert image.shape[:2] == (expected_width * 72, expected_width * 72) + + +def test_one_cell_grid_axis_titles_stay_inside_canvas(): + name = "sequence_a" + kwargs = _grid_kwargs( + singles=[_records(name, name, [(10, 10, 100)])], + doubles=[], + single_names=[name], + double_names=[], + ) + kwargs["width"] = 4 + + figure, axes = _build_grid_figure(**kwargs) + try: + vertical_title = next( + text for text in axes[0, 0].texts if text.get_gid() == "grid-y-axis-title" + ) + for family in (None, FALLBACK_FONT_FAMILY): + if family is not None: + set_figure_font_family(figure, family) + for dpi in (72, 100, 300, 600): + figure.set_dpi(dpi) + figure.canvas.draw() + renderer = figure.canvas.get_renderer() + for title in (axes[0, 0].xaxis.label, vertical_title): + bounds = title.get_window_extent(renderer) + assert figure.bbox.contains(bounds.x0, bounds.y0) + assert figure.bbox.contains(bounds.x1, bounds.y1) finally: plt.close(figure) diff --git a/tests/test_native_full_plot.py b/tests/test_native_full_plot.py new file mode 100644 index 0000000..d9d0a6c --- /dev/null +++ b/tests/test_native_full_plot.py @@ -0,0 +1,224 @@ +import matplotlib.pyplot as plt +from matplotlib.colors import to_rgba +import numpy as np +import pandas as pd +from PIL import Image +import pytest + +import moddotplot.static_plots as static_plots + + +def _processed_tiles(rows): + frame = pd.DataFrame.from_records( + rows, + columns=[ + "q", + "q_st", + "q_en", + "r", + "r_st", + "r_en", + "perID_by_events", + "discrete", + ], + ) + frame["discrete"] = pd.Categorical( + frame["discrete"], categories=[0, 1], ordered=True + ) + return frame + + +@pytest.mark.parametrize("deraster", [False, True]) +def test_native_full_self_plot_mirrors_only_missing_triangle(deraster): + frame = _processed_tiles( + [ + ("chr1", 10, 14, "chr1", 10, 14, 100.0, 1), + ("chr1", 10, 14, "chr1", 20, 24, 95.0, 0), + ] + ) + + figure = static_plots._build_full_figure( + sdf=frame, + name_x="chr1:1-30", + name_y="chr1:1-30", + palette="Spectral_11", + palette_orientation="+", + custom_colors=["#010203", "#abcdef"], + axes_labels=[1, 15, 30], + xlim=(1, 30), + deraster=deraster, + width=3, + is_pairwise=False, + ) + try: + axis = figure.axes[0] + assert [len(collection.get_paths()) for collection in axis.collections] == [ + 2, + 1, + ] + assert [collection.get_rasterized() for collection in axis.collections] == [ + not deraster, + not deraster, + ] + np.testing.assert_allclose( + axis.collections[1].get_paths()[0].vertices[:4], + [[18, 8], [22, 8], [22, 12], [18, 12]], + ) + np.testing.assert_allclose( + axis.collections[1].get_facecolors(), [to_rgba("#010203")] + ) + assert axis.get_xlim() == pytest.approx((1, 30)) + assert axis.get_ylim() == pytest.approx((1, 30)) + assert axis.get_xlabel() == "Genomic Position (Kbp)" + assert axis.get_title() == "chr1" + assert axis.get_ylabel() == "chr1" + assert figure._suptitle.get_text() == "Self-Identity Plot: chr1" + finally: + plt.close(figure) + + +def test_native_full_self_plot_does_not_duplicate_loaded_symmetric_rows(): + frame = _processed_tiles( + [ + ("chr1", 10, 14, "chr1", 10, 14, 100.0, 1), + ("chr1", 10, 14, "chr1", 20, 24, 95.0, 0), + ("chr1", 20, 24, "chr1", 10, 14, 95.0, 0), + ] + ) + + figure = static_plots._build_full_figure( + sdf=frame, + name_x="chr1", + name_y="chr1", + palette="Spectral_11", + palette_orientation="+", + custom_colors=None, + axes_labels=None, + xlim=(1, 30), + deraster=False, + width=3, + is_pairwise=False, + ) + try: + assert len(figure.axes[0].collections) == 1 + assert len(figure.axes[0].collections[0].get_paths()) == 3 + finally: + plt.close(figure) + + +def test_native_full_plot_supports_empty_sparse_data_with_explicit_bounds(): + frame = _processed_tiles([]) + + figure = static_plots._build_full_figure( + sdf=frame, + name_x="chrM", + name_y="chrM", + palette="Spectral_11", + palette_orientation="+", + custom_colors=None, + axes_labels=None, + xlim=(1, 16_569), + deraster=False, + width=2, + is_pairwise=False, + ) + try: + axis = figure.axes[0] + assert axis.get_xlim() == pytest.approx((1, 16_569)) + assert axis.get_ylim() == pytest.approx((1, 16_569)) + assert len(axis.collections) == 1 + assert len(axis.collections[0].get_paths()) == 0 + finally: + plt.close(figure) + + +@pytest.mark.parametrize("savefig_bbox", [None, "tight"]) +def test_create_plots_uses_native_comparative_renderer_and_exact_canvas( + tmp_path, monkeypatch, savefig_bbox +): + bed = [ + ( + "#query_name", + "query_start", + "query_end", + "reference_name", + "reference_start", + "reference_end", + "perID_by_events", + ), + ("alpha", 1, 100, "beta", 101, 200, 95.0), + ] + + def fail_plotnine_full(*_args, **_kwargs): + raise AssertionError("the full/compare path must not call make_dot") + + monkeypatch.setattr(static_plots, "make_dot", fail_plotnine_full) + monkeypatch.setattr(static_plots, "make_hist", lambda *_args, **_kwargs: None) + + with plt.rc_context({"savefig.bbox": savefig_bbox}): + outputs = static_plots.create_plots( + sdf=[bed], + directory=str(tmp_path), + name_x="alpha", + name_y="beta", + palette="Spectral_11", + palette_orientation="+", + no_hist=True, + width=2, + dpi=40, + is_freq=False, + xlim=(1, 200), + custom_colors=None, + custom_breakpoints=None, + from_file=None, + is_pairwise=True, + axes_labels=None, + axes_tick_number=7, + vector_format="svg", + deraster=False, + annotation=None, + ) + + png = tmp_path / "alpha_beta_COMPARE.png" + svg = tmp_path / "alpha_beta_COMPARE.svg" + assert set(outputs) == {str(png), str(svg)} + assert png.read_bytes().startswith(b"\x89PNG\r\n\x1a\n") + assert svg.read_bytes().lstrip().startswith(b" colors[0, 0] + assert colors[1, 0] > colors[1, 2] + assert np.mean(colors[1, :3]) < np.mean(colors[0, :3]) + assert figure.axes[0].get_title() == "query" + assert figure.axes[0].get_ylabel() == "reference" + finally: + plt.close(figure) diff --git a/tests/test_native_sequence_sketches.py b/tests/test_native_sequence_sketches.py new file mode 100644 index 0000000..905ce97 --- /dev/null +++ b/tests/test_native_sequence_sketches.py @@ -0,0 +1,153 @@ +import numpy as np +import pytest + +from moddotplot import _nthash +from moddotplot.estimate_identity import ( + prepare_modimizer_sketches, + prepare_sequence_sketches, +) +from moddotplot.parse_fasta import _hash_sequence + + +def _legacy_sketches( + sequence, window_size, sparsity, delta, k, ambiguous, expectation, canonical +): + hashes = _hash_sequence( + sequence, k, fw_only=not canonical, ambiguous=ambiguous + ) + return prepare_modimizer_sketches( + len(hashes), + hashes, + window_size, + sparsity, + delta, + k, + ambiguous, + expectation, + ) + + +def _assert_prepared_equal(actual, expected): + assert len(actual.core) == len(expected.core) + assert len(actual.neighbors) == len(expected.neighbors) + for actual_sketch, expected_sketch in zip(actual.core, expected.core): + np.testing.assert_array_equal(actual_sketch, expected_sketch) + for actual_sketch, expected_sketch in zip(actual.neighbors, expected.neighbors): + np.testing.assert_array_equal(actual_sketch, expected_sketch) + + +@pytest.mark.parametrize("canonical", [False, True]) +@pytest.mark.parametrize("ambiguous", [False, True]) +@pytest.mark.parametrize("delta", [0, 0.35, 0.5]) +@pytest.mark.parametrize("sparsity", [1, 8, 64]) +def test_native_sequence_sketches_match_legacy_pipeline( + canonical, ambiguous, delta, sparsity +): + rng = np.random.default_rng(20260929) + sequence = "".join(rng.choice(list("ACGT"), size=713)) + sequence = sequence[:91] + "nRy" + sequence[94:351].lower() + "U" + sequence[352:] + parameters = dict( + window_size=73, + sparsity=sparsity, + delta=delta, + k=11, + ambiguous=ambiguous, + expectation=91, + canonical=canonical, + ) + + actual = prepare_sequence_sketches(sequence, **parameters) + expected = _legacy_sketches(sequence, **parameters) + + _assert_prepared_equal(actual, expected) + + +@pytest.mark.parametrize( + ("sequence", "window_size", "k", "expectation"), + [ + ("A" * 401, 83, 11, 200), + ("ACGT" * 10, 17, 21, 100), + ("ACGT", 10, 21, 100), + ("N" * 101, 31, 7, 100), + ], +) +def test_native_sequence_sketches_match_adaptive_and_empty_edge_cases( + sequence, window_size, k, expectation +): + parameters = dict( + window_size=window_size, + sparsity=64, + delta=0.5, + k=k, + ambiguous=False, + expectation=expectation, + canonical=True, + ) + + _assert_prepared_equal( + prepare_sequence_sketches(sequence, **parameters), + _legacy_sketches(sequence, **parameters), + ) + + +def test_sequence_sketch_path_does_not_materialize_positional_hashes(monkeypatch): + def fail_if_called(*_args, **_kwargs): + raise AssertionError("the chromosome-wide positional hash API was used") + + monkeypatch.setattr(_nthash, "hash_kmers", fail_if_called) + + prepared = prepare_sequence_sketches( + "ACGT" * 10_000, + window_size=1_000, + sparsity=64, + delta=0.5, + k=21, + ambiguous=False, + expectation=16, + ) + + assert prepared.core + assert prepared.neighbors + assert all(sketch.dtype == np.uint64 for sketch in prepared.core) + + +@pytest.mark.parametrize("sparsity", [0, -1, 3, 12]) +def test_native_sequence_sketches_require_power_of_two_sparsity(sparsity): + with pytest.raises(ValueError, match="positive power of two"): + prepare_sequence_sketches( + "ACGTACGT", + window_size=4, + sparsity=sparsity, + delta=0.5, + k=3, + ambiguous=False, + expectation=1, + ) + + +def test_native_sketch_api_rejects_out_of_range_interval_bounds(): + with pytest.raises(ValueError, match="interval bounds"): + _nthash.sketch_kmers("ACGT", 3, True, [(0, 3)], 1, 1, False) + + +def test_native_intersections_read_each_sequence_length_only_once(): + packed = np.asarray([7], dtype=np.uint64).tobytes() + + class ChangingLengthSequence: + def __init__(self): + self.length_calls = 0 + + def __len__(self): + self.length_calls += 1 + return self.length_calls + + def __getitem__(self, index): + if index == 0: + return packed + raise IndexError(index) + + left = ChangingLengthSequence() + result = _nthash.intersection_counts(left, [packed]) + + np.testing.assert_array_equal(np.frombuffer(result, dtype=np.int32), [1]) + assert left.length_calls == 1 diff --git a/tests/test_plot_fonts.py b/tests/test_plot_fonts.py index e2ca36e..53aa958 100644 --- a/tests/test_plot_fonts.py +++ b/tests/test_plot_fonts.py @@ -62,7 +62,11 @@ def test_plotnine_outputs_retry_with_dejavu_on_glyph_failure(monkeypatch): attempted_families = [] def fail_for_helvetica(plot, **_kwargs): - family = plot.theme.getp(("text", "family"))[0] + if hasattr(plot.theme, "getp"): + family = plot.theme.getp(("text", "family"))[0] + else: + # Plotnine <0.15 stores resolved themeable properties directly. + family = plot.theme.themeables["text"].properties["family"][0] attempted_families.append(family) if family == DEFAULT_FONT_FAMILY: raise RuntimeError("failed to load glyph") @@ -71,3 +75,42 @@ def fail_for_helvetica(plot, **_kwargs): static_plots._save_plot(ggplot(), filename="unused.png") assert attempted_families == [DEFAULT_FONT_FAMILY, FALLBACK_FONT_FAMILY] + + +def test_plotnine_pair_draws_once_for_png_and_vector(monkeypatch, tmp_path): + figure = plt.figure() + + class FakePlot: + draw_count = 0 + + def __add__(self, _other): + return self + + def draw(self, show=False): + assert show is False + self.draw_count += 1 + return figure + + saved = [] + + def fake_save_figure_pair(current, prefix, vector_format, dpi, **kwargs): + saved.append((current, prefix, vector_format, dpi, kwargs)) + return tmp_path / "plot.png", tmp_path / "plot.svg" + + monkeypatch.setattr(static_plots, "save_figure_pair", fake_save_figure_pair) + plot = FakePlot() + + static_plots._draw_and_save_plot_pair( + plot, + tmp_path / "plot", + width=9, + height=9, + dpi=300, + vector_format="svg", + ) + + assert plot.draw_count == 1 + assert saved[0][0] is figure + assert saved[0][2:4] == ("svg", 300) + assert saved[0][4] == {"bbox_inches": figure.bbox_inches} + assert not plt.fignum_exists(figure.number) diff --git a/tests/test_sparse_containment.py b/tests/test_sparse_containment.py index afbe808..51b97ea 100644 --- a/tests/test_sparse_containment.py +++ b/tests/test_sparse_containment.py @@ -1,6 +1,7 @@ import numpy as np import pytest +from moddotplot import _nthash import moddotplot.estimate_identity as estimate_identity from moddotplot.estimate_identity import ( _sketch_intersection_counts, @@ -112,6 +113,113 @@ def test_sparse_counts_use_wide_accumulator(): assert counts[0, 0] == 300 +def test_native_sorted_uint64_intersections_match_scalar_reference(monkeypatch): + rng = np.random.default_rng(616) + sketches_a = [ + np.sort( + rng.choice(5_000, size=int(rng.integers(0, 250)), replace=False) + ).astype(np.uint64) + for _ in range(17) + ] + sketches_b = [ + np.sort( + rng.choice(5_000, size=int(rng.integers(0, 250)), replace=False) + ).astype(np.uint64) + for _ in range(13) + ] + expected = np.asarray( + [ + [len(set(left.tolist()) & set(right.tolist())) for right in sketches_b] + for left in sketches_a + ], + dtype=np.int32, + ) + + def fail_if_called(*_args, **_kwargs): + raise AssertionError("the CSR compatibility path was used") + + monkeypatch.setattr(estimate_identity, "csr_matrix", fail_if_called) + + actual = _sketch_intersection_counts(sketches_a, sketches_b) + + np.testing.assert_array_equal(actual, expected) + assert actual.dtype == np.int32 + + +def test_native_intersection_merge_handles_hash_shared_by_every_window(): + shared = np.uint64(2**63 + 17) + sketches_a = [np.array([index, shared], dtype=np.uint64) for index in range(20)] + sketches_b = [ + np.array([index + 100, shared], dtype=np.uint64) for index in range(30) + ] + # Keep the production precondition explicit: arrays are sorted and unique. + sketches_a = [np.sort(sketch) for sketch in sketches_a] + sketches_b = [np.sort(sketch) for sketch in sketches_b] + + np.testing.assert_array_equal( + _sketch_intersection_counts(sketches_a, sketches_b), + np.ones((20, 30), dtype=np.int32), + ) + + +def test_native_intersection_retains_ephemeral_sequence_items(): + released = [] + + class TrackedBytes(bytes): + def __new__(cls, payload): + instance = super().__new__(cls, payload) + instance.release_events = released + return instance + + def __del__(self): + self.release_events.append(True) + + class EphemeralSketches: + def __init__(self, sketches): + self.sketches = sketches + + def __len__(self): + return len(self.sketches) + + def __getitem__(self, index): + # A sequence implementation is allowed to return a newly-created + # object for each item. No item may be released while the native + # function is still materializing or reading the sequence. + assert not released + return TrackedBytes(self.sketches[index]) + + left = EphemeralSketches( + [ + np.array([1, 3], dtype=np.uint64).tobytes(), + np.array([2, 3], dtype=np.uint64).tobytes(), + ] + ) + right = EphemeralSketches( + [ + np.array([3, 4], dtype=np.uint64).tobytes(), + np.array([1, 2], dtype=np.uint64).tobytes(), + ] + ) + + packed = _nthash.intersection_counts(left, right) + + np.testing.assert_array_equal( + np.frombuffer(packed, dtype=np.int32).reshape(2, 2), + np.array([[1, 1], [1, 1]], dtype=np.int32), + ) + assert len(released) == 4 + + +def test_unsorted_uint64_arrays_retain_compatibility_path(): + sketches_a = [np.array([9, 1, 5], dtype=np.uint64)] + sketches_b = [np.array([5, 2, 9], dtype=np.uint64)] + + np.testing.assert_array_equal( + _sketch_intersection_counts(sketches_a, sketches_b), + np.array([[2]], dtype=np.int32), + ) + + def test_matrix_path_does_not_fall_back_to_per_cell_set_intersections(monkeypatch): def fail_if_called(*_args, **_kwargs): raise AssertionError("per-cell Python containment was used") diff --git a/tests/test_static_customization.py b/tests/test_static_customization.py index 9916790..fbcdb23 100644 --- a/tests/test_static_customization.py +++ b/tests/test_static_customization.py @@ -256,3 +256,15 @@ def test_get_colors_rejects_invalid_custom_breakpoints(breakpoints, message): is_freq=False, custom_breakpoints=breakpoints, ) + + +def test_get_colors_rejects_breakpoints_that_do_not_cover_observed_values(): + scores = pd.DataFrame({"perID_by_events": [86.0, 100.0]}) + + with pytest.raises(ValueError, match="cover all finite identity values"): + get_colors( + scores, + ncolors=2, + is_freq=False, + custom_breakpoints=[86, 90, 95], + ) From 69728ba941b56f42937aa92c9f46e16096408a9d Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Wed, 30 Sep 2026 10:24:25 -0400 Subject: [PATCH 04/16] Format Python code with Black --- benchmarks/benchmark_hashing.py | 9 +++---- setup.py | 1 - src/moddotplot/annotations.py | 1 - src/moddotplot/moddotplot.py | 29 ++++++---------------- src/moddotplot/native_render.py | 1 - src/moddotplot/parse_fasta.py | 8 ++---- src/moddotplot/static_plots.py | 20 ++++++--------- tests/test_cli_integration.py | 1 - tests/test_entrypoints.py | 1 - tests/test_fasta_parser.py | 4 +-- tests/test_grid.py | 1 - tests/test_hash_benchmark.py | 1 - tests/test_issue53_memory_safe_plotting.py | 1 - tests/test_native_full_plot.py | 2 ++ tests/test_native_sequence_sketches.py | 4 +-- tests/test_nthash.py | 1 - tests/test_packaging_metadata.py | 1 - 17 files changed, 26 insertions(+), 60 deletions(-) diff --git a/benchmarks/benchmark_hashing.py b/benchmarks/benchmark_hashing.py index 0c30371..27a2667 100644 --- a/benchmarks/benchmark_hashing.py +++ b/benchmarks/benchmark_hashing.py @@ -25,7 +25,6 @@ from pathlib import Path from typing import Callable, Dict, Iterable, List, Optional, Sequence, Tuple - try: from moddotplot.parse_fasta import _hash_sequence except ModuleNotFoundError: # Permit running from an uninstalled source tree. @@ -143,10 +142,10 @@ def benchmark(sequence: str, k: int, repeats: int, seed: int) -> Dict[str, objec ) } if mmh3_module is not None: - implementations[ - "legacy_mmh3" - ] = lambda canonical=canonical: _legacy_mmh3_hashes( - sequence, k, canonical, mmh3_module + implementations["legacy_mmh3"] = ( + lambda canonical=canonical: _legacy_mmh3_hashes( + sequence, k, canonical, mmh3_module + ) ) raw: Dict[str, List[float]] = {name: [] for name in implementations} diff --git a/setup.py b/setup.py index 4123808..79c595a 100644 --- a/setup.py +++ b/setup.py @@ -2,7 +2,6 @@ from setuptools import Extension, setup - if sys.platform == "win32": compile_args = ["/O2", "/std:c++17"] else: diff --git a/src/moddotplot/annotations.py b/src/moddotplot/annotations.py index 33af709..5f3eb88 100644 --- a/src/moddotplot/annotations.py +++ b/src/moddotplot/annotations.py @@ -3,7 +3,6 @@ import numpy as np import pandas as pd - DEFAULT_ANNOTATION_COLOR = "#4C72B0" BED_COLUMNS = [ "chrom", diff --git a/src/moddotplot/moddotplot.py b/src/moddotplot/moddotplot.py index d7bd529..3158f26 100755 --- a/src/moddotplot/moddotplot.py +++ b/src/moddotplot/moddotplot.py @@ -46,7 +46,6 @@ from moddotplot.plot_summary import PlotSummaryWriter - # Static plotting pulls in the Plotnine and Matplotlib stacks. Keep those # imports behind the static command boundary so ``--help`` and interactive # mode do not pay their startup cost. @@ -797,9 +796,7 @@ def _matrix_config_for_length(kmer_count, args): requested_resolution = int(args.resolution) if requested_resolution <= 0: raise ValueError("resolution must be greater than zero") - window_size = max( - int(args.kmer), math.ceil(kmer_count / requested_resolution) - ) + window_size = max(int(args.kmer), math.ceil(kmer_count / requested_resolution)) resolution = math.ceil(kmer_count / window_size) if window_size < 10: @@ -830,9 +827,7 @@ def _write_bedpe(path, rows): row_iterator = iter(rows) with open(path, "w") as bedfile: while batch := list(islice(row_iterator, 8192)): - bedfile.writelines( - "\t".join(map(str, row)) + "\n" for row in batch - ) + bedfile.writelines("\t".join(map(str, row)) + "\n" for row in batch) def _write_matrix_bedpe(path, chunks): @@ -869,9 +864,7 @@ def _annotate_bed_direction_frame( try: values = np.asarray(forward_matrix)[query_indices, reference_indices] except IndexError as error: - raise ValueError( - "BEDPE coordinates fall outside direction matrices" - ) from error + raise ValueError("BEDPE coordinates fall outside direction matrices") from error annotated["direction"] = np.where(values > 0, "Forward", "Reverse") return annotated @@ -1034,9 +1027,7 @@ def _process_static_self_record( chromsizes=kmer_count, output_cool=cooler_output, ) - print( - f"Saved self-identity matrix as a cooler file to {cooler_output}\n" - ) + print(f"Saved self-identity matrix as a cooler file to {cooler_output}\n") except Exception as error: print(f"Error creating cooler file: {error}") @@ -1201,9 +1192,8 @@ def _run_streaming_static_self( for sequence_id in fasta_headers[fasta_path] ] unique_record_ids = {task[2] for task in tasks} - indexed_access = ( - len(unique_record_ids) == len(tasks) - and all(supports_indexed_fasta_access(path) for path in fasta_list) + indexed_access = len(unique_record_ids) == len(tasks) and all( + supports_indexed_fasta_access(path) for path in fasta_list ) process_count = _streaming_process_count(args, len(tasks), indexed_access) @@ -1575,8 +1565,7 @@ def main(): and not args.grid_only and args.compare_order == "sequential" and all( - os.path.isfile(path) and os.path.getsize(path) > 0 - for path in fasta_list + os.path.isfile(path) and os.path.getsize(path) > 0 for path in fasta_list ) ) if streaming_static_self: @@ -2367,9 +2356,7 @@ def main(): ) try: - matrix_config = _matrix_config_for_length( - smaller_length, args - ) + matrix_config = _matrix_config_for_length(smaller_length, args) except ValueError as error: print(f"Error: {error}.\n") sys.exit(2) diff --git a/src/moddotplot/native_render.py b/src/moddotplot/native_render.py index a27d73d..d8a3f07 100644 --- a/src/moddotplot/native_render.py +++ b/src/moddotplot/native_render.py @@ -21,7 +21,6 @@ import numpy as np import pandas as pd - ColorSource = Union[Sequence[str], Mapping[object, str]] TickFormatter = Callable[[float, int], str] diff --git a/src/moddotplot/parse_fasta.py b/src/moddotplot/parse_fasta.py index 9d531d0..7418718 100644 --- a/src/moddotplot/parse_fasta.py +++ b/src/moddotplot/parse_fasta.py @@ -286,9 +286,7 @@ def _fetch_indexed_region( raise ValueError( f"Compressed FASTA {filename!s} does not have a usable BGZF .gzi index" ) - raw_sequence = _read_bgzf_range( - filename, bgzf_index, start_byte, byte_count - ) + raw_sequence = _read_bgzf_range(filename, bgzf_index, start_byte, byte_count) else: with open(filename, "rb") as fasta: fasta.seek(start_byte) @@ -447,9 +445,7 @@ def iter_fasta_records( {entry.name: entry for entry in fasta_index} if fasta_index else None ) if indexed_entries is not None: - selected_ids = ( - list(indexed_entries) if requested_ids is None else requested_ids - ) + selected_ids = list(indexed_entries) if requested_ids is None else requested_ids missing_ids = [ sequence_id for sequence_id in selected_ids diff --git a/src/moddotplot/static_plots.py b/src/moddotplot/static_plots.py index ae0b8ec..6797569 100755 --- a/src/moddotplot/static_plots.py +++ b/src/moddotplot/static_plots.py @@ -64,7 +64,6 @@ visible_annotation_intervals as _visible_annotation_intervals, ) - REGION_SUFFIX_PATTERN = re.compile(r"(?::\d+-\d+)+$") @@ -1549,19 +1548,14 @@ def _missing_symmetric_rows(dataframe): # MultiIndexes for that common case; a strict start-coordinate ordering # proves that no transposed off-diagonal row can already be present. start_difference = ( - q_start.loc[off_diagonal].to_numpy() - - r_start.loc[off_diagonal].to_numpy() + q_start.loc[off_diagonal].to_numpy() - r_start.loc[off_diagonal].to_numpy() ) if np.all(start_difference < 0) or np.all(start_difference > 0): return dataframe.loc[off_diagonal, render_columns] existing = pd.MultiIndex.from_arrays([q_start, q_end, r_start, r_end]) - mirrored = pd.MultiIndex.from_arrays( - [r_start, r_end, q_start, q_end] - ) - return dataframe.loc[ - off_diagonal & ~mirrored.isin(existing), render_columns - ] + mirrored = pd.MultiIndex.from_arrays([r_start, r_end, q_start, q_end]) + return dataframe.loc[off_diagonal & ~mirrored.isin(existing), render_columns] def _full_plot_limits(dataframe, requested_limit): @@ -2140,9 +2134,11 @@ def create_grid( ) grid_size = axes.shape[0] directional = any( - "direction" in matrix.columns - if isinstance(matrix, pd.DataFrame) - else bool(matrix) and "direction" in matrix[0] + ( + "direction" in matrix.columns + if isinstance(matrix, pd.DataFrame) + else bool(matrix) and "direction" in matrix[0] + ) for matrix in [*singles, *doubles] ) grid_label = "DIRECTION_GRID" if directional else "GRID" diff --git a/tests/test_cli_integration.py b/tests/test_cli_integration.py index fd32770..90266b4 100644 --- a/tests/test_cli_integration.py +++ b/tests/test_cli_integration.py @@ -5,7 +5,6 @@ import sys import xml.etree.ElementTree as ET - PROJECT_ROOT = Path(__file__).resolve().parents[1] diff --git a/tests/test_entrypoints.py b/tests/test_entrypoints.py index 8d12f06..d9c0217 100644 --- a/tests/test_entrypoints.py +++ b/tests/test_entrypoints.py @@ -3,7 +3,6 @@ import subprocess import sys - PROJECT_ROOT = Path(__file__).resolve().parents[1] diff --git a/tests/test_fasta_parser.py b/tests/test_fasta_parser.py index 9267416..20ce4e4 100644 --- a/tests/test_fasta_parser.py +++ b/tests/test_fasta_parser.py @@ -257,9 +257,7 @@ def fail_indexed_fetch(*_args, **_kwargs): monkeypatch.setattr(fasta_parser, "_fetch_indexed_region", fail_indexed_fetch) - assert list(iter_fasta_records(fasta)) == [ - ("alpha", records[0][1], "alpha") - ] + assert list(iter_fasta_records(fasta)) == [("alpha", records[0][1], "alpha")] def test_public_record_iterator_preserves_requested_order_and_reports_missing(tmp_path): diff --git a/tests/test_grid.py b/tests/test_grid.py index a744e6e..278f4a2 100644 --- a/tests/test_grid.py +++ b/tests/test_grid.py @@ -8,7 +8,6 @@ from moddotplot.native_render import FALLBACK_FONT_FAMILY, set_figure_font_family from moddotplot.static_plots import _build_grid_figure, create_grid - BED_HEADER = ( "#query_name", "query_start", diff --git a/tests/test_hash_benchmark.py b/tests/test_hash_benchmark.py index f48f2ea..4cd92bd 100644 --- a/tests/test_hash_benchmark.py +++ b/tests/test_hash_benchmark.py @@ -1,7 +1,6 @@ import importlib.util from pathlib import Path - PROJECT_ROOT = Path(__file__).resolve().parents[1] BENCHMARK_PATH = PROJECT_ROOT / "benchmarks" / "benchmark_hashing.py" diff --git a/tests/test_issue53_memory_safe_plotting.py b/tests/test_issue53_memory_safe_plotting.py index b12b8d5..5062251 100644 --- a/tests/test_issue53_memory_safe_plotting.py +++ b/tests/test_issue53_memory_safe_plotting.py @@ -4,7 +4,6 @@ from moddotplot.static_plots import make_dot, make_dot_final, make_dot_grid, make_tri - GENOME_SIZE = 496_000_000 WINDOW_SIZE = 2_000 diff --git a/tests/test_native_full_plot.py b/tests/test_native_full_plot.py index d9d0a6c..7e8d834 100644 --- a/tests/test_native_full_plot.py +++ b/tests/test_native_full_plot.py @@ -186,6 +186,8 @@ def fail_plotnine_full(*_args, **_kwargs): assert svg.read_bytes().lstrip().startswith(b" Date: Wed, 30 Sep 2026 10:50:41 -0400 Subject: [PATCH 05/16] Fix cross-platform CI validation --- .github/workflows/ci.yml | 2 +- src/moddotplot/static_plots.py | 12 +++++++++--- tests/test_grid.py | 11 ++++++++--- tests/test_packaging_metadata.py | 7 +++++++ 4 files changed, 25 insertions(+), 7 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index e00998a..3d169c8 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -121,7 +121,7 @@ jobs: run: | python -c "import pathlib; wheel = next(pathlib.Path('dist').glob('*.whl')); assert 'cp38-abi3' in wheel.name, wheel.name" python -m pip install --no-deps --force-reinstall dist/*.whl - python -c "import importlib.metadata as m, pathlib, tomllib; expected = tomllib.loads(pathlib.Path('pyproject.toml').read_text())['project']; actual = m.metadata('ModDotPlot'); assert m.version('ModDotPlot') == expected['version']; assert actual['Requires-Python'] == expected['requires-python']" + python -c "import importlib.metadata as m, pathlib, tomllib; from packaging.specifiers import SpecifierSet; expected = tomllib.loads(pathlib.Path('pyproject.toml').read_text())['project']; actual = m.metadata('ModDotPlot'); assert m.version('ModDotPlot') == expected['version']; assert SpecifierSet(actual['Requires-Python']) == SpecifierSet(expected['requires-python'])" python -c "from moddotplot import _nthash; hashes, mask = _nthash.hash_kmers('ACGTACGT', 3, True); assert len(hashes) == 48 and mask == b''" - name: Store distributions diff --git a/src/moddotplot/static_plots.py b/src/moddotplot/static_plots.py index 6797569..0f81c1f 100755 --- a/src/moddotplot/static_plots.py +++ b/src/moddotplot/static_plots.py @@ -156,7 +156,11 @@ def _fit_grid_sequence_labels(figure, axes): title_box = title.get_window_extent(renderer=renderer) axis_box = axis.get_window_extent(renderer=renderer) if title_box.width: - ratios.append(axis_box.width * 0.9 / title_box.width) + # Font fallback and hinting can change the final extent by a + # fraction of a pixel on another backend. Leave enough + # headroom that fitted labels remain inside their panels + # after the renderer rounds the scaled font size. + ratios.append(axis_box.width * 0.88 / title_box.width) for axis in axes[:, 0]: label = axis.yaxis.label @@ -164,7 +168,7 @@ def _fit_grid_sequence_labels(figure, axes): label_box = label.get_window_extent(renderer=renderer) axis_box = axis.get_window_extent(renderer=renderer) if label_box.height: - ratios.append(axis_box.height * 0.9 / label_box.height) + ratios.append(axis_box.height * 0.88 / label_box.height) artists = [axis.title for axis in axes[0, :]] + [ axis.yaxis.label for axis in axes[:, 0] @@ -2077,7 +2081,9 @@ def _build_grid_figure( axis_title, xy=(0, 0.5), xycoords=bottom_left_axis.yaxis.label, - xytext=(-4, 0), + # Keep the shared title visibly separate from the adjacent row + # label under both Helvetica and Matplotlib's fallback fonts. + xytext=(-6, 0), textcoords="offset points", ha="center", va="center", diff --git a/tests/test_grid.py b/tests/test_grid.py index 278f4a2..7ddb513 100644 --- a/tests/test_grid.py +++ b/tests/test_grid.py @@ -4,6 +4,7 @@ from pathlib import Path import pytest +import moddotplot.static_plots as static_plots from moddotplot.const import DIRECTION_COLORS from moddotplot.native_render import FALLBACK_FONT_FAMILY, set_figure_font_family from moddotplot.static_plots import _build_grid_figure, create_grid @@ -331,7 +332,8 @@ def test_full_self_comparison_input_is_not_mirrored_twice(): plt.close(figure) -def test_grid_region_names_fit_panels_and_numeric_labels_are_doubled(): +def test_grid_region_names_fit_panels_and_numeric_labels_are_doubled(monkeypatch): + monkeypatch.setattr(static_plots, "DEFAULT_FONT_FAMILY", FALLBACK_FONT_FAMILY) names = [ "PAN010.chr14.haplotype1.paternal:1-4000000", "PAN010.chr14.haplotype2.maternal:1-4000000", @@ -430,7 +432,10 @@ def test_grid_exact_bounds_do_not_expand_to_next_nice_tick(): ("axis_end", "unit"), [(100_000, "Kbp"), (103_156_783, "Mbp"), (500_000_000, "Gbp")], ) -def test_grid_labels_genomic_units_only_on_bottom_left_cell(axis_end, unit): +def test_grid_labels_genomic_units_only_on_bottom_left_cell( + axis_end, unit, monkeypatch +): + monkeypatch.setattr(static_plots, "DEFAULT_FONT_FAMILY", FALLBACK_FONT_FAMILY) kwargs = _basic_two_sequence_grid( xlim=(1, axis_end), axes_label=None, @@ -470,7 +475,7 @@ def test_grid_labels_genomic_units_only_on_bottom_left_cell(axis_end, unit): assert horizontal_bounds.x0 + horizontal_bounds.width / 2 <= cell_bounds.x1 assert cell_bounds.y0 <= vertical_bounds.y0 + vertical_bounds.height / 2 assert vertical_bounds.y0 + vertical_bounds.height / 2 <= cell_bounds.y1 - assert vertical_bounds.x1 <= row_label_bounds.x0 + assert vertical_bounds.x1 + 1 <= row_label_bounds.x0 for bounds in (horizontal_bounds, vertical_bounds): assert figure_bounds.contains(bounds.x0, bounds.y0) assert figure_bounds.contains(bounds.x1, bounds.y1) diff --git a/tests/test_packaging_metadata.py b/tests/test_packaging_metadata.py index 0425535..664d7d8 100644 --- a/tests/test_packaging_metadata.py +++ b/tests/test_packaging_metadata.py @@ -92,6 +92,13 @@ def test_ci_covers_every_supported_python_minor(): assert ' - "3.10"' not in workflow +def test_package_ci_compares_python_specifiers_semantically(): + workflow = (PROJECT_ROOT / ".github/workflows/ci.yml").read_text() + + assert "SpecifierSet(actual['Requires-Python'])" in workflow + assert "SpecifierSet(expected['requires-python'])" in workflow + + def test_release_workflow_is_tag_gated_and_uses_trusted_publishing(): workflow = (PROJECT_ROOT / ".github/workflows/publish-to-pypi.yml").read_text() setup_config = (PROJECT_ROOT / "setup.cfg").read_text() From 80d2c72d85f7ed24b234ff9979f3375c7f02ca92 Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Wed, 30 Sep 2026 14:25:05 -0400 Subject: [PATCH 06/16] Refine full plot typography --- benchmarks/benchmark_hashing.py | 8 +++--- src/moddotplot/static_plots.py | 14 +++++----- tests/test_native_full_plot.py | 47 +++++++++++++++++++++++++++++++-- 3 files changed, 57 insertions(+), 12 deletions(-) diff --git a/benchmarks/benchmark_hashing.py b/benchmarks/benchmark_hashing.py index 27a2667..b620997 100644 --- a/benchmarks/benchmark_hashing.py +++ b/benchmarks/benchmark_hashing.py @@ -142,10 +142,10 @@ def benchmark(sequence: str, k: int, repeats: int, seed: int) -> Dict[str, objec ) } if mmh3_module is not None: - implementations["legacy_mmh3"] = ( - lambda canonical=canonical: _legacy_mmh3_hashes( - sequence, k, canonical, mmh3_module - ) + implementations[ + "legacy_mmh3" + ] = lambda canonical=canonical: _legacy_mmh3_hashes( + sequence, k, canonical, mmh3_module ) raw: Dict[str, List[float]] = {name: [] for name in implementations} diff --git a/src/moddotplot/static_plots.py b/src/moddotplot/static_plots.py index 0f81c1f..86f4149 100755 --- a/src/moddotplot/static_plots.py +++ b/src/moddotplot/static_plots.py @@ -1647,9 +1647,11 @@ def _build_full_figure( breaks=breaks, ) _divisor, unit = genomic_scale(region_end) + genomic_axis_title_size = clamped_font_size(width, 1.575) + sequence_axis_title_size = clamped_font_size(width, 3.0) axis.set_xlabel( f"Genomic Position ({unit})", - fontsize=clamped_font_size(width, 2.8), + fontsize=genomic_axis_title_size, fontfamily=DEFAULT_FONT_FAMILY, ) axis.tick_params( @@ -1668,16 +1670,16 @@ def _build_full_figure( # placing the descriptive title independently above them. axis.set_title( display_x, - fontsize=clamped_font_size(width, 1.2), + fontsize=sequence_axis_title_size, fontfamily=DEFAULT_FONT_FAMILY, - pad=5, + pad=9, ) axis.set_ylabel( display_y, - fontsize=clamped_font_size(width, 1.2), + fontsize=sequence_axis_title_size, fontfamily=DEFAULT_FONT_FAMILY, rotation=-90, - labelpad=16, + labelpad=28, ) axis.yaxis.set_label_position("right") @@ -1690,7 +1692,7 @@ def _build_full_figure( title, fontsize=max(MIN_TITLE_SIZE, title_size), fontfamily=DEFAULT_FONT_FAMILY, - y=0.975, + y=0.92, ) figure.subplots_adjust( left=0.14, diff --git a/tests/test_native_full_plot.py b/tests/test_native_full_plot.py index 7e8d834..b2f1471 100644 --- a/tests/test_native_full_plot.py +++ b/tests/test_native_full_plot.py @@ -77,6 +77,48 @@ def test_native_full_self_plot_mirrors_only_missing_triangle(deraster): plt.close(figure) +@pytest.mark.parametrize("is_pairwise", [False, True]) +def test_native_full_plot_uses_requested_typography(is_pairwise): + name_y = "chr21" if is_pairwise else "chr20" + frame = _processed_tiles([("chr20", 0, 1_000_000, name_y, 0, 1_000_000, 100.0, 1)]) + + figure = static_plots._build_full_figure( + sdf=frame, + name_x="chr20", + name_y=name_y, + palette="Spectral_11", + palette_orientation="+", + custom_colors=None, + axes_labels=None, + xlim=(0, 66_000_000), + deraster=False, + width=9, + is_pairwise=is_pairwise, + ) + try: + axis = figure.axes[0] + assert axis.xaxis.label.get_fontsize() == pytest.approx(14.175) + assert axis.title.get_fontsize() == pytest.approx(27.0) + assert axis.yaxis.label.get_fontsize() == pytest.approx(27.0) + assert axis.yaxis.labelpad == pytest.approx(28) + assert figure._suptitle.get_position()[1] == pytest.approx(0.92) + title_offset = axis.title.get_transform().transform((0, 0))[1] + axes_origin = axis.transAxes.transform((0, 0))[1] + assert title_offset - axes_origin == pytest.approx(9 * figure.dpi / 72) + + figure.canvas.draw() + renderer = figure.canvas.get_renderer() + axis_box = axis.get_window_extent(renderer) + assert axis.title.get_window_extent(renderer).y0 > axis_box.y1 + assert axis.yaxis.label.get_window_extent(renderer).x0 > axis_box.x1 + assert ( + figure._suptitle.get_window_extent(renderer).y0 + > axis.title.get_window_extent(renderer).y1 + ) + finally: + plt.close(figure) + + def test_native_full_self_plot_does_not_duplicate_loaded_symmetric_rows(): frame = _processed_tiles( [ @@ -220,7 +262,8 @@ def test_native_comparative_direction_colors_preserve_hue_and_ani_strength(): assert colors[0, 2] > colors[0, 0] assert colors[1, 0] > colors[1, 2] assert np.mean(colors[1, :3]) < np.mean(colors[0, :3]) - assert figure.axes[0].get_title() == "query" - assert figure.axes[0].get_ylabel() == "reference" + axis = figure.axes[0] + assert axis.get_title() == "query" + assert axis.get_ylabel() == "reference" finally: plt.close(figure) From 37dc34da9bae92f1f627bf17e692825dcc4307d8 Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Wed, 30 Sep 2026 17:27:41 -0400 Subject: [PATCH 07/16] Restore Python 3.10 and make interactive dependencies optional --- .github/workflows/ci.yml | 33 ++++++- .github/workflows/publish-to-pypi.yml | 6 +- README.md | 20 ++-- pyproject.toml | 23 +++-- src/moddotplot/estimate_identity.py | 13 ++- src/moddotplot/interactive.py | 38 +++++++- src/moddotplot/moddotplot.py | 54 ++++++++++- src/moddotplot/optional_dependencies.py | 13 +++ tests/test_entrypoints.py | 8 +- tests/test_optional_dependencies.py | 121 ++++++++++++++++++++++++ tests/test_packaging_metadata.py | 52 ++++++++-- 11 files changed, 350 insertions(+), 31 deletions(-) create mode 100644 src/moddotplot/optional_dependencies.py create mode 100644 tests/test_optional_dependencies.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 3d169c8..1328cf0 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -15,6 +15,36 @@ concurrency: cancel-in-progress: true jobs: + base-install: + name: Base install without optional dependencies + runs-on: ubuntu-latest + timeout-minutes: 15 + + steps: + - name: Check out source + uses: actions/checkout@v7 + with: + persist-credentials: false + + - name: Set up Python 3.10 + uses: actions/setup-python@v7 + with: + python-version: "3.10" + cache: pip + cache-dependency-path: pyproject.toml + + - name: Install the base package + run: | + python -m pip install --upgrade pip + python -m pip install . + + - name: Verify optional packages were not installed + run: | + python -m pip check + python -c "import importlib.util; assert all(importlib.util.find_spec(name) is None for name in ('cooler', 'dash', 'plotly'))" + python -c "import sys; import moddotplot.moddotplot; assert not {'cooler', 'dash', 'plotly', 'moddotplot.interactive'}.intersection(sys.modules)" + python -m moddotplot --help + tests: name: Tests (Python ${{ matrix.python-version }}) runs-on: ubuntu-latest @@ -23,6 +53,7 @@ jobs: fail-fast: false matrix: python-version: + - "3.10" - "3.11" - "3.12" - "3.13" @@ -46,7 +77,7 @@ jobs: - name: Install package and test dependencies run: | python -m pip install --upgrade pip - python -m pip install --editable ".[test]" + python -m pip install --editable ".[test,interactive]" - name: Run unit and integration tests run: >- diff --git a/.github/workflows/publish-to-pypi.yml b/.github/workflows/publish-to-pypi.yml index 903deed..d35a953 100644 --- a/.github/workflows/publish-to-pypi.yml +++ b/.github/workflows/publish-to-pypi.yml @@ -42,7 +42,7 @@ jobs: - name: Install package and tests run: | python -m pip install --upgrade pip - python -m pip install ".[test]" + python -m pip install ".[test,interactive]" - name: Run the release test suite run: python -m pytest @@ -100,13 +100,13 @@ jobs: with: persist-credentials: false - # The extension uses CPython's stable ABI. Building once with CPython 3.11 + # The extension uses CPython's stable ABI. Building once with CPython 3.10 # produces a cp38-abi3 wheel that supports every Python version declared # by the package without compiling the same binary repeatedly. - name: Build stable-ABI wheel uses: pypa/cibuildwheel@v4.2.0 env: - CIBW_BUILD: "cp311-*" + CIBW_BUILD: "cp310-*" CIBW_SKIP: "*-musllinux_*" CIBW_ARCHS_MACOS: universal2 CIBW_TEST_COMMAND: >- diff --git a/README.md b/README.md index fa9cbe8..73fd3d2 100644 --- a/README.md +++ b/README.md @@ -1,9 +1,9 @@ ![](images/logo.png) + --- [![PyPI](https://img.shields.io/pypi/v/ModDotPlot?color=blue&label=PyPI)](https://pypi.org/project/ModDotPlot/) [![CI](https://github.com/marbl/ModDotPlot/actions/workflows/ci.yml/badge.svg)](https://github.com/marbl/ModDotPlot/actions/workflows/ci.yml) -- [](#) - [Cite](#cite) - [About](#about) - [Installation](#installation) @@ -51,7 +51,7 @@ If you're interested in learning more about _ModDotPlot_ and how to visualize ta ## Installation -_ModDotPlot_ can be installed by running `pip install moddotplot`. Version 1.0.0 supports Python 3.11 through 3.14 and uses the current Matplotlib 3.11 and Plotnine 0.15 release lines. Alternatively, you can download the current release from GitHub by using: +_ModDotPlot_ can be installed for static plotting by running `pip install moddotplot`. Interactive plotting and Cooler export use optional dependencies; install them with `pip install "ModDotPlot[interactive]"`. ModDotPlot supports Python 3.10 through 3.14, using Matplotlib 3.10.9 or newer on Python 3.10, Matplotlib 3.11.2 or newer on later Python versions, and the Plotnine 0.15 release line. Alternatively, you can download the current release from GitHub by using: ``` git clone https://github.com/marbl/ModDotPlot.git @@ -65,12 +65,18 @@ python -m venv venv source venv/bin/activate ``` -Once activated, you can install the required dependencies: +Once activated, install either the base package for static plotting: ``` python -m pip install . ``` +or include the optional interactive plotting and Cooler dependencies: + +``` +python -m pip install ".[interactive]" +``` + Finally, confirm that the installation was installed correctly and that your version is up to date by running `moddotplot -h`: ``` __ __ _ _____ _ _____ _ _ @@ -138,7 +144,9 @@ moddotplot interactive Interactive mode is deprecated and maintenance-only. It remains available, but will not receive new features. It runs only when the `interactive` subcommand is -explicitly provided. +explicitly provided. Install its optional dependencies with +`pip install "ModDotPlot[interactive]"`, or `pip install ".[interactive]"` from +a source checkout. Running _ModDotPlot_ in interactive mode will launch a [Dash application](https://plotly.com/dash/) on your machine's localhost. Open any web browser and go to `http://127.0.0.1:` to view the interactive plot (this should happen automatically, but depending on your environment you might need to copy and paste this URL into your web browser). Running `Ctrl+C` on the command line will exit the Dash application. The default port number used by Dash is `8050`, but this can be customized using the `--port` command (see [interactive mode commands](#interactive-mode-commands) for further info, and [Sample run - Port Forwarding](#sample-run---port-forwarding) for tips on running interactive mode on an HPC environment). @@ -224,7 +232,7 @@ IDs are errors. When using a config file, provide the same list under the `--cooler ` -If set, will output a matrix as a cooler file for each input sequence, in addition to a bedpe file. +If set, will output a matrix as a Cooler file for each input sequence, in addition to a BEDPE file. Cooler support is part of the optional dependency set installed with `pip install "ModDotPlot[interactive]"` (or `pip install ".[interactive]"` from a source checkout). `--no-bedpe ` @@ -546,4 +554,4 @@ For bug reports or general usage questions, please raise a GitHub issue, or emai - Mac users might encounter the following unexpected command line output: `/bin/sh: lscpu: command not found`. This is a known issue with Plotnine, the Python plotting library used by ModDotPlot. This can be safely ignored. -- The error ` UserWarning: h5py is running against HDF5 1.xx.x when it was built against 1.xx.x, this may cause problems` is due to the h5py library used by cooler having conflicting versions in the dependency tree. This can also be safely ignored, but if you want to remove this message run `pip uninstall -y h5py` `pip install --no-binary=h5py h5py` +- When the optional `ModDotPlot[interactive]` dependencies are installed, Cooler may report `UserWarning: h5py is running against HDF5 1.xx.x when it was built against 1.xx.x, this may cause problems`. This can be safely ignored. To remove the warning, reinstall h5py against the local HDF5 library with `pip uninstall -y h5py` followed by `pip install --no-binary=h5py h5py`. diff --git a/pyproject.toml b/pyproject.toml index 6bceec0..b12e4d0 100755 --- a/pyproject.toml +++ b/pyproject.toml @@ -5,19 +5,16 @@ build-backend = "setuptools.build_meta" [project] name = "ModDotPlot" version = "1.0.0" -requires-python = ">=3.11,<3.15" +requires-python = ">=3.10,<3.15" dependencies = [ "pandas", - "matplotlib>=3.11.2", - "plotly", - "dash", + "matplotlib>=3.10.9; python_version < '3.11'", + "matplotlib>=3.11.2; python_version >= '3.11'", "plotnine>=0.15.8,<0.16", "palettable", "setproctitle", "numpy", "scipy", - "pillow", - "cooler", ] authors = [ {name = "Alex Sweeten", email = "alex.sweeten@nih.gov"}, @@ -30,13 +27,27 @@ maintainers = [ readme = {file = "README.md", content-type = "text/markdown"} license = {file = "LICENSE"} keywords = ["dotplot", "sketching", "modimizer", "heatmap"] +classifiers = [ + "Programming Language :: Python :: 3 :: Only", + "Programming Language :: Python :: 3.10", + "Programming Language :: Python :: 3.11", + "Programming Language :: Python :: 3.12", + "Programming Language :: Python :: 3.13", + "Programming Language :: Python :: 3.14", +] [project.scripts] moddotplot = "moddotplot.__main__:main" [project.optional-dependencies] +interactive = [ + "cooler", + "dash>=2.9", + "plotly", +] # development dependency groups test = [ "pytest", "pytest-cov", + "tomli>=1.1; python_version < '3.11'", ] diff --git a/src/moddotplot/estimate_identity.py b/src/moddotplot/estimate_identity.py index 094af69..3dbc470 100644 --- a/src/moddotplot/estimate_identity.py +++ b/src/moddotplot/estimate_identity.py @@ -11,10 +11,10 @@ from palettable import colorbrewer from typing import Collection, Hashable, List, Set, Dict, Tuple import pandas as pd -import cooler from scipy.sparse import csr_matrix from moddotplot import _nthash +from moddotplot.optional_dependencies import OptionalDependencyError from moddotplot.parse_fasta import printProgressBar @@ -713,6 +713,16 @@ def convertMatrixToBed( return bed +def require_cooler_dependency(): + """Return Cooler or explain how to install the optional export support.""" + + try: + import cooler + except ModuleNotFoundError as error: + raise OptionalDependencyError("Cooler export") from error + return cooler + + def convertMatrixToCool( matrix, window_size, @@ -740,6 +750,7 @@ def convertMatrixToCool( chromsizes (dict): Dict of chromosome lengths, e.g. {"chr1": 248956422}. output_cool (str): Path to save cooler file. """ + cooler = require_cooler_dependency() rows, cols = matrix.shape # ---- build bin table ---- diff --git a/src/moddotplot/interactive.py b/src/moddotplot/interactive.py index 6173a8d..3f2ea73 100644 --- a/src/moddotplot/interactive.py +++ b/src/moddotplot/interactive.py @@ -1,4 +1,3 @@ -import plotly.express as px from moddotplot.estimate_identity import ( getInteractiveColor, getMatchingColors, @@ -9,13 +8,11 @@ ) import numpy as np -import dash -from dash import Input, Output, html, dcc, State import math -import plotly.graph_objs as go import logging import os from moddotplot.annotations import visible_annotation_intervals +from moddotplot.optional_dependencies import OptionalDependencyError from moddotplot.parse_fasta import extractRegion INTERACTIVE_FONT_FAMILY = "Helvetica, 'DejaVu Sans', sans-serif" @@ -25,6 +22,29 @@ log.setLevel(logging.ERROR) +def require_interactive_dependencies(): + """Load the optional Dash and Plotly stack on demand.""" + + try: + import dash + from dash import Input, Output, State, dcc, html + import plotly.express as px + import plotly.graph_objs as go + except ModuleNotFoundError as error: + raise OptionalDependencyError("Interactive mode") from error + + return { + "dash": dash, + "Input": Input, + "Output": Output, + "State": State, + "dcc": dcc, + "html": html, + "px": px, + "go": go, + } + + def _plotly_annotation_color(color): """Convert a shared annotation color into a Plotly-compatible value.""" @@ -263,6 +283,16 @@ def run_dash( output_dir, annotations=None, ): + dependencies = require_interactive_dependencies() + dash = dependencies["dash"] + Input = dependencies["Input"] + Output = dependencies["Output"] + State = dependencies["State"] + dcc = dependencies["dcc"] + html = dependencies["html"] + px = dependencies["px"] + go = dependencies["go"] + # Run Dash app app = dash.Dash(__name__, prevent_initial_callbacks="initial_duplicate") app.title = "ModDotPlot" diff --git a/src/moddotplot/moddotplot.py b/src/moddotplot/moddotplot.py index 3158f26..683fded 100755 --- a/src/moddotplot/moddotplot.py +++ b/src/moddotplot/moddotplot.py @@ -20,6 +20,7 @@ convertMatrixToBedDataFrame, iterMatrixToBedChunks, convertMatrixToCool, + require_cooler_dependency, createSelfMatrix, createPairwiseMatrix, create_self_matrix_from_sketches, @@ -28,9 +29,9 @@ ModimizerSketchCache, partitionOverlaps, ) -from moddotplot.interactive import interactive_axis_bounds, run_dash from moddotplot.annotations import read_annotation_beds from moddotplot.const import ASCII_ART, VERSION +from moddotplot.optional_dependencies import OptionalDependencyError import argparse from concurrent.futures import ProcessPoolExecutor, as_completed @@ -52,6 +53,9 @@ read_df_from_file = None create_plots = None create_grid = None +interactive_axis_bounds = None +run_dash = None +require_interactive_dependencies = None COMMANDS = frozenset({"interactive", "static"}) INTERACTIVE_DEPRECATION_MESSAGE = ( @@ -73,6 +77,21 @@ def _load_static_plotting(): create_grid = static_plots.create_grid +def _load_interactive_plotting(): + """Load Dash/Plotly integration only when an interactive UI is requested.""" + + global interactive_axis_bounds, run_dash, require_interactive_dependencies + + from moddotplot import interactive + + if interactive_axis_bounds is None: + interactive_axis_bounds = interactive.interactive_axis_bounds + if run_dash is None: + run_dash = interactive.run_dash + if require_interactive_dependencies is None: + require_interactive_dependencies = interactive.require_interactive_dependencies + + def get_parser(): """ Argument parsing for stand-alone runs. @@ -359,7 +378,12 @@ def get_parser(): ) static_parser.add_argument( - "--cooler", action="store_true", help="Output matrix to cooler file." + "--cooler", + action="store_true", + help=( + "Output matrix to a Cooler file. Requires the optional " + "ModDotPlot[interactive] dependencies." + ), ) static_parser.add_argument( @@ -1248,6 +1272,25 @@ def main(): print(ASCII_ART) print(f"v{VERSION} \n") args = parse_args() + + # Matrix-only interactive exports do not use Dash or Plotly. Every path + # that opens the interactive UI validates its extra before doing expensive + # sequence work. + matrix_only_interactive = ( + args.command == "interactive" + and args.save + and args.no_plot + and not getattr(args, "load", None) + ) + needs_interactive_ui = args.command == "interactive" and not matrix_only_interactive + if needs_interactive_ui: + _load_interactive_plotting() + try: + require_interactive_dependencies() + except OptionalDependencyError as error: + print(f"Error: {error}", file=sys.stderr) + sys.exit(2) + summary_writer = PlotSummaryWriter(shlex.join(sys.argv)) annotation_df = None if args.command == "interactive" and args.bed: @@ -1305,6 +1348,13 @@ def main(): config = json.load(f) _apply_static_config(args, config) + if args.cooler: + try: + require_cooler_dependency() + except OptionalDependencyError as error: + print(f"Error: {error}", file=sys.stderr) + sys.exit(2) + try: args.processes = _validated_process_count(args.processes) except ValueError as error: diff --git a/src/moddotplot/optional_dependencies.py b/src/moddotplot/optional_dependencies.py new file mode 100644 index 0000000..23bf138 --- /dev/null +++ b/src/moddotplot/optional_dependencies.py @@ -0,0 +1,13 @@ +"""Shared errors and installation guidance for optional features.""" + +INTERACTIVE_INSTALL_COMMAND = 'python -m pip install "ModDotPlot[interactive]"' + + +class OptionalDependencyError(ImportError): + """Raised when a requested feature is missing its optional dependencies.""" + + def __init__(self, feature): + super().__init__( + f"{feature} requires the optional interactive dependencies. " + f"Install them with: {INTERACTIVE_INSTALL_COMMAND}" + ) diff --git a/tests/test_entrypoints.py b/tests/test_entrypoints.py index d9c0217..c601c62 100644 --- a/tests/test_entrypoints.py +++ b/tests/test_entrypoints.py @@ -19,14 +19,18 @@ def _run_module(*arguments): ) -def test_help_does_not_import_static_rendering_stack(): +def test_base_entrypoint_imports_neither_rendering_nor_optional_stacks(): result = subprocess.run( [ sys.executable, "-c", ( "import sys; import moddotplot.moddotplot; " - "assert 'moddotplot.static_plots' not in sys.modules" + "unexpected = {" + "'moddotplot.static_plots', 'moddotplot.interactive', " + "'cooler', 'dash', 'plotly'" + "}.intersection(sys.modules); " + "assert not unexpected, sorted(unexpected)" ), ], cwd=PROJECT_ROOT, diff --git a/tests/test_optional_dependencies.py b/tests/test_optional_dependencies.py new file mode 100644 index 0000000..26eb943 --- /dev/null +++ b/tests/test_optional_dependencies.py @@ -0,0 +1,121 @@ +import builtins +import sys + +import numpy as np +import pytest + +import moddotplot.interactive as interactive +import moddotplot.moddotplot as cli +from moddotplot.estimate_identity import convertMatrixToCool, require_cooler_dependency +from moddotplot.optional_dependencies import OptionalDependencyError + + +INSTALL_HINT = 'python -m pip install "ModDotPlot[interactive]"' + + +def _block_import(monkeypatch, missing_package): + real_import = builtins.__import__ + + def import_without_optional_package( + name, globals=None, locals=None, fromlist=(), level=0 + ): + if name == missing_package or name.startswith(f"{missing_package}."): + raise ModuleNotFoundError( + f"No module named '{missing_package}'", name=missing_package + ) + return real_import(name, globals, locals, fromlist, level) + + monkeypatch.setattr(builtins, "__import__", import_without_optional_package) + + +def test_missing_cooler_reports_the_interactive_extra_install_command(monkeypatch): + _block_import(monkeypatch, "cooler") + + with pytest.raises(OptionalDependencyError) as exc_info: + require_cooler_dependency() + + assert "Cooler export requires" in str(exc_info.value) + assert INSTALL_HINT in str(exc_info.value) + assert isinstance(exc_info.value.__cause__, ModuleNotFoundError) + + +@pytest.mark.parametrize("missing_package", ["dash", "plotly"]) +def test_missing_interactive_package_reports_the_extra_install_command( + monkeypatch, missing_package +): + _block_import(monkeypatch, missing_package) + + with pytest.raises(OptionalDependencyError) as exc_info: + interactive.require_interactive_dependencies() + + assert "Interactive mode requires" in str(exc_info.value) + assert INSTALL_HINT in str(exc_info.value) + assert isinstance(exc_info.value.__cause__, ModuleNotFoundError) + + +def test_static_cooler_request_exits_with_actionable_error(monkeypatch, capsys): + def missing_cooler(): + raise OptionalDependencyError("Cooler export") + + monkeypatch.setattr(cli, "require_cooler_dependency", missing_cooler) + monkeypatch.setattr( + sys, + "argv", + [ + "moddotplot", + "static", + "--fasta", + "sequence.fa", + "--cooler", + "--no-plot", + ], + ) + + with pytest.raises(SystemExit) as exc_info: + cli.main() + + assert exc_info.value.code == 2 + assert INSTALL_HINT in capsys.readouterr().err + + +def test_cooler_extra_writes_a_readable_comparative_matrix(tmp_path): + cooler = pytest.importorskip("cooler") + output = tmp_path / "comparison.cool" + + convertMatrixToCool( + matrix=np.asarray([[1.0, 0.9], [0.9, 1.0]]), + window_size=10, + id_threshold=80, + x_name="chr1", + y_name="chr2", + self_identity=False, + x_offset=0, + y_offset=0, + chromsizes={"chr1": 20, "chr2": 20}, + output_cool=str(output), + ) + + matrix = cooler.Cooler(str(output)) + assert tuple(matrix.shape) == (4, 4) + assert matrix.info["nnz"] == 4 + + +def test_interactive_request_exits_with_actionable_error(monkeypatch, capsys): + def missing_interactive_dependencies(): + raise OptionalDependencyError("Interactive mode") + + monkeypatch.setattr(cli, "_load_interactive_plotting", lambda: None) + monkeypatch.setattr( + cli, "require_interactive_dependencies", missing_interactive_dependencies + ) + monkeypatch.setattr( + sys, + "argv", + ["moddotplot", "interactive", "--fasta", "sequence.fa"], + ) + + with pytest.raises(SystemExit) as exc_info: + cli.main() + + assert exc_info.value.code == 2 + assert INSTALL_HINT in capsys.readouterr().err diff --git a/tests/test_packaging_metadata.py b/tests/test_packaging_metadata.py index 664d7d8..65b375a 100644 --- a/tests/test_packaging_metadata.py +++ b/tests/test_packaging_metadata.py @@ -21,15 +21,46 @@ def test_runtime_and_distribution_versions_match(): def test_declared_python_floor_matches_documentation(): - assert project_metadata()["requires-python"] == ">=3.11,<3.15" + metadata = project_metadata() + assert metadata["requires-python"] == ">=3.10,<3.15" + for minor in range(10, 15): + assert f"Programming Language :: Python :: 3.{minor}" in metadata["classifiers"] readme = (PROJECT_ROOT / "README.md").read_text() - assert "supports Python 3.11 through 3.14" in readme + assert "supports Python 3.10 through 3.14" in readme def test_plotnine_supports_declared_python_floor(): dependencies = project_metadata()["dependencies"] assert "plotnine>=0.15.8,<0.16" in dependencies - assert "matplotlib>=3.11.2" in dependencies + assert "matplotlib>=3.10.9; python_version < '3.11'" in dependencies + assert "matplotlib>=3.11.2; python_version >= '3.11'" in dependencies + + +def test_interactive_dependencies_are_not_installed_with_the_core_package(): + metadata = project_metadata() + core_dependencies = metadata["dependencies"] + core_names = { + dependency.split(";", 1)[0] + .split("[", 1)[0] + .split("=", 1)[0] + .split("<", 1)[0] + .split(">", 1)[0] + .strip() + .lower() + for dependency in core_dependencies + } + + assert core_names.isdisjoint({"cooler", "dash", "plotly", "pillow"}) + assert set(metadata["optional-dependencies"]["interactive"]) == { + "cooler", + "dash>=2.9", + "plotly", + } + + +def test_python_310_test_dependencies_include_tomli(): + test_dependencies = project_metadata()["optional-dependencies"]["test"] + assert "tomli>=1.1; python_version < '3.11'" in test_dependencies def test_mmh3_is_not_a_runtime_dependency(): @@ -87,9 +118,8 @@ def test_svg_composition_dependencies_are_not_runtime_dependencies(): def test_ci_covers_every_supported_python_minor(): workflow = (PROJECT_ROOT / ".github/workflows/ci.yml").read_text() - for minor in range(11, 15): + for minor in range(10, 15): assert f' - "3.{minor}"' in workflow - assert ' - "3.10"' not in workflow def test_package_ci_compares_python_specifiers_semantically(): @@ -99,6 +129,16 @@ def test_package_ci_compares_python_specifiers_semantically(): assert "SpecifierSet(expected['requires-python'])" in workflow +def test_full_test_workflows_install_interactive_test_dependencies(): + ci_workflow = (PROJECT_ROOT / ".github/workflows/ci.yml").read_text() + release_workflow = ( + PROJECT_ROOT / ".github/workflows/publish-to-pypi.yml" + ).read_text() + + assert 'python -m pip install --editable ".[test,interactive]"' in ci_workflow + assert 'python -m pip install ".[test,interactive]"' in release_workflow + + def test_release_workflow_is_tag_gated_and_uses_trusted_publishing(): workflow = (PROJECT_ROOT / ".github/workflows/publish-to-pypi.yml").read_text() setup_config = (PROJECT_ROOT / "setup.cfg").read_text() @@ -108,7 +148,7 @@ def test_release_workflow_is_tag_gated_and_uses_trusted_publishing(): assert "id-token: write" in workflow assert "pypa/gh-action-pypi-publish@release/v1" in workflow assert "pypa/cibuildwheel@v4.2.0" in workflow - assert 'CIBW_BUILD: "cp311-*"' in workflow + assert 'CIBW_BUILD: "cp310-*"' in workflow assert "py_limited_api = cp38" in setup_config for runner in ("ubuntu-latest", "macos-15-intel", "windows-latest"): assert f" - {runner}" in workflow From 1320b8b3d4418924c664436060b7aee9358825f0 Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Thu, 1 Oct 2026 22:59:37 -0400 Subject: [PATCH 08/16] Optimize comparative plotting runtime and memory --- README.md | 78 +- benchmarks/benchmark_containment_matrix.py | 1 - src/moddotplot/__main__.py | 32 +- src/moddotplot/estimate_identity.py | 95 +-- src/moddotplot/moddotplot.py | 838 ++++++++++++++++++++- src/moddotplot/parse_fasta.py | 62 +- tests/test_algorithms.py | 45 +- tests/test_cli_integration.py | 138 ++++ tests/test_cli_runtime.py | 92 +++ tests/test_entrypoints.py | 23 + tests/test_sketch_cache.py | 2 +- tests/test_sparse_containment.py | 5 +- 12 files changed, 1253 insertions(+), 158 deletions(-) diff --git a/README.md b/README.md index 73fd3d2..406096e 100644 --- a/README.md +++ b/README.md @@ -175,6 +175,11 @@ K-mer size to use. This should be large enough to distinguish unique k-mers with Name of output directory for bed file & plots. Default is current working directory. +`--quiet` + +Suppress all console output, including warnings and errors. The process exit +status still indicates whether the run succeeded. + `-id / --identity ` Minimum sequence identity cutoff threshold. Default is 86. While it is possible to go as low as 50% sequence identity, anything below 80% is not recommended. @@ -230,6 +235,15 @@ fallback (so `chr1` selects `Chr1`). Unknown, ambiguous, and duplicate requested IDs are errors. When using a config file, provide the same list under the `sequence` key, for example `"sequence": ["chr1", "chr2"]`. +`--pairs ` + +Limit `--compare` or `--compare-only` to explicitly requested sequence pairs. +The file contains two whitespace-delimited FASTA identifiers per line, in +x-axis then y-axis order. Blank lines and lines beginning with `#` are ignored. +Duplicate, self, unknown, and ambiguous pairs are errors. Indexed FASTA input +is required so each requested record can be fetched without scanning the full +genome. + `--cooler ` If set, will output a matrix as a Cooler file for each input sequence, in addition to a BEDPE file. Cooler support is part of the optional dependency set installed with `pip install "ModDotPlot[interactive]"` (or `pip install ".[interactive]"` from a source checkout). @@ -248,13 +262,28 @@ Save .bedpe to file, but skip rendering of plots. `--processes <1-4>` -Set the number of independent chromosome workers for self-only static runs. -When omitted, ModDotPlot uses two workers while rendering or up to four for a +Set the number of independent chromosome or comparison-group workers. When +omitted, ModDotPlot uses two workers while rendering or up to four for a `--no-plot` run on multi-record FASTA files that have random-access indexes -(`.fai`, plus `.gzi` for BGZF). This keeps default aggregate memory bounded; -`--processes 4` opts into maximum plotting throughput. Ordinary gzip and -unindexed inputs remain single-pass and sequential so they are not scanned -once per worker. Use `--processes 1` for explicitly serial execution. +(`.fai`, plus `.gzi` for BGZF). Comparative workers group pairs by their y-axis +record and reuse that record's exact sketch. This keeps default aggregate +memory bounded; `--processes 4` opts into maximum throughput. Ordinary gzip and +unindexed inputs remain single-pass and sequential so they are not scanned once +per worker. Use `--processes 1` for explicitly serial execution. + +`--memory-limit ` + +Set an aggregate memory budget for comparative workers. ModDotPlot estimates +the largest pair-local sequence, sketch, matrix, and rendering footprint and +reduces the worker count when necessary. When omitted, available memory is used +when the operating system exposes it. + +`--sketch-cache ` + +Persist compact, prepared comparison sketches for reuse by later indexed runs. +Cache entries are keyed by the FASTA path, size, modification time, record and +region, strand mode, and every sketch parameter. Positional chromosome-wide +k-mer arrays are never stored in this cache. `--width ` @@ -369,9 +398,7 @@ $ moddotplot static -c config/config.json Running ModDotPlot in static mode -Retrieving k-mers from Chr1:14000001-18000000.... - -Progress: |████████████████████████████████████████| 100.0% Completed +Retrieving k-mers from Chr1:14000001-18000000.... Chr1:14000001-18000000 k-mers retrieved! @@ -385,9 +412,6 @@ Computing self identity matrix for Chr1:14000001-18000000... Plot Resolution r: 1000 -Progress: |████████████████████████████████████████| 100.0% Completed - - Saved self-identity matrix as a paired-end bed file to Arabadopsis/Chr1:14000001-18000000/Chr1:14000001-18000000.bedpe Triangle plots, full plots, and histogram for Arabadopsis/Chr1:14000001-18000000/Chr1:14000001-18000000 saved sucessfully. @@ -434,6 +458,26 @@ ModDotPlot can produce an a vs. b style dotplot for each pairwise combination of moddotplot static -f sequences/*_MATERNAL*.fa --compare-only ``` +For a diploid multi-record assembly, a pair manifest avoids comparing every +chromosome and unplaced contig against every other record: + +```text +# homologs.tsv +chr1_mat_hsa1 chr1_pat_hsa1 +chr2_mat_hsa3 chr2_pat_hsa3 +``` + +```bash +moddotplot static -f diploid.fa.gz --compare-only \ + --pairs homologs.tsv --processes 4 --memory-limit 32 \ + --sketch-cache .moddotplot-sketches +``` + +For BGZF-compressed FASTA, both `.fai` and `.gzi` indexes must accompany the +input. Indexed comparative mode fetches one record at a time, sketches it +directly, and releases pair-local data before continuing; it does not retain a +`uint64` positional hash for every base in the genome. + ![](images/chr13_MATERNAL:1-4000000_chr14_MATERNAL:1-4000000_COMPARE.png) --- @@ -495,9 +539,7 @@ $ moddotplot interactive -f sequences/Chr1_cen.fa Running ModDotPlot in interactive mode -Retrieving k-mers from Chr1:14000000-18000000.... - -Progress: |████████████████████████████████████████| 100.0% Completed +Retrieving k-mers from Chr1:14000000-18000000.... Chr1:14000000-18000000 k-mers retrieved! @@ -505,14 +547,8 @@ Building self-identity matrices for Chr1:14000000-18000000, using a minimum wind Layer 1 using window length 2000 -Progress: |████████████████████████████████████████| 100.0% Completed - - Layer 2 using window length 4000 -Progress: |████████████████████████████████████████| 100.0% Completed - - ModDotPlot interactive mode is successfully running on http://127.0.0.1:8050/ Dash is running on http://127.0.0.1:8050/ diff --git a/benchmarks/benchmark_containment_matrix.py b/benchmarks/benchmark_containment_matrix.py index 8e3df71..008b826 100644 --- a/benchmarks/benchmark_containment_matrix.py +++ b/benchmarks/benchmark_containment_matrix.py @@ -81,7 +81,6 @@ def main(argv=None): expanded, args.identity, args.kmer, - supress_progress=True, ) pair_seconds = time.perf_counter() - started diff --git a/src/moddotplot/__main__.py b/src/moddotplot/__main__.py index 80bf3d0..59c51bc 100644 --- a/src/moddotplot/__main__.py +++ b/src/moddotplot/__main__.py @@ -1,13 +1,29 @@ #!/usr/bin/env python3 +from contextlib import redirect_stderr, redirect_stdout +import os import sys -from moddotplot.estimate_identity import * -from moddotplot.moddotplot import main -from moddotplot.parse_fasta import * -import setproctitle -if __name__ == "__main__": - # Keeping execution behind the standard guard makes this module safe to - # import in spawn-based chromosome worker processes and as a console-script - # entry point. + +def _run(arguments): + import setproctitle + + from moddotplot.moddotplot import main as run_moddotplot + setproctitle.setproctitle("ModDotPlot") + return run_moddotplot(arguments) + + +def main(arguments=None): + """Load the CLI under quiet redirection to cover import-time output.""" + + raw_arguments = list(sys.argv[1:] if arguments is None else arguments) + if "--quiet" not in raw_arguments: + return _run(raw_arguments) + + with open(os.devnull, "w", encoding="utf-8") as sink: + with redirect_stdout(sink), redirect_stderr(sink): + return _run(raw_arguments) + + +if __name__ == "__main__": sys.exit(main()) diff --git a/src/moddotplot/estimate_identity.py b/src/moddotplot/estimate_identity.py index 3dbc470..004556e 100644 --- a/src/moddotplot/estimate_identity.py +++ b/src/moddotplot/estimate_identity.py @@ -8,14 +8,12 @@ DIVERGING_PALETTES, QUALITATIVE_PALETTES, ) -from palettable import colorbrewer from typing import Collection, Hashable, List, Set, Dict, Tuple import pandas as pd from scipy.sparse import csr_matrix from moddotplot import _nthash from moddotplot.optional_dependencies import OptionalDependencyError -from moddotplot.parse_fasta import printProgressBar @dataclass(frozen=True) @@ -781,8 +779,6 @@ def convertMatrixToCool( # ---- write cooler ---- cooler.create_cooler(output_cool, bins=bins, pixels=pixels, ordered=True) - print(bins) - print(pixels) return output_cool @@ -1028,7 +1024,6 @@ def selfContainmentMatrix( n = len(mod_set) if len(mod_set_neighbors) != n: raise IndexError("core and expanded self sketches must have equal lengths") - printProgressBar(0, n, prefix="Progress:", suffix="Complete", length=40) intersection_counts = _sketch_intersection_counts(mod_set, mod_set_neighbors) core_sizes = np.fromiter((len(sketch) for sketch in mod_set), dtype=float, count=n) directional_containment = np.zeros((n, n), dtype=float) @@ -1038,11 +1033,14 @@ def selfContainmentMatrix( out=directional_containment, where=core_sizes[:, np.newaxis] != 0, ) - symmetric_containment = np.maximum( - directional_containment, directional_containment.T + del intersection_counts + np.maximum( + directional_containment, + directional_containment.T, + out=directional_containment, ) containment_matrix = _identity_matrix_from_containment( - symmetric_containment, identity, k + directional_containment, identity, k ) diagonal = np.full(n, 100.0) @@ -1050,10 +1048,6 @@ def selfContainmentMatrix( diagonal[core_sizes == 0] = 0.0 np.fill_diagonal(containment_matrix, diagonal) - printProgressBar( - n, n, prefix="Progress:", suffix="Completed", length=40 - ) # show completed progress bar - print("\n") return containment_matrix @@ -1064,7 +1058,7 @@ def pairwiseContainmentMatrix( mod_set_y_neighbors: List[Set[int]], identity: int, k: int, - supress_progress: bool, + supress_progress: bool = False, ) -> np.ndarray: """ Calculate an updated identity matrix using specified parameters. @@ -1076,8 +1070,7 @@ def pairwiseContainmentMatrix( mod_set_y_neighbors (List[Set[int]]): Neighbor sets for y-axis windows. identity (int): Resolution parameter. k (int): Value for the k parameter in the binomial_distance function. - supress_progress (bool): if true supresses the progress bar - + supress_progress (bool): Retained for compatibility; has no effect. Returns: np.ndarray: A ``(len(mod_set_y), len(mod_set_x))`` identity matrix. """ @@ -1085,8 +1078,6 @@ def pairwiseContainmentMatrix( cols = len(mod_set_x) if len(mod_set_x_neighbors) != cols or len(mod_set_y_neighbors) != rows: raise IndexError("core and expanded pairwise sketches must have equal lengths") - if not supress_progress: - printProgressBar(0, rows, prefix="Progress:", suffix="Complete", length=40) x_core_sizes = np.fromiter( (len(sketch) for sketch in mod_set_x), dtype=float, count=cols ) @@ -1094,39 +1085,51 @@ def pairwiseContainmentMatrix( (len(sketch) for sketch in mod_set_y), dtype=float, count=rows ) - # core X against expanded Y, transposed into the public (Y, X) layout. - x_to_y_counts = _sketch_intersection_counts(mod_set_x, mod_set_y_neighbors).T - x_to_y = np.zeros((rows, cols), dtype=float) - np.divide( - x_to_y_counts, - x_core_sizes[np.newaxis, :], - out=x_to_y, - where=x_core_sizes[np.newaxis, :] != 0, - ) - if not supress_progress: - printProgressBar( - rows // 2, rows, prefix="Progress:", suffix="Complete", length=40 + # Keep only the final containment matrix at full size. Both directional + # count calculations and the second floating-point direction are bounded + # row blocks, avoiding several simultaneous dense matrix-sized arrays. + containment_matrix = np.zeros((rows, cols), dtype=float) + # A default 1000x1000 plot stays in one native call per direction. Larger + # or rectangular matrices are split so their count temporaries remain + # bounded without penalizing the common resolution. + max_block_cells = 1_048_576 + rows_per_block = max(1, max_block_cells // max(cols, 1)) + for row_start in range(0, rows, rows_per_block): + row_end = min(row_start + rows_per_block, rows) + block = containment_matrix[row_start:row_end] + + # core X against expanded Y, transposed into public (Y, X) layout. + x_to_y_counts = _sketch_intersection_counts( + mod_set_x, mod_set_y_neighbors[row_start:row_end] + ).T + np.divide( + x_to_y_counts, + x_core_sizes[np.newaxis, :], + out=block, + where=x_core_sizes[np.newaxis, :] != 0, ) + del x_to_y_counts + + # core Y against expanded X already has public (Y, X) orientation. + y_to_x_counts = _sketch_intersection_counts( + mod_set_y[row_start:row_end], mod_set_x_neighbors + ) + y_to_x = np.zeros(block.shape, dtype=float) + block_y_sizes = y_core_sizes[row_start:row_end, np.newaxis] + np.divide( + y_to_x_counts, + block_y_sizes, + out=y_to_x, + where=block_y_sizes != 0, + ) + del y_to_x_counts + np.maximum(block, y_to_x, out=block) + del y_to_x - # core Y against expanded X already has the public (Y, X) orientation. - y_to_x_counts = _sketch_intersection_counts(mod_set_y, mod_set_x_neighbors) - y_to_x = np.zeros((rows, cols), dtype=float) - np.divide( - y_to_x_counts, - y_core_sizes[:, np.newaxis], - out=y_to_x, - where=y_core_sizes[:, np.newaxis] != 0, - ) - symmetric_containment = np.maximum(x_to_y, y_to_x) containment_matrix = _identity_matrix_from_containment( - symmetric_containment, identity, k + containment_matrix, identity, k ) - if not supress_progress: - printProgressBar( - rows, rows, prefix="Progress:", suffix="Completed", length=40 - ) # show completed progress bar - print("\n") return containment_matrix @@ -1140,6 +1143,8 @@ def findElementsWithPrefix(lst, prefix): def getInteractiveColor(palette_name, palette_orientation): + from palettable import colorbrewer + palettes = colorbrewer.COLOR_MAPS tmp_color = [] new_palette = palette_name.split("_") diff --git a/src/moddotplot/moddotplot.py b/src/moddotplot/moddotplot.py index 683fded..f2c1572 100755 --- a/src/moddotplot/moddotplot.py +++ b/src/moddotplot/moddotplot.py @@ -6,6 +6,7 @@ iter_fasta_records, supports_indexed_fasta_access, getInputHeaders, + getInputSeqLength, isValidFasta, extractFiles, extractRegion, @@ -26,6 +27,7 @@ create_self_matrix_from_sketches, create_pairwise_matrix_from_sketches, prepare_sequence_sketches, + PreparedModimizerSketches, ModimizerSketchCache, partitionOverlaps, ) @@ -35,9 +37,11 @@ import argparse from concurrent.futures import ProcessPoolExecutor, as_completed +from contextlib import contextmanager, redirect_stderr, redirect_stdout from dataclasses import dataclass from itertools import islice import math +import hashlib import json import numpy as np import pickle @@ -64,6 +68,19 @@ ) +@contextmanager +def _silence_output(enabled): + """Discard Python-level stdout and stderr while a quiet command runs.""" + + if not enabled: + yield + return + + with open(os.devnull, "w", encoding="utf-8") as sink: + with redirect_stdout(sink), redirect_stderr(sink): + yield + + def _load_static_plotting(): global read_df_from_file, create_plots, create_grid @@ -100,6 +117,12 @@ def get_parser(): parser = argparse.ArgumentParser( formatter_class=argparse.ArgumentDefaultsHelpFormatter, description="ModDotPlot: Visualization of Tandem Repeats", + allow_abbrev=False, + ) + parser.add_argument( + "--quiet", + action="store_true", + help="Suppress all console output, including warnings and errors.", ) subparsers = parser.add_subparsers( dest="command", @@ -107,17 +130,26 @@ def get_parser(): help="Choose mode; static is used when omitted", ) static_parser = subparsers.add_parser( - "static", help="Static mode commands (default)" + "static", help="Static mode commands (default)", allow_abbrev=False ) interactive_parser = subparsers.add_parser( "interactive", help="Interactive mode commands (deprecated; explicit use only)", + allow_abbrev=False, description=( "Deprecated interactive mode. This mode remains available but is " "maintenance-only and will not receive new features." ), ) + for mode_parser in (static_parser, interactive_parser): + mode_parser.add_argument( + "--quiet", + action="store_true", + default=argparse.SUPPRESS, + help="Suppress all console output, including warnings and errors.", + ) + # -----------INTERACTIVE MODE SUBCOMMANDS----------- interactive_input_group = interactive_parser.add_mutually_exclusive_group( required=True @@ -292,6 +324,17 @@ def get_parser(): ), ) + static_parser.add_argument( + "--pairs", + default=None, + metavar="FILE", + help=( + "Compare only the sequence pairs listed in a two-column text file. " + "Blank lines and lines beginning with '#' are ignored. Requires " + "--compare or --compare-only." + ), + ) + # Add a mutually exclusive group for compare and compare only. static_compare_group = static_parser.add_mutually_exclusive_group(required=False) static_window_size_group = static_parser.add_mutually_exclusive_group( @@ -405,9 +448,31 @@ def get_parser(): choices=range(1, 5), metavar="N", help=( - "Number of independent chromosome workers (1-4). The default " - "automatically uses up to four workers for indexed multi-record " - "FASTA input." + "Number of independent chromosome or comparison-group workers " + "(1-4). The default automatically uses a bounded worker count for " + "indexed multi-record FASTA input." + ), + ) + + static_parser.add_argument( + "--memory-limit", + default=None, + type=float, + metavar="GIB", + help=( + "Maximum aggregate memory budget for comparison workers in GiB. " + "When omitted, available memory is detected when the platform " + "exposes it." + ), + ) + + static_parser.add_argument( + "--sketch-cache", + default=None, + metavar="DIRECTORY", + help=( + "Persist compact prepared sketches in this directory for reuse " + "across indexed comparative runs." ), ) @@ -533,7 +598,10 @@ def _arguments_with_default_command(arguments=None): normalized = list(sys.argv[1:] if arguments is None else arguments) if normalized[:1] and normalized[0] in ("-h", "--help"): return normalized - if not normalized or normalized[0] not in COMMANDS: + explicit_command = bool(normalized and normalized[0] in COMMANDS) + if not explicit_command and normalized[:1] == ["--quiet"] and len(normalized) > 1: + explicit_command = normalized[1] in COMMANDS + if not explicit_command: normalized.insert(0, "static") return normalized @@ -557,6 +625,7 @@ def _apply_static_config(args, config): args.resolution = config.get("resolution", args.resolution) args.window = config.get("window", args.window) args.sequence = config.get("sequence", args.sequence) + args.pairs = config.get("pairs", args.pairs) args.region = config.get("region", args.region) args.identity = config.get("identity", args.identity) args.delta = config.get("delta", args.delta) @@ -570,6 +639,8 @@ def _apply_static_config(args, config): args.no_plot = config.get("no_plot", args.no_plot) args.no_hist = config.get("no_hist", args.no_hist) args.processes = config.get("processes", args.processes) + args.memory_limit = config.get("memory_limit", args.memory_limit) + args.sketch_cache = config.get("sketch_cache", args.sketch_cache) args.width = config.get("width", args.width) args.axes_limits = config.get("axes_limits", args.axes_limits) args.dpi = config.get("dpi", args.dpi) @@ -798,6 +869,128 @@ class MatrixConfig: expectation: int +@dataclass(frozen=True) +class StaticSequenceRecord: + """Indexed FASTA metadata needed by the record-local static pipeline.""" + + source_path: str + name: str + length: int + region: tuple | None = None + + @property + def selected_length(self): + if self.region is None: + return self.length + return self.region[2] - self.region[1] + 1 + + +def _selected_sequence_records(fasta_headers, region_by_name): + """Return selected records with lengths obtained without reading sequences.""" + + records = [] + for source_path, selected_headers in fasta_headers.items(): + all_headers = getInputHeaders(source_path) + all_lengths = getInputSeqLength(source_path) + lengths = dict(zip(all_headers, all_lengths)) + for name in selected_headers: + records.append( + StaticSequenceRecord( + source_path=os.fspath(source_path), + name=name, + length=lengths[name], + region=region_by_name.get(name), + ) + ) + return records + + +def _resolve_pair_selector(selector, records): + """Resolve one pair-file identifier with CLI sequence matching semantics.""" + + exact = [record for record in records if record.name == selector] + matches = exact or [ + record for record in records if record.name.casefold() == selector.casefold() + ] + if not matches: + raise ValueError( + f"pair identifier {selector!r} does not match a selected FASTA record" + ) + if len(matches) > 1: + formatted = ", ".join( + f"{record.name!r} in {record.source_path!r}" for record in matches + ) + raise ValueError(f"pair identifier {selector!r} is ambiguous: {formatted}") + return matches[0] + + +def _read_pair_file(path, records): + """Read and validate an explicitly ordered two-column comparison plan.""" + + pairs = [] + seen = set() + with open(path, "rt", encoding="utf-8") as pair_file: + for line_number, raw_line in enumerate(pair_file, start=1): + line = raw_line.strip() + if not line or line.startswith("#"): + continue + fields = line.split() + if len(fields) != 2: + raise ValueError( + f"{path!s}:{line_number}: expected exactly two sequence identifiers" + ) + x_record = _resolve_pair_selector(fields[0], records) + y_record = _resolve_pair_selector(fields[1], records) + if x_record == y_record: + raise ValueError( + f"{path!s}:{line_number}: a sequence cannot be compared with itself" + ) + unordered_key = frozenset((x_record, y_record)) + if unordered_key in seen: + raise ValueError( + f"{path!s}:{line_number}: duplicate comparison for " + f"{x_record.name!r} and {y_record.name!r}" + ) + seen.add(unordered_key) + pairs.append((x_record, y_record)) + if not pairs: + raise ValueError(f"pair file {path!s} does not contain any comparisons") + return pairs + + +def _plan_static_pairs(args, records): + """Build oriented ``(x, y)`` comparisons in deterministic output order.""" + + if getattr(args, "pairs", None): + pairs = _read_pair_file(args.pairs, records) + if args.compare_order == "size": + pairs = [ + (y_record, x_record) + if x_record.selected_length < y_record.selected_length + else (x_record, y_record) + for x_record, y_record in pairs + ] + return pairs + + ordered = list(records) + if args.compare_order == "size": + ordered.sort(key=lambda record: record.selected_length, reverse=True) + return [ + (ordered[i], ordered[j]) + for i in range(len(ordered)) + for j in range(i + 1, len(ordered)) + ] + + +def _group_static_pairs(pairs): + """Group pairs by y-axis record so its exact sketch is prepared once.""" + + groups = {} + for x_record, y_record in pairs: + groups.setdefault(y_record, []).append(x_record) + return [(y_record, tuple(x_records)) for y_record, x_records in groups.items()] + + def _matrix_config_for_length(kmer_count, args): """Resolve matrix parameters without mutating the parsed CLI arguments. @@ -930,6 +1123,14 @@ def _validated_process_count(value): def _process_static_self_task(task): """Reopen and process one indexed FASTA record in a spawned worker.""" + args = task[0] + with _silence_output(getattr(args, "quiet", False)): + return _process_static_self_task_inner(task) + + +def _process_static_self_task_inner(task): + """Implement a worker task under its requested output policy.""" + args, fasta_path, sequence_id, selected_region, summary_command = task if not args.no_plot: _load_static_plotting() @@ -1191,6 +1392,431 @@ def _process_static_self_record( ) +def _fetch_static_record(record): + """Fetch one selected FASTA record or region through the indexed reader.""" + + regions = {record.name: record.region} if record.region is not None else None + records = iter_fasta_records( + record.source_path, + regions=regions, + record_ids=[record.name], + ) + try: + sequence_id, sequence, sequence_label = next(records) + except StopIteration as error: + raise ValueError( + f"sequence {record.name!r} was not found in {record.source_path!r}" + ) from error + if sequence_id != record.name: + raise ValueError( + f"indexed FASTA returned {sequence_id!r} while fetching {record.name!r}" + ) + if record.region is None: + start, end = 1, len(sequence) + else: + _name, start, end = record.region + return sequence, sequence_label, start, end + + +def _sketch_cache_path(record, config, args, canonical): + """Return a content-addressed persistent sketch-cache path.""" + + if not args.sketch_cache: + return None + source = os.path.abspath(record.source_path) + source_stat = os.stat(source) + cache_key = repr( + ( + "prepared-sketch-v1", + VERSION, + HASH_ALGORITHM, + source, + source_stat.st_size, + source_stat.st_mtime_ns, + record.name, + record.region, + config.window_size, + config.sparsity, + config.expectation, + args.delta, + args.kmer, + args.ambiguous, + bool(canonical), + ) + ).encode("utf-8") + digest = hashlib.sha256(cache_key).hexdigest() + return os.path.join(os.fspath(args.sketch_cache), digest + ".npz") + + +def _pack_sketch_arrays(sketches): + offsets = np.zeros(len(sketches) + 1, dtype=np.int64) + if sketches: + offsets[1:] = np.cumsum( + np.fromiter((len(sketch) for sketch in sketches), dtype=np.int64) + ) + values = np.concatenate(sketches).astype(np.uint64, copy=False) + else: + values = np.empty(0, dtype=np.uint64) + return values, offsets + + +def _unpack_sketch_arrays(values, offsets): + return [values[offsets[i] : offsets[i + 1]] for i in range(len(offsets) - 1)] + + +def _load_cached_sketch(path): + if path is None or not os.path.isfile(path): + return None + try: + with np.load(path, allow_pickle=False) as cached: + core_values = cached["core_values"] + core_offsets = cached["core_offsets"] + neighbor_values = cached["neighbor_values"] + neighbor_offsets = cached["neighbor_offsets"] + except (KeyError, OSError, ValueError): + return None + return PreparedModimizerSketches( + core=_unpack_sketch_arrays(core_values, core_offsets), + neighbors=_unpack_sketch_arrays(neighbor_values, neighbor_offsets), + ) + + +def _store_cached_sketch(path, prepared): + if path is None: + return + cache_directory = os.path.dirname(path) + os.makedirs(cache_directory, exist_ok=True) + core_values, core_offsets = _pack_sketch_arrays(prepared.core) + neighbor_values, neighbor_offsets = _pack_sketch_arrays(prepared.neighbors) + temporary_path = f"{path}.{os.getpid()}.tmp.npz" + np.savez( + temporary_path, + core_values=core_values, + core_offsets=core_offsets, + neighbor_values=neighbor_values, + neighbor_offsets=neighbor_offsets, + ) + os.replace(temporary_path, path) + + +def _prepare_static_record_sketches(record, config, args, direction_rendering): + """Fetch a record once and return its main and optional direction sketches.""" + + main_canonical = not args.forward + main_cache_path = _sketch_cache_path(record, config, args, main_canonical) + prepared = _load_cached_sketch(main_cache_path) + alternate_canonical = bool(args.forward) + alternate_cache_path = ( + _sketch_cache_path(record, config, args, alternate_canonical) + if direction_rendering + else None + ) + alternate = _load_cached_sketch(alternate_cache_path) + if prepared is not None and (not direction_rendering or alternate is not None): + if record.region is None: + sequence_label, start, end = record.name, 1, record.length + else: + _name, start, end = record.region + sequence_label = f"{record.name}:{start}-{end}" + return prepared, alternate, sequence_label, start, end + + sequence, sequence_label, start, end = _fetch_static_record(record) + if len(sequence) < args.kmer: + raise ValueError( + f"sequence {sequence_label!r} is shorter than k-mer size {args.kmer}" + ) + if prepared is None: + prepared = prepare_sequence_sketches( + sequence, + config.window_size, + config.sparsity, + args.delta, + args.kmer, + args.ambiguous, + config.expectation, + canonical=main_canonical, + ) + _store_cached_sketch(main_cache_path, prepared) + if direction_rendering and alternate is None: + alternate = prepare_sequence_sketches( + sequence, + config.window_size, + config.sparsity, + args.delta, + args.kmer, + args.ambiguous, + config.expectation, + canonical=alternate_canonical, + ) + _store_cached_sketch(alternate_cache_path, alternate) + del sequence + return prepared, alternate, sequence_label, start, end + + +def _emit_static_pair( + *, + args, + x_record, + y_record, + x_name, + y_name, + x_start, + x_end, + y_start, + y_end, + config, + pair_mat, + direction_pair_mat, + summary_writer, +): + """Write and render one pair without constructing Python BEDPE row lists.""" + + direction_rendering = direction_pair_mat is not None + canonical_matrix = ( + direction_pair_mat if direction_rendering and args.forward else pair_mat + ) + if np.all(canonical_matrix == 0): + print( + f"The pairwise identity matrix for {x_name} and {y_name} is empty. " + "Skipping.\n" + ) + return + + output_prefix = f"{x_name}_{y_name}" + output_directory = os.path.join(args.output_dir or ".", output_prefix) + if (not args.no_bedpe) or (not args.no_plot): + os.makedirs(output_directory, exist_ok=True) + + if args.cooler: + try: + os.makedirs(output_directory, exist_ok=True) + cooler_output = os.path.join(output_directory, output_prefix + ".cooler") + convertMatrixToCool( + matrix=pair_mat, + window_size=config.window_size, + id_threshold=args.identity, + x_name=x_name, + y_name=y_name, + self_identity=False, + x_offset=x_start, + y_offset=y_start, + chromsizes=max(x_record.selected_length - args.kmer + 1, 0), + output_cool=cooler_output, + ) + print(f"Saved comparative matrix as a cooler file to {cooler_output}\n") + except Exception as error: + print(f"Error creating pairwise cooler file: {error}") + + bed_frame = None + direction_frame = None + if not args.no_plot: + bed_frame = convertMatrixToBedDataFrame( + pair_mat, + config.window_size, + args.identity, + x_name, + y_name, + False, + x_start, + y_start, + x_end, + y_end, + ) + if direction_rendering: + if args.forward: + canonical_frame = convertMatrixToBedDataFrame( + direction_pair_mat, + config.window_size, + args.identity, + x_name, + y_name, + False, + x_start, + y_start, + x_end, + y_end, + ) + forward_matrix = pair_mat + else: + canonical_frame = bed_frame + forward_matrix = direction_pair_mat + direction_frame = _annotate_bed_direction_frame( + canonical_frame, + forward_matrix, + config.window_size, + x_start, + y_start, + ) + + if not args.no_bedpe: + bedfile_output = os.path.join( + output_directory, output_prefix + "_COMPARE.bedpe" + ) + chunks = ( + [bed_frame] + if bed_frame is not None + else iterMatrixToBedChunks( + pair_mat, + config.window_size, + args.identity, + x_name, + y_name, + False, + x_start, + y_start, + x_end, + y_end, + ) + ) + _write_matrix_bedpe(bedfile_output, chunks) + print( + "Saved comparative matrix as a paired-end bed file to " + f"{bedfile_output}\n" + ) + + if args.no_plot: + return + + pair_axis_bounds = args.axes_limits or ( + min(x_start, y_start), + max(x_end, y_end), + ) + plot_files = create_plots( + sdf=None, + directory=output_directory, + name_x=x_name, + name_y=y_name, + palette=args.palette, + palette_orientation=args.palette_orientation, + no_hist=args.no_hist, + width=args.width, + dpi=args.dpi, + is_freq=args.bin_freq, + xlim=pair_axis_bounds, + custom_colors=args.colors, + custom_breakpoints=args.breakpoints, + from_file=bed_frame, + is_pairwise=True, + axes_labels=args.axes_ticks, + axes_tick_number=args.axes_number, + vector_format=args.vector, + deraster=args.deraster, + annotation=args.bed, + ) + summary_writer.add( + output_directory, + plot_files or [], + fasta_files=[x_record.source_path, y_record.source_path], + window_sizes=[config.window_size], + regions=_regions_from_names([x_name, y_name]), + bed_file=args.bed, + ) + + if direction_rendering: + direction_directory = os.path.join(output_directory, "directionality") + direction_files = create_plots( + sdf=None, + directory=direction_directory, + name_x=x_name, + name_y=y_name, + palette=args.palette, + palette_orientation=args.palette_orientation, + no_hist=args.no_hist, + width=args.width, + dpi=args.dpi, + is_freq=args.bin_freq, + xlim=pair_axis_bounds, + custom_colors=args.colors, + custom_breakpoints=args.breakpoints, + from_file=direction_frame, + is_pairwise=True, + axes_labels=args.axes_ticks, + axes_tick_number=args.axes_number, + vector_format=args.vector, + deraster=args.deraster, + annotation=None, + ) + summary_writer.add( + direction_directory, + direction_files or [], + fasta_files=[x_record.source_path, y_record.source_path], + window_sizes=[config.window_size], + regions=_regions_from_names([x_name, y_name]), + bed_file=args.bed, + ) + + +def _process_static_pair_group(task): + """Process one y-axis record and every x-axis partner assigned to it.""" + + args = task[0] + with _silence_output(getattr(args, "quiet", False)): + return _process_static_pair_group_inner(task) + + +def _process_static_pair_group_inner(task): + args, y_record, x_records, summary_command = task + if not args.no_plot: + _load_static_plotting() + + y_kmer_count = y_record.selected_length - args.kmer + 1 + config = _matrix_config_for_length(y_kmer_count, args) + direction_rendering = args.plot_direction and not args.no_plot + prepared_y, alternate_y, y_name, y_start, y_end = _prepare_static_record_sketches( + y_record, config, args, direction_rendering + ) + summary_writer = PlotSummaryWriter(summary_command) + + for x_record in x_records: + print( + f"Computing pairwise identity matrix for {x_record.name} and {y_name}... \n" + ) + ( + prepared_x, + alternate_x, + x_name, + x_start, + x_end, + ) = _prepare_static_record_sketches(x_record, config, args, direction_rendering) + print(f"\tSequence length {x_name}: {x_record.selected_length}\n") + print(f"\tSequence length {y_name}: {y_record.selected_length}\n") + print(f"\tWindow size w: {config.window_size}\n") + print(f"\tModimizer sketch size: {config.expectation}\n") + print(f"\tPlot Resolution r: {config.resolution}\n") + + # Preserve the established matrix orientation: the y-axis record is + # the first sketch argument and the x-axis record is the second. + pair_mat = create_pairwise_matrix_from_sketches( + prepared_y, prepared_x, args.identity, args.kmer + ) + direction_pair_mat = None + if direction_rendering: + direction_pair_mat = create_pairwise_matrix_from_sketches( + alternate_y, alternate_x, args.identity, args.kmer + ) + del prepared_x, alternate_x + + _emit_static_pair( + args=args, + x_record=x_record, + y_record=y_record, + x_name=x_name, + y_name=y_name, + x_start=x_start, + x_end=x_end, + y_start=y_start, + y_end=y_end, + config=config, + pair_mat=pair_mat, + direction_pair_mat=direction_pair_mat, + summary_writer=summary_writer, + ) + del pair_mat, direction_pair_mat + + del prepared_y, alternate_y + return y_record.name + + def _run_streaming_static_self( args, fasta_list, fasta_headers, region_by_name, summary_writer ): @@ -1268,10 +1894,155 @@ def _run_streaming_static_self( ) -def main(): +def _available_memory_bytes(): + """Return currently available physical memory when the OS exposes it.""" + + try: + page_size = int(os.sysconf("SC_PAGE_SIZE")) + available_pages = int(os.sysconf("SC_AVPHYS_PAGES")) + except (AttributeError, OSError, TypeError, ValueError): + return None + if page_size <= 0 or available_pages <= 0: + return None + return page_size * available_pages + + +def _estimate_pair_group_peak_bytes(args, group): + """Conservatively estimate one comparison group's largest pair footprint.""" + + y_record, x_records = group + y_kmers = max(y_record.selected_length - args.kmer + 1, 1) + config = _matrix_config_for_length(y_kmers, args) + largest_x = max(record.selected_length for record in x_records) + x_kmers = max(largest_x - args.kmer + 1, 1) + x_windows = math.ceil(x_kmers / config.window_size) + y_windows = math.ceil(y_kmers / config.window_size) + matrix_cells = x_windows * y_windows + + # The native sketcher may briefly encode the selected sequence while the + # Python string is live. Prepared core and expanded sketches average at + # most ``expectation`` uint64 hashes per window. Rendering additionally + # retains compact BEDPE columns and figure buffers. + sequence_bytes = largest_x * 3 + sketch_bytes = (x_windows + y_windows) * config.expectation * 16 + matrix_bytes = matrix_cells * (24 if args.no_plot else 72) + rendering_bytes = 0 if args.no_plot else 256 * 1024**2 + return max(1, sequence_bytes + sketch_bytes + matrix_bytes + rendering_bytes) + + +def _comparison_process_count(args, groups): + """Choose pair workers using CPU, task count, and aggregate memory budget.""" + + process_count = _streaming_process_count(args, len(groups), indexed_access=True) + if process_count <= 1: + return process_count + + if args.memory_limit is not None: + budget = int(float(args.memory_limit) * 1024**3) + else: + available = _available_memory_bytes() + budget = int(available * 0.75) if available is not None else None + if budget is None: + return process_count + + largest_peak = max(_estimate_pair_group_peak_bytes(args, group) for group in groups) + memory_workers = max(1, budget // largest_peak) + return min(process_count, memory_workers) + + +def _run_streaming_static_compare( + args, + fasta_list, + fasta_headers, + region_by_name, + summary_writer, +): + """Run indexed pairwise plots with record-local sequence memory.""" + + records = _selected_sequence_records(fasta_headers, region_by_name) + pairs = _plan_static_pairs(args, records) + if not pairs: + if args.compare_only: + raise ValueError( + "can't create a comparative plot with fewer than two sequences" + ) + _run_streaming_static_self( + args, + fasta_list, + fasta_headers, + region_by_name, + summary_writer, + ) + return + groups = _group_static_pairs(pairs) + print( + f"Planning {len(pairs)} pairwise comparison" + f"{'s' if len(pairs) != 1 else ''} across {len(groups)} sketch group" + f"{'s' if len(groups) != 1 else ''}.\n" + ) + + if not args.compare_only: + _run_streaming_static_self( + args, + fasta_list, + fasta_headers, + region_by_name, + summary_writer, + ) + + tasks = [ + (args, y_record, x_records, summary_writer.command) + for y_record, x_records in groups + ] + process_count = _comparison_process_count(args, groups) + if process_count > 1: + print( + f"Processing {len(pairs)} comparisons with {process_count} " + "bounded-memory workers.\n" + ) + context = multiprocessing.get_context("spawn") + try: + with ProcessPoolExecutor( + max_workers=process_count, + mp_context=context, + ) as executor: + future_records = { + executor.submit(_process_static_pair_group, task): task[1].name + for task in tasks + } + for future in as_completed(future_records): + sequence_id = future_records[future] + try: + future.result() + except Exception as error: + for pending in future_records: + pending.cancel() + raise ValueError( + "failed while processing comparisons grouped by " + f"{sequence_id!r}: {error}" + ) from error + except ValueError: + raise + except (OSError, RuntimeError) as error: + raise ValueError(f"unable to run comparison workers: {error}") from error + return + + for task in tasks: + _process_static_pair_group(task) + + +def main(arguments=None): + """Run ModDotPlot, suppressing every output stream when requested.""" + + raw_arguments = list(sys.argv[1:] if arguments is None else arguments) + with _silence_output("--quiet" in raw_arguments): + return _main(raw_arguments) + + +def _main(arguments): print(ASCII_ART) print(f"v{VERSION} \n") - args = parse_args() + args = parse_args(arguments) # Matrix-only interactive exports do not use Dash or Plotly. Every path # that opens the interactive UI validates its extra before doing expensive @@ -1360,6 +2131,9 @@ def main(): except ValueError as error: print(f"Error: {error}.", file=sys.stderr) sys.exit(2) + if args.memory_limit is not None and args.memory_limit <= 0: + print("Error: --memory-limit must be greater than zero.", file=sys.stderr) + sys.exit(2) # Plotting imports are comparatively expensive. Compute-only static # runs should not import Plotnine or initialize Matplotlib at all. @@ -1378,6 +2152,20 @@ def main(): ) sys.exit(2) + if args.pairs and getattr(args, "load", None): + print("Error: --pairs requires FASTA input; it cannot be used with --load.") + sys.exit(2) + + if args.pairs and not (args.compare or args.compare_only): + print("Error: --pairs requires --compare or --compare-only.") + sys.exit(2) + + if args.pairs and (args.grid or args.grid_only): + print( + "Error: --pairs currently targets individual comparative plots, not grids." + ) + sys.exit(2) + if args.plot_direction and getattr(args, "load", None): print( "Error: --plot-direction requires FASTA input because strand " @@ -1602,6 +2390,35 @@ def main(): print(f"Error: {error}.\n") sys.exit(2) + streaming_static_compare = ( + args.command == "static" + and bool(fasta_list) + and (args.compare or args.compare_only) + and not args.grid + and not args.grid_only + and all(supports_indexed_fasta_access(path) for path in fasta_list) + ) + if streaming_static_compare: + try: + _run_streaming_static_compare( + args, + fasta_list, + fasta_headers, + region_by_name, + summary_writer, + ) + except (OSError, UnicodeError, ValueError) as error: + print(f"Error processing comparative FASTA input: {error}", file=sys.stderr) + sys.exit(2) + return + if getattr(args, "pairs", None): + print( + "Error: --pairs requires random-access FASTA input (.fai, plus " + ".gzi for BGZF-compressed files).", + file=sys.stderr, + ) + sys.exit(2) + # Independent static self plots have no cross-record dependency. Stream # them directly instead of retaining every positional hash in ``k_list``. # The file-existence condition preserves unit tests and third-party callers @@ -1640,7 +2457,7 @@ def main(): readKmersFromFile( i, args.kmer, - False, + args.quiet, True, args.ambiguous, region_by_name, @@ -1652,7 +2469,7 @@ def main(): readKmersFromFile( i, args.kmer, - False, + args.quiet, False, args.ambiguous, region_by_name, @@ -1670,7 +2487,7 @@ def main(): readKmersFromFile( path, args.kmer, - False, + args.quiet, not args.forward, args.ambiguous, region_by_name, @@ -1896,7 +2713,6 @@ def main(): smaller_mods_neigh, args.identity, args.kmer, - False, ) image_pyramid.insert(0, matrix_layer) matrices.append(image_pyramid) diff --git a/src/moddotplot/parse_fasta.py b/src/moddotplot/parse_fasta.py index 7418718..356dd1e 100644 --- a/src/moddotplot/parse_fasta.py +++ b/src/moddotplot/parse_fasta.py @@ -660,36 +660,13 @@ def generateKmersFromFasta( The public iterator remains compatible with existing callers. Ambiguous windows are yielded as ``None`` unless ``ambiguous`` is enabled, while the FASTA reader below keeps the compact bulk NumPy representation in memory. + ``quiet`` is retained for call compatibility; hashing no longer emits a + progress display in either mode. """ - total_kmers = max(len(seq) - k + 1, 0) - if not quiet: - printProgressBar( - 0, total_kmers, prefix="Progress:", suffix="Complete", length=40 - ) - hashes = _hash_sequence(seq, k, fw_only, ambiguous) - progress_threshold = max(round(total_kmers / 77), 1) - for index, kmer_hash in enumerate(hashes): - if not quiet and index % progress_threshold == 0: - printProgressBar( - index, - total_kmers, - prefix="Progress:", - suffix="Complete", - length=40, - ) - + for kmer_hash in hashes: yield None if np.ma.is_masked(kmer_hash) else int(kmer_hash) - if not quiet and index == total_kmers - 1: - printProgressBar( - total_kmers, - total_kmers, - prefix="Progress:", - suffix="Completed", - length=40, - ) - def isValidFasta(file_path): try: @@ -759,17 +736,9 @@ def printProgressBar( fill="█", printEnd="\r", ): - if total <= 0: - percent = f"{100:.{decimals}f}" - filledLength = length - else: - percent = f"{100 * (iteration / total):.{decimals}f}" - filledLength = int(length * iteration // total) - bar = [fill] * filledLength + ["-"] * (length - filledLength) - bar_str = "".join(bar) - print(f"\r{prefix} |{bar_str}| {percent}% {suffix}", end=printEnd) - if iteration == total: - print() + """Compatibility no-op retained after removal of progress displays.""" + + return None def readKmersFromFile( @@ -788,7 +757,8 @@ def readKmersFromFile( for seq_id, sequence, sequence_label in iter_fasta_records( filename, regions=regions, record_ids=record_ids ): - print(f"Retrieving k-mers from {sequence_label}.... \n") + if not quiet: + print(f"Retrieving k-mers from {sequence_label}.... \n") if len(sequence) < ksize: if regions and seq_id in regions: _name, start, end = regions[seq_id] @@ -797,22 +767,10 @@ def readKmersFromFile( f"{ksize}" ) - total_kmers = max(len(sequence) - ksize + 1, 0) - if not quiet: - printProgressBar( - 0, total_kmers, prefix="Progress:", suffix="Complete", length=40 - ) kmers_for_seq = _hash_sequence(sequence, ksize, fw_only, ambiguous) - if not quiet: - printProgressBar( - total_kmers, - total_kmers, - prefix="Progress:", - suffix="Completed", - length=40, - ) all_kmers.append(kmers_for_seq) - print(f"\n{sequence_label} k-mers retrieved! \n") + if not quiet: + print(f"\n{sequence_label} k-mers retrieved! \n") return all_kmers diff --git a/tests/test_algorithms.py b/tests/test_algorithms.py index 7d25283..a404720 100644 --- a/tests/test_algorithms.py +++ b/tests/test_algorithms.py @@ -406,7 +406,6 @@ def test_pairwise_containment_matrix_is_rectangular_and_keeps_axis_orientation() mod_set_y_neighbors=[{1}, {3}], identity=0, k=1, - supress_progress=True, ) np.testing.assert_array_equal( @@ -429,7 +428,6 @@ def test_pairwise_containment_matrix_supports_more_rows_than_columns(): mod_set_y_neighbors=[{1}, {2}, {3}], identity=0, k=1, - supress_progress=True, ) np.testing.assert_array_equal(matrix, np.array([[0.0], [100.0], [0.0]])) @@ -454,7 +452,6 @@ def test_pairwise_containment_matrix_preserves_empty_axis_dimensions( mod_set_y_neighbors=list(mod_set_y), identity=0, k=1, - supress_progress=True, ) assert matrix.shape == expected_shape @@ -469,28 +466,46 @@ def test_pairwise_containment_matrix_does_not_hide_misaligned_neighbor_data(): mod_set_y_neighbors=[{1}], identity=0, k=1, - supress_progress=True, ) +@pytest.mark.parametrize("legacy_progress_setting", [False, True]) +def test_pairwise_legacy_progress_argument_is_silent(legacy_progress_setting, capsys): + pairwiseContainmentMatrix( + mod_set_x=[{1}], + mod_set_y=[{1}], + mod_set_x_neighbors=[{1}], + mod_set_y_neighbors=[{1}], + identity=0, + k=1, + supress_progress=legacy_progress_setting, + ) + + captured = capsys.readouterr() + assert captured.out == "" + assert captured.err == "" + + @pytest.mark.parametrize("sequence", ["", "A", "AC"]) -def test_generate_kmers_shorter_than_k_with_progress_returns_empty(sequence, capsys): +def test_generate_kmers_shorter_than_k_returns_empty_without_progress(sequence, capsys): assert list(generateKmersFromFasta(sequence, 3, quiet=False, fw_only=True)) == [] - assert "100.0%" in capsys.readouterr().out + captured = capsys.readouterr() + assert captured.out == "" + assert captured.err == "" -def test_generate_one_kmer_with_progress_does_not_use_zero_modulus(capsys): +def test_generate_one_kmer_does_not_emit_progress(capsys): result = list(generateKmersFromFasta("ACG", 3, quiet=False, fw_only=True)) assert result == [np.uint64(0xB13A5310100F646E)] - output = capsys.readouterr().out - assert "100.0%" in output - assert "Completed" in output + captured = capsys.readouterr() + assert captured.out == "" + assert captured.err == "" -def test_print_progress_bar_accepts_zero_total(capsys): - printProgressBar(0, 0, prefix="Progress:", suffix="Completed", length=4) +def test_legacy_progress_helper_is_a_silent_noop(capsys): + assert printProgressBar(1, 1, prefix="Progress:", suffix="Completed") is None - output = capsys.readouterr().out - assert "|████|" in output - assert "100.0%" in output + captured = capsys.readouterr() + assert captured.out == "" + assert captured.err == "" diff --git a/tests/test_cli_integration.py b/tests/test_cli_integration.py index 90266b4..93f7129 100644 --- a/tests/test_cli_integration.py +++ b/tests/test_cli_integration.py @@ -94,6 +94,109 @@ def test_indexed_self_plots_are_identical_with_spawned_chromosome_workers(tmp_pa assert parallel_files == serial_files +def test_indexed_pairwise_plots_are_identical_with_bounded_workers(tmp_path): + fasta = tmp_path / "indexed.fa" + serial_output = tmp_path / "pair-serial" + parallel_output = tmp_path / "pair-parallel" + _write_indexed_multifasta(fasta) + common = ( + "static", + "--fasta", + fasta, + "--window", + 100, + "--modimizer", + 10, + "--identity", + 80, + "--compare-only", + "--no-plot", + ) + + serial = _run_cli(*common, "--processes", 1, "--output-dir", serial_output) + parallel = _run_cli(*common, "--processes", 2, "--output-dir", parallel_output) + + assert serial.returncode == 0, serial.stderr + serial.stdout + assert parallel.returncode == 0, parallel.stderr + parallel.stdout + assert "3 pairwise comparisons across 2 sketch groups" in serial.stdout + assert "with 2 bounded-memory workers" in parallel.stdout + serial_files = { + path.relative_to(serial_output): path.read_bytes() + for path in serial_output.rglob("*.bedpe") + } + parallel_files = { + path.relative_to(parallel_output): path.read_bytes() + for path in parallel_output.rglob("*.bedpe") + } + assert parallel_files == serial_files + + +def test_pair_manifest_limits_indexed_comparisons(tmp_path): + fasta = tmp_path / "indexed.fa" + output = tmp_path / "selected-pairs" + pair_file = tmp_path / "pairs.tsv" + _write_indexed_multifasta(fasta) + pair_file.write_text("# selected comparisons\nalpha beta\nbeta gamma\n") + + result = _run_cli( + "static", + "--fasta", + fasta, + "--compare-only", + "--pairs", + pair_file, + "--window", + 100, + "--modimizer", + 10, + "--identity", + 80, + "--no-plot", + "--processes", + 1, + "--output-dir", + output, + ) + + assert result.returncode == 0, result.stderr + result.stdout + assert sorted(path.relative_to(output) for path in output.rglob("*.bedpe")) == [ + Path("alpha_beta/alpha_beta_COMPARE.bedpe"), + Path("beta_gamma/beta_gamma_COMPARE.bedpe"), + ] + + +def test_quiet_static_cli_suppresses_parent_and_worker_output(tmp_path): + fasta = tmp_path / "indexed.fa" + output = tmp_path / "quiet" + _write_indexed_multifasta(fasta) + + result = _run_cli( + "--fasta", + fasta, + "--window", + 100, + "--modimizer", + 10, + "--identity", + 80, + "--no-plot", + "--processes", + 2, + "--quiet", + "--output-dir", + output, + ) + + assert result.returncode == 0 + assert result.stdout == "" + assert result.stderr == "" + assert sorted(path.name for path in output.rglob("*.bedpe")) == [ + "alpha.bedpe", + "beta.bedpe", + "gamma.bedpe", + ] + + def test_static_cli_computes_all_self_and_pairwise_outputs(tmp_path): fasta = tmp_path / "three.fa" output = tmp_path / "static" @@ -352,11 +455,14 @@ def test_static_cli_reads_gzip_and_renders_bed_annotations(tmp_path): "--identity", 80, "--no-hist", + "--quiet", "--output-dir", output, ) assert result.returncode == 0, result.stderr + result.stdout + assert result.stdout == "" + assert result.stderr == "" sequence_output = output / "alpha" expected = [ sequence_output / "alpha_ANNOTATION_TRACK.svg", @@ -400,3 +506,35 @@ def test_interactive_cli_forward_mode_saves_matrix_without_launching_server(tmp_ assert (saved / "metadata.pkl").is_file() assert "Saved matrices" in result.stdout assert "interactive mode is deprecated and maintenance-only" in result.stderr + + +def test_quiet_interactive_matrix_export_has_no_console_output(tmp_path): + fasta = tmp_path / "one.fa" + fasta.write_text(">alpha\n" + "ACGT" * 300 + "\n") + output = tmp_path / "interactive-quiet" + + result = _run_cli( + "--quiet", + "interactive", + "--fasta", + fasta, + "--window", + 100, + "--resolution", + 10, + "--modimizer", + 10, + "--quick", + "--forward", + "--save", + "--no-plot", + "--output-dir", + output, + ) + + assert result.returncode == 0 + assert result.stdout == "" + assert result.stderr == "" + saved = output / "interactive_matrices" + assert (saved / "alpha_0.npz").is_file() + assert (saved / "metadata.pkl").is_file() diff --git a/tests/test_cli_runtime.py b/tests/test_cli_runtime.py index 0ee3de1..9c5d8d9 100644 --- a/tests/test_cli_runtime.py +++ b/tests/test_cli_runtime.py @@ -42,6 +42,60 @@ def test_matrix_config_preserves_default_nuclear_chromosome_parameters(): assert config.expectation == 1941 +def test_pair_plan_groups_by_y_axis_for_sketch_reuse(tmp_path): + records = [ + cli.StaticSequenceRecord("input.fa", "alpha", 300), + cli.StaticSequenceRecord("input.fa", "beta", 200), + cli.StaticSequenceRecord("input.fa", "gamma", 100), + ] + args = SimpleNamespace(pairs=None, compare_order="sequential") + + pairs = cli._plan_static_pairs(args, records) + groups = cli._group_static_pairs(pairs) + + assert [(x.name, y.name) for x, y in pairs] == [ + ("alpha", "beta"), + ("alpha", "gamma"), + ("beta", "gamma"), + ] + assert [(y.name, [x.name for x in xs]) for y, xs in groups] == [ + ("beta", ["alpha"]), + ("gamma", ["alpha", "beta"]), + ] + + +def test_persistent_sketch_cache_avoids_refetch_and_rehash(monkeypatch, tmp_path): + record = cli.StaticSequenceRecord(str(tmp_path / "input.fa"), "alpha", 400) + (tmp_path / "input.fa").write_text(">alpha\n" + "ACGT" * 100 + "\n") + args = SimpleNamespace( + sketch_cache=str(tmp_path / "cache"), + forward=False, + delta=0.5, + kmer=21, + ambiguous=False, + ) + config = cli.MatrixConfig(100, 4, 10, 8, 12) + calls = [] + + def fetch(_record): + calls.append("fetch") + return "ACGT" * 100, "alpha", 1, 400 + + def prepare(*_args, **_kwargs): + calls.append("prepare") + values = np.asarray([1, 2, 3], dtype=np.uint64) + return cli.PreparedModimizerSketches([values], [values]) + + monkeypatch.setattr(cli, "_fetch_static_record", fetch) + monkeypatch.setattr(cli, "prepare_sequence_sketches", prepare) + + first = cli._prepare_static_record_sketches(record, config, args, False)[0] + second = cli._prepare_static_record_sketches(record, config, args, False)[0] + + assert calls == ["fetch", "prepare"] + np.testing.assert_array_equal(first.core[0], second.core[0]) + + def test_streaming_self_runner_finishes_one_record_before_requesting_next( monkeypatch, ): @@ -247,6 +301,44 @@ def test_static_parser_accepts_bounded_process_request(): assert args.processes == 3 +def test_quiet_is_available_in_static_and_interactive_modes(): + static_args = cli.parse_args(["--fasta", "sequence.fa", "--quiet"]) + leading_static_args = cli.parse_args( + ["--quiet", "static", "--fasta", "sequence.fa"] + ) + interactive_args = cli.parse_args( + ["interactive", "--fasta", "sequence.fa", "--quiet"] + ) + leading_interactive_args = cli.parse_args( + ["--quiet", "interactive", "--fasta", "sequence.fa"] + ) + + assert static_args.quiet is True + assert leading_static_args.quiet is True + assert leading_static_args.command == "static" + assert interactive_args.quiet is True + assert leading_interactive_args.quiet is True + assert leading_interactive_args.command == "interactive" + + +def test_quiet_long_option_cannot_be_abbreviated(): + with pytest.raises(SystemExit) as exc_info: + cli.parse_args(["--fasta", "sequence.fa", "--qui"]) + + assert exc_info.value.code == 2 + + +def test_quiet_output_redirection_is_restored(capsys): + with cli._silence_output(True): + print("hidden") + print("also hidden", file=sys.stderr) + + print("visible") + captured = capsys.readouterr() + assert captured.out == "visible\n" + assert captured.err == "" + + @pytest.mark.parametrize("value", [0, 5, True, "many"]) def test_process_count_validation_rejects_invalid_config_values(value): with pytest.raises(ValueError, match="integer from 1 through 4"): diff --git a/tests/test_entrypoints.py b/tests/test_entrypoints.py index c601c62..706caa8 100644 --- a/tests/test_entrypoints.py +++ b/tests/test_entrypoints.py @@ -63,3 +63,26 @@ def test_module_without_arguments_defaults_to_static_parser(): result.stderr ) assert "the following arguments are required: command" not in result.stderr + + +def test_quiet_suppresses_parser_errors(): + result = _run_module("--quiet") + + assert result.returncode == 2 + assert result.stdout == "" + assert result.stderr == "" + + +def test_quiet_suppresses_explicit_help_output(): + result = _run_module("--quiet", "--help") + + assert result.returncode == 0 + assert result.stdout == "" + assert result.stderr == "" + + +def test_abbreviated_quiet_option_is_rejected_without_silencing_the_error(): + result = _run_module("--fasta", "sequence.fa", "--qui") + + assert result.returncode == 2 + assert "unrecognized arguments: --qui" in result.stderr diff --git a/tests/test_sketch_cache.py b/tests/test_sketch_cache.py index 4b8f9a4..cc846df 100644 --- a/tests/test_sketch_cache.py +++ b/tests/test_sketch_cache.py @@ -63,7 +63,7 @@ def test_prepared_sketch_matrix_results_match_compatibility_apis(): parameters["expectation"], ) actual_pair = create_pairwise_matrix_from_sketches( - prepared_first, prepared_second, 0, parameters["k"], True + prepared_first, prepared_second, 0, parameters["k"] ) np.testing.assert_array_equal(actual_pair, expected_pair) diff --git a/tests/test_sparse_containment.py b/tests/test_sparse_containment.py index 51b97ea..81fb596 100644 --- a/tests/test_sparse_containment.py +++ b/tests/test_sparse_containment.py @@ -69,7 +69,6 @@ def test_sparse_pairwise_matches_scalar_reference(identity, k): expanded_y, identity, k, - supress_progress=True, ) np.testing.assert_allclose(actual, expected) @@ -228,9 +227,7 @@ def fail_if_called(*_args, **_kwargs): core = [{index, index + 1} for index in range(250)] expanded = [sketch | {index + 2} for index, sketch in enumerate(core)] - matrix = pairwiseContainmentMatrix( - core, core, expanded, expanded, 0, 21, supress_progress=True - ) + matrix = pairwiseContainmentMatrix(core, core, expanded, expanded, 0, 21) assert matrix.shape == (250, 250) From 2a9ed40b4bff48125246f738e8dc25e93622c717 Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Fri, 2 Oct 2026 16:58:53 -0400 Subject: [PATCH 09/16] Allow config with loaded BEDPE input --- src/moddotplot/moddotplot.py | 63 ++++++++++++++++++++++----- src/moddotplot/static_plots.py | 40 ++++++++++++++++++ tests/test_cli_runtime.py | 68 ++++++++++++++++++++++++++++++ tests/test_static_customization.py | 30 +++++++++++++ 4 files changed, 191 insertions(+), 10 deletions(-) diff --git a/src/moddotplot/moddotplot.py b/src/moddotplot/moddotplot.py index f2c1572..31220b9 100755 --- a/src/moddotplot/moddotplot.py +++ b/src/moddotplot/moddotplot.py @@ -288,15 +288,18 @@ def get_parser(): ) # -----------STATIC MODE SUBCOMMANDS----------- - static_input_group = static_parser.add_mutually_exclusive_group(required=True) - static_input_group.add_argument( + static_parser.add_argument( "-c", "--config", default=None, type=str, - help="Config file to use. Takes precedence over any other competing command line arguments.", + help=( + "Config file to use. Explicit input and output paths take precedence; " + "config values take precedence for other competing arguments." + ), ) + static_input_group = static_parser.add_mutually_exclusive_group(required=False) static_input_group.add_argument( "-l", "--load", @@ -609,14 +612,38 @@ def _arguments_with_default_command(arguments=None): def parse_args(arguments=None): """Parse command-line arguments, defaulting omitted subcommands to static.""" - return get_parser().parse_args(_arguments_with_default_command(arguments)) + parser = get_parser() + args = parser.parse_args(_arguments_with_default_command(arguments)) + if args.command == "static" and not any( + ( + args.config, + getattr(args, "load", None), + getattr(args, "fasta", None), + ) + ): + subparsers_action = next( + action + for action in parser._actions + if isinstance(action, argparse._SubParsersAction) + ) + subparsers_action.choices["static"].error( + "one of the arguments -c/--config -l/--load -f/--fasta is required" + ) + return args def _apply_static_config(args, config): """Apply static-mode JSON configuration values to parsed arguments.""" # TODO: Remove args that are interactive only - args.fasta = config.get("fasta") - args.load = config.get("load") + cli_fasta = getattr(args, "fasta", None) + cli_load = getattr(args, "load", None) + cli_output_dir = args.output_dir + if cli_fasta is not None or cli_load is not None: + args.fasta = cli_fasta + args.load = cli_load + else: + args.fasta = config.get("fasta") + args.load = config.get("load") args.bed = config.get("bed") # Distance matrix commands @@ -624,12 +651,21 @@ def _apply_static_config(args, config): args.modimizer = config.get("modimizer", args.modimizer) args.resolution = config.get("resolution", args.resolution) args.window = config.get("window", args.window) - args.sequence = config.get("sequence", args.sequence) - args.pairs = config.get("pairs", args.pairs) - args.region = config.get("region", args.region) + if args.load: + args.sequence = None + args.pairs = None + args.region = None + else: + args.sequence = config.get("sequence", args.sequence) + args.pairs = config.get("pairs", args.pairs) + args.region = config.get("region", args.region) args.identity = config.get("identity", args.identity) args.delta = config.get("delta", args.delta) - args.output_dir = config.get("output_dir", args.output_dir) + args.output_dir = ( + cli_output_dir + if cli_output_dir is not None + else config.get("output_dir", args.output_dir) + ) args.compare = config.get("compare", args.compare) args.compare_only = config.get("compare_only", args.compare_only) args.compare_order = config.get("compare_order", args.compare_order) @@ -799,6 +835,12 @@ def _bedpe_window_sizes(dataframe): return [] +def _filter_loaded_identity(dataframe, identity): + """Apply the requested identity threshold to loaded BEDPE rows.""" + + return dataframe.loc[dataframe["perID_by_events"] >= identity].copy() + + def _regions_from_names(names): """Return normalized region strings embedded in sequence names.""" @@ -2200,6 +2242,7 @@ def _main(arguments): for bed in args.load: # If args.load is provided as input, run static mode directly from the paired-end bed file. Skip counting input k-mers. df = read_df_from_file(bed) + df = _filter_loaded_identity(df, args.identity) unique_query_names = df["#query_name"].unique() unique_reference_names = df["reference_name"].unique() diff --git a/src/moddotplot/static_plots.py b/src/moddotplot/static_plots.py index 86f4149..1684bc1 100755 --- a/src/moddotplot/static_plots.py +++ b/src/moddotplot/static_plots.py @@ -441,6 +441,46 @@ def get_colors(sdf, ncolors, is_freq, custom_breakpoints): # TODO: Remove pandas dependency def read_df_from_file(file_path): + browser_columns = None + with open(file_path, "r", encoding="utf-8") as bedpe: + for line in bedpe: + if line.startswith("#chrom1\t"): + browser_columns = line[1:].rstrip("\r\n").split("\t") + break + if not line.startswith("#"): + break + + if browser_columns is not None: + data = pd.read_csv( + file_path, + delimiter="\t", + comment="#", + names=browser_columns, + usecols=( + "chrom1", + "start1", + "end1", + "chrom2", + "start2", + "end2", + "ani_c", + ), + ) + data.rename( + columns={ + "chrom1": "#query_name", + "start1": "query_start", + "end1": "query_end", + "chrom2": "reference_name", + "start2": "reference_start", + "end2": "reference_end", + "ani_c": "perID_by_events", + }, + inplace=True, + ) + data["perID_by_events"] *= 100 + return data + data = pd.read_csv(file_path, delimiter="\t") return data diff --git a/tests/test_cli_runtime.py b/tests/test_cli_runtime.py index 9c5d8d9..d4632a4 100644 --- a/tests/test_cli_runtime.py +++ b/tests/test_cli_runtime.py @@ -3,6 +3,7 @@ from types import SimpleNamespace import numpy as np +import pandas as pd import pytest import moddotplot.moddotplot as cli @@ -281,6 +282,31 @@ def test_parser_defaults_omitted_subcommand_to_static(): assert args.no_plot +def test_static_parser_accepts_config_with_explicit_load_and_output(): + args = cli.parse_args( + [ + "--load", + "matrix.bedpe", + "--config", + "plot.json", + "--output-dir", + "plots", + ] + ) + + assert args.command == "static" + assert args.load == ["matrix.bedpe"] + assert args.config == "plot.json" + assert args.output_dir == "plots" + + +def test_static_parser_still_rejects_load_with_fasta(): + with pytest.raises(SystemExit) as exc_info: + cli.parse_args(["--load", "matrix.bedpe", "--fasta", "sequence.fa"]) + + assert exc_info.value.code == 2 + + def test_parser_preserves_explicit_interactive_subcommand(): args = cli.parse_args(["interactive", "--fasta", "sequence.fa"]) @@ -449,6 +475,48 @@ def test_static_config_accepts_sequence_selection(): assert args.sequence == ["chr1", "chr2"] +def test_static_config_preserves_explicit_input_and_output_paths(): + args = cli.parse_args( + [ + "--config", + "plot.json", + "--load", + "provided/matrix.bedpe", + "--output-dir", + "provided/plots", + ] + ) + + cli._apply_static_config( + args, + { + "fasta": ["configured.fa"], + "load": ["configured/matrix.bedpe"], + "output_dir": "configured/plots", + "palette": "Blues_7", + "sequence": ["chr1"], + "pairs": "pairs.tsv", + "region": ["chr1:1-100"], + }, + ) + + assert args.fasta is None + assert args.load == ["provided/matrix.bedpe"] + assert args.output_dir == "provided/plots" + assert args.palette == "Blues_7" + assert args.sequence is None + assert args.pairs is None + assert args.region is None + + +def test_loaded_bedpe_rows_are_filtered_at_identity_threshold(): + data = pd.DataFrame({"perID_by_events": [73.0, 91.699, 91.7, 100.0]}) + + filtered = cli._filter_loaded_identity(data, 91.7) + + assert filtered["perID_by_events"].tolist() == [91.7, 100.0] + + def test_static_delta_defaults_to_half_window(): args = cli.get_parser().parse_args(["static", "--fasta", "sequence.fa"]) diff --git a/tests/test_static_customization.py b/tests/test_static_customization.py index fbcdb23..b798e09 100644 --- a/tests/test_static_customization.py +++ b/tests/test_static_customization.py @@ -8,9 +8,39 @@ generate_breaks, get_colors, make_dot, + read_df_from_file, ) +def test_read_df_from_file_normalizes_browser_bedpe_export(tmp_path): + bedpe = tmp_path / "browser.bedpe" + bedpe.write_text( + "# moddotplot-interactive current-view BEDPE export\n" + '# provenance={"software":"0.9.5"}\n' + "#chrom1\tstart1\tend1\tchrom2\tstart2\tend2\tname\tscore\t" + "strand1\tstrand2\tani_c\tdirection\tdirection_support\n" + "chr1\t10\t20\tchr1\t30\t40\tani_c=0.9750\t975\t+\t+\t" + "0.9750\t1.000000\t12\n" + ) + + data = read_df_from_file(bedpe) + + assert data.columns.tolist() == [ + "#query_name", + "query_start", + "query_end", + "reference_name", + "reference_start", + "reference_end", + "perID_by_events", + ] + assert data.loc[0, "#query_name"] == "chr1" + assert data.loc[0, "query_start"] == 10 + assert data.loc[0, "reference_name"] == "chr1" + assert data.loc[0, "reference_start"] == 30 + assert data.loc[0, "perID_by_events"] == pytest.approx(97.5) + + def test_get_colors_uses_string_custom_breakpoints(): identity_scores = pd.DataFrame( { From 00ed120ebb9860b0f8a05b3149f4d29067a61658 Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Sat, 3 Oct 2026 23:27:37 -0400 Subject: [PATCH 10/16] Improve static input type errors --- src/moddotplot/moddotplot.py | 100 +++++++++++++++++++++++++++++++++++ tests/test_cli_runtime.py | 15 ++++++ tests/test_entrypoints.py | 27 ++++++++++ 3 files changed, 142 insertions(+) diff --git a/src/moddotplot/moddotplot.py b/src/moddotplot/moddotplot.py index 31220b9..009da59 100755 --- a/src/moddotplot/moddotplot.py +++ b/src/moddotplot/moddotplot.py @@ -39,6 +39,7 @@ from concurrent.futures import ProcessPoolExecutor, as_completed from contextlib import contextmanager, redirect_stderr, redirect_stdout from dataclasses import dataclass +import gzip from itertools import islice import math import hashlib @@ -62,6 +63,11 @@ require_interactive_dependencies = None COMMANDS = frozenset({"interactive", "static"}) +STATIC_INPUT_OPTIONS = { + "bedpe": ("--load", "BEDPE file"), + "fasta": ("--fasta", "FASTA file"), + "json": ("--config", "JSON config file"), +} INTERACTIVE_DEPRECATION_MESSAGE = ( "Warning: interactive mode is deprecated and maintenance-only. " "It remains available, but will not receive new features." @@ -700,6 +706,90 @@ def _apply_static_config(args, config): return args +def _static_input_kind(path): + """Identify known static input formats from content, then filename.""" + + path_string = os.fspath(path) + lines = [] + try: + with open(path_string, "rb") as probe: + compressed = probe.read(2) == b"\x1f\x8b" + opener = gzip.open if compressed else open + with opener( + path_string, + "rt", + encoding="utf-8", + errors="replace", + ) as input_file: + lines = list(islice(input_file, 16)) + except (OSError, EOFError): + pass + + for line in lines: + stripped = line.strip() + if not stripped: + continue + if stripped.startswith(">"): + return "fasta" + if stripped.startswith(("{", "[")): + return "json" + columns = set(stripped.removeprefix("#").split("\t")) + if { + "query_name", + "query_start", + "query_end", + "reference_name", + "reference_start", + "reference_end", + "perID_by_events", + }.issubset(columns) or { + "chrom1", + "start1", + "end1", + "chrom2", + "start2", + "end2", + "ani_c", + }.issubset( + columns + ): + return "bedpe" + + lower_path = path_string.lower() + if lower_path.endswith(".json"): + return "json" + if lower_path.endswith(".bedpe"): + return "bedpe" + fasta_path = lower_path + for compression_suffix in (".gz", ".bgz", ".bgzf"): + if fasta_path.endswith(compression_suffix): + fasta_path = fasta_path[: -len(compression_suffix)] + break + if fasta_path.endswith((".fa", ".fasta", ".fna", ".fas")): + return "fasta" + return None + + +def _validate_static_input_types(args): + """Reject known input formats passed through the wrong static option.""" + + checks = [] + if getattr(args, "config", None): + checks.append(("json", args.config)) + checks.extend(("bedpe", path) for path in (getattr(args, "load", None) or [])) + checks.extend(("fasta", path) for path in (getattr(args, "fasta", None) or [])) + for expected, path in checks: + detected = _static_input_kind(path) + if detected is None or detected == expected: + continue + expected_option, expected_name = STATIC_INPUT_OPTIONS[expected] + detected_option, detected_name = STATIC_INPUT_OPTIONS[detected] + raise ValueError( + f"{expected_option} expects a {expected_name}, but {path!r} appears " + f"to be a {detected_name}. Use {detected_option} for this file" + ) + + def _select_fasta_headers(fasta_headers, requested_sequences): """Filter FASTA headers using exact-first, case-insensitive selectors. @@ -2154,12 +2244,22 @@ def _main(arguments): sys.exit(0) elif args.command == "static": print(f"Running ModDotPlot in static mode\n") + try: + _validate_static_input_types(args) + except ValueError as error: + print(f"Error: {error}.", file=sys.stderr) + sys.exit(2) # -----------CONFIG PARSING----------- # TODO: Change to yml file, add readme to config folder if args.config: with open(args.config, "r") as f: config = json.load(f) _apply_static_config(args, config) + try: + _validate_static_input_types(args) + except ValueError as error: + print(f"Error: {error}.", file=sys.stderr) + sys.exit(2) if args.cooler: try: diff --git a/tests/test_cli_runtime.py b/tests/test_cli_runtime.py index d4632a4..8ea2e64 100644 --- a/tests/test_cli_runtime.py +++ b/tests/test_cli_runtime.py @@ -1,6 +1,7 @@ import sys import shlex from types import SimpleNamespace +import gzip import numpy as np import pandas as pd @@ -20,6 +21,20 @@ def _matrix_args(**overrides): return SimpleNamespace(**values) +def test_static_input_kind_detects_content_including_compressed_fasta(tmp_path): + config = tmp_path / "settings.data" + config.write_text('{"identity": 90}\n') + bedpe = tmp_path / "matrix.data" + bedpe.write_text("#chrom1\tstart1\tend1\tchrom2\tstart2\tend2\tani_c\n") + fasta = tmp_path / "sequence.data.gz" + with gzip.open(fasta, "wt") as output: + output.write(">chr1\nACGT\n") + + assert cli._static_input_kind(config) == "json" + assert cli._static_input_kind(bedpe) == "bedpe" + assert cli._static_input_kind(fasta) == "fasta" + + def test_matrix_config_caps_resolution_at_one_valid_kmer_per_window(): args = _matrix_args() diff --git a/tests/test_entrypoints.py b/tests/test_entrypoints.py index 706caa8..03dbbd0 100644 --- a/tests/test_entrypoints.py +++ b/tests/test_entrypoints.py @@ -86,3 +86,30 @@ def test_abbreviated_quiet_option_is_rejected_without_silencing_the_error(): assert result.returncode == 2 assert "unrecognized arguments: --qui" in result.stderr + + +def test_static_options_explain_when_a_known_input_type_uses_the_wrong_flag(tmp_path): + config = tmp_path / "settings.json" + config.write_text('{"identity": 90}\n') + fasta = tmp_path / "sequence.fa" + fasta.write_text(">chr1\nACGT\n") + bedpe = tmp_path / "matrix.bedpe" + bedpe.write_text( + "#query_name\tquery_start\tquery_end\treference_name\t" + "reference_start\treference_end\tperID_by_events\n" + ) + cases = [ + ("--load", config, "--load expects a BEDPE file", "Use --config"), + ("--load", fasta, "--load expects a BEDPE file", "Use --fasta"), + ("--fasta", config, "--fasta expects a FASTA file", "Use --config"), + ("--fasta", bedpe, "--fasta expects a FASTA file", "Use --load"), + ("--config", fasta, "--config expects a JSON config file", "Use --fasta"), + ("--config", bedpe, "--config expects a JSON config file", "Use --load"), + ] + + for option, path, expectation, suggestion in cases: + result = _run_module(option, str(path)) + assert result.returncode == 2 + assert expectation in result.stderr + assert suggestion in result.stderr + assert "Traceback" not in result.stderr From aea29f6819f123f1ae0f65b0e8ead5e405ea7fc6 Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Sat, 3 Oct 2026 23:34:30 -0400 Subject: [PATCH 11/16] Apply current Black formatting --- benchmarks/benchmark_hashing.py | 8 ++++---- src/moddotplot/moddotplot.py | 8 +++++--- tests/test_optional_dependencies.py | 1 - 3 files changed, 9 insertions(+), 8 deletions(-) diff --git a/benchmarks/benchmark_hashing.py b/benchmarks/benchmark_hashing.py index b620997..27a2667 100644 --- a/benchmarks/benchmark_hashing.py +++ b/benchmarks/benchmark_hashing.py @@ -142,10 +142,10 @@ def benchmark(sequence: str, k: int, repeats: int, seed: int) -> Dict[str, objec ) } if mmh3_module is not None: - implementations[ - "legacy_mmh3" - ] = lambda canonical=canonical: _legacy_mmh3_hashes( - sequence, k, canonical, mmh3_module + implementations["legacy_mmh3"] = ( + lambda canonical=canonical: _legacy_mmh3_hashes( + sequence, k, canonical, mmh3_module + ) ) raw: Dict[str, List[float]] = {name: [] for name in implementations} diff --git a/src/moddotplot/moddotplot.py b/src/moddotplot/moddotplot.py index 009da59..4341544 100755 --- a/src/moddotplot/moddotplot.py +++ b/src/moddotplot/moddotplot.py @@ -1097,9 +1097,11 @@ def _plan_static_pairs(args, records): pairs = _read_pair_file(args.pairs, records) if args.compare_order == "size": pairs = [ - (y_record, x_record) - if x_record.selected_length < y_record.selected_length - else (x_record, y_record) + ( + (y_record, x_record) + if x_record.selected_length < y_record.selected_length + else (x_record, y_record) + ) for x_record, y_record in pairs ] return pairs diff --git a/tests/test_optional_dependencies.py b/tests/test_optional_dependencies.py index 26eb943..3e7cca6 100644 --- a/tests/test_optional_dependencies.py +++ b/tests/test_optional_dependencies.py @@ -9,7 +9,6 @@ from moddotplot.estimate_identity import convertMatrixToCool, require_cooler_dependency from moddotplot.optional_dependencies import OptionalDependencyError - INSTALL_HINT = 'python -m pip install "ModDotPlot[interactive]"' From 00a73f7c856c9eb1309e8ee6ee04cdc977f75cc8 Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Sun, 4 Oct 2026 23:14:36 -0400 Subject: [PATCH 12/16] Reduce core runtime dependencies --- .github/workflows/ci.yml | 2 +- .github/workflows/publish-to-pypi.yml | 2 +- README.md | 29 +- THIRD_PARTY_LICENSES.md | 33 + pyproject.toml | 7 +- setup.cfg | 1 + src/moddotplot/__main__.py | 3 - src/moddotplot/color_palettes.py | 1744 ++++++++++++++++++++ src/moddotplot/estimate_identity.py | 29 +- src/moddotplot/moddotplot.py | 2 +- src/moddotplot/optional_dependencies.py | 12 +- src/moddotplot/static_plots.py | 1106 +++---------- tests/test_issue53_memory_safe_plotting.py | 32 +- tests/test_optional_dependencies.py | 13 +- tests/test_packaging_metadata.py | 37 +- tests/test_plot_fonts.py | 38 +- tests/test_static_customization.py | 46 +- 17 files changed, 2106 insertions(+), 1030 deletions(-) create mode 100644 THIRD_PARTY_LICENSES.md create mode 100644 src/moddotplot/color_palettes.py diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 1328cf0..e54a30a 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -77,7 +77,7 @@ jobs: - name: Install package and test dependencies run: | python -m pip install --upgrade pip - python -m pip install --editable ".[test,interactive]" + python -m pip install --editable ".[test,interactive,cooler]" - name: Run unit and integration tests run: >- diff --git a/.github/workflows/publish-to-pypi.yml b/.github/workflows/publish-to-pypi.yml index d35a953..69146d8 100644 --- a/.github/workflows/publish-to-pypi.yml +++ b/.github/workflows/publish-to-pypi.yml @@ -42,7 +42,7 @@ jobs: - name: Install package and tests run: | python -m pip install --upgrade pip - python -m pip install ".[test,interactive]" + python -m pip install ".[test,interactive,cooler]" - name: Run the release test suite run: python -m pytest diff --git a/README.md b/README.md index 406096e..3fec657 100644 --- a/README.md +++ b/README.md @@ -37,12 +37,6 @@ If you use ModDotPlot for your research, please cite our software! _ModDotPlot_ is a dot plot visualization tool designed to be used at scale, both for smaller sequences and whole genomes. _ModDotPlot_ is the spiritual successor to [StainedGlass](https://mrvollger.github.io/StainedGlass/). The core algorithm breaks an input sequence down into intervals of sketched *k*-mers called **mod**imizers. This enables the rapid approximation of the Average Nucleotide Identity between combinations of intervals! -Version 1.0.0 uses a bundled [ntHash2](https://github.com/BirolLab/ntHash) implementation for k-mer hashing, replacing the previous `mmh3` runtime dependency. Hash values and exact sketches therefore differ from pre-1.0 releases; regenerate data instead of mixing sketches produced by the two algorithms. Previously saved interactive matrices remain loadable because they contain completed matrices rather than raw hashes. - -FASTA parsing and static BED annotation rendering are also built into ModDotPlot in version 1.0.0, replacing the previous `pysam` and `pyGenomeTracks` runtime dependencies. Plain FASTA, gzip-compressed FASTA, and BGZF-compressed FASTA inputs remain supported, and static annotations produce both PNG and the selected SVG, PDF, or PostScript vector format. - -Static triangle plots, annotation layouts, and multi-sequence grids are now composed directly with Matplotlib. This replaces the previous `CairoSVG`, `svgutils`, and `patchworklib` image-conversion and SVG-composition dependencies while retaining raster and vector output formats. - ![](images/demo.gif) If you're interested in learning more about _ModDotPlot_ and how to visualize tandem repeats, we have an in-depth [YouTube video tutorial](https://www.youtube.com/watch?v=_7sQaljB_ys&t=2321s&pp=ygUXYWxleCBzd2VldGVuIG1vZGRvdHBsb3Q%3D) hosted by the [BioDiversity Genomics Academy](https://thebgacademy.org). @@ -51,7 +45,7 @@ If you're interested in learning more about _ModDotPlot_ and how to visualize ta ## Installation -_ModDotPlot_ can be installed for static plotting by running `pip install moddotplot`. Interactive plotting and Cooler export use optional dependencies; install them with `pip install "ModDotPlot[interactive]"`. ModDotPlot supports Python 3.10 through 3.14, using Matplotlib 3.10.9 or newer on Python 3.10, Matplotlib 3.11.2 or newer on later Python versions, and the Plotnine 0.15 release line. Alternatively, you can download the current release from GitHub by using: +_ModDotPlot_ can be installed for static plotting by running `pip install moddotplot`. Alternatively, you can download the current release from GitHub by using: ``` git clone https://github.com/marbl/ModDotPlot.git @@ -121,7 +115,7 @@ moddotplot static -f sequence.fa moddotplot static ``` -Running _ModDotPlot_ in static mode quickly create plots under the specified output directory `-o`. By default, running _ModDotPlot_ in static mode this will produce the following files: +Running _ModDotPlot_ quickly create plots under the specified output directory `-o`. By default, running _ModDotPlot_ will produce the following files: - A paired-end bed file `.bedpe`, containing intervals alongside their corresponding identity estimates. - A self-identity dotplot for each sequence, as both an upper triangle matrix `_TRI` and full matrix `_FULL` representation. @@ -129,15 +123,19 @@ Running _ModDotPlot_ in static mode quickly create plots under the specified out ![](images/moddotplot_output.png) -Plots and histograms are output as both rasterized `.png` images and vector graphics (default: `.svg`). [Plotnine](https://plotnine.org/) provides the primary plotting interface, while Matplotlib directly renders triangle plots, annotation layouts, multi-sequence grids, and each requested output format. Grid axes state their genomic unit (Kbp, Mbp, or Gbp). Plot text uses Helvetica by default with an automatic DejaVu Sans fallback if Helvetica cannot render a glyph. +Plots and histograms are rendered with Matplotlib and output as both rasterized +`.png` images and vector graphics (default: `.svg`). Grid axes state their +genomic unit (Kbp, Mbp, or Gbp). Every directory containing generated static plots also receives a `plot_summary.txt` reproducibility record. It lists the creation time, absolute plot and input paths, window size, any selected region or annotation BED file, and the exact command used for the run. -_ModDotPlot_ supports highly customizable plotting features in static mode. See [static mode commands](#static-mode-commands) for a complete list of features. +_ModDotPlot_ supports highly customizable plotting features in static mode. See [plot customization](#static-mode-commands) for a complete list of features. ### Interactive Mode +**As of ModDotPlot v1.0.0, interactive mode has been deprecated!** + ``` moddotplot interactive ``` @@ -246,7 +244,7 @@ genome. `--cooler ` -If set, will output a matrix as a Cooler file for each input sequence, in addition to a BEDPE file. Cooler support is part of the optional dependency set installed with `pip install "ModDotPlot[interactive]"` (or `pip install ".[interactive]"` from a source checkout). +If set, will output a matrix as a Cooler file for each input sequence, in addition to a BEDPE file. Install Cooler support with `pip install "ModDotPlot[cooler]"` (or `pip install ".[cooler]"` from a source checkout). `--no-bedpe ` @@ -313,7 +311,10 @@ Plot only the requested 1-based, inclusive range for each named sequence. Syntax `--palette ` -List of accepted palettes can be found [here](https://jiffyclub.github.io/palettable/colorbrewer/). Palettes are segregated into 3 types: _Diverging_, _Qualitative_, and _Sequential_. Syntax is the name of the palette, followed by an underscore and the number of colors, eg. `OrRd_8`. Default is `Spectral_11`. +The accepted palettes use [ColorBrewer](https://colorbrewer2.org/) color +specifications and are segregated into 3 types: _Diverging_, _Qualitative_, +and _Sequential_. Syntax is the name of the palette, followed by an underscore +and the number of colors, for example `OrRd_8`. The default is `Spectral_11`. `--breakpoints ` @@ -588,6 +589,4 @@ For bug reports or general usage questions, please raise a GitHub issue, or emai ## Known Issues -- Mac users might encounter the following unexpected command line output: `/bin/sh: lscpu: command not found`. This is a known issue with Plotnine, the Python plotting library used by ModDotPlot. This can be safely ignored. - -- When the optional `ModDotPlot[interactive]` dependencies are installed, Cooler may report `UserWarning: h5py is running against HDF5 1.xx.x when it was built against 1.xx.x, this may cause problems`. This can be safely ignored. To remove the warning, reinstall h5py against the local HDF5 library with `pip uninstall -y h5py` followed by `pip install --no-binary=h5py h5py`. +- When the optional `ModDotPlot[cooler]` dependencies are installed, Cooler may report `UserWarning: h5py is running against HDF5 1.xx.x when it was built against 1.xx.x, this may cause problems`. This can be safely ignored. To remove the warning, reinstall h5py against the local HDF5 library with `pip uninstall -y h5py` followed by `pip install --no-binary=h5py h5py`. diff --git a/THIRD_PARTY_LICENSES.md b/THIRD_PARTY_LICENSES.md new file mode 100644 index 0000000..b435e48 --- /dev/null +++ b/THIRD_PARTY_LICENSES.md @@ -0,0 +1,33 @@ +# Third-party licenses + +## ColorBrewer color schemes + +This product includes color specifications and designs developed by Cynthia +Brewer ([colorbrewer.org](https://colorbrewer.org/)). + +Copyright © 2002 Cynthia Brewer, Mark Harrower, and The Pennsylvania State +University. + +Licensed under the Apache License, Version 2.0 (the “License”); you may not use +these color specifications except in compliance with the License. You may +obtain a copy of the License at +. + +Unless required by applicable law or agreed to in writing, software distributed +under the License is distributed on an “AS IS” BASIS, WITHOUT WARRANTIES OR +CONDITIONS OF ANY KIND, either express or implied. See the License for the +specific language governing permissions and limitations under the License. + +Redistribution and use in source and binary forms, with or without +modification, are permitted provided that: + +1. Redistributions as source code retain this copyright notice, this list of + conditions, and the preceding disclaimer. +2. End-user documentation includes the acknowledgment at the beginning of this + section, or that acknowledgment appears wherever third-party acknowledgments + normally appear. +3. The name “ColorBrewer” is not used to endorse or promote derived products + without prior written permission from Cynthia Brewer. +4. Derived products are not called “ColorBrewer” and do not include + “ColorBrewer” in their names without prior written permission from Cynthia + Brewer. diff --git a/pyproject.toml b/pyproject.toml index b12e4d0..ecfc6c1 100755 --- a/pyproject.toml +++ b/pyproject.toml @@ -10,9 +10,6 @@ dependencies = [ "pandas", "matplotlib>=3.10.9; python_version < '3.11'", "matplotlib>=3.11.2; python_version >= '3.11'", - "plotnine>=0.15.8,<0.16", - "palettable", - "setproctitle", "numpy", "scipy", ] @@ -41,10 +38,12 @@ moddotplot = "moddotplot.__main__:main" [project.optional-dependencies] interactive = [ - "cooler", "dash>=2.9", "plotly", ] +cooler = [ + "cooler", +] # development dependency groups test = [ "pytest", diff --git a/setup.cfg b/setup.cfg index a57e107..d142cd3 100644 --- a/setup.cfg +++ b/setup.cfg @@ -2,6 +2,7 @@ license_files = LICENSE LICENSE.ntHash + THIRD_PARTY_LICENSES.md [bdist_wheel] py_limited_api = cp38 diff --git a/src/moddotplot/__main__.py b/src/moddotplot/__main__.py index 59c51bc..3da72e9 100644 --- a/src/moddotplot/__main__.py +++ b/src/moddotplot/__main__.py @@ -5,11 +5,8 @@ def _run(arguments): - import setproctitle - from moddotplot.moddotplot import main as run_moddotplot - setproctitle.setproctitle("ModDotPlot") return run_moddotplot(arguments) diff --git a/src/moddotplot/color_palettes.py b/src/moddotplot/color_palettes.py new file mode 100644 index 0000000..9ddf2c4 --- /dev/null +++ b/src/moddotplot/color_palettes.py @@ -0,0 +1,1744 @@ +"""Bundled ColorBrewer palettes used by ModDotPlot. + +Color specifications and designs are by Cynthia Brewer, Mark Harrower, and +The Pennsylvania State University (http://colorbrewer.org/). + +Copyright (c) 2002 Cynthia Brewer, Mark Harrower, and The Pennsylvania State +University. Licensed under the Apache License, Version 2.0, and redistributed +under the additional terms in this project's THIRD_PARTY_LICENSES.md. +""" + +PALETTES = { + "Accent_3": ("#7FC97F", "#BEAED4", "#FDC086"), + "Accent_4": ("#7FC97F", "#BEAED4", "#FDC086", "#FFFF99"), + "Accent_5": ("#7FC97F", "#BEAED4", "#FDC086", "#FFFF99", "#386CB0"), + "Accent_6": ("#7FC97F", "#BEAED4", "#FDC086", "#FFFF99", "#386CB0", "#F0027F"), + "Accent_7": ( + "#7FC97F", + "#BEAED4", + "#FDC086", + "#FFFF99", + "#386CB0", + "#F0027F", + "#BF5B17", + ), + "Accent_8": ( + "#7FC97F", + "#BEAED4", + "#FDC086", + "#FFFF99", + "#386CB0", + "#F0027F", + "#BF5B17", + "#666666", + ), + "Blues_3": ("#DEEBF7", "#9ECAE1", "#3182BD"), + "Blues_4": ("#EFF3FF", "#BDD7E7", "#6BAED6", "#2171B5"), + "Blues_5": ("#EFF3FF", "#BDD7E7", "#6BAED6", "#3182BD", "#08519C"), + "Blues_6": ("#EFF3FF", "#C6DBEF", "#9ECAE1", "#6BAED6", "#3182BD", "#08519C"), + "Blues_7": ( + "#EFF3FF", + "#C6DBEF", + "#9ECAE1", + "#6BAED6", + "#4292C6", + "#2171B5", + "#084594", + ), + "Blues_8": ( + "#F7FBFF", + "#DEEBF7", + "#C6DBEF", + "#9ECAE1", + "#6BAED6", + "#4292C6", + "#2171B5", + "#084594", + ), + "Blues_9": ( + "#F7FBFF", + "#DEEBF7", + "#C6DBEF", + "#9ECAE1", + "#6BAED6", + "#4292C6", + "#2171B5", + "#08519C", + "#08306B", + ), + "BrBG_10": ( + "#543005", + "#8C510A", + "#BF812D", + "#DFC27D", + "#F6E8C3", + "#C7EAE5", + "#80CDC1", + "#35978F", + "#01665E", + "#003C30", + ), + "BrBG_11": ( + "#543005", + "#8C510A", + "#BF812D", + "#DFC27D", + "#F6E8C3", + "#F5F5F5", + "#C7EAE5", + "#80CDC1", + "#35978F", + "#01665E", + "#003C30", + ), + "BrBG_3": ("#D8B365", "#F5F5F5", "#5AB4AC"), + "BrBG_4": ("#A6611A", "#DFC27D", "#80CDC1", "#018571"), + "BrBG_5": ("#A6611A", "#DFC27D", "#F5F5F5", "#80CDC1", "#018571"), + "BrBG_6": ("#8C510A", "#D8B365", "#F6E8C3", "#C7EAE5", "#5AB4AC", "#01665E"), + "BrBG_7": ( + "#8C510A", + "#D8B365", + "#F6E8C3", + "#F5F5F5", + "#C7EAE5", + "#5AB4AC", + "#01665E", + ), + "BrBG_8": ( + "#8C510A", + "#BF812D", + "#DFC27D", + "#F6E8C3", + "#C7EAE5", + "#80CDC1", + "#35978F", + "#01665E", + ), + "BrBG_9": ( + "#8C510A", + "#BF812D", + "#DFC27D", + "#F6E8C3", + "#F5F5F5", + "#C7EAE5", + "#80CDC1", + "#35978F", + "#01665E", + ), + "BuGn_3": ("#E5F5F9", "#99D8C9", "#2CA25F"), + "BuGn_4": ("#EDF8FB", "#B2E2E2", "#66C2A4", "#238B45"), + "BuGn_5": ("#EDF8FB", "#B2E2E2", "#66C2A4", "#2CA25F", "#006D2C"), + "BuGn_6": ("#EDF8FB", "#CCECE6", "#99D8C9", "#66C2A4", "#2CA25F", "#006D2C"), + "BuGn_7": ( + "#EDF8FB", + "#CCECE6", + "#99D8C9", + "#66C2A4", + "#41AE76", + "#238B45", + "#005824", + ), + "BuGn_8": ( + "#F7FCFD", + "#E5F5F9", + "#CCECE6", + "#99D8C9", + "#66C2A4", + "#41AE76", + "#238B45", + "#005824", + ), + "BuGn_9": ( + "#F7FCFD", + "#E5F5F9", + "#CCECE6", + "#99D8C9", + "#66C2A4", + "#41AE76", + "#238B45", + "#006D2C", + "#00441B", + ), + "BuPu_3": ("#E0ECF4", "#9EBCDA", "#8856A7"), + "BuPu_4": ("#EDF8FB", "#B3CDE3", "#8C96C6", "#88419D"), + "BuPu_5": ("#EDF8FB", "#B3CDE3", "#8C96C6", "#8856A7", "#810F7C"), + "BuPu_6": ("#EDF8FB", "#BFD3E6", "#9EBCDA", "#8C96C6", "#8856A7", "#810F7C"), + "BuPu_7": ( + "#EDF8FB", + "#BFD3E6", + "#9EBCDA", + "#8C96C6", + "#8C6BB1", + "#88419D", + "#6E016B", + ), + "BuPu_8": ( + "#F7FCFD", + "#E0ECF4", + "#BFD3E6", + "#9EBCDA", + "#8C96C6", + "#8C6BB1", + "#88419D", + "#6E016B", + ), + "BuPu_9": ( + "#F7FCFD", + "#E0ECF4", + "#BFD3E6", + "#9EBCDA", + "#8C96C6", + "#8C6BB1", + "#88419D", + "#810F7C", + "#4D004B", + ), + "Dark2_3": ("#1B9E77", "#D95F02", "#7570B3"), + "Dark2_4": ("#1B9E77", "#D95F02", "#7570B3", "#E7298A"), + "Dark2_5": ("#1B9E77", "#D95F02", "#7570B3", "#E7298A", "#66A61E"), + "Dark2_6": ("#1B9E77", "#D95F02", "#7570B3", "#E7298A", "#66A61E", "#E6AB02"), + "Dark2_7": ( + "#1B9E77", + "#D95F02", + "#7570B3", + "#E7298A", + "#66A61E", + "#E6AB02", + "#A6761D", + ), + "Dark2_8": ( + "#1B9E77", + "#D95F02", + "#7570B3", + "#E7298A", + "#66A61E", + "#E6AB02", + "#A6761D", + "#666666", + ), + "GnBu_3": ("#E0F3DB", "#A8DDB5", "#43A2CA"), + "GnBu_4": ("#F0F9E8", "#BAE4BC", "#7BCCC4", "#2B8CBE"), + "GnBu_5": ("#F0F9E8", "#BAE4BC", "#7BCCC4", "#43A2CA", "#0868AC"), + "GnBu_6": ("#F0F9E8", "#CCEBC5", "#A8DDB5", "#7BCCC4", "#43A2CA", "#0868AC"), + "GnBu_7": ( + "#F0F9E8", + "#CCEBC5", + "#A8DDB5", + "#7BCCC4", + "#4EB3D3", + "#2B8CBE", + "#08589E", + ), + "GnBu_8": ( + "#F7FCF0", + "#E0F3DB", + "#CCEBC5", + "#A8DDB5", + "#7BCCC4", + "#4EB3D3", + "#2B8CBE", + "#08589E", + ), + "GnBu_9": ( + "#F7FCF0", + "#E0F3DB", + "#CCEBC5", + "#A8DDB5", + "#7BCCC4", + "#4EB3D3", + "#2B8CBE", + "#0868AC", + "#084081", + ), + "Greens_3": ("#E5F5E0", "#A1D99B", "#31A354"), + "Greens_4": ("#EDF8E9", "#BAE4B3", "#74C476", "#238B45"), + "Greens_5": ("#EDF8E9", "#BAE4B3", "#74C476", "#31A354", "#006D2C"), + "Greens_6": ("#EDF8E9", "#C7E9C0", "#A1D99B", "#74C476", "#31A354", "#006D2C"), + "Greens_7": ( + "#EDF8E9", + "#C7E9C0", + "#A1D99B", + "#74C476", + "#41AB5D", + "#238B45", + "#005A32", + ), + "Greens_8": ( + "#F7FCF5", + "#E5F5E0", + "#C7E9C0", + "#A1D99B", + "#74C476", + "#41AB5D", + "#238B45", + "#005A32", + ), + "Greens_9": ( + "#F7FCF5", + "#E5F5E0", + "#C7E9C0", + "#A1D99B", + "#74C476", + "#41AB5D", + "#238B45", + "#006D2C", + "#00441B", + ), + "Greys_3": ("#F0F0F0", "#BDBDBD", "#636363"), + "Greys_4": ("#F7F7F7", "#CCCCCC", "#969696", "#525252"), + "Greys_5": ("#F7F7F7", "#CCCCCC", "#969696", "#636363", "#252525"), + "Greys_6": ("#F7F7F7", "#D9D9D9", "#BDBDBD", "#969696", "#636363", "#252525"), + "Greys_7": ( + "#F7F7F7", + "#D9D9D9", + "#BDBDBD", + "#969696", + "#737373", + "#525252", + "#252525", + ), + "Greys_8": ( + "#FFFFFF", + "#F0F0F0", + "#D9D9D9", + "#BDBDBD", + "#969696", + "#737373", + "#525252", + "#252525", + ), + "Greys_9": ( + "#FFFFFF", + "#F0F0F0", + "#D9D9D9", + "#BDBDBD", + "#969696", + "#737373", + "#525252", + "#252525", + "#000000", + ), + "OrRd_3": ("#FEE8C8", "#FDBB84", "#E34A33"), + "OrRd_4": ("#FEF0D9", "#FDCC8A", "#FC8D59", "#D7301F"), + "OrRd_5": ("#FEF0D9", "#FDCC8A", "#FC8D59", "#E34A33", "#B30000"), + "OrRd_6": ("#FEF0D9", "#FDD49E", "#FDBB84", "#FC8D59", "#E34A33", "#B30000"), + "OrRd_7": ( + "#FEF0D9", + "#FDD49E", + "#FDBB84", + "#FC8D59", + "#EF6548", + "#D7301F", + "#990000", + ), + "OrRd_8": ( + "#FFF7EC", + "#FEE8C8", + "#FDD49E", + "#FDBB84", + "#FC8D59", + "#EF6548", + "#D7301F", + "#990000", + ), + "OrRd_9": ( + "#FFF7EC", + "#FEE8C8", + "#FDD49E", + "#FDBB84", + "#FC8D59", + "#EF6548", + "#D7301F", + "#B30000", + "#7F0000", + ), + "Oranges_3": ("#FEE6CE", "#FDAE6B", "#E6550D"), + "Oranges_4": ("#FEEDDE", "#FDBE85", "#FD8D3C", "#D94701"), + "Oranges_5": ("#FEEDDE", "#FDBE85", "#FD8D3C", "#E6550D", "#A63603"), + "Oranges_6": ("#FEEDDE", "#FDD0A2", "#FDAE6B", "#FD8D3C", "#E6550D", "#A63603"), + "Oranges_7": ( + "#FEEDDE", + "#FDD0A2", + "#FDAE6B", + "#FD8D3C", + "#F16913", + "#D94801", + "#8C2D04", + ), + "Oranges_8": ( + "#FFF5EB", + "#FEE6CE", + "#FDD0A2", + "#FDAE6B", + "#FD8D3C", + "#F16913", + "#D94801", + "#8C2D04", + ), + "Oranges_9": ( + "#FFF5EB", + "#FEE6CE", + "#FDD0A2", + "#FDAE6B", + "#FD8D3C", + "#F16913", + "#D94801", + "#A63603", + "#7F2704", + ), + "PRGn_10": ( + "#40004B", + "#762A83", + "#9970AB", + "#C2A5CF", + "#E7D4E8", + "#D9F0D3", + "#A6DBA0", + "#5AAE61", + "#1B7837", + "#00441B", + ), + "PRGn_11": ( + "#40004B", + "#762A83", + "#9970AB", + "#C2A5CF", + "#E7D4E8", + "#F7F7F7", + "#D9F0D3", + "#A6DBA0", + "#5AAE61", + "#1B7837", + "#00441B", + ), + "PRGn_3": ("#AF8DC3", "#F7F7F7", "#7FBF7B"), + "PRGn_4": ("#7B3294", "#C2A5CF", "#A6DBA0", "#008837"), + "PRGn_5": ("#7B3294", "#C2A5CF", "#F7F7F7", "#A6DBA0", "#008837"), + "PRGn_6": ("#762A83", "#AF8DC3", "#E7D4E8", "#D9F0D3", "#7FBF7B", "#1B7837"), + "PRGn_7": ( + "#762A83", + "#AF8DC3", + "#E7D4E8", + "#F7F7F7", + "#D9F0D3", + "#7FBF7B", + "#1B7837", + ), + "PRGn_8": ( + "#762A83", + "#9970AB", + "#C2A5CF", + "#E7D4E8", + "#D9F0D3", + "#A6DBA0", + "#5AAE61", + "#1B7837", + ), + "PRGn_9": ( + "#762A83", + "#9970AB", + "#C2A5CF", + "#E7D4E8", + "#F7F7F7", + "#D9F0D3", + "#A6DBA0", + "#5AAE61", + "#1B7837", + ), + "Paired_10": ( + "#A6CEE3", + "#1F78B4", + "#B2DF8A", + "#33A02C", + "#FB9A99", + "#E31A1C", + "#FDBF6F", + "#FF7F00", + "#CAB2D6", + "#6A3D9A", + ), + "Paired_11": ( + "#A6CEE3", + "#1F78B4", + "#B2DF8A", + "#33A02C", + "#FB9A99", + "#E31A1C", + "#FDBF6F", + "#FF7F00", + "#CAB2D6", + "#6A3D9A", + "#FFFF99", + ), + "Paired_12": ( + "#A6CEE3", + "#1F78B4", + "#B2DF8A", + "#33A02C", + "#FB9A99", + "#E31A1C", + "#FDBF6F", + "#FF7F00", + "#CAB2D6", + "#6A3D9A", + "#FFFF99", + "#B15928", + ), + "Paired_3": ("#A6CEE3", "#1F78B4", "#B2DF8A"), + "Paired_4": ("#A6CEE3", "#1F78B4", "#B2DF8A", "#33A02C"), + "Paired_5": ("#A6CEE3", "#1F78B4", "#B2DF8A", "#33A02C", "#FB9A99"), + "Paired_6": ("#A6CEE3", "#1F78B4", "#B2DF8A", "#33A02C", "#FB9A99", "#E31A1C"), + "Paired_7": ( + "#A6CEE3", + "#1F78B4", + "#B2DF8A", + "#33A02C", + "#FB9A99", + "#E31A1C", + "#FDBF6F", + ), + "Paired_8": ( + "#A6CEE3", + "#1F78B4", + "#B2DF8A", + "#33A02C", + "#FB9A99", + "#E31A1C", + "#FDBF6F", + "#FF7F00", + ), + "Paired_9": ( + "#A6CEE3", + "#1F78B4", + "#B2DF8A", + "#33A02C", + "#FB9A99", + "#E31A1C", + "#FDBF6F", + "#FF7F00", + "#CAB2D6", + ), + "Pastel1_3": ("#FBB4AE", "#B3CDE3", "#CCEBC5"), + "Pastel1_4": ("#FBB4AE", "#B3CDE3", "#CCEBC5", "#DECBE4"), + "Pastel1_5": ("#FBB4AE", "#B3CDE3", "#CCEBC5", "#DECBE4", "#FED9A6"), + "Pastel1_6": ("#FBB4AE", "#B3CDE3", "#CCEBC5", "#DECBE4", "#FED9A6", "#FFFFCC"), + "Pastel1_7": ( + "#FBB4AE", + "#B3CDE3", + "#CCEBC5", + "#DECBE4", + "#FED9A6", + "#FFFFCC", + "#E5D8BD", + ), + "Pastel1_8": ( + "#FBB4AE", + "#B3CDE3", + "#CCEBC5", + "#DECBE4", + "#FED9A6", + "#FFFFCC", + "#E5D8BD", + "#FDDAEC", + ), + "Pastel1_9": ( + "#FBB4AE", + "#B3CDE3", + "#CCEBC5", + "#DECBE4", + "#FED9A6", + "#FFFFCC", + "#E5D8BD", + "#FDDAEC", + "#F2F2F2", + ), + "Pastel2_3": ("#B3E2CD", "#FDCDAC", "#CBD5E8"), + "Pastel2_4": ("#B3E2CD", "#FDCDAC", "#CBD5E8", "#F4CAE4"), + "Pastel2_5": ("#B3E2CD", "#FDCDAC", "#CBD5E8", "#F4CAE4", "#E6F5C9"), + "Pastel2_6": ("#B3E2CD", "#FDCDAC", "#CBD5E8", "#F4CAE4", "#E6F5C9", "#FFF2AE"), + "Pastel2_7": ( + "#B3E2CD", + "#FDCDAC", + "#CBD5E8", + "#F4CAE4", + "#E6F5C9", + "#FFF2AE", + "#F1E2CC", + ), + "Pastel2_8": ( + "#B3E2CD", + "#FDCDAC", + "#CBD5E8", + "#F4CAE4", + "#E6F5C9", + "#FFF2AE", + "#F1E2CC", + "#CCCCCC", + ), + "PiYG_10": ( + "#8E0152", + "#C51B7D", + "#DE77AE", + "#F1B6DA", + "#FDE0EF", + "#E6F5D0", + "#B8E186", + "#7FBC41", + "#4D9221", + "#276419", + ), + "PiYG_11": ( + "#8E0152", + "#C51B7D", + "#DE77AE", + "#F1B6DA", + "#FDE0EF", + "#F7F7F7", + "#E6F5D0", + "#B8E186", + "#7FBC41", + "#4D9221", + "#276419", + ), + "PiYG_3": ("#E9A3C9", "#F7F7F7", "#A1D76A"), + "PiYG_4": ("#D01C8B", "#F1B6DA", "#B8E186", "#4DAC26"), + "PiYG_5": ("#D01C8B", "#F1B6DA", "#F7F7F7", "#B8E186", "#4DAC26"), + "PiYG_6": ("#C51B7D", "#E9A3C9", "#FDE0EF", "#E6F5D0", "#A1D76A", "#4D9221"), + "PiYG_7": ( + "#C51B7D", + "#E9A3C9", + "#FDE0EF", + "#F7F7F7", + "#E6F5D0", + "#A1D76A", + "#4D9221", + ), + "PiYG_8": ( + "#C51B7D", + "#DE77AE", + "#F1B6DA", + "#FDE0EF", + "#E6F5D0", + "#B8E186", + "#7FBC41", + "#4D9221", + ), + "PiYG_9": ( + "#C51B7D", + "#DE77AE", + "#F1B6DA", + "#FDE0EF", + "#F7F7F7", + "#E6F5D0", + "#B8E186", + "#7FBC41", + "#4D9221", + ), + "PuBuGn_3": ("#ECE2F0", "#A6BDDB", "#1C9099"), + "PuBuGn_4": ("#F6EFF7", "#BDC9E1", "#67A9CF", "#02818A"), + "PuBuGn_5": ("#F6EFF7", "#BDC9E1", "#67A9CF", "#1C9099", "#016C59"), + "PuBuGn_6": ("#F6EFF7", "#D0D1E6", "#A6BDDB", "#67A9CF", "#1C9099", "#016C59"), + "PuBuGn_7": ( + "#F6EFF7", + "#D0D1E6", + "#A6BDDB", + "#67A9CF", + "#3690C0", + "#02818A", + "#016450", + ), + "PuBuGn_8": ( + "#FFF7FB", + "#ECE2F0", + "#D0D1E6", + "#A6BDDB", + "#67A9CF", + "#3690C0", + "#02818A", + "#016450", + ), + "PuBuGn_9": ( + "#FFF7FB", + "#ECE2F0", + "#D0D1E6", + "#A6BDDB", + "#67A9CF", + "#3690C0", + "#02818A", + "#016C59", + "#014636", + ), + "PuBu_3": ("#ECE7F2", "#A6BDDB", "#2B8CBE"), + "PuBu_4": ("#F1EEF6", "#BDC9E1", "#74A9CF", "#0570B0"), + "PuBu_5": ("#F1EEF6", "#BDC9E1", "#74A9CF", "#2B8CBE", "#045A8D"), + "PuBu_6": ("#F1EEF6", "#D0D1E6", "#A6BDDB", "#74A9CF", "#2B8CBE", "#045A8D"), + "PuBu_7": ( + "#F1EEF6", + "#D0D1E6", + "#A6BDDB", + "#74A9CF", + "#3690C0", + "#0570B0", + "#034E7B", + ), + "PuBu_8": ( + "#FFF7FB", + "#ECE7F2", + "#D0D1E6", + "#A6BDDB", + "#74A9CF", + "#3690C0", + "#0570B0", + "#034E7B", + ), + "PuBu_9": ( + "#FFF7FB", + "#ECE7F2", + "#D0D1E6", + "#A6BDDB", + "#74A9CF", + "#3690C0", + "#0570B0", + "#045A8D", + "#023858", + ), + "PuOr_10": ( + "#7F3B08", + "#B35806", + "#E08214", + "#FDB863", + "#FEE0B6", + "#D8DAEB", + "#B2ABD2", + "#8073AC", + "#542788", + "#2D004B", + ), + "PuOr_11": ( + "#7F3B08", + "#B35806", + "#E08214", + "#FDB863", + "#FEE0B6", + "#F7F7F7", + "#D8DAEB", + "#B2ABD2", + "#8073AC", + "#542788", + "#2D004B", + ), + "PuOr_3": ("#F1A340", "#F7F7F7", "#998EC3"), + "PuOr_4": ("#E66101", "#FDB863", "#B2ABD2", "#5E3C99"), + "PuOr_5": ("#E66101", "#FDB863", "#F7F7F7", "#B2ABD2", "#5E3C99"), + "PuOr_6": ("#B35806", "#F1A340", "#FEE0B6", "#D8DAEB", "#998EC3", "#542788"), + "PuOr_7": ( + "#B35806", + "#F1A340", + "#FEE0B6", + "#F7F7F7", + "#D8DAEB", + "#998EC3", + "#542788", + ), + "PuOr_8": ( + "#B35806", + "#E08214", + "#FDB863", + "#FEE0B6", + "#D8DAEB", + "#B2ABD2", + "#8073AC", + "#542788", + ), + "PuOr_9": ( + "#B35806", + "#E08214", + "#FDB863", + "#FEE0B6", + "#F7F7F7", + "#D8DAEB", + "#B2ABD2", + "#8073AC", + "#542788", + ), + "PuRd_3": ("#E7E1EF", "#C994C7", "#DD1C77"), + "PuRd_4": ("#F1EEF6", "#D7B5D8", "#DF65B0", "#CE1256"), + "PuRd_5": ("#F1EEF6", "#D7B5D8", "#DF65B0", "#DD1C77", "#980043"), + "PuRd_6": ("#F1EEF6", "#D4B9DA", "#C994C7", "#DF65B0", "#DD1C77", "#980043"), + "PuRd_7": ( + "#F1EEF6", + "#D4B9DA", + "#C994C7", + "#DF65B0", + "#E7298A", + "#CE1256", + "#91003F", + ), + "PuRd_8": ( + "#F7F4F9", + "#E7E1EF", + "#D4B9DA", + "#C994C7", + "#DF65B0", + "#E7298A", + "#CE1256", + "#91003F", + ), + "PuRd_9": ( + "#F7F4F9", + "#E7E1EF", + "#D4B9DA", + "#C994C7", + "#DF65B0", + "#E7298A", + "#CE1256", + "#980043", + "#67001F", + ), + "Purples_3": ("#EFEDF5", "#BCBDDC", "#756BB1"), + "Purples_4": ("#F2F0F7", "#CBC9E2", "#9E9AC8", "#6A51A3"), + "Purples_5": ("#F2F0F7", "#CBC9E2", "#9E9AC8", "#756BB1", "#54278F"), + "Purples_6": ("#F2F0F7", "#DADAEB", "#BCBDDC", "#9E9AC8", "#756BB1", "#54278F"), + "Purples_7": ( + "#F2F0F7", + "#DADAEB", + "#BCBDDC", + "#9E9AC8", + "#807DBA", + "#6A51A3", + "#4A1486", + ), + "Purples_8": ( + "#FCFBFD", + "#EFEDF5", + "#DADAEB", + "#BCBDDC", + "#9E9AC8", + "#807DBA", + "#6A51A3", + "#4A1486", + ), + "Purples_9": ( + "#FCFBFD", + "#EFEDF5", + "#DADAEB", + "#BCBDDC", + "#9E9AC8", + "#807DBA", + "#6A51A3", + "#54278F", + "#3F007D", + ), + "RdBu_10": ( + "#67001F", + "#B2182B", + "#D6604D", + "#F4A582", + "#FDDBC7", + "#D1E5F0", + "#92C5DE", + "#4393C3", + "#2166AC", + "#053061", + ), + "RdBu_11": ( + "#67001F", + "#B2182B", + "#D6604D", + "#F4A582", + "#FDDBC7", + "#F7F7F7", + "#D1E5F0", + "#92C5DE", + "#4393C3", + "#2166AC", + "#053061", + ), + "RdBu_3": ("#EF8A62", "#F7F7F7", "#67A9CF"), + "RdBu_4": ("#CA0020", "#F4A582", "#92C5DE", "#0571B0"), + "RdBu_5": ("#CA0020", "#F4A582", "#F7F7F7", "#92C5DE", "#0571B0"), + "RdBu_6": ("#B2182B", "#EF8A62", "#FDDBC7", "#D1E5F0", "#67A9CF", "#2166AC"), + "RdBu_7": ( + "#B2182B", + "#EF8A62", + "#FDDBC7", + "#F7F7F7", + "#D1E5F0", + "#67A9CF", + "#2166AC", + ), + "RdBu_8": ( + "#B2182B", + "#D6604D", + "#F4A582", + "#FDDBC7", + "#D1E5F0", + "#92C5DE", + "#4393C3", + "#2166AC", + ), + "RdBu_9": ( + "#B2182B", + "#D6604D", + "#F4A582", + "#FDDBC7", + "#F7F7F7", + "#D1E5F0", + "#92C5DE", + "#4393C3", + "#2166AC", + ), + "RdGy_10": ( + "#67001F", + "#B2182B", + "#D6604D", + "#F4A582", + "#FDDBC7", + "#E0E0E0", + "#BABABA", + "#878787", + "#4D4D4D", + "#1A1A1A", + ), + "RdGy_11": ( + "#67001F", + "#B2182B", + "#D6604D", + "#F4A582", + "#FDDBC7", + "#FFFFFF", + "#E0E0E0", + "#BABABA", + "#878787", + "#4D4D4D", + "#1A1A1A", + ), + "RdGy_3": ("#EF8A62", "#FFFFFF", "#999999"), + "RdGy_4": ("#CA0020", "#F4A582", "#BABABA", "#404040"), + "RdGy_5": ("#CA0020", "#F4A582", "#FFFFFF", "#BABABA", "#404040"), + "RdGy_6": ("#B2182B", "#EF8A62", "#FDDBC7", "#E0E0E0", "#999999", "#4D4D4D"), + "RdGy_7": ( + "#B2182B", + "#EF8A62", + "#FDDBC7", + "#FFFFFF", + "#E0E0E0", + "#999999", + "#4D4D4D", + ), + "RdGy_8": ( + "#B2182B", + "#D6604D", + "#F4A582", + "#FDDBC7", + "#E0E0E0", + "#BABABA", + "#878787", + "#4D4D4D", + ), + "RdGy_9": ( + "#B2182B", + "#D6604D", + "#F4A582", + "#FDDBC7", + "#FFFFFF", + "#E0E0E0", + "#BABABA", + "#878787", + "#4D4D4D", + ), + "RdPu_3": ("#FDE0DD", "#FA9FB5", "#C51B8A"), + "RdPu_4": ("#FEEBE2", "#FBB4B9", "#F768A1", "#AE017E"), + "RdPu_5": ("#FEEBE2", "#FBB4B9", "#F768A1", "#C51B8A", "#7A0177"), + "RdPu_6": ("#FEEBE2", "#FCC5C0", "#FA9FB5", "#F768A1", "#C51B8A", "#7A0177"), + "RdPu_7": ( + "#FEEBE2", + "#FCC5C0", + "#FA9FB5", + "#F768A1", + "#DD3497", + "#AE017E", + "#7A0177", + ), + "RdPu_8": ( + "#FFF7F3", + "#FDE0DD", + "#FCC5C0", + "#FA9FB5", + "#F768A1", + "#DD3497", + "#AE017E", + "#7A0177", + ), + "RdPu_9": ( + "#FFF7F3", + "#FDE0DD", + "#FCC5C0", + "#FA9FB5", + "#F768A1", + "#DD3497", + "#AE017E", + "#7A0177", + "#49006A", + ), + "RdYlBu_10": ( + "#A50026", + "#D73027", + "#F46D43", + "#FDAE61", + "#FEE090", + "#E0F3F8", + "#ABD9E9", + "#74ADD1", + "#4575B4", + "#313695", + ), + "RdYlBu_11": ( + "#A50026", + "#D73027", + "#F46D43", + "#FDAE61", + "#FEE090", + "#FFFFBF", + "#E0F3F8", + "#ABD9E9", + "#74ADD1", + "#4575B4", + "#313695", + ), + "RdYlBu_3": ("#FC8D59", "#FFFFBF", "#91BFDB"), + "RdYlBu_4": ("#D7191C", "#FDAE61", "#ABD9E9", "#2C7BB6"), + "RdYlBu_5": ("#D7191C", "#FDAE61", "#FFFFBF", "#ABD9E9", "#2C7BB6"), + "RdYlBu_6": ("#D73027", "#FC8D59", "#FEE090", "#E0F3F8", "#91BFDB", "#4575B4"), + "RdYlBu_7": ( + "#D73027", + "#FC8D59", + "#FEE090", + "#FFFFBF", + "#E0F3F8", + "#91BFDB", + "#4575B4", + ), + "RdYlBu_8": ( + "#D73027", + "#F46D43", + "#FDAE61", + "#FEE090", + "#E0F3F8", + "#ABD9E9", + "#74ADD1", + "#4575B4", + ), + "RdYlBu_9": ( + "#D73027", + "#F46D43", + "#FDAE61", + "#FEE090", + "#FFFFBF", + "#E0F3F8", + "#ABD9E9", + "#74ADD1", + "#4575B4", + ), + "RdYlGn_10": ( + "#A50026", + "#D73027", + "#F46D43", + "#FDAE61", + "#FEE08B", + "#D9EF8B", + "#A6D96A", + "#66BD63", + "#1A9850", + "#006837", + ), + "RdYlGn_11": ( + "#A50026", + "#D73027", + "#F46D43", + "#FDAE61", + "#FEE08B", + "#FFFFBF", + "#D9EF8B", + "#A6D96A", + "#66BD63", + "#1A9850", + "#006837", + ), + "RdYlGn_3": ("#FC8D59", "#FFFFBF", "#91CF60"), + "RdYlGn_4": ("#D7191C", "#FDAE61", "#A6D96A", "#1A9641"), + "RdYlGn_5": ("#D7191C", "#FDAE61", "#FFFFBF", "#A6D96A", "#1A9641"), + "RdYlGn_6": ("#D73027", "#FC8D59", "#FEE08B", "#D9EF8B", "#91CF60", "#1A9850"), + "RdYlGn_7": ( + "#D73027", + "#FC8D59", + "#FEE08B", + "#FFFFBF", + "#D9EF8B", + "#91CF60", + "#1A9850", + ), + "RdYlGn_8": ( + "#D73027", + "#F46D43", + "#FDAE61", + "#FEE08B", + "#D9EF8B", + "#A6D96A", + "#66BD63", + "#1A9850", + ), + "RdYlGn_9": ( + "#D73027", + "#F46D43", + "#FDAE61", + "#FEE08B", + "#FFFFBF", + "#D9EF8B", + "#A6D96A", + "#66BD63", + "#1A9850", + ), + "Reds_3": ("#FEE0D2", "#FC9272", "#DE2D26"), + "Reds_4": ("#FEE5D9", "#FCAE91", "#FB6A4A", "#CB181D"), + "Reds_5": ("#FEE5D9", "#FCAE91", "#FB6A4A", "#DE2D26", "#A50F15"), + "Reds_6": ("#FEE5D9", "#FCBBA1", "#FC9272", "#FB6A4A", "#DE2D26", "#A50F15"), + "Reds_7": ( + "#FEE5D9", + "#FCBBA1", + "#FC9272", + "#FB6A4A", + "#EF3B2C", + "#CB181D", + "#99000D", + ), + "Reds_8": ( + "#FFF5F0", + "#FEE0D2", + "#FCBBA1", + "#FC9272", + "#FB6A4A", + "#EF3B2C", + "#CB181D", + "#99000D", + ), + "Reds_9": ( + "#FFF5F0", + "#FEE0D2", + "#FCBBA1", + "#FC9272", + "#FB6A4A", + "#EF3B2C", + "#CB181D", + "#A50F15", + "#67000D", + ), + "Set1_3": ("#E41A1C", "#377EB8", "#4DAF4A"), + "Set1_4": ("#E41A1C", "#377EB8", "#4DAF4A", "#984EA3"), + "Set1_5": ("#E41A1C", "#377EB8", "#4DAF4A", "#984EA3", "#FF7F00"), + "Set1_6": ("#E41A1C", "#377EB8", "#4DAF4A", "#984EA3", "#FF7F00", "#FFFF33"), + "Set1_7": ( + "#E41A1C", + "#377EB8", + "#4DAF4A", + "#984EA3", + "#FF7F00", + "#FFFF33", + "#A65628", + ), + "Set1_8": ( + "#E41A1C", + "#377EB8", + "#4DAF4A", + "#984EA3", + "#FF7F00", + "#FFFF33", + "#A65628", + "#F781BF", + ), + "Set1_9": ( + "#E41A1C", + "#377EB8", + "#4DAF4A", + "#984EA3", + "#FF7F00", + "#FFFF33", + "#A65628", + "#F781BF", + "#999999", + ), + "Set2_3": ("#66C2A5", "#FC8D62", "#8DA0CB"), + "Set2_4": ("#66C2A5", "#FC8D62", "#8DA0CB", "#E78AC3"), + "Set2_5": ("#66C2A5", "#FC8D62", "#8DA0CB", "#E78AC3", "#A6D854"), + "Set2_6": ("#66C2A5", "#FC8D62", "#8DA0CB", "#E78AC3", "#A6D854", "#FFD92F"), + "Set2_7": ( + "#66C2A5", + "#FC8D62", + "#8DA0CB", + "#E78AC3", + "#A6D854", + "#FFD92F", + "#E5C494", + ), + "Set2_8": ( + "#66C2A5", + "#FC8D62", + "#8DA0CB", + "#E78AC3", + "#A6D854", + "#FFD92F", + "#E5C494", + "#B3B3B3", + ), + "Set3_10": ( + "#8DD3C7", + "#FFFFB3", + "#BEBADA", + "#FB8072", + "#80B1D3", + "#FDB462", + "#B3DE69", + "#FCCDE5", + "#D9D9D9", + "#BC80BD", + ), + "Set3_11": ( + "#8DD3C7", + "#FFFFB3", + "#BEBADA", + "#FB8072", + "#80B1D3", + "#FDB462", + "#B3DE69", + "#FCCDE5", + "#D9D9D9", + "#BC80BD", + "#CCEBC5", + ), + "Set3_12": ( + "#8DD3C7", + "#FFFFB3", + "#BEBADA", + "#FB8072", + "#80B1D3", + "#FDB462", + "#B3DE69", + "#FCCDE5", + "#D9D9D9", + "#BC80BD", + "#CCEBC5", + "#FFED6F", + ), + "Set3_3": ("#8DD3C7", "#FFFFB3", "#BEBADA"), + "Set3_4": ("#8DD3C7", "#FFFFB3", "#BEBADA", "#FB8072"), + "Set3_5": ("#8DD3C7", "#FFFFB3", "#BEBADA", "#FB8072", "#80B1D3"), + "Set3_6": ("#8DD3C7", "#FFFFB3", "#BEBADA", "#FB8072", "#80B1D3", "#FDB462"), + "Set3_7": ( + "#8DD3C7", + "#FFFFB3", + "#BEBADA", + "#FB8072", + "#80B1D3", + "#FDB462", + "#B3DE69", + ), + "Set3_8": ( + "#8DD3C7", + "#FFFFB3", + "#BEBADA", + "#FB8072", + "#80B1D3", + "#FDB462", + "#B3DE69", + "#FCCDE5", + ), + "Set3_9": ( + "#8DD3C7", + "#FFFFB3", + "#BEBADA", + "#FB8072", + "#80B1D3", + "#FDB462", + "#B3DE69", + "#FCCDE5", + "#D9D9D9", + ), + "Spectral_10": ( + "#9E0142", + "#D53E4F", + "#F46D43", + "#FDAE61", + "#FEE08B", + "#E6F598", + "#ABDDA4", + "#66C2A5", + "#3288BD", + "#5E4FA2", + ), + "Spectral_11": ( + "#9E0142", + "#D53E4F", + "#F46D43", + "#FDAE61", + "#FEE08B", + "#FFFFBF", + "#E6F598", + "#ABDDA4", + "#66C2A5", + "#3288BD", + "#5E4FA2", + ), + "Spectral_3": ("#FC8D59", "#FFFFBF", "#99D594"), + "Spectral_4": ("#D7191C", "#FDAE61", "#ABDDA4", "#2B83BA"), + "Spectral_5": ("#D7191C", "#FDAE61", "#FFFFBF", "#ABDDA4", "#2B83BA"), + "Spectral_6": ("#D53E4F", "#FC8D59", "#FEE08B", "#E6F598", "#99D594", "#3288BD"), + "Spectral_7": ( + "#D53E4F", + "#FC8D59", + "#FEE08B", + "#FFFFBF", + "#E6F598", + "#99D594", + "#3288BD", + ), + "Spectral_8": ( + "#D53E4F", + "#F46D43", + "#FDAE61", + "#FEE08B", + "#E6F598", + "#ABDDA4", + "#66C2A5", + "#3288BD", + ), + "Spectral_9": ( + "#D53E4F", + "#F46D43", + "#FDAE61", + "#FEE08B", + "#FFFFBF", + "#E6F598", + "#ABDDA4", + "#66C2A5", + "#3288BD", + ), + "YlGnBu_3": ("#EDF8B1", "#7FCDBB", "#2C7FB8"), + "YlGnBu_4": ("#FFFFCC", "#A1DAB4", "#41B6C4", "#225EA8"), + "YlGnBu_5": ("#FFFFCC", "#A1DAB4", "#41B6C4", "#2C7FB8", "#253494"), + "YlGnBu_6": ("#FFFFCC", "#C7E9B4", "#7FCDBB", "#41B6C4", "#2C7FB8", "#253494"), + "YlGnBu_7": ( + "#FFFFCC", + "#C7E9B4", + "#7FCDBB", + "#41B6C4", + "#1D91C0", + "#225EA8", + "#0C2C84", + ), + "YlGnBu_8": ( + "#FFFFD9", + "#EDF8B1", + "#C7E9B4", + "#7FCDBB", + "#41B6C4", + "#1D91C0", + "#225EA8", + "#0C2C84", + ), + "YlGnBu_9": ( + "#FFFFD9", + "#EDF8B1", + "#C7E9B4", + "#7FCDBB", + "#41B6C4", + "#1D91C0", + "#225EA8", + "#253494", + "#081D58", + ), + "YlGn_3": ("#F7FCB9", "#ADDD8E", "#31A354"), + "YlGn_4": ("#FFFFCC", "#C2E699", "#78C679", "#238443"), + "YlGn_5": ("#FFFFCC", "#C2E699", "#78C679", "#31A354", "#006837"), + "YlGn_6": ("#FFFFCC", "#D9F0A3", "#ADDD8E", "#78C679", "#31A354", "#006837"), + "YlGn_7": ( + "#FFFFCC", + "#D9F0A3", + "#ADDD8E", + "#78C679", + "#41AB5D", + "#238443", + "#005A32", + ), + "YlGn_8": ( + "#FFFFE5", + "#F7FCB9", + "#D9F0A3", + "#ADDD8E", + "#78C679", + "#41AB5D", + "#238443", + "#005A32", + ), + "YlGn_9": ( + "#FFFFE5", + "#F7FCB9", + "#D9F0A3", + "#ADDD8E", + "#78C679", + "#41AB5D", + "#238443", + "#006837", + "#004529", + ), + "YlOrBr_3": ("#FFF7BC", "#FEC44F", "#D95F0E"), + "YlOrBr_4": ("#FFFFD4", "#FED98E", "#FE9929", "#CC4C02"), + "YlOrBr_5": ("#FFFFD4", "#FED98E", "#FE9929", "#D95F0E", "#993404"), + "YlOrBr_6": ("#FFFFD4", "#FEE391", "#FEC44F", "#FE9929", "#D95F0E", "#993404"), + "YlOrBr_7": ( + "#FFFFD4", + "#FEE391", + "#FEC44F", + "#FE9929", + "#EC7014", + "#CC4C02", + "#8C2D04", + ), + "YlOrBr_8": ( + "#FFFFE5", + "#FFF7BC", + "#FEE391", + "#FEC44F", + "#FE9929", + "#EC7014", + "#CC4C02", + "#8C2D04", + ), + "YlOrBr_9": ( + "#FFFFE5", + "#FFF7BC", + "#FEE391", + "#FEC44F", + "#FE9929", + "#EC7014", + "#CC4C02", + "#993404", + "#662506", + ), + "YlOrRd_3": ("#FFEDA0", "#FEB24C", "#F03B20"), + "YlOrRd_4": ("#FFFFB2", "#FECC5C", "#FD8D3C", "#E31A1C"), + "YlOrRd_5": ("#FFFFB2", "#FECC5C", "#FD8D3C", "#F03B20", "#BD0026"), + "YlOrRd_6": ("#FFFFB2", "#FED976", "#FEB24C", "#FD8D3C", "#F03B20", "#BD0026"), + "YlOrRd_7": ( + "#FFFFB2", + "#FED976", + "#FEB24C", + "#FD8D3C", + "#FC4E2A", + "#E31A1C", + "#B10026", + ), + "YlOrRd_8": ( + "#FFFFCC", + "#FFEDA0", + "#FED976", + "#FEB24C", + "#FD8D3C", + "#FC4E2A", + "#E31A1C", + "#B10026", + ), + "YlOrRd_9": ( + "#FFFFCC", + "#FFEDA0", + "#FED976", + "#FEB24C", + "#FD8D3C", + "#FC4E2A", + "#E31A1C", + "#BD0026", + "#800026", + ), +} + +PALETTE_CATEGORIES = { + "Accent_3": "qualitative", + "Accent_4": "qualitative", + "Accent_5": "qualitative", + "Accent_6": "qualitative", + "Accent_7": "qualitative", + "Accent_8": "qualitative", + "Blues_3": "sequential", + "Blues_4": "sequential", + "Blues_5": "sequential", + "Blues_6": "sequential", + "Blues_7": "sequential", + "Blues_8": "sequential", + "Blues_9": "sequential", + "BrBG_10": "diverging", + "BrBG_11": "diverging", + "BrBG_3": "diverging", + "BrBG_4": "diverging", + "BrBG_5": "diverging", + "BrBG_6": "diverging", + "BrBG_7": "diverging", + "BrBG_8": "diverging", + "BrBG_9": "diverging", + "BuGn_3": "sequential", + "BuGn_4": "sequential", + "BuGn_5": "sequential", + "BuGn_6": "sequential", + "BuGn_7": "sequential", + "BuGn_8": "sequential", + "BuGn_9": "sequential", + "BuPu_3": "sequential", + "BuPu_4": "sequential", + "BuPu_5": "sequential", + "BuPu_6": "sequential", + "BuPu_7": "sequential", + "BuPu_8": "sequential", + "BuPu_9": "sequential", + "Dark2_3": "qualitative", + "Dark2_4": "qualitative", + "Dark2_5": "qualitative", + "Dark2_6": "qualitative", + "Dark2_7": "qualitative", + "Dark2_8": "qualitative", + "GnBu_3": "sequential", + "GnBu_4": "sequential", + "GnBu_5": "sequential", + "GnBu_6": "sequential", + "GnBu_7": "sequential", + "GnBu_8": "sequential", + "GnBu_9": "sequential", + "Greens_3": "sequential", + "Greens_4": "sequential", + "Greens_5": "sequential", + "Greens_6": "sequential", + "Greens_7": "sequential", + "Greens_8": "sequential", + "Greens_9": "sequential", + "Greys_3": "sequential", + "Greys_4": "sequential", + "Greys_5": "sequential", + "Greys_6": "sequential", + "Greys_7": "sequential", + "Greys_8": "sequential", + "Greys_9": "sequential", + "OrRd_3": "sequential", + "OrRd_4": "sequential", + "OrRd_5": "sequential", + "OrRd_6": "sequential", + "OrRd_7": "sequential", + "OrRd_8": "sequential", + "OrRd_9": "sequential", + "Oranges_3": "sequential", + "Oranges_4": "sequential", + "Oranges_5": "sequential", + "Oranges_6": "sequential", + "Oranges_7": "sequential", + "Oranges_8": "sequential", + "Oranges_9": "sequential", + "PRGn_10": "diverging", + "PRGn_11": "diverging", + "PRGn_3": "diverging", + "PRGn_4": "diverging", + "PRGn_5": "diverging", + "PRGn_6": "diverging", + "PRGn_7": "diverging", + "PRGn_8": "diverging", + "PRGn_9": "diverging", + "Paired_10": "qualitative", + "Paired_11": "qualitative", + "Paired_12": "qualitative", + "Paired_3": "qualitative", + "Paired_4": "qualitative", + "Paired_5": "qualitative", + "Paired_6": "qualitative", + "Paired_7": "qualitative", + "Paired_8": "qualitative", + "Paired_9": "qualitative", + "Pastel1_3": "qualitative", + "Pastel1_4": "qualitative", + "Pastel1_5": "qualitative", + "Pastel1_6": "qualitative", + "Pastel1_7": "qualitative", + "Pastel1_8": "qualitative", + "Pastel1_9": "qualitative", + "Pastel2_3": "qualitative", + "Pastel2_4": "qualitative", + "Pastel2_5": "qualitative", + "Pastel2_6": "qualitative", + "Pastel2_7": "qualitative", + "Pastel2_8": "qualitative", + "PiYG_10": "diverging", + "PiYG_11": "diverging", + "PiYG_3": "diverging", + "PiYG_4": "diverging", + "PiYG_5": "diverging", + "PiYG_6": "diverging", + "PiYG_7": "diverging", + "PiYG_8": "diverging", + "PiYG_9": "diverging", + "PuBuGn_3": "sequential", + "PuBuGn_4": "sequential", + "PuBuGn_5": "sequential", + "PuBuGn_6": "sequential", + "PuBuGn_7": "sequential", + "PuBuGn_8": "sequential", + "PuBuGn_9": "sequential", + "PuBu_3": "sequential", + "PuBu_4": "sequential", + "PuBu_5": "sequential", + "PuBu_6": "sequential", + "PuBu_7": "sequential", + "PuBu_8": "sequential", + "PuBu_9": "sequential", + "PuOr_10": "diverging", + "PuOr_11": "diverging", + "PuOr_3": "diverging", + "PuOr_4": "diverging", + "PuOr_5": "diverging", + "PuOr_6": "diverging", + "PuOr_7": "diverging", + "PuOr_8": "diverging", + "PuOr_9": "diverging", + "PuRd_3": "sequential", + "PuRd_4": "sequential", + "PuRd_5": "sequential", + "PuRd_6": "sequential", + "PuRd_7": "sequential", + "PuRd_8": "sequential", + "PuRd_9": "sequential", + "Purples_3": "sequential", + "Purples_4": "sequential", + "Purples_5": "sequential", + "Purples_6": "sequential", + "Purples_7": "sequential", + "Purples_8": "sequential", + "Purples_9": "sequential", + "RdBu_10": "diverging", + "RdBu_11": "diverging", + "RdBu_3": "diverging", + "RdBu_4": "diverging", + "RdBu_5": "diverging", + "RdBu_6": "diverging", + "RdBu_7": "diverging", + "RdBu_8": "diverging", + "RdBu_9": "diverging", + "RdGy_10": "diverging", + "RdGy_11": "diverging", + "RdGy_3": "diverging", + "RdGy_4": "diverging", + "RdGy_5": "diverging", + "RdGy_6": "diverging", + "RdGy_7": "diverging", + "RdGy_8": "diverging", + "RdGy_9": "diverging", + "RdPu_3": "sequential", + "RdPu_4": "sequential", + "RdPu_5": "sequential", + "RdPu_6": "sequential", + "RdPu_7": "sequential", + "RdPu_8": "sequential", + "RdPu_9": "sequential", + "RdYlBu_10": "diverging", + "RdYlBu_11": "diverging", + "RdYlBu_3": "diverging", + "RdYlBu_4": "diverging", + "RdYlBu_5": "diverging", + "RdYlBu_6": "diverging", + "RdYlBu_7": "diverging", + "RdYlBu_8": "diverging", + "RdYlBu_9": "diverging", + "RdYlGn_10": "diverging", + "RdYlGn_11": "diverging", + "RdYlGn_3": "diverging", + "RdYlGn_4": "diverging", + "RdYlGn_5": "diverging", + "RdYlGn_6": "diverging", + "RdYlGn_7": "diverging", + "RdYlGn_8": "diverging", + "RdYlGn_9": "diverging", + "Reds_3": "sequential", + "Reds_4": "sequential", + "Reds_5": "sequential", + "Reds_6": "sequential", + "Reds_7": "sequential", + "Reds_8": "sequential", + "Reds_9": "sequential", + "Set1_3": "qualitative", + "Set1_4": "qualitative", + "Set1_5": "qualitative", + "Set1_6": "qualitative", + "Set1_7": "qualitative", + "Set1_8": "qualitative", + "Set1_9": "qualitative", + "Set2_3": "qualitative", + "Set2_4": "qualitative", + "Set2_5": "qualitative", + "Set2_6": "qualitative", + "Set2_7": "qualitative", + "Set2_8": "qualitative", + "Set3_10": "qualitative", + "Set3_11": "qualitative", + "Set3_12": "qualitative", + "Set3_3": "qualitative", + "Set3_4": "qualitative", + "Set3_5": "qualitative", + "Set3_6": "qualitative", + "Set3_7": "qualitative", + "Set3_8": "qualitative", + "Set3_9": "qualitative", + "Spectral_10": "diverging", + "Spectral_11": "diverging", + "Spectral_3": "diverging", + "Spectral_4": "diverging", + "Spectral_5": "diverging", + "Spectral_6": "diverging", + "Spectral_7": "diverging", + "Spectral_8": "diverging", + "Spectral_9": "diverging", + "YlGnBu_3": "sequential", + "YlGnBu_4": "sequential", + "YlGnBu_5": "sequential", + "YlGnBu_6": "sequential", + "YlGnBu_7": "sequential", + "YlGnBu_8": "sequential", + "YlGnBu_9": "sequential", + "YlGn_3": "sequential", + "YlGn_4": "sequential", + "YlGn_5": "sequential", + "YlGn_6": "sequential", + "YlGn_7": "sequential", + "YlGn_8": "sequential", + "YlGn_9": "sequential", + "YlOrBr_3": "sequential", + "YlOrBr_4": "sequential", + "YlOrBr_5": "sequential", + "YlOrBr_6": "sequential", + "YlOrBr_7": "sequential", + "YlOrBr_8": "sequential", + "YlOrBr_9": "sequential", + "YlOrRd_3": "sequential", + "YlOrRd_4": "sequential", + "YlOrRd_5": "sequential", + "YlOrRd_6": "sequential", + "YlOrRd_7": "sequential", + "YlOrRd_8": "sequential", + "YlOrRd_9": "sequential", +} + + +def palette_colors(name): + """Return the exact discrete colors for a supported palette.""" + + try: + return list(PALETTES[name]) + except KeyError: + return list(PALETTES["Spectral_11"]) + + +def palette_category(name): + """Return the palette category name.""" + + return PALETTE_CATEGORIES.get(name, "diverging") diff --git a/src/moddotplot/estimate_identity.py b/src/moddotplot/estimate_identity.py index 004556e..6d2b9a2 100644 --- a/src/moddotplot/estimate_identity.py +++ b/src/moddotplot/estimate_identity.py @@ -13,6 +13,7 @@ from scipy.sparse import csr_matrix from moddotplot import _nthash +from moddotplot.color_palettes import palette_colors from moddotplot.optional_dependencies import OptionalDependencyError @@ -1143,32 +1144,26 @@ def findElementsWithPrefix(lst, prefix): def getInteractiveColor(palette_name, palette_orientation): - from palettable import colorbrewer - - palettes = colorbrewer.COLOR_MAPS - tmp_color = [] - new_palette = palette_name.split("_") + colors = palette_colors(palette_name) if palette_name in DIVERGING_PALETTES: - tmp_color = palettes["Diverging"][new_palette[0]][new_palette[1]]["Colors"] if palette_orientation == "+": palette_orientation = "-" else: palette_orientation = "+" - elif palette_name in SEQUENTIAL_PALETTES: - tmp_color = palettes["Sequential"][new_palette[0]][new_palette[1]]["Colors"] - elif palette_name in QUALITATIVE_PALETTES: - tmp_color = palettes["Qualitative"][new_palette[0]][new_palette[1]]["Colors"] - else: + elif palette_name not in SEQUENTIAL_PALETTES + QUALITATIVE_PALETTES: print("Unable to determine color palette. Selecting default \n") - tmp_color = palettes["Diverging"]["Spectral"]["11"]["Colors"] palette_orientation = "-" if palette_orientation == "-": - tmp_color = tmp_color[::-1] - tmp_color = [[255, 255, 255]] + tmp_color - total_values = len(tmp_color) + colors = colors[::-1] + colors = ["#FFFFFF", *colors] + total_values = len(colors) formatted_values = [ - [i / (total_values - 1), f"rgb({r}, {g}, {b})"] - for i, (r, g, b) in enumerate(tmp_color) + [ + i / (total_values - 1), + f"rgb({int(color[1:3], 16)}, {int(color[3:5], 16)}, " + f"{int(color[5:7], 16)})", + ] + for i, color in enumerate(colors) ] return formatted_values diff --git a/src/moddotplot/moddotplot.py b/src/moddotplot/moddotplot.py index 4341544..14adf04 100755 --- a/src/moddotplot/moddotplot.py +++ b/src/moddotplot/moddotplot.py @@ -434,7 +434,7 @@ def get_parser(): action="store_true", help=( "Output matrix to a Cooler file. Requires the optional " - "ModDotPlot[interactive] dependencies." + "ModDotPlot[cooler] dependencies." ), ) diff --git a/src/moddotplot/optional_dependencies.py b/src/moddotplot/optional_dependencies.py index 23bf138..fc67bdc 100644 --- a/src/moddotplot/optional_dependencies.py +++ b/src/moddotplot/optional_dependencies.py @@ -1,13 +1,19 @@ """Shared errors and installation guidance for optional features.""" -INTERACTIVE_INSTALL_COMMAND = 'python -m pip install "ModDotPlot[interactive]"' +INSTALL_COMMANDS = { + "Cooler export": 'python -m pip install "ModDotPlot[cooler]"', + "Interactive mode": 'python -m pip install "ModDotPlot[interactive]"', +} class OptionalDependencyError(ImportError): """Raised when a requested feature is missing its optional dependencies.""" def __init__(self, feature): + install_command = INSTALL_COMMANDS.get( + feature, 'python -m pip install "ModDotPlot[interactive]"' + ) super().__init__( - f"{feature} requires the optional interactive dependencies. " - f"Install them with: {INTERACTIVE_INSTALL_COMMAND}" + f"{feature} requires optional dependencies. " + f"Install them with: {install_command}" ) diff --git a/src/moddotplot/static_plots.py b/src/moddotplot/static_plots.py index 1684bc1..6143896 100755 --- a/src/moddotplot/static_plots.py +++ b/src/moddotplot/static_plots.py @@ -1,28 +1,3 @@ -from plotnine import ( - ggsave, - ggplot, - aes, - geom_histogram, - scale_color_discrete, - element_blank, - theme, - xlab, - scale_fill_manual, - scale_color_cmap, - coord_cartesian, - ylab, - scale_x_continuous, - scale_y_continuous, - geom_tile, - coord_fixed, - facet_grid, - labs, - element_line, - element_text, - theme_light, - geom_blank, - theme_minimal, -) import pandas as pd import numpy as np import math @@ -30,11 +5,12 @@ import re import matplotlib.pyplot as plt from matplotlib.colors import to_hex, to_rgb +from matplotlib.figure import Figure from matplotlib.patches import Rectangle from matplotlib.ticker import ScalarFormatter +from moddotplot.color_palettes import palette_category, palette_colors from moddotplot.native_render import ( DEFAULT_FONT_FAMILY, - FALLBACK_FONT_FAMILY, MIN_TEXT_SIZE, MIN_TITLE_SIZE, clamped_font_size, @@ -45,7 +21,6 @@ draw_triangle_tiles, genomic_scale, genomic_tick_formatter, - is_glyph_loading_error, save_figure_pair, save_with_font_fallback, set_figure_font_family, @@ -56,7 +31,6 @@ QUALITATIVE_PALETTES, SEQUENTIAL_PALETTES, ) -from palettable.colorbrewer import qualitative, sequential, diverging from moddotplot.annotations import ( DEFAULT_ANNOTATION_COLOR, annotation_color as _annotation_color, @@ -67,30 +41,6 @@ REGION_SUFFIX_PATTERN = re.compile(r"(?::\d+-\d+)+$") -def _plot_font_theme(family=DEFAULT_FONT_FAMILY): - """Apply one family to every Plotnine text themeable.""" - - font = element_text(family=[family]) - return theme( - text=font, - title=element_text(family=[family]), - axis_text=element_text(family=[family]), - strip_text=element_text(family=[family]), - legend_text=element_text(family=[family]), - ) - - -def _save_plot(plot, **kwargs): - """Save a Plotnine plot in Helvetica, retrying on glyph-load failure.""" - - try: - ggsave(plot + _plot_font_theme(), **kwargs) - except RuntimeError as error: - if not is_glyph_loading_error(error): - raise - ggsave(plot + _plot_font_theme(FALLBACK_FONT_FAMILY), **kwargs) - - def _draw_and_save_plot_pair( plot, output_prefix, @@ -100,37 +50,17 @@ def _draw_and_save_plot_pair( dpi, vector_format, ): - """Build one Plotnine figure and save both raster and vector outputs. - - ``ggsave`` redraws a plot for every requested format. Large tile plots and - histograms therefore paid their complete scale/layout/rasterization cost - twice. Drawing once also guarantees that both files contain the same axes, - labels, and tile realization. - """ + """Save one Matplotlib figure to raster and vector outputs.""" - def draw(family): - styled = ( - plot - + _plot_font_theme(family) - + theme(figure_size=(float(width), float(height)), dpi=int(dpi)) - ) - return styled.draw(show=False) - - try: - figure = draw(DEFAULT_FONT_FAMILY) - except RuntimeError as error: - if not is_glyph_loading_error(error): - raise - figure = draw(FALLBACK_FONT_FAMILY) + figure = plot if isinstance(plot, Figure) else plot.draw(show=False) + figure.set_size_inches(float(width), float(height), forward=True) + figure.set_dpi(int(dpi)) try: return save_figure_pair( figure, output_prefix, vector_format, dpi, - # Match plotnine/ggsave's requested physical canvas exactly. A - # tight bounding box changes both the raster dimensions and plot - # framing (for example, 3 in at 96 dpi no longer yields 288 px). bbox_inches=figure.bbox_inches, ) finally: @@ -185,20 +115,21 @@ def _fit_grid_sequence_labels(figure, axes): def _resolve_native_colors(palette, palette_orientation, custom_colors=None): - """Resolve plot colors with the same orientation rules as plotnine paths.""" - if palette in DIVERGING_PALETTES: - palette_colors = getattr(diverging, palette).hex_colors + """Resolve plot colors while preserving historical orientation rules.""" + + colors = palette_colors(palette) + supported = ( + palette in DIVERGING_PALETTES + or palette in QUALITATIVE_PALETTES + or palette in SEQUENTIAL_PALETTES + ) + if palette_category(palette) == "diverging": palette_orientation = "-" if palette_orientation == "+" else "+" - elif palette in QUALITATIVE_PALETTES: - palette_colors = getattr(qualitative, palette).hex_colors - elif palette in SEQUENTIAL_PALETTES: - palette_colors = getattr(sequential, palette).hex_colors - else: - palette_colors = diverging.Spectral_11.hex_colors + if not supported: palette_orientation = "-" - colors = palette_colors[::-1] if palette_orientation == "-" else palette_colors - return list(custom_colors) if custom_colors else list(colors) + oriented = colors[::-1] if palette_orientation == "-" else colors + return list(custom_colors) if custom_colors else oriented DIRECTION_ANI_COLUMN = "direction_ani" @@ -244,11 +175,6 @@ def _direction_ani_style(dataframe): return styled, colors, DIRECTION_ANI_COLUMN -def is_plot_empty(p): - # Check if the plot has data or any layers - return len(p.layers) == 0 and p.data.empty - - def draw_annotation_track( axis, bed_df, @@ -373,21 +299,6 @@ def make_scale(vals: list) -> list: return make_m(scaled) -def _dotplot_tiles(mapping, deraster=False, **kwargs): - """Create tiles without materializing a genomic-coordinate-sized image. - - ``plotnine.geom_raster`` expands sparse coordinates into an RGBA array whose - dimensions are derived from the smallest coordinate spacing. A 496 Mb - sequence plotted in 2 kb windows can therefore request roughly - 248,000-by-248,000 pixels even when only a small fraction of those cells - contain matches. ``geom_tile`` draws only the cells present in the input - dataframe. Setting ``raster=True`` keeps the default compact, rasterized - appearance in vector output, while ``--deraster`` leaves the tiles as - vectors. - """ - return geom_tile(mapping, raster=not deraster, **kwargs) - - def get_colors(sdf, ncolors, is_freq, custom_breakpoints): if ncolors < 1: raise ValueError("At least one color is required") @@ -500,36 +411,15 @@ def read_df( else: data = pj[0] df = pd.DataFrame(data[1:], columns=data[0]) - hexcodes = [] - new_hexcodes = [] - if palette in DIVERGING_PALETTES: - function_name = getattr(diverging, palette) - hexcodes = function_name.hex_colors - if palette_orientation == "+": - palette_orientation = "-" - else: - palette_orientation = "+" - elif palette in QUALITATIVE_PALETTES: - function_name = getattr(qualitative, palette) - hexcodes = function_name.hex_colors - elif palette in SEQUENTIAL_PALETTES: - function_name = getattr(sequential, palette) - hexcodes = function_name.hex_colors - else: + supported = ( + palette in DIVERGING_PALETTES + or palette in QUALITATIVE_PALETTES + or palette in SEQUENTIAL_PALETTES + ) + if not supported: print(f"Palette {palette} not found. Defaulting to Spectral_11.\n") - function_name = getattr(diverging, "Spectral_11") - palette_orientation = "-" - hexcodes = function_name.hex_colors - - if palette_orientation == "-": - new_hexcodes = hexcodes[::-1] - else: - new_hexcodes = hexcodes - - if custom_colors: - new_hexcodes = custom_colors - - ncolors = len(new_hexcodes) + colors = _resolve_native_colors(palette, palette_orientation, custom_colors) + ncolors = len(colors) # Get colors for each row based on the values in the dataframe df["discrete"] = get_colors(df, ncolors, is_freq, custom_breakpoints) # Rename columns if they have different names in the dataframe @@ -631,124 +521,23 @@ def make_dot( width, is_pairwise, ): - display_x = display_sequence_name(name_x) - display_y = display_sequence_name(name_y) - if is_pairwise: - title_name = f"Comparative Plot: {display_x} vs {display_y}" - else: - title_name = f"Self-Identity Plot: {display_x}" - title_length = 2 * width - if len(title_name) > 50: - title_length = 1.5 * width - elif len(title_name) > 80: - title_length = width - sdf, direction_colors, direction_column = _direction_ani_style(sdf) - direction_coloring = direction_colors is not None - # Select the color palette - if hasattr(diverging, palette): - function_name = getattr(diverging, palette) - elif hasattr(qualitative, palette): - function_name = getattr(qualitative, palette) - elif hasattr(sequential, palette): - function_name = getattr(sequential, palette) - else: - function_name = diverging.Spectral_11 # Default palette - palette_orientation = "-" + """Build a full dotplot with the native Matplotlib renderer.""" - hexcodes = function_name.hex_colors - - # Adjust palette orientation - if palette in diverging.__dict__: - palette_orientation = "-" if palette_orientation == "+" else "+" - - new_hexcodes = hexcodes[::-1] if palette_orientation == "-" else hexcodes - if colors: - new_hexcodes = colors # Override colors if provided - fill_column = direction_column if direction_coloring else "discrete" - fill_colors = direction_colors if direction_coloring else new_hexcodes - # Determine the exact genomic interval. A two-value limit is supplied by - # FASTA mode so blank edge windows do not shrink or extend the plot. - min_val, max_val = _data_axis_limits(sdf, xlim) - - # If user provides breaks, convert to ints - if not breaks: - breaks = generate_breaks(int(min_val), int(max_val)) - else: - breaks = [int(x) for x in breaks] - # Compute window size (handling exceptions) - try: - window = max(sdf["q_en"] - sdf["q_st"]) - except ValueError: # Empty dataframe case - return ggplot(aes(x=[], y=[])) + theme_minimal() - - # Region-qualified names remain in BEDPE data and filenames, but plot - # headings should show only the underlying FASTA identifier. - sdf = sdf.copy() - sdf["q"] = sdf["q"].map(display_sequence_name) - sdf["r"] = sdf["r"].map(display_sequence_name) - - # Determine axis label scale based on genomic position size - if max_val < 200_000: - x_label = "Genomic Position (Kbp)" - elif max_val < 200_000_000: - x_label = "Genomic Position (Mbp)" - else: - x_label = "Genomic Position (Gbp)" - - # Create the plot - common_theme = theme( - legend_position="none", - panel_grid_major=element_blank(), - panel_grid_minor=element_blank(), - plot_background=element_blank(), - panel_background=element_blank(), - axis_line=element_line(color="black"), - axis_text=element_text( - family=[DEFAULT_FONT_FAMILY], - size=clamped_font_size(width, 2.0), - ), - axis_ticks_major=element_line( - size=(width), color="black" - ), # Increased tick length - title=element_text( - family=[DEFAULT_FONT_FAMILY], - size=max(MIN_TITLE_SIZE, title_length), - hjust=0.5, - ), # Center title - axis_title_x=element_text( - size=clamped_font_size(width, 2.8), - family=[DEFAULT_FONT_FAMILY], - ), - strip_background=element_blank(), # Remove facet strip background - strip_text=element_text( - size=clamped_font_size(width, 1.2), family=[DEFAULT_FONT_FAMILY] - ), # Customize facet label text size (optional) - ) - - # Construct the plot arguments - ggplot_args = ( - ggplot(sdf) - + scale_color_discrete(guide=None) - + scale_fill_manual(values=fill_colors, guide=None) - + common_theme - + scale_x_continuous( - labels=make_scale, limits=[min_val, max_val], breaks=breaks - ) - + scale_y_continuous( - labels=make_scale, limits=[min_val, max_val], breaks=breaks - ) - + coord_fixed(ratio=1) - + facet_grid("r ~ q") - + labs(x=x_label, y="", title=title_name) - ) - - p = ggplot_args + _dotplot_tiles( - aes(x="q_st", y="r_st", fill=fill_column, height=window, width=window), - deraster, + del num_ticks + return _build_full_figure( + sdf=sdf, + name_x=name_x, + name_y=name_y, + palette=palette, + palette_orientation=palette_orientation, + custom_colors=colors, + axes_labels=breaks, + xlim=xlim, + deraster=deraster, + width=width, + is_pairwise=is_pairwise, ) - return p - def make_dot_grid( sdf, @@ -762,101 +551,22 @@ def make_dot_grid( deraster, width, ): - title_name = display_sequence_name(title_name) - # Select the color palette - if hasattr(diverging, palette): - function_name = getattr(diverging, palette) - elif hasattr(qualitative, palette): - function_name = getattr(qualitative, palette) - elif hasattr(sequential, palette): - function_name = getattr(sequential, palette) - else: - function_name = diverging.Spectral_11 # Default palette - palette_orientation = "-" - - hexcodes = function_name.hex_colors - - # Adjust palette orientation - if palette in diverging.__dict__: - palette_orientation = "-" if palette_orientation == "+" else "+" + """Build a standalone grid cell with the native Matplotlib renderer.""" - new_hexcodes = hexcodes[::-1] if palette_orientation == "-" else hexcodes - if colors: - new_hexcodes = colors # Override colors if provided - if not xlim: - xlim = 0 - # Determine maximum genomic position for scaling - min_val = max(sdf["q_st"].min(), sdf["r_st"].min()) - max_val = max(sdf["q_en"].max(), sdf["r_en"].max(), xlim) - - # If user provides breaks, convert to ints - if not breaks: - breaks = generate_breaks(int(min_val), int(max_val)) - else: - breaks = [int(x) for x in breaks] - xlim = xlim or 0 - # Compute window size (handling exceptions) - try: - window = max(sdf["q_en"] - sdf["q_st"]) - except ValueError: # Empty dataframe case - return ggplot(aes(x=[], y=[])) + theme_minimal() - - # Determine axis label scale based on genomic position size - if max_val < 200_000: - x_label = "Genomic Position (Kbp)" - elif max_val < 200_000_000: - x_label = "Genomic Position (Mbp)" - else: - x_label = "Genomic Position (Gbp)" - - # Create the plot - common_theme = theme( - legend_position="none", - panel_grid_major=element_blank(), - panel_grid_minor=element_blank(), - plot_background=element_blank(), - panel_background=element_blank(), - axis_line=element_line(color="black"), - axis_text=element_text( - family=[DEFAULT_FONT_FAMILY], size=clamped_font_size(width, 1.0) - ), - axis_ticks_major=element_line( - size=(width), color="black" - ), # Increased tick length - title=element_text( - size=clamped_font_size(width, 1.2, MIN_TITLE_SIZE), - family=[DEFAULT_FONT_FAMILY], - alpha=0, - ), - axis_title_x=element_text( - size=clamped_font_size(width, 1.2), - family=[DEFAULT_FONT_FAMILY], - ), - strip_background=element_blank(), # Remove facet strip background - strip_text=element_text( - size=clamped_font_size(width, 1.2), family=[DEFAULT_FONT_FAMILY] - ), # Customize facet label text size (optional) - ) - - # Construct the plot arguments - ggplot_args = ( - ggplot(sdf) - + scale_color_discrete(guide=None) - + scale_fill_manual(values=new_hexcodes, guide=None) - + common_theme - + scale_x_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) - + scale_y_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) - + coord_fixed(ratio=1) - + labs(x="", y="", title="") - ) - - p = ggplot_args + _dotplot_tiles( - aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), - deraster, + return _build_full_figure( + sdf=sdf, + name_x=title_name, + name_y=title_name, + palette=palette, + palette_orientation=palette_orientation, + custom_colors=colors, + axes_labels=breaks, + xlim=xlim, + deraster=deraster, + width=width, + is_pairwise=not on_diagonal, ) - return p - def direction_dataframe( canonical_matrix, @@ -868,12 +578,8 @@ def direction_dataframe( x_offset=0, y_offset=0, ): - """Build plotting records that distinguish forward and reverse matches. + """Build plotting records that distinguish forward and reverse matches.""" - Canonical k-mers match in either orientation, while forward-only k-mers - match only same-strand sequence. A canonical hit missing from the - forward-only matrix therefore represents a reverse-orientation match. - """ canonical_matrix = np.asarray(canonical_matrix, dtype=float) forward_matrix = np.asarray(forward_matrix, dtype=float) if canonical_matrix.shape != forward_matrix.shape: @@ -927,6 +633,7 @@ def create_direction_plot( y_offset=0, ): """Save a blue/pink plot showing match orientation.""" + dataframe = direction_dataframe( canonical_matrix, forward_matrix, @@ -941,10 +648,6 @@ def create_direction_plot( print(f"No directional matches found for {name_x} and {name_y}. Skipping.\n") return None - dataframe = dataframe.assign( - q_position=dataframe["q_st"] + window_size / 2, - r_position=dataframe["r_st"] + window_size / 2, - ) requested_bounds = _requested_axis_bounds(xlim) if requested_bounds is not None: min_val, max_val = requested_bounds @@ -966,81 +669,59 @@ def create_direction_plot( else f"Direction Plot: {display_sequence_name(name_x)} vs " f"{display_sequence_name(name_y)}" ) - plot = ( - ggplot(dataframe) - + _dotplot_tiles( - aes( - x="q_position", - y="r_position", - fill="direction", - height=window_size, - width=window_size, - ), - deraster, + + render_data = dataframe.copy() + render_data["q_en"] = render_data["q_st"] + window_size + render_data["r_en"] = render_data["r_st"] + window_size + figure, axis = plt.subplots(figsize=(float(width), float(width))) + try: + draw_rectangular_tiles( + axis, + render_data, + DIRECTION_COLORS, + color_column="direction", + rasterized=not deraster, ) - + scale_fill_manual( - values={"Forward": "#2166AC", "Reverse": "#D01C8B"}, - name="Direction", + configure_dotplot_axis(axis, min_val, max_val, breaks=breaks) + axis.set_xlabel( + "Genomic Position", + fontsize=clamped_font_size(width, 1.2), + fontfamily=DEFAULT_FONT_FAMILY, ) - + scale_x_continuous( - labels=make_scale, limits=[min_val, max_val], breaks=breaks + axis.set_title( + title, + fontsize=clamped_font_size(width, 1.4, MIN_TITLE_SIZE), + fontfamily=DEFAULT_FONT_FAMILY, ) - + scale_y_continuous( - labels=make_scale, limits=[min_val, max_val], breaks=breaks + axis.grid(False) + figure.text( + 0.5, + 0.02, + "Blue: forward Pink: reverse", + ha="center", + fontfamily=DEFAULT_FONT_FAMILY, + fontsize=clamped_font_size(width, 0.9), ) - + coord_fixed(ratio=1) - + labs( - x="Genomic Position", - y="", - title=title, - caption="Blue: forward Pink: reverse", + figure.subplots_adjust(left=0.16, right=0.96, bottom=0.18, top=0.88) + set_figure_font_family(figure, DEFAULT_FONT_FAMILY) + + os.makedirs(directory, exist_ok=True) + filename = ( + f"{name_x}_DIRECTION" if self_identity else f"{name_x}_{name_y}_DIRECTION" ) - + theme_light() - + theme( - legend_position="none", - panel_grid_major=element_blank(), - panel_grid_minor=element_blank(), - axis_text=element_text( - family=[DEFAULT_FONT_FAMILY], - size=clamped_font_size(width, 1.0), - ), - title=element_text( - family=[DEFAULT_FONT_FAMILY], - size=clamped_font_size(width, 1.4, MIN_TITLE_SIZE), - hjust=0.5, - ), - axis_title_x=element_text( - size=clamped_font_size(width, 1.2), - family=[DEFAULT_FONT_FAMILY], - ), + prefix = os.path.join(directory, filename) + save_figure_pair( + figure, + prefix, + vector_format, + dpi, + bbox_inches=figure.bbox_inches, ) - ) + finally: + plt.close(figure) - os.makedirs(directory, exist_ok=True) - filename = ( - f"{name_x}_DIRECTION" if self_identity else f"{name_x}_{name_y}_DIRECTION" - ) - prefix = os.path.join(directory, filename) - _save_plot( - plot, - width=width, - height=width, - dpi=dpi, - format=vector_format, - filename=f"{prefix}.{vector_format}", - verbose=False, - ) - _save_plot( - plot, - width=width, - height=width, - dpi=dpi, - format="png", - filename=f"{prefix}.png", - verbose=False, - ) print(f"Direction plots saved to {prefix}.png and {prefix}.{vector_format}.\n") - return plot + return figure def make_dot_final( @@ -1054,143 +735,49 @@ def make_dot_final( transpose=False, deraster=False, ): - if sdf.empty: - max_val = xlim or 1 - if not breaks: - breaks = generate_breaks(0, int(max_val)) - else: - breaks = [int(x) for x in breaks] - return ( - ggplot(sdf) - + geom_blank() - + scale_x_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) - + scale_y_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) - + coord_fixed(ratio=1) - + labs(x=None, y=None, title=None) - + theme( - panel_grid_major=element_blank(), - panel_grid_minor=element_blank(), - plot_background=element_blank(), - panel_background=element_blank(), - axis_line=element_line(color="black"), - axis_text=element_text( - family=[DEFAULT_FONT_FAMILY], - size=clamped_font_size(width, 1.0), - ), - axis_ticks_major=element_line(), - axis_title_x=element_blank(), - axis_title_y=element_blank(), - ) - ) - - if hasattr(diverging, palette): - function_name = getattr(diverging, palette) - elif hasattr(qualitative, palette): - function_name = getattr(qualitative, palette) - elif hasattr(sequential, palette): - function_name = getattr(sequential, palette) - else: - function_name = diverging.Spectral_11 # Default palette - palette_orientation = "-" - - hexcodes = function_name.hex_colors - - # Adjust palette orientation - if palette in diverging.__dict__: - palette_orientation = "-" if palette_orientation == "+" else "+" - - new_hexcodes = hexcodes[::-1] if palette_orientation == "-" else hexcodes - if colors: - new_hexcodes = colors # Override colors if provided - if not xlim: - xlim = 0 - # Determine maximum genomic position for scaling - min_val = min(sdf["q_st"].min(), sdf["r_st"].min()) - max_val = max(sdf["q_en"].max(), sdf["r_en"].max(), xlim) - - # If user provides breaks, convert to ints - if not breaks: - breaks = generate_breaks(int(min_val), int(max_val)) - else: - breaks = [int(x) for x in breaks] - xlim = xlim or 0 - - max_val = max(sdf["q_en"].max(), sdf["r_en"].max(), xlim) - try: - window = max(sdf["q_en"] - sdf["q_st"]) - except: - p = ( - ggplot(aes(x=[], y=[])) - + theme_minimal() - + theme( - panel_grid_major=element_blank(), - panel_grid_minor=element_blank(), - ) - ) - return p - - x_col, y_col = ("r_st", "q_st") if transpose else ("q_st", "r_st") + """Build a single grid-style dotplot panel.""" - if deraster: - p = ( - ggplot(sdf) - + _dotplot_tiles( - aes(x=x_col, y=y_col, fill="discrete", height=window, width=window), - deraster, - ) - + scale_color_discrete(guide=None) - + scale_fill_manual(values=new_hexcodes, guide=None) - + theme( - legend_position="none", - panel_grid_major=element_blank(), - panel_grid_minor=element_blank(), - plot_background=element_blank(), - panel_background=element_blank(), - axis_line=element_line(color="black"), - axis_text=element_text( - family=[DEFAULT_FONT_FAMILY], - size=clamped_font_size(width, 1.0), - ), - axis_ticks_major=element_line(), - title=element_text(family=[DEFAULT_FONT_FAMILY]), - ) - + scale_x_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) - + scale_y_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) - + coord_fixed(ratio=1) - + labs(x=None, y=None, title=None) - ) - else: - p = ( - ggplot(sdf) - + _dotplot_tiles( - aes(x=x_col, y=y_col, fill="discrete", height=window, width=window), - deraster, - ) - + scale_color_discrete(guide=None) - + scale_fill_manual(values=new_hexcodes, guide=None) - + theme( - legend_position="none", - panel_grid_major=element_blank(), - panel_grid_minor=element_blank(), - plot_background=element_blank(), - panel_background=element_blank(), - axis_line=element_line(color="black"), - axis_text=element_text( - family=[DEFAULT_FONT_FAMILY], - size=clamped_font_size(width, 1.0), - ), - axis_ticks_major=element_line(), - title=element_text(family=[DEFAULT_FONT_FAMILY]), - ) - + scale_x_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) - + scale_y_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) - + coord_fixed(ratio=1) - + labs(x=None, y=None, title=None) + dataframe = sdf.copy() + if transpose: + dataframe = dataframe.rename( + columns={ + "q": "_r", + "q_st": "_r_st", + "q_en": "_r_en", + "r": "q", + "r_st": "q_st", + "r_en": "q_en", + } + ).rename( + columns={ + "_r": "r", + "_r_st": "r_st", + "_r_en": "r_en", + } ) - - p += theme(axis_title_x=element_blank(), axis_title_y=element_blank()) - - return p + name_x = ( + display_sequence_name(dataframe["q"].iloc[0]) + if not dataframe.empty and "q" in dataframe + else "" + ) + name_y = ( + display_sequence_name(dataframe["r"].iloc[0]) + if not dataframe.empty and "r" in dataframe + else name_x + ) + return _build_full_figure( + sdf=dataframe, + name_x=name_x, + name_y=name_y, + palette=palette, + palette_orientation=palette_orientation, + custom_colors=colors, + axes_labels=breaks, + xlim=xlim, + deraster=deraster, + width=width, + is_pairwise=True, + ) def make_tri( @@ -1205,356 +792,79 @@ def make_tri( deraster, width, ): - title_name = display_sequence_name(title_name) - # Select the color palette - if hasattr(diverging, palette): - function_name = getattr(diverging, palette) - elif hasattr(qualitative, palette): - function_name = getattr(qualitative, palette) - elif hasattr(sequential, palette): - function_name = getattr(sequential, palette) - else: - function_name = diverging.Spectral_11 # Default palette - palette_orientation = "-" + """Build a native Matplotlib triangle plot and return its main axis.""" - hexcodes = function_name.hex_colors - - # Adjust palette orientation - if palette in diverging.__dict__: - palette_orientation = "-" if palette_orientation == "+" else "+" - - new_hexcodes = hexcodes[::-1] if palette_orientation == "-" else hexcodes - if colors: - new_hexcodes = colors # Override colors if provided - if not xlim: - xlim = 0 - # Determine maximum genomic position for scaling - min_val = max(sdf["q_st"].min(), sdf["r_st"].min()) - max_val = max(sdf["q_en"].max(), sdf["r_en"].max(), xlim) - - # If user provides breaks, convert to ints - if not breaks: - breaks = generate_breaks(int(min_val), int(max_val)) - else: - breaks = [int(x) for x in breaks] - xlim = xlim or 0 - # Compute window size (handling exceptions) - try: - window = max(sdf["q_en"] - sdf["q_st"]) - except ValueError: # Empty dataframe case - return ggplot(aes(x=[], y=[])) + theme_minimal() - - # Determine axis label scale based on genomic position size - if max_val < 200_000: - x_label = "Genomic Position (Kbp)" - elif max_val < 200_000_000: - x_label = "Genomic Position (Mbp)" - else: - x_label = "Genomic Position (Gbp)" - - if not deraster: - tri = ( - ggplot(sdf) - + _dotplot_tiles( - aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), - deraster, - alpha=1.0, - ) # Ensure full opacity - + scale_fill_manual(values=new_hexcodes, guide=None) - + scale_color_discrete(guide=None) - + scale_x_continuous( - labels=make_scale, limits=[min_val, max_val], breaks=breaks - ) - + scale_y_continuous( - labels=make_scale, limits=[min_val, max_val], breaks=breaks - ) - + coord_fixed(ratio=1) - + labs(x=x_label, y="", title=title_name) - + theme( - legend_position="none", - panel_grid_major=element_blank(), - panel_grid_minor=element_blank(), - plot_background=element_blank(), - panel_background=element_blank(), - axis_text=element_text( - family=[DEFAULT_FONT_FAMILY], - size=clamped_font_size(width, 1.0), - ), - axis_line_x=element_line(), - axis_line_y=element_blank(), - axis_ticks_major_x=element_line(), - axis_ticks_major_y=element_blank(), - axis_ticks_major=element_line(size=(width)), - title=element_text( - family=[DEFAULT_FONT_FAMILY], - size=clamped_font_size(width, 1.4, MIN_TITLE_SIZE), - hjust=0.5, - ), - axis_title_x=element_text( - size=clamped_font_size(width, 1.4), - family=[DEFAULT_FONT_FAMILY], - ), - axis_text_y=element_blank(), - ) - ) - axis = ( - ggplot(sdf) - + geom_tile( - aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), - alpha=0, - ) - + scale_color_discrete(guide=None) - + scale_fill_manual(values=new_hexcodes, guide=None) - + scale_x_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) - + scale_y_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) - + coord_fixed(ratio=1) - + labs(x="", y="", title=title_name) - + theme( - legend_position="none", - panel_grid_major=element_blank(), - panel_grid_minor=element_blank(), - plot_background=element_blank(), - panel_background=element_blank(), - axis_line=element_line(color="black"), - axis_text=element_text( - family=[DEFAULT_FONT_FAMILY], - size=clamped_font_size(width, 1.0), - ), - axis_ticks_major=element_line(), - axis_line_x=element_line(), - axis_line_y=element_blank(), - axis_ticks_major_x=element_line(), - axis_ticks_major_y=element_blank(), - axis_text_x=element_line(), - axis_text_y=element_blank(), - plot_title=element_blank(), - axis_title_x=element_text( - size=clamped_font_size(width, 1.2), - family=[DEFAULT_FONT_FAMILY], - ), - ) - ) - else: - tri = ( - ggplot(sdf) - + _dotplot_tiles( - aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), - deraster, - alpha=1.0, - ) # Ensure full opacity - + scale_fill_manual(values=new_hexcodes, guide=None) - + scale_color_discrete(guide=None) - + scale_x_continuous( - labels=make_scale, limits=[min_val, max_val], breaks=breaks - ) - + scale_y_continuous( - labels=make_scale, limits=[min_val, max_val], breaks=breaks - ) - + coord_fixed(ratio=1) - + labs(x=x_label, y="", title=title_name) - + theme( - legend_position="none", - panel_grid_major=element_blank(), - panel_grid_minor=element_blank(), - plot_background=element_blank(), - panel_background=element_blank(), - axis_text=element_text( - family=[DEFAULT_FONT_FAMILY], - size=clamped_font_size(width, 1.0), - ), - axis_line_x=element_line(), - axis_line_y=element_blank(), - axis_ticks_major_x=element_line(), - axis_ticks_major_y=element_blank(), - axis_ticks_major=element_line(), - axis_text_y=element_blank(), - title=element_blank(), - axis_title_x=element_text( - size=clamped_font_size(width, 1.2), - family=[DEFAULT_FONT_FAMILY], - ), - ) - ) - axis = ( - ggplot(sdf) - + geom_tile( - aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), - alpha=0, - ) - + scale_color_discrete(guide=None) - + scale_fill_manual(values=new_hexcodes, guide=None) - + scale_x_continuous( - labels=make_scale, limits=[min_val, max_val], breaks=breaks - ) - + scale_y_continuous( - labels=make_scale, limits=[min_val, max_val], breaks=breaks - ) - + coord_fixed(ratio=1) - + labs(x="", y="", title="") - + theme( - legend_position="none", - panel_grid_major=element_blank(), - panel_grid_minor=element_blank(), - plot_background=element_blank(), - panel_background=element_blank(), - axis_line=element_line(color="black"), - axis_text=element_text(family=[DEFAULT_FONT_FAMILY]), - axis_ticks_major=element_line(), - axis_line_x=element_line(), - axis_line_y=element_blank(), - axis_ticks_major_x=element_line(), - axis_ticks_major_y=element_blank(), - axis_text_x=element_line(), - axis_text_y=element_blank(), - plot_title=element_blank(), - axis_title_x=element_text( - size=clamped_font_size(width, 1.2), - family=[DEFAULT_FONT_FAMILY], - ), - ) - ) - - return tri, axis - - -def make_tri_axis(sdf, title_name, palette, palette_orientation, colors, breaks, xlim): - title_name = display_sequence_name(title_name) - if not breaks: - breaks = True - else: - breaks = [float(number) for number in breaks] - if not xlim: - xlim = 0 - hexcodes = [] - new_hexcodes = [] - if palette in DIVERGING_PALETTES: - function_name = getattr(diverging, palette) - hexcodes = function_name.hex_colors - if palette_orientation == "+": - palette_orientation = "-" - else: - palette_orientation = "+" - elif palette in QUALITATIVE_PALETTES: - function_name = getattr(qualitative, palette) - hexcodes = function_name.hex_colors - elif palette in SEQUENTIAL_PALETTES: - function_name = getattr(sequential, palette) - hexcodes = function_name.hex_colors - else: - function_name = getattr(sequential, "Spectral_11") - palette_orientation = "-" - hexcodes = function_name.hex_colors - - if palette_orientation == "-": - new_hexcodes = hexcodes[::-1] - else: - new_hexcodes = hexcodes - if colors: - new_hexcodes = colors - max_val = max(sdf["q_en"].max(), sdf["r_en"].max(), xlim) - window = max(sdf["q_en"] - sdf["q_st"]) - if max_val < 100000: - x_label = "Genomic Position (Kbp)" - elif max_val < 100000000: - x_label = "Genomic Position (Mbp)" - else: - x_label = "Genomic Position (Gbp)" - p = ( - ggplot(sdf) - + geom_tile( - aes(x="q_st", y="r_st", fill="discrete", height=window, width=window), - alpha=0, - ) - + scale_color_discrete(guide=None) - + scale_fill_manual( - values=new_hexcodes, - guide=None, - ) - + theme( - legend_position="none", - panel_grid_major=element_blank(), - panel_grid_minor=element_blank(), - plot_background=element_blank(), - panel_background=element_blank(), - axis_line=element_line(color="black"), # Adjust axis line size - axis_text=element_text( - family=[DEFAULT_FONT_FAMILY] - ), # Change axis text font and size - axis_ticks_major=element_line(), - axis_line_x=element_line(), # Keep the x-axis line - axis_line_y=element_blank(), # Remove the y-axis line - axis_ticks_major_x=element_line(), # Keep x-axis ticks - axis_ticks_major_y=element_blank(), # Remove y-axis ticks - axis_text_x=element_line(), # Keep x-axis text - axis_text_y=element_blank(), - plot_title=element_blank(), - ) - + scale_x_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) - + scale_y_continuous(labels=make_scale, limits=[0, max_val], breaks=breaks) - + coord_fixed(ratio=1) - + labs(x="", y="", title=title_name) + del num_ticks + figure = _build_triangle_figure( + sdf=sdf, + title=title_name, + palette=palette, + palette_orientation=palette_orientation, + custom_colors=colors, + axes_labels=breaks, + xlim=xlim, + deraster=deraster, + width=width, ) - - # Adjust x-axis label size - p += theme(axis_title_x=element_text()) - - return p + return figure, figure.axes[0] def make_hist(sdf, palette, palette_orientation, custom_colors, custom_breakpoints): - hexcodes = [] - new_hexcodes = [] - if palette in DIVERGING_PALETTES: - function_name = getattr(diverging, palette) - hexcodes = function_name.hex_colors - if palette_orientation == "+": - palette_orientation = "-" - else: - palette_orientation = "+" - elif palette in QUALITATIVE_PALETTES: - function_name = getattr(qualitative, palette) - hexcodes = function_name.hex_colors - elif palette in SEQUENTIAL_PALETTES: - function_name = getattr(sequential, palette) - hexcodes = function_name.hex_colors - else: - function_name = getattr(diverging, "Spectral_11") - palette_orientation = "-" - hexcodes = function_name.hex_colors - - if palette_orientation == "-": - new_hexcodes = hexcodes[::-1] - else: - new_hexcodes = hexcodes + """Build the identity histogram with native Matplotlib.""" - if custom_colors: - new_hexcodes = custom_colors + del custom_breakpoints + colors = _resolve_native_colors(palette, palette_orientation, custom_colors) try: - bot = np.quantile(sdf["perID_by_events"], q=0.001) - except IndexError: - bot = 0 - count = sdf.shape[0] - extra = "" - - if count > 1e6: - extra = "\n(thousands)" + lower_bound = float(np.quantile(sdf["perID_by_events"], q=0.001)) + except (IndexError, ValueError): + lower_bound = 0.0 - sdf, direction_colors, direction_column = _direction_ani_style(sdf) + dataframe, direction_colors, direction_column = _direction_ani_style(sdf) fill_column = direction_column or "discrete" - fill_colors = direction_colors or new_hexcodes - p = ( - ggplot(data=sdf, mapping=aes(x="perID_by_events", fill=fill_column)) - + geom_histogram(bins=300) - + scale_color_cmap(cmap_name="plasma") - + scale_fill_manual(fill_colors) - + theme_light() - + _plot_font_theme() - + theme(legend_position="none") - + coord_cartesian(xlim=(bot, 100)) - + xlab("% Identity Estimate") - + ylab("# of Estimates{}".format(extra)) + fill_colors = direction_colors or colors + categories = ( + list(dataframe[fill_column].cat.categories) + if isinstance(dataframe[fill_column].dtype, pd.CategoricalDtype) + else list(pd.unique(dataframe[fill_column].dropna())) ) - return p + + samples = [] + sample_colors = [] + for index, category in enumerate(categories): + values = dataframe.loc[ + dataframe[fill_column] == category, "perID_by_events" + ].dropna() + if values.empty: + continue + samples.append(values.to_numpy()) + if isinstance(fill_colors, dict): + sample_colors.append(fill_colors[category]) + else: + sample_colors.append(fill_colors[index % len(fill_colors)]) + + figure, axis = plt.subplots(figsize=(3.0, 3.0)) + try: + if samples: + axis.hist( + samples, + bins=300, + range=(lower_bound, 100.0), + stacked=True, + color=sample_colors, + ) + axis.set_xlim(lower_bound, 100.0) + axis.set_xlabel("% Identity Estimate") + suffix = "\n(thousands)" if dataframe.shape[0] > 1e6 else "" + axis.set_ylabel(f"# of Estimates{suffix}") + axis.grid(False) + for spine in axis.spines.values(): + spine.set_color("#BDBDBD") + set_figure_font_family(figure, DEFAULT_FONT_FAMILY) + figure.tight_layout() + except Exception: + plt.close(figure) + raise + return figure def _missing_symmetric_rows(dataframe): @@ -2247,9 +1557,11 @@ def create_plots( if is_pairwise: plot_filename = os.path.join(directory, f"{name_x}_{name_y}") - histy = make_hist( - sdf, palette, palette_orientation, custom_colors, custom_breakpoints - ) + histy = None + if not no_hist: + histy = make_hist( + sdf, palette, palette_orientation, custom_colors, custom_breakpoints + ) annotation_track_created = False annotation_bed_df = None diff --git a/tests/test_issue53_memory_safe_plotting.py b/tests/test_issue53_memory_safe_plotting.py index 5062251..a2882d4 100644 --- a/tests/test_issue53_memory_safe_plotting.py +++ b/tests/test_issue53_memory_safe_plotting.py @@ -1,6 +1,5 @@ import matplotlib.pyplot as plt import pandas as pd -from plotnine.geoms.geom_tile import geom_tile from moddotplot.static_plots import make_dot, make_dot_final, make_dot_grid, make_tri @@ -11,7 +10,7 @@ def _large_sparse_dotplot_data(): # The adjacent first two coordinates establish a 2 kb resolution while the # final coordinate establishes the ~496 Mb extent from issue #53. A - # geom_raster layer would try to allocate about 248,000**2 RGBA pixels. + # A coordinate-sized raster would allocate about 248,000**2 RGBA pixels. starts = [0, WINDOW_SIZE, GENOME_SIZE - WINDOW_SIZE] return pd.DataFrame( { @@ -43,34 +42,30 @@ def _make_full_plot(data, deraster=False): ) -def _assert_tile_layer(plot, rasterized): - layer = plot.layers[0] - assert isinstance(layer.geom, geom_tile) - assert layer.geom._kwargs["raster"] is rasterized +def _assert_tile_collection(figure, rasterized): + axis = figure.axes[0] + assert not axis.images + assert axis.collections + assert axis.collections[0].get_rasterized() is rasterized def test_large_sparse_plot_does_not_build_coordinate_sized_raster(): - plot = _make_full_plot(_large_sparse_dotplot_data()) + figure = _make_full_plot(_large_sparse_dotplot_data()) - _assert_tile_layer(plot, rasterized=True) - figure = plot.draw(show=False) try: + _assert_tile_collection(figure, rasterized=True) axis = figure.axes[0] - assert not axis.images assert len(axis.collections) == 1 - assert axis.collections[0].get_rasterized() is True assert len(axis.collections[0].get_paths()) == 3 finally: plt.close(figure) def test_deraster_keeps_memory_safe_tiles_as_vectors(): - plot = _make_full_plot(_large_sparse_dotplot_data(), deraster=True) + figure = _make_full_plot(_large_sparse_dotplot_data(), deraster=True) - _assert_tile_layer(plot, rasterized=False) - figure = plot.draw(show=False) try: - assert figure.axes[0].collections[0].get_rasterized() is False + _assert_tile_collection(figure, rasterized=False) finally: plt.close(figure) @@ -100,5 +95,8 @@ def test_grid_and_triangle_paths_use_the_same_memory_safe_geometry(): num_ticks=3, ) - for plot in (grid_cell, grid_plot, triangle): - _assert_tile_layer(plot, rasterized=True) + for figure in (grid_cell, grid_plot, triangle): + try: + _assert_tile_collection(figure, rasterized=True) + finally: + plt.close(figure) diff --git a/tests/test_optional_dependencies.py b/tests/test_optional_dependencies.py index 3e7cca6..34d882e 100644 --- a/tests/test_optional_dependencies.py +++ b/tests/test_optional_dependencies.py @@ -9,7 +9,8 @@ from moddotplot.estimate_identity import convertMatrixToCool, require_cooler_dependency from moddotplot.optional_dependencies import OptionalDependencyError -INSTALL_HINT = 'python -m pip install "ModDotPlot[interactive]"' +INTERACTIVE_INSTALL_HINT = 'python -m pip install "ModDotPlot[interactive]"' +COOLER_INSTALL_HINT = 'python -m pip install "ModDotPlot[cooler]"' def _block_import(monkeypatch, missing_package): @@ -27,14 +28,14 @@ def import_without_optional_package( monkeypatch.setattr(builtins, "__import__", import_without_optional_package) -def test_missing_cooler_reports_the_interactive_extra_install_command(monkeypatch): +def test_missing_cooler_reports_the_cooler_extra_install_command(monkeypatch): _block_import(monkeypatch, "cooler") with pytest.raises(OptionalDependencyError) as exc_info: require_cooler_dependency() assert "Cooler export requires" in str(exc_info.value) - assert INSTALL_HINT in str(exc_info.value) + assert COOLER_INSTALL_HINT in str(exc_info.value) assert isinstance(exc_info.value.__cause__, ModuleNotFoundError) @@ -48,7 +49,7 @@ def test_missing_interactive_package_reports_the_extra_install_command( interactive.require_interactive_dependencies() assert "Interactive mode requires" in str(exc_info.value) - assert INSTALL_HINT in str(exc_info.value) + assert INTERACTIVE_INSTALL_HINT in str(exc_info.value) assert isinstance(exc_info.value.__cause__, ModuleNotFoundError) @@ -74,7 +75,7 @@ def missing_cooler(): cli.main() assert exc_info.value.code == 2 - assert INSTALL_HINT in capsys.readouterr().err + assert COOLER_INSTALL_HINT in capsys.readouterr().err def test_cooler_extra_writes_a_readable_comparative_matrix(tmp_path): @@ -117,4 +118,4 @@ def missing_interactive_dependencies(): cli.main() assert exc_info.value.code == 2 - assert INSTALL_HINT in capsys.readouterr().err + assert INTERACTIVE_INSTALL_HINT in capsys.readouterr().err diff --git a/tests/test_packaging_metadata.py b/tests/test_packaging_metadata.py index 65b375a..a2dfb35 100644 --- a/tests/test_packaging_metadata.py +++ b/tests/test_packaging_metadata.py @@ -20,18 +20,26 @@ def test_runtime_and_distribution_versions_match(): assert project_metadata()["version"] == VERSION -def test_declared_python_floor_matches_documentation(): +def test_declared_python_floor_matches_classifiers(): metadata = project_metadata() assert metadata["requires-python"] == ">=3.10,<3.15" for minor in range(10, 15): assert f"Programming Language :: Python :: 3.{minor}" in metadata["classifiers"] - readme = (PROJECT_ROOT / "README.md").read_text() - assert "supports Python 3.10 through 3.14" in readme -def test_plotnine_supports_declared_python_floor(): +def test_core_dependencies_are_the_required_numeric_and_plotting_stack(): dependencies = project_metadata()["dependencies"] - assert "plotnine>=0.15.8,<0.16" in dependencies + names = { + dependency.split(";", 1)[0] + .split("=", 1)[0] + .split("<", 1)[0] + .split(">", 1)[0] + .strip() + .lower() + for dependency in dependencies + } + + assert names == {"numpy", "pandas", "scipy", "matplotlib"} assert "matplotlib>=3.10.9; python_version < '3.11'" in dependencies assert "matplotlib>=3.11.2; python_version >= '3.11'" in dependencies @@ -50,12 +58,14 @@ def test_interactive_dependencies_are_not_installed_with_the_core_package(): for dependency in core_dependencies } - assert core_names.isdisjoint({"cooler", "dash", "plotly", "pillow"}) + assert core_names.isdisjoint( + {"cooler", "dash", "plotly", "pillow", "plotnine", "palettable", "setproctitle"} + ) assert set(metadata["optional-dependencies"]["interactive"]) == { - "cooler", "dash>=2.9", "plotly", } + assert metadata["optional-dependencies"]["cooler"] == ["cooler"] def test_python_310_test_dependencies_include_tomli(): @@ -116,6 +126,11 @@ def test_svg_composition_dependencies_are_not_runtime_dependencies(): assert "matplotlib" in normalized_names +def test_colorbrewer_license_is_included_in_distributions(): + setup_config = (PROJECT_ROOT / "setup.cfg").read_text() + assert "THIRD_PARTY_LICENSES.md" in setup_config + + def test_ci_covers_every_supported_python_minor(): workflow = (PROJECT_ROOT / ".github/workflows/ci.yml").read_text() for minor in range(10, 15): @@ -129,14 +144,16 @@ def test_package_ci_compares_python_specifiers_semantically(): assert "SpecifierSet(expected['requires-python'])" in workflow -def test_full_test_workflows_install_interactive_test_dependencies(): +def test_full_test_workflows_install_all_optional_test_dependencies(): ci_workflow = (PROJECT_ROOT / ".github/workflows/ci.yml").read_text() release_workflow = ( PROJECT_ROOT / ".github/workflows/publish-to-pypi.yml" ).read_text() - assert 'python -m pip install --editable ".[test,interactive]"' in ci_workflow - assert 'python -m pip install ".[test,interactive]"' in release_workflow + assert ( + 'python -m pip install --editable ".[test,interactive,cooler]"' in ci_workflow + ) + assert 'python -m pip install ".[test,interactive,cooler]"' in release_workflow def test_release_workflow_is_tag_gated_and_uses_trusted_publishing(): diff --git a/tests/test_plot_fonts.py b/tests/test_plot_fonts.py index 53aa958..a9adbf9 100644 --- a/tests/test_plot_fonts.py +++ b/tests/test_plot_fonts.py @@ -1,5 +1,4 @@ import matplotlib.pyplot as plt -from plotnine import ggplot import moddotplot.static_plots as static_plots from moddotplot.native_render import ( @@ -58,39 +57,9 @@ def fail_for_helvetica(*args, **kwargs): plt.close(figure) -def test_plotnine_outputs_retry_with_dejavu_on_glyph_failure(monkeypatch): - attempted_families = [] - - def fail_for_helvetica(plot, **_kwargs): - if hasattr(plot.theme, "getp"): - family = plot.theme.getp(("text", "family"))[0] - else: - # Plotnine <0.15 stores resolved themeable properties directly. - family = plot.theme.themeables["text"].properties["family"][0] - attempted_families.append(family) - if family == DEFAULT_FONT_FAMILY: - raise RuntimeError("failed to load glyph") - - monkeypatch.setattr(static_plots, "ggsave", fail_for_helvetica) - static_plots._save_plot(ggplot(), filename="unused.png") - - assert attempted_families == [DEFAULT_FONT_FAMILY, FALLBACK_FONT_FAMILY] - - -def test_plotnine_pair_draws_once_for_png_and_vector(monkeypatch, tmp_path): +def test_matplotlib_pair_uses_one_figure_for_png_and_vector(monkeypatch, tmp_path): figure = plt.figure() - class FakePlot: - draw_count = 0 - - def __add__(self, _other): - return self - - def draw(self, show=False): - assert show is False - self.draw_count += 1 - return figure - saved = [] def fake_save_figure_pair(current, prefix, vector_format, dpi, **kwargs): @@ -98,10 +67,8 @@ def fake_save_figure_pair(current, prefix, vector_format, dpi, **kwargs): return tmp_path / "plot.png", tmp_path / "plot.svg" monkeypatch.setattr(static_plots, "save_figure_pair", fake_save_figure_pair) - plot = FakePlot() - static_plots._draw_and_save_plot_pair( - plot, + figure, tmp_path / "plot", width=9, height=9, @@ -109,7 +76,6 @@ def fake_save_figure_pair(current, prefix, vector_format, dpi, **kwargs): vector_format="svg", ) - assert plot.draw_count == 1 assert saved[0][0] is figure assert saved[0][2:4] == ("svg", 300) assert saved[0][4] == {"bbox_inches": figure.bbox_inches} diff --git a/tests/test_static_customization.py b/tests/test_static_customization.py index b798e09..c8c0d40 100644 --- a/tests/test_static_customization.py +++ b/tests/test_static_customization.py @@ -1,7 +1,9 @@ +import matplotlib.pyplot as plt import pandas as pd import pytest from matplotlib.colors import to_rgba +from moddotplot.color_palettes import palette_colors from moddotplot.const import DIRECTION_COLORS from moddotplot.static_plots import ( display_sequence_name, @@ -12,6 +14,10 @@ ) +def test_bundled_colorbrewer_palette_preserves_exact_discrete_colors(): + assert palette_colors("Spectral_3") == ["#FC8D59", "#FFFFBF", "#99D594"] + + def test_read_df_from_file_normalizes_browser_bedpe_export(tmp_path): bedpe = tmp_path / "browser.bedpe" bedpe.write_text( @@ -95,7 +101,7 @@ def test_make_dot_uses_custom_color_scale(): } ) - plot = make_dot( + figure = make_dot( sdf=plot_data, name_x="query", name_y="reference", @@ -110,8 +116,15 @@ def test_make_dot_uses_custom_color_scale(): is_pairwise=True, ) - fill_scale = plot.scales.get_scales("fill") - assert fill_scale.palette(len(custom_colors)) == custom_colors + try: + rendered = { + tuple(color) + for collection in figure.axes[0].collections + for color in collection.get_facecolors() + } + assert rendered == {to_rgba(color) for color in custom_colors} + finally: + plt.close(figure) def test_make_dot_uses_direction_colors_when_orientation_is_present(): @@ -128,7 +141,7 @@ def test_make_dot_uses_direction_colors_when_orientation_is_present(): } ) - plot = make_dot( + figure = make_dot( sdf=plot_data, name_x="query", name_y="reference", @@ -143,7 +156,6 @@ def test_make_dot_uses_direction_colors_when_orientation_is_present(): is_pairwise=True, ) - figure = plot.draw(show=False) try: rendered = { tuple(color) @@ -155,8 +167,6 @@ def test_make_dot_uses_direction_colors_when_orientation_is_present(): assert len(rendered) == 4 assert to_rgba("#000000") not in rendered finally: - import matplotlib.pyplot as plt - plt.close(figure) @@ -173,7 +183,7 @@ def test_make_dot_honors_exact_region_bounds(): } ) - plot = make_dot( + figure = make_dot( sdf=plot_data, name_x="query", name_y="reference", @@ -188,8 +198,11 @@ def test_make_dot_honors_exact_region_bounds(): is_pairwise=True, ) - assert plot.scales.get_scales("x").limits == (101.0, 400.0) - assert plot.scales.get_scales("y").limits == (101.0, 400.0) + try: + assert figure.axes[0].get_xlim() == (101.0, 400.0) + assert figure.axes[0].get_ylim() == (101.0, 400.0) + finally: + plt.close(figure) def test_display_names_omit_region_and_full_axes_are_twice_as_large(): @@ -206,7 +219,7 @@ def test_display_names_omit_region_and_full_axes_are_twice_as_large(): } ) - plot = make_dot( + figure = make_dot( sdf=plot_data, name_x=name, name_y=name, @@ -222,18 +235,13 @@ def test_display_names_omit_region_and_full_axes_are_twice_as_large(): ) assert display_sequence_name(name) == "PAN010.chr14.haplotype1.paternal" - assert ":1-4000000" not in plot.labels.title - assert plot.data["q"].unique().tolist() == ["PAN010.chr14.haplotype1.paternal"] - - figure = plot.draw(show=False) try: axis = figure.axes[0] + assert ":1-4000000" not in figure._suptitle.get_text() + assert ":1-4000000" not in axis.get_title() assert axis.get_xticklabels()[0].get_fontsize() == pytest.approx(8) - # Plotnine rounds text sizes to whole points: 2 * (width * 1.4) = 11.2. - assert axis.xaxis.label.get_fontsize() == pytest.approx(11) + assert axis.xaxis.label.get_fontsize() >= 8 finally: - import matplotlib.pyplot as plt - plt.close(figure) From 38c3849e9025745a429e28a7f2cbbd29a93144e8 Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Mon, 5 Oct 2026 00:15:06 -0400 Subject: [PATCH 13/16] Refresh README command reference --- README.md | 462 +++++++++++-------------------- tests/test_packaging_metadata.py | 41 +++ 2 files changed, 197 insertions(+), 306 deletions(-) diff --git a/README.md b/README.md index 3fec657..c2cf19f 100644 --- a/README.md +++ b/README.md @@ -8,12 +8,13 @@ - [About](#about) - [Installation](#installation) - [Usage](#usage) - - [Static Mode](#static-mode) - - [Interactive Mode](#interactive-mode) - - [Standard arguments](#standard-arguments) - - [Static Mode Commands](#static-mode-commands) - - [Input/Output \& Formatting Commands](#inputoutput--formatting-commands) - - [Plot Customization Commands](#plot-customization-commands) + - [Command Line Arguments](#command-line-arguments) + - [General Options](#general-options) + - [Input Options](#input-options) + - [Analysis Options](#analysis-options) + - [Output Options](#output-options) + - [Plot Formatting Options](#plot-formatting-options) + - [Plot Customization Options](#plot-customization-options) - [Sample run - Static Plots](#sample-run---static-plots) - [Using a config file](#using-a-config-file) - [Adding custom bed file annotations](#adding-custom-bed-file-annotations) @@ -45,14 +46,14 @@ If you're interested in learning more about _ModDotPlot_ and how to visualize ta ## Installation -_ModDotPlot_ can be installed for static plotting by running `pip install moddotplot`. Alternatively, you can download the current release from GitHub by using: +_ModDotPlot_ can be installed by running `pip install moddotplot`. Alternatively, you can download the current release from GitHub by using: ``` git clone https://github.com/marbl/ModDotPlot.git cd ModDotPlot ``` -Although optional, it's recommended to setup a virtual environment before using _ModDotPlot_: +Although optional, it's recommended to set up a virtual environment before using _ModDotPlot_: ``` python -m venv venv @@ -65,13 +66,15 @@ Once activated, install either the base package for static plotting: python -m pip install . ``` -or include the optional interactive plotting and Cooler dependencies: +If desired, add the deprecated interactive plotting dependencies, Cooler export support, or both only when needed: ``` python -m pip install ".[interactive]" +python -m pip install ".[cooler]" +python -m pip install ".[interactive,cooler]" ``` -Finally, confirm that the installation was installed correctly and that your version is up to date by running `moddotplot -h`: +Finally, confirm that the installation completed correctly and that your version is up to date by running `moddotplot -h`: ``` __ __ _ _____ _ _____ _ _ | \/ | | | | __ \ | | | __ \| | | | @@ -82,7 +85,7 @@ Finally, confirm that the installation was installed correctly and that your ver v1.0.0 -usage: moddotplot [-h] [{static,interactive}] ... +usage: moddotplot [-h] [--quiet] [{static,interactive}] ... ModDotPlot: Visualization of Tandem Repeats @@ -93,6 +96,7 @@ positional arguments: options: -h, --help show this help message and exit + --quiet suppress all console output, including warnings and errors ``` Note that running `moddotplot -h` might take a while at first! This is because the Python interpreter is compiling source code into the __pycache__ directory. Subsequent runs will use the pre-compiled code and load much faster! @@ -101,269 +105,114 @@ Note that running `moddotplot -h` might take a while at first! This is because t ## Usage -_ModDotPlot_ runs in `static` mode by default. The `static` subcommand remains -available for compatibility and clarity, so these forms are equivalent: +_ModDotPlot_ runs in static mode by default. The explicit `static` subcommand +is retained for clarity and compatibility, so these commands are equivalent: -``` +```bash moddotplot -f sequence.fa moddotplot static -f sequence.fa ``` -### Static Mode - -``` -moddotplot static -``` - -Running _ModDotPlot_ quickly create plots under the specified output directory `-o`. By default, running _ModDotPlot_ will produce the following files: - -- A paired-end bed file `.bedpe`, containing intervals alongside their corresponding identity estimates. -- A self-identity dotplot for each sequence, as both an upper triangle matrix `_TRI` and full matrix `_FULL` representation. -- A histogram of identity values for each sequence. +A standard run writes a BEDPE identity table, full and triangle dotplots, +and an identity histogram beneath the selected output directory. Plots and +histograms are rendered with Matplotlib as both PNG and the selected vector +format (SVG by default). Every plot directory also receives a +`plot_summary.txt` reproducibility record containing input paths, plotting +parameters, and the command used for the run. ![](images/moddotplot_output.png) -Plots and histograms are rendered with Matplotlib and output as both rasterized -`.png` images and vector graphics (default: `.svg`). Grid axes state their -genomic unit (Kbp, Mbp, or Gbp). - -Every directory containing generated static plots also receives a `plot_summary.txt` reproducibility record. It lists the creation time, absolute plot and input paths, window size, any selected region or annotation BED file, and the exact command used for the run. - -_ModDotPlot_ supports highly customizable plotting features in static mode. See [plot customization](#static-mode-commands) for a complete list of features. - - -### Interactive Mode - -**As of ModDotPlot v1.0.0, interactive mode has been deprecated!** - -``` -moddotplot interactive -``` - -Interactive mode is deprecated and maintenance-only. It remains available, but -will not receive new features. It runs only when the `interactive` subcommand is -explicitly provided. Install its optional dependencies with -`pip install "ModDotPlot[interactive]"`, or `pip install ".[interactive]"` from -a source checkout. - -Running _ModDotPlot_ in interactive mode will launch a [Dash application](https://plotly.com/dash/) on your machine's localhost. Open any web browser and go to `http://127.0.0.1:` to view the interactive plot (this should happen automatically, but depending on your environment you might need to copy and paste this URL into your web browser). Running `Ctrl+C` on the command line will exit the Dash application. The default port number used by Dash is `8050`, but this can be customized using the `--port` command (see [interactive mode commands](#interactive-mode-commands) for further info, and [Sample run - Port Forwarding](#sample-run---port-forwarding) for tips on running interactive mode on an HPC environment). - ---- - -### Standard arguments - -The following arguments are the same in both interactive and static mode: - -`-f / --fasta ` - -FASTA files to input. Multi-FASTA files are accepted. By default, every record -is analyzed; static mode can limit a run to named records with -`-s/--sequence`. Interactive mode will only support a maximum of two sequences -at a time. - -`-b / --bed <.bed file(s)>` - -Input BED3-BED9 annotation file used for dotplot annotation (this is not the paired-end BEDPE file produced by ModDotPlot). The BED chromosome field must match a FASTA header, excluding any trailing `:start-end` region suffix. Static mode accepts one BED file and produces an annotation track plus annotated triangle output. Interactive mode accepts one or more BED files, combines their matching intervals, and displays a collapsed track beneath the x axis. Comparative interactive plots also display a track beside the y axis when that sequence has matching annotations. BED `itemRgb` colors are used when present. - -`-k / --kmer ` - -K-mer size to use. This should be large enough to distinguish unique k-mers with enough specificity, but not too large that sensitivity is removed. Default: 21. - -`-o / --output-dir ` - -Name of output directory for bed file & plots. Default is current working directory. - -`--quiet` - -Suppress all console output, including warnings and errors. The process exit -status still indicates whether the run succeeded. - -`-id / --identity ` - -Minimum sequence identity cutoff threshold. Default is 86. While it is possible to go as low as 50% sequence identity, anything below 80% is not recommended. - -`--delta ` - -Each partition includes a fraction of the adjacent windows' k-mers when estimating identity. This recovers repetitive matches that straddle different window boundaries. The default is 0.5, and the accepted range is between 0 and 1; values greater than 0.5 are not recommended. Set this to 0 only when strictly core-local comparisons are desired. - -`-m / --modimizer ` - -Modimizer sketch size. Must be lower than window size `w`. A lower sketch size means less k-mers to compare (and faster runtime), at the expense of lower accuracy. Recommended to be kept >= 1000. - -`--forward ` - -Use forward k-mers only, instead of the default of canonical k-mers. Warning: this will give strand specific output. - -`-r / --resolution ` - -Dotplot resolution. This corresponds to the number of windows each input sequence is partitioned into. Default is 1000. Overrides the `--window` parameter. - -`--compare ` - -If set when 2 or more sequences are input into ModDotPlot, this will show an A vs. B style plot, in addition to a self-identity plot. Note that interactive mode currently only supports a maximum of two sequences. If more than two sequences are input, only the first two will be shown. - -`--compare-only ` - -If set when 2 or more sequences are input into ModDotPlot, this will show an A vs. B style plot, without showing self-identity plots. - -`--ambiguous ` - -By default, every k-mer window containing a non-ACGTU character is excluded from identity estimation without changing its genomic position. This produces gaps through regions containing ambiguous IUPAC bases. To include deterministic hashes for those windows, set the `--ambiguous` flag in either interactive or static mode. - ---- - -### Static Mode Commands - -#### Input/Output & Formatting Commands - -`-l / --load <.bedpe file>` - -Create a plot from a previously computed pairwise bed file. Skips Average Nucleotide Identity computation. Used instead of `-f/--fasta`. Will only accept paired-end bed files produced by ModDotPlot. - -`-c / --config <.json file>` - -Run moddotplot static with a config file instead of command line args. Example syntax in `config/config.json`. Recommended when creating a really customized plot. Used instead of -f/--fasta. - -`-s / --sequence [ ...]` - -In static mode, analyze only the requested records from the input FASTA -file(s). Each ID is matched to the first whitespace-delimited token in a FASTA -header. Exact matches are preferred, with an unambiguous case-insensitive -fallback (so `chr1` selects `Chr1`). Unknown, ambiguous, and duplicate requested -IDs are errors. When using a config file, provide the same list under the -`sequence` key, for example `"sequence": ["chr1", "chr2"]`. - -`--pairs ` - -Limit `--compare` or `--compare-only` to explicitly requested sequence pairs. -The file contains two whitespace-delimited FASTA identifiers per line, in -x-axis then y-axis order. Blank lines and lines beginning with `#` are ignored. -Duplicate, self, unknown, and ambiguous pairs are errors. Indexed FASTA input -is required so each requested record can be fetched without scanning the full -genome. - -`--cooler ` - -If set, will output a matrix as a Cooler file for each input sequence, in addition to a BEDPE file. Install Cooler support with `pip install "ModDotPlot[cooler]"` (or `pip install ".[cooler]"` from a source checkout). - -`--no-bedpe ` - -Skip output of bed file. - -`--no-hist ` - -Skip output of histogram legend. - -`--no-plot ` - -Save .bedpe to file, but skip rendering of plots. - -`--processes <1-4>` - -Set the number of independent chromosome or comparison-group workers. When -omitted, ModDotPlot uses two workers while rendering or up to four for a -`--no-plot` run on multi-record FASTA files that have random-access indexes -(`.fai`, plus `.gzi` for BGZF). Comparative workers group pairs by their y-axis -record and reuse that record's exact sketch. This keeps default aggregate -memory bounded; `--processes 4` opts into maximum throughput. Ordinary gzip and -unindexed inputs remain single-pass and sequential so they are not scanned once -per worker. Use `--processes 1` for explicitly serial execution. - -`--memory-limit ` - -Set an aggregate memory budget for comparative workers. ModDotPlot estimates -the largest pair-local sequence, sketch, matrix, and rendering footprint and -reduces the worker count when necessary. When omitted, available memory is used -when the operating system exposes it. - -`--sketch-cache ` - -Persist compact, prepared comparison sketches for reuse by later indexed runs. -Cache entries are keyed by the FASTA path, size, modification time, record and -region, strand mode, and every sketch parameter. Positional chromosome-wide -k-mer arrays are never stored in this cache. - -`--width ` - -Adjust the output figure width. For a grid this is the width of the complete grid, not each cell. Default is 9 inches. - -`--dpi ` - -Image resolution in dots per inch (not to be confused with dotplot resolution). Default is `300`. - -`--vector ` - -Vectorized image format to output to. Must be one of ["svg", "pdf", "ps"]. Default: `svg` - -`--deraster ` - -By default, vectorized outputs rasterize the actual dotplot (not the axis). This is done to save space, as a high-resolution dotplot can be extremely space inefficient and prevent use of image manipulation software. This plot rasterization can be removed using this flag. - -#### Plot Customization Commands - -`-w / --window ` - -Window size. Unlike interactive mode, only one matrix will be created, so this represents the *only* window size. Default is set to `n/1000` (eg. 3000bp for a 3Mbp sequence). - -`--region ` - -Plot only the requested 1-based, inclusive range for each named sequence. Syntax is `FASTA_ID:start-end`; the identifier must exactly match the FASTA header's first whitespace-delimited token. Supply one value per sequence when every grid row and column should be cropped, for example `--region sample.hap1:1-4000000 sample.hap2:1-4000000`. Region limits apply to self plots, pairwise plots, BEDPE coordinates, and grid axes. - -`--palette ` - -The accepted palettes use [ColorBrewer](https://colorbrewer2.org/) color -specifications and are segregated into 3 types: _Diverging_, _Qualitative_, -and _Sequential_. Syntax is the name of the palette, followed by an underscore -and the number of colors, for example `OrRd_8`. The default is `Spectral_11`. - -`--breakpoints ` - -Add custom identity threshold breakpoints. Note that the number of breakpoints must be equal to the number of colors + 1, otherwise an error will occur. - -`--palette-orientation ` - -Flip sequential order of color palette. Set to `-` by default for divergent palettes. - -`--colors ` (legacy alias: `--color`) - -List of custom colors in hexcode format can be entered sequentially, mapped from low to high identity. - -`--plot-direction ` - -With FASTA input, retain the standard ANI-colored plots and additionally create a `directionality` subfolder. Direction plots use blue for same-orientation matches and pink for reverse-orientation matches, with darker shades representing stronger ANI. Self-comparisons are named `_DIRECTION_FULL`, `_DIRECTION_TRI`, and `_DIRECTION_HIST`; a requested grid is named `_DIRECTION_GRID`. This option reads each input in both canonical and forward-only modes and cannot be reconstructed from a loaded BEDPE file. - -`--grid ` - -Create a square grid containing every self comparison on the bottom-left-to-top-right diagonal and every pairwise comparison off the diagonal. Self-comparison cells are rendered as full symmetric dotplots. The shared genomic-axis titles appear only around the bottom-left cell without enlarging the output canvas. The grid is rendered as one Matplotlib figure and supports three or more input sequences, although large grids become visually dense. - -For example, select only `chr1` and `chr2` from a multi-FASTA input and create -their grid: - -``` -moddotplot -f ../../moddotplot-interactive/Col-CEN_v1.2.fasta -s chr1 chr2 --grid -``` - -`--grid-only ` - -Create the comparison grid without writing the individual dotplots. - -`-t / --axes-ticks ` - -Custom tickmarks for x and y axis. Values outside of the `--axes-limits` will not be shown. - -`-a / --axes-limits ` - -Change axis limits for x and y axis. Useful when comparing multiple plots, allowing them to stay in scale. - -`--bin-freq ` - -By default, histograms are evenly spaced based on the number of colors and the identity threshold. Select this argument to bin based on the frequency of observed identity values. +The deprecated interactive application remains available only through the +explicit `interactive` subcommand. Its commands are documented under +[Interactive Mode Commands](#interactive-mode-commands). + +### Command Line Arguments + +The default/static command accepts every option in the following two sections. +Options marked **shared** are also accepted by the deprecated interactive +command; interactive-specific behavior is summarized separately below. +Flags such as `--grid` and `--no-plot` are switches and do not take a +`true` or `false` value. + +#### General Options + +| Argument | Description | +| --- | --- | +| `-h, --help` | Show help for the root command or selected subcommand and exit. | +| `--quiet` | Suppress console output, including warnings and errors. The exit status still reports success or failure. May appear before or after the subcommand. | + +#### Input Options + +| Argument | Description | +| --- | --- | +| `-f, --fasta FILE [FILE ...]` | Read one or more FASTA, gzip-compressed FASTA, or BGZF FASTA files. Static mode analyzes every record unless `--sequence` limits the selection. Mutually exclusive with static `--load`. | +| `-l, --load BEDPE [BEDPE ...]` | Plot one or more ModDotPlot BEDPE files without recomputing identity. Mutually exclusive with `--fasta`, but may be combined with `--config`; explicit input paths override input paths in the config. Interactive `--load` has different behavior described below. | +| `-c, --config JSON` | Load command settings from a JSON config. Explicit FASTA, BEDPE, config, and output paths are resolved independently; explicit input/output paths take precedence, while config values supply other settings. | +| `-s, --sequence ID [ID ...]` | Analyze only named records from FASTA input. IDs match the first whitespace-delimited FASTA token, preferring exact matches and then unambiguous case-insensitive matches. | +| `--region ID:START-END [ID:START-END ...]` | Analyze 1-based inclusive regions. Region limits propagate to identity computation, BEDPE coordinates, plots, and grid axes. | +| `--pairs FILE` | Restrict `--compare` or `--compare-only` to pairs in a two-column text file. Blank lines and `#` comments are ignored. Requires indexed FASTA input; duplicate, self, unknown, and ambiguous pairs are errors. | +| `-b, --bed BED` | Add one BED3–BED9 annotation file. BED chromosome names must match FASTA identifiers after any `:start-end` suffix is removed. Valid `itemRgb` values are retained. Interactive mode accepts multiple BED files. | + +#### Analysis Options + +| Argument | Description | +| --- | --- | +| `--compare` | Add pairwise comparisons while retaining self-comparisons. Static mode considers all selected pairs unless `--pairs` restricts them. | +| `--compare-only` | Produce pairwise comparisons without self-comparisons. Mutually exclusive with `--compare`. | +| `-k, --kmer INT` | Set k-mer length. Default: `21`. | +| `-m, --modimizer INT` | Set the modimizer sketch target. Must be smaller than the window length. Smaller values are faster but reduce accuracy. Default: `1000`. | +| `-r, --resolution INT` | Set the approximate number of sequence windows. Mutually exclusive with `--window` in static mode. Default: `1000`. | +| `-w, --window INT` | Set window length in base pairs. When omitted, it is inferred from sequence length and resolution. Mutually exclusive with `--resolution` in static mode. | +| `-id, --identity FLOAT` | Set the minimum estimated identity percentage. Default: `86.0`. Values below 80 are generally not recommended. | +| `-d, --delta FLOAT` | Include this fraction of each neighboring window when estimating identity. Accepted range: 0–1. Default: `0.5`. | +| `--forward` | Hash forward k-mers only instead of canonical k-mers, producing strand-specific output. | +| `--ambiguous` | Include deterministic hashes for windows containing non-ACGTU IUPAC bases instead of masking those windows. | +| `--processes N` | Use 1–4 independent chromosome or comparison-group workers. When omitted, indexed multi-record input uses a bounded automatic worker count; unindexed and ordinary gzip input remain sequential. | +| `--memory-limit GIB` | Set an aggregate GiB memory budget for comparison workers. When omitted, available memory is used when the platform exposes it. | +| `--sketch-cache DIRECTORY` | Persist compact prepared sketches for reuse in later indexed comparative runs. Entries are keyed by input identity, record/region, strand mode, and sketch parameters. | + +#### Output Options + +| Argument | Description | +| --- | --- | +| `-o, --output-dir DIRECTORY` | Set the output directory. Static mode writes BEDPE, plots, and summaries; interactive mode writes saved matrices and coordinate logs. Default: the current directory. | +| `--cooler` | Write Cooler matrices in addition to BEDPE. Install with `pip install "ModDotPlot[cooler]"` or `pip install ".[cooler]"`. | +| `--no-bedpe` | Skip BEDPE output. | +| `--no-plot` | Skip all plot rendering in static mode. In interactive mode, prevent Dash from launching; must be combined with `--save`. | +| `--no-hist` | Skip identity histograms. | + +#### Plot Formatting Options + +| Argument | Description | +| --- | --- | +| `--grid` | Render selected self- and pairwise comparisons in a single square grid, in addition to individual plots. | +| `--grid-only` | Render only the comparison grid and skip individual plots. | +| `--compare-order {sequential,size}` | Choose comparative axis order. `sequential` preserves input order; `size` places the larger sequence on the x-axis. Default: `sequential`. | +| `-a, --axes-limits FLOAT` | Set common x/y axis limits for self-identity plots. The value cannot be shorter than the sequence. | +| `-t, --axes-ticks INT [INT ...]` | Set explicit x/y tick positions. Ticks outside the visible limits are omitted. | +| `--axes-number VALUE` | Retained for configuration compatibility as the requested number of axis ticks; currently unused by the Matplotlib renderer. Default: `7`. | +| `--width FLOAT` | Set plot width in inches. For grids, this is the total grid width, not the width of each cell. Default: `9`. | +| `--dpi INT` | Set raster resolution in dots per inch. Default: `300`. | +| `--vector {svg,pdf,ps}` | Select the vector output format. Default: `svg`. | +| `--deraster` | Keep dotplot tiles as vector geometry instead of rasterizing them inside vector output. This can produce very large files. | + +#### Plot Customization Options + +| Argument | Description | +| --- | --- | +| `--palette NAME_COUNT` | Select an exact discrete [ColorBrewer](https://colorbrewer2.org/) palette, such as `OrRd_8`. Default: `Spectral_11`. | +| `--palette-orientation {+,-}` | Select forward or reversed palette order. Diverging palettes retain ModDotPlot's historical orientation convention. Default: `+`. | +| `--colors COLOR [COLOR ...]`, `--color ...` | Supply a custom low-to-high color sequence in hexadecimal or RGB form. `--color` is a legacy alias. | +| `--breakpoints VALUE [VALUE ...]` | Supply custom identity thresholds between the identity cutoff and 100. The number of breakpoints must equal the number of colors plus one. | +| `--bin-freq` | Derive identity color bins from the observed value distribution instead of evenly spacing them between the identity cutoff and 100. | +| `--plot-direction` | Compute strand direction and color matches blue for the same orientation and pink for reverse orientation, with ANI represented by shade intensity. Available only with FASTA input. | +--- ### Sample run - Static Plots #### Using a config file -When running _ModDotPlot_ to produce static plots, it is recommended to use a config file. The config file is provided in JSON, and accepts the same syntax as the command line arguments shown above. Here is an sample run using a centromeric sequence of _Arabadopsis thaliana_: +When running _ModDotPlot_ to produce static plots, it is recommended to use a config file. The config file is provided in JSON, and accepts the same syntax as the command line arguments shown above. Here is a sample run using a centromeric sequence of _Arabidopsis thaliana_: ``` $ cat config/config.json @@ -381,9 +230,9 @@ $ cat config/config.json 99, 100 ], - "output_dir": "Arabadopsis", + "output_dir": "Arabidopsis", "fasta": [ - "sequences/Arabadopsis_chr1_centromere.fa" + "sequences/Arabidopsis_chr1_centromere.fa" ] } ``` @@ -413,20 +262,20 @@ Computing self identity matrix for Chr1:14000001-18000000... Plot Resolution r: 1000 -Saved self-identity matrix as a paired-end bed file to Arabadopsis/Chr1:14000001-18000000/Chr1:14000001-18000000.bedpe +Saved self-identity matrix as a paired-end bed file to Arabidopsis/Chr1:14000001-18000000/Chr1:14000001-18000000.bedpe -Triangle plots, full plots, and histogram for Arabadopsis/Chr1:14000001-18000000/Chr1:14000001-18000000 saved sucessfully. +Triangle plots, full plots, and histogram for Arabidopsis/Chr1:14000001-18000000/Chr1:14000001-18000000 saved successfully. ``` ![](images/Chr1:14000001-18000000_FULL.png) -Using `samtools faidx` will result in a genomic range being added to a fasta file's header (eg. in the above sequence, the header is Chr1:14000001-18000000). _ModDotPlot_ will parse this syntax to add the appropriate axis. +Using `samtools faidx` will result in a genomic range being added to a FASTA file's header (e.g., in the above sequence, the header is Chr1:14000001-18000000). _ModDotPlot_ will parse this syntax to add the appropriate axis. #### Adding custom bed file annotations If providing a custom BED3-BED9 annotation file using `--bed/-b`, _ModDotPlot_ will output additional files: - A collapsed annotation track `_ANNOTATION_TRACK` in PNG and the selected SVG, PDF, or PostScript vector format. Interval colors use the BED `itemRgb` value in column 9 when present, with a default color for BED3-BED8 records or invalid RGB values. -- The annotation track overlayed with a self-identity dotplot `_ANNOTATED` for each sequence present in the annotation track. +- The annotation track overlaid with a self-identity dotplot `_ANNOTATED` for each sequence present in the annotation track. ``` $ moddotplot static -f sequences/HG002_chr13_MATERNAL:1-4000000.fa -b config/hg002v1.1.cenSatv2.0.bed @@ -443,7 +292,7 @@ Running ModDotPlot in static mode Annotation track saved to chr13_MATERNAL:1-4000000/chr13_MATERNAL:1-4000000_ANNOTATION_TRACK -Triangle plots, full plots, and histogram for chr13_MATERNAL:1-4000000/chr13_MATERNAL:1-4000000 saved sucessfully. +Triangle plots, full plots, and histogram for chr13_MATERNAL:1-4000000/chr13_MATERNAL:1-4000000 saved successfully. ``` ![](images/chr13_MATERNAL:1-4000000_TRI_ANNOTATED.png) @@ -485,46 +334,47 @@ directly, and releases pair-local data before continuing; it does not retain a ### Interactive Mode Commands -`-b / --bed <.bed file> [<.bed file> ...]` +**Deprecated:** interactive mode is maintenance-only and will not receive new +features. New browser-based work should use +[ModDotPlot Browser](https://marbl.github.io/ModDotPlot-Browser/). The legacy +Dash application remains available through the explicit subcommand: -Add one or more BED3-BED9 annotation files. A self-identity plot shows the -matching track beneath its x axis. A comparative plot shows independent x- and -y-axis tracks when BED chromosome names match both FASTA headers. A FASTA -header such as `chr14_MATERNAL:1-4000000` matches BED chromosome -`chr14_MATERNAL`, and the interactive axes retain those genomic coordinates. -For example: - -``` -moddotplot interactive -f sample1.fa sample2.fa --compare \ - --bed sample1.bed sample2.bed +```bash +moddotplot interactive ``` -`--port ` - -Port to display ModDotPlot on. Default is 8050, this can be changed to any accepted port. - -`-w / --window ` - -Minimum window size. By default, interactive mode sets a minimum window size based on the sequence length `n/2000` (eg. a 3Mbp sequence will have a 1500bp window). The maximum window size will always be set to `n/1000` (3000bp under the same example). This means that 2 matrices will be created. - -`-q / --quick ` - -This will automatically run interactive mode with a minimum window size equal to the maximum window size (`n/1000`). This will result in a quick launch, however the resolution of the plot will not improve upon zooming in. - -`-s / --save ` +Install its optional dependencies with +`pip install "ModDotPlot[interactive]"`, or +`pip install ".[interactive]"` from a source checkout. The application +listens on `http://127.0.0.1:8050` by default and exits when the process is +stopped with `Ctrl+C`. + +Interactive mode accepts the following arguments. Some share names with the +default/static command but have interactive-specific behavior. + +| Argument | Description | +| --- | --- | +| `--quiet` | Suppress all console output, including warnings and errors. It may appear before or after `interactive`. | +| `-f, --fasta FILE [FILE ...]` | Read FASTA input and compute an interactive matrix hierarchy. Mutually exclusive with `--load`; interactive displays support at most two sequences. | +| `-l, --load DIRECTORY` | Load a previously saved `interactive_matrices` directory containing compressed matrices and `metadata.pkl`. Mutually exclusive with `--fasta`. This is not the static BEDPE loader. | +| `-b, --bed BED [BED ...]` | Add one or more BED3-BED9 annotation files. Tracks are aligned to matching FASTA identifiers on each matrix axis. | +| `-o, --output-dir DIRECTORY` | Set the directory used for saved matrices and coordinate logs. Default: current directory. | +| `-k, --kmer INT` | Set k-mer length. Default: `21`. | +| `-m, --modimizer INT` | Set the modimizer sketch target. Default: `1000`. | +| `-r, --resolution INT` | Set interactive dotplot resolution. Default: `1000`. | +| `-w, --window INT` | Set the minimum interactive window length. When omitted it is inferred from sequence length and resolution. | +| `-id, --identity FLOAT` | Set the minimum estimated identity percentage. Default: `86.0`. | +| `-d, --delta FLOAT` | Include this fraction of neighboring windows during identity estimation. Default: `0.5`. | +| `--compare` | Add a pairwise comparison while retaining self-comparisons. | +| `--compare-only` | Produce the pairwise comparison without self-comparisons. Mutually exclusive with `--compare`. | +| `--ambiguous` | Include deterministic hashes for windows containing non-ACGTU IUPAC bases. | +| `--forward` | Hash forward k-mers only instead of canonical k-mers. | +| `-s, --save` | Save the matrix hierarchy under `OUTPUT_DIR/interactive_matrices` as compressed NumPy arrays plus `metadata.pkl`. | +| `--port INT` | Set the localhost port used by Dash. Default: `8050`. | +| `-q, --quick` | Build a single matrix layer instead of the normal hierarchy for a faster launch without progressively finer zoom resolution. | +| `--no-plot` | Save matrices without launching Dash. Must be combined with `--save`. | -Save the matrices produced in interactive mode. By default, a folder called `interactive_matrices` will be saved in `--output_dir`, containing each matrix in compressed NumPy format, as well as metadata for each matrix in a pickle. Modifying the files in `interactive_matrices` will cause errors when attempting to load them in the future. - -`--no-plot ` - -Save .bedpe to file, but skip rendering of plots. Must be used with `--save`. - -`-l / --load ` - -Load previously saved matrices. Used instead of `-f/--fasta`. - - ---- +--- ### Sample run - Interactive Mode @@ -557,7 +407,7 @@ Dash is running on http://127.0.0.1:8050/ ![](images/chr1_screenshot.png) -The plotly plot can be navigated using the zoom (magnifying glass) and pan (hand) icons. The plot can be reset by double-clicking or selecting the home button. The identity threshold can be modified by seelcting the slider. Colors can be readjusted according to the same gradient based on the new identity levels. +The Plotly plot can be navigated using the zoom (magnifying glass) and pan (hand) icons. The plot can be reset by double-clicking or selecting the home button. The identity threshold can be modified by selecting the slider. Colors can be readjusted according to the same gradient based on the new identity levels. ### Sample run - Port Forwarding @@ -575,7 +425,7 @@ ssh -N -f -L :127.0.0.1: HPC@LOGIN.CREDENTIA You should now be able to view interactive mode using `http://127.0.0.1:`. Note that your own HPC environment may have specific instructions and/or restrictions for setting up port forwarding. -VSCode now has automatic port forwarding built into the terminal menu. See [VSCode documentation](https://code.visualstudio.com/docs/editor/port-forwarding) for further details +VS Code now has automatic port forwarding built into the terminal menu. See [VS Code documentation](https://code.visualstudio.com/docs/editor/port-forwarding) for further details. ![](images/portforwarding.png) diff --git a/tests/test_packaging_metadata.py b/tests/test_packaging_metadata.py index a2dfb35..1b48a94 100644 --- a/tests/test_packaging_metadata.py +++ b/tests/test_packaging_metadata.py @@ -1,3 +1,4 @@ +import argparse from pathlib import Path try: @@ -7,6 +8,7 @@ import tomli as tomllib from moddotplot.const import VERSION +from moddotplot.moddotplot import get_parser PROJECT_ROOT = Path(__file__).resolve().parents[1] @@ -16,6 +18,45 @@ def project_metadata(): return tomllib.load(pyproject)["project"] +def parser_option_strings(parser): + """Return every explicit option from a parser and its subcommands.""" + + options = set() + for action in parser._actions: + options.update(action.option_strings) + if isinstance(action, argparse._SubParsersAction): + for subparser in action.choices.values(): + options.update(parser_option_strings(subparser)) + return options + + +def test_readme_usage_represents_every_cli_argument(): + readme = (PROJECT_ROOT / "README.md").read_text() + usage = readme.split("## Usage", 1)[1].split("## Questions", 1)[0] + + missing = sorted( + option for option in parser_option_strings(get_parser()) if option not in usage + ) + assert missing == [] + + +def test_readme_usage_toc_has_the_current_command_structure(): + readme = (PROJECT_ROOT / "README.md").read_text() + toc = readme.split("## Cite", 1)[0] + + assert " - [Command Line Arguments](#command-line-arguments)" in toc + assert " - [General Options](#general-options)" in toc + assert " - [Input Options](#input-options)" in toc + assert " - [Analysis Options](#analysis-options)" in toc + assert " - [Output Options](#output-options)" in toc + assert " - [Plot Formatting Options](#plot-formatting-options)" in toc + assert " - [Plot Customization Options](#plot-customization-options)" in toc + assert " - [Interactive Mode Commands](#interactive-mode-commands)" in toc + assert "[| `--plot-direction` |" not in toc + assert "[Static Mode](#static-mode)" not in toc + assert "[Static Mode Commands](#static-mode-commands)" not in toc + + def test_runtime_and_distribution_versions_match(): assert project_metadata()["version"] == VERSION From cd5e2d7ca979f21ddb66fac06969d219cbb0447a Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Mon, 5 Oct 2026 00:18:47 -0400 Subject: [PATCH 14/16] Remove benchmark scripts --- benchmarks/benchmark_containment_matrix.py | 101 -------- benchmarks/benchmark_hashing.py | 255 --------------------- benchmarks/benchmark_sketch_reuse.py | 153 ------------- tests/test_hash_benchmark.py | 63 ----- 4 files changed, 572 deletions(-) delete mode 100644 benchmarks/benchmark_containment_matrix.py delete mode 100644 benchmarks/benchmark_hashing.py delete mode 100644 benchmarks/benchmark_sketch_reuse.py delete mode 100644 tests/test_hash_benchmark.py diff --git a/benchmarks/benchmark_containment_matrix.py b/benchmarks/benchmark_containment_matrix.py deleted file mode 100644 index 008b826..0000000 --- a/benchmarks/benchmark_containment_matrix.py +++ /dev/null @@ -1,101 +0,0 @@ -#!/usr/bin/env python3 -"""Benchmark the exact sparse containment-matrix implementation. - -The generated sketches model the default ``delta=0.5`` layout: each expanded -window contains its core plus half of each adjacent core. This exercises the -same dimensions and sketch cardinalities as a roughly 100 Mb sequence at the -default 1,000-bin resolution without allocating a genome-sized hash array. - -Example:: - - python benchmarks/benchmark_containment_matrix.py - python benchmarks/benchmark_containment_matrix.py --resolution 2000 -""" - -from __future__ import annotations - -import argparse -import resource -import sys -import time -from pathlib import Path - -import numpy as np - -try: - from moddotplot.estimate_identity import ( - pairwiseContainmentMatrix, - selfContainmentMatrix, - ) -except ModuleNotFoundError: # Permit running from an uninstalled source tree. - sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) - from moddotplot.estimate_identity import ( - pairwiseContainmentMatrix, - selfContainmentMatrix, - ) - - -def build_sketches(resolution: int, sketch_size: int, seed: int): - rng = np.random.default_rng(seed) - core = [ - np.unique( - rng.integers(0, np.iinfo(np.uint64).max, sketch_size, dtype=np.uint64) - ) - for _ in range(resolution) - ] - expanded = [] - halfway = sketch_size // 2 - for index, sketch in enumerate(core): - pieces = [sketch] - if index: - pieces.append(core[index - 1][halfway:]) - if index + 1 < resolution: - pieces.append(core[index + 1][:halfway]) - expanded.append(np.unique(np.concatenate(pieces))) - return core, expanded - - -def main(argv=None): - parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--resolution", type=int, default=1000) - parser.add_argument("--sketch-size", type=int, default=1612) - parser.add_argument("--identity", type=float, default=86) - parser.add_argument("--kmer", type=int, default=21) - parser.add_argument("--seed", type=int, default=519) - args = parser.parse_args(argv) - - core, expanded = build_sketches(args.resolution, args.sketch_size, args.seed) - compact_bytes = sum(sketch.nbytes for sketch in core + expanded) - - started = time.perf_counter() - self_matrix = selfContainmentMatrix( - core, expanded, args.kmer, args.identity, ambiguous=False - ) - self_seconds = time.perf_counter() - started - - started = time.perf_counter() - pair_matrix = pairwiseContainmentMatrix( - core, - core, - expanded, - expanded, - args.identity, - args.kmer, - ) - pair_seconds = time.perf_counter() - started - - print(f"resolution: {args.resolution}") - print(f"core hashes/window: {args.sketch_size}") - print(f"expanded hashes/window: about {args.sketch_size * 2}") - print(f"prepared sketch memory: {compact_bytes / 2**20:.2f} MiB") - print(f"self matrix: {self_seconds:.3f} s ({self_matrix.shape})") - print(f"pair matrix: {pair_seconds:.3f} s ({pair_matrix.shape})") - print("two-self-plus-pair estimate: " f"{2 * self_seconds + pair_seconds:.3f} s") - peak_rss = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss - # macOS reports bytes; Linux and other supported Unix platforms report KiB. - peak_rss_mib = peak_rss / (2**20 if sys.platform == "darwin" else 2**10) - print(f"peak process RSS: {peak_rss_mib:.2f} MiB") - - -if __name__ == "__main__": - main() diff --git a/benchmarks/benchmark_hashing.py b/benchmarks/benchmark_hashing.py deleted file mode 100644 index 27a2667..0000000 --- a/benchmarks/benchmark_hashing.py +++ /dev/null @@ -1,255 +0,0 @@ -#!/usr/bin/env python3 -"""Compare ModDotPlot's ntHash2 path with the removed mmh3 implementation. - -``mmh3`` is deliberately optional. When it is installed, this utility -recreates the legacy ModDotPlot loop for an apples-to-apples migration -benchmark. Otherwise it reports ntHash2 throughput on its own. - -Examples:: - - python benchmarks/benchmark_hashing.py --length 1000000 --repeats 7 - python benchmarks/benchmark_hashing.py --fasta sequence.fa --kmer 21 - python benchmarks/benchmark_hashing.py --json results.json -""" - -from __future__ import annotations - -import argparse -import gc -import gzip -import json -import random -import statistics -import sys -import time -from pathlib import Path -from typing import Callable, Dict, Iterable, List, Optional, Sequence, Tuple - -try: - from moddotplot.parse_fasta import _hash_sequence -except ModuleNotFoundError: # Permit running from an uninstalled source tree. - sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) - from moddotplot.parse_fasta import _hash_sequence - - -DNA_ALPHABET = "ACGT" -REVERSE_COMPLEMENT = str.maketrans("ACGT", "TGCA") - - -def _load_mmh3(): - """Return the optional legacy module without making it a dependency.""" - try: - import mmh3 # type: ignore[import-not-found] - except ImportError: - return None - return mmh3 - - -def _legacy_mmh3_hashes( - sequence: str, k: int, canonical: bool, mmh3_module -) -> List[int]: - """Reproduce ModDotPlot's pre-ntHash2 per-k-mer implementation.""" - result = [] - for start in range(max(len(sequence) - k + 1, 0)): - kmer = sequence[start : start + k].upper() - forward = mmh3_module.hash(kmer) - if canonical: - reverse = mmh3_module.hash(kmer[::-1].translate(REVERSE_COMPLEMENT)) - result.append(min(forward, reverse)) - else: - result.append(forward) - return result - - -def _moddotplot_nthash2_hashes(sequence: str, k: int, canonical: bool): - """Exercise the compact batch hashing path used by ModDotPlot's CLI.""" - return _hash_sequence( - sequence, - k, - fw_only=not canonical, - ambiguous=False, - ) - - -def _read_first_fasta(path: Path) -> Tuple[str, str]: - opener = gzip.open if path.suffix == ".gz" else open - name: Optional[str] = None - chunks: List[str] = [] - with opener(path, "rt") as handle: - for line in handle: - line = line.strip() - if line.startswith(">"): - if name is not None: - break - name = line[1:].split()[0] or path.name - elif name is not None: - chunks.append(line) - if name is None: - raise ValueError(f"No FASTA record found in {path}") - return name, "".join(chunks) - - -def _time_call(function: Callable[[], Sequence[int]]) -> Tuple[float, int, int]: - gc.collect() - gc.disable() - try: - started = time.perf_counter_ns() - result = function() - elapsed = (time.perf_counter_ns() - started) / 1e9 - finally: - gc.enable() - count = len(result) - # Touch the result without adding a full O(n) checksum to the timing. - raw_result = getattr(result, "data", result) - checksum = 0 if count == 0 else int(raw_result[0]) ^ int(raw_result[-1]) - return elapsed, count, checksum - - -def _summary(samples: Iterable[float], count: int) -> Dict[str, float]: - values = list(samples) - median = statistics.median(values) - return { - "minimum_seconds": min(values), - "median_seconds": median, - "mean_seconds": statistics.mean(values), - "stdev_seconds": statistics.stdev(values) if len(values) > 1 else 0.0, - "maximum_seconds": max(values), - "median_million_hashes_per_second": count / median / 1_000_000, - } - - -def benchmark(sequence: str, k: int, repeats: int, seed: int) -> Dict[str, object]: - mmh3_module = _load_mmh3() - rng = random.Random(seed) - report: Dict[str, object] = { - "sequence_length": len(sequence), - "kmer_length": k, - "repeats": repeats, - "legacy_mmh3_available": mmh3_module is not None, - "modes": {}, - } - - # Warm native code, imports, and allocators before recording samples. - warm_sequence = sequence[: max(k, min(len(sequence), 10_000))] - _moddotplot_nthash2_hashes(warm_sequence, k, canonical=True) - if mmh3_module is not None: - _legacy_mmh3_hashes(warm_sequence, k, True, mmh3_module) - - for mode, canonical in (("forward", False), ("canonical", True)): - implementations: Dict[str, Callable[[], Sequence[int]]] = { - "nthash2": lambda canonical=canonical: _moddotplot_nthash2_hashes( - sequence, k, canonical - ) - } - if mmh3_module is not None: - implementations["legacy_mmh3"] = ( - lambda canonical=canonical: _legacy_mmh3_hashes( - sequence, k, canonical, mmh3_module - ) - ) - - raw: Dict[str, List[float]] = {name: [] for name in implementations} - counts: Dict[str, int] = {} - checksums: Dict[str, int] = {} - for _ in range(repeats): - order = list(implementations) - rng.shuffle(order) - for name in order: - elapsed, count, checksum = _time_call(implementations[name]) - raw[name].append(elapsed) - counts[name] = count - checksums[name] = checksum - - expected_count = max(len(sequence) - k + 1, 0) - if any(count != expected_count for count in counts.values()): - raise RuntimeError( - f"Unexpected k-mer cardinality: expected {expected_count}, got {counts}" - ) - - mode_report: Dict[str, object] = { - "hash_count": expected_count, - "implementations": { - name: { - **_summary(samples, counts[name]), - "raw_seconds": samples, - "checksum": checksums[name], - } - for name, samples in raw.items() - }, - } - if mmh3_module is not None: - summaries = mode_report["implementations"] - mode_report["speedup_over_legacy_mmh3"] = ( - summaries["legacy_mmh3"]["median_seconds"] - / summaries["nthash2"]["median_seconds"] - ) - report["modes"][mode] = mode_report - - return report - - -def _print_report(label: str, report: Dict[str, object]) -> None: - print( - f"Input: {label}; {report['sequence_length']:,} bases; " - f"k={report['kmer_length']}; {report['repeats']} repeats" - ) - print("ntHash2 result: compact NumPy uint64 array; legacy result: Python int list") - if not report["legacy_mmh3_available"]: - print("Legacy comparison: skipped (optional mmh3 is not installed)") - print( - f"{'mode':<10} {'implementation':<14} {'median (s)':>12} " - f"{'Mhash/s':>10} {'speedup':>10}" - ) - for mode, mode_report in report["modes"].items(): - speedup = mode_report.get("speedup_over_legacy_mmh3") - for implementation, stats in mode_report["implementations"].items(): - shown_speedup = ( - f"{speedup:.2f}x" - if implementation == "nthash2" and speedup is not None - else "-" - ) - print( - f"{mode:<10} {implementation:<14} " - f"{stats['median_seconds']:>12.6f} " - f"{stats['median_million_hashes_per_second']:>10.2f} " - f"{shown_speedup:>10}" - ) - - -def main(argv: Optional[Sequence[str]] = None) -> int: - parser = argparse.ArgumentParser(description=__doc__) - source = parser.add_mutually_exclusive_group() - source.add_argument("--fasta", type=Path, help="benchmark the first FASTA record") - source.add_argument( - "--length", type=int, default=1_000_000, help="synthetic sequence length" - ) - parser.add_argument("--kmer", type=int, default=21) - parser.add_argument("--repeats", type=int, default=7) - parser.add_argument("--seed", type=int, default=20260927) - parser.add_argument("--json", type=Path, help="also write raw results as JSON") - args = parser.parse_args(argv) - - if args.kmer <= 0: - parser.error("--kmer must be positive") - if args.length < 0: - parser.error("--length cannot be negative") - if args.repeats <= 0: - parser.error("--repeats must be positive") - - if args.fasta: - label, sequence = _read_first_fasta(args.fasta) - else: - label = f"synthetic(seed={args.seed})" - sequence = "".join( - random.Random(args.seed).choices(DNA_ALPHABET, k=args.length) - ) - - report = benchmark(sequence, args.kmer, args.repeats, args.seed) - _print_report(label, report) - if args.json: - args.json.write_text(json.dumps(report, indent=2) + "\n") - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/benchmarks/benchmark_sketch_reuse.py b/benchmarks/benchmark_sketch_reuse.py deleted file mode 100644 index 484bc21..0000000 --- a/benchmarks/benchmark_sketch_reuse.py +++ /dev/null @@ -1,153 +0,0 @@ -#!/usr/bin/env python3 -"""Benchmark prepared-sketch reuse for a two-sequence static grid. - -The benchmark isolates the work before matrix comparison. With a fixed window, -the uncached workflow prepares both sequences for their self matrices and then -prepares both again for the pairwise matrix (four preparations). The cache -performs the same access pattern with two preparations and two exact hits. - -Example:: - - python benchmarks/benchmark_sketch_reuse.py --length 1000000 --repeats 5 -""" - -from __future__ import annotations - -import argparse -import gc -import math -import statistics -import sys -import time -from pathlib import Path -from typing import Callable, Dict, List, Sequence - -import numpy as np - -try: - from moddotplot.estimate_identity import ( - ModimizerSketchCache, - prepare_modimizer_sketches, - ) -except ModuleNotFoundError: # Permit running from an uninstalled source tree. - sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src")) - from moddotplot.estimate_identity import ( - ModimizerSketchCache, - prepare_modimizer_sketches, - ) - - -def _time(function: Callable[[], None]) -> float: - gc.collect() - started = time.perf_counter() - function() - return time.perf_counter() - started - - -def _configuration(length: int, resolution: int, modimizer: int): - window = math.ceil(length / resolution) - effective_modimizer = min(window, modimizer) - raw_sparsity = round(window / effective_modimizer) - if raw_sparsity <= effective_modimizer: - sparsity = 2 ** int(math.log2(raw_sparsity)) - else: - sparsity = 2 ** (int(math.log2(raw_sparsity - 1)) + 1) - return window, sparsity, round(window / sparsity) - - -def benchmark( - length: int, - resolution: int, - modimizer: int, - delta: float, - kmer: int, - repeats: int, -) -> None: - rng = np.random.default_rng(56) - sequences = [ - rng.integers(0, np.iinfo(np.uint64).max, length, dtype=np.uint64) - for _ in range(2) - ] - window, sparsity, expectation = _configuration(length, resolution, modimizer) - - def prepare(sequence): - return prepare_modimizer_sketches( - length, - sequence, - window, - sparsity, - delta, - kmer, - False, - expectation, - ) - - def uncached(): - for index in (0, 1, 1, 0): - prepare(sequences[index]) - - def cached(): - cache = ModimizerSketchCache(max_entries=2) - for index in (0, 1, 1, 0): - cache.get_or_prepare( - index, - length, - sequences[index], - window, - sparsity, - delta, - kmer, - False, - expectation, - ) - cache.clear() - - # Warm NumPy dispatch and allocators before recording samples. - prepare(sequences[0][: min(length, window)]) - samples: Dict[str, List[float]] = {"uncached": [], "cached": []} - for repeat in range(repeats): - order: Sequence[str] = ( - ("uncached", "cached") if repeat % 2 == 0 else ("cached", "uncached") - ) - for name in order: - samples[name].append(_time(uncached if name == "uncached" else cached)) - - uncached_median = statistics.median(samples["uncached"]) - cached_median = statistics.median(samples["cached"]) - print( - f"{length:,} hashes/sequence; window={window:,}; delta={delta}; " - f"resolution={resolution}; repeats={repeats}" - ) - print(f"uncached (4 preparations): {uncached_median:.6f} s") - print(f"cached (2 preparations): {cached_median:.6f} s") - print(f"preparation speedup: {uncached_median / cached_median:.2f}x") - - -def main(argv=None) -> int: - parser = argparse.ArgumentParser(description=__doc__) - parser.add_argument("--length", type=int, default=100_000) - parser.add_argument("--resolution", type=int, default=100) - parser.add_argument("--modimizer", type=int, default=100) - parser.add_argument("--delta", type=float, default=0.5) - parser.add_argument("--kmer", type=int, default=21) - parser.add_argument("--repeats", type=int, default=5) - args = parser.parse_args(argv) - if min(args.length, args.resolution, args.modimizer, args.kmer, args.repeats) <= 0: - parser.error( - "length, resolution, modimizer, kmer, and repeats must be positive" - ) - if args.delta < 0: - parser.error("delta must be non-negative") - benchmark( - args.length, - args.resolution, - args.modimizer, - args.delta, - args.kmer, - args.repeats, - ) - return 0 - - -if __name__ == "__main__": - raise SystemExit(main()) diff --git a/tests/test_hash_benchmark.py b/tests/test_hash_benchmark.py deleted file mode 100644 index 4cd92bd..0000000 --- a/tests/test_hash_benchmark.py +++ /dev/null @@ -1,63 +0,0 @@ -import importlib.util -from pathlib import Path - -PROJECT_ROOT = Path(__file__).resolve().parents[1] -BENCHMARK_PATH = PROJECT_ROOT / "benchmarks" / "benchmark_hashing.py" - - -def _load_benchmark_module(): - spec = importlib.util.spec_from_file_location( - "moddotplot_hash_benchmark", BENCHMARK_PATH - ) - module = importlib.util.module_from_spec(spec) - spec.loader.exec_module(module) - return module - - -def test_benchmark_runs_without_optional_legacy_dependency(monkeypatch): - benchmark = _load_benchmark_module() - monkeypatch.setattr(benchmark, "_load_mmh3", lambda: None) - - report = benchmark.benchmark("ACGT" * 50, k=5, repeats=1, seed=7) - - assert report["legacy_mmh3_available"] is False - for mode in ("forward", "canonical"): - mode_report = report["modes"][mode] - assert mode_report["hash_count"] == 196 - assert set(mode_report["implementations"]) == {"nthash2"} - assert "speedup_over_legacy_mmh3" not in mode_report - - -def test_benchmark_reports_mmh3_only_when_optional_module_is_available(monkeypatch): - benchmark = _load_benchmark_module() - - class FakeMmh3: - @staticmethod - def hash(value): - return sum(map(ord, value)) - - monkeypatch.setattr(benchmark, "_load_mmh3", lambda: FakeMmh3()) - - report = benchmark.benchmark("ACGT" * 50, k=5, repeats=1, seed=7) - - assert report["legacy_mmh3_available"] is True - for mode in ("forward", "canonical"): - mode_report = report["modes"][mode] - assert set(mode_report["implementations"]) == { - "nthash2", - "legacy_mmh3", - } - assert mode_report["speedup_over_legacy_mmh3"] > 0 - - -def test_benchmark_cli_clearly_reports_skipped_legacy_comparison(monkeypatch, capsys): - benchmark = _load_benchmark_module() - monkeypatch.setattr(benchmark, "_load_mmh3", lambda: None) - - assert benchmark.main(["--length", "100", "--kmer", "5", "--repeats", "1"]) == 0 - - output = capsys.readouterr().out - assert "Legacy comparison: skipped (optional mmh3 is not installed)" in output - assert "forward" in output - assert "canonical" in output - assert "nthash2" in output From 0aeb07e0ff4b55ac665766d31087d603da6b06b3 Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Mon, 5 Oct 2026 00:27:18 -0400 Subject: [PATCH 15/16] Document dependencies and fix formatting CI --- .github/workflows/ci.yml | 2 +- README.md | 34 +++++++++++++++++++++++++++++++++- 2 files changed, 34 insertions(+), 2 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index e54a30a..ab3a1a6 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -119,7 +119,7 @@ jobs: run: python -m pip install black - name: Check formatting - run: python -m black --check src tests benchmarks setup.py + run: python -m black --check src tests setup.py package: name: Build and validate package diff --git a/README.md b/README.md index c2cf19f..47d8486 100644 --- a/README.md +++ b/README.md @@ -7,6 +7,7 @@ - [Cite](#cite) - [About](#about) - [Installation](#installation) +- [Dependencies](#dependencies) - [Usage](#usage) - [Command Line Arguments](#command-line-arguments) - [General Options](#general-options) @@ -101,7 +102,37 @@ options: Note that running `moddotplot -h` might take a while at first! This is because the Python interpreter is compiling source code into the __pycache__ directory. Subsequent runs will use the pre-compiled code and load much faster! ---- +--- + +## Dependencies + +_ModDotPlot_ requires Python 3.10 through 3.14 (`>=3.10,<3.15`). The base +installation includes everything needed for static plotting. `pip` installs +these required runtime dependencies automatically: + +| Dependency | Version requirement | Purpose | +| --- | --- | --- | +| [NumPy](https://numpy.org/) | No explicit minimum | Numerical arrays and matrix operations. | +| [pandas](https://pandas.pydata.org/) | No explicit minimum | BEDPE and annotation table handling. | +| [Matplotlib](https://matplotlib.org/) | `>=3.10.9` on Python 3.10; `>=3.11.2` on Python 3.11–3.14 | Static plot, grid, and histogram rendering. | +| [SciPy](https://scipy.org/) | No explicit minimum | Numerical and statistical utilities. | + +Building from source also requires Setuptools 61 or newer. This build +dependency is installed automatically by modern versions of `pip`. + +Optional dependencies are grouped by feature and are not needed for standard +static plots: + +| Extra | Dependencies | Purpose | Installation | +| --- | --- | --- | --- | +| `interactive` | Dash `>=2.9`, Plotly | Deprecated interactive Dash application. | `python -m pip install "ModDotPlot[interactive]"` | +| `cooler` | Cooler | Cooler matrix export with `--cooler`. | `python -m pip install "ModDotPlot[cooler]"` | +| `test` | pytest, pytest-cov; tomli `>=1.1` on Python 3.10 | Development and test suite. | `python -m pip install ".[test]"` | + +From a source checkout, multiple extras can be installed together, for +example: `python -m pip install ".[interactive,cooler,test]"`. + +--- ## Usage @@ -206,6 +237,7 @@ Flags such as `--grid` and `--no-plot` are switches and do not take a | `--breakpoints VALUE [VALUE ...]` | Supply custom identity thresholds between the identity cutoff and 100. The number of breakpoints must equal the number of colors plus one. | | `--bin-freq` | Derive identity color bins from the observed value distribution instead of evenly spacing them between the identity cutoff and 100. | | `--plot-direction` | Compute strand direction and color matches blue for the same orientation and pink for reverse orientation, with ANI represented by shade intensity. Available only with FASTA input. | + --- ### Sample run - Static Plots From 2aafdba233aed82aa7682bd9fb230204ed2be641 Mon Sep 17 00:00:00 2001 From: Alex Sweeten Date: Mon, 5 Oct 2026 00:43:57 -0400 Subject: [PATCH 16/16] Potential fix for pull request finding 'Transpose heatmap matrix before exporting coordinates' Co-authored-by: Copilot Autofix powered by AI <175728472+Copilot@users.noreply.github.com> --- src/moddotplot/interactive.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/moddotplot/interactive.py b/src/moddotplot/interactive.py index 3f2ea73..0312d32 100644 --- a/src/moddotplot/interactive.py +++ b/src/moddotplot/interactive.py @@ -230,7 +230,7 @@ def figure_to_bed(figure, default_identity=86.0): filename = f"{x_name}.bedpe" if self_identity else f"{x_name}-{y_name}.bedpe" rows = convertMatrixToBed( - matrix, + matrix.T, window_size, identity, x_name,