diff --git a/.github/workflows/publish.yml b/.github/workflows/publish.yml new file mode 100644 index 0000000..ee744e1 --- /dev/null +++ b/.github/workflows/publish.yml @@ -0,0 +1,52 @@ +name: publish to PyPI + +on: + release: + types: [published] + +jobs: + test: + uses: ./.github/workflows/test.yml + + build: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + + - name: Install uv + uses: astral-sh/setup-uv@v5 + + - name: Check version matches the release tag + run: | + PKG_VERSION=$(grep -oP '__version__ = "\K[^"]+' nt2/__init__.py) + TAG="${{ github.event.release.tag_name }}" + TAG="${TAG#v}" # allow tags like v0.1.0 + if [ "$PKG_VERSION" != "$TAG" ]; then + echo "::error::package version ($PKG_VERSION) != release tag ($TAG)" + exit 1 + fi + + - name: Build sdist and wheel + run: uv build + + - uses: actions/upload-artifact@v4 + with: + name: dist + path: dist/ + + publish: + needs: build + runs-on: ubuntu-latest + environment: + name: pypi + url: https://pypi.org/p/nt2py + permissions: + id-token: write # required for PyPI Trusted Publishing (OIDC) + steps: + - uses: actions/download-artifact@v4 + with: + name: dist + path: dist/ + + - name: Publish to PyPI + uses: pypa/gh-action-pypi-publish@release/v1 diff --git a/.github/workflows/unittests.yml b/.github/workflows/test.yml similarity index 56% rename from .github/workflows/unittests.yml rename to .github/workflows/test.yml index 005af95..cb5a96a 100644 --- a/.github/workflows/unittests.yml +++ b/.github/workflows/test.yml @@ -1,17 +1,19 @@ name: Unit tests -on: [push] +on: + push: + workflow_call: jobs: - build: + test: runs-on: ubuntu-latest strategy: fail-fast: false matrix: - python-version: ["3.8", "3.9", "3.10", "3.11", "3.12", "3.13", "3.14"] + python-version: ["3.9", "3.10", "3.11", "3.12", "3.13", "3.14"] steps: - - uses: actions/checkout@v3 + - uses: actions/checkout@v4 with: lfs: true @@ -28,9 +30,3 @@ jobs: - name: Test with `pytest` run: | pytest - - - name: Publish package - if: github.event_name == 'push' && startsWith(github.ref, 'refs/tags') && matrix.python-version == '3.14' - uses: pypa/gh-action-pypi-publish@release/v1 - with: - password: ${{ secrets.PYPI_API_TOKEN }} diff --git a/.gitignore b/.gitignore index 17eb961..5567ed1 100644 --- a/.gitignore +++ b/.gitignore @@ -47,7 +47,7 @@ coverage.xml *.cover *.py,cover .hypothesis/ -.pytest_cache/ +.*cache/ cover/ # Translations @@ -155,3 +155,15 @@ nt2/tests/testdata test/ temp/ *.bak +.*venv + +# LLM +.codex +.agents +.claude + +# devenv +.devenv* +devenv.local.nix +devenv.local.yaml +.direnv diff --git a/.vscode/extensions.json b/.vscode/extensions.json index 00563da..0b56f9d 100644 --- a/.vscode/extensions.json +++ b/.vscode/extensions.json @@ -1,6 +1,8 @@ { "recommendations": [ "ms-python.python", - "ms-python.black-formatter" + "meta.pyrefly", + "charliermarsh.ruff", + "github.vscode-github-actions" ] } \ No newline at end of file diff --git a/LICENSE b/LICENSE index 3ac8c81..9f07b06 100644 --- a/LICENSE +++ b/LICENSE @@ -1,6 +1,6 @@ BSD 3-Clause License -Copyright (c) 2025, Entity development team +Copyright (c) 2026, Entity development team Redistribution and use in source and binary forms, with or without modification, are permitted provided that the following conditions are met: diff --git a/devenv.lock b/devenv.lock new file mode 100644 index 0000000..85e76bb --- /dev/null +++ b/devenv.lock @@ -0,0 +1,65 @@ +{ + "nodes": { + "devenv": { + "locked": { + "dir": "src/modules", + "lastModified": 1789480904, + "narHash": "sha256-anHWHTkxI7423/eUroN9Rc4+H8VIGyIUWq/hV1egBDQ=", + "owner": "cachix", + "repo": "devenv", + "rev": "0fd5e6d3a9b2ae6f10205a3e95e0ca58d526c5a6", + "type": "github" + }, + "original": { + "dir": "src/modules", + "owner": "cachix", + "repo": "devenv", + "type": "github" + } + }, + "nixpkgs": { + "inputs": { + "nixpkgs-src": "nixpkgs-src" + }, + "locked": { + "lastModified": 1787753358, + "narHash": "sha256-Tl77VbWyAKrOfRNQhL6JbQtb/MLzbYa/1RG1gWWfICk=", + "owner": "cachix", + "repo": "devenv-nixpkgs", + "rev": "256551e45f6303e142ab4a98be1bf243feb77dc0", + "type": "github" + }, + "original": { + "owner": "cachix", + "ref": "rolling", + "repo": "devenv-nixpkgs", + "type": "github" + } + }, + "nixpkgs-src": { + "flake": false, + "locked": { + "lastModified": 1787394516, + "narHash": "sha256-pRGOQSClnXNI2iLUG6DYpsGvYcuw0drOutVZFTJNw90=", + "owner": "NixOS", + "repo": "nixpkgs", + "rev": "c8f90650c15282fa8656a041bfbbd2403997a9a7", + "type": "github" + }, + "original": { + "owner": "NixOS", + "ref": "nixpkgs-unstable", + "repo": "nixpkgs", + "type": "github" + } + }, + "root": { + "inputs": { + "devenv": "devenv", + "nixpkgs": "nixpkgs" + } + } + }, + "root": "root", + "version": 7 +} \ No newline at end of file diff --git a/devenv.nix b/devenv.nix new file mode 100644 index 0000000..e6f6ca2 --- /dev/null +++ b/devenv.nix @@ -0,0 +1,43 @@ +{ pkgs, lib, ... }: + +let + # override with + # devenv shell -O languages.python.package:pkg python312 + py = "313"; +in +{ + name = "nt2dev"; + + languages.python = { + enable = true; + package = pkgs."python${py}"; + + venv = { + enable = true; + requirements = '' + ipykernel + jupyterlab + pytest + -e . + ''; + }; + }; + + # https://devenv.sh/packages/ + packages = with pkgs; [ + black + pyright + taplo + vscode-langservers-extracted + zlib + ]; + + env.LD_LIBRARY_PATH = lib.makeLibraryPath [ + pkgs.stdenv.cc.cc + pkgs.zlib + ]; + + enterShell = '' + echo "nt2dev devenv activated: $DEVENV_STATE/venv/bin/python" + ''; +} diff --git a/devenv.yaml b/devenv.yaml new file mode 100644 index 0000000..68616a4 --- /dev/null +++ b/devenv.yaml @@ -0,0 +1,4 @@ +# yaml-language-server: $schema=https://devenv.sh/devenv.schema.json +inputs: + nixpkgs: + url: github:cachix/devenv-nixpkgs/rolling diff --git a/dist/nt2py-0.2.1.tar.gz b/dist/nt2py-0.2.1.tar.gz deleted file mode 100644 index 925ffea..0000000 Binary files a/dist/nt2py-0.2.1.tar.gz and /dev/null differ diff --git a/dist/nt2py-0.2.tar.gz b/dist/nt2py-0.2.tar.gz deleted file mode 100644 index 669f22b..0000000 Binary files a/dist/nt2py-0.2.tar.gz and /dev/null differ diff --git a/dist/nt2py-0.3.0.tar.gz b/dist/nt2py-0.3.0.tar.gz deleted file mode 100644 index 9f9bb67..0000000 Binary files a/dist/nt2py-0.3.0.tar.gz and /dev/null differ diff --git a/dist/nt2py-0.3.1.tar.gz b/dist/nt2py-0.3.1.tar.gz deleted file mode 100644 index 6757d0f..0000000 Binary files a/dist/nt2py-0.3.1.tar.gz and /dev/null differ diff --git a/dist/nt2py-0.4.0.tar.gz b/dist/nt2py-0.4.0.tar.gz deleted file mode 100644 index 78a8ba2..0000000 Binary files a/dist/nt2py-0.4.0.tar.gz and /dev/null differ diff --git a/dist/nt2py-0.4.1.tar.gz b/dist/nt2py-0.4.1.tar.gz deleted file mode 100644 index 6f4b247..0000000 Binary files a/dist/nt2py-0.4.1.tar.gz and /dev/null differ diff --git a/dist/nt2py-0.5.0-py3-none-any.whl b/dist/nt2py-0.5.0-py3-none-any.whl deleted file mode 100644 index 97af17a..0000000 Binary files a/dist/nt2py-0.5.0-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-0.5.0.tar.gz b/dist/nt2py-0.5.0.tar.gz deleted file mode 100644 index 9fdcccc..0000000 Binary files a/dist/nt2py-0.5.0.tar.gz and /dev/null differ diff --git a/dist/nt2py-0.5.1-py3-none-any.whl b/dist/nt2py-0.5.1-py3-none-any.whl deleted file mode 100644 index acd59f8..0000000 Binary files a/dist/nt2py-0.5.1-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-0.5.1.tar.gz b/dist/nt2py-0.5.1.tar.gz deleted file mode 100644 index 04d0c23..0000000 Binary files a/dist/nt2py-0.5.1.tar.gz and /dev/null differ diff --git a/dist/nt2py-0.5.2-py3-none-any.whl b/dist/nt2py-0.5.2-py3-none-any.whl deleted file mode 100644 index 970f14b..0000000 Binary files a/dist/nt2py-0.5.2-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-0.5.2.tar.gz b/dist/nt2py-0.5.2.tar.gz deleted file mode 100644 index 9033fb8..0000000 Binary files a/dist/nt2py-0.5.2.tar.gz and /dev/null differ diff --git a/dist/nt2py-0.5.3-py3-none-any.whl b/dist/nt2py-0.5.3-py3-none-any.whl deleted file mode 100644 index f73e439..0000000 Binary files a/dist/nt2py-0.5.3-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-0.5.3.tar.gz b/dist/nt2py-0.5.3.tar.gz deleted file mode 100644 index b8477a0..0000000 Binary files a/dist/nt2py-0.5.3.tar.gz and /dev/null differ diff --git a/dist/nt2py-0.6.0-py3-none-any.whl b/dist/nt2py-0.6.0-py3-none-any.whl deleted file mode 100644 index 04d80fc..0000000 Binary files a/dist/nt2py-0.6.0-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-0.6.0.tar.gz b/dist/nt2py-0.6.0.tar.gz deleted file mode 100644 index 51cbc0b..0000000 Binary files a/dist/nt2py-0.6.0.tar.gz and /dev/null differ diff --git a/dist/nt2py-1.0.1-py3-none-any.whl b/dist/nt2py-1.0.1-py3-none-any.whl deleted file mode 100644 index ff87313..0000000 Binary files a/dist/nt2py-1.0.1-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-1.0.1.tar.gz b/dist/nt2py-1.0.1.tar.gz deleted file mode 100644 index 52abe0f..0000000 Binary files a/dist/nt2py-1.0.1.tar.gz and /dev/null differ diff --git a/dist/nt2py-1.1.0-py3-none-any.whl b/dist/nt2py-1.1.0-py3-none-any.whl deleted file mode 100644 index 580236c..0000000 Binary files a/dist/nt2py-1.1.0-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-1.1.0.tar.gz b/dist/nt2py-1.1.0.tar.gz deleted file mode 100644 index f41e835..0000000 Binary files a/dist/nt2py-1.1.0.tar.gz and /dev/null differ diff --git a/dist/nt2py-1.2.0-py3-none-any.whl b/dist/nt2py-1.2.0-py3-none-any.whl deleted file mode 100644 index c1aae6e..0000000 Binary files a/dist/nt2py-1.2.0-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-1.2.0.tar.gz b/dist/nt2py-1.2.0.tar.gz deleted file mode 100644 index 43eec3e..0000000 Binary files a/dist/nt2py-1.2.0.tar.gz and /dev/null differ diff --git a/dist/nt2py-1.2.1-py3-none-any.whl b/dist/nt2py-1.2.1-py3-none-any.whl deleted file mode 100644 index 18f65fc..0000000 Binary files a/dist/nt2py-1.2.1-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-1.2.1.tar.gz b/dist/nt2py-1.2.1.tar.gz deleted file mode 100644 index ca69169..0000000 Binary files a/dist/nt2py-1.2.1.tar.gz and /dev/null differ diff --git a/dist/nt2py-1.3.0-py3-none-any.whl b/dist/nt2py-1.3.0-py3-none-any.whl deleted file mode 100644 index 62f8ebd..0000000 Binary files a/dist/nt2py-1.3.0-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-1.3.0.tar.gz b/dist/nt2py-1.3.0.tar.gz deleted file mode 100644 index 9470de8..0000000 Binary files a/dist/nt2py-1.3.0.tar.gz and /dev/null differ diff --git a/dist/nt2py-1.4.0-py3-none-any.whl b/dist/nt2py-1.4.0-py3-none-any.whl deleted file mode 100644 index 7344d38..0000000 Binary files a/dist/nt2py-1.4.0-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-1.4.0.tar.gz b/dist/nt2py-1.4.0.tar.gz deleted file mode 100644 index 6ad5285..0000000 Binary files a/dist/nt2py-1.4.0.tar.gz and /dev/null differ diff --git a/dist/nt2py-1.5.0-py3-none-any.whl b/dist/nt2py-1.5.0-py3-none-any.whl deleted file mode 100644 index 23f57bb..0000000 Binary files a/dist/nt2py-1.5.0-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-1.5.0.tar.gz b/dist/nt2py-1.5.0.tar.gz deleted file mode 100644 index 2b748bc..0000000 Binary files a/dist/nt2py-1.5.0.tar.gz and /dev/null differ diff --git a/dist/nt2py-1.5.1-py3-none-any.whl b/dist/nt2py-1.5.1-py3-none-any.whl deleted file mode 100644 index 9dfd59d..0000000 Binary files a/dist/nt2py-1.5.1-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-1.5.1.tar.gz b/dist/nt2py-1.5.1.tar.gz deleted file mode 100644 index 941a45e..0000000 Binary files a/dist/nt2py-1.5.1.tar.gz and /dev/null differ diff --git a/dist/nt2py-1.5.2-py3-none-any.whl b/dist/nt2py-1.5.2-py3-none-any.whl deleted file mode 100644 index 129371c..0000000 Binary files a/dist/nt2py-1.5.2-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-1.5.2.tar.gz b/dist/nt2py-1.5.2.tar.gz deleted file mode 100644 index 15209cf..0000000 Binary files a/dist/nt2py-1.5.2.tar.gz and /dev/null differ diff --git a/dist/nt2py-1.5.3-py3-none-any.whl b/dist/nt2py-1.5.3-py3-none-any.whl deleted file mode 100644 index 01be5e5..0000000 Binary files a/dist/nt2py-1.5.3-py3-none-any.whl and /dev/null differ diff --git a/dist/nt2py-1.5.3.tar.gz b/dist/nt2py-1.5.3.tar.gz deleted file mode 100644 index 51a2f28..0000000 Binary files a/dist/nt2py-1.5.3.tar.gz and /dev/null differ diff --git a/legacy/containers/__init__.py b/legacy/containers/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/nt2/containers/container.py b/legacy/containers/container.py similarity index 59% rename from nt2/containers/container.py rename to legacy/containers/container.py index 430051b..dcb4092 100644 --- a/nt2/containers/container.py +++ b/legacy/containers/container.py @@ -1,4 +1,4 @@ -from typing import Callable, Optional, Dict, Tuple +from typing import Callable, Optional, Dict, Tuple, List from nt2.readers.base import BaseReader @@ -10,6 +10,9 @@ class BaseContainer: __reader: BaseReader __remap: Optional[Dict[str, Callable[[str], str]]] + __valid_steps: Dict[str, List[int]] + __valid_files: Dict[str, List[str]] + def __init__( self, path: str, @@ -48,6 +51,40 @@ def remap(self) -> Optional[Dict[str, Callable[[str], str]]]: """{ str: (str) -> str } : The coordinate/field remap dictionary.""" return self.__remap + @property + def valid_steps(self) -> Dict[str, List[int]]: + """dict[str, list[int]]: The valid steps for each category.""" + return self.__valid_steps + + @property + def valid_files(self) -> Dict[str, List[str]]: + """dict[str, list[str]]: The valid files for each category.""" + return self.__valid_files + + @valid_steps.setter + def valid_steps(self, value: Dict[str, List[int]]) -> None: + """Set the valid steps for each category. + + Parameters + ---------- + value : dict[str, list[int]] + The valid steps for each category. + + """ + self.__valid_steps = value + + @valid_files.setter + def valid_files(self, value: Dict[str, List[str]]) -> None: + """Set the valid files for each category. + + Parameters + ---------- + value : dict[str, list[str]] + The valid files for each category. + + """ + self.__valid_files = value + def __dask_tokenize__(self) -> Tuple[str, str, str]: """Provide a deterministic Dask token for container instances.""" return ( diff --git a/legacy/containers/data.py b/legacy/containers/data.py new file mode 100644 index 0000000..570f872 --- /dev/null +++ b/legacy/containers/data.py @@ -0,0 +1,417 @@ +import logging +import sys +from typing import Any, Callable, Dict, List, Optional, Union + +if sys.version_info >= (3, 12): + from typing import override +else: + + def override(method): + return method + + +import pandas as pd +import xarray as xr + +import nt2.plotters.inspect as acc_inspect +import nt2.plotters.movie as acc_movie +import nt2.plotters.particles as acc_particles +import nt2.plotters.polar as acc_polar +from nt2.containers.diagnostics import Diagnostics +from nt2.containers.fields import Fields +from nt2.containers.particles import Particles +from nt2.containers.spectra import Spectra +from nt2.plotters.export import makeFramesAndMovie +from nt2.readers.adios2 import Reader as BP5Reader +from nt2.readers.base import BaseReader +from nt2.readers.hdf5 import Reader as HDF5Reader +from nt2.utils import ( + CoordinateSystem, + DetermineDataFormat, + Format, + InheritClassDocstring, + ToHumanReadable, +) + + +@xr.register_dataset_accessor("polar") +@InheritClassDocstring +class DatasetPolarPlotAccessor(acc_polar.ds_accessor): + pass + + +@xr.register_dataset_accessor("particles") +@InheritClassDocstring +class DatasetParticlesPlotAccessor(acc_particles.ds_accessor): + pass + + +@xr.register_dataarray_accessor("polar") +@InheritClassDocstring +class PolarPlotAccessor(acc_polar.accessor): + pass + + +@xr.register_dataset_accessor("inspect") +@InheritClassDocstring +class DatasetInspectPlotAccessor(acc_inspect.ds_accessor): + pass + + +@xr.register_dataarray_accessor("movie") +@InheritClassDocstring +class MoviePlotAccessor(acc_movie.accessor): + pass + + +# Cartesian remapping functions +def remap_fields_cart(name: str) -> str: + name = name[1:] + fieldname = name.split("_")[0] + fieldname = fieldname.replace("0", "t") + fieldname = fieldname.replace("1", "x") + fieldname = fieldname.replace("2", "y") + fieldname = fieldname.replace("3", "z") + suffix = "_".join(name.split("_")[1:]) + return f"{fieldname}{'_' + suffix if suffix != '' else ''}" + + +def remap_coords_cart(name: str) -> str: + return { + "X1": "x", + "X2": "y", + "X3": "z", + }.get(name, name) + + +def remap_prtl_quantities_cart(name: str) -> str: + shortname = name[1:] + return { + "X1": "x", + "X2": "y", + "X3": "z", + "U1": "ux", + "U2": "uy", + "U3": "uz", + "W": "w", + }.get(shortname, shortname) + + +# Spherical remapping functions +def remap_fields_sph(name: str) -> str: + name = name[1:] + fieldname = name.split("_")[0] + fieldname = fieldname.replace("0", "t") + fieldname = fieldname.replace("1", "r") + fieldname = fieldname.replace("2", "th") + fieldname = fieldname.replace("3", "ph") + suffix = "_".join(name.split("_")[1:]) + return f"{fieldname}{'_' + suffix if suffix != '' else ''}" + + +def remap_coords_sph(name: str) -> str: + return { + "X1": "r", + "X2": "th", + "X3": "ph", + }.get(name, name) + + +def remap_prtl_quantities_sph(name: str) -> str: + shortname = name[1:] + return { + "X1": "r", + "X2": "th", + "X3": "ph", + "U1": "ur", + "U2": "uth", + "U3": "uph", + "W": "w", + }.get(shortname, shortname) + + +def compactify(lst: Union[List[Any], Any]) -> str: + c = "" + cntr = 0 + for l_ in lst: + if cntr > 5: + c += "\n| " + cntr = 0 + c += f"{l_}, " + cntr += 1 + return c[:-2] + + +class Data(Fields, Particles, Spectra): + """Main class to manage all the data containers. + + Inherits from all category-specific containers. + + """ + + __reader: BaseReader + __diagnostics: Optional[Diagnostics] + + def __init__( + self, + path: str, + reader: Optional[BaseReader] = None, + remap: Optional[Dict[str, Callable[[str], str]]] = None, + coord_system: Optional[CoordinateSystem] = None, + ): + """Initializer for the Data class. + + Parameters + ---------- + path : str + Main path to the data + reader : BaseReader, optional + Reader to use to read the data. If None, it will be determined + based on the file format. + remap : dict[str, Callable[[str], str]], optional + Remap dictionary to use to remap the data names (coords, fields, etc.). + coord_system : CoordinateSystem, optional + Coordinate system of the data. If None, it will be determined + based on the data attrs (if remap is also None). + + Raises + ------ + NotImplementedError + If the data format or coordinate system support is not implemented. + ValueError + If the reader format does not match the data format or if coordinate system cannot be inferred. + """ + # determine the reader from the format + fmt = DetermineDataFormat(path) + if reader is None: + if fmt == Format.HDF5: + self.__reader = HDF5Reader() + elif fmt == Format.BP5: + self.__reader = BP5Reader() + else: + raise NotImplementedError( + "Only HDF5 & BP5 formats are supported at the moment." + ) + else: + if fmt != reader.format: + raise ValueError( + f"Reader format {reader.format} does not match data format {fmt}." + ) + self.__reader = reader + + # determine valid files & steps for each category + for category in ["fields", "particles", "spectra"]: + self.valid_files[category] = self.__reader.GetValidFiles( + path=path, category=category + ) + if self.__reader.DefinesCategory( + path, category, self.valid_files[category] + ): + self.valid_steps[category] = self.__reader.GetValidSteps( + path=path, category=category + ) + if len(self.valid_steps[category]) == 0: + raise ValueError(f"No valid steps found for category {category}.") + + # determine the coordinate system and remapping + self.__attrs: Dict[str, Any] = {} + for category in ["fields", "particles", "spectra"]: + if ( + category in self.valid_files + and len(self.valid_files[category]) > 0 + and category in self.valid_steps + and len(self.valid_steps[category]) > 0 + ): + first_step = self.valid_steps[category][0] + attrs = self.__reader.ReadAttrsAtTimestep(path, category, first_step) + self.__attrs.update(**attrs) + if "Coordinates" not in attrs: + raise ValueError( + f"Coordinates not found in attributes for category {category}." + ) + else: + if attrs["Coordinates"] in [b"cart", "cart"]: + coord_system = CoordinateSystem.XYZ + elif attrs["Coordinates"] in [b"sph", "sph", b"qsph", "qsph"]: + coord_system = CoordinateSystem.SPH + + else: + raise NotImplementedError( + f"Coordinate system {attrs['Coordinates']} not supported." + ) + if remap is None: + remap = { + "coords": ( + remap_coords_cart + if coord_system == CoordinateSystem.XYZ + else remap_coords_sph + ), + "fields": ( + remap_fields_cart + if coord_system == CoordinateSystem.XYZ + else remap_fields_sph + ), + "particles": ( + remap_prtl_quantities_cart + if coord_system == CoordinateSystem.XYZ + else remap_prtl_quantities_sph + ), + } + break + + if coord_system is None: + raise ValueError("No coordinate system found in the data.") + + self.__coordinate_system = coord_system + + super(Data, self).__init__(path=path, reader=self.__reader, remap=remap) + try: + self.__diagnostics = Diagnostics(path) + except Exception as e: + logging.warning(f"Failed to read diagnostics: {e}") + self.__diagnostics = None + + def makeMovie( + self, + plot: Callable, + time: Optional[List[float]] = None, + num_cpus: Optional[int] = None, + **movie_kwargs: Any, + ) -> bool: + """Create animation with provided plot function. + + Parameters + ---------- + plot : callable + A function that takes a single argument (time in physical units) and produces a plot. + time : array_like, optional + An array of time values to use for the animation. If not provided, the entire time range will be used. + num_cpus : int, optional + The number of CPUs to use for parallel processing. If None, it will use all available CPUs. + **movie_kwargs : dict + Additional keyword arguments to pass to the movie creation function. + + Returns + ------- + bool + True if the movie was created successfully, False otherwise. + """ + if time is None: + if self.fields_defined: + time = self.fields.t.values + elif self.particles_defined and self.particles is not None: + time = list(self.particles.times) + else: + raise ValueError("No time values found.") + assert time is not None, "Time values must be provided." + name: str = "" + if self.attrs.get("simulation.name", None) is None: + name = movie_kwargs.pop("name", "movie") + else: + name_b = self.attrs.get("simulation.name") + if isinstance(name_b, bytes): + name = name_b.decode("utf-8") + else: + name = str(name_b) + return makeFramesAndMovie( + name=name, + data=self, + plot=plot, + times=time, + num_cpus=num_cpus, + **movie_kwargs, + ) + + @property + def coordinate_system(self) -> CoordinateSystem: + """CoordinateSystem: The coordinate system of the data.""" + return self.__coordinate_system + + @property + def attrs(self) -> Dict[str, Any]: + """dict[str, Any]: The attributes of the data.""" + return self.__attrs + + @property + def diagnostics(self) -> Union[pd.DataFrame, None]: + """pd.DataFrame or None: The diagnostics output if .out file is found, None otherwise.""" + if self.__diagnostics is None: + return None + return self.__diagnostics.df + + def to_str(self) -> str: + """str: String representation of the all the enclosed dataframes.""" + + string = "" + if self.fields_defined: + string += "FieldsDataset:\n" + string += "==============\n" + string += f"| Coordinates:\n| {self.coordinate_system.value}\n|\n" + string += f"| Data axes:\n| {compactify(self.fields.indexes.keys())}\n|\n" + delta_t = ( + self.fields.coords["t"].values[1] - self.fields.coords["t"].values[0] + ) / (self.fields.coords["s"].values[1] - self.fields.coords["s"].values[0]) + string += f"| - dt: {delta_t:.2e}\n" + for key in self.fields.coords.keys(): + crd = self.fields.coords[key].values + fmt = "" + if key != "s": + fmt = ".2f" + string += f"| - {key}: {crd.min():{fmt}} -> {crd.max():{fmt}} [{len(crd)}]\n" + string += "|\n" + string += f"| Quantities:\n| {compactify(sorted(map(str, self.fields.data_vars.keys())))}\n|\n" + string += f"| Total size: {ToHumanReadable(self.fields.nbytes)}\n\n" + else: + string += "FieldsDataset:\n" + string += "==============\n" + string += " empty\n\n" + if self.particles_defined and self.particles is not None: + species = sorted(self.particles.species) + string += "ParticleDataset:\n" + string += "================\n" + string += f"| Species:\n| {compactify(species)}\n|\n" + string += f"| Timesteps:\n| {len(self.particles.times)}\n|\n" + string += f"| Quantities:\n| {compactify(self.particles.columns)}\n|\n" + string += f"| Total size: {ToHumanReadable(self.particles.nbytes)}\n|\n" + string += self.help_particles("| ") + string += "\n" + else: + string += "ParticleDataset:\n" + string += "================\n" + string += " empty\n\n" + if self.spectra_defined and self.spectra is not None: + string += "SpectraDataset:\n" + string += "===============\n" + string += ( + f"| Data axes:\n| {compactify(self.spectra.indexes.keys())}\n|\n" + ) + delta_t = ( + self.spectra.coords["t"].values[1] - self.spectra.coords["t"].values[0] + ) / ( + self.spectra.coords["s"].values[1] - self.spectra.coords["s"].values[0] + ) + string += f"| - dt: {delta_t:.2e}\n" + for key in self.spectra.coords.keys(): + crd = self.spectra.coords[key].values + fmt = "" + if key != "s": + fmt = ".2f" + string += f"| - {key}: {crd.min():{fmt}} -> {crd.max():{fmt}} [{len(crd)}]\n" + string += "|\n" + string += f"| Quantities:\n| {compactify(sorted(map(str, self.spectra.data_vars.keys())))}\n|\n" + string += f"| Total size: {ToHumanReadable(self.spectra.nbytes)}\n|\n" + string += self.help_spectra("| ") + else: + string += "SpectraDataset:\n" + string += "===============\n" + string += " empty\n\n" + + return string + + @override + def __str__(self) -> str: + return self.to_str() + + @override + def __repr__(self) -> str: + return self.to_str() diff --git a/legacy/containers/diagnostics.py b/legacy/containers/diagnostics.py new file mode 100644 index 0000000..98fe33b --- /dev/null +++ b/legacy/containers/diagnostics.py @@ -0,0 +1,108 @@ +from typing import Union +import pandas as pd + + +class Diagnostics: + df: Union[pd.DataFrame, None] + + def __init__(self, path: str): + import os + import logging + import re + + outfiles = [o for o in os.listdir(path) if o.endswith(".out")] + if len(outfiles) == 0: + logging.warning(f"No .out files found in {path}") + self.df = None + else: + self.outfile = os.path.join(path, outfiles[0]) + + data = {} + + with open(self.outfile, "r") as f: + content = f.read() + steps = re.findall(r"Step:\s+(\d+)\.+\[", content) + times = re.findall(r"Time:\s+([\d.]+\d)\.+\[", content) + substeps = re.findall(r"\s+([A-Za-z]+)\.+([\d.]+)\s+([mµn]?s)", content) + species = re.findall( + r"\s+species\s+(\d+)\s+\(.+\)\.+([\deE+-.]+)(\s+\d+\%\s:\s\d+\%\s+)?([\deE+-.]+)?( : )?([\deE+-.]+)?", + content, + ) + + data["steps"] = [] + for step in steps: + data["steps"].append(int(step)) + + data["times"] = [] + for time in times: + data["times"].append(float(time)) + + assert len(data["steps"]) == len( + data["times"] + ), "Number of steps and times do not match" + + data["substeps"] = {} + for substep in substeps: + if substep[0] not in data["substeps"].keys(): + data["substeps"][substep[0]] = [] + + def to_ns(value: float, unit: str) -> float: + if unit == "s": + return value * 1e9 + elif unit == "ms": + return value * 1e6 + elif unit == "µs": + return value * 1e3 + elif unit == "ns": + return value + else: + raise ValueError(f"Unknown time unit: {unit}") + + data["substeps"][substep[0]].append( + to_ns(float(substep[1]), substep[2]) + ) + + for key in data["substeps"].keys(): + assert len(data["substeps"][key]) == len( + data["steps"] + ), f"Number of substep entries for {key} does not match number of steps" + + data["species"] = {} + data["species_min"] = {} + data["species_max"] = {} + for specie in species: + if specie[0] not in data["species"].keys(): + data["species"][specie[0]] = [] + data["species_min"][specie[0]] = [] + data["species_max"][specie[0]] = [] + data["species"][specie[0]].append(int(float(specie[1]))) + if len(specie) == 6 and specie[3] != "" and specie[5] != "": + data["species_min"][specie[0]].append(int(float(specie[3]))) + data["species_max"][specie[0]].append(int(float(specie[5]))) + + for key in data["species"].keys(): + assert len(data["species"][key]) == len( + data["steps"] + ), f"Number of species entries for {key} does not match number of steps" + assert (len(data["species_min"][key]) == len(data["steps"])) or ( + len(data["species_min"][key]) == 0 + ), f"Number of species min entries for {key} does not match number of steps" + assert (len(data["species_max"][key]) == len(data["steps"])) or ( + len(data["species_max"][key]) == 0 + ), f"Number of species max entries for {key} does not match number of steps" + + self.df = pd.DataFrame(index=data["steps"]) + self.df["Step"] = data["steps"] + self.df["Time"] = data["times"] + for key in data["substeps"].keys(): + self.df[key] = data["substeps"][key] + for key in data["species"].keys(): + self.df[f"species_{key}"] = data["species"][key] + if ( + len(data["species_min"][key]) > 0 + and len(data["species_max"][key]) > 0 + ): + self.df[f"species_{key}_min"] = data["species_min"][key] + self.df[f"species_{key}_max"] = data["species_max"][key] + + del data diff --git a/legacy/containers/fields.py b/legacy/containers/fields.py new file mode 100644 index 0000000..6cbb5b6 --- /dev/null +++ b/legacy/containers/fields.py @@ -0,0 +1,161 @@ +from typing import Any + +import dask +import dask.array as da +import xarray as xr +from tqdm import tqdm + +from nt2.containers.container import BaseContainer +from nt2.utils import Layout + + +class Fields(BaseContainer): + """Parent class to manage the fields dataframe.""" + + def _read_field(self, layout: Layout, field: str, step: int) -> Any: + """Reads a field from the data. + + This is a dask-delayed function used further to build the dataset. + + Parameters + ---------- + layout : Layout + Layout of the field. + field : str + Field to read. + step : int + Step to read. + + Returns + ------- + Any + Field data. + + """ + if layout == Layout.L: + return self.reader.ReadArrayAtTimestep(self.path, "fields", field, step) + else: + return self.reader.ReadArrayAtTimestep(self.path, "fields", field, step).T + + def __init__( + self, + **kwargs: Any, + ) -> None: + """Initializer for the Fields class. + + Parameters + ---------- + **kwargs : dict + Keyword arguments to be passed to the parent BaseContainer class. + + """ + super(Fields, self).__init__(**kwargs) + if self.reader.DefinesCategory(self.path, "fields", self.valid_files["fields"]): + self.__fields_defined = True + self.__fields = self._read_fields() + else: + self.__fields_defined = False + self.__fields = xr.Dataset() + + @property + def fields_defined(self) -> bool: + """bool: Whether the fields category is defined.""" + return self.__fields_defined + + @property + def fields(self) -> xr.Dataset: + """xr.Dataset: The fields dataframe.""" + return self.__fields + + def _read_fields(self) -> xr.Dataset: + """Helper function to read the fields dataframe.""" + self.reader.VerifySameCategoryNames( + self.path, "fields", "f", self.valid_steps["fields"] + ) + self.reader.VerifySameFieldShapes(self.path, self.valid_steps["fields"]) + self.reader.VerifySameFieldLayouts(self.path, self.valid_steps["fields"]) + + field_names = self.reader.ReadCategoryNamesAtTimestep( + self.path, "fields", "f", self.valid_steps["fields"][0] + ) + + first_step = self.valid_steps["fields"][0] + first_name = next(iter(field_names)) + layout = self.reader.ReadFieldLayoutAtTimestep(self.path, first_step) + shape = self.reader.ReadArrayShapeAtTimestep( + self.path, "fields", first_name, first_step + ) + coords = self.reader.ReadFieldCoordsAtTimestep(self.path, first_step) + coords = {k: coords[k] for k in sorted(coords.keys())[::-1]} + # rename coordinates if remap is provided + if self.remap is not None and "coords" in self.remap: + new_coords = {} + for coord in coords.keys(): + new_coords[self.remap["coords"](coord)] = coords[coord] + coords = new_coords + + times = self.reader.ReadPerTimestepVariable( + self.path, + "fields", + "Time", + "t", + self.valid_files["fields"], + ) + steps = self.reader.ReadPerTimestepVariable( + self.path, + "fields", + "Step", + "s", + self.valid_files["fields"], + ) + + edge_coords = self.reader.ReadEdgeCoordsAtTimestep(self.path, first_step) + new_edge_coords = {} + for coord in edge_coords.keys(): + assoc_x = ( + coord[:-1] + if (self.remap is None or "coords" not in self.remap) + else self.remap["coords"](coord[:-1]) + ) + new_edge_coords[assoc_x + "_min"] = (assoc_x, edge_coords[coord][:-1]) + new_edge_coords[assoc_x + "_max"] = (assoc_x, edge_coords[coord][1:]) + edge_coords = new_edge_coords + + all_dims = {**times, **coords}.keys() + all_coords = {**times, **coords, "s": ("t", steps["s"]), **edge_coords} + + return xr.Dataset( + { + ( + remapped_name := ( + self.remap["fields"](name) + if (self.remap is not None and "fields" in self.remap) + else name + ) + ): xr.DataArray( + da.stack( + [ + da.from_delayed( + dask.delayed(self._read_field)(layout, name, step), + shape=shape[:: -1 if layout == Layout.R else 1], + dtype="float", + ) + for step in tqdm( + self.valid_steps["fields"], + desc="steps", + position=1, + leave=False, + ) + ], + axis=0, + ), + name=remapped_name, + dims=all_dims, + coords=all_coords, + ) + for name in tqdm(field_names, desc="fields", position=0, leave=False) + }, + attrs=self.reader.ReadAttrsAtTimestep( + path=self.path, category="fields", step=first_step + ), + ) diff --git a/legacy/containers/particles.py b/legacy/containers/particles.py new file mode 100644 index 0000000..910cfbd --- /dev/null +++ b/legacy/containers/particles.py @@ -0,0 +1,814 @@ +from typing import ( + Any, + Callable, + List, + Optional, + Sequence, + Tuple, + Literal, + Union, + Dict, + Type, +) +import numpy.typing as npt +from copy import copy + +import dask +import dask.dataframe as dd +import pandas as pd +import numpy as np +from tqdm import tqdm + +import matplotlib.pyplot as plt +import matplotlib.axes as maxes + +from nt2.containers.container import BaseContainer + + +IntSelector = Union[int, Sequence[int], slice, Tuple[int, int]] +FloatSelector = Union[float, slice, Sequence[float], Tuple[float, float]] + + +class Selection: + def __init__( + self, + type: Literal["value", "range", "list"], + value: Optional[Union[int, float, list, tuple]] = None, + ): + self.type = type + self.value = value + + def intersect(self, other: "Selection") -> "Selection": + if self.value is None: + return copy(other) + elif other.value is None: + return copy(self) + if self.type == "value" and other.type == "value": + if self.value == other.value: + return Selection("value", self.value) + else: + return Selection("value") + elif self.type == "value" and other.type == "list": + assert isinstance(other.value, list), "other.value must be a list" + if self.value in other.value: + return Selection("value", self.value) + else: + return Selection("value") + elif self.type == "value" and other.type == "range": + assert isinstance(other.value, tuple) and len(other.value) == 2, ( + "other.value must be a tuple of length 2" + ) + lo, hi = other.value + if lo <= self.value < hi: + return Selection("value", self.value) + else: + return Selection("value") + elif self.type == "list" and other.type == "value": + return other.intersect(self) + elif self.type == "list" and other.type == "list": + assert isinstance(self.value, list), "self.value must be a list" + assert isinstance(other.value, list), "other.value must be a list" + new_values = [v for v in self.value if v in other.value] + return Selection("list", new_values) + elif self.type == "list" and other.type == "range": + assert isinstance(other.value, tuple) and len(other.value) == 2, ( + "other.value must be a tuple of length 2" + ) + assert isinstance(self.value, list), "self.value must be a list" + lo, hi = other.value + new_values = [v for v in self.value if lo <= v <= hi] + return Selection("list", new_values) + elif self.type == "range" and other.type == "value": + return other.intersect(self) + elif self.type == "range" and other.type == "list": + return other.intersect(self) + elif self.type == "range" and other.type == "range": + assert isinstance(self.value, tuple) and len(self.value) == 2, ( + "self.value must be a tuple of length 2" + ) + assert isinstance(other.value, tuple) and len(other.value) == 2, ( + "other.value must be a tuple of length 2" + ) + lo1, hi1 = self.value + lo2, hi2 = other.value + new_lo = max(lo1, lo2) + new_hi = min(hi1, hi2) + if new_lo <= new_hi: + return Selection("range", (new_lo, new_hi)) + else: + return Selection("value") + else: + raise ValueError(f"Unknown selection types: {self.type}, {other.type}") + + def __repr__(self) -> str: + if self.type == "value": + return "all" if self.value is None else f"{self.value:.3g}" + elif self.type == "range": + if self.value is None: + return "all" + else: + assert isinstance(self.value, tuple) and len(self.value) == 2, ( + "value must be a tuple of length 2" + ) + lo, hi = self.value + lo_str = "..." if lo is None or lo == -np.inf else f"{lo:.3g}" + hi_str = "..." if hi is None or hi == np.inf else f"{hi:.3g}" + return f"[ {lo_str} -> {hi_str} ]" + elif self.type == "list": + assert isinstance(self.value, list), "value must be a list" + return "{ " + ", ".join(f"{v:.3g}" for v in self.value) + " }" + else: + return "InvalidSelection" + + def __str__(self) -> str: + return self.__repr__() + + +def _coerce_selector_to_mask( + s: Union[IntSelector, FloatSelector], + series: Any, + inclusive_tuple: bool = True, + method="exact", +): + from operator import ior + from functools import reduce + + if isinstance(s, slice): + lo = s.start if s.start is not None else -np.inf + hi = s.stop if s.stop is not None else np.inf + step = s.step + mask = (series >= lo) & (series <= hi) + if step not in (None, 1): + mask = mask & (((series - lo) % step) == 0) + return mask, ("range", (lo, hi)) + elif isinstance(s, tuple) and len(s) == 2 and inclusive_tuple: + lo, hi = s + if lo is None: + lo = -np.inf + if hi is None: + hi = np.inf + return (series >= lo) & (series <= hi), ("range", (lo, hi)) + elif isinstance(s, (list, tuple, np.ndarray, pd.Index, pd.Series)): + if method == "exact": + return series.isin(list(s)), ("list", list(s)) + else: + return reduce( + ior, [np.abs(series - v) == np.abs(series - v).min() for v in s] + ), ("list", list(s)) + else: + if method == "exact": + return series == s, ("value", s) + else: + return np.abs(series - s) == np.abs(series - s).min(), ("value", s) + + +def _attach_columns( + part: pd.DataFrame, + cols_tuple, + read_column, + metadtypes, +) -> pd.DataFrame: + if len(part) == 0: + for c in cols_tuple: + part[c] = np.array([], dtype=metadtypes[c]) + return part + st_val = int(part["st"].iloc[0]) + + arrays = {c: read_column(st_val, c) for c in cols_tuple} + + sel = part["row"].to_numpy() + for c in cols_tuple: + part[c] = np.asarray(arrays[c])[sel] + return part + + +class ParticleDataset: + steps: npt.NDArray[np.int64] + times: npt.NDArray[np.float64] + colnames: List[str] + + def __init__( + self, + species: List[int], + steps: npt.NDArray[np.int64], + times: npt.NDArray[np.float64], + colnames: List[str], + read_column: Callable[ + [int, str], npt.NDArray[Union[np.float64, np.int64, np.float32, np.int32]] + ], + fprec: Optional[Type] = np.float32, + selection: Optional[Dict[str, Selection]] = None, + ddf_index: Optional[dd.DataFrame] = None, + ): + self.species = species + self.steps = steps + self.times = times + self.colnames = colnames + + self.read_column = read_column + self.fprec = fprec + self.index_cols = ("id", "sp") + self._all_columns_cache: Optional[List[str]] = None + + if selection is not None: + self.selection = selection + else: + self.selection = { + "t": Selection("range"), + "st": Selection("range"), + "sp": Selection("range"), + "id": Selection("range"), + } + + self._dtypes = { + "id": np.int64, + "sp": np.int32, + "row": np.int64, + "st": np.int64, + "t": fprec, + "x": fprec, + "y": fprec, + "z": fprec, + "ux": fprec, + "uy": fprec, + "uz": fprec, + "r": fprec, + "th": fprec, + "ph": fprec, + "ur": fprec, + "uth": fprec, + "uph": fprec, + } + + if ddf_index is not None: + self._ddf_index = ddf_index + else: + self._ddf_index = self._build_index_ddf() + + @property + def ddf(self) -> dd.DataFrame: + return self._ddf_index + + @property + def nbytes(self) -> int: + return self.ddf.memory_usage(index=True, deep=True).sum().compute() + + @property + def columns(self) -> List[str]: + if self._all_columns_cache is None: + self._all_columns_cache = self.colnames + return self._all_columns_cache + + def sel( + self, + t: Optional[Union[IntSelector, FloatSelector]] = None, + st: Optional[IntSelector] = None, + sp: Optional[IntSelector] = None, + id: Optional[IntSelector] = None, + method: str = "exact", + ) -> "ParticleDataset": + ddf = self._ddf_index + new_selection = {k: copy(v) for k, v in self.selection.items()} + if st is not None: + ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( + st, ddf["st"], method="exact" + ) + ddf = ddf[ddf_sel] + new_selection["st"] = new_selection["st"].intersect( + Selection(sel_type, sel_value) + ) + if t is not None: + ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( + t, ddf["t"], method=method + ) + ddf = ddf[ddf_sel] + new_selection["t"] = new_selection["t"].intersect( + Selection(sel_type, sel_value) + ) + if sp is not None: + ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( + sp, ddf["sp"], method="exact" + ) + ddf = ddf[ddf_sel] + new_selection["sp"] = new_selection["sp"].intersect( + Selection(sel_type, sel_value) + ) + if id is not None: + ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( + id, ddf["id"], method="exact" + ) + ddf = ddf[ddf_sel] + new_selection["id"] = new_selection["id"].intersect( + Selection(sel_type, sel_value) + ) + + return ParticleDataset( + species=self.species, + steps=self.steps, + times=self.times, + colnames=self.colnames, + read_column=self.read_column, + fprec=self.fprec, + selection=new_selection, + ddf_index=ddf, + ) + + def isel( + self, t: Optional[IntSelector] = None, st: Optional[IntSelector] = None + ) -> "ParticleDataset": + ddf = self._ddf_index + new_selection = {k: v for k, v in self.selection.items()} + for t_or_s, t_or_s_str, t_or_s_arr in zip( + [t, st], ["t", "st"], [self.times, self.steps] + ): + if t_or_s is not None: + if isinstance(t_or_s, slice): + lo = t_or_s.start if t_or_s.start is not None else 0 + hi = t_or_s.stop if t_or_s.stop is not None else -1 + ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( + slice(t_or_s_arr[lo], t_or_s_arr[hi]), + ddf[t_or_s_str], + method="exact", + ) + ddf = ddf[ddf_sel] + new_selection[t_or_s_str] = new_selection[t_or_s_str].intersect( + Selection(sel_type, sel_value) + ) + elif isinstance(t_or_s, (list, tuple, np.ndarray, pd.Index, pd.Series)): + ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( + [t_or_s_arr[ti] for ti in t_or_s], + ddf[t_or_s_str], + method="exact", + ) + ddf = ddf[ddf_sel] + new_selection[t_or_s_str] = new_selection[t_or_s_str].intersect( + Selection(sel_type, sel_value) + ) + else: + ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( + t_or_s_arr[t_or_s], ddf[t_or_s_str], method="exact" + ) + ddf = ddf[ddf_sel] + new_selection[t_or_s_str] = new_selection[t_or_s_str].intersect( + Selection(sel_type, sel_value) + ) + return ParticleDataset( + species=self.species, + steps=self.steps, + times=self.times, + colnames=self.colnames, + read_column=self.read_column, + fprec=self.fprec, + selection=new_selection, + ddf_index=ddf, + ) + + def _load_index_partition(self, st: int, t: float, index_cols: Tuple[str, ...]): + cols = {c: self.read_column(st, c) for c in index_cols} + n = len(next(iter(cols.values()))) + df = pd.DataFrame(cols) + df["st"] = np.asarray(st, dtype=np.int64) + df["t"] = np.asarray(t, dtype=float) + df["row"] = np.arange(n, dtype=np.int64) + return df + + def _build_index_ddf(self) -> dd.DataFrame: + delayed_parts = [ + dask.delayed(self._load_index_partition)(st, t, self.index_cols) + for st, t in zip(self.steps, self.times) + ] + + meta = pd.DataFrame( + { + **{ + c: np.array([], dtype=self._dtypes.get(c, "O")) + for c in self.index_cols + }, + "st": np.array([], dtype=self._dtypes.get("st", np.int64)), + "t": np.array([], dtype=self._dtypes.get("t", np.int64)), + "row": np.array([], dtype=self._dtypes.get("row", np.int64)), + } + ) + + ddf = dd.from_delayed(delayed_parts, meta=meta) + return ddf + + def load(self, cols: Optional[Sequence[str]] = None) -> pd.DataFrame: + if cols is None: + cols = self.columns + + cols = [c for c in cols if c not in ("t", "st", "row")] + + meta_dict = { + c: np.array([], dtype=self._dtypes.get(c, np.float64)) for c in cols + } + meta = self._ddf_index._meta.assign(**meta_dict) + + cols_tuple = tuple(cols) + + return ( + self._ddf_index.map_partitions( + _attach_columns, + cols_tuple=cols_tuple, + read_column=self.read_column, + metadtypes=meta.dtypes, + meta=meta, + ) + .compute() + .drop(columns=["row"]) + ) + + def help(self, prepend="") -> str: + ret = f"{prepend}- use .sel(...) to select particles based on criteria:\n" + ret += f"{prepend} t : time (float)\n" + ret += f"{prepend} st : step (int)\n" + ret += f"{prepend} sp : species (int)\n" + ret += f"{prepend} id : particle id (int)\n{prepend}\n" + ret += f"{prepend} # example:\n" + ret += f"{prepend} # .sel(t=slice(10.0, 20.0), sp=[1, 2, 3], id=[42, 22])\n{prepend}\n" + ret += f"{prepend}- use .isel(...) to select particles based on output step:\n" + ret += f"{prepend} t : timestamp index (int)\n" + ret += f"{prepend} st : step index (int)\n{prepend}\n" + ret += f"{prepend} # example:\n" + ret += f"{prepend} # .isel(t=-1)\n" + ret += f"{prepend}\n" + ret += f"{prepend}- .sel and .isel can be chained together:\n{prepend}\n" + ret += f"{prepend} # example:\n" + ret += f"{prepend} # .isel(t=-1).sel(sp=1).sel(id=[55, 66])\n{prepend}\n" + ret += f"{prepend}- use .load(cols=[...]) to load data into a pandas DataFrame (`cols` defaults to all columns)\n{prepend}\n" + ret += f"{prepend} # example:\n" + ret += f"{prepend} # .sel(...).load()\n" + return ret + + def __repr__(self) -> str: + ret = "ParticleDataset:\n" + ret += "================\n" + ret += f"Variables:\n {self.columns}\n\n" + ret += "Current selection:\n" + for k, v in self.selection.items(): + ret += f" {k:<5} : {v}\n" + ret += "\nHelp:\n" + ret += "-----\n" + ret += f"{self.help()}" + return ret + + def __str__(self) -> str: + return self.__repr__() + + def spectrum_plot( + self, + ax: Optional[maxes.Axes] = None, + bins: Optional[npt.NDArray] = None, + quantity: Optional[Callable[[pd.DataFrame], npt.NDArray]] = None, + ): + if ax is None: + ax = plt.gca() + + def _uSqr_cart(df: pd.DataFrame): + return np.sum( + [ + np.asarray(df[c].to_numpy(), dtype=np.float64) ** 2 + for c in ["ux", "uy", "uz"] + ], + axis=0, + ) + + def _quantity_cart(df: pd.DataFrame): + uSqr = _uSqr_cart(df) + return uSqr * np.sqrt(1.0 + uSqr) + + def _uSqr_sph(df: pd.DataFrame): + return np.sum( + [ + np.asarray(df[c].to_numpy(), dtype=np.float64) ** 2 + for c in ["ur", "uth", "uph"] + ], + axis=0, + ) + + def _quantity_sph(df: pd.DataFrame): + uSqr = _uSqr_sph(df) + return uSqr * np.sqrt(1.0 + uSqr) + + if "ux" in self.columns: + cols = ["ux", "uy", "uz"] + if quantity is None: + quantity = _quantity_cart + else: + cols = ["ur", "uth", "uph"] + if quantity is None: + quantity = _quantity_sph + df = self.load(cols=["sp", *cols]) + species = sorted(df["sp"].unique()) + arrays = df.groupby("sp").apply(quantity, include_groups=False) + if bins is None: + bins = np.logspace(0, 4, 100) + hists = {sp: np.histogram(arrays[sp], bins=bins)[0] for sp in species} + bins = 0.5 * (bins[1:] + bins[:-1]) + for sp in species: + ax.loglog(bins, hists[sp], label=f"{sp}") + if bins.min() > 0 and bins.max() / bins.min() > 100: + ax.set(xscale="log", yscale="log") + + def phase_plot( + self, + ax: Optional[maxes.Axes] = None, + x_quantity: Optional[Callable[[pd.DataFrame], npt.NDArray]] = None, + y_quantity: Optional[Callable[[pd.DataFrame], npt.NDArray]] = None, + xy_bins: Optional[Tuple[npt.NDArray, npt.NDArray]] = None, + **kwargs: Any, + ): + if ax is None: + ax = plt.gca() + + def _xquantity_cart(df: pd.DataFrame): + return np.asarray(df["x"].to_numpy(), dtype=np.float64) + + def _yquantity_cart(df: pd.DataFrame): + return np.asarray(df["ux"].to_numpy(), dtype=np.float64) + + def _xquantity_sph(df: pd.DataFrame): + return np.asarray(df["r"].to_numpy(), dtype=np.float64) + + def _yquantity_sph(df: pd.DataFrame): + return np.asarray(df["ur"].to_numpy(), dtype=np.float64) + + if "ux" in self.columns: + cols = ["ux", "uy", "uz"] + for c in "xyz": + if c in self.columns: + cols.append(c) + if x_quantity is None: + x_quantity = _xquantity_cart + if y_quantity is None: + y_quantity = _yquantity_cart + else: + cols = ["ur", "uth", "uph"] + for c in ["r", "th", "ph"]: + if c in self.columns: + cols.append(c) + if x_quantity is None: + x_quantity = _xquantity_sph + if y_quantity is None: + y_quantity = _yquantity_sph + + df = self.load(cols=[*cols]) + x_array = x_quantity(df) + y_array = y_quantity(df) + + if xy_bins is None: + x_bins = np.linspace(x_array.min(), x_array.max(), 100) + y_bins = np.linspace(y_array.min(), y_array.max(), 100) + xy_bins = (x_bins, y_bins) + else: + x_bins, y_bins = xy_bins + + h2d, xedges, yedges = np.histogram2d(x_array, y_array, bins=[x_bins, y_bins]) + X, Y = np.meshgrid( + 0.5 * (xedges[1:] + xedges[:-1]), 0.5 * (yedges[1:] + yedges[:-1]) + ) + pcm = ax.pcolormesh( + X, + Y, + h2d.T, + shading="auto", + rasterized=True, + **kwargs, + ) + return pcm + + +class Particles(BaseContainer): + """Parent class to manage the particles dataframe.""" + + __particles_defined: bool + __particles: Optional[ParticleDataset] + quantities: List[str] + sp_with_idx: List[int] + sp_without_idx: List[int] + + def __init__(self, **kwargs: Any) -> None: + """Initializer for the Particles class. + + Parameters + ---------- + **kwargs : dict + Keyword arguments to be passed to the parent BaseContainer class. + + """ + super(Particles, self).__init__(**kwargs) + if ( + self.reader.DefinesCategory( + self.path, "particles", self.valid_files["particles"] + ) + and self.particles_present + ): + self.__particles_defined = True + + valid_steps = self.nonempty_steps + quantities_ = [ + self.reader.ReadCategoryNamesAtTimestep( + self.path, "particles", "p", step + ) + for step in valid_steps + ] + self.quantities = sorted( + np.unique([q for qtys in quantities_ for q in qtys]) + ) + + unique_quantities = sorted( + list( + set( + f"{q}".split("_")[0] + for q in self.quantities + if not q.startswith("pIDX") and not q.startswith("pRNK") + ) + ) + ) + all_species = sorted( + list(set([int(f"{q}".split("_")[1]) for q in self.quantities])) + ) + + self.sp_with_idx = sorted( + [ + int(f"{q}".split("_")[1]) + for q in self.quantities + if f"{q}".startswith("pIDX") + ] + ) + self.sp_without_idx = sorted( + [sp for sp in all_species if sp not in self.sp_with_idx] + ) + + self.__particles = ParticleDataset( + species=all_species, + steps=self.valid_steps["particles"], + times=self.reader.ReadPerTimestepVariable( + self.path, + "particles", + "Time", + "t", + self.valid_files["particles"], + )["t"], + colnames=[ + ( + self.remap["particles"](q) + if (self.remap is not None and "particles" in self.remap) + else q + ) + for q in unique_quantities + ] + + ["id", "sp"], + read_column=self._read_column, + ) + else: + self.__particles_defined = False + self.__particles = None + + @property + def particles_present(self) -> bool: + """bool: Whether the particles are present in any of the timesteps.""" + return len(self.nonempty_steps) > 0 + + @property + def nonempty_steps(self) -> List[int]: + """list[int]: List of timesteps that contain particles data.""" + return [ + step + for step in self.valid_steps["particles"] + if len( + set( + q.split("_")[0] + for q in self.reader.ReadCategoryNamesAtTimestep( + self.path, "particles", "p", step + ) + if q.startswith("p") + ) + ) + > 0 + ] + + @property + def particles_defined(self) -> bool: + """bool: Whether the particles category is defined.""" + return self.__particles_defined + + @property + def particles(self) -> Optional[ParticleDataset]: + """Returns the particles data. + + Returns + ------- + ParticleDataset + Dictionary of datasets for each step. + + """ + return self.__particles + + def help_particles(self, prepend: str = "") -> str: + return self.particles.help(prepend) if self.particles is not None else "" + + def _get_count(self, step: int, sp: int) -> np.int64: + try: + return np.int64( + self.reader.ReadArrayShapeAtTimestep( + self.path, "particles", f"pX1_{sp}", step + )[0] + ) + except Exception: + return np.int64(0) + + def _species_has_quantity(self, read_colname: str, step: int, sp: int) -> bool: + return f"{read_colname}_{sp}" in self.reader.ReadCategoryNamesAtTimestep( + self.path, "particles", "p", step + ) + + def _get_quantity_for_species( + self, + read_colname: str, + step: int, + sp: int, + ) -> npt.NDArray[Union[np.float64, np.int64]]: + if f"{read_colname}_{sp}" in self.quantities: + return self.reader.ReadArrayAtTimestep( + self.path, "particles", f"{read_colname}_{sp}", step + ) + else: + return np.zeros(self._get_count(step, sp)) * np.nan + + def _read_column( + self, step: int, colname: str + ) -> npt.NDArray[Union[np.float64, np.int64, np.float32, np.int32]]: + read_colname = None + if colname == "id": + idx = np.concatenate( + [ + self.reader.ReadArrayAtTimestep( + self.path, "particles", f"pIDX_{sp}", step + ).astype(np.int64) + for sp in self.sp_with_idx + ] + + [ + np.zeros(self._get_count(step, sp), dtype=np.int64) - 100 + for sp in self.sp_without_idx + ] + ) + if ( + len(self.sp_with_idx) > 0 + and f"pRNK_{self.sp_with_idx[0]}" in self.quantities + ): + rnk = np.concatenate( + [ + self.reader.ReadArrayAtTimestep( + self.path, "particles", f"pRNK_{sp}", step + ).astype(np.int64) + for sp in self.sp_with_idx + ] + + [ + np.zeros(self._get_count(step, sp), dtype=np.int64) - 100 + for sp in self.sp_without_idx + ] + ) + return (idx + rnk) * (idx + rnk + 1) // 2 + rnk + else: + return idx + elif colname == "x" or colname == "r": + read_colname = "pX1" + elif colname == "y" or colname == "th": + read_colname = "pX2" + elif colname == "z" or colname == "ph": + read_colname = "pX3" + elif colname == "ux" or colname == "ur": + read_colname = "pU1" + elif colname == "uy" or colname == "uth": + read_colname = "pU2" + elif colname == "uz" or colname == "uph": + read_colname = "pU3" + elif colname == "w": + read_colname = "pW" + elif colname == "sp": + return np.concatenate( + [ + np.zeros(self._get_count(step, sp), dtype=np.int32) + sp + for sp in self.sp_with_idx + ] + + [ + np.zeros(self._get_count(step, sp), dtype=np.int32) + sp + for sp in self.sp_without_idx + ] + ) + else: + read_colname = f"p{colname}" + + return np.concatenate( + [ + self._get_quantity_for_species(read_colname, step, sp) + for sp in self.sp_with_idx + if self._species_has_quantity(read_colname, step, sp) + ] + + [ + self._get_quantity_for_species(read_colname, step, sp) + for sp in self.sp_without_idx + if self._species_has_quantity(read_colname, step, sp) + ] + ) diff --git a/legacy/containers/spectra.py b/legacy/containers/spectra.py new file mode 100644 index 0000000..ff47dd1 --- /dev/null +++ b/legacy/containers/spectra.py @@ -0,0 +1,164 @@ +from typing import Any + +import dask +import dask.array as da +import xarray as xr +import numpy as np +from tqdm import tqdm + +from nt2.containers.container import BaseContainer +from nt2.readers.base import BaseReader + + +class Spectra(BaseContainer): + """Parent class to manager the spectra dataframe.""" + + @staticmethod + def read_spectrum(path: str, reader: BaseReader, spectrum: str, step: int) -> Any: + """Reads a spectrum from the data. + + This is a dask-delayed function used further to build the dataset. + + Parameters + ---------- + path : str + Main path to the data. + reader : BaseReader + Reader to use to read the data. + spectrum : str + Spectrum array to read. + step : int + Step to read. + + Returns + ------- + Any + Spectrum data. + + """ + return reader.ReadArrayAtTimestep(path, "spectra", spectrum, step) + + def __init__(self, **kwargs: Any) -> None: + super(Spectra, self).__init__(**kwargs) + if self.reader.DefinesCategory( + self.path, + "spectra", + self.valid_steps["spectra"], + ): + self.__spectra_defined = True + self.__spectra = self.__read_spectra() + else: + self.__spectra_defined = False + self.__spectra = xr.Dataset() + + @property + def spectra_defined(self) -> bool: + """bool: Whether the spectra category is defined.""" + return self.__spectra_defined + + @property + def spectra(self) -> xr.Dataset: + """xr.Dataset: The spectra dataframe.""" + return self.__spectra + + def __read_spectra(self) -> xr.Dataset: + self.reader.VerifySameCategoryNames( + self.path, + "spectra", + "s", + self.valid_steps["spectra"], + ) + first_step = self.valid_steps["spectra"] + spectra_names = self.reader.ReadCategoryNamesAtTimestep( + self.path, "spectra", "s", first_step + ) + spectra_names = set(s for s in sorted(spectra_names) if s.startswith("sN")) + ebin_name = "sEbn" + first_spectrum_name = next(iter(spectra_names)) + shape = self.reader.ReadArrayShapeExplicitlyAtTimestep( + self.path, "spectra", first_spectrum_name, first_step + ) + times = self.reader.ReadPerTimestepVariable( + self.path, + "spectra", + "Time", + "t", + self.valid_files["spectra"], + ) + steps = self.reader.ReadPerTimestepVariable( + self.path, + "spectra", + "Step", + "s", + self.valid_files["spectra"], + ) + + ebins = self.reader.ReadArrayAtTimestep( + self.path, "spectra", ebin_name, first_step + ) + + diffs = np.diff(ebins) + if np.isclose(diffs[1] - diffs[0], diffs[-1] - diffs[-2], atol=1e-2): + ebins = 0.5 * (ebins[1:] + ebins[:-1]) + else: + ebins = (ebins[1:] * ebins[:-1]) ** 0.5 + + all_dims = {**times, "E": ebins} + all_coords = {**all_dims, "s": ("t", steps["s"])} + + def remap_name(name: str) -> str: + return name[1:] + + return xr.Dataset( + { + remap_name(spectrum): xr.DataArray( + da.stack( + [ + da.from_delayed( + dask.delayed(self.read_spectrum)( + path=self.path, + reader=self.reader, + spectrum=spectrum, + step=step, + ), + shape=shape, + dtype="float", + ) + for step in tqdm( + self.valid_steps["spectra"], + desc="steps", + position=1, + leave=False, + ) + ], + ), + name=remap_name(spectrum), + dims=all_dims, + coords=all_coords, + ) + for spectrum in tqdm( + spectra_names, desc="spectra", position=0, leave=False + ) + }, + attrs=self.reader.ReadAttrsAtTimestep( + path=self.path, category="spectra", step=first_step + ), + ) + + def help_spectra(self, prepend="") -> str: + ret = f"{prepend}- use .sel(...) to select specific energy or time intervals\n" + ret += f"{prepend} t : time (float)\n" + ret += f"{prepend} st : step (int)\n" + ret += f"{prepend} E : energy bin (float)\n{prepend}\n" + ret += f"{prepend} # example:\n" + ret += f"{prepend} # .sel(E=slice(10.0, 20.0)).sel(t=0, method='nearest')\n{prepend}\n" + ret += f"{prepend}- use .isel(...) to select spectra based on energy bin or time index:\n" + ret += f"{prepend} t : timestamp index (int)\n" + ret += f"{prepend} st : step index (int)\n" + ret += f"{prepend} E : energy bin index (int)\n{prepend}\n" + ret += f"{prepend} # example:\n" + ret += f"{prepend} # .isel(t=-1, E=11)\n" + ret += f"{prepend}\n" + ret += f"{prepend} # example:\n" + ret += f"{prepend} # .spectra.N_1.sel(E=slice(None, 50)).isel(t=5).plot()\n" + return ret diff --git a/nt2/__init__.py b/nt2/__init__.py index 1f5aee9..4a1f655 100644 --- a/nt2/__init__.py +++ b/nt2/__init__.py @@ -1,7 +1,63 @@ -__version__ = "1.5.3" +__version__ = "1.6.0" -import nt2.containers.data as nt2_data +import xarray as xr +from .containers.data import Data as nt2Data +from .plotters import inspect as acc_inspect +from .plotters import movie as acc_movie +from .plotters import particles as acc_particles +from .plotters import polar as acc_polar +from .utils import InheritClassDocstring -class Data(nt2_data.Data): + +class Data(nt2Data): + pass + + +# patch until xarray decides to actually fix this stupid bug +def __patch_xarray_dark(cls): + STYLE = """ + + """ + cls._repr_html_orig_ = cls._repr_html_ + cls._repr_html_ = lambda self: STYLE + cls._repr_html_orig_(self) + + +__patch_xarray_dark(xr.DataArray) +__patch_xarray_dark(xr.Dataset) + + +# register custom plotters +@xr.register_dataset_accessor("polar") +@InheritClassDocstring +class DatasetPolarPlotAccessor(acc_polar.ds_accessor): + pass + + +@xr.register_dataset_accessor("particles") +@InheritClassDocstring +class DatasetParticlesPlotAccessor(acc_particles.ds_accessor): + pass + + +@xr.register_dataarray_accessor("polar") +@InheritClassDocstring +class PolarPlotAccessor(acc_polar.accessor): + pass + + +@xr.register_dataset_accessor("inspect") +@InheritClassDocstring +class DatasetInspectPlotAccessor(acc_inspect.ds_accessor): + pass + + +@xr.register_dataarray_accessor("movie") +@InheritClassDocstring +class MoviePlotAccessor(acc_movie.accessor): pass diff --git a/nt2/cli/main.py b/nt2/cli/main.py index a7ca8cf..718fe99 100644 --- a/nt2/cli/main.py +++ b/nt2/cli/main.py @@ -1,7 +1,10 @@ -from typing import Union, Dict -import typer, nt2, os -from typing_extensions import Annotated +import os +from typing import Annotated, Dict, Union + import matplotlib.pyplot as plt +import typer + +import nt2 app = typer.Typer() @@ -23,18 +26,18 @@ def check_path(path: str) -> str: return path -def check_sel(sel: str) -> Dict[str, Union[int, float, slice]]: +def check_sel(sel: str) -> Dict[str, Union[float, slice]]: if sel == "": return {} sel_list = sel.strip().split(";") - sel_dict: Dict[str, Union[int, float, slice]] = {} + sel_dict: Dict[str, Union[float, slice]] = {} for _, s in enumerate(sel_list): coord, arg = s.strip().split("=", 1) coord = coord.strip() arg_exec = eval(arg.strip()) - assert isinstance( - arg_exec, (int, float, slice) - ), f"Invalid selection argument for '{coord}': {arg_exec}. Must be int, float, or slice." + assert isinstance(arg_exec, (int, float, slice)), ( + f"Invalid selection argument for '{coord}': {arg_exec}. Must be int, float, or slice." + ) sel_dict[coord] = arg_exec return sel_dict @@ -122,19 +125,19 @@ def plot( ): fname = os.path.basename(path.strip("/")) data = nt2.Data(path) - assert isinstance( - sel, dict - ), f"Invalid selection format: {sel}. Must be a dictionary." + assert isinstance(sel, dict), ( + f"Invalid selection format: {sel}. Must be a dictionary." + ) assert isinstance(isel, dict), f"Invalid isel format: {isel}. Must be a dictionary." if what == "fields": d = data.fields if sel != {}: slices = {} sels = {} - slices: Dict[str, Union[slice, float, int]] = { + slices: Dict[str, Union[slice, float]] = { k: v for k, v in sel.items() if isinstance(v, slice) } - sels: Dict[str, Union[slice, float, int]] = { + sels: Dict[str, Union[slice, float]] = { k: v for k, v in sel.items() if not isinstance(v, slice) } d = d.sel(**sels, method="nearest") diff --git a/nt2/containers/base.py b/nt2/containers/base.py new file mode 100644 index 0000000..69b7a3c --- /dev/null +++ b/nt2/containers/base.py @@ -0,0 +1,220 @@ +from __future__ import annotations + +from typing import Callable + +import numpy as np +import numpy.typing as npt + +from ..readers.base import BaseReader +from ..utils import CoordinateSystem + + +class BaseContainer: + """Parent container class for holding any category data.""" + + __path: str + __category: str + __reader: BaseReader + __verify: bool + __timerange: tuple[float | None, float | None] | None + __steprange: tuple[int | None, int | None] | None + __remap: dict[str, Callable[[str], str]] | None + __coordinate_system: CoordinateSystem | None + __num_cpus: int | None + + __valid_steps: list[int] + __valid_files: list[str] + + __times: npt.NDArray + __steps: npt.NDArray + + def __init__( + self, + path: str, + category: str, + reader: BaseReader, + verify: bool = False, + timerange: tuple[float | None, float | None] | None = None, + steprange: tuple[int | None, int | None] | None = None, + remap: dict[str, Callable[[str], str]] | None = None, + coord_system: CoordinateSystem | None = None, + num_cpus: int | None = None, + ): + """Initializer for the BaseContainer class. + + Parameters + ---------- + path : str + The path to the data. + category : str + The category of the data. + reader : BaseReader + The reader to be used for reading the data. + verify : Optional[bool] + Whether to verify the data. If None, it will use the reader's default. + timerange : Optional[tuple[Union[float, None], Union[float, None]]] + Time range to load. If None, all times will be loaded. + steprange : Optional[tuple[Union[int, None], Union[int, None]]] + Step range to load. If None, all steps will be loaded. + remap : Optional[dict[str, Callable[[str], str]]] + Remap dictionary to use to remap the data names (coords, fields, etc.). + coord_system : Optional[CoordinateSystem] + The coordinate system of the data. + num_cpus : Optional[int] + The number of CPUs to use for parallel processing. If None, it will use all available CPUs. + + """ + self.__path = path + self.__category = category + self.__reader = reader + self.__timerange = timerange + self.__steprange = steprange + self.__verify = verify + self.__remap = remap + self.__coordinate_system = coord_system + self.__num_cpus = num_cpus + + self.__valid_files, self.__valid_steps = self.__reader.GetValidFilesAndSteps( + self.path, + self.category, + self.steprange, + self.num_cpus, + ) + self.__times, self.__steps = self.read_times_and_steps() + if self.steprange is None: + self.narrow_timerange() + + @property + def path(self) -> str: + """str: The main path of the data.""" + return self.__path + + @property + def category(self) -> str: + """str: The category of the data.""" + return self.__category + + @property + def reader(self) -> BaseReader: + """BaseReader: The reader used to read the data.""" + return self.__reader + + @property + def verify(self) -> bool: + """bool: Whether to verify the data.""" + return self.__verify + + @property + def remap(self) -> dict[str, Callable[[str], str]] | None: + """{ str: (str) -> str } : The coordinate/field remap dictionary.""" + return self.__remap + + @property + def coordinate_system(self) -> CoordinateSystem | None: + """CoordinateSystem: The coordinate system of the data.""" + return self.__coordinate_system + + @property + def num_cpus(self) -> int | None: + """int: The number of CPUs to use for parallel processing.""" + return self.__num_cpus + + @property + def timerange(self) -> tuple[float | None, float | None] | None: + """tuple[float | None, float | None]: The time range of the data.""" + return self.__timerange + + @property + def steprange(self) -> tuple[int | None, int | None] | None: + """tuple[int | None, int | None]: The step range of the data.""" + return self.__steprange + + @property + def valid_files(self) -> list[str]: + """list[str]: The valid files of the data.""" + return self.__valid_files + + @property + def valid_steps(self) -> list[int]: + """list[int]: The valid steps of the data.""" + return self.__valid_steps + + @property + def times(self) -> npt.NDArray: + """npt.NDArray: The times of the data.""" + return self.__times + + @property + def steps(self) -> npt.NDArray: + """npt.NDArray: The steps of the data.""" + return self.__steps + + def read_times_and_steps(self) -> tuple[npt.NDArray, npt.NDArray]: + """Reads the times and steps for the given category. + + Parameters + ---------- + category : str + The category to read the times and steps for. + + Returns + ------- + tuple[npt.NDArray, npt.NDArray] + A tuple containing the times and steps for the given category. + + """ + vars = self.reader.ReadPerTimestepVariables( + path=self.path, + category=self.category, + varnames=["Time", "Step"], + newnames=["t", "s"], + valid_files=self.valid_files, + ) + return (vars["t"], vars["s"]) + + def narrow_timerange(self): + if self.timerange is not None: + start_time, end_time = self.timerange + start_idx, end_idx = 0, -1 + if start_time is None: + start_idx = 0 + else: + start_idx = int(np.searchsorted(self.__times, start_time)) + if end_time is None: + end_idx = len(self.__times) - 1 + else: + end_idx = int(np.searchsorted(self.__times, end_time)) + self.__times = self.__times[start_idx : end_idx + 1] + self.__steps = self.__steps[start_idx : end_idx + 1] + self.__valid_files = self.__valid_files[start_idx : end_idx + 1] + self.__valid_steps = self.__valid_steps[start_idx : end_idx + 1] + + def set_remap(self, remap: dict[str, Callable[[str], str]]) -> None: + """Set the remap dictionary for the container. + + Parameters + ---------- + remap : dict[str, Callable[[str], str]] + The remap dictionary to set. + + """ + self.__remap = remap + + def set_coordinate_system(self, coord_system: CoordinateSystem) -> None: + """Set the coordinate system of the data. + + Parameters + ---------- + coord_system : CoordinateSystem + The coordinate system to set. + + """ + self.__coordinate_system = coord_system + + def __dask_tokenize__(self) -> tuple[str, str, str]: + """Provide a deterministic Dask token for container instances.""" + return ( + self.__class__.__name__, + self.__path, + self.__reader.format.value, + ) diff --git a/nt2/containers/data.py b/nt2/containers/data.py index dec97cd..0fd67a7 100644 --- a/nt2/containers/data.py +++ b/nt2/containers/data.py @@ -1,139 +1,30 @@ -from typing import Callable, Any, Union, Optional, List, Dict +from __future__ import annotations -import sys -import logging +import os +from typing import Any, Callable -if sys.version_info >= (3, 12): - from typing import override -else: - - def override(method): - return method - - -from nt2.utils import ToHumanReadable - -import xarray as xr import pandas as pd +import xarray as xr -from nt2.utils import ( +from ..plotters.export import makeFramesAndMovie +from ..readers.adios2 import Reader as BP5Reader +from ..readers.base import BaseReader +from ..readers.hdf5 import Reader as HDF5Reader +from ..utils import ( + CoordinateSystem, + CoordinateSystemType, DetermineDataFormat, - InheritClassDocstring, Format, - CoordinateSystem, + ToHumanReadable, ) -from nt2.readers.base import BaseReader -from nt2.readers.hdf5 import Reader as HDF5Reader -from nt2.readers.adios2 import Reader as BP5Reader -from nt2.containers.fields import Fields -from nt2.containers.particles import Particles -from nt2.containers.spectra import Spectra -from nt2.containers.diagnostics import Diagnostics - -import nt2.plotters.polar as acc_polar -import nt2.plotters.particles as acc_particles -import nt2.plotters.inspect as acc_inspect -import nt2.plotters.movie as acc_movie -from nt2.plotters.export import makeFramesAndMovie - - -@xr.register_dataset_accessor("polar") -@InheritClassDocstring -class DatasetPolarPlotAccessor(acc_polar.ds_accessor): - pass - - -@xr.register_dataset_accessor("particles") -@InheritClassDocstring -class DatasetParticlesPlotAccessor(acc_particles.ds_accessor): - pass - - -@xr.register_dataarray_accessor("polar") -@InheritClassDocstring -class PolarPlotAccessor(acc_polar.accessor): - pass - - -@xr.register_dataset_accessor("inspect") -@InheritClassDocstring -class DatasetInspectPlotAccessor(acc_inspect.ds_accessor): - pass - +from .diagnostics import Diagnostics +from .fields import FieldContainer +from .particle_dataset import ParticleDataset +from .particles import ParticleContainer +from .spectra import SpectraContainer -@xr.register_dataarray_accessor("movie") -@InheritClassDocstring -class MoviePlotAccessor(acc_movie.accessor): - pass - -# Cartesian remapping functions -def remap_fields_cart(name: str) -> str: - name = name[1:] - fieldname = name.split("_")[0] - fieldname = fieldname.replace("0", "t") - fieldname = fieldname.replace("1", "x") - fieldname = fieldname.replace("2", "y") - fieldname = fieldname.replace("3", "z") - suffix = "_".join(name.split("_")[1:]) - return f"{fieldname}{'_' + suffix if suffix != '' else ''}" - - -def remap_coords_cart(name: str) -> str: - return { - "X1": "x", - "X2": "y", - "X3": "z", - }.get(name, name) - - -def remap_prtl_quantities_cart(name: str) -> str: - shortname = name[1:] - return { - "X1": "x", - "X2": "y", - "X3": "z", - "U1": "ux", - "U2": "uy", - "U3": "uz", - "W": "w", - }.get(shortname, shortname) - - -# Spherical remapping functions -def remap_fields_sph(name: str) -> str: - name = name[1:] - fieldname = name.split("_")[0] - fieldname = fieldname.replace("0", "t") - fieldname = fieldname.replace("1", "r") - fieldname = fieldname.replace("2", "th") - fieldname = fieldname.replace("3", "ph") - suffix = "_".join(name.split("_")[1:]) - return f"{fieldname}{'_' + suffix if suffix != '' else ''}" - - -def remap_coords_sph(name: str) -> str: - return { - "X1": "r", - "X2": "th", - "X3": "ph", - }.get(name, name) - - -def remap_prtl_quantities_sph(name: str) -> str: - shortname = name[1:] - return { - "X1": "r", - "X2": "th", - "X3": "ph", - "U1": "ur", - "U2": "uth", - "U3": "uph", - "W": "w", - }.get(shortname, shortname) - - -def compactify(lst: Union[List[Any], Any]) -> str: +def compactify(lst: list[Any] | Any) -> str: c = "" cntr = 0 for l_ in lst: @@ -145,19 +36,28 @@ def compactify(lst: Union[List[Any], Any]) -> str: return c[:-2] -class Data(Fields, Particles, Spectra): - """Main class to manage all the data containers. +class Data: + """Main class to manage all the data containers.""" - Inherits from all category-specific containers. - - """ + _fields: FieldContainer | None = None + _particles: ParticleContainer | None = None + _spectra: SpectraContainer | None = None + _diagnostics: Diagnostics | None = None def __init__( self, path: str, - reader: Optional[BaseReader] = None, - remap: Optional[Dict[str, Callable[[str], str]]] = None, - coord_system: Optional[CoordinateSystem] = None, + fields: bool = True, + particles: bool = True, + spectra: bool = True, + diagnostics: bool = False, + verify: bool = False, + timerange: tuple[float | None, float | None] | None = None, + steprange: tuple[int | None, int | None] | None = None, + reader: BaseReader | None = None, + remap: dict[str, Callable[[str], str]] | None = None, + coord_system: CoordinateSystemType | None = None, + num_cpus: int | None = min(os.cpu_count() or 1, 16), ): """Initializer for the Data class. @@ -165,14 +65,32 @@ def __init__( ---------- path : str Main path to the data + components : list[OutputComponentType], optional + List of components to load. If None, all components will be loaded. + fields : bool, optional + Whether to load the fields component. Default is True. + particles : bool, optional + Whether to load the particles component. Default is True. + spectra : bool, optional + Whether to load the spectra component. Default is True. + diagnostics : bool, optional + Whether to load the diagnostics component. Default is False. + verify : bool, optional + Whether to verify the data. Default is False. + timerange : tuple[float | None, float | None], optional + Time range to load. If None, all times will be loaded. + steprange : tuple[int | None, int | None], optional + Step range to load. If None, all steps will be loaded. reader : BaseReader, optional Reader to use to read the data. If None, it will be determined based on the file format. remap : dict[str, Callable[[str], str]], optional Remap dictionary to use to remap the data names (coords, fields, etc.). - coord_system : CoordinateSystem, optional + coord_system : Literal["XYZ", "SPH"], optional Coordinate system of the data. If None, it will be determined based on the data attrs (if remap is also None). + num_cpus : int, optional + Number of CPUs to use for parallel processing. If None, it will use all available CPUs. Raises ------ @@ -185,9 +103,9 @@ def __init__( fmt = DetermineDataFormat(path) if reader is None: if fmt == Format.HDF5: - self.__reader = HDF5Reader() + __reader = HDF5Reader() elif fmt == Format.BP5: - self.__reader = BP5Reader() + __reader = BP5Reader() else: raise NotImplementedError( "Only HDF5 & BP5 formats are supported at the moment." @@ -197,73 +115,131 @@ def __init__( raise ValueError( f"Reader format {reader.format} does not match data format {fmt}." ) - self.__reader = reader + __reader = reader + + if fields: + self._fields = FieldContainer( + path=path, + reader=__reader, + verify=verify, + timerange=timerange, + steprange=steprange, + remap=remap, + coord_system=CoordinateSystem.from_str(coord_system) + if coord_system + else None, + num_cpus=num_cpus, + ) + if particles: + self._particles = ParticleContainer( + path=path, + reader=__reader, + verify=verify, + timerange=timerange, + steprange=steprange, + remap=remap, + coord_system=CoordinateSystem.from_str(coord_system) + if coord_system + else None, + num_cpus=num_cpus, + ) + if spectra: + self._spectra = SpectraContainer( + path=path, + reader=__reader, + verify=verify, + timerange=timerange, + steprange=steprange, + remap=remap, + coord_system=CoordinateSystem.from_str(coord_system) + if coord_system + else None, + num_cpus=num_cpus, + ) + if diagnostics: + self._diagnostics = Diagnostics(path=path) + self.__attrs: dict[str, Any] = {} + if self.fields_defined: + self.__attrs.update(**self._fields.attrs) + if self.particles_defined: + self.__attrs.update(**self._particles.attrs) + if self.spectra_defined: + self.__attrs.update(**self._spectra.attrs) - # determine the coordinate system and remapping - self.__attrs: Dict[str, Any] = {} - for category in ["fields", "particles", "spectra"]: - if self.__reader.DefinesCategory(path, category): - valid_steps = self.__reader.GetValidSteps(path, category) - if len(valid_steps) == 0: - raise ValueError(f"No valid steps found for category {category}.") - first_step = valid_steps[0] - attrs = self.__reader.ReadAttrsAtTimestep(path, category, first_step) - self.__attrs.update(**attrs) - if "Coordinates" not in attrs: - raise ValueError( - f"Coordinates not found in attributes for category {category}." - ) - else: - if attrs["Coordinates"] in [b"cart", "cart"]: - coord_system = CoordinateSystem.XYZ - elif attrs["Coordinates"] in [b"sph", "sph", b"qsph", "qsph"]: - coord_system = CoordinateSystem.SPH + @property + def fields_defined(self) -> bool: + """bool: Whether fields are defined in the data.""" + return ( + self._fields is not None + and self._fields.fields_defined + and self._fields.fields is not None + ) - else: - raise NotImplementedError( - f"Coordinate system {attrs['Coordinates']} not supported." - ) - if remap is None: - remap = { - "coords": ( - remap_coords_cart - if coord_system == CoordinateSystem.XYZ - else remap_coords_sph - ), - "fields": ( - remap_fields_cart - if coord_system == CoordinateSystem.XYZ - else remap_fields_sph - ), - "particles": ( - remap_prtl_quantities_cart - if coord_system == CoordinateSystem.XYZ - else remap_prtl_quantities_sph - ), - } - break + @property + def particles_defined(self) -> bool: + """bool: Whether particles are defined in the data.""" + return ( + self._particles is not None + and self._particles.particles_defined + and self._particles.particles is not None + ) - if coord_system is None: - raise ValueError("No coordinate system found in the data.") + @property + def spectra_defined(self) -> bool: + """bool: Whether spectra are defined in the data.""" + return ( + self._spectra is not None + and self._spectra.spectra_defined + and self._spectra.spectra is not None + ) - self.__coordinate_system = coord_system + @property + def diagnostics_defined(self) -> bool: + """bool: Whether diagnostics are defined in the data.""" + return self._diagnostics is not None and self._diagnostics.df is not None + + @property + def fields(self) -> xr.Dataset: + """xr.Dataset: The fields dataset.""" + if not self.fields_defined: + raise ValueError("Fields are not defined in the data.") + return self._fields.fields + + @property + def particles(self) -> ParticleDataset: + """ParticleDataset: The particles dataset.""" + if not self.particles_defined: + raise ValueError("Particles are not defined in the data.") + return self._particles.particles + + @property + def spectra(self) -> xr.Dataset: + """xr.Dataset: The spectra dataset.""" + if not self.spectra_defined: + raise ValueError("Spectra are not defined in the data.") + return self._spectra.spectra + + @property + def diagnostics(self) -> pd.DataFrame: + """Diagnostics: The diagnostics dataset.""" + if not self.diagnostics_defined: + raise ValueError("Diagnostics are not defined in the data.") + return self._diagnostics.df - super(Data, self).__init__(path=path, reader=self.__reader, remap=remap) - try: - self.__diagnostics = Diagnostics(path) - except Exception as e: - logging.warning(f"Failed to read diagnostics: {e}") - self.__diagnostics = None + @property + def attrs(self) -> dict[str, Any]: + """dict: The attributes of the data.""" + return self.__attrs def makeMovie( self, plot: Callable, - time: Optional[List[float]] = None, - num_cpus: Optional[int] = None, + time: list[float] | None = None, + num_cpus: int | None = None, **movie_kwargs: Any, ) -> bool: - f"""Create animation with provided plot function. - + """Create animation with provided plot function. + Parameters ---------- plot : callable @@ -274,7 +250,7 @@ def makeMovie( The number of CPUs to use for parallel processing. If None, it will use all available CPUs. **movie_kwargs : dict Additional keyword arguments to pass to the movie creation function. - + Returns ------- bool @@ -283,20 +259,27 @@ def makeMovie( if time is None: if self.fields_defined: time = self.fields.t.values - elif self.particles_defined and self.particles is not None: + elif self.particles_defined: time = list(self.particles.times) else: raise ValueError("No time values found.") - assert time is not None, "Time values must be provided." + if time is None: + raise ValueError("No time values found.") name: str = "" - if self.attrs.get("simulation.name", None) == None: - name = movie_kwargs.pop("name", "movie") + provided_name = movie_kwargs.pop("name", None) + if ( + provided_name is not None + or self.attrs.get("simulation.name", None) is not None + ): + name = provided_name else: name_b = self.attrs.get("simulation.name") if isinstance(name_b, bytes): name = name_b.decode("utf-8") else: name = str(name_b) + if name is None: + raise ValueError("No name provided for the movie.") return makeFramesAndMovie( name=name, data=self, @@ -306,23 +289,6 @@ def makeMovie( **movie_kwargs, ) - @property - def coordinate_system(self) -> CoordinateSystem: - """CoordinateSystem: The coordinate system of the data.""" - return self.__coordinate_system - - @property - def attrs(self) -> Dict[str, Any]: - """dict[str, Any]: The attributes of the data.""" - return self.__attrs - - @property - def diagnostics(self) -> Union[pd.DataFrame, None]: - """pd.DataFrame or None: The diagnostics output if .out file is found, None otherwise.""" - if self.__diagnostics is None: - return None - return self.__diagnostics.df - def to_str(self) -> str: """str: String representation of the all the enclosed dataframes.""" @@ -330,34 +296,44 @@ def to_str(self) -> str: if self.fields_defined: string += "FieldsDataset:\n" string += "==============\n" - string += f"| Coordinates:\n| {self.coordinate_system.value}\n|\n" + string += f"| Coordinates:\n| {self._fields.coordinate_system.value}\n|\n" string += f"| Data axes:\n| {compactify(self.fields.indexes.keys())}\n|\n" - delta_t = ( - self.fields.coords["t"].values[1] - self.fields.coords["t"].values[0] - ) / (self.fields.coords["s"].values[1] - self.fields.coords["s"].values[0]) - string += f"| - dt: {delta_t:.2e}\n" - for key in self.fields.coords.keys(): + if ( + len(self.fields.coords["t"].values) > 1 + and len(self.fields.coords["s"].values) > 1 + ): + delta_t = ( + self.fields.coords["t"].values[1] + - self.fields.coords["t"].values[0] + ) / ( + self.fields.coords["s"].values[1] + - self.fields.coords["s"].values[0] + ) + string += f"| - dt: {delta_t:.2e}\n" + for key in self.fields.coords: crd = self.fields.coords[key].values fmt = "" if key != "s": fmt = ".2f" string += f"| - {key}: {crd.min():{fmt}} -> {crd.max():{fmt}} [{len(crd)}]\n" string += "|\n" - string += f"| Quantities:\n| {compactify(sorted(self.fields.data_vars.keys()))}\n|\n" + string += f"| Quantities:\n| {compactify(sorted(map(str, self.fields.data_vars.keys())))}\n|\n" string += f"| Total size: {ToHumanReadable(self.fields.nbytes)}\n\n" else: string += "FieldsDataset:\n" string += "==============\n" string += " empty\n\n" - if self.particles_defined and self.particles is not None: + if self.particles_defined: species = sorted(self.particles.species) string += "ParticleDataset:\n" string += "================\n" string += f"| Species:\n| {compactify(species)}\n|\n" string += f"| Timesteps:\n| {len(self.particles.times)}\n|\n" string += f"| Quantities:\n| {compactify(self.particles.columns)}\n|\n" - string += f"| Total size: {ToHumanReadable(self.particles.nbytes)}\n|\n" - string += self.help_particles("| ") + string += ( + f"| Estimated index size: {ToHumanReadable(self.particles.nbytes)}\n|\n" + ) + string += self._particles.help_particles("| ") string += "\n" else: string += "ParticleDataset:\n" @@ -375,16 +351,16 @@ def to_str(self) -> str: self.spectra.coords["s"].values[1] - self.spectra.coords["s"].values[0] ) string += f"| - dt: {delta_t:.2e}\n" - for key in self.spectra.coords.keys(): + for key in self.spectra.coords: crd = self.spectra.coords[key].values fmt = "" if key != "s": fmt = ".2f" string += f"| - {key}: {crd.min():{fmt}} -> {crd.max():{fmt}} [{len(crd)}]\n" string += "|\n" - string += f"| Quantities:\n| {compactify(sorted(self.spectra.data_vars.keys()))}\n|\n" + string += f"| Quantities:\n| {compactify(sorted(map(str, self.spectra.data_vars.keys())))}\n|\n" string += f"| Total size: {ToHumanReadable(self.spectra.nbytes)}\n|\n" - string += self.help_spectra("| ") + string += self._spectra.help_spectra("| ") else: string += "SpectraDataset:\n" string += "===============\n" @@ -392,10 +368,8 @@ def to_str(self) -> str: return string - @override def __str__(self) -> str: return self.to_str() - @override def __repr__(self) -> str: return self.to_str() diff --git a/nt2/containers/diagnostics.py b/nt2/containers/diagnostics.py index 98fe33b..0a8d001 100644 --- a/nt2/containers/diagnostics.py +++ b/nt2/containers/diagnostics.py @@ -1,108 +1,301 @@ -from typing import Union +"""Parse Entity diagnostic log files into pandas dataframes.""" + +from __future__ import annotations + +import json +import logging +import os +import re +from collections.abc import Iterator +from pathlib import Path + import pandas as pd +_NUMBER = r"[+-]?(?:\d+(?:\.\d*)?|\.\d+)(?:[eE][+-]?\d+)?" +_STEP_RE = re.compile(r"^Step:\s+(\d+)\.+\[") +_TIME_RE = re.compile(rf"^Time:\s+({_NUMBER})\.+\[") +_SUBSTEP_RE = re.compile( + rf"^\s+(?P[A-Za-z]+)\.+(?P{_NUMBER})\s+" + rf"(?Pms|[µμu]s|ns|s)\b" +) +_SPECIES_RE = re.compile( + rf"^\s+species\s+(?P\d+)\s+\([^)]*\)\.+" + rf"(?P{_NUMBER})" + rf"(?:\s+\d+%\s*:\s*\d+%\s+(?P{_NUMBER})" + rf"\s*:\s*(?P{_NUMBER}))?" +) + +_UNIT_TO_NS = { + "s": 1e9, + "ms": 1e6, + "µs": 1e3, + "μs": 1e3, + "us": 1e3, + "ns": 1.0, +} + class Diagnostics: - df: Union[pd.DataFrame, None] + """Diagnostic data parsed from an Entity ``.out`` file. + + The source is read one line at a time and converted in bounded-size chunks. + By default, the parsed dataframe is cached next to the source as Parquet. + A valid cache is used on later construction instead of scanning the log. + + Parameters + ---------- + path + Directory containing a ``.out`` file, or the path to one ``.out`` file. + cache + Whether to read and write the Parquet cache. + chunk_size + Number of timestep records per Parquet row group. + cache_path + Optional cache filename. The default is + ``.diagnostics.parquet``. + """ + + _CACHE_VERSION = 1 + + df: pd.DataFrame | None + outfile: str - def __init__(self, path: str): - import os - import logging - import re + def __init__( + self, + path: str | os.PathLike[str], + *, + cache: bool = True, + chunk_size: int = 100_000, + cache_path: str | os.PathLike[str] | None = None, + ) -> None: + if chunk_size <= 0: + raise ValueError("chunk_size must be greater than zero") - outfiles = [o for o in os.listdir(path) if o.endswith(".out")] - if len(outfiles) == 0: - logging.warning(f"No .out files found in {path}") + outfile = self._find_outfile(Path(path)) + logger = logging.getLogger(__name__) + if outfile is None: + logger.warning("No .out files found in %s", path) self.df = None - else: - self.outfile = os.path.join(path, outfiles[0]) - - data = {} - - with open(self.outfile, "r") as f: - content = f.read() - steps = re.findall(r"Step:\s+(\d+)\.+\[", content) - times = re.findall(r"Time:\s+([\d.]+\d)\.+\[", content) - substeps = re.findall(r"\s+([A-Za-z]+)\.+([\d.]+)\s+([mµn]?s)", content) - species = re.findall( - r"\s+species\s+(\d+)\s+\(.+\)\.+([\deE+-.]+)(\s+\d+\%\s:\s\d+\%\s+)?([\deE+-.]+)?( : )?([\deE+-.]+)?", - content, + return + + self.outfile = str(outfile) + parquet_path = ( + Path(cache_path) + if cache_path is not None + else Path(f"{outfile}.diagnostics.parquet") + ) + metadata_path = Path(f"{parquet_path}.json") + if parquet_path.resolve() == outfile.resolve(): + raise ValueError("cache_path must not overwrite the source .out file") + + if cache and self._cache_is_valid(outfile, parquet_path, metadata_path): + try: + self.df = pd.read_parquet(parquet_path) + return + except (OSError, ValueError): + logger.warning( + "Failed to read diagnostics cache %s; rebuilding it", + parquet_path, + exc_info=True, ) - data["steps"] = [] - for step in steps: - data["steps"].append(int(step)) - - data["times"] = [] - for time in times: - data["times"].append(float(time)) - - assert len(data["steps"]) == len( - data["times"] - ), "Number of steps and times do not match" - - data["substeps"] = {} - for substep in substeps: - if substep[0] not in data["substeps"].keys(): - data["substeps"][substep[0]] = [] - - def to_ns(value: float, unit: str) -> float: - if unit == "s": - return value * 1e9 - elif unit == "ms": - return value * 1e6 - elif unit == "µs": - return value * 1e3 - elif unit == "ns": - return value - else: - raise ValueError(f"Unknown time unit: {unit}") - - data["substeps"][substep[0]].append( - to_ns(float(substep[1]), substep[2]) + if cache: + self.df = self._parse_to_cache( + outfile, parquet_path, metadata_path, chunk_size + ) + else: + chunks = list(self._dataframe_chunks(outfile, chunk_size)) + self.df = self._combine_chunks(chunks) + + @staticmethod + def _find_outfile(path: Path) -> Path | None: + if path.is_file(): + if path.suffix != ".out": + raise ValueError(f"Expected a .out file, got {path}") + return path + if not path.exists(): + raise FileNotFoundError(path) + if not path.is_dir(): + raise ValueError(f"Expected a directory or .out file, got {path}") + + outfiles = sorted(path.glob("*.out")) + return outfiles[0] if outfiles else None + + @classmethod + def _cache_is_valid( + cls, source: Path, parquet_path: Path, metadata_path: Path + ) -> bool: + if not parquet_path.is_file() or not metadata_path.is_file(): + return False + try: + with metadata_path.open("r", encoding="utf-8") as stream: + metadata = json.load(stream) + stat = source.stat() + return metadata == { + "parser_version": cls._CACHE_VERSION, + "source": str(source.resolve()), + "source_size": stat.st_size, + "source_mtime_ns": stat.st_mtime_ns, + } + except (OSError, ValueError, TypeError): + return False + + @classmethod + def _cache_metadata(cls, source: Path) -> dict[str, int | str]: + stat = source.stat() + return { + "parser_version": cls._CACHE_VERSION, + "source": str(source.resolve()), + "source_size": stat.st_size, + "source_mtime_ns": stat.st_mtime_ns, + } + + @classmethod + def _parse_to_cache( + cls, + source: Path, + parquet_path: Path, + metadata_path: Path, + chunk_size: int, + ) -> pd.DataFrame: + # ParquetWriter creates multiple row groups in one file without retaining + # the full parsed table in memory. Import lazily so cache=False only needs + # pandas. + import pyarrow as pa + import pyarrow.parquet as pq + + parquet_path.parent.mkdir(parents=True, exist_ok=True) + temporary_parquet = parquet_path.with_name( + f".{parquet_path.name}.{os.getpid()}.tmp" + ) + temporary_metadata = metadata_path.with_name( + f".{metadata_path.name}.{os.getpid()}.tmp" + ) + + writer = None + source_metadata = cls._cache_metadata(source) + try: + for chunk in cls._dataframe_chunks(source, chunk_size): + table = pa.Table.from_pandas(chunk, preserve_index=True) + if writer is None: + writer = pq.ParquetWriter(temporary_parquet, table.schema) + elif table.schema != writer.schema: + raise ValueError( + "Diagnostics columns or dtypes changed within the log" + ) + writer.write_table(table) + + if writer is not None: + writer.close() + writer = None + else: + # Preserve the historical result for a log with no records. + empty = pd.DataFrame(columns=["Step", "Time"]) + empty.index = pd.Index([], name="Step", dtype="int64") + empty.to_parquet(temporary_parquet) + + with temporary_metadata.open("w", encoding="utf-8") as stream: + # Use the source state from before parsing. If a running + # simulation appended to the file meanwhile, this cache will + # intentionally be stale and rebuilt on the next construction. + json.dump(source_metadata, stream) + + os.replace(temporary_parquet, parquet_path) + os.replace(temporary_metadata, metadata_path) + finally: + if writer is not None: + writer.close() + temporary_parquet.unlink(missing_ok=True) + temporary_metadata.unlink(missing_ok=True) + + return pd.read_parquet(parquet_path) + + @classmethod + def _dataframe_chunks(cls, source: Path, chunk_size: int) -> Iterator[pd.DataFrame]: + records: list[dict[str, float]] = [] + columns: list[str] | None = None + + for record in cls._records(source): + if columns is None: + columns = list(record) + else: + missing = set(columns) - set(record) + extra = set(record) - set(columns) + if missing or extra: + raise ValueError( + f"Inconsistent diagnostics record at step {record['Step']}: " + f"missing={sorted(missing)}, extra={sorted(extra)}" ) + records.append(record) + if len(records) >= chunk_size: + yield cls._make_dataframe(records, columns) + records = [] + + if records: + # columns is necessarily set when records is non-empty. + yield cls._make_dataframe(records, columns or []) + + @staticmethod + def _make_dataframe( + records: list[dict[str, float]], columns: list[str] + ) -> pd.DataFrame: + dataframe = pd.DataFrame.from_records(records, columns=columns) + return dataframe.set_index("Step", drop=False) + + @staticmethod + def _combine_chunks(chunks: list[pd.DataFrame]) -> pd.DataFrame: + if chunks: + return pd.concat(chunks) + dataframe = pd.DataFrame(columns=["Step", "Time"]) + dataframe.index = pd.Index([], name="Step", dtype="int64") + return dataframe + + @staticmethod + def _records(source: Path) -> Iterator[dict[str, float]]: + record: dict[str, float] | None = None + + with source.open("r", encoding="utf-8-sig", errors="replace") as stream: + for line in stream: + if line.startswith("Step:"): + if record is not None: + yield record + match = _STEP_RE.match(line) + record = {"Step": int(match.group(1))} if match else None + continue + + if record is None: + continue + + # Entity terminates each complete timestep with a dotted line. + # Yielding here means a final record interrupted while the + # simulation is still writing is not exposed as valid data. + if line.startswith("........................................"): + yield record + record = None + continue + + if line.startswith("Time:"): + match = _TIME_RE.match(line) + if match: + record["Time"] = float(match.group(1)) + continue + + if line.startswith(" species"): + match = _SPECIES_RE.match(line) + if match: + species = match.group("species") + record[f"species_{species}"] = int(float(match.group("total"))) + minimum = match.group("minimum") + maximum = match.group("maximum") + if minimum is not None and maximum is not None: + record[f"species_{species}_min"] = int(float(minimum)) + record[f"species_{species}_max"] = int(float(maximum)) + continue - for key in data["substeps"].keys(): - assert len(data["substeps"][key]) == len( - data["steps"] - ), f"Number of substep entries for {key} does not match number of steps" - - data["species"] = {} - data["species_min"] = {} - data["species_max"] = {} - for specie in species: - if specie[0] not in data["species"].keys(): - data["species"][specie[0]] = [] - data["species_min"][specie[0]] = [] - data["species_max"][specie[0]] = [] - data["species"][specie[0]].append(int(float(specie[1]))) - if len(specie) == 6 and specie[3] != "" and specie[5] != "": - data["species_min"][specie[0]].append(int(float(specie[3]))) - data["species_max"][specie[0]].append(int(float(specie[5]))) - - for key in data["species"].keys(): - assert len(data["species"][key]) == len( - data["steps"] - ), f"Number of species entries for {key} does not match number of steps" - assert (len(data["species_min"][key]) == len(data["steps"])) or ( - len(data["species_min"][key]) == 0 - ), f"Number of species min entries for {key} does not match number of steps" - assert (len(data["species_max"][key]) == len(data["steps"])) or ( - len(data["species_max"][key]) == 0 - ), f"Number of species max entries for {key} does not match number of steps" - - self.df = pd.DataFrame(index=data["steps"]) - self.df["Step"] = data["steps"] - self.df["Time"] = data["times"] - for key in data["substeps"].keys(): - self.df[key] = data["substeps"][key] - for key in data["species"].keys(): - self.df[f"species_{key}"] = data["species"][key] - if ( - len(data["species_min"][key]) > 0 - and len(data["species_max"][key]) > 0 - ): - self.df[f"species_{key}_min"] = data["species_min"][key] - self.df[f"species_{key}_max"] = data["species_max"][key] - - del data + if line.startswith(" "): + match = _SUBSTEP_RE.match(line) + if match: + record[match.group("name")] = ( + float(match.group("value")) + * _UNIT_TO_NS[match.group("unit")] + ) diff --git a/nt2/containers/fields.py b/nt2/containers/fields.py index 5f323a3..7636244 100644 --- a/nt2/containers/fields.py +++ b/nt2/containers/fields.py @@ -1,40 +1,59 @@ +from __future__ import annotations + from typing import Any import dask import dask.array as da import xarray as xr +from tqdm import tqdm -from nt2.containers.container import BaseContainer -from nt2.utils import Layout +from ..utils import CoordinateSystem, Layout +from .base import BaseContainer -class Fields(BaseContainer): - """Parent class to manage the fields dataframe.""" +def remap_fields_cart(name: str) -> str: + name = name[1:] + fieldname = name.split("_")[0] + fieldname = fieldname.replace("0", "t") + fieldname = fieldname.replace("1", "x") + fieldname = fieldname.replace("2", "y") + fieldname = fieldname.replace("3", "z") + suffix = "_".join(name.split("_")[1:]) + return f"{fieldname}{'_' + suffix if suffix != '' else ''}" - def _read_field(self, layout: Layout, field: str, step: int) -> Any: - """Reads a field from the data. - This is a dask-delayed function used further to build the dataset. +def remap_coords_cart(name: str) -> str: + return { + "X1": "x", + "X2": "y", + "X3": "z", + }.get(name, name) - Parameters - ---------- - layout : Layout - Layout of the field. - field : str - Field to read. - step : int - Step to read. - Returns - ------- - Any - Field data. +def remap_fields_sph(name: str) -> str: + name = name[1:] + fieldname = name.split("_")[0] + fieldname = fieldname.replace("0", "t") + fieldname = fieldname.replace("1", "r") + fieldname = fieldname.replace("2", "th") + fieldname = fieldname.replace("3", "ph") + suffix = "_".join(name.split("_")[1:]) + return f"{fieldname}{'_' + suffix if suffix != '' else ''}" - """ - if layout == Layout.L: - return self.reader.ReadArrayAtTimestep(self.path, "fields", field, step) - else: - return self.reader.ReadArrayAtTimestep(self.path, "fields", field, step).T + +def remap_coords_sph(name: str) -> str: + return { + "X1": "r", + "X2": "th", + "X3": "ph", + }.get(name, name) + + +class FieldContainer(BaseContainer): + """Parent class to manage the fields dataframe.""" + + __fields_defined: bool = False + __fields: xr.Dataset | None = None def __init__( self, @@ -48,13 +67,10 @@ def __init__( Keyword arguments to be passed to the parent BaseContainer class. """ - super(Fields, self).__init__(**kwargs) - if self.reader.DefinesCategory(self.path, "fields"): + super().__init__(category="fields", **kwargs) + if self.reader.DefinesCategory(self.path, "fields", self.valid_files): self.__fields_defined = True self.__fields = self._read_fields() - else: - self.__fields_defined = False - self.__fields = xr.Dataset() @property def fields_defined(self) -> bool: @@ -62,22 +78,63 @@ def fields_defined(self) -> bool: return self.__fields_defined @property - def fields(self) -> xr.Dataset: + def fields(self) -> xr.Dataset | None: """xr.Dataset: The fields dataframe.""" return self.__fields + def _read_field(self, layout: Layout, field: str, step: int) -> Any: + """Reads a field from the data. + + This is a dask-delayed function used further to build the dataset. + + Parameters + ---------- + layout : Layout + Layout of the field. + field : str + Field to read. + step : int + Step to read. + + Returns + ------- + Any + Field data. + + """ + if layout == Layout.L: + return self.reader.ReadArrayAtTimestep(self.path, "fields", field, step) + else: + return self.reader.ReadArrayAtTimestep(self.path, "fields", field, step).T + def _read_fields(self) -> xr.Dataset: """Helper function to read the fields dataframe.""" - self.reader.VerifySameCategoryNames(self.path, "fields", "f") - self.reader.VerifySameFieldShapes(self.path) - self.reader.VerifySameFieldLayouts(self.path) - valid_steps = self.reader.GetValidSteps(self.path, "fields") + # ensure that the category names, shapes, and layouts are consistent across all steps + if self.verify: + self.reader.VerifySameCategoryNames( + self.path, + "fields", + "f", + self.valid_steps, + self.num_cpus, + ) + self.reader.VerifySameFieldShapes( + self.path, + self.valid_steps, + self.num_cpus, + ) + self.reader.VerifySameFieldLayouts( + self.path, + self.valid_steps, + self.num_cpus, + ) + + # read the field names, layout, shape, coordinates, and attributes from the first step + first_step = self.valid_steps[0] field_names = self.reader.ReadCategoryNamesAtTimestep( - self.path, "fields", "f", valid_steps[0] + self.path, "fields", "f", first_step ) - - first_step = valid_steps[0] first_name = next(iter(field_names)) layout = self.reader.ReadFieldLayoutAtTimestep(self.path, first_step) shape = self.reader.ReadArrayShapeAtTimestep( @@ -85,19 +142,51 @@ def _read_fields(self) -> xr.Dataset: ) coords = self.reader.ReadFieldCoordsAtTimestep(self.path, first_step) coords = {k: coords[k] for k in sorted(coords.keys())[::-1]} + attributes = self.reader.ReadAttrsAtTimestep( + path=self.path, category="fields", step=first_step + ) + if self.coordinate_system is None: + if "Coordinates" not in attributes: + raise ValueError("Coordinates not found in attributes for fields.") + if attributes["Coordinates"] in [b"cart", "cart"]: + self.set_coordinate_system(CoordinateSystem.XYZ) + elif attributes["Coordinates"] in [b"sph", "sph", b"qsph", "qsph"]: + self.set_coordinate_system(CoordinateSystem.SPH) + else: + raise NotImplementedError( + f"Coordinate system {attributes['Coordinates']} not supported." + ) + + if ( + self.remap is None + or self.remap.get("coords", None) is None + or self.remap.get("fields", None) is None + ): + self.set_remap( + { + "coords": ( + remap_coords_cart + if self.coordinate_system == CoordinateSystem.XYZ + else remap_coords_sph + ), + "fields": ( + remap_fields_cart + if self.coordinate_system == CoordinateSystem.XYZ + else remap_fields_sph + ), + } + ) + # rename coordinates if remap is provided if self.remap is not None and "coords" in self.remap: new_coords = {} - for coord in coords.keys(): + for coord in coords: new_coords[self.remap["coords"](coord)] = coords[coord] coords = new_coords - times = self.reader.ReadPerTimestepVariable(self.path, "fields", "Time", "t") - steps = self.reader.ReadPerTimestepVariable(self.path, "fields", "Step", "s") - edge_coords = self.reader.ReadEdgeCoordsAtTimestep(self.path, first_step) new_edge_coords = {} - for coord in edge_coords.keys(): + for coord in edge_coords: assoc_x = ( coord[:-1] if (self.remap is None or "coords" not in self.remap) @@ -107,8 +196,8 @@ def _read_fields(self) -> xr.Dataset: new_edge_coords[assoc_x + "_max"] = (assoc_x, edge_coords[coord][1:]) edge_coords = new_edge_coords - all_dims = {**times, **coords}.keys() - all_coords = {**times, **coords, "s": ("t", steps["s"]), **edge_coords} + all_dims = {"t": self.times, **coords}.keys() + all_coords = {"t": self.times, **coords, "s": ("t", self.steps), **edge_coords} return xr.Dataset( { @@ -126,7 +215,12 @@ def _read_fields(self) -> xr.Dataset: shape=shape[:: -1 if layout == Layout.R else 1], dtype="float", ) - for step in valid_steps + for step in tqdm( + self.valid_steps, + desc="steps", + position=1, + leave=False, + ) ], axis=0, ), @@ -134,9 +228,20 @@ def _read_fields(self) -> xr.Dataset: dims=all_dims, coords=all_coords, ) - for name in field_names + for name in tqdm( + field_names, + desc="fields", + position=0, + leave=False, + ) }, - attrs=self.reader.ReadAttrsAtTimestep( - path=self.path, category="fields", step=first_step - ), + attrs=attributes, ) + + @property + def attrs(self) -> dict[str, Any]: + """dict: The attributes of the fields dataframe.""" + if self.fields_defined: + return self.fields.attrs + else: + return {} diff --git a/nt2/containers/particle_dataset.py b/nt2/containers/particle_dataset.py new file mode 100644 index 0000000..70cdaf6 --- /dev/null +++ b/nt2/containers/particle_dataset.py @@ -0,0 +1,666 @@ +from __future__ import annotations + +import sys +from collections.abc import Sequence +from copy import copy +from typing import ( + Any, + Callable, + Literal, + cast, +) + +import dask +import dask.dataframe as dd +import matplotlib.axes as maxes +import matplotlib.pyplot as plt +import numpy as np +import numpy.typing as npt +import pandas as pd +from dask.delayed import Delayed +from dask.optimization import cull + +if sys.version_info >= (3, 10): + IntSelector = int | Sequence[int] | slice | tuple[int, int] + FloatSelector = float | slice | Sequence[float] | tuple[float, float] +else: + IntSelector = Any + FloatSelector = Any + + +def _cull_dataframe_graph(ddf: dd.DataFrame) -> dd.DataFrame: + keys = ddf.__dask_keys__() + graph, _ = cull(ddf.__dask_graph__(), keys) + partitions = [Delayed(key, graph) for key in keys] + return cast( + dd.DataFrame, + dd.from_delayed( + partitions, + meta=ddf._meta, + divisions=ddf.divisions, + ), + ) + + +class Selection: + def __init__( + self, + type: Literal["value", "range", "list"], + value: float | list | tuple | None = None, + ): + self.type = type + self.value = value + + def intersect(self, other: Selection) -> Selection: + if self.value is None: + return copy(other) + elif other.value is None: + return copy(self) + if self.type == "value" and other.type == "value": + if self.value == other.value: + return Selection("value", self.value) + else: + return Selection("value") + elif self.type == "value" and other.type == "list": + assert isinstance(other.value, list), "other.value must be a list" + if self.value in other.value: + return Selection("value", self.value) + else: + return Selection("value") + elif self.type == "value" and other.type == "range": + assert isinstance(other.value, tuple) and len(other.value) == 2, ( + "other.value must be a tuple of length 2" + ) + lo, hi = other.value + if lo <= self.value < hi: + return Selection("value", self.value) + else: + return Selection("value") + elif self.type == "list" and other.type == "value": + return other.intersect(self) + elif self.type == "list" and other.type == "list": + assert isinstance(self.value, list), "self.value must be a list" + assert isinstance(other.value, list), "other.value must be a list" + new_values = [v for v in self.value if v in other.value] + return Selection("list", new_values) + elif self.type == "list" and other.type == "range": + assert isinstance(other.value, tuple) and len(other.value) == 2, ( + "other.value must be a tuple of length 2" + ) + assert isinstance(self.value, list), "self.value must be a list" + lo, hi = other.value + new_values = [v for v in self.value if lo <= v <= hi] + return Selection("list", new_values) + elif (self.type == "range" and other.type == "value") or ( + self.type == "range" and other.type == "list" + ): + return other.intersect(self) + elif self.type == "range" and other.type == "range": + assert isinstance(self.value, tuple) and len(self.value) == 2, ( + "self.value must be a tuple of length 2" + ) + assert isinstance(other.value, tuple) and len(other.value) == 2, ( + "other.value must be a tuple of length 2" + ) + lo1, hi1 = self.value + lo2, hi2 = other.value + new_lo = max(lo1, lo2) + new_hi = min(hi1, hi2) + if new_lo <= new_hi: + return Selection("range", (new_lo, new_hi)) + else: + return Selection("value") + else: + raise ValueError(f"Unknown selection types: {self.type}, {other.type}") + + def __repr__(self) -> str: + if self.type == "value": + return "all" if self.value is None else f"{self.value:.3g}" + elif self.type == "range": + if self.value is None: + return "all" + else: + assert isinstance(self.value, tuple) and len(self.value) == 2, ( + "value must be a tuple of length 2" + ) + lo, hi = self.value + lo_str = "..." if lo is None or lo == -np.inf else f"{lo:.3g}" + hi_str = "..." if hi is None or hi == np.inf else f"{hi:.3g}" + return f"[ {lo_str} -> {hi_str} ]" + elif self.type == "list": + assert isinstance(self.value, list), "value must be a list" + return "{ " + ", ".join(f"{v:.3g}" for v in self.value) + " }" + else: + return "InvalidSelection" + + def __str__(self) -> str: + return self.__repr__() + + +def _coerce_selector_to_mask( + s: IntSelector | FloatSelector, + series: Any, + inclusive_tuple: bool = True, + method="exact", +): + from functools import reduce + from operator import ior + + if isinstance(s, slice): + lo = s.start if s.start is not None else -np.inf + hi = s.stop if s.stop is not None else np.inf + step = s.step + mask = (series >= lo) & (series <= hi) + if step not in (None, 1): + mask = mask & (((series - lo) % step) == 0) + return mask, ("range", (lo, hi)) + elif isinstance(s, tuple) and len(s) == 2 and inclusive_tuple: + lo, hi = s + if lo is None: + lo = -np.inf + if hi is None: + hi = np.inf + return (series >= lo) & (series <= hi), ("range", (lo, hi)) + elif isinstance(s, (list, tuple, np.ndarray, pd.Index, pd.Series)): + if method == "exact": + return series.isin(list(s)), ("list", list(s)) + else: + return reduce( + ior, [np.abs(series - v) == np.abs(series - v).min() for v in s] + ), ("list", list(s)) + else: + if method == "exact": + return series == s, ("value", s) + else: + return np.abs(series - s) == np.abs(series - s).min(), ("value", s) + + +def _attach_columns( + part: pd.DataFrame, + cols_tuple: tuple[str, ...], + read_column: Callable[[int, str], npt.NDArray[Any]], + metadtypes: Any, +) -> pd.DataFrame: + if len(part) == 0: + return part.assign(**{c: pd.Series(dtype=metadtypes[c]) for c in cols_tuple}) + st_val = int(part["st"].iloc[0]) + + columns_to_read = tuple(c for c in cols_tuple if c not in part.columns) + arrays = {c: read_column(st_val, c) for c in columns_to_read} + + sel = part["row"].to_numpy() + for c in columns_to_read: + part[c] = np.asarray(arrays[c])[sel] + return part + + +def _load_index_partition( + st: int, + t: float, + index_cols: tuple[str, ...], + read_column: Callable[[int, str], npt.NDArray[Any]], +) -> pd.DataFrame: + cols = {c: read_column(st, c) for c in index_cols} + n = len(next(iter(cols.values()))) + return pd.DataFrame( + { + **cols, + "st": np.full(n, st, dtype=np.int64), + "t": np.full(n, t, dtype=float), + "row": np.arange(n, dtype=np.int64), + } + ) + + +class ParticleDataset: + steps: npt.NDArray[np.int64] + times: npt.NDArray[np.float64] + colnames: list[str] + _ddf_index: dd.DataFrame + + def __init__( + self, + species: list[int], + steps: npt.NDArray[np.int64], + times: npt.NDArray[np.float64], + colnames: list[str], + read_column: Callable[ + [int, str], npt.NDArray[np.float64 | np.int64 | np.float32 | np.int32] + ], + fprec: type | None = np.float32, + selection: dict[str, Selection] | None = None, + ddf_index: dd.DataFrame | None = None, + partition_lengths: Sequence[int] | None = None, + ): + self.species = species + self.steps = steps + self.times = times + self.colnames = colnames + + self.read_column = read_column + self.fprec = fprec + self.index_cols = ("id", "sp") + self._all_columns_cache: list[str] | None = None + self._partition_lengths = ( + tuple(int(length) for length in partition_lengths) + if partition_lengths is not None + else None + ) + + if selection is not None: + self.selection = selection + else: + self.selection = { + "t": Selection("range"), + "st": Selection("range"), + "sp": Selection("range"), + "id": Selection("range"), + } + + self._dtypes = { + "id": np.int64, + "sp": np.int32, + "row": np.int64, + "st": np.int64, + "t": np.float64, + "x": fprec, + "y": fprec, + "z": fprec, + "ux": fprec, + "uy": fprec, + "uz": fprec, + "r": fprec, + "th": fprec, + "ph": fprec, + "ur": fprec, + "uth": fprec, + "uph": fprec, + } + + if ddf_index is not None: + self._ddf_index = ddf_index + self._partition_indices: tuple[int, ...] | None = None + else: + self._ddf_index = self._build_index_ddf() + self._partition_indices = tuple(range(self._ddf_index.npartitions)) + + if ( + self._partition_lengths is not None + and len(self._partition_lengths) != self._ddf_index.npartitions + ): + raise ValueError( + "partition_lengths must contain one value per Dask partition" + ) + + @property + def ddf(self) -> dd.DataFrame: + return self._ddf_index + + @property + def nbytes(self) -> int: + """Estimated bytes occupied by the particle index's NumPy buffers. + + The estimate excludes small pandas/Dask object overhead. For datasets + filtered by particle values (for example ``sel(sp=1)``), it is an upper + bound because finding the exact surviving row count would require + computing the lazy index. + """ + if self._partition_lengths is None: + # Compatibility fallback for ParticleDataset instances constructed + # directly by downstream code without reader-provided metadata. + return int(self.ddf.memory_usage(index=True, deep=True).sum().compute()) + + bytes_per_row = sum( + np.dtype(self._dtypes[column]).itemsize + for column in (*self.index_cols, "st", "t", "row") + ) + return int(sum(self._partition_lengths) * bytes_per_row) + + @property + def columns(self) -> list[str]: + if self._all_columns_cache is None: + self._all_columns_cache = self.colnames + return self._all_columns_cache + + def sel( + self, + t: IntSelector | FloatSelector | None = None, + st: IntSelector | None = None, + sp: IntSelector | None = None, + id: IntSelector | None = None, + method: str = "exact", + ) -> ParticleDataset: + ddf: dd.DataFrame = self._ddf_index + new_selection = {k: copy(v) for k, v in self.selection.items()} + if st is not None: + ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( + st, ddf["st"], method="exact" + ) + ddf = cast(dd.DataFrame, ddf[ddf_sel]) + new_selection["st"] = new_selection["st"].intersect( + Selection(sel_type, sel_value) + ) + if t is not None: + ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( + t, ddf["t"], method=method + ) + ddf = cast(dd.DataFrame, ddf[ddf_sel]) + new_selection["t"] = new_selection["t"].intersect( + Selection(sel_type, sel_value) + ) + if sp is not None: + ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( + sp, ddf["sp"], method="exact" + ) + ddf = cast(dd.DataFrame, ddf[ddf_sel]) + new_selection["sp"] = new_selection["sp"].intersect( + Selection(sel_type, sel_value) + ) + if id is not None: + ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( + id, ddf["id"], method="exact" + ) + ddf = cast(dd.DataFrame, ddf[ddf_sel]) + new_selection["id"] = new_selection["id"].intersect( + Selection(sel_type, sel_value) + ) + + result = ParticleDataset( + species=self.species, + steps=self.steps, + times=self.times, + colnames=self.colnames, + read_column=self.read_column, + fprec=self.fprec, + selection=new_selection, + ddf_index=ddf, + partition_lengths=self._partition_lengths, + ) + result._partition_indices = self._partition_indices + return result + + def isel( + self, t: IntSelector | None = None, st: IntSelector | None = None + ) -> ParticleDataset: + ddf: dd.DataFrame = self._ddf_index + partition_indices = self._partition_indices + partition_lengths = self._partition_lengths + new_selection = {k: v for k, v in self.selection.items()} + for t_or_s, t_or_s_str, t_or_s_arr in zip( + [t, st], ["t", "st"], [self.times, self.steps] + ): + if t_or_s is not None: + selector: Any + if isinstance(t_or_s, slice): + lo = t_or_s.start if t_or_s.start is not None else 0 + hi = t_or_s.stop if t_or_s.stop is not None else -1 + selector = slice(t_or_s_arr[lo], t_or_s_arr[hi]) + elif isinstance(t_or_s, (list, tuple, np.ndarray, pd.Index, pd.Series)): + selector = [t_or_s_arr[ti] for ti in t_or_s] + else: + selector = t_or_s_arr[t_or_s] + + selector_for_mask = cast(Any, selector) + if partition_indices is not None: + partition_mask, _ = _coerce_selector_to_mask( + selector_for_mask, + pd.Series(t_or_s_arr.tolist()), + method="exact", + ) + matching_indices = set(np.flatnonzero(np.asarray(partition_mask))) + selected_partitions = [ + i + for i, original_index in enumerate(partition_indices) + if original_index in matching_indices + ] + if not selected_partitions: + # Keep one partition for the filter to turn into an empty frame. + selected_partitions = [0] + ddf = cast(dd.DataFrame, ddf.partitions[selected_partitions]) + partition_indices = tuple( + partition_indices[i] for i in selected_partitions + ) + if partition_lengths is not None: + partition_lengths = tuple( + partition_lengths[i] for i in selected_partitions + ) + + ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( + selector_for_mask, + ddf[t_or_s_str], + method="exact", + ) + ddf = cast(dd.DataFrame, ddf[ddf_sel]) + new_selection[t_or_s_str] = new_selection[t_or_s_str].intersect( + Selection(sel_type, sel_value) + ) + + if partition_indices is not None: + ddf = _cull_dataframe_graph(ddf) + + result = ParticleDataset( + species=self.species, + steps=self.steps, + times=self.times, + colnames=self.colnames, + read_column=self.read_column, + fprec=self.fprec, + selection=new_selection, + ddf_index=ddf, + partition_lengths=partition_lengths, + ) + result._partition_indices = partition_indices + return result + + def _build_index_ddf(self) -> dd.DataFrame: + delayed_parts = [ + dask.delayed(_load_index_partition)( + st, + t, + self.index_cols, + self.read_column, + ) + for st, t in zip(self.steps, self.times) + ] + + meta = pd.DataFrame( + { + **{ + c: np.array([], dtype=self._dtypes.get(c, "O")) + for c in self.index_cols + }, + "st": np.array([], dtype=self._dtypes.get("st", np.int64)), + "t": np.array([], dtype=self._dtypes.get("t", np.int64)), + "row": np.array([], dtype=self._dtypes.get("row", np.int64)), + } + ) + + ddf = cast(dd.DataFrame, dd.from_delayed(delayed_parts, meta=meta)) + return ddf + + def load(self, cols: Sequence[str] | None = None) -> pd.DataFrame: + if cols is None: + cols = self.columns + + cols = [c for c in cols if c not in ("t", "st", "row")] + + meta_dict = { + c: np.array([], dtype=self._dtypes.get(c, np.float64)) for c in cols + } + meta = self._ddf_index._meta.assign(**meta_dict) + + cols_tuple = tuple(cols) + + return ( + self._ddf_index.map_partitions( + _attach_columns, + cols_tuple=cols_tuple, + read_column=self.read_column, + metadtypes=meta.dtypes, + meta=meta, + ) + .compute() + .drop(columns=["row"]) + ) + + def help(self, prepend="") -> str: + ret = f"{prepend}- use .sel(...) to select particles based on criteria:\n" + ret += f"{prepend} t : time (float)\n" + ret += f"{prepend} st : step (int)\n" + ret += f"{prepend} sp : species (int)\n" + ret += f"{prepend} id : particle id (int)\n{prepend}\n" + ret += f"{prepend} # example:\n" + ret += f"{prepend} # .sel(t=slice(10.0, 20.0), sp=[1, 2, 3], id=[42, 22])\n{prepend}\n" + ret += f"{prepend}- use .isel(...) to select particles based on output step:\n" + ret += f"{prepend} t : timestamp index (int)\n" + ret += f"{prepend} st : step index (int)\n{prepend}\n" + ret += f"{prepend} # example:\n" + ret += f"{prepend} # .isel(t=-1)\n" + ret += f"{prepend}\n" + ret += f"{prepend}- .sel and .isel can be chained together:\n{prepend}\n" + ret += f"{prepend} # example:\n" + ret += f"{prepend} # .isel(t=-1).sel(sp=1).sel(id=[55, 66])\n{prepend}\n" + ret += f"{prepend}- use .load(cols=[...]) to load data into a pandas DataFrame (`cols` defaults to all columns)\n{prepend}\n" + ret += f"{prepend} # example:\n" + ret += f"{prepend} # .sel(...).load()\n" + return ret + + def __repr__(self) -> str: + ret = "ParticleDataset:\n" + ret += "================\n" + ret += f"Variables:\n {self.columns}\n\n" + ret += "Current selection:\n" + for k, v in self.selection.items(): + ret += f" {k:<5} : {v}\n" + ret += "\nHelp:\n" + ret += "-----\n" + ret += f"{self.help()}" + return ret + + def __str__(self) -> str: + return self.__repr__() + + def spectrum_plot( + self, + ax: maxes.Axes | None = None, + bins: npt.NDArray | None = None, + quantity: Callable[[pd.DataFrame], npt.NDArray] | None = None, + ): + if ax is None: + ax = plt.gca() + + def _uSqr_cart(df: pd.DataFrame): + return np.sum( + [ + np.asarray(df[c].to_numpy(), dtype=np.float64) ** 2 + for c in ["ux", "uy", "uz"] + ], + axis=0, + ) + + def _quantity_cart(df: pd.DataFrame): + uSqr = _uSqr_cart(df) + return uSqr * np.sqrt(1.0 + uSqr) + + def _uSqr_sph(df: pd.DataFrame): + return np.sum( + [ + np.asarray(df[c].to_numpy(), dtype=np.float64) ** 2 + for c in ["ur", "uth", "uph"] + ], + axis=0, + ) + + def _quantity_sph(df: pd.DataFrame): + uSqr = _uSqr_sph(df) + return uSqr * np.sqrt(1.0 + uSqr) + + if "ux" in self.columns: + cols = ["ux", "uy", "uz"] + if quantity is None: + quantity = _quantity_cart + else: + cols = ["ur", "uth", "uph"] + if quantity is None: + quantity = _quantity_sph + assert quantity is not None + df = self.load(cols=["sp", *cols]) + species = sorted(df["sp"].unique()) + arrays = { + sp: quantity(group.drop(columns=["sp"])) for sp, group in df.groupby("sp") + } + if bins is None: + bins = np.logspace(0, 4, 100) + hists = {sp: np.histogram(arrays[sp], bins=bins)[0] for sp in species} + bins = 0.5 * (bins[1:] + bins[:-1]) + for sp in species: + ax.loglog(bins, hists[sp], label=f"{sp}") + if bins.min() > 0 and bins.max() / bins.min() > 100: + ax.set(xscale="log", yscale="log") + + def phase_plot( + self, + ax: maxes.Axes | None = None, + x_quantity: Callable[[pd.DataFrame], npt.NDArray] | None = None, + y_quantity: Callable[[pd.DataFrame], npt.NDArray] | None = None, + xy_bins: tuple[npt.NDArray, npt.NDArray] | None = None, + **kwargs: Any, + ): + if ax is None: + ax = plt.gca() + + def _xquantity_cart(df: pd.DataFrame): + return np.asarray(df["x"].to_numpy(), dtype=np.float64) + + def _yquantity_cart(df: pd.DataFrame): + return np.asarray(df["ux"].to_numpy(), dtype=np.float64) + + def _xquantity_sph(df: pd.DataFrame): + return np.asarray(df["r"].to_numpy(), dtype=np.float64) + + def _yquantity_sph(df: pd.DataFrame): + return np.asarray(df["ur"].to_numpy(), dtype=np.float64) + + if "ux" in self.columns: + cols = ["ux", "uy", "uz"] + for c in "xyz": + if c in self.columns: + cols.append(c) + if x_quantity is None: + x_quantity = _xquantity_cart + if y_quantity is None: + y_quantity = _yquantity_cart + else: + cols = ["ur", "uth", "uph"] + for c in ["r", "th", "ph"]: + if c in self.columns: + cols.append(c) + if x_quantity is None: + x_quantity = _xquantity_sph + if y_quantity is None: + y_quantity = _yquantity_sph + + df = self.load(cols=[*cols]) + x_array = x_quantity(df) + y_array = y_quantity(df) + + if xy_bins is None: + x_bins = np.linspace(x_array.min(), x_array.max(), 100) + y_bins = np.linspace(y_array.min(), y_array.max(), 100) + xy_bins = (x_bins, y_bins) + else: + x_bins, y_bins = xy_bins + + h2d, xedges, yedges = np.histogram2d(x_array, y_array, bins=[x_bins, y_bins]) + X, Y = np.meshgrid( + 0.5 * (xedges[1:] + xedges[:-1]), 0.5 * (yedges[1:] + yedges[:-1]) + ) + pcm = ax.pcolormesh( + X, + Y, + h2d.T, + shading="auto", + rasterized=True, + **kwargs, + ) + return pcm diff --git a/nt2/containers/particles.py b/nt2/containers/particles.py index a352e4e..e19c76d 100644 --- a/nt2/containers/particles.py +++ b/nt2/containers/particles.py @@ -1,612 +1,194 @@ -from typing import ( - Any, - Callable, - List, - Optional, - Sequence, - Tuple, - Literal, - Union, - Dict, - Type, -) -import numpy.typing as npt -from copy import copy - -import dask -import dask.dataframe as dd -import pandas as pd -import numpy as np - -import matplotlib.pyplot as plt -import matplotlib.axes as maxes +from __future__ import annotations -from nt2.containers.container import BaseContainer - - -IntSelector = Union[int, Sequence[int], slice, Tuple[int, int]] -FloatSelector = Union[float, slice, Sequence[float], Tuple[float, float]] - - -class Selection: - def __init__( - self, - type: Literal["value", "range", "list"], - value: Optional[Union[int, float, list, tuple]] = None, - ): - self.type = type - self.value = value - - def intersect(self, other: "Selection") -> "Selection": - if self.value is None: - return copy(other) - elif other.value is None: - return copy(self) - if self.type == "value" and other.type == "value": - if self.value == other.value: - return Selection("value", self.value) - else: - return Selection("value") - elif self.type == "value" and other.type == "list": - assert isinstance(other.value, list), "other.value must be a list" - if self.value in other.value: - return Selection("value", self.value) - else: - return Selection("value") - elif self.type == "value" and other.type == "range": - assert ( - isinstance(other.value, tuple) and len(other.value) == 2 - ), "other.value must be a tuple of length 2" - lo, hi = other.value - if lo <= self.value < hi: - return Selection("value", self.value) - else: - return Selection("value") - elif self.type == "list" and other.type == "value": - return other.intersect(self) - elif self.type == "list" and other.type == "list": - assert isinstance(self.value, list), "self.value must be a list" - assert isinstance(other.value, list), "other.value must be a list" - new_values = [v for v in self.value if v in other.value] - return Selection("list", new_values) - elif self.type == "list" and other.type == "range": - assert ( - isinstance(other.value, tuple) and len(other.value) == 2 - ), "other.value must be a tuple of length 2" - assert isinstance(self.value, list), "self.value must be a list" - lo, hi = other.value - new_values = [v for v in self.value if lo <= v <= hi] - return Selection("list", new_values) - elif self.type == "range" and other.type == "value": - return other.intersect(self) - elif self.type == "range" and other.type == "list": - return other.intersect(self) - elif self.type == "range" and other.type == "range": - assert ( - isinstance(self.value, tuple) and len(self.value) == 2 - ), "self.value must be a tuple of length 2" - assert ( - isinstance(other.value, tuple) and len(other.value) == 2 - ), "other.value must be a tuple of length 2" - lo1, hi1 = self.value - lo2, hi2 = other.value - new_lo = max(lo1, lo2) - new_hi = min(hi1, hi2) - if new_lo <= new_hi: - return Selection("range", (new_lo, new_hi)) - else: - return Selection("value") - else: - raise ValueError(f"Unknown selection types: {self.type}, {other.type}") - - def __repr__(self) -> str: - if self.type == "value": - return "all" if self.value is None else f"{self.value:.3g}" - elif self.type == "range": - if self.value is None: - return "all" - else: - assert ( - isinstance(self.value, tuple) and len(self.value) == 2 - ), "value must be a tuple of length 2" - lo, hi = self.value - lo_str = "..." if lo is None or lo == -np.inf else f"{lo:.3g}" - hi_str = "..." if hi is None or hi == np.inf else f"{hi:.3g}" - return f"[ {lo_str} -> {hi_str} ]" - elif self.type == "list": - assert isinstance(self.value, list), "value must be a list" - return "{ " + ", ".join(f"{v:.3g}" for v in self.value) + " }" - else: - return "InvalidSelection" - - def __str__(self) -> str: - return self.__repr__() - - -def _coerce_selector_to_mask( - s: Union[IntSelector, FloatSelector], - series: Any, - inclusive_tuple: bool = True, - method="exact", -): - from operator import ior - from functools import reduce - - if isinstance(s, slice): - lo = s.start if s.start is not None else -np.inf - hi = s.stop if s.stop is not None else np.inf - step = s.step - mask = (series >= lo) & (series <= hi) - if step not in (None, 1): - mask = mask & (((series - lo) % step) == 0) - return mask, ("range", (lo, hi)) - elif isinstance(s, tuple) and len(s) == 2 and inclusive_tuple: - lo, hi = s - if lo is None: - lo = -np.inf - if hi is None: - hi = np.inf - return (series >= lo) & (series <= hi), ("range", (lo, hi)) - elif isinstance(s, (list, tuple, np.ndarray, pd.Index, pd.Series)): - if method == "exact": - return series.isin(list(s)), ("list", list(s)) - else: - return reduce( - ior, [np.abs(series - v) == np.abs(series - v).min() for v in s] - ), ("list", list(s)) - else: - if method == "exact": - return series == s, ("value", s) - else: - return np.abs(series - s) == np.abs(series - s).min(), ("value", s) - - -def _attach_columns( - part: pd.DataFrame, - cols_tuple, - read_column, - metadtypes, -) -> pd.DataFrame: - if len(part) == 0: - for c in cols_tuple: - part[c] = np.array([], dtype=metadtypes[c]) - return part - st_val = int(part["st"].iloc[0]) - - arrays = {c: read_column(st_val, c) for c in cols_tuple} - - sel = part["row"].to_numpy() - for c in cols_tuple: - part[c] = np.asarray(arrays[c])[sel] - return part +from typing import Any +import numpy as np +import numpy.typing as npt -class ParticleDataset: - steps: npt.NDArray[np.int64] - times: npt.NDArray[np.float64] - colnames: List[str] +from ..utils import CoordinateSystem +from .base import BaseContainer +from .particle_dataset import ParticleDataset + + +def remap_prtl_quantities_cart(name: str) -> str: + shortname = name[1:] + return { + "X1": "x", + "X2": "y", + "X3": "z", + "U1": "ux", + "U2": "uy", + "U3": "uz", + "W": "w", + }.get(shortname, shortname) + + +def remap_prtl_quantities_sph(name: str) -> str: + shortname = name[1:] + return { + "X1": "r", + "X2": "th", + "X3": "ph", + "U1": "ur", + "U2": "uth", + "U3": "uph", + "W": "w", + }.get(shortname, shortname) + + +class ParticleContainer(BaseContainer): + """Parent class to manage the particles dataframe.""" - def __init__( - self, - species: List[int], - steps: npt.NDArray[np.int64], - times: npt.NDArray[np.float64], - colnames: List[str], - read_column: Callable[ - [int, str], npt.NDArray[Union[np.float64, np.int64, np.float32, np.int32]] - ], - fprec: Optional[Type] = np.float32, - selection: Optional[Dict[str, Selection]] = None, - ddf_index: Optional[dd.DataFrame] = None, - ): - self.species = species - self.steps = steps - self.times = times - self.colnames = colnames - - self.read_column = read_column - self.fprec = fprec - self.index_cols = ("id", "sp") - self._all_columns_cache: Optional[List[str]] = None - - if selection is not None: - self.selection = selection - else: - self.selection = { - "t": Selection("range"), - "st": Selection("range"), - "sp": Selection("range"), - "id": Selection("range"), + __particles_defined: bool = False + __particles: ParticleDataset | None = None + + nonempty_steps: list[int] + attributes: dict[str, Any] + quantities: list[str] + sp_with_idx: list[int] + sp_without_idx: list[int] + quantity_names_by_step: dict[int, set[str]] + + def __getstate__(self) -> dict[str, Any]: + state = self.__dict__.copy() + particles = state.get("_ParticleContainer__particles") + if particles is not None: + state["_ParticleContainer__particle_dataset_state"] = { + "species": particles.species, + "steps": particles.steps, + "times": particles.times, + "colnames": particles.colnames, + "fprec": particles.fprec, + "selection": particles.selection, + "partition_lengths": particles._partition_lengths, } + state.pop("_ParticleContainer__particles", None) + return state - self._dtypes = { - "id": np.int64, - "sp": np.int32, - "row": np.int64, - "st": np.int64, - "t": fprec, - "x": fprec, - "y": fprec, - "z": fprec, - "ux": fprec, - "uy": fprec, - "uz": fprec, - "r": fprec, - "th": fprec, - "ph": fprec, - "ur": fprec, - "uth": fprec, - "uph": fprec, - } - - if ddf_index is not None: - self._ddf_index = ddf_index - else: - self._ddf_index = self._build_index_ddf() + def __setstate__(self, state: dict[str, Any]) -> None: + particle_dataset_state = state.pop( + "_ParticleContainer__particle_dataset_state", None + ) + self.__dict__.update(state) + if self.__particles_defined: + if particle_dataset_state is None: + ( + self.quantities, + self.sp_with_idx, + self.sp_without_idx, + self.attributes, + self.__particles, + ) = self._read_particles() + else: + self.__particles = ParticleDataset( + **particle_dataset_state, + read_column=self._read_column, + ) - @property - def ddf(self) -> dd.DataFrame: - return self._ddf_index + def __init__(self, **kwargs: Any) -> None: + """Initializer for the ParticleContainer class. - @property - def nbytes(self) -> int: - return self.ddf.memory_usage(index=True, deep=True).sum().compute() + Parameters + ---------- + **kwargs : dict + Keyword arguments to be passed to the parent BaseContainer class. - @property - def columns(self) -> List[str]: - if self._all_columns_cache is None: - self._all_columns_cache = self.colnames - return self._all_columns_cache + """ + super().__init__(category="particles", **kwargs) - def sel( - self, - t: Optional[Union[IntSelector, FloatSelector]] = None, - st: Optional[IntSelector] = None, - sp: Optional[IntSelector] = None, - id: Optional[IntSelector] = None, - method: str = "exact", - ) -> "ParticleDataset": - ddf = self._ddf_index - new_selection = {k: copy(v) for k, v in self.selection.items()} - if st is not None: - ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( - st, ddf["st"], method="exact" - ) - ddf = ddf[ddf_sel] - new_selection["st"] = new_selection["st"].intersect( - Selection(sel_type, sel_value) - ) - if t is not None: - ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( - t, ddf["t"], method=method - ) - ddf = ddf[ddf_sel] - new_selection["t"] = new_selection["t"].intersect( - Selection(sel_type, sel_value) - ) - if sp is not None: - ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( - sp, ddf["sp"], method="exact" + # @TODO: parallelize + self.quantity_names_by_step = { + step: self.reader.ReadCategoryNamesAtTimestep( + self.path, "particles", "p", step ) - ddf = ddf[ddf_sel] - new_selection["sp"] = new_selection["sp"].intersect( - Selection(sel_type, sel_value) - ) - if id is not None: - ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( - id, ddf["id"], method="exact" - ) - ddf = ddf[ddf_sel] - new_selection["id"] = new_selection["id"].intersect( - Selection(sel_type, sel_value) - ) - - return ParticleDataset( - species=self.species, - steps=self.steps, - times=self.times, - colnames=self.colnames, - read_column=self.read_column, - fprec=self.fprec, - selection=new_selection, - ddf_index=ddf, - ) + for step in self.valid_steps + } + self.nonempty_steps = [ + step + for step, names in self.quantity_names_by_step.items() + if any(q.startswith("p") for q in names) + ] - def isel( - self, t: Optional[IntSelector] = None, st: Optional[IntSelector] = None - ) -> "ParticleDataset": - ddf = self._ddf_index - new_selection = {k: v for k, v in self.selection.items()} - for t_or_s, t_or_s_str, t_or_s_arr in zip( - [t, st], ["t", "st"], [self.times, self.steps] + if ( + self.reader.DefinesCategory(self.path, "particles", self.valid_files) + and len(self.nonempty_steps) > 0 ): - if t_or_s is not None: - if isinstance(t_or_s, slice): - lo = t_or_s.start if t_or_s.start is not None else 0 - hi = t_or_s.stop if t_or_s.stop is not None else -1 - ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( - slice(t_or_s_arr[lo], t_or_s_arr[hi]), - ddf[t_or_s_str], - method="exact", - ) - ddf = ddf[ddf_sel] - new_selection[t_or_s_str] = new_selection[t_or_s_str].intersect( - Selection(sel_type, sel_value) - ) - elif isinstance(t_or_s, (list, tuple, np.ndarray, pd.Index, pd.Series)): - ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( - [t_or_s_arr[ti] for ti in t_or_s], - ddf[t_or_s_str], - method="exact", - ) - ddf = ddf[ddf_sel] - new_selection[t_or_s_str] = new_selection[t_or_s_str].intersect( - Selection(sel_type, sel_value) - ) - else: - ddf_sel, (sel_type, sel_value) = _coerce_selector_to_mask( - t_or_s_arr[t_or_s], ddf[t_or_s_str], method="exact" - ) - ddf = ddf[ddf_sel] - new_selection[t_or_s_str] = new_selection[t_or_s_str].intersect( - Selection(sel_type, sel_value) - ) - return ParticleDataset( - species=self.species, - steps=self.steps, - times=self.times, - colnames=self.colnames, - read_column=self.read_column, - fprec=self.fprec, - selection=new_selection, - ddf_index=ddf, - ) - - def _load_index_partition(self, st: int, t: float, index_cols: Tuple[str, ...]): - cols = {c: self.read_column(st, c) for c in index_cols} - n = len(next(iter(cols.values()))) - df = pd.DataFrame(cols) - df["st"] = np.asarray(st, dtype=np.int64) - df["t"] = np.asarray(t, dtype=float) - df["row"] = np.arange(n, dtype=np.int64) - return df - - def _build_index_ddf(self) -> dd.DataFrame: - delayed_parts = [ - dask.delayed(self._load_index_partition)(st, t, self.index_cols) - for st, t in zip(self.steps, self.times) + self.__particles_defined = True + ( + self.quantities, + self.sp_with_idx, + self.sp_without_idx, + self.attributes, + self.__particles, + ) = self._read_particles() + + def _read_particles(self): + # read unique quantities and species + quantities_ = [ + self.quantity_names_by_step[step] for step in self.nonempty_steps ] + quantities = sorted(np.unique([q for qtys in quantities_ for q in qtys])) - meta = pd.DataFrame( + unique_quantities = sorted( { - **{ - c: np.array([], dtype=self._dtypes.get(c, "O")) - for c in self.index_cols - }, - "st": np.array([], dtype=self._dtypes.get("st", np.int64)), - "t": np.array([], dtype=self._dtypes.get("t", np.int64)), - "row": np.array([], dtype=self._dtypes.get("row", np.int64)), + f"{q}".split("_")[0] + for q in quantities + if not q.startswith("pIDX") and not q.startswith("pRNK") } ) + all_species = sorted({int(f"{q}".split("_")[1]) for q in quantities}) - ddf = dd.from_delayed(delayed_parts, meta=meta) - return ddf - - def load(self, cols: Optional[Sequence[str]] = None) -> pd.DataFrame: - if cols is None: - cols = self.columns - - cols = [c for c in cols if c not in ("t", "st", "row")] - - meta_dict = { - c: np.array([], dtype=self._dtypes.get(c, np.float64)) for c in cols - } - meta = self._ddf_index._meta.assign(**meta_dict) - - cols_tuple = tuple(cols) + sp_with_idx = sorted( + [int(f"{q}".split("_")[1]) for q in quantities if f"{q}".startswith("pIDX")] + ) + sp_without_idx = sorted([sp for sp in all_species if sp not in sp_with_idx]) - return ( - self._ddf_index.map_partitions( - _attach_columns, - cols_tuple=cols_tuple, - read_column=self.read_column, - metadtypes=meta.dtypes, - meta=meta, + partition_lengths = tuple( + sum( + self.reader.ReadParticleCountsAtTimestep( + self.path, step, all_species + ).values() ) - .compute() - .drop(columns=["row"]) + for step in self.valid_steps ) - def help(self, prepend="") -> str: - ret = f"{prepend}- use .sel(...) to select particles based on criteria:\n" - ret += f"{prepend} t : time (float)\n" - ret += f"{prepend} st : step (int)\n" - ret += f"{prepend} sp : species (int)\n" - ret += f"{prepend} id : particle id (int)\n{prepend}\n" - ret += f"{prepend} # example:\n" - ret += f"{prepend} # .sel(t=slice(10.0, 20.0), sp=[1, 2, 3], id=[42, 22])\n{prepend}\n" - ret += f"{prepend}- use .isel(...) to select particles based on output step:\n" - ret += f"{prepend} t : timestamp index (int)\n" - ret += f"{prepend} st : step index (int)\n{prepend}\n" - ret += f"{prepend} # example:\n" - ret += f"{prepend} # .isel(t=-1)\n" - ret += f"{prepend}\n" - ret += f"{prepend}- .sel and .isel can be chained together:\n{prepend}\n" - ret += f"{prepend} # example:\n" - ret += f"{prepend} # .isel(t=-1).sel(sp=1).sel(id=[55, 66])\n{prepend}\n" - ret += f"{prepend}- use .load(cols=[...]) to load data into a pandas DataFrame (`cols` defaults to all columns)\n{prepend}\n" - ret += f"{prepend} # example:\n" - ret += f"{prepend} # .sel(...).load()\n" - return ret - - def __repr__(self) -> str: - ret = "ParticleDataset:\n" - ret += "================\n" - ret += f"Variables:\n {self.columns}\n\n" - ret += "Current selection:\n" - for k, v in self.selection.items(): - ret += f" {k:<5} : {v}\n" - ret += "\nHelp:\n" - ret += "-----\n" - ret += f"{self.help()}" - return ret - - def __str__(self) -> str: - return self.__repr__() - - def spectrum_plot( - self, - ax: Optional[maxes.Axes] = None, - bins: Optional[npt.NDArray] = None, - quantity: Optional[Callable[[pd.DataFrame], npt.NDArray]] = None, - ): - if ax is None: - ax = plt.gca() - - if "ux" in self.columns: - cols = ["ux", "uy", "uz"] - if quantity is None: - uSqr = lambda df: np.sum( - [df[c].to_numpy(dtype=np.float64) ** 2 for c in cols], axis=0 - ) - quantity = lambda df: uSqr(df) * np.sqrt(1.0 + uSqr(df)) - else: - cols = ["ur", "uth", "uph"] - if quantity is None: - uSqr = lambda df: np.sum( - [df[c].to_numpy(dtype=np.float64) ** 2 for c in cols], axis=0 - ) - quantity = lambda df: uSqr(df) * np.sqrt(1.0 + uSqr(df)) - df = self.load(cols=["sp", *cols]) - species = sorted(df["sp"].unique()) - arrays = df.groupby("sp").apply(quantity, include_groups=False) - if bins is None: - bins = np.logspace(0, 4, 100) - hists = {sp: np.histogram(arrays[sp], bins=bins)[0] for sp in species} - bins = 0.5 * (bins[1:] + bins[:-1]) - for sp in species: - ax.loglog(bins, hists[sp], label=f"{sp}") - if bins.min() > 0 and bins.max() / bins.min() > 100: - ax.set(xscale="log", yscale="log") - - def phase_plot( - self, - ax: Optional[maxes.Axes] = None, - x_quantity: Optional[Callable[[pd.DataFrame], np.ndarray]] = None, - y_quantity: Optional[Callable[[pd.DataFrame], np.ndarray]] = None, - xy_bins: Optional[Tuple[npt.NDArray, npt.NDArray]] = None, - **kwargs: Any, - ): - if ax is None: - ax = plt.gca() - - if "ux" in self.columns: - cols = ["ux", "uy", "uz"] - for c in "xyz": - if c in self.columns: - cols.append(c) - if x_quantity is None: - x_quantity = lambda df: df["x"].to_numpy(dtype=np.float64) - if y_quantity is None: - y_quantity = lambda df: df["ux"].to_numpy(dtype=np.float64) - else: - cols = ["ur", "uth", "uph"] - for c in ["r", "th", "ph"]: - if c in self.columns: - cols.append(c) - if x_quantity is None: - x_quantity = lambda df: df["r"].to_numpy(dtype=np.float64) - if y_quantity is None: - y_quantity = lambda df: df["ur"].to_numpy(dtype=np.float64) - - df = self.load(cols=[*cols]) - x_array = x_quantity(df) - y_array = y_quantity(df) - - if xy_bins is None: - x_bins = np.linspace(x_array.min(), x_array.max(), 100) - y_bins = np.linspace(y_array.min(), y_array.max(), 100) - xy_bins = (x_bins, y_bins) - else: - x_bins, y_bins = xy_bins - - h2d, xedges, yedges = np.histogram2d(x_array, y_array, bins=[x_bins, y_bins]) - X, Y = np.meshgrid( - 0.5 * (xedges[1:] + xedges[:-1]), 0.5 * (yedges[1:] + yedges[:-1]) + # determine coordinate system and remap functions + first_step = self.valid_steps[0] + attributes = self.reader.ReadAttrsAtTimestep( + path=self.path, category="particles", step=first_step ) - pcm = ax.pcolormesh( - X, - Y, - h2d.T, - shading="auto", - rasterized=True, - **kwargs, - ) - return pcm - - -class Particles(BaseContainer): - """Parent class to manage the particles dataframe.""" - - __particles_defined: bool - __particles: Optional[ParticleDataset] - quantities: List[str] - sp_with_idx: List[int] - sp_without_idx: List[int] - - def __init__(self, **kwargs: Any) -> None: - """Initializer for the Particles class. - - Parameters - ---------- - **kwargs : dict - Keyword arguments to be passed to the parent BaseContainer class. - - """ - super(Particles, self).__init__(**kwargs) - if ( - self.reader.DefinesCategory(self.path, "particles") - and self.particles_present - ): - self.__particles_defined = True - - valid_steps = self.nonempty_steps - quantities_ = [ - self.reader.ReadCategoryNamesAtTimestep( - self.path, "particles", "p", step - ) - for step in valid_steps - ] - self.quantities = sorted( - np.unique([q for qtys in quantities_ for q in qtys]) - ) - - unique_quantities = sorted( - list( - set( - str(q).split("_")[0] - for q in self.quantities - if not q.startswith("pIDX") and not q.startswith("pRNK") - ) + if self.coordinate_system is None: + if "Coordinates" not in attributes: + raise ValueError("Coordinates not found in attributes for particles.") + if attributes["Coordinates"] in [b"cart", "cart"]: + self.set_coordinate_system(CoordinateSystem.XYZ) + elif attributes["Coordinates"] in [b"sph", "sph", b"qsph", "qsph"]: + self.set_coordinate_system(CoordinateSystem.SPH) + else: + raise NotImplementedError( + f"Coordinate system {attributes['Coordinates']} not supported." ) - ) - all_species = sorted( - list(set([int(str(q).split("_")[1]) for q in self.quantities])) - ) - self.sp_with_idx = sorted( - [int(q.split("_")[1]) for q in self.quantities if q.startswith("pIDX")] - ) - self.sp_without_idx = sorted( - [sp for sp in all_species if sp not in self.sp_with_idx] + if self.remap is None: + self.set_remap( + { + "particles": ( + remap_prtl_quantities_cart + if self.coordinate_system == CoordinateSystem.XYZ + else remap_prtl_quantities_sph + ), + } ) - self.__particles = ParticleDataset( + return ( + quantities, + sp_with_idx, + sp_without_idx, + attributes, + ParticleDataset( species=all_species, - steps=np.array(self.reader.GetValidSteps(self.path, "particles")), - times=self.reader.ReadPerTimestepVariable( - self.path, "particles", "Time", "t" - )["t"], + steps=self.steps, + times=self.times, colnames=[ ( self.remap["particles"](q) @@ -617,34 +199,9 @@ def __init__(self, **kwargs: Any) -> None: ] + ["id", "sp"], read_column=self._read_column, - ) - else: - self.__particles_defined = False - self.__particles = None - - @property - def particles_present(self) -> bool: - """bool: Whether the particles are present in any of the timesteps.""" - return len(self.nonempty_steps) > 0 - - @property - def nonempty_steps(self) -> List[int]: - """list[int]: List of timesteps that contain particles data.""" - valid_steps = self.reader.GetValidSteps(self.path, "particles") - return [ - step - for step in valid_steps - if len( - set( - q.split("_")[0] - for q in self.reader.ReadCategoryNamesAtTimestep( - self.path, "particles", "p", step - ) - if q.startswith("p") - ) - ) - > 0 - ] + partition_lengths=partition_lengths, + ), + ) @property def particles_defined(self) -> bool: @@ -652,17 +209,25 @@ def particles_defined(self) -> bool: return self.__particles_defined @property - def particles(self) -> Optional[ParticleDataset]: + def particles(self) -> ParticleDataset | None: """Returns the particles data. Returns ------- - ParticleDataset - Dictionary of datasets for each step. + ParticleDataset | None + The particles data if defined, otherwise None. """ return self.__particles + @property + def attrs(self) -> dict[str, Any]: + """dict: The attributes of the particles dataframe.""" + if self.particles_defined: + return self.attributes + else: + return {} + def help_particles(self, prepend: str = "") -> str: return self.particles.help(prepend) if self.particles is not None else "" @@ -673,7 +238,7 @@ def _get_count(self, step: int, sp: int) -> np.int64: self.path, "particles", f"pX1_{sp}", step )[0] ) - except: + except Exception: return np.int64(0) def _species_has_quantity(self, read_colname: str, step: int, sp: int) -> bool: @@ -686,7 +251,7 @@ def _get_quantity_for_species( read_colname: str, step: int, sp: int, - ) -> npt.NDArray[Union[np.float64, np.int64]]: + ) -> npt.NDArray[np.float64 | np.int64]: if f"{read_colname}_{sp}" in self.quantities: return self.reader.ReadArrayAtTimestep( self.path, "particles", f"{read_colname}_{sp}", step @@ -696,7 +261,7 @@ def _get_quantity_for_species( def _read_column( self, step: int, colname: str - ) -> npt.NDArray[Union[np.float64, np.int64, np.float32, np.int32]]: + ) -> npt.NDArray[np.float64 | np.int64 | np.float32 | np.int32]: read_colname = None if colname == "id": idx = np.concatenate( diff --git a/nt2/containers/spectra.py b/nt2/containers/spectra.py index 044656d..6add10c 100644 --- a/nt2/containers/spectra.py +++ b/nt2/containers/spectra.py @@ -1,50 +1,49 @@ +from __future__ import annotations + from typing import Any import dask import dask.array as da -import xarray as xr import numpy as np +import xarray as xr +from tqdm import tqdm -from nt2.containers.container import BaseContainer -from nt2.readers.base import BaseReader +from ..utils import CoordinateSystem +from .base import BaseContainer -class Spectra(BaseContainer): - """Parent class to manager the spectra dataframe.""" +def remap_coords_cart(name: str) -> str: + return { + "X1": "x", + "X2": "y", + "X3": "z", + }.get(name, name) - @staticmethod - def read_spectrum(path: str, reader: BaseReader, spectrum: str, step: int) -> Any: - """Reads a spectrum from the data. - This is a dask-delayed function used further to build the dataset. +def remap_coords_sph(name: str) -> str: + return { + "X1": "r", + "X2": "th", + "X3": "ph", + }.get(name, name) - Parameters - ---------- - path : str - Main path to the data. - reader : BaseReader - Reader to use to read the data. - spectrum : str - Spectrum array to read. - step : int - Step to read. - Returns - ------- - Any - Spectrum data. +class SpectraContainer(BaseContainer): + """Parent class to manage the spectra dataframe.""" - """ - return reader.ReadArrayAtTimestep(path, "spectra", spectrum, step) + __spectra_defined: bool = False + __spectra: xr.Dataset | None = None def __init__(self, **kwargs: Any) -> None: - super(Spectra, self).__init__(**kwargs) - if self.reader.DefinesCategory(self.path, "spectra"): + super().__init__(category="spectra", **kwargs) + + if self.reader.DefinesCategory( + self.path, + "spectra", + self.valid_files, + ): self.__spectra_defined = True - self.__spectra = self.__read_spectra() - else: - self.__spectra_defined = False - self.__spectra = xr.Dataset() + self.__spectra = self._read_spectra() @property def spectra_defined(self) -> bool: @@ -52,38 +51,117 @@ def spectra_defined(self) -> bool: return self.__spectra_defined @property - def spectra(self) -> xr.Dataset: + def spectra(self) -> xr.Dataset | None: """xr.Dataset: The spectra dataframe.""" return self.__spectra - def __read_spectra(self) -> xr.Dataset: - self.reader.VerifySameCategoryNames(self.path, "spectra", "s") - valid_steps = sorted(self.reader.GetValidSteps(self.path, "spectra")) + def _read_spectrum(self, spectrum: str, step: int) -> Any: + """Reads a spectrum from the data. + + This is a dask-delayed function used further to build the dataset. + + Parameters + ---------- + spectrum : str + Spectrum array to read. + step : int + Step to read. + + Returns + ------- + Any + Spectrum data. + + """ + return self.reader.ReadArrayAtTimestep(self.path, "spectra", spectrum, step) + + def _read_spectra(self) -> xr.Dataset: + if self.verify: + self.reader.VerifySameCategoryNames( + self.path, + "spectra", + "s", + self.valid_steps, + ) + first_step = self.valid_steps[0] spectra_names = self.reader.ReadCategoryNamesAtTimestep( - self.path, "spectra", "s", valid_steps[0] + self.path, "spectra", "s", first_step ) - spectra_names = set(s for s in sorted(spectra_names) if s.startswith("sN")) + spectra_names = {s for s in sorted(spectra_names) if s.startswith("sN")} ebin_name = "sEbn" - first_step = valid_steps[0] first_spectrum_name = next(iter(spectra_names)) shape = self.reader.ReadArrayShapeExplicitlyAtTimestep( self.path, "spectra", first_spectrum_name, first_step ) - times = self.reader.ReadPerTimestepVariable(self.path, "spectra", "Time", "t") - steps = self.reader.ReadPerTimestepVariable(self.path, "spectra", "Step", "s") - ebins = self.reader.ReadArrayAtTimestep( + energy_binedges = self.reader.ReadArrayAtTimestep( self.path, "spectra", ebin_name, first_step ) + num_spatial_dims = len(shape) - 1 + + x_binedges = [ + self.reader.ReadArrayAtTimestep( + self.path, "spectra", f"sX{i + 1}bn", first_step + ) + for i in range(num_spatial_dims) + ] + + def edges_to_bins(edges: np.ndarray) -> np.ndarray: + diffs = np.diff(edges) + if len(diffs) == 1 or np.isclose( + diffs[1] - diffs[0], diffs[-1] - diffs[-2], atol=1e-2 + ): + return 0.5 * (edges[1:] + edges[:-1]) + else: + return (edges[1:] * edges[:-1]) ** 0.5 + + attributes = self.reader.ReadAttrsAtTimestep( + path=self.path, category="spectra", step=first_step + ) + if self.coordinate_system is None: + if "Coordinates" not in attributes: + raise ValueError("Coordinates not found in attributes for particles.") + if attributes["Coordinates"] in [b"cart", "cart"]: + self.set_coordinate_system(CoordinateSystem.XYZ) + elif attributes["Coordinates"] in [b"sph", "sph", b"qsph", "qsph"]: + self.set_coordinate_system(CoordinateSystem.SPH) + else: + raise NotImplementedError( + f"Coordinate system {attributes['Coordinates']} not supported." + ) + + if self.remap is None or self.remap.get("coords", None) is None: + self.set_remap( + { + "coords": ( + remap_coords_cart + if self.coordinate_system == CoordinateSystem.XYZ + else remap_coords_sph + ), + } + ) - diffs = np.diff(ebins) - if np.isclose(diffs[1] - diffs[0], diffs[-1] - diffs[-2], atol=1e-2): - ebins = 0.5 * (ebins[1:] + ebins[:-1]) + ebins = edges_to_bins(energy_binedges) + xbins = [edges_to_bins(xb) for xb in x_binedges] + + if self.remap is not None and "coords" in self.remap: + new_xbins = {} + for i in range(num_spatial_dims): + new_xbins[self.remap["coords"](f"X{i + 1}")] = xbins[i] + xbins = new_xbins else: - ebins = (ebins[1:] * ebins[:-1]) ** 0.5 + xbins = {f"X{i + 1}": xbins[i] for i in range(num_spatial_dims)} + + all_dims = { + "t": self.times, + **xbins, + "E": ebins, + } + all_coords = {**all_dims, "s": ("t", self.steps)} - all_dims = {**times, "E": ebins} - all_coords = {**all_dims, "s": ("t", steps["s"])} + attributes = self.reader.ReadAttrsAtTimestep( + path=self.path, category="spectra", step=first_step + ) def remap_name(name: str) -> str: return name[1:] @@ -94,29 +172,43 @@ def remap_name(name: str) -> str: da.stack( [ da.from_delayed( - dask.delayed(self.read_spectrum)( - path=self.path, - reader=self.reader, + dask.delayed(self._read_spectrum)( spectrum=spectrum, step=step, ), shape=shape, dtype="float", ) - for step in valid_steps + for step in tqdm( + self.valid_steps, + desc="steps", + position=1, + leave=False, + ) ], ), name=remap_name(spectrum), dims=all_dims, coords=all_coords, ) - for spectrum in spectra_names + for spectrum in tqdm( + spectra_names, + desc="spectra", + position=0, + leave=False, + ) }, - attrs=self.reader.ReadAttrsAtTimestep( - path=self.path, category="spectra", step=first_step - ), + attrs=attributes, ) + @property + def attrs(self) -> dict[str, Any]: + """dict: The attributes of the spectra dataframe.""" + if self.spectra_defined: + return self.spectra.attrs + else: + return {} + def help_spectra(self, prepend="") -> str: ret = f"{prepend}- use .sel(...) to select specific energy or time intervals\n" ret += f"{prepend} t : time (float)\n" diff --git a/nt2/plotters/annotations.py b/nt2/plotters/annotations.py index 4a768b4..1373ab1 100644 --- a/nt2/plotters/annotations.py +++ b/nt2/plotters/annotations.py @@ -2,11 +2,24 @@ def annotatePulsar( - ax, data, rmax, rstar=1.1, ti=None, time=None, attrs={}, ax_props={}, star_props={} + ax, + data, + rmax, + rstar=1.1, + ti=None, + time=None, + attrs=None, + ax_props=None, + star_props=None, ): import numpy as np - from matplotlib import lines - from matplotlib import patches + from matplotlib import lines, patches + + logger = logging.getLogger(__name__) + + attrs = attrs or {} + ax_props = ax_props or {} + star_props = star_props or {} if ti is None and time is None: raise ValueError("Must provide either ti or time") @@ -21,7 +34,7 @@ def annotatePulsar( ) ) ) is None: - logging.warning( + logger.warning( "No spinup time or spin period found, please specify explicitly as `attrs = {'psr_omega': ..., 'psr_spinup_time': ...}`" ) demo_rotation = False @@ -45,7 +58,7 @@ def annotatePulsar( xy=(0.0, rmax * 0.95), xytext=(0.0, -rmax * 0.95), zorder=4, - arrowprops=dict(arrowstyle="->", color=ax_props.get("color", "k"), lw=0.5), + arrowprops={"arrowstyle": "->", "color": ax_props.get("color", "k"), "lw": 0.5}, ) for i in range(-int(rmax * 0.8) // 2 - 1, int(rmax * 0.8) // 2): if i != -1: diff --git a/nt2/plotters/export.py b/nt2/plotters/export.py index eb7dafd..2c89712 100644 --- a/nt2/plotters/export.py +++ b/nt2/plotters/export.py @@ -1,11 +1,14 @@ -from typing import Any, Callable, Union, Optional, List +from __future__ import annotations + +from typing import Any, Callable + import matplotlib.pyplot as plt def makeFramesAndMovie( name: str, plot: Callable, - times: List[float], + times: list[float], data: Any = None, **kwargs: Any, ) -> bool: @@ -36,7 +39,7 @@ def makeFramesAndMovie( raise ValueError("Failed to make frames") -def makeMovie(**ffmpeg_kwargs: Union[str, int, float]) -> bool: +def makeMovie(**ffmpeg_kwargs: str | float) -> bool: """ Create a movie from frames using the `ffmpeg` command-line tool. @@ -62,9 +65,7 @@ def makeMovie(**ffmpeg_kwargs: Union[str, int, float]) -> bool: """ import subprocess - input_pattern: str = ( - f"{ffmpeg_kwargs.get('input', 'step_')}%0{ffmpeg_kwargs.get('number', 3)}d.{ffmpeg_kwargs.get('extension', 'png')}" - ) + input_pattern: str = f"{ffmpeg_kwargs.get('input', 'step_')}%0{ffmpeg_kwargs.get('number', 3)}d.{ffmpeg_kwargs.get('extension', 'png')}" command = [ ffmpeg_kwargs.get("ffmpeg", "ffmpeg"), @@ -86,7 +87,12 @@ def makeMovie(**ffmpeg_kwargs: Union[str, int, float]) -> bool: ] command = [str(c) for c in command if c is not None] print("Command:\n", " ".join(command)) - result = subprocess.run(command, capture_output=True, text=True) + result = subprocess.run( + command, + capture_output=True, + text=True, + check=False, + ) if result.returncode == 0: print("ffmpeg -- [OK]") @@ -112,11 +118,11 @@ def _plot_and_save(ti: int, t: float, fpath: str, plot: Callable, data: Any) -> def makeFrames( plot: Callable, - times: List[float], + times: list[float], fpath: str, data: Any = None, - num_cpus: Optional[int] = None, -) -> List[bool]: + num_cpus: int | None = None, +) -> list[bool]: """ Create plot frames from a set of timesteps of the same dataset. @@ -160,9 +166,10 @@ def makeFrames( >>> makeFrames(plot_func, range(100), 'output/', num_cpus=16) """ + import os + from loky import get_reusable_executor from tqdm import tqdm - import os os.makedirs(fpath, exist_ok=True) diff --git a/nt2/plotters/inspect.py b/nt2/plotters/inspect.py index 2e64ec5..c4b55d8 100644 --- a/nt2/plotters/inspect.py +++ b/nt2/plotters/inspect.py @@ -1,9 +1,13 @@ -from typing import Any, Callable, Optional, Union, List, Dict, Tuple -import matplotlib.pyplot as plt +from __future__ import annotations + +from typing import Any, Callable + import matplotlib.figure as mfigure +import matplotlib.pyplot as plt import xarray as xr -from nt2.utils import DataIs2DPolar + from nt2.plotters.export import makeFramesAndMovie +from nt2.utils import DataIs2DPolar class ds_accessor: @@ -12,7 +16,7 @@ def __init__(self, xarray_obj: xr.Dataset): def __axes_grid( self, - grouped_fields: Dict[str, List[str]], + grouped_fields: dict[str, list[str]], makeplot: Callable, nrows: int, ncols: int, @@ -20,7 +24,7 @@ def __axes_grid( size: float, aspect: float, **fig_kwargs: Any, - ) -> Tuple[mfigure.Figure, List[plt.Axes]]: + ) -> tuple[mfigure.Figure, list[plt.Axes]]: vpad = fig_kwargs.pop("vpad", 0.5) hpad = fig_kwargs.pop("hpad", 0.5) if aspect > 1: @@ -51,7 +55,7 @@ def __axes_grid( @staticmethod def _fixed_axes_grid_with_cbars( - fields: List[str], + fields: list[str], makeplot: Callable, makecbar: Callable, nrows: int, @@ -61,7 +65,7 @@ def _fixed_axes_grid_with_cbars( aspect: float, cbar_w: float, **fig_kwargs: Any, - ) -> Tuple[mfigure.Figure, List[plt.Axes]]: + ) -> tuple[mfigure.Figure, list[plt.Axes]]: from mpl_toolkits.axes_grid1 import Divider, Size vpad = fig_kwargs.pop("vpad", 0.5) @@ -89,7 +93,7 @@ def _fixed_axes_grid_with_cbars( v += [Size.Fixed(vpad)] divider = Divider(fig, (0, 0, 1, 1), h, v, aspect=False) - axes: List[plt.Axes] = [] + axes: list[plt.Axes] = [] cntr = 0 for i in range(nrows): @@ -117,15 +121,15 @@ def _fixed_axes_grid_with_cbars( def plot( self, - fig: Optional[mfigure.Figure] = None, - name: Optional[str] = None, - skip_fields: Optional[List[str]] = None, - only_fields: Optional[List[str]] = None, - fig_kwargs: Optional[Dict[str, Any]] = None, - plot_kwargs: Optional[Dict[str, Any]] = None, - movie_kwargs: Optional[Dict[str, Any]] = None, - set_aspect: Optional[str] = "equal", - ) -> Union[mfigure.Figure, bool]: + fig: mfigure.Figure | None = None, + name: str | None = None, + skip_fields: list[str] | None = None, + only_fields: list[str] | None = None, + fig_kwargs: dict[str, Any] | None = None, + plot_kwargs: dict[str, Any] | None = None, + movie_kwargs: dict[str, Any] | None = None, + set_aspect: str | None = "equal", + ) -> mfigure.Figure | bool: """ Plots the overview plot for fields at a given time or step (or as a movie). @@ -239,20 +243,20 @@ def plot_func(ti: int, _): @staticmethod def _get_fields_to_plot( - data: xr.Dataset, skip_fields: List[str], only_fields: List[str] - ) -> List[str]: + data: xr.Dataset, skip_fields: list[str], only_fields: list[str] + ) -> list[str]: import re nfields = len(data.data_vars) if nfields > 0: - keys: List[str] = [str(k) for k in data.keys()] + keys: list[str] = [str(k) for k in data] if len(only_fields) == 0: fields_to_plot = [ - f for f in keys if not any([re.match(sf, f) for sf in skip_fields]) + f for f in keys if not any(re.match(sf, f) for sf in skip_fields) ] else: fields_to_plot = [ - f for f in keys if any([re.match(sf, f) for sf in only_fields]) + f for f in keys if any(re.match(sf, f) for sf in only_fields) ] else: fields_to_plot = [] @@ -265,9 +269,9 @@ def _get_fields_to_plot( @staticmethod def _get_fields_minmax( - data: xr.Dataset, fields: List[str] - ) -> Dict[str, Optional[Tuple[float, float]]]: - minmax: Dict[str, Optional[Tuple[float, float]]] = { + data: xr.Dataset, fields: list[str] + ) -> dict[str, tuple[float, float] | None]: + minmax: dict[str, tuple[float, float] | None] = { "E": None, "B": None, "J": None, @@ -304,22 +308,23 @@ def _get_fields_minmax( def plot_frame_1d( self, data: xr.Dataset, - fig: Optional[mfigure.Figure], - skip_fields: List[str], - only_fields: List[str], - fig_kwargs: Dict[str, Any], - plot_kwargs: Dict[str, Any], + fig: mfigure.Figure | None, + skip_fields: list[str], + only_fields: list[str], + fig_kwargs: dict[str, Any], + plot_kwargs: dict[str, Any], ) -> mfigure.Figure: if len(data.dims) != 1: raise ValueError("Pass 1D data; use .sel or .isel to reduce dimension.") - import math, re + import math + import re # count the number of subplots fields_to_plot = self._get_fields_to_plot(data, skip_fields, only_fields) # group fields by their first letter - grouped_fields: Dict[str, List[str]] = {} + grouped_fields: dict[str, list[str]] = {} for f in fields_to_plot: key = f[0] if key not in grouped_fields: @@ -329,8 +334,8 @@ def plot_frame_1d( nplots = len(grouped_fields) aspect = 0.5 - ncols = max(1, int(math.floor(nplots * 1.5 * aspect / (1 + 1.5 * aspect)))) - nrows = max(1, int(math.ceil(nplots / ncols))) + ncols = max(1, math.floor(nplots * 1.5 * aspect / (1 + 1.5 * aspect))) + nrows = max(1, math.ceil(nplots / ncols)) figsize0 = 3.0 @@ -338,9 +343,9 @@ def plot_frame_1d( kwargs = {} for fld in fields_to_plot: kwargs[fld] = {} - for fld_kwargs in plot_kwargs: + for fld_kwargs, val in plot_kwargs.items(): if re.match(fld_kwargs, fld): - kwargs[fld] = {**plot_kwargs[fld_kwargs]} + kwargs[fld] = {**val} break def make_plot(ax: plt.Axes, fld: str): @@ -373,21 +378,23 @@ def make_plot(ax: plt.Axes, fld: str): def plot_frame_2d( self, data: xr.Dataset, - fig: Optional[mfigure.Figure], - skip_fields: List[str], - only_fields: List[str], - fig_kwargs: Dict[str, Any], - plot_kwargs: Dict[str, Any], - set_aspect: Optional[str], + fig: mfigure.Figure | None, + skip_fields: list[str], + only_fields: list[str], + fig_kwargs: dict[str, Any], + plot_kwargs: dict[str, Any], + set_aspect: str | None, ) -> mfigure.Figure: if len(data.dims) != 2: raise ValueError("Pass 2D data; use .sel or .isel to reduce dimension.") x1, x2 = data.dims + import math + import re + import matplotlib.colors as mcolors import numpy as np - import math, re # count the number of subplots fields_to_plot = self._get_fields_to_plot(data, skip_fields, only_fields) @@ -402,8 +409,8 @@ def plot_frame_2d( else: aspect = 1.5 - ncols = max(1, int(math.floor(nfields * 1.5 * aspect / (1 + 1.5 * aspect)))) - nrows = max(1, int(math.ceil(nfields / ncols))) + ncols = max(1, math.floor(nfields * 1.5 * aspect / (1 + 1.5 * aspect))) + nrows = max(1, math.ceil(nfields / ncols)) figsize0 = 3.0 @@ -443,9 +450,9 @@ def plot_frame_2d( "vmax": vmax, } kwargs[fld] = default_kwargs - for fld_kwargs in plot_kwargs: + for fld_kwargs, val in plot_kwargs.items(): if re.match(fld_kwargs, fld): - kwargs[fld] = {**default_kwargs, **plot_kwargs[fld_kwargs]} + kwargs[fld] = {**default_kwargs, **val} break if "norm" in kwargs[fld]: vmin = kwargs[fld].pop("vmin") diff --git a/nt2/plotters/movie.py b/nt2/plotters/movie.py index 9481c4e..6fb0b2c 100644 --- a/nt2/plotters/movie.py +++ b/nt2/plotters/movie.py @@ -1,8 +1,12 @@ -from typing import Any, Optional, Dict +from __future__ import annotations + +from typing import Any + +import xarray as xr + from nt2.plotters.export import ( makeFramesAndMovie, ) -import xarray as xr class accessor: @@ -21,8 +25,8 @@ def __init__(self, xarray_obj: xr.DataArray) -> None: def plot( self, name: str, - movie_kwargs: Dict[str, Any] = {}, - fig_kwargs: Dict[str, Any] = {}, + movie_kwargs: dict[str, Any] | None = None, + fig_kwargs: dict[str, Any] | None = None, aspect_equal: bool = False, **kwargs: Any, ) -> bool: @@ -58,12 +62,16 @@ def plot( import matplotlib.pyplot as plt + movie_kwargs = movie_kwargs or {} + fig_kwargs = fig_kwargs or {} + def plot_func(ti: int, _: Any) -> None: if len(self._obj.isel(t=ti).dims) == 2: if aspect_equal: x1, x2 = self._obj.isel(t=ti).dims - nx1, nx2 = len(self._obj.isel(t=ti)[x1]), len( - self._obj.isel(t=ti)[x2] + nx1, nx2 = ( + len(self._obj.isel(t=ti)[x1]), + len(self._obj.isel(t=ti)[x2]), ) aspect = nx1 / nx2 figsize = fig_kwargs.get("figsize", (6, 4 * aspect)) @@ -75,7 +83,7 @@ def plot_func(ti: int, _: Any) -> None: plt.gca().set_aspect("equal") plt.tight_layout() - num_cpus: Optional[int] = movie_kwargs.pop("num_cpus", None) + num_cpus: int | None = movie_kwargs.pop("num_cpus", None) return makeFramesAndMovie( name=name, data=self._obj, diff --git a/nt2/plotters/particles.py b/nt2/plotters/particles.py index a001e9a..f9557dd 100644 --- a/nt2/plotters/particles.py +++ b/nt2/plotters/particles.py @@ -1,7 +1,8 @@ -import xarray as xr +from __future__ import annotations + import numpy as np import numpy.typing as npt -from typing import Optional, Tuple +import xarray as xr class ds_accessor: @@ -12,10 +13,10 @@ def phaseplot( self, x: str = "x", y: str = "ux", - xbins: Optional[npt.NDArray] = None, - ybins: Optional[npt.NDArray] = None, - xlims: Optional[Tuple[float, float]] = None, - ylims: Optional[Tuple[float, float]] = None, + xbins: npt.NDArray | None = None, + ybins: npt.NDArray | None = None, + xlims: tuple[float, float] | None = None, + ylims: tuple[float, float] | None = None, xnbins: int = 100, ynbins: int = 100, **kwargs, @@ -57,12 +58,12 @@ def phaseplot( -------- >>> ds.phaseplot(x='x', y='ux', xbins=np.linspace(0, 1000, 100), ybins=np.linspace(-5, 5, 50)) """ - assert x in list(self._obj.keys()) and y in list( - self._obj.keys() - ), "x and y must be valid variable names in the dataset" - assert ( - len(self._obj[x].dims) == 1 and len(self._obj[y].dims) == 1 - ), "x and y must be 1D variables" + assert x in list(self._obj.keys()) and y in list(self._obj.keys()), ( + "x and y must be valid variable names in the dataset" + ) + assert len(self._obj[x].dims) == 1 and len(self._obj[y].dims) == 1, ( + "x and y must be 1D variables" + ) assert "t" not in self._obj.dims, "Dataset must not have time dimension" import matplotlib.pyplot as plt diff --git a/nt2/plotters/polar.py b/nt2/plotters/polar.py index a9dfbeb..8488cd0 100644 --- a/nt2/plotters/polar.py +++ b/nt2/plotters/polar.py @@ -1,5 +1,8 @@ +from __future__ import annotations + +from typing import Any + import numpy as np -from typing import Any, Dict from nt2.utils import DataIs2DPolar @@ -176,7 +179,7 @@ def fieldlines(self, fr, fth, start_points, **kwargs): fxs = self._obj[fr] * np.sin(ths) + self._obj[fth] * np.cos(ths) fys = self._obj[fr] * np.cos(ths) - self._obj[fth] * np.sin(ths) - props: Dict[str, Any] = { + props: dict[str, Any] = { "method": "nearest", "bounds_error": False, "fill_value": 0, @@ -188,9 +191,10 @@ def fieldlines(self, fr, fth, start_points, **kwargs): ] def _fieldline(self, interp_fx, interp_fy, r_th_start, **kwargs): - import numpy as np from copy import copy + import numpy as np + direction = kwargs.pop("direction", "both") stopWhen = kwargs.pop("stopWhen", lambda _, __: False) ds = kwargs.pop("ds", 0.1) @@ -292,10 +296,9 @@ def pcolor(self, **kwargs) -> Any: Additional keyword arguments are passed to `pcolormesh`. """ - import matplotlib.pyplot as plt - from matplotlib import colors - from matplotlib import tri import matplotlib as mpl + import matplotlib.pyplot as plt + from matplotlib import colors, tri from mpl_toolkits.axes_grid1 import make_axes_locatable ax = kwargs.pop("ax", plt.gca()) @@ -434,6 +437,7 @@ def contour(self, **kwargs): """ import warnings + import matplotlib.pyplot as plt ax = kwargs.pop("ax", plt.gca()) diff --git a/nt2/readers/adios2.py b/nt2/readers/adios2.py index 3b22ccf..5086793 100644 --- a/nt2/readers/adios2.py +++ b/nt2/readers/adios2.py @@ -1,6 +1,9 @@ -from typing import Any, List, Dict, Tuple, Set +from __future__ import annotations import sys +from typing import Any + +from tqdm import tqdm if sys.version_info >= (3, 12): from typing import override @@ -10,15 +13,15 @@ def override(method): return method -import re import os -import numpy as np -import numpy.typing as npt +import re import adios2 as bp +import numpy as np +import numpy.typing as npt -from nt2.utils import Format, Layout from nt2.readers.base import BaseReader +from nt2.utils import Format, Layout class Reader(BaseReader): @@ -41,15 +44,18 @@ def ReadPerTimestepVariable( category: str, varname: str, newname: str, - ) -> Dict[str, npt.NDArray[Any]]: - variables: List[float] = [] - for filename in self.GetValidFiles( - path=path, - category=category, + valid_files: list[str], + ) -> dict[str, npt.NDArray[Any]]: + variables: list[float] = [] + for filename in tqdm( + valid_files, + desc=f"Reading {category} {varname}", + position=0, + leave=False, ): with bp.FileReader(os.path.join(path, category, filename)) as f: - avail: Dict[str, Any] = f.available_variables() - vars: List[str] = list(avail.keys()) + avail: dict[str, Any] = f.available_variables() + vars: list[str] = list(avail.keys()) if varname in vars: var = f.inquire_variable(varname) if var is not None: @@ -62,16 +68,65 @@ def ReadPerTimestepVariable( raise ValueError(f"{varname} not found in the BP file {filename}") return {newname: np.array(variables)} + @override + def ReadPerTimestepVariables( + self, + path: str, + category: str, + varnames: list[str], + newnames: list[str], + valid_files: list[str], + ) -> dict[str, npt.NDArray[Any]]: + variables = {newname: [] for newname in newnames} + for filename in tqdm( + valid_files, + desc=f"Reading {category} {varnames}", + position=0, + leave=False, + ): + with bp.FileReader(os.path.join(path, category, filename)) as f: + avail: dict[str, Any] = f.available_variables() + vars: list[str] = list(avail.keys()) + for varname, newname in zip(varnames, newnames): + if varname in vars: + var = f.inquire_variable(varname) + if var is not None: + variables[newname].append(f.read(var)) + else: + raise ValueError( + f"{varname} is not a variable in the BP file {filename}" + ) + else: + raise ValueError( + f"{varname} not found in the BP file {filename}" + ) + return {newname: np.array(variables[newname]) for newname in newnames} + + @override + def ReadParticleCountsAtTimestep( + self, path: str, step: int, species: list[int] + ) -> dict[int, int]: + """Read all per-species counts from one BP file's metadata.""" + with bp.FileReader(self.FullPath(path, "particles", step)) as f: + available = f.available_variables() + counts: dict[int, int] = {} + for sp in species: + name = f"pX1_{sp}" + var = f.inquire_variable(name) if name in available else None + shape = var.shape() if var is not None else [] + counts[sp] = int(shape[0]) if shape else 0 + return counts + @override def ReadEdgeCoordsAtTimestep( self, path: str, step: int, - ) -> Dict[str, Any]: - dct: Dict[str, npt.NDArray[Any]] = {} + ) -> dict[str, Any]: + dct: dict[str, npt.NDArray[Any]] = {} with bp.FileReader(self.FullPath(path, "fields", step)) as f: - avail: Dict[str, Any] = f.available_variables() - vars: List[str] = list(avail.keys()) + avail: dict[str, Any] = f.available_variables() + vars: list[str] = list(avail.keys()) for var in vars: if var.startswith("X") and var.endswith("e"): var_obj = f.inquire_variable(var) @@ -85,7 +140,7 @@ def ReadAttrsAtTimestep( path: str, category: str, step: int, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: with bp.FileReader(self.FullPath(path, category, step)) as f: return {k: f.read_attribute(k) for k in f.available_attributes()} @@ -117,9 +172,9 @@ def ReadCategoryNamesAtTimestep( category: str, prefix: str, step: int, - ) -> Set[str]: + ) -> set[str]: with bp.FileReader(self.FullPath(path, category, step)) as f: - keys: List[str] = f.available_variables() + keys: list[str] = f.available_variables() return set( filter( lambda c: c.startswith(prefix), @@ -130,7 +185,7 @@ def ReadCategoryNamesAtTimestep( @override def ReadArrayShapeAtTimestep( self, path: str, category: str, quantity: str, step: int - ) -> Tuple[int, ...]: + ) -> tuple[int, ...]: with bp.FileReader(filename := self.FullPath(path, category, step)) as f: if quantity in f.available_variables(): var = f.inquire_variable(quantity) @@ -148,7 +203,7 @@ def ReadArrayShapeAtTimestep( @override def ReadArrayShapeExplicitlyAtTimestep( self, path: str, category: str, quantity: str, step: int - ) -> Tuple[int, ...]: + ) -> tuple[int, ...]: with bp.FileReader(filename := self.FullPath(path, category, step)) as f: if quantity in f.available_variables(): var = f.inquire_variable(quantity) @@ -166,7 +221,7 @@ def ReadArrayShapeExplicitlyAtTimestep( @override def ReadFieldCoordsAtTimestep( self, path: str, step: int - ) -> Dict[str, npt.NDArray[Any]]: + ) -> dict[str, npt.NDArray[Any]]: with bp.FileReader(filename := self.FullPath(path, "fields", step)) as f: def get_coord(c: str) -> npt.NDArray[Any]: @@ -176,13 +231,13 @@ def get_coord(c: str) -> npt.NDArray[Any]: else: raise ValueError(f"Field {c} is not a group in the {filename}") - keys: List[str] = list(f.available_variables()) + keys: list[str] = list(f.available_variables()) return {c: get_coord(c) for c in keys if re.match(r"^X[1|2|3]$", c)} @override def ReadFieldLayoutAtTimestep(self, path: str, step: int) -> Layout: with bp.FileReader(filename := self.FullPath(path, "fields", step)) as f: - attrs: Dict[str, Any] = f.available_attributes() + attrs: dict[str, Any] = f.available_attributes() keys = list(attrs.keys()) if "LayoutRight" not in keys: raise ValueError(f"LayoutRight attribute not found in the {filename}") diff --git a/nt2/readers/base.py b/nt2/readers/base.py index d53280d..501e0a7 100644 --- a/nt2/readers/base.py +++ b/nt2/readers/base.py @@ -1,10 +1,69 @@ -from typing import Any, List, Tuple, Dict, Set +from __future__ import annotations + +import logging +import os +import re +from concurrent.futures import as_completed +from typing import Any, Callable + import numpy.typing as npt -import os, re, logging +from loky import get_reusable_executor +from tqdm import tqdm from nt2.utils import Format, Layout +def _check_file(enter_file: Callable[[str], Any], filename: str) -> None: + with enter_file(filename): + pass + + +def _get_field_shapes( + reader: BaseReader, path: str, step: int +) -> dict[str, tuple[int, ...]]: + names = reader.ReadCategoryNamesAtTimestep( + path=path, + category="fields", + prefix="f", + step=step, + ) + return { + name: reader.ReadArrayShapeAtTimestep( + path=path, + category="fields", + quantity=name, + step=step, + ) + for name in names + } + + +def _verify_particle_shapes(reader: BaseReader, path: str, step: int) -> None: + prtl_species = reader.ReadParticleSpeciesAtTimestep(path=path, step=step) + quantities = reader.ReadCategoryNamesAtTimestep( + path=path, + category="particles", + prefix="p", + step=step, + ) + quantities = {q.split("_")[0] for q in quantities if q.startswith("p")} + for species in prtl_species: + shape = None + for quantity in quantities: + current_shape = reader.ReadArrayShapeAtTimestep( + path=path, + category="particles", + quantity=f"{quantity}_{species}", + step=step, + ) + if shape is None: + shape = current_shape + elif shape != current_shape: + raise ValueError( + f"Different particle shapes found in the {reader.format.value} files for species {species} and quantity {quantity} in step {step}" + ) + + class BaseReader: """Base virtual class for arbitrary format readers. @@ -12,7 +71,7 @@ class BaseReader: """ - skipped_files: List[str] + skipped_files: list[str] def __init__(self) -> None: """Initializer for the BaseReader class.""" @@ -52,7 +111,8 @@ def ReadPerTimestepVariable( category: str, varname: str, newname: str, - ) -> Dict[str, npt.NDArray[Any]]: + valid_files: list[str], + ) -> dict[str, npt.NDArray[Any]]: """Read a variable at each timestep and return a dictionary with the new name. Parameters @@ -65,6 +125,8 @@ def ReadPerTimestepVariable( The name of the variable to be read. newname : str The new name of the variable to be returned. + valid_files : list[str] + The valid files to be read. Returns ------- @@ -74,12 +136,70 @@ def ReadPerTimestepVariable( """ raise NotImplementedError("ReadPerTimestepVariable is not implemented") + def ReadPerTimestepVariables( + self, + path: str, + category: str, + varnames: list[str], + newnames: list[str], + valid_files: list[str], + ) -> dict[str, npt.NDArray[Any]]: + """Read multiple variables at each timestep and return a dictionary with the new names. + + Parameters + ---------- + path : str + The path to the files. + category : str + The category of the files. + varnames : list[str] + The names of the variables to be read. + newnames : list[str] + The new names of the variables to be returned. + valid_files : list[str] + The valid files to be read. + + Returns + ------- + dict[str, NDArray[Any]] + A dictionary with the new names and the variables at each timestep. + + """ + raise NotImplementedError("ReadPerTimestepVariables is not implemented") + + def ReadParticleCountsAtTimestep( + self, + path: str, + step: int, + species: list[int], + ) -> dict[int, int]: + """Return particle counts by species without reading particle arrays. + + Readers may override this method to collect all counts while opening the + timestep only once. The default implementation uses array-shape + metadata and is kept for third-party readers. + """ + counts: dict[int, int] = {} + for sp in species: + try: + counts[sp] = int( + self.ReadArrayShapeAtTimestep( + path=path, + category="particles", + quantity=f"pX1_{sp}", + step=step, + )[0] + ) + except (IndexError, KeyError, OSError, ValueError): + counts[sp] = 0 + return counts + def ReadAttrsAtTimestep( self, path: str, category: str, step: int, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: """Read the attributes of a given timestep. Parameters @@ -103,7 +223,7 @@ def ReadEdgeCoordsAtTimestep( self, path: str, step: int, - ) -> Dict[str, npt.NDArray[Any]]: + ) -> dict[str, npt.NDArray[Any]]: """Read the coordinates of cell edges at a given timestep. Parameters @@ -155,7 +275,7 @@ def ReadCategoryNamesAtTimestep( category: str, prefix: str, step: int, - ) -> Set[str]: + ) -> set[str]: """Read the names of the variables in a given category and timestep. Parameters @@ -177,7 +297,7 @@ def ReadCategoryNamesAtTimestep( """ raise NotImplementedError("ReadCategoryNamesAtTimestep is not implemented") - def ReadParticleSpeciesAtTimestep(self, path: str, step: int) -> Set[int]: + def ReadParticleSpeciesAtTimestep(self, path: str, step: int) -> set[int]: """Read the particle species indices at a given timestep. Parameters @@ -193,10 +313,10 @@ def ReadParticleSpeciesAtTimestep(self, path: str, step: int) -> Set[int]: A set of particle species indices at a given timestep. """ - return set( + return { int(f.split("_")[1]) for f in self.ReadCategoryNamesAtTimestep(path, "particles", "p", step) - ) + } def ReadArrayShapeAtTimestep( self, @@ -204,7 +324,7 @@ def ReadArrayShapeAtTimestep( category: str, quantity: str, step: int, - ) -> Tuple[int, ...]: + ) -> tuple[int, ...]: """Read the shape of an array at a given timestep. Parameters @@ -232,7 +352,7 @@ def ReadArrayShapeExplicitlyAtTimestep( category: str, quantity: str, step: int, - ) -> Tuple[int, ...]: + ) -> tuple[int, ...]: """Read the shape of an array at a given timestep, without relying on metadata. Parameters @@ -260,7 +380,7 @@ def ReadFieldCoordsAtTimestep( self, path: str, step: int, - ) -> Dict[str, npt.NDArray[Any]]: + ) -> dict[str, npt.NDArray[Any]]: """Read the coordinates of the fields at a given timestep. Parameters @@ -301,7 +421,7 @@ def ReadFieldLayoutAtTimestep(self, path: str, step: int) -> Layout: # # # # # # # # # # # # # # # # # # # # # # # # @staticmethod - def CategoryFiles(path: str, category: str, format: str) -> List[str]: + def CategoryFiles(path: str, category: str, format: str) -> list[str]: """Get the list of files in a given category and format. Parameters @@ -324,6 +444,8 @@ def CategoryFiles(path: str, category: str, format: str) -> List[str]: If no files are found. """ + if not os.path.exists(os.path.join(path, category)): + return [] files = [ f for f in os.listdir(os.path.join(path, category)) @@ -356,51 +478,14 @@ def FullPath(self, path: str, category: str, step: int) -> str: path, category, f"{category}.{step:08d}.{self.format.value}" ) - def GetValidSteps( - self, - path: str, - category: str, - ) -> List[int]: - """Get valid timesteps (sorted) in a given path and category. - - Parameters - ---------- - path : str - The path to the files. - category : str - The category of the files. - - Returns - ------- - list[int] - A list of valid timesteps in the given path and category. - - """ - steps: List[int] = [] - for filename in BaseReader.CategoryFiles( - path=path, - category=category, - format=self.format.value, - ): - try: - with self.EnterFile(os.path.join(path, category, filename)): - step = int(filename.split(".")[1]) - steps.append(step) - except OSError: - if filename not in self.skipped_files: - self.skipped_files.append(filename) - logging.warning(f"Could not read {filename}, skipping it") - except Exception as e: - raise e - steps.sort() - return steps - - def GetValidFiles( + def GetValidFilesAndSteps( self, path: str, category: str, - ) -> List[str]: - """Get valid files (sorted by timestep) in a given path and category. + steprange: tuple[int | None, int | None] | None = None, + num_cpus: int | None = None, + ) -> tuple[list[str], list[int]]: + """Get valid files (sorted by timestep) and steps in a given path and category. Parameters ---------- @@ -408,36 +493,75 @@ def GetValidFiles( The path to the files. category : str The category of the files. + steprange : tuple[int | None, int | None] | None + The range of timesteps to be considered. If None, all timesteps are considered. + num_cpus : int | None + The number of CPU cores to use for parallel processing. Returns ------- - list[str] - A list of valid files in the given path and category. + tuple[list[str], list[int]] + A tuple containing a list of valid files and a list of valid timesteps in the given path and category. """ - files: List[str] = [] - for filename in BaseReader.CategoryFiles( + category_files = BaseReader.CategoryFiles( path=path, category=category, format=self.format.value, + ) + num_cpus = num_cpus if num_cpus is not None else (os.cpu_count() or 1) + executor = get_reusable_executor(max_workers=num_cpus) + + def is_inrange(filename: str) -> bool: + step = int(filename.split(".")[1]) + if steprange is None: + return True + start, end = steprange + if start is not None and step < start: + return False + return not (end is not None and step >= end) + + futures = { + executor.submit( + _check_file, + self.EnterFile, + os.path.join(path, category, filename), + ): filename + for filename in category_files + if is_inrange(filename) + } + logger = logging.getLogger(__name__) + + files: list[str] = [] + steps: list[int] = [] + for future in tqdm( + as_completed(futures), + total=len(futures), + desc=f"getting valid files & steps for {category}", + leave=False, ): + filename = futures[future] try: - with self.EnterFile(os.path.join(path, category, filename)): - files.append(filename) + future.result() + files.append(filename) + steps.append(int(filename.split(".")[1])) except OSError: if filename not in self.skipped_files: self.skipped_files.append(filename) - logging.warning(f"Could not read {filename}, skipping it") - except Exception as e: - raise e + logger.warning(f"Could not read {filename}, skipping it") + except Exception: + raise files.sort(key=lambda x: int(x.split(".")[1])) - return files + steps.sort() + return (files, steps) def VerifySameCategoryNames( self, path: str, category: str, prefix: str, + valid_steps: list[int], + num_cpus: int | None = None, ): """Verify that all files in a given category have the same names. @@ -449,6 +573,10 @@ def VerifySameCategoryNames( The category of the files. prefix : str The prefix of the variables to be read. + valid_steps : list[int] + The valid timesteps to be checked. + num_cpus : int | None + The number of CPU cores to use for parallel processing. Raises ------ @@ -456,32 +584,41 @@ def VerifySameCategoryNames( If different names are found. """ - names = None - for step in self.GetValidSteps( - path=path, - category=category, + num_cpus = num_cpus if num_cpus is not None else (os.cpu_count() or 1) + executor = get_reusable_executor(max_workers=num_cpus) + futures = { + executor.submit( + self.ReadCategoryNamesAtTimestep, + path=path, + category=category, + prefix=prefix, + step=step, + ): step + for step in valid_steps + } + names_by_step = {} + for future in tqdm( + as_completed(futures), + total=len(futures), + desc=f"verifying same names for {category}", + leave=False, ): + names_by_step[futures[future]] = future.result() + + names = None + for step in valid_steps: if names is None: - names = self.ReadCategoryNamesAtTimestep( - path=path, - category=category, - prefix=prefix, - step=step, + names = names_by_step[step] + elif names != names_by_step[step]: + raise ValueError( + f"Different field names found in the {self.format.value} files for step {step}" ) - else: - if names != self.ReadCategoryNamesAtTimestep( - path=path, - category=category, - prefix=prefix, - step=step, - ): - raise ValueError( - f"Different field names found in the {self.format.value} files for step {step}" - ) def VerifySameFieldShapes( self, path: str, + valid_steps: list[int], + num_cpus: int | None = None, ): """Verify that all fields in a given path have the same shape. @@ -489,6 +626,10 @@ def VerifySameFieldShapes( ---------- path : str The path to the files. + valid_steps : list[int] + The valid timesteps to be checked. + num_cpus : int | None + The number of CPU cores to use for parallel processing. Raises ------ @@ -496,43 +637,49 @@ def VerifySameFieldShapes( If different shapes are found. """ - shape = None - for step in self.GetValidSteps( - path=path, - category="fields", + num_cpus = num_cpus if num_cpus is not None else (os.cpu_count() or 1) + executor = get_reusable_executor(max_workers=num_cpus) + futures = { + executor.submit(_get_field_shapes, self, path, step): step + for step in valid_steps + } + shapes_by_step = {} + for future in tqdm( + as_completed(futures), + total=len(futures), + desc="verifying same shapes for fields", + leave=False, ): - names = self.ReadCategoryNamesAtTimestep( - path=path, - category="fields", - prefix="f", - step=step, - ) + shapes_by_step[futures[future]] = future.result() + + shape = None + for step in valid_steps: + names = set(shapes_by_step[step]) if shape is None: name = names.pop() - shape = self.ReadArrayShapeAtTimestep( - path=path, - category="fields", - quantity=name, - step=step, - ) + shape = shapes_by_step[step][name] for name in names: - if shape != self.ReadArrayShapeAtTimestep( - path=path, - category="fields", - quantity=name, - step=step, - ): + if shape != shapes_by_step[step][name]: raise ValueError( f"Different field shapes found in the {self.format.value} files for field {name} in step {step}" ) - def VerifySameFieldLayouts(self, path: str): + def VerifySameFieldLayouts( + self, + path: str, + valid_steps: list[int], + num_cpus: int | None = None, + ): """Verify that all timesteps in a given path have the same layout. Parameters ---------- path : str The path to the files. + valid_steps : list[int] + The valid timesteps to be checked. + num_cpus : int | None + The number of CPU cores to use for parallel processing. Raises ------ @@ -540,32 +687,50 @@ def VerifySameFieldLayouts(self, path: str): If different layouts are found. """ - layout = None - for step in self.GetValidSteps( - path=path, - category="fields", + num_cpus = num_cpus if num_cpus is not None else (os.cpu_count() or 1) + executor = get_reusable_executor(max_workers=num_cpus) + futures = { + executor.submit( + self.ReadFieldLayoutAtTimestep, + path=path, + step=step, + ): step + for step in valid_steps + } + layouts_by_step = {} + for future in tqdm( + as_completed(futures), + total=len(futures), + desc="verifying same layouts for fields", + leave=False, ): + layouts_by_step[futures[future]] = future.result() + + layout = None + for step in valid_steps: if layout is None: - layout = self.ReadFieldLayoutAtTimestep( - path=path, - step=step, + layout = layouts_by_step[step] + elif layout != layouts_by_step[step]: + raise ValueError( + f"Different field layouts found in the {self.format.value} files for step {step}" ) - else: - if layout != self.ReadFieldLayoutAtTimestep( - path=path, - step=step, - ): - raise ValueError( - f"Different field layouts found in the {self.format.value} files for step {step}" - ) - def VerifySameParticleShapes(self, path: str): + def VerifySameParticleShapes( + self, + path: str, + valid_steps: list[int], + num_cpus: int | None = None, + ): """Verify that all particle quantities in a given path have the same shape at specific timesteps. Parameters ---------- path : str The path to the files. + valid_steps : list[int] + The valid timesteps to be checked. + num_cpus : int | None + The number of CPU cores to use for parallel processing. Raises ------ @@ -573,40 +738,21 @@ def VerifySameParticleShapes(self, path: str): If different shapes are found. """ - for step in self.GetValidSteps( - path=path, - category="particles", + num_cpus = num_cpus if num_cpus is not None else (os.cpu_count() or 1) + executor = get_reusable_executor(max_workers=num_cpus) + futures = [ + executor.submit(_verify_particle_shapes, self, path, step) + for step in valid_steps + ] + for future in tqdm( + as_completed(futures), + total=len(futures), + desc="verifying same shapes for particles", + leave=False, ): - prtl_species = self.ReadParticleSpeciesAtTimestep(path=path, step=step) - quantities = self.ReadCategoryNamesAtTimestep( - path=path, - category="particles", - prefix="p", - step=step, - ) - quantities = set(q.split("_")[0] for q in quantities if q.startswith("p")) - for sp in prtl_species: - shape = None - for q in quantities: - if shape is None: - shape = self.ReadArrayShapeAtTimestep( - path=path, - category="particles", - quantity=f"{q}_{sp}", - step=step, - ) - else: - if shape != self.ReadArrayShapeAtTimestep( - path=path, - category="particles", - quantity=f"{q}_{sp}", - step=step, - ): - raise ValueError( - f"Different particle shapes found in the {self.format.value} files for species {sp} and quantity {q} in step {step}" - ) - - def DefinesCategory(self, path: str, category: str) -> bool: + future.result() + + def DefinesCategory(self, path: str, category: str, valid_files: list[str]) -> bool: """Check whether a given category is defined in the path. Parameters @@ -622,6 +768,4 @@ def DefinesCategory(self, path: str, category: str) -> bool: True if the category is defined, False otherwise. """ - return os.path.exists(os.path.join(path, category)) and ( - len(self.GetValidFiles(path=path, category=category)) > 0 - ) + return os.path.exists(os.path.join(path, category)) and (len(valid_files) > 0) diff --git a/nt2/readers/hdf5.py b/nt2/readers/hdf5.py index 0ba198b..ae82e36 100644 --- a/nt2/readers/hdf5.py +++ b/nt2/readers/hdf5.py @@ -1,8 +1,9 @@ from __future__ import annotations -from typing import Any, TYPE_CHECKING, List, Dict, Tuple, Set - import sys +from typing import TYPE_CHECKING, Any + +from tqdm import tqdm if sys.version_info >= (3, 12): from typing import override @@ -12,8 +13,9 @@ def override(method): return method -import re import os +import re + import numpy as np import numpy.typing as npt @@ -25,8 +27,8 @@ def override(method): if TYPE_CHECKING: import h5py as _h5py -from nt2.utils import Format, Layout from nt2.readers.base import BaseReader +from nt2.utils import Format, Layout def _require_h5py(): @@ -40,16 +42,16 @@ def _require_h5py(): class Reader(BaseReader): @staticmethod - def __extract_step0(f: "_h5py.File") -> "_h5py.Group": + def __extract_step0(f: _h5py.File) -> _h5py.Group: h5 = _require_h5py() - if "Step0" in f.keys(): + if "Step0" in f: f0 = f["Step0"] if isinstance(f0, h5.Group): return f0 else: - raise ValueError(f"Step0 is not a group in the HDF5 file") + raise ValueError("Step0 is not a group in the HDF5 file") else: - raise ValueError(f"Wrong structure of the hdf5 file") + raise ValueError("Wrong structure of the hdf5 file") @property @override @@ -60,7 +62,7 @@ def format(self) -> Format: @override def EnterFile( filename: str, - ) -> "_h5py.File": + ) -> _h5py.File: h5 = _require_h5py() return h5.File(filename, "r") @@ -71,16 +73,19 @@ def ReadPerTimestepVariable( category: str, varname: str, newname: str, - ) -> Dict[str, npt.NDArray[Any]]: - variables: List[Any] = [] + valid_files: list[str], + ) -> dict[str, npt.NDArray[Any]]: + variables: list[Any] = [] h5 = _require_h5py() - for filename in self.GetValidFiles( - path=path, - category=category, + for filename in tqdm( + valid_files, + desc=f"Reading {category}/{varname}", + position=0, + leave=False, ): with h5.File(os.path.join(path, category, filename), "r") as f: f0 = Reader.__extract_step0(f) - if varname in f0.keys(): + if varname in f0: var = f0[varname] if isinstance(var, h5.Dataset): variables.append(var[()]) @@ -93,13 +98,63 @@ def ReadPerTimestepVariable( return {newname: np.array(variables)} + @override + def ReadPerTimestepVariables( + self, + path: str, + category: str, + varnames: list[str], + newnames: list[str], + valid_files: list[str], + ) -> dict[str, npt.NDArray[Any]]: + variables = {newname: [] for newname in newnames} + h5 = _require_h5py() + for filename in tqdm( + valid_files, + desc=f"Reading {category} {varnames}", + position=0, + leave=False, + ): + with h5.File(os.path.join(path, category, filename), "r") as f: + f0 = Reader.__extract_step0(f) + for varname, newname in zip(varnames, newnames): + if varname in f0: + var = f0[varname] + if isinstance(var, h5.Dataset): + variables[newname].append(var[()]) + else: + raise ValueError( + f"{varname} is not a group in the HDF5 file {filename}" + ) + else: + raise ValueError( + f"{varname} not found in the HDF5 file {filename}" + ) + + return {newname: np.array(variables[newname]) for newname in newnames} + + @override + def ReadParticleCountsAtTimestep( + self, path: str, step: int, species: list[int] + ) -> dict[int, int]: + """Read all per-species counts from one HDF5 file's metadata.""" + h5 = _require_h5py() + with h5.File(self.FullPath(path, "particles", step), "r") as f: + f0 = Reader.__extract_step0(f) + return { + sp: int(f0[name].shape[0]) + if (name := f"pX1_{sp}") in f0 and isinstance(f0[name], h5.Dataset) + else 0 + for sp in species + } + @override def ReadAttrsAtTimestep( self, path: str, category: str, step: int, - ) -> Dict[str, Any]: + ) -> dict[str, Any]: h5 = _require_h5py() with h5.File(self.FullPath(path, category, step), "r") as f: return {k: v for k, v in f.attrs.items()} @@ -109,7 +164,7 @@ def ReadEdgeCoordsAtTimestep( self, path: str, step: int, - ) -> Dict[str, npt.NDArray[Any]]: + ) -> dict[str, npt.NDArray[Any]]: h5 = _require_h5py() with h5.File(self.FullPath(path, "fields", step), "r") as f: f0 = Reader.__extract_step0(f) @@ -126,7 +181,7 @@ def ReadArrayAtTimestep( h5 = _require_h5py() with h5.File(filename := self.FullPath(path, category, step), "r") as f: f0 = Reader.__extract_step0(f) - if quantity in f0.keys(): + if quantity in f0: var = f0[quantity] if isinstance(var, h5.Dataset): return np.array(var[:]) @@ -142,21 +197,21 @@ def ReadCategoryNamesAtTimestep( category: str, prefix: str, step: int, - ) -> Set[str]: + ) -> set[str]: h5 = _require_h5py() with h5.File(self.FullPath(path, category, step), "r") as f: f0 = Reader.__extract_step0(f) - keys: List[str] = list(f0.keys()) - return set(c for c in keys if c.startswith(prefix)) + keys: list[str] = list(f0.keys()) + return {c for c in keys if c.startswith(prefix)} @override def ReadArrayShapeAtTimestep( self, path: str, category: str, quantity: str, step: int - ) -> Tuple[int, ...]: + ) -> tuple[int, ...]: h5 = _require_h5py() with h5.File(filename := self.FullPath(path, category, step), "r") as f: f0 = Reader.__extract_step0(f) - if quantity in f0.keys(): + if quantity in f0: var = f0[quantity] if isinstance(var, h5.Dataset): return var.shape @@ -172,11 +227,11 @@ def ReadArrayShapeAtTimestep( @override def ReadArrayShapeExplicitlyAtTimestep( self, path: str, category: str, quantity: str, step: int - ) -> Tuple[int, ...]: + ) -> tuple[int, ...]: h5 = _require_h5py() with h5.File(self.FullPath(path, category, step), "r") as f: f0 = Reader.__extract_step0(f) - if quantity in f0.keys(): + if quantity in f0: var = f0[quantity] if isinstance(var, h5.Dataset) and (read := var[:]) is not None: return read.shape @@ -192,7 +247,7 @@ def ReadArrayShapeExplicitlyAtTimestep( @override def ReadFieldCoordsAtTimestep( self, path: str, step: int - ) -> Dict[str, npt.NDArray[Any]]: + ) -> dict[str, npt.NDArray[Any]]: h5 = _require_h5py() with h5.File(filename := self.FullPath(path, "fields", step), "r") as f: f0 = Reader.__extract_step0(f) @@ -202,9 +257,9 @@ def get_coord(c: str) -> Any: if isinstance(f0_c, h5.Dataset): return f0_c[:] else: - raise ValueError(f"Field {c} is not a group in the {filename}") + raise TypeError(f"Field {c} is not a group in the {filename}") - keys: List[str] = list(f0.keys()) + keys: list[str] = list(f0.keys()) return {c: get_coord(c) for c in keys if re.match(r"^X[1|2|3]$", c)} @override diff --git a/nt2/tests/test_cli.py b/nt2/tests/test_cli.py index 731c2be..f43a9da 100644 --- a/nt2/tests/test_cli.py +++ b/nt2/tests/test_cli.py @@ -1,6 +1,9 @@ +from unittest.mock import Mock + +import matplotlib.pyplot as plt +import numpy as np import pytest from typer.testing import CliRunner -import matplotlib.pyplot as plt import os import nt2 @@ -13,77 +16,116 @@ def test_version(): result = runner.invoke(app, ["version"]) assert result.exit_code == 0, f"Expected exit code 0, got {result.exit_code}" - assert ( - nt2.__version__ in result.output - ), f"Expected version {nt2.__version__} in output, got {result.output}" + assert nt2.__version__ in result.output, ( + f"Expected version {nt2.__version__} in output, got {result.output}" + ) @pytest.mark.parametrize( "test", [test for test in TESTS], ) -def test_show(test): - PATH = test["path"] - data = nt2.Data(PATH) - result = runner.invoke(app, ["show", PATH]) +def test_show(test, monkeypatch): + path = test["path"] + expected = f"Data summary for {path}" + data = Mock() + data.to_str.return_value = expected + data_factory = Mock(return_value=data) + monkeypatch.setattr(nt2, "Data", data_factory) + + result = runner.invoke(app, ["show", path]) + assert result.exit_code == 0, f"Expected exit code 0, got {result.exit_code}" - assert ( - data.to_str() in result.output - ), f"Expected data info in output, got {result.output}" + data_factory.assert_called_once_with(path) + data.to_str.assert_called_once_with() + assert expected in result.output @pytest.mark.parametrize( "test", - [test for test in TESTS], + [test for test in TESTS if test["fields"]], ) -def test_plot_png(test): - PATH = test["path"] - if test["fields"] == {}: - return - if test.get("coords", "cart") == "cart": - result = runner.invoke( - app, - [ - "plot", - PATH, - "--what", - "fields", - "--sel", - "x=slice(None, 5);y=slice(-5.0, 5.0)", - "--isel", - f"t=0{';z=0' if test['dim'] == '3D' else ''}", - ], +def test_plot_png(test, monkeypatch, tmp_path): + path = test["path"] + data_class = nt2.Data + + def field_data(data_path): + return data_class( + data_path, + fields=True, + particles=False, + spectra=False, + num_cpus=1, ) - assert result.exit_code == 0, f"Expected exit code 0, got {result.exit_code}" - - data = nt2.Data(PATH) - fname = os.path.basename(PATH.strip("/")) - - d = data.fields.sel(x=slice(None, 5), y=slice(-5, 5)).isel(t=0) - if test["dim"] == "3D": - d = d.isel(z=0) - d.inspect.plot(fig_kwargs={"dpi": 200}) - plt.savefig(fname=f"{fname}-2.png") - - def files_are_identical(path1, path2): - with open(path1, "rb") as f1, open(path2, "rb") as f2: - return f1.read() == f2.read() - - assert files_are_identical( - f"{fname}-2.png", f"{fname}.png" - ), f"Files {fname}-2.png and {fname}.png are not identical." - - os.remove(f"{fname}-2.png") - os.remove(f"{fname}.png") - # else: - # result = runner.invoke( - # app, - # [ - # "plot", - # PATH, - # "--sel", - # "r=slice(None, 5);th=slice(1.5, 2.5)", - # "--isel", - # "t=0", - # ], - # ) + + monkeypatch.setattr(nt2, "Data", field_data) + monkeypatch.chdir(tmp_path) + + is_cartesian = test.get("coords", "cart") == "cart" + selection = ( + "x=slice(None, 5);y=slice(-5.0, 5.0)" + if is_cartesian + else "r=slice(None, 5);th=slice(1.5, 2.5)" + ) + index_selection = f"t=0{';z=0' if test['dim'] == '3D' else ''}" + result = runner.invoke( + app, + [ + "plot", + path, + "--what", + "fields", + "--sel", + selection, + "--isel", + index_selection, + ], + ) + assert result.exit_code == 0, f"Expected exit code 0, got {result.exit_code}" + + fname = os.path.basename(path.strip("/")) + actual_path = tmp_path / f"{fname}.png" + expected_path = tmp_path / f"{fname}-expected.png" + assert actual_path.is_file() + + data = field_data(path) + if is_cartesian: + selected = data.fields.sel(x=slice(None, 5), y=slice(-5, 5)).isel(t=0) + else: + selected = data.fields.sel(r=slice(None, 5), th=slice(1.5, 2.5)).isel(t=0) + if test["dim"] == "3D": + selected = selected.isel(z=0) + + plt.close("all") + selected.inspect.plot(name=fname, fig_kwargs={"dpi": 200}) + plt.savefig(expected_path) + plt.close("all") + + np.testing.assert_array_equal( + plt.imread(actual_path), + plt.imread(expected_path), + ) + + +@pytest.mark.parametrize("test", [test for test in TESTS if not test["fields"]]) +def test_plot_without_fields_fails(test, monkeypatch, tmp_path): + path = test["path"] + data_class = nt2.Data + + def field_data(data_path): + return data_class( + data_path, + fields=True, + particles=False, + spectra=False, + num_cpus=1, + ) + + monkeypatch.setattr(nt2, "Data", field_data) + monkeypatch.chdir(tmp_path) + + result = runner.invoke(app, ["plot", path, "--what", "fields", "--isel", "t=0"]) + + assert result.exit_code == 1 + assert isinstance(result.exception, ValueError) + assert "Fields are not defined" in str(result.exception) diff --git a/nt2/tests/test_containers.py b/nt2/tests/test_containers.py index 9e8f058..a770a40 100644 --- a/nt2/tests/test_containers.py +++ b/nt2/tests/test_containers.py @@ -1,13 +1,19 @@ +from typing import List, Type, Union + import pytest -from typing import Union, List, Type from nt2.readers.base import BaseReader -from nt2.containers.fields import Fields -from nt2.containers.particles import Particles +from nt2.containers.fields import FieldContainer +from nt2.containers.particles import ParticleContainer from nt2.containers.data import Data from nt2.tests.cases import TESTS +FIELD_TESTS = [test for test in TESTS if test["fields"]] +EMPTY_FIELD_TESTS = [test for test in TESTS if not test["fields"]] +PARTICLE_TESTS = [test for test in TESTS if test["particles"]] + + def check_shape(shape1, shape2): """ Check if two shapes are equal @@ -16,45 +22,69 @@ def check_shape(shape1, shape2): @pytest.mark.parametrize( - "test,field_container", [[test, fc] for test in TESTS for fc in [Data, Fields]] + "test,field_container", + [[test, container] for test in FIELD_TESTS for container in [Data, FieldContainer]], ) -def test_fields(test, field_container: Union[Type[Data], Type[Fields]]): +def test_fields(test, field_container: Union[Type[Data], Type[FieldContainer]]): reader: BaseReader = test["reader"]() - PATH = test["path"] - if test["fields"] == {}: - return + path = test["path"] coords: List[str] = ["x", "y", "z"] flds: List[str] = ["Ex", "Ey", "Ez", "Bx", "By", "Bz"] - def coord_remap(Xold: str) -> str: + def coord_remap_cart(Xold: str) -> str: return { "X1": "x", "X2": "y", "X3": "z", }.get(Xold, Xold) - if test.get("coords", "cart") != "cart": - coords = ["r", "th", "ph"] - flds = ["Er", "Eth", "Eph", "Br", "Bth", "Bph"] - coord_remap = lambda Xold: { + def coord_remap_sph(Xold: str) -> str: + return { "X1": "r", "X2": "th", "X3": "ph", }.get(Xold, Xold) + if test.get("coords", "cart") != "cart": + coords = ["r", "th", "ph"] + flds = ["Er", "Eth", "Eph", "Br", "Bth", "Bph"] + def field_remap(Fold: str): return { - f"f{F}{i+1}": f"{F}{x}" for i, x in enumerate(coords) for F in "EB" + f"f{F}{i + 1}": f"{F}{x}" for i, x in enumerate(coords) for F in "EB" }.get(Fold, Fold) - fields = field_container( - path=PATH, - reader=reader, - remap={"coords": coord_remap, "fields": field_remap}, - ) + remap = { + "coords": ( + coord_remap_cart + if test.get("coords", "cart") == "cart" + else coord_remap_sph + ), + "fields": field_remap, + } + if field_container is Data: + fields = field_container( + path=path, + fields=True, + particles=False, + spectra=False, + reader=reader, + remap=remap, + num_cpus=1, + ) + else: + fields = field_container( + path=path, + reader=reader, + remap=remap, + coord_system=None, + num_cpus=1, + ) - steps = reader.GetValidSteps(path=PATH, category="fields") + _, steps = reader.GetValidFilesAndSteps( + path=path, category="fields", num_cpus=1 + ) nx1 = test["fields"]["nx1"] nx2 = test["fields"]["nx2"] assert fields.fields is not None, "Fields are None" @@ -106,17 +136,56 @@ def field_remap(Fold: str): ) +@pytest.mark.parametrize( + "test,field_container", + [ + [test, container] + for test in EMPTY_FIELD_TESTS + for container in [Data, FieldContainer] + ], +) +def test_missing_field_output_is_undefined( + test, field_container: Union[Type[Data], Type[FieldContainer]] +): + reader: BaseReader = test["reader"]() + kwargs = { + "path": test["path"], + "reader": reader, + "remap": None, + "num_cpus": 1, + } + + if field_container is Data: + container = field_container( + **kwargs, + fields=True, + particles=False, + spectra=False, + ) + assert not container.fields_defined, "Fields are unexpectedly defined" + with pytest.raises(ValueError, match="Fields are not defined"): + _ = container.fields + else: + container = field_container(**kwargs, coord_system=None) + assert not container.fields_defined, "Fields are unexpectedly defined" + assert container.fields is None, "Fields are not None" + + @pytest.mark.parametrize( "test,particle_container", - [[test, fc] for test in TESTS for fc in [Data, Particles]], + [ + [test, container] + for test in PARTICLE_TESTS + for container in [Data, ParticleContainer] + ], ) -def test_particles(test, particle_container: Union[Type[Data], Type[Particles]]): +def test_particles( + test, particle_container: Union[Type[Data], Type[ParticleContainer]] +): reader: BaseReader = test["reader"]() - PATH = test["path"] - if test["particles"] == {}: - return + path = test["path"] - def prtl_remap(Xold: str) -> str: + def prtl_remap_cart(Xold: str) -> str: return { "pX1": "x", "pX2": "y", @@ -127,8 +196,8 @@ def prtl_remap(Xold: str) -> str: "pW": "w", }.get(Xold, Xold) - if test.get("coords", "cart") != "cart": - prtl_remap = lambda Xold: { + def prtl_remap_sph(Xold: str) -> str: + return { "pX1": "r", "pX2": "th", "pX3": "ph", @@ -137,9 +206,47 @@ def prtl_remap(Xold: str) -> str: "pU3": "uph", "pW": "w", }.get(Xold, Xold) - particles = particle_container( - path=PATH, - reader=reader, - remap={"particles": prtl_remap}, - ) - assert particles.particles is not None, "Particles are None" + + remap = { + "particles": ( + prtl_remap_cart if test.get("coords", "cart") == "cart" else prtl_remap_sph + ) + } + if particle_container is Data: + particles = particle_container( + path=path, + fields=False, + particles=True, + spectra=False, + reader=reader, + remap=remap, + num_cpus=1, + ) + else: + particles = particle_container( + path=path, + reader=reader, + remap=remap, + coord_system=None, + num_cpus=1, + ) + + dataset = particles.particles + assert dataset is not None, "Particles are None" + assert len(dataset.species) == test["particles"].get("nspec", 4) + + selected = dataset.isel(t=-1).load(cols=["id", "sp", "w"]) + last_step = int(dataset.steps[-1]) + expected_counts = [ + reader.ReadArrayShapeAtTimestep( + path=path, + category="particles", + quantity=f"pW_{species}", + step=last_step, + )[0] + for species in dataset.species + ] + + assert selected["st"].unique().tolist() == [last_step] + assert selected.groupby("sp", sort=True).size().tolist() == expected_counts + assert selected["w"].notna().all() diff --git a/nt2/tests/test_diagnostics.py b/nt2/tests/test_diagnostics.py new file mode 100644 index 0000000..abc3cdd --- /dev/null +++ b/nt2/tests/test_diagnostics.py @@ -0,0 +1,117 @@ +import os +from pathlib import Path + +import pandas as pd +import pytest + +from nt2.containers.diagnostics import Diagnostics + + +def _write_log(path: Path, steps: int = 3) -> None: + preamble = "Entity diagnostic output\n\n" + records = [] + for step in range(1, steps + 1): + records.append( + f"""Step: {step}....................[of {steps}] +Time: {step * 0.25:.2f}................[Δt = 0.25] + +[SUBSTEP]..................[DURATION] + Communications..............{step}.00 ms 3% + FieldSolver................{step + 1}.00 µs 1% + Injector.....................0.00 ns 0% + +[PARTICLE SPECIES] [TOTAL] [% TOT] + species 1 (e-)............{step}.20e+03 1% + species 2 (i+)............2.00e+03 1% + +................................................................................ + +""" + ) + path.write_text(preamble + "".join(records), encoding="utf-8") + + +def test_streaming_parser_without_cache(tmp_path: Path) -> None: + outfile = tmp_path / "simulation.out" + _write_log(outfile) + + diagnostics = Diagnostics(tmp_path, cache=False, chunk_size=2) + + expected = pd.DataFrame( + { + "Step": [1, 2, 3], + "Time": [0.25, 0.50, 0.75], + "Communications": [1e6, 2e6, 3e6], + "FieldSolver": [2e3, 3e3, 4e3], + "Injector": [0.0, 0.0, 0.0], + "species_1": [1200, 2200, 3200], + "species_2": [2000, 2000, 2000], + } + ).set_index("Step", drop=False) + pd.testing.assert_frame_equal(diagnostics.df, expected) + + +def test_parquet_cache_is_reused_and_invalidated( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + outfile = tmp_path / "simulation.out" + cache = tmp_path / "diagnostics.parquet" + _write_log(outfile, steps=2) + first = Diagnostics(outfile, cache_path=cache, chunk_size=1) + assert first.df is not None + assert cache.is_file() + assert Path(f"{cache}.json").is_file() + + def fail_if_parsed(source: Path): + raise AssertionError(f"Unexpected parse of {source}") + + monkeypatch.setattr(Diagnostics, "_records", fail_if_parsed) + cached = Diagnostics(outfile, cache_path=cache) + pd.testing.assert_frame_equal(cached.df, first.df) + + monkeypatch.undo() + with outfile.open("a", encoding="utf-8") as stream: + stream.write("\n") + os.utime(outfile, None) + reparsed = Diagnostics(outfile, cache_path=cache) + pd.testing.assert_frame_equal(reparsed.df, first.df) + + +def test_inconsistent_record_raises(tmp_path: Path) -> None: + outfile = tmp_path / "simulation.out" + _write_log(outfile, steps=2) + text = outfile.read_text(encoding="utf-8") + outfile.write_text( + text.replace(" Injector.....................0.00 ns 0%\n", "", 1), + encoding="utf-8", + ) + + with pytest.raises(ValueError, match="Inconsistent diagnostics record at step 2"): + Diagnostics(outfile, cache=False) + + +def test_incomplete_final_record_is_ignored(tmp_path: Path) -> None: + outfile = tmp_path / "simulation.out" + _write_log(outfile, steps=2) + with outfile.open("a", encoding="utf-8") as stream: + stream.write("Step: 3....................[of 3]\nTime: 0.75") + + diagnostics = Diagnostics(outfile, cache=False) + + assert diagnostics.df is not None + assert diagnostics.df["Step"].tolist() == [1, 2] + + +def test_no_outfile(tmp_path: Path, caplog: pytest.LogCaptureFixture) -> None: + diagnostics = Diagnostics(tmp_path) + + assert diagnostics.df is None + assert "No .out files found" in caplog.text + + +def test_cache_cannot_overwrite_source(tmp_path: Path) -> None: + outfile = tmp_path / "simulation.out" + _write_log(outfile) + + with pytest.raises(ValueError, match="must not overwrite"): + Diagnostics(outfile, cache_path=outfile) diff --git a/nt2/tests/test_particle_container_serialization.py b/nt2/tests/test_particle_container_serialization.py new file mode 100644 index 0000000..768a723 --- /dev/null +++ b/nt2/tests/test_particle_container_serialization.py @@ -0,0 +1,34 @@ +import pickle + +import numpy as np + +from nt2.containers.particle_dataset import ParticleDataset +from nt2.containers.particles import ParticleContainer + + +class _RestoringParticleContainer(ParticleContainer): + def _read_particles(self): + raise AssertionError("deserialization should not rescan particle files") + + def _read_column(self, step, colname): + return np.array([], dtype=np.int64) + + +def test_particle_dataset_is_restored_after_serialization(): + container = _RestoringParticleContainer.__new__(_RestoringParticleContainer) + container.__dict__["_ParticleContainer__particles_defined"] = True + container.__dict__["_ParticleContainer__particles"] = ParticleDataset( + species=[1], + steps=np.array([10]), + times=np.array([2.5]), + colnames=["x", "id", "sp"], + read_column=container._read_column, + partition_lengths=(0,), + ) + + restored = pickle.loads(pickle.dumps(container)) + + assert restored.particles_defined + assert restored.particles is not None + assert restored.particles.species == [1] + assert restored.particles.times.tolist() == [2.5] diff --git a/nt2/tests/test_particle_dataset_nbytes.py b/nt2/tests/test_particle_dataset_nbytes.py new file mode 100644 index 0000000..74cd07e --- /dev/null +++ b/nt2/tests/test_particle_dataset_nbytes.py @@ -0,0 +1,41 @@ +import numpy as np + +from nt2.containers.particle_dataset import ParticleDataset + + +def _unexpected_read(step: int, column: str): + raise AssertionError(f"nbytes unexpectedly read {column} at step {step}") + + +def test_nbytes_uses_cached_partition_lengths_without_computing(): + particles = ParticleDataset( + species=[1], + steps=np.array([10, 20], dtype=np.int64), + times=np.array([1.0, 2.0], dtype=np.float64), + colnames=["x", "id", "sp"], + read_column=_unexpected_read, + partition_lengths=(1, 2), + ) + + bytes_per_row = sum( + np.dtype(dtype).itemsize + for dtype in (np.int64, np.int32, np.int64, np.float64, np.int64) + ) + assert particles.nbytes == 3 * bytes_per_row + + +def test_nbytes_tracks_timestep_partition_selection_without_computing(): + particles = ParticleDataset( + species=[1], + steps=np.array([10, 20], dtype=np.int64), + times=np.array([1.0, 2.0], dtype=np.float64), + colnames=["x", "id", "sp"], + read_column=_unexpected_read, + partition_lengths=(1, 2), + ) + + bytes_per_row = sum( + np.dtype(dtype).itemsize + for dtype in (np.int64, np.int32, np.int64, np.float64, np.int64) + ) + assert particles.isel(t=-1).nbytes == 2 * bytes_per_row diff --git a/nt2/tests/test_reader.py b/nt2/tests/test_reader.py index 17a889e..a0d6582 100644 --- a/nt2/tests/test_reader.py +++ b/nt2/tests/test_reader.py @@ -12,16 +12,16 @@ def pytest_generate_tests(metafunc): def check_equal_arrays(arr1, arr2): if isinstance(arr1, set): - assert len(arr1) == len( - arr2 - ), f"Set lengths do not match: {len(arr1)} != {len(arr2)}" + assert len(arr1) == len(arr2), ( + f"Set lengths do not match: {len(arr1)} != {len(arr2)}" + ) assert arr1 == arr2, f"Sets do not match: {arr1} != {arr2}" else: arr1 = np.array(arr1) arr2 = np.array(arr2) - assert ( - arr1.shape == arr2.shape - ), f"Shapes do not match: {arr1.shape} != {arr2.shape}" + assert arr1.shape == arr2.shape, ( + f"Shapes do not match: {arr1.shape} != {arr2.shape}" + ) assert np.all(np.isclose(arr1, arr2)), f"Arrays do not match: {arr1} != {arr2}" @@ -78,9 +78,9 @@ def test_reader(test): field_names = test["fields"].get( "quantities", - [f"{f}{i+1}" for i in range(3) for f in "BE"] + [f"{f}{i + 1}" for i in range(3) for f in "BE"] + [f"N_{i}" for i in ["1_2", "3_4"]] - + [f"T0{c+1}_{i+1}" for i in range(4) for c in range(3)], + + [f"T0{c + 1}_{i + 1}" for i in range(4) for c in range(3)], ) field_names = set(f"f{f}" for f in field_names) # Check that invalid_tstep raises OSError in fields @@ -92,13 +92,20 @@ def test_reader(test): OSError, ) + valid_files, valid_steps = reader.GetValidFilesAndSteps( + path=PATH, category="fields" + ) + # Check that timesteps are read correctly from fields - times = reader.ReadPerTimestepVariable( - path=PATH, category="fields", varname="Time", newname="t" - )["t"] - steps = reader.ReadPerTimestepVariable( - path=PATH, category="fields", varname="Step", newname="s" - )["s"] + timestep_variables = reader.ReadPerTimestepVariables( + path=PATH, + category="fields", + varnames=["Time", "Step"], + newnames=["t", "s"], + valid_files=valid_files, + ) + times = timestep_variables["t"] + steps = timestep_variables["s"] check_equal_arrays( times, np.array([s * dt for s in steps]), @@ -119,7 +126,6 @@ def test_reader(test): check_equal_arrays(coords["X2"], x2) if test["dim"] == "3D": - sx3 = test["fields"]["sx3"] x3min = dx / 2 nx3 = test["fields"]["nx3"] x3 = np.array([x3min + i * dx for i in range(int(nx3))]) @@ -140,18 +146,26 @@ def test_reader(test): shape, (nx1, nx2, nx3) if layout == Layout.R else (nx3, nx2, nx1) ) - for step in reader.GetValidSteps(path=PATH, category="fields"): + for step in valid_steps: for f in field_names: field = reader.ReadArrayAtTimestep( - path=PATH, category="fields", quantity=f, step=step + path=PATH, + category="fields", + quantity=f, + step=step, ) check_equal_arrays(field.shape, shape) - reader.VerifySameCategoryNames(path=PATH, category="fields", prefix="f") - reader.VerifySameFieldLayouts(path=PATH) + reader.VerifySameCategoryNames( + path=PATH, + category="fields", + prefix="f", + valid_steps=valid_steps, + ) + reader.VerifySameFieldLayouts(path=PATH, valid_steps=valid_steps) # Check that the shapes of the fields are read correctly - reader.VerifySameFieldShapes(path=PATH) + reader.VerifySameFieldShapes(path=PATH, valid_steps=valid_steps) if test["particles"] != {}: dt = 0 @@ -171,9 +185,9 @@ def test_reader(test): nspec: int = test["particles"].get("nspec", 4) prtl_names = ( - [f"U{i+1}_{j+1}" for i in range(3) for j in range(nspec)] + [f"U{i + 1}_{j + 1}" for i in range(3) for j in range(nspec)] + [ - f"X{i+1}_{j+1}" + f"X{i + 1}_{j + 1}" for i in range( 2 if test["dim"] == "2D" and test.get("coords", "cart") == "cart" @@ -181,17 +195,23 @@ def test_reader(test): ) for j in range(nspec) ] - + [f"W_{i+1}" for i in range(nspec)] + + [f"W_{i + 1}" for i in range(nspec)] ) prtl_names = set(f"p{p}" for p in prtl_names) + valid_files, valid_steps = reader.GetValidFilesAndSteps( + path=PATH, category="particles" + ) # Check that timesteps are read correctly from particles - times = reader.ReadPerTimestepVariable( - path=PATH, category="particles", varname="Time", newname="t" - )["t"] - steps = reader.ReadPerTimestepVariable( - path=PATH, category="particles", varname="Step", newname="s" - )["s"] + timestep_variables = reader.ReadPerTimestepVariables( + path=PATH, + category="particles", + varnames=["Time", "Step"], + newnames=["t", "s"], + valid_files=valid_files, + ) + times = timestep_variables["t"] + steps = timestep_variables["s"] if dt is not None: check_equal_arrays( @@ -207,11 +227,11 @@ def test_reader(test): check_equal_arrays(names, prtl_names) # Check prtl shapes - for step in reader.GetValidSteps(path=PATH, category="particles"): + for step in valid_steps: for sp in range(nspec): reader.ReadArrayShapeAtTimestep( path=PATH, category="particles", - quantity=f"pW_{sp+1}", + quantity=f"pW_{sp + 1}", step=step, ) diff --git a/nt2/tests/test_tokenize.py b/nt2/tests/test_tokenize.py index 7e49f74..be73014 100644 --- a/nt2/tests/test_tokenize.py +++ b/nt2/tests/test_tokenize.py @@ -1,6 +1,7 @@ +import numpy as np from dask.base import tokenize -from nt2.containers.container import BaseContainer +from nt2.containers.base import BaseContainer from nt2.readers.base import BaseReader from nt2.utils import Format @@ -10,9 +11,24 @@ class _Reader(BaseReader): def format(self) -> Format: return Format.HDF5 + def GetValidFilesAndSteps(self, path, category, steprange=None, num_cpus=None): + return (["fields.00000001.h5"], [1]) + + def ReadPerTimestepVariables( + self, path, category, varnames, newnames, valid_files + ): + return {"t": np.array([0.5]), "s": np.array([1])} + def test_base_container_has_deterministic_dask_token(): - container = BaseContainer(path="/tmp/sim", reader=_Reader(), remap=None) + container = BaseContainer( + path="/tmp/sim", + category="fields", + reader=_Reader(), + remap=None, + coord_system=None, + num_cpus=1, + ) token1 = tokenize(container) token2 = tokenize(container) diff --git a/nt2/utils.py b/nt2/utils.py index 12ea889..8509d71 100644 --- a/nt2/utils.py +++ b/nt2/utils.py @@ -1,10 +1,12 @@ -from typing import Union -from enum import Enum +from __future__ import annotations + +import inspect import os import re -import inspect -import numpy as np +from enum import Enum +from typing import Literal +import numpy as np import xarray as xr @@ -26,6 +28,19 @@ class CoordinateSystem(Enum): XYZ = "Cartesian" SPH = "Spherical" + @staticmethod + def from_str(s): + s = s.lower() + if s in ("cartesian", "xyz"): + return CoordinateSystem.XYZ + elif s in ("spherical", "sph"): + return CoordinateSystem.SPH + else: + raise ValueError(f"Unknown coordinate system: {s}") + + +CoordinateSystemType = Literal["XYZ", "SPH"] + def DetermineDataFormat(path: str) -> Format: """Determine the data format for the files in the given path. @@ -64,7 +79,7 @@ def DetermineDataFormat(path: str) -> Format: raise ValueError("Could not determine file format.") -def ToHumanReadable(num: Union[float, int], suffix: str = "B") -> str: +def ToHumanReadable(num: float, suffix: str = "B") -> str: """Convert a number to a human-readable format with SI prefixes. Parameters diff --git a/pyproject.toml b/pyproject.toml index fef7be7..49ada55 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,8 +1,8 @@ [project] name = "nt2py" -description = "Post-processing & visualization toolkit for the Entity PIC code" +description = "Post-processing & visualization package for the Entity PIC code" readme = "README.md" -requires-python = ">=3.8" +requires-python = ">=3.9" license-files = ["LICENSE"] authors = [{ name = "Hayk", email = "haykh.astro@gmail.com" }] maintainers = [{ name = "Hayk", email = "haykh.astro@gmail.com" }] @@ -12,7 +12,6 @@ classifiers = [ "Intended Audience :: Science/Research", "License :: OSI Approved :: BSD License", "Programming Language :: Python :: 3 :: Only", - "Programming Language :: Python :: 3.8", "Programming Language :: Python :: 3.9", "Programming Language :: Python :: 3.10", "Programming Language :: Python :: 3.11", @@ -47,7 +46,7 @@ Repository = "https://github.com/entity-toolkit/nt2py" nt2 = "nt2.cli.main:app" [project.optional-dependencies] -dev = ["black", "pytest"] +dev = ["ruff", "pyrefly", "pytest", "types-tqdm"] hdf5 = ["h5py"] [build-system] @@ -59,7 +58,7 @@ path = "nt2/__init__.py" [tool.hatch.build.targets.wheel] packages = ["nt2"] -exclude = ["/.git", "/.venv", "/dist", "/temp", "/nt2/tests/"] +exclude = ["/nt2/tests/"] [tool.hatch.build.targets.sdist] -exclude = ["/.git", "/.venv", "/dist", "/temp", "/nt2/tests/"] +exclude = ["/legacy/"] \ No newline at end of file diff --git a/pyrightconfig.json b/pyrightconfig.json deleted file mode 100644 index cb36191..0000000 --- a/pyrightconfig.json +++ /dev/null @@ -1,12 +0,0 @@ -{ - "extraPath": [ - "./" - ], - "reportAny": false, - "reportExplicitAny": false, - "reportUnknownVariableType": false, - "reportUnknownMemberType": false, - "reportUnknownArgumentType": false, - "reportArgumentType": false, - "reportPrivateImportUsage": false -}