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