diff --git a/CLAUDE.md b/CLAUDE.md index 667be1b2..04644da0 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -109,45 +109,60 @@ VS Code: install the official Ruff extension (`charliermarsh.ruff`) and add to y ## Architecture -The code lives under `src/lib/` and is organized around three concepts: **sources**, **adaptors**, and **plots**. Argument parsing wires them together. +The code lives under `src/lib/` and is organized around **loaders** (sources), **adaptors** (transforms), and **plots** (renderers). A lazy **node graph** wraps a **`DataWorld`** value that flows through them; argument parsing wires it all together. ### Data flow -1. `parsing.get_parsed_args()` (`src/lib/parsing/parse.py`) builds a flat argparse parser with positional `prefix` (choices-constrained to known prefixes) and optional `variable`, plus all registered adaptor/hook flags. Returns an `Args` namespace (`src/lib/parsing/args.py`) directly. -2. `parsing.get_parsed_args()` also calls `discover_loaders(CONFIG.data_dir)` (`src/lib/data/loader.py`) to map each prefix to its loader class, then instantiates the chosen loader as `args.loader`. `args.get_animation()` calls `compile_source` (`src/lib/data/compile.py`) to wrap `args.loader` in a `DataSourceWithPipeline` made of the user-supplied `Adaptor` list. If no `Versus` adaptor is present, a default one is appended (`y,z` vs `t`) — this is what selects axes and time dim. -3. `source.get_data()` loads raw data and runs the pipeline, returning a `DataWithAttrs` (a `Field` wrapping `xr.Dataset`, or a `List` wrapping `pd.DataFrame` / `dd.DataFrame`). `DataSource` is an ABC with a single abstract method `get_data()`. -4. `get_plot(data)` (`src/lib/plotting/get_plot.py`) picks a `Renderer` based on data type and geometry (`Field1dRenderer`, `Field2dRenderer`, `PolarFieldRenderer`, `ScatterRenderer`), then wraps it in `StaticPlot` or `AnimatedPlot` based on whether `time_dim` is present. `ScatterRenderer` requires exactly 2 `spatial_dims`. Each renderer defines `make_init_data`/`init` (one-time setup with frame 0) and `make_update_data`/`draw` (per-frame update for animations). Hooks are added to the plot object after construction. -5. Hooks (`src/lib/plotting/hooks/`) such as `--scale log`, `--grid`, `--vline`, `--fit` are appended onto the chosen plot before `show()`/`save()`. +1. `cli.main()` (`src/lib/cli.py`) configures dask, calls `parse_args()` (`src/lib/parsing/parse.py`) — a flat argparse parser with positional `prefix` + optional `variable` plus all registered adaptor/hook flags, returning an `Args` namespace (`src/lib/parsing/args.py`). It then calls `compile_action_nodes(args, config)` (`src/lib/data/compile.py`) and `.pull()`s each returned action node. (The `prefix` positional is **not** choices-constrained; an unknown prefix fails later in `get_loader`.) +2. Everything runs through a lazy, memoized **node graph** (`src/lib/data/node.py`). `compile_action_nodes` chains: `RootNode(config)` → `AdaptorNode(loader)` → one `AdaptorNode` per user adaptor → `PlotNode(hooks)` → action node(s): `ShowPlotNode` (`--show`, default on), `SavePlotNode` (`--save`), and/or `DaskGraphNode` (`--dask-graph`, renders the dask graph as SVG instead of plotting). Each `DataProcessingNode` has a `@cache`d `pull()` and accumulates `name_fragments` (used to derive save filenames). If the adaptor list contains no `Versus`, a default `Versus(["y","z"], time_dim_rule="guess", color_dim=None)` is appended (`_with_versus`) — this is what selects axes/color/time dims and appends the plot target. +3. Data flows as a **`DataWorld`** (`src/lib/data/data_world.py`): a frozen dataclass holding `datas: dict[str, DataWithAttrs]`, an `active_key: str | None`, `plot_targets: list[PlotTarget]`, and `config`. `RootNode.pull()` returns an empty world; each `AdaptorNode` applies a `WorldAdaptor` to it. The **loader is itself a `WorldAdaptor`** (`Loader.apply_world` inserts the loaded `DataWithAttrs` under key = `prefix` and makes it active). `DataWorld.active_data` / `with_active_data()` operate on `datas[active_key]`. Holding **multiple named datas + multiple plot targets** is the "split vars" capability this branch is named for. +4. `PlotNode.pull()` calls `get_plot(world)` (`src/lib/plotting/get_plot.py`), which builds one `Renderer` **per `PlotTarget`**: `Field1dRenderer` when the target has no `color_dim`; `PolarFieldRenderer` / `Field2dRenderer` chosen by `SpatialDimsRTheta` / `SpatialDimsXY`; `ScatterRenderer` for `List` data with exactly 2 spatial dims. It wraps the renderer **list** in `StaticPlot` or `AnimatedPlot` based on `n_frames = max(r.get_n_frames())` — **animated iff `n_frames > 1`**, not on `time_dim` presence. Hooks are then attached via `plot.add_hook()`. +5. `ShowPlotNode`/`SavePlotNode.pull()` call `plot.show()` / `plot.save_to_path()`, which lazily `_initialize()` the figure (see the PlotInfo/AxesManager layer below). + +### PlotTarget + +`src/lib/data/plot_target.py`: a `PlotTarget` names one thing to draw — a `prefix` (which entry of `DataWorld.datas`), a `spatial_dims` (`SpatialDimsXY(x_dim, y_dim)` or `SpatialDimsRTheta(r_dim, theta_dim)`), an optional `color_dim` and `time_dim`, and an `axes_index: (col, row)` (1-based) selecting which subplot it lands in. `Versus.apply_world` is what constructs and appends targets; multiple targets sharing an `axes_index` are overlaid on one axes. + +### WorldAdaptor / Adaptor class hierarchy + +`src/lib/data/adaptor.py`: +- `WorldAdaptor` (ABC) — single abstract `apply_world(world) -> DataWorld`. The shared node-graph interface for **both loaders and adaptors**. Adaptors that must touch the whole world (e.g. `Versus`, which reads `active_data` and appends a `PlotTarget`) override this directly. +- `Adaptor(WorldAdaptor)` — default `apply_world` = `world.with_active_data(self.apply(world.active_data))`. Override `apply_field`/`apply_list`; the unused one raises a friendly "use `--bin`/`--scatter`" error. +- `MetadataAdaptor(Adaptor)` — wraps `apply` to also modify the active variable's `VarInfo` in `var_infos` (used to derive axis labels/filenames). Override `get_modified_display_latex(metadata)` and/or `get_modified_unit_latex(metadata)`; both receive the current `metadata` so they can inspect e.g. `active_key` and `active_var_info`. +- `BareAdaptor(MetadataAdaptor)` — operates on the raw active variable (a single `xr.DataArray` for fields, a single `pd.Series`/`dd.Series` for lists) and doesn't touch metadata; override `apply_field_bare`/`apply_list_bare`. ### Auto-registration of loaders, adaptors, and hooks `src/lib/__init__.py` imports `lib.data.loaders`, `lib.data.adaptors`, and `lib.plotting.hooks`, whose `__init__.py` files glob and `importlib.import_module` every sibling `*.py`. -**Loaders** (`src/lib/data/loaders/`) register themselves via the bare `@loader` decorator from `src/lib/data/loader.py`, which appends to `LOADERS: list[type[Loader]]`. Each loader class exposes a `discover_prefixes(cls, data_dir: Path) -> list[str]` classmethod that returns the prefixes it claims in that directory. +**Loaders** (`src/lib/data/loaders/`) are `Loader` (a `WorldAdaptor` subclass) classes registered via the bare `@loader` decorator from `src/lib/data/loader.py`, which appends to `LOADERS: list[type[Loader]]`. Each loader exposes `discover_prefixes(cls, data_dir: Path) -> list[str]` and `suffix(cls)` classmethods, plus `get_data(config) -> DataWithAttrs`. -`discover_loaders(data_dir)` polls every registered loader's `discover_prefixes()` and returns a `dict[str, type[Loader]]`; the resulting keys become argparse's prefix choices. On prefix conflicts, the later-registered loader wins with a `UserWarning`, so user-defined loaders can shadow built-ins. To add a new data source, drop a file into `src/lib/data/loaders/`, decorate the class with `@loader`, and implement `discover_prefixes`. +`discover_loaders(data_dir)` polls every registered loader's `discover_prefixes()` and returns a `dict[str, type[Loader]]`; `get_loader(data_dir, prefix, active_key)` instantiates the one matching `prefix`. On prefix conflicts, the later-registered loader wins with a `UserWarning`, so user-defined loaders can shadow built-ins. To add a new data source, drop a file into `src/lib/data/loaders/`, decorate the class with `@loader`, and implement `discover_prefixes`/`suffix`/`get_data`. **Adaptors and hooks** register their argparse flags via the `@arg_parser(...)` / `@const_arg(...)` decorators in `src/lib/parsing/args_registry.py`, which append to the module-level `CUSTOM_ARGS` list. `parse._get_parser()` then iterates `CUSTOM_ARGS` and adds them to the parser. To add a new adaptor or hook, drop a new file into `src/lib/data/adaptors/` (or `src/lib/plotting/hooks/`) and decorate its parse function. **Consequence:** a loader, adaptor, or hook that fails to register (e.g. import error in that file) will silently disappear from the CLI; suspect the auto-import if a flag or prefix goes missing. -### Adaptor class hierarchy - -`src/lib/data/adaptor.py`: -- `Adaptor` — base. Override `apply_field`/`apply_list`; the unused one raises a friendly "use `--bin`/`--scatter`" error. -- `MetadataAdaptor` — wraps `apply` to also append name fragments and modify the active variable's `VarInfo` in `var_infos` (used to derive saved filenames and axis labels). Override `get_modified_display_latex(metadata)` and/or `get_modified_unit_latex(metadata)`; both receive the current `metadata` so they can inspect e.g. `active_key` and `active_var_info`. -- `BareAdaptor` — for adaptors that operate on the active variable (a single `xr.DataArray` for fields, a single `pd.Series`/`dd.Series` for lists) and don't touch metadata. - -`Pipeline` (`src/lib/data/pipeline.py`) is itself an `Adaptor` that chains a list of adaptors. - ### Data wrapper -`src/lib/data/data_with_attrs.py` defines `DataWithAttrs[D, MD]` and concrete `Field` (`xr.Dataset`-backed), `FullList` (pandas), `LazyList` (dask). Frozen dataclasses; mutate via `assign_data` / `assign_metadata` / `assign`. `Metadata` carries `active_key` (`str | None`), `var_infos` (`dict[str, VarInfo]` — maps all known variable/dimension keys to `VarInfo` objects), `name_fragments`, `spatial_dims`, `time_dim`, and `color_dim`. `active_key` defaults to `None` — particle data may have no active variable (e.g. pure scatter of positions). The convenience property `active_var_info` returns `var_infos[active_key]`. `var_infos` is populated at load time from `src/lib/var_info_registry.py` via `lookup(prefix, key)` for every coordinate and the active variable. `FieldMetadata` also carries `prefix` (the file prefix, e.g. `"pfd_moments"`). `ListMetadata` also carries `subject: Latex | None` — describes what the list contains (e.g. "Particles", "Ions", "Electrons"); set by `ParticleLoader`, refined by `SpeciesFilter`, and used by `Bin` (for distribution function subscripts) and `ScatterRenderer` (for plot titles). `ListMetadata` also carries optional `partition_dim: str | None` and `partition_ranges: list[tuple[int,int]] | None` — when set (currently by both particle loaders, with `partition_dim="t"`), they let `Idx.apply_list` prune by `df.partitions[...]` instead of a `df[df[dim] == pos]` predicate filter. **Loader invariant:** `partition_ranges` must describe the actual partition layout of the `dd.DataFrame` returned (one entry per value of `partition_dim`, each `(start, end)` matching the per-step `npartitions`). `LazyList.compute()` clears these fields because they describe the dask layout and become meaningless after materialization. The unusual `**` unpacking via `__getitem__` + `keys()` is what `Metadata.create_from` and `assign` use to round-trip values between subclasses (`FieldMetadata` vs `ListMetadata`). +`src/lib/data/data_with_attrs.py` defines `DataWithAttrs[D, MD]` and concrete `Field` (dict of `xr.DataArray`), `FullList` (pandas), `LazyList` (dask). Frozen dataclasses; mutate via `assign_data` / `assign_metadata` / `assign`. `Metadata` now carries only `active_key` (`str | None`) and `var_infos` (`dict[str, VarInfo]` — maps all known variable/dimension keys to `VarInfo` objects). `active_key` defaults to `None` — particle data may have no active variable (e.g. pure scatter of positions). The convenience property `active_var_info` returns `var_infos[active_key]`. `var_infos` is populated at load time from `src/lib/var_info_registry.py` via `lookup(prefix, key)` for every coordinate and the active variable. `FieldMetadata` also carries `prefix` (the file prefix, e.g. `"pfd_moments"`). `ListMetadata` also carries `subject: Latex | None` — describes what the list contains (e.g. "Particles", "Ions", "Electrons"); set by `ParticleLoader`, refined by `SpeciesFilter`, and used by `Bin` (for distribution function subscripts) and `ScatterRenderer` (for plot titles). `ListMetadata` also carries optional `partition_dim: str | None` and `partition_ranges: list[tuple[int,int]] | None` — when set (currently by both particle loaders, with `partition_dim="t"`), they let `Idx.apply_list` prune by `df.partitions[...]` instead of a `df[df[dim] == pos]` predicate filter. **Loader invariant:** `partition_ranges` must describe the actual partition layout of the `dd.DataFrame` returned (one entry per value of `partition_dim`, each `(start, end)` matching the per-step `npartitions`). `LazyList.compute()` clears these fields because they describe the dask layout and become meaningless after materialization. The unusual `**` unpacking via `__getitem__` + `keys()` is what `Metadata.create_from` and `assign` use to round-trip values between subclasses (`FieldMetadata` vs `ListMetadata`). + +> **Note (split-vars):** the `spatial_dims` / `time_dim` / `color_dim` axis-selection fields and the `name_fragments` that `Metadata` used to carry have moved out — geometry/axis selection now lives on `PlotTarget` (inside `DataWorld`), and `name_fragments` are accumulated by the node graph (`DataProcessingNode.name_fragments` / `HasNameFragments`). Both `Field` and `List` expose an `active_data` property and `with_active_data()` method. For `Field`, `active_data` returns the `xr.DataArray` for `metadata.active_key`; `with_active_data(da)` replaces it and drops grid-incompatible siblings. For `List`, `active_data` returns the `pd.Series`/`dd.Series` column for `metadata.active_key`; `with_active_data(series)` replaces that column. Both raise `ValueError` if `active_key` is `None`. Most code should use `active_data` rather than `data` directly. `BareAdaptor` handles this automatically via the shims in `adaptor.py`. The class-level `data: ...`/`metadata: ...` annotations on the subclasses look redundant but are intentional — see the comment in `DataWithAttrs.__init__`. They are needed so `isinstance`-narrowed code gets the concrete types; don't "clean them up." +### PlotInfo, renderers, and AxesManager + +Renderers don't draw to matplotlib directly — they produce a `PlotInfo` (`src/lib/plotting/plot_info.py`: `LineInfo` / `ImageInfo` / `ScatterInfo` / `PolarMeshInfo`, all `PlotInfo`/`PlotInfo2D`) describing *what* to draw: data arrays, per-dim `dim_scales` / `dim_bounds` / `dim_displays` / `dim_units`, `scalar_coord_values` (for titles), `axes_index`, and `projection`. `Renderer.__init__` calls `init_plot_info()`; `update_plot_info(frame)` mutates the `PlotInfo` for animation via `PlotInfo.set(key, value)`, which fires registered `_setter_callbacks` that update the matplotlib artists in place. + +`setup_fig(plot_infos)` (`src/lib/plotting/setup_fig.py`) is called once from `Plot._initialize()`. It groups infos by `axes_index` into a subplot grid, then picks an `AxesManager` per axes: single info → `AxesManagerSingleLine` / `…Image` / `…Scatter` / `…PolarMesh`; several `LineInfo`s → `AxesManagerMultiLine` (shared axes + legend); one `ImageInfo` + `LineInfo`s → `AxesManagerImageAndLines` (lines on a twinned y-axis). Each manager wires `PlotInfo._setter_callbacks` so later `.set()` calls re-render without rebuilding the figure. Figures use `layout="constrained"`. + +### Hooks + +Hooks (`src/lib/plotting/hooks/`) such as `--grid`, `--vline`, `--fit`, `--show-com`, `--show-initial` subclass `Hook` (`src/lib/plotting/hook.py`) and implement `post_init_fig(message)` / `post_update_fig(message)`, receiving a `DrawMessage(plot_info, axes, frame_data)`. `PlotNode` attaches them and `Plot._initialize()` calls `post_init_fig` after building the figure. **Currently hooks are applied to the first renderer/axes only** — see the TODO in `plot.py`. + ### Dimensions and var_infos `src/lib/var_info.py` defines `VarInfo` as a frozen value (`display: Latex`, `unit: Latex`, `geometry`, `key`). `src/lib/var_info_registry.py` provides a single `_REGISTRY` dict keyed by `(prefix | None, key)` — `None`-prefix entries are shared dimensions (x, y, z, t); string-prefix entries are per-file-type variables (e.g. `("pfd", "hx_fc")` → `VarInfo(display="B_x", ...)`). diff --git a/src/lib/cli.py b/src/lib/cli.py index a7d4f479..d0b71faa 100644 --- a/src/lib/cli.py +++ b/src/lib/cli.py @@ -1,117 +1,25 @@ -import sys -import warnings -import webbrowser -from pathlib import Path - import dask -import matplotlib.pyplot as plt - -from lib import parsing -from lib.config import CONFIG -from lib.parsing.args import Args -from lib.plotting.plot import SaveFormat - - -def _resolve_save_format(args: Args) -> SaveFormat | None: - if args.save is None: - if args.save_format is not None: - print("error: --save-format requires --save", file=sys.stderr) - sys.exit(1) - return None - - if args.save_format == "mp4": - if not CONFIG.ffmpeg_bin: - print("error: --save-format mp4 requires ffmpeg", file=sys.stderr) - sys.exit(1) - return "mp4" - - if args.save_format == "gif": - return "gif" - - # save_format is None: try mp4, fall back to gif - if CONFIG.ffmpeg_bin: - return "mp4" - - message = "ffmpeg not found; will save animations as gif instead of mp4" - warnings.warn(message) - return "gif" - - -def _run_dask_graph(args: Args) -> None: - data = args.get_data() - collections = data.dask_collections() - if not collections: - print( - f"error: --dask-graph requires dask-backed data; pipeline produced eager {type(data).__name__}", - file=sys.stderr, - ) - sys.exit(1) - - try: - import graphviz # noqa: F401 - except ImportError: - print( - "error: --dask-graph requires the 'graphviz' package; install with `pip install -e \".[dask-graph]\"`", - file=sys.stderr, - ) - sys.exit(1) - - save_dir = args.save or Path.cwd() - save_dir.mkdir(exist_ok=True, parents=True) - path = save_dir / f"{args.get_save_file_stem()}.daskgraph.svg" - # dask.visualize's optimize_graph flag only lowers legacy HLG collections - # (e.g. dask Arrays), not new-style Expr ones (dask DataFrames) — without - # pre-optimizing the latter, un-lowered nodes (e.g. Concat from dd.concat) - # fail with NotImplementedError in _layer. - collections = [c.optimize() if hasattr(c, "optimize") else c for c in collections] - dask.visualize(*collections, filename=str(path), optimize_graph=True) - print(f"wrote to {path}") - - if args.show: - webbrowser.open(path.absolute().as_uri()) +from lib.config import PscPlotConfig +from lib.data.compile import compile_action_nodes +from lib.parsing.parse import parse_args def main(): - dask.config.set(num_workers=CONFIG.dask_num_workers) - if CONFIG.dask_scheduler == "distributed": + config = PscPlotConfig.from_env() + + dask.config.set(num_workers=config.dask_num_workers) + if config.dask_scheduler == "distributed": from dask.distributed import Client, LocalCluster - cluster = LocalCluster(n_workers=CONFIG.dask_num_workers, threads_per_worker=1, processes=True) + cluster = LocalCluster(n_workers=config.dask_num_workers, threads_per_worker=1, processes=True) Client(cluster) - elif CONFIG.dask_scheduler: - dask.config.set(scheduler=CONFIG.dask_scheduler) - - args = parsing.get_parsed_args() - - if args.dask_graph: - if args.save_format is not None: - warnings.warn("--save-format is ignored with --dask-graph") - _run_dask_graph(args) - return - - # resolve format BEFORE applying pipeline in order to fail early - format = _resolve_save_format(args) - - if format == "mp4": - plt.rcParams["animation.ffmpeg_path"] = str(CONFIG.ffmpeg_bin) - - plot = args.get_animation() - - if args.show: - plot.show() - if args.save is not None: - args.save.mkdir(exist_ok=True, parents=True) + elif config.dask_scheduler: + dask.config.set(scheduler=config.dask_scheduler) - if format not in plot.allowed_save_formats(): - if format == args.save_format: # user actually specified this format - message = f"{format} is incompatible with the data; reverting to default ({plot.default_save_format()})" - warnings.warn(message) - else: - assert args.save_format is None + args = parse_args() - format = plot.default_save_format() + actions = compile_action_nodes(args, config) - path = args.save / f"{args.get_save_file_stem()}.{format}" - plot.save_to_path(path, dpi=args.save_dpi) - print(f"wrote to {path}") + for action in actions: + action.pull() diff --git a/src/lib/config.py b/src/lib/config.py index 820c664f..ae3e6910 100644 --- a/src/lib/config.py +++ b/src/lib/config.py @@ -1,6 +1,6 @@ import os import shutil -from dataclasses import dataclass +from dataclasses import KW_ONLY, dataclass, field from pathlib import Path from typing import Callable, Self @@ -19,32 +19,21 @@ def parse_optional[T](s: str | None, parser: Callable[[str], T]) -> T | None: @dataclass class PscPlotConfig: - data_dir: Path - ffmpeg_bin: Path | None - dask_num_workers: int - dask_chunk_size: int - dask_scheduler: str | None + _: KW_ONLY + data_dir: Path = field(default_factory=Path.cwd) + ffmpeg_bin: Path | None = None + dask_num_workers: int = 1 + dask_chunk_size: int = 1_000_000 + dask_scheduler: str | None = None @classmethod - def _load(cls) -> Self: - data_dir = parse_optional(os.environ.get(_DATA_DIR_KEY), Path) - if not data_dir: - message = f"Path to data not specified. Set the {_DATA_DIR_KEY} environment variable to specify." - raise RuntimeError(message) + def from_env(cls) -> Self: + config = cls() - ffmpeg_bin = parse_optional(os.environ.get(_FFMPEG_BIN_KEY, shutil.which("ffmpeg")), Path) + config.data_dir = parse_optional(os.environ.get(_DATA_DIR_KEY), Path) or config.data_dir + config.ffmpeg_bin = parse_optional(os.environ.get(_FFMPEG_BIN_KEY, shutil.which("ffmpeg")), Path) or config.ffmpeg_bin + config.dask_num_workers = parse_optional(os.environ.get(_DASK_NUM_WORKERS_KEY), int) or os.cpu_count() or config.dask_num_workers + config.dask_chunk_size = parse_optional(os.environ.get(_DASK_CHUNK_SIZE_KEY), int) or config.dask_chunk_size + config.dask_scheduler = os.environ.get(_DASK_SCHEDULER_KEY) or config.dask_scheduler - dask_num_workers = parse_optional(os.environ.get(_DASK_NUM_WORKERS_KEY), int) - if not dask_num_workers: - dask_num_workers = os.cpu_count() or 1 - - dask_chunk_size = parse_optional(os.environ.get(_DASK_CHUNK_SIZE_KEY), int) - if not dask_chunk_size: - dask_chunk_size = 1_000_000 - - dask_scheduler = os.environ.get(_DASK_SCHEDULER_KEY) or None - - return cls(data_dir, ffmpeg_bin, dask_num_workers, dask_chunk_size, dask_scheduler) - - -CONFIG = PscPlotConfig._load() + return config diff --git a/src/lib/data/adaptor.py b/src/lib/data/adaptor.py index 07ddcbe2..801575f5 100644 --- a/src/lib/data/adaptor.py +++ b/src/lib/data/adaptor.py @@ -1,10 +1,13 @@ from __future__ import annotations +from abc import ABC, abstractmethod + import dask.dataframe as dd import pandas as pd import xarray as xr from lib.data.data_with_attrs import DataWithAttrs, Field, List, Metadata +from lib.data.data_world import DataWorld from lib.has_name_fragments import HasNameFragments from lib.latex import Latex @@ -19,14 +22,22 @@ def _fail_apply_list(adaptor_type: type[Adaptor]): raise RuntimeError(message) -class Adaptor(HasNameFragments): +class WorldAdaptor(ABC, HasNameFragments): + @abstractmethod + def apply_world(self, world: DataWorld) -> DataWorld: ... + + +class Adaptor(WorldAdaptor): + def apply_world(self, world: DataWorld) -> DataWorld: + return world.with_active_data(self.apply(world.active_data)) + def apply(self, data: DataWithAttrs) -> DataWithAttrs: if isinstance(data, List): return self.apply_list(data) elif isinstance(data, Field): return self.apply_field(data) else: - message = f"unrecognized data type: {data.__class__:r}" + message = f"unrecognized data type: {data.__class__!r}" raise Exception(message) def apply_list(self, data: List) -> DataWithAttrs: diff --git a/src/lib/data/adaptors/bin.py b/src/lib/data/adaptors/bin.py index 30885b5e..a75b9695 100644 --- a/src/lib/data/adaptors/bin.py +++ b/src/lib/data/adaptors/bin.py @@ -10,7 +10,6 @@ from lib.data.data_with_attrs import Field, FieldMetadata, List from lib.parsing import parse_util from lib.parsing.args_registry import arg_parser -from lib.var_info import VarInfo def _guess_bin_edgess(data: List, varname_to_nbins: dict[str, int | None]) -> list: @@ -81,7 +80,6 @@ def __init__(self, varname_to_nbins: dict[str, int | None]): self.varname_to_nbins = varname_to_nbins def apply_field(self, data: Field) -> Field: - da = data.data dim_names_to_bin_size = {} for dim_name, nbins in self.varname_to_nbins.items(): if not nbins: @@ -95,8 +93,7 @@ def apply_field(self, data: Field) -> Field: dim_names_to_bin_size[dim_name] = bin_size - da = da.coarsen(dim_names_to_bin_size, boundary="pad").mean() - return data.assign_data(da) + return data.with_active_data(data.active_data.coarsen(dim_names_to_bin_size, boundary="pad").mean()) def apply_list(self, data: List) -> Field: bin_edgess = _guess_bin_edgess(data, self.varname_to_nbins) @@ -137,7 +134,7 @@ def apply_list(self, data: List) -> Field: # want: psc-plot prt.i --derive K="ux^2+uy^2+uz^2" --bin K=128 -v K --scale log new_var_infos["f"] = f_dim - return Field(da.to_dataset(name="f"), FieldMetadata.create_from(data.metadata, active_key="f", var_infos=new_var_infos)) + return Field({"f": da}, FieldMetadata.create_from(data.metadata, active_key="f", var_infos=new_var_infos)) def get_name_fragments(self) -> list[str]: subfrags = "_".join(f"{varname}={nbins}" if nbins else varname for varname, nbins in self.varname_to_nbins.items()) @@ -163,7 +160,7 @@ def parse_bin(args: list[str]) -> Bin: if len(split_arg) == 2 and not split_arg[1]: # arg is "t=", i.e., disable implicit binning along t - parse_util.check_value(split_arg[0], "active_key", ["t"]) + parse_util.parse_value(split_arg[0], "active_key", ["t"]) insert_bin_t = False continue elif len(split_arg) > 2: @@ -171,7 +168,7 @@ def parse_bin(args: list[str]) -> Bin: [active_key, nbins_arg, *_] = split_arg + [""] - parse_util.check_identifier(active_key, "active_key") + parse_util.parse_identifier(active_key, "active_key") nbins = parse_util.parse_optional_number(nbins_arg, "nbins", int) varname_to_nbins[active_key] = nbins diff --git a/src/lib/data/adaptors/compute.py b/src/lib/data/adaptors/compute.py index 01bc7871..1fc252df 100644 --- a/src/lib/data/adaptors/compute.py +++ b/src/lib/data/adaptors/compute.py @@ -13,4 +13,4 @@ def apply_list(self, data: List) -> FullList: return data.compute() def apply_field(self, data: Field) -> Field: - return data.map_data(lambda da: da.compute()) + return data.with_active_data(data.active_data.compute()) diff --git a/src/lib/data/adaptors/derive.py b/src/lib/data/adaptors/derive.py index 4e49b286..dc0c5eb5 100644 --- a/src/lib/data/adaptors/derive.py +++ b/src/lib/data/adaptors/derive.py @@ -96,7 +96,7 @@ def variable(self, toks: list): [tok] = toks name = str(tok) ds = self._data.data - if name not in ds.variables: + if name not in ds: self._resolve_from_registry(name) return self._data.data[name] @@ -129,7 +129,7 @@ def assign_default(self, toks: list): def assignment(self, toks: list): [new_variable, val] = toks - new_ds = self._data.data.assign({new_variable: val}) + new_ds = self._data.data | {new_variable: val} dim = var_info_registry.lookup(self._data.metadata.prefix, new_variable) new_var_infos = {**self._data.metadata.var_infos, new_variable: dim} return self._data.assign(new_ds, active_key=new_variable, var_infos=new_var_infos) diff --git a/src/lib/data/adaptors/diff.py b/src/lib/data/adaptors/diff.py index 940549e1..d212199f 100644 --- a/src/lib/data/adaptors/diff.py +++ b/src/lib/data/adaptors/diff.py @@ -81,11 +81,11 @@ def parse(args: list[str]) -> Diff: dims_arg, dir_arg = parse_util.parse_assignment(arg, DIFF_FORMAT) - parse_util.check_value(dir_arg, "dir", DIR_TO_SHIFT.keys()) + parse_util.parse_value(dir_arg, "dir", DIR_TO_SHIFT.keys()) dir = DIR_TO_SHIFT[dir_arg] for dim in parse_util.parse_comma_separated_list(dims_arg): - parse_util.check_identifier(dim, "dim_key") + parse_util.parse_identifier(dim, "dim_key") diffs_1d.append(_Diff1d(dim, dir, boundary)) return Diff(diffs_1d) diff --git a/src/lib/data/adaptors/downsample.py b/src/lib/data/adaptors/downsample.py index eb34031a..7eeb314f 100644 --- a/src/lib/data/adaptors/downsample.py +++ b/src/lib/data/adaptors/downsample.py @@ -33,7 +33,7 @@ def parse(args: list[str]) -> Downsample: for arg in args: [dim_name, bin_size_arg] = parse_util.parse_assignment(arg, DOWNSAMPLE_FORMAT) - parse_util.check_identifier(dim_name, "dim_name") + parse_util.parse_identifier(dim_name, "dim_name") bin_size = parse_util.parse_number(bin_size_arg, BIN_SIZE, int) dim_names_to_bin_size[dim_name] = bin_size diff --git a/src/lib/data/adaptors/fourier.py b/src/lib/data/adaptors/fourier.py index 128d675e..60019c9b 100644 --- a/src/lib/data/adaptors/fourier.py +++ b/src/lib/data/adaptors/fourier.py @@ -9,20 +9,20 @@ from lib.var_info import VarInfo -def toggle_fourier(da: xr.DataArray, dim: VarInfo) -> xr.DataArray: +def toggle_fourier(da: xr.DataArray, info: VarInfo) -> xr.DataArray: temp_prefix = "temp_" - f_dim = dim.toggle_fourier() + f_info = info.toggle_fourier() # multiply/or divide coords by 2pi to go from frequency <-> angular frequency - if dim.is_fourier(): - da = da.assign_coords({dim.key: da.coords[dim.key] / (2 * np.pi)}) - da = xrft.ifft(da, dim=dim.key, prefix=temp_prefix, lag=0.0) - da = da.rename({temp_prefix + dim.key: f_dim.key}) + if info.is_fourier(): + da = da.assign_coords({info.key: da.coords[info.key] / (2 * np.pi)}) + da = xrft.ifft(da, dim=info.key, prefix=temp_prefix, lag=0.0) + da = da.rename({temp_prefix + info.key: f_info.key}) else: - da = xrft.fft(da, dim=dim.key, prefix=temp_prefix) - da = da.rename({temp_prefix + dim.key: f_dim.key}) - da = da.assign_coords({f_dim.key: da.coords[f_dim.key] * (2 * np.pi)}) + da = xrft.fft(da, dim=info.key, prefix=temp_prefix) + da = da.rename({temp_prefix + info.key: f_info.key}) + da = da.assign_coords({f_info.key: da.coords[f_info.key] * (2 * np.pi)}) return da @@ -37,18 +37,17 @@ def apply_field(self, data: Field) -> Field: pre_dim_latexs = [data.metadata.var_infos[key].display.latex for key in self.dim_keys] da = data.active_data - new_var_infos = dict(data.metadata.var_infos) + new_var_infos = data.metadata.var_infos.copy() + for key in self.dim_keys: - dim = new_var_infos[key] - f_dim = dim.toggle_fourier() - da = toggle_fourier(da, dim) - del new_var_infos[key] - new_var_infos[f_dim.key] = f_dim + info = new_var_infos[key] + f_info = info.toggle_fourier() + new_var_infos[f_info.key] = f_info + da = toggle_fourier(da, info) - if data.metadata.active_key is not None and data.metadata.active_key in new_var_infos: - old_active = new_var_infos[data.metadata.active_key] - new_display = f"\\mathcal{{F}}_{{{','.join(pre_dim_latexs)}}}[{old_active.display}]" - new_var_infos[data.metadata.active_key] = old_active.assign(display=new_display) + old_active_info = data.metadata.active_var_info + new_display = f"\\mathcal{{F}}_{{{','.join(pre_dim_latexs)}}}[{old_active_info.display}]" + new_var_infos[data.metadata.active_key] = old_active_info.assign(display=new_display) return data.with_active_data(da).assign_metadata(var_infos=new_var_infos) @@ -68,6 +67,6 @@ def get_name_fragments(self) -> list[str]: ) def parse_fourier(args: list[str]) -> Fourier: for dim_name in args: - parse_util.check_identifier(dim_name, "dim_name") + parse_util.parse_identifier(dim_name, "dim_name") return Fourier(args) diff --git a/src/lib/data/adaptors/idx.py b/src/lib/data/adaptors/idx.py index ce707786..2c087057 100644 --- a/src/lib/data/adaptors/idx.py +++ b/src/lib/data/adaptors/idx.py @@ -10,7 +10,7 @@ def __init__(self, dim_names_to_isel: dict[str, int | slice]): self.dim_names_to_isel = dim_names_to_isel def apply_field(self, data: Field) -> Field: - return data.assign_data(data.data.isel(self.dim_names_to_isel)) + return data.with_active_data(data.active_data.isel(self.dim_names_to_isel)) def apply_list(self, data: List) -> List: coordss = data.coordss.copy() @@ -28,11 +28,11 @@ def apply_list(self, data: List) -> List: selected_steps = [selected_steps] partition_indices = [p for step in selected_steps for p in range(*data.metadata.partition_ranges[step])] df = df.partitions[partition_indices] - coordss[dim] = coordss[dim][isel] if isinstance(isel, slice) else float(coordss[dim][isel]) + coordss[dim] = coordss[dim][isel] if isinstance(isel, slice) else coordss[dim][isel] continue if isinstance(isel, int): - pos = float(coordss[dim][isel]) + pos = coordss[dim][isel] df = df[df[dim] == pos] if len(df) == 0: import warnings @@ -42,11 +42,11 @@ def apply_list(self, data: List) -> List: coordss[dim] = pos else: if isel.start not in [None, 0]: - pos_lower = float(coordss[dim][isel.start]) + pos_lower = coordss[dim][isel.start] df = df[df[dim] >= pos_lower] if isel.stop is not None: - pos_upper = float(coordss[dim][isel.stop]) + pos_upper = coordss[dim][isel.stop] df = df[df[dim] < pos_upper] coordss[dim] = coordss[dim][isel] @@ -73,7 +73,7 @@ def parse_idx(args: list[str]) -> Idx: for arg in args: [dim_name, isel_arg] = parse_util.parse_assignment(arg, IDX_FORMAT) - parse_util.check_identifier(dim_name, "dim_name") + parse_util.parse_identifier(dim_name, "dim_name") if ":" in isel_arg: dim_names_to_isel[dim_name] = parse_util.parse_slice(isel_arg, int) else: diff --git a/src/lib/data/adaptors/pos.py b/src/lib/data/adaptors/pos.py index 56350253..cfdf2043 100644 --- a/src/lib/data/adaptors/pos.py +++ b/src/lib/data/adaptors/pos.py @@ -33,7 +33,7 @@ def __init__( def apply_field(self, data: Field) -> Field: dim_names_to_pos = {dim_name: pos for dim_name, pos in self.dim_names_to_sel.items() if isinstance(pos, float)} dim_names_to_slice = {dim_name: s for dim_name, s in self.dim_names_to_sel.items() if isinstance(s, slice)} - return data.assign_data(data.data.sel(dim_names_to_pos, method="nearest").sel(dim_names_to_slice)) + return data.with_active_data(data.active_data.sel(dim_names_to_pos, method="nearest").sel(dim_names_to_slice)) def apply_list(self, data: List) -> List: # Lazy-import Idx to avoid a circular import via lib.plotting.animated_plot. @@ -84,7 +84,7 @@ def parse_pos(args: list[str]) -> Pos: for arg in args: [dim_name, sel_arg] = parse_util.parse_assignment(arg, POS_FORMAT) - parse_util.check_identifier(dim_name, "dim_name") + parse_util.parse_identifier(dim_name, "dim_name") if ":" in sel_arg: dim_names_to_sel[dim_name] = parse_util.parse_slice(sel_arg, float) else: diff --git a/src/lib/data/adaptors/quantile.py b/src/lib/data/adaptors/quantile.py index 7af0b8f4..2b10a596 100644 --- a/src/lib/data/adaptors/quantile.py +++ b/src/lib/data/adaptors/quantile.py @@ -49,7 +49,7 @@ def parse_quantile(args: list[str]) -> Quantile: for arg in args: [dim_name, sel_arg] = parse_util.parse_assignment(arg, QUANTILE_FORMAT) - parse_util.check_identifier(dim_name, "dim_name") + parse_util.parse_identifier(dim_name, "dim_name") dim_names_to_quants[dim_name] = parse_util.parse_slice(sel_arg, float) return Quantile(dim_names_to_quants) diff --git a/src/lib/data/adaptors/recenter.py b/src/lib/data/adaptors/recenter.py index e94fef43..bcfccb41 100644 --- a/src/lib/data/adaptors/recenter.py +++ b/src/lib/data/adaptors/recenter.py @@ -68,11 +68,11 @@ def parse(args: list[str]) -> Recenter: dims_arg, interp_dir_arg = parse_util.parse_assignment(arg, RECENTER_FORMAT) - parse_util.check_value(interp_dir_arg, "dir", DIR_TO_SHIFT.keys()) + parse_util.parse_value(interp_dir_arg, "dir", DIR_TO_SHIFT.keys()) interp_dir = DIR_TO_SHIFT[interp_dir_arg] for dim in parse_util.parse_comma_separated_list(dims_arg): - parse_util.check_identifier(dim, "dim_name") + parse_util.parse_identifier(dim, "dim_name") specs.append((dim, interp_dir, boundary)) return Recenter(specs) diff --git a/src/lib/data/adaptors/reduce.py b/src/lib/data/adaptors/reduce.py index 3d402432..88dee022 100644 --- a/src/lib/data/adaptors/reduce.py +++ b/src/lib/data/adaptors/reduce.py @@ -74,8 +74,8 @@ def parse_reduce(arg: str) -> Reduce: dim_names = parse_util.parse_comma_separated_list(dim_names_arg) for dim_name in dim_names: - parse_util.check_identifier(dim_name, "dim_name") + parse_util.parse_identifier(dim_name, "dim_name") - parse_util.check_value(func_name, "reduce_func", REDUCE_FUNCS) + parse_util.parse_value(func_name, "reduce_func", REDUCE_FUNCS) return Reduce(dim_names, func_name) diff --git a/src/lib/data/adaptors/roll.py b/src/lib/data/adaptors/roll.py index 389dface..7ed2077f 100644 --- a/src/lib/data/adaptors/roll.py +++ b/src/lib/data/adaptors/roll.py @@ -32,7 +32,7 @@ def parse(args: list[str]) -> Roll: for arg in args: [dim_name, window_size_arg] = parse_util.parse_assignment(arg, ROLL_FORMAT) - parse_util.check_identifier(dim_name, "dim_name") + parse_util.parse_identifier(dim_name, "dim_name") window_size = parse_util.parse_number(window_size_arg, "window_size", int) dim_names_to_window_size[dim_name] = window_size diff --git a/src/lib/data/adaptors/set_scale.py b/src/lib/data/adaptors/set_scale.py new file mode 100644 index 00000000..fffbdd10 --- /dev/null +++ b/src/lib/data/adaptors/set_scale.py @@ -0,0 +1,49 @@ +from dataclasses import replace + +from lib.data.adaptor import MetadataAdaptor +from lib.data.data_with_attrs import DataWithAttrs +from lib.parsing import parse_util +from lib.parsing.args_registry import arg_parser +from lib.scale import SCALE_TYPES, Scale + + +class SetScale(MetadataAdaptor): + def __init__(self, dim_name: str | None, scale: Scale): + self.dim_name = dim_name + self.scale = scale + + def apply(self, data: DataWithAttrs) -> DataWithAttrs: + dim_name = self.dim_name or data.metadata.active_key + new_var_infos = data.metadata.var_infos.copy() + new_var_infos[dim_name] = replace(new_var_infos[dim_name], scale=self.scale) + return data.assign_metadata(var_infos=new_var_infos) + + def get_name_fragments(self) -> list[str]: + maybe_dim_name = f"{self.dim_name}=" if self.dim_name is not None else "" + return [f"scale_{maybe_dim_name}{self.scale.to_name_fragment_part()}"] + + +ANY_SCALE_ARGS_FORMAT = "{" + ",".join(scale_type.to_argparse_format() for scale_type in SCALE_TYPES) + "}" +SCALE_FORMAT = f"[var_key=]{ANY_SCALE_ARGS_FORMAT}" + + +@arg_parser( + flags="--scale", + metavar=SCALE_FORMAT, + help="set the axis scale or color normalization of the given quantity (default is 'linear')", + dest="adaptors", +) +def parse_scale(arg: str) -> Scale: + if "=" in arg: + var_key, scale_arg = parse_util.parse_assignment(arg, SCALE_FORMAT) + parse_util.parse_identifier(var_key, "var_key") + else: + var_key = None + scale_arg = arg + + for scale_type in SCALE_TYPES: + maybe_scale = scale_type.try_from_argparse_format(scale_arg) + if maybe_scale: + return SetScale(var_key, maybe_scale) + + parse_util.fail_format(scale_arg, ANY_SCALE_ARGS_FORMAT) diff --git a/src/lib/data/adaptors/transform_polar.py b/src/lib/data/adaptors/transform_polar.py index 3aaa84cb..968c3a49 100644 --- a/src/lib/data/adaptors/transform_polar.py +++ b/src/lib/data/adaptors/transform_polar.py @@ -8,11 +8,11 @@ from lib.latex import Latex from lib.parsing import parse_util from lib.parsing.args_registry import arg_parser -from lib.var_info import RADIAN, VarInfo, check_unit_compatability +from lib.var_info import RADIAN, VarInfo, check_unit_compatibility def _build_polar_dims(dim_x: VarInfo, dim_y: VarInfo) -> tuple[VarInfo, VarInfo]: - check_unit_compatability(dim_x, dim_y, "polar") + check_unit_compatibility(dim_x, dim_y, "polar") r_symbol = "k" if dim_x.is_fourier() else "r" dim_r = VarInfo(Latex(f"{r_symbol}_\\text{{polar}}"), dim_x.unit, "polar:r", key=f"{r_symbol}_p") dim_theta = VarInfo(Latex("\\theta"), RADIAN, "polar:theta") @@ -111,7 +111,7 @@ def get_name_fragments(self) -> list[str]: ) def parse_transform_polar(args: list[str]) -> TransformPolar: for i, arg in enumerate(args, start=1): - parse_util.check_identifier(arg, f"dim_{i}") + parse_util.parse_identifier(arg, f"dim_{i}") try: return TransformPolar(args[0], args[1]) except ValueError as e: diff --git a/src/lib/data/adaptors/transform_spherical.py b/src/lib/data/adaptors/transform_spherical.py index c9c07b26..e597f4d7 100644 --- a/src/lib/data/adaptors/transform_spherical.py +++ b/src/lib/data/adaptors/transform_spherical.py @@ -8,12 +8,12 @@ from lib.latex import Latex from lib.parsing import parse_util from lib.parsing.args_registry import arg_parser -from lib.var_info import RADIAN, VarInfo, check_unit_compatability +from lib.var_info import RADIAN, VarInfo, check_unit_compatibility def _build_spherical_dims(dim_x: VarInfo, dim_y: VarInfo, dim_z: VarInfo) -> tuple[VarInfo, VarInfo, VarInfo]: - check_unit_compatability(dim_x, dim_y, "spherical") - check_unit_compatability(dim_x, dim_z, "spherical") + check_unit_compatibility(dim_x, dim_y, "spherical") + check_unit_compatibility(dim_x, dim_z, "spherical") r_symbol = "k" if dim_x.is_fourier() else "r" dim_r = VarInfo(Latex(f"{r_symbol}_\\text{{spherical}}"), dim_x.unit, "spherical:r", key=f"{r_symbol}_s") dim_theta = VarInfo(Latex("\\theta"), RADIAN, "spherical:theta") @@ -132,7 +132,7 @@ def get_name_fragments(self) -> list[str]: ) def parse_transform_spherical(args: list[str]) -> TransformSpherical: for i, arg in enumerate(args, start=1): - parse_util.check_identifier(arg, f"dim_{i}") + parse_util.parse_identifier(arg, f"dim_{i}") try: return TransformSpherical(args[0], args[1], args[2]) except ValueError as e: diff --git a/src/lib/data/adaptors/versus.py b/src/lib/data/adaptors/versus.py index 819f2f6c..b2690317 100644 --- a/src/lib/data/adaptors/versus.py +++ b/src/lib/data/adaptors/versus.py @@ -1,9 +1,11 @@ +from dataclasses import replace from typing import Literal from lib.data.adaptor import MetadataAdaptor from lib.data.adaptors.fourier import Fourier from lib.data.adaptors.reduce import Reduce from lib.data.data_with_attrs import DataWithAttrs, Field, List +from lib.data.plot_target import PlotTarget, SpatialDims, SpatialDimsRTheta, SpatialDimsXY from lib.parsing import parse_util from lib.parsing.args_registry import arg_parser @@ -15,10 +17,12 @@ def __init__( *, time_dim_rule: str | None | Literal["guess"], color_dim: str | None, + axes_idx: tuple[int, int] = (1, 1), ): self.spatial_dims = spatial_dims self.time_dim_rule = time_dim_rule self.color_dim = color_dim + self.axes_idx = axes_idx def _get_retained_dim_keys(self, data: DataWithAttrs) -> list[str]: retained_dims = self.spatial_dims.copy() @@ -43,6 +47,43 @@ def _get_time_dim(self, data: DataWithAttrs) -> str | None: return None + def _get_spatial_dims(self, data: DataWithAttrs) -> SpatialDims: + spatial_dims = self.spatial_dims.copy() + + if len(spatial_dims) == 1 and data.metadata.active_key: + spatial_dims.append(data.metadata.active_key) + + if data.metadata.var_infos[spatial_dims[0]].geometry == "polar:r" and data.metadata.var_infos[spatial_dims[1]].geometry == "polar:theta": + return SpatialDimsRTheta(*spatial_dims) + + return SpatialDimsXY(*spatial_dims) + + def _get_color_dim(self, data: DataWithAttrs) -> str | None: + if isinstance(data, Field): + if self.color_dim: + message = "Can't set color dim of field data" + raise ValueError(message) + if len(self.spatial_dims) == 2: + return data.metadata.active_key + return None + + return self.color_dim + + def apply_world(self, world): + data = self.apply(world.active_data) + new_plot_target = PlotTarget( + world.active_key, + spatial_dims=self._get_spatial_dims(data), + color_dim=self._get_color_dim(data), + time_dim=self._get_time_dim(data), + axes_index=self.axes_idx, + ) + return replace( + world, + plot_targets=world.plot_targets + [new_plot_target], + datas=world.datas | {world.active_key: data}, + ) + def apply_field(self, data: Field) -> Field: # 1. apply implicit coordinate transforms, as necessary retained_dims = self._get_retained_dim_keys(data) @@ -70,11 +111,7 @@ def apply_field(self, data: Field) -> Field: reduce = Reduce(reduce_dims, "mean") data = reduce.apply(data) - return data.assign_metadata( - spatial_dims=self.spatial_dims.copy(), - time_dim=self._get_time_dim(data), - color_dim=self.color_dim, - ) + return data def apply_list(self, data: List) -> List: # 1. coordinate transform @@ -89,11 +126,7 @@ def apply_list(self, data: List) -> List: if len(spatial_dims) == 1 and data.metadata.active_key is not None and data.metadata.active_key not in spatial_dims: spatial_dims.append(data.metadata.active_key) - return data.assign_metadata( - spatial_dims=spatial_dims, - time_dim=self._get_time_dim(data), - color_dim=self.color_dim, - ) + return data def get_name_fragments(self) -> list[str]: dims = ",".join(self.spatial_dims) @@ -106,29 +139,38 @@ def get_name_fragments(self) -> list[str]: _TIME_PREFIX = "time=" _COLOR_PREFIX = "color=" -_VERSUS_FORMAT = f"dim_key | {_TIME_PREFIX}[dim_key] | {_COLOR_PREFIX}dim_key" +_AXES_IDX_PREFIX = "loc=" +_AXES_IDX_FORMAT = f"{_AXES_IDX_PREFIX}i,j" +_VERSUS_FORMAT = f"dim_key | {_TIME_PREFIX}[dim_key] | {_COLOR_PREFIX}dim_key | {_AXES_IDX_FORMAT}" @arg_parser( dest="adaptors", flags=["--versus", "-v"], metavar=_VERSUS_FORMAT, - help=f"Specifies the independent axes of the plot. Remaining dimensions are reduced via arithmetic mean. Time has a special behavior: if {_TIME_PREFIX}[dim_key] is omitted, it is set to 't' if 't' is present in the data and isn't being used as a different axis. Disable this guessing by passing {_TIME_PREFIX} (with no dim_key).", + help=f"Specifies the independent axes of the plot. Remaining dimensions are reduced via arithmetic mean. Time has a special behavior: if {_TIME_PREFIX}[dim_key] is omitted, it is set to 't' if 't' is present in the data and isn't being used as a different axis. Disable this guessing by passing '{_TIME_PREFIX}' (with no dim_key). The special '{_AXES_IDX_FORMAT}' argument sets the 1-indexed location of the subplot in the figure grid, which is 1,1 by default.", nargs="+", ) def parse_versus(args: list[str]) -> Versus: spatial_dims = [] time_dim_rule = "guess" color_dim = None + axes_idx = 1, 1 for arg in args: if arg.startswith(_TIME_PREFIX): time_dim_rule = arg.removeprefix(_TIME_PREFIX) or None - parse_util.check_optional_identifier(time_dim_rule, "time dim_key") + parse_util.parse_optional_identifier(time_dim_rule, "time dim_key") elif arg.startswith(_COLOR_PREFIX): color_dim = arg.removeprefix(_COLOR_PREFIX) - parse_util.check_identifier(color_dim, "color dim_key") + parse_util.parse_identifier(color_dim, "color dim_key") + elif arg.startswith(_AXES_IDX_PREFIX): + axes_idx_arg = arg.removeprefix(_AXES_IDX_PREFIX) + i_arg, j_arg = parse_util.parse_assignment(axes_idx_arg, "i,j", delim=",") # bit of a hack + i = parse_util.parse_number(i_arg, "i", int) + j = parse_util.parse_number(j_arg, "j", int) + axes_idx = i, j else: - parse_util.check_identifier(arg, "dim_key") + parse_util.parse_identifier(arg, "dim_key") spatial_dims.append(arg) - return Versus(spatial_dims, time_dim_rule=time_dim_rule, color_dim=color_dim) + return Versus(spatial_dims, time_dim_rule=time_dim_rule, color_dim=color_dim, axes_idx=axes_idx) diff --git a/src/lib/data/adaptors/window.py b/src/lib/data/adaptors/window.py index e932cfe7..d4270a00 100644 --- a/src/lib/data/adaptors/window.py +++ b/src/lib/data/adaptors/window.py @@ -61,7 +61,7 @@ def get_name_fragment_fragment(self) -> str: def parse_window(arg: str) -> Window: [dim_name, beta] = parse_util.parse_assignment(arg, KAISER_FORMAT) - parse_util.check_identifier(dim_name, "dim_name") + parse_util.parse_identifier(dim_name, "dim_name") beta = parse_util.parse_number(beta, "beta", float) return Kaiser(dim_name, beta) diff --git a/src/lib/data/adaptors/with_.py b/src/lib/data/adaptors/with_.py index c7f66a6c..01d5e9be 100644 --- a/src/lib/data/adaptors/with_.py +++ b/src/lib/data/adaptors/with_.py @@ -1,25 +1,55 @@ -from lib.data.adaptor import Adaptor -from lib.data.data_with_attrs import DataWithAttrs +from lib.data.adaptor import WorldAdaptor +from lib.data.loader import get_loader +from lib.parsing import parse_util from lib.parsing.args_registry import arg_parser -class With(Adaptor): - def __init__(self, key: str): +class With(WorldAdaptor): + def __init__(self, prefix_or_key: str, key: str | None = None): + self.prefix_or_key = prefix_or_key self.key = key - def apply(self, data: DataWithAttrs) -> DataWithAttrs: - return data.assign_metadata(active_key=self.key) + def apply_world(self, world): + # case 1: prefix_or_key is a key within the active prefix + if not self.key and world.active_key and self.prefix_or_key in world.active_data.metadata.var_infos: + key = self.prefix_or_key + return world.with_active_data(world.active_data.assign_metadata(active_key=key)) + + # case 2: prefix_or_key is a prefix + prefix = self.prefix_or_key + key = self.key + + if prefix in world.datas: + return world.with_active_data(world.active_data.assign_metadata(active_key=key), prefix) + + loader = get_loader(world.config.data_dir, prefix, key) + return loader.apply_world(world) def get_name_fragments(self) -> list[str]: - return [f"with_{self.key}"] + maybe_prefix = f"{self.prefix_or_key}{SCOPE_OP}" if self.prefix_or_key else "" + return [f"with_{maybe_prefix}{self.key or ''}"] + + +SCOPE_OP = "::" +WITH_FORMAT = f"prefix[{SCOPE_OP}[key]] | key" @arg_parser( dest="adaptors", flags=["--with", "-w"], - metavar="key", - help="set the active variable to the given key", + metavar=WITH_FORMAT, + help="switch to a different prefix and/or variable", nargs="just one", ) def parse_with(arg: str) -> With: - return With(arg) + split_arg = arg.split(SCOPE_OP) + + if len(split_arg) == 2: + prefix = parse_util.parse_identifier(split_arg[0], "prefix") + key = parse_util.parse_optional_identifier(split_arg[1] or None, "key") + return With(prefix, key) + elif len(split_arg) == 1: + prefix_or_key = parse_util.parse_identifier(split_arg[0], "prefix | key") + return With(prefix_or_key) + else: + parse_util.fail_format(arg, WITH_FORMAT) diff --git a/src/lib/data/compatability.py b/src/lib/data/compatability.py deleted file mode 100644 index 879a52ab..00000000 --- a/src/lib/data/compatability.py +++ /dev/null @@ -1,31 +0,0 @@ -from types import GenericAlias, UnionType -from typing import Any, TypeAliasType, _LiteralGenericAlias, get_args, get_origin - - -def isinstance2(val: Any, typelike: Any) -> bool: - # TODO use PEP 747's TypeForm[T] and make this a TypeGuard[T] - if isinstance(typelike, type): - return isinstance(val, typelike) - - elif isinstance(typelike, UnionType): - return any(isinstance2(val, t) for t in get_args(typelike)) - - elif isinstance(typelike, GenericAlias): - origin = get_origin(typelike) - args = get_args(typelike) - - if origin is list and len(args) == 1: - (elem_type,) = args - return isinstance(val, list) and all(isinstance2(elem, elem_type) for elem in val) - - raise NotImplementedError(f"Unsupported generic: {typelike!r}") - - elif isinstance(typelike, TypeAliasType): - return isinstance2(val, typelike.__value__) - - elif isinstance(typelike, _LiteralGenericAlias): - return val in get_args(typelike) - - else: - message = f"Unsupported type expression: {typelike!r}, of class {typelike.__class__!r}" - raise NotImplementedError(message) diff --git a/src/lib/data/compile.py b/src/lib/data/compile.py index a4281490..a5cc8200 100644 --- a/src/lib/data/compile.py +++ b/src/lib/data/compile.py @@ -1,17 +1,69 @@ +import sys + +from lib.config import PscPlotConfig from lib.data.adaptor import Adaptor from lib.data.adaptors.versus import Versus -from lib.data.data_source import DataSource, DataSourceWithPipeline -from lib.data.pipeline import Pipeline +from lib.data.loader import get_loader +from lib.data.node import AdaptorNode, DaskGraphNode, DataProcessingNode, PlotNode, RootNode, SavePlotNode, ShowPlotNode +from lib.parsing.args import Args -def compile_source(loader: DataSource, adaptors: list[Adaptor]) -> DataSource: +def _with_versus(adaptors: list[Adaptor]) -> list[Adaptor]: + adaptors = adaptors.copy() for adaptor in adaptors: if isinstance(adaptor, Versus): break else: adaptors.append(Versus(["y", "z"], time_dim_rule="guess", color_dim=None)) + return adaptors + + +def compile_data_node(args: Args, config: PscPlotConfig): + node = RootNode(config) + + node = AdaptorNode(node, get_loader(config.data_dir, args.prefix, args.variable)) + + for adaptor in _with_versus(args.adaptors): + node = AdaptorNode(node, adaptor) + + return node + + +def compile_plot_node(args: Args, config: PscPlotConfig) -> PlotNode: + node = compile_data_node(args, config) + + node = PlotNode(node, args.hooks) + + return node + + +def compile_action_nodes(args: Args, config: PscPlotConfig) -> list[DataProcessingNode[None]]: + plot_node = compile_plot_node(args, config) + action_nodes = [] + + if args.dask_graph: + action_nodes.append(DaskGraphNode(plot_node.input_node, save_dir=args.save, show=args.show)) + return action_nodes + + if args.show: + action_nodes.append(ShowPlotNode(plot_node)) + + if args.save is None and args.save_format: + print("error: --save-format requires --save", file=sys.stderr) + sys.exit(1) + + if args.save_format == "mp4" and not config.ffmpeg_bin: + print("error: --save-format mp4 requires ffmpeg", file=sys.stderr) + sys.exit(1) - pipeline = Pipeline(*adaptors) - source = DataSourceWithPipeline(loader, pipeline) + if args.save is not None: + action_nodes.append( + SavePlotNode( + plot_node, + save_dir=args.save, + save_format=args.save_format, + save_dpi=args.save_dpi, + ) + ) - return source + return action_nodes diff --git a/src/lib/data/data_source.py b/src/lib/data/data_source.py deleted file mode 100644 index 0ac71b8e..00000000 --- a/src/lib/data/data_source.py +++ /dev/null @@ -1,22 +0,0 @@ -import typing -from abc import ABC, abstractmethod - -from lib.data.data_with_attrs import DataWithAttrs - -from .pipeline import Pipeline - - -class DataSource(ABC): - @abstractmethod - def get_data(self) -> DataWithAttrs: ... - - -class DataSourceWithPipeline(DataSource): - def __init__(self, source: DataSource, pipeline: Pipeline): - self.source = source - self.pipeline = pipeline - - def get_data(self) -> typing.Any: - da = self.source.get_data() - da = self.pipeline.apply(da) - return da diff --git a/src/lib/data/data_with_attrs.py b/src/lib/data/data_with_attrs.py index 6af3e15f..7944b09a 100644 --- a/src/lib/data/data_with_attrs.py +++ b/src/lib/data/data_with_attrs.py @@ -20,10 +20,6 @@ class Metadata: active_key: str | None = None - spatial_dims: list[str] = field(default_factory=list) - time_dim: str | None = None - color_dim: str | None = None - var_infos: dict[str, VarInfo] = field(default_factory=dict) species: dict[str, SpeciesInfo] = field(default_factory=dict) @@ -62,7 +58,7 @@ def assign(self, **vals: Any) -> Self: @dataclass(frozen=True, init=False) -class DataWithAttrs[D: xr.DataArray | pd.DataFrame | dd.DataFrame, MD: Metadata](ABC): +class DataWithAttrs[D: dict[str, xr.DataArray] | pd.DataFrame | dd.DataFrame, MD: Metadata](ABC): """A data wrapper to provide a uniform, typed, and reliable metadata interface.""" # The type checker ignores type bounds when no generic argument is present, e.g. after `isinstance` (and function parameters). @@ -70,7 +66,7 @@ class DataWithAttrs[D: xr.DataArray | pd.DataFrame | dd.DataFrame, MD: Metadata] # Thus, it's necessary to annotate __init__ parameters via generics and the fields themselves with concrete types. # Unfortunately, annotating a field in a superclass requires also annotating it in each subclass that refines that field's type. # And with all this, other methods still don't get type hints :( - data: xr.DataArray | pd.DataFrame | dd.DataFrame + data: dict[str, xr.DataArray] | pd.DataFrame | dd.DataFrame metadata: Metadata _caches: dict[str, dict[str, Any]] @@ -119,8 +115,8 @@ class FieldMetadata(Metadata): prefix: str | None = None -class Field(DataWithAttrs[xr.Dataset, FieldMetadata]): - data: xr.Dataset +class Field(DataWithAttrs[dict[str, xr.DataArray], FieldMetadata]): + data: dict[str, xr.DataArray] metadata: FieldMetadata @property @@ -130,19 +126,8 @@ def active_data(self) -> xr.DataArray: return self.data[self.metadata.active_key] def with_active_data(self, new_da: xr.DataArray) -> Self: - """Returns a copy with the active variable replaced by `new_da`. Sibling variables that - are no longer compatible with the new active grid (e.g. share a dim name with different - coordinate values) are dropped.""" - active_key = self.metadata.active_key - new_ds = new_da.to_dataset(name=active_key) - for sib in self.data.data_vars: - if sib == active_key: - continue - try: - new_ds = xr.merge([new_ds, self.data[[sib]]], join="exact", compat="no_conflicts") - except (xr.MergeError, ValueError): - pass - return self.assign_data(new_ds) + """Returns a shallow copy with the active variable replaced by `new_da`.""" + return self.assign_data(self.data | {self.metadata.active_key: new_da}) @cached_property def coordss(self) -> dict[str, np.ndarray]: @@ -170,7 +155,7 @@ def var_bounds(self) -> tuple[float, float]: return dask.compute(np.min(active), np.max(active)) def dask_collections(self) -> list: - return [da.data for da in self.data.data_vars.values() if dask.is_dask_collection(da.data)] + return [da.data for da in self.data.values() if dask.is_dask_collection(da.data)] @dataclass(kw_only=True, frozen=True) diff --git a/src/lib/data/data_world.py b/src/lib/data/data_world.py new file mode 100644 index 00000000..68507030 --- /dev/null +++ b/src/lib/data/data_world.py @@ -0,0 +1,41 @@ +from __future__ import annotations + +from dataclasses import KW_ONLY, dataclass, field, replace + +from lib.config import PscPlotConfig +from lib.data.data_with_attrs import DataWithAttrs +from lib.data.plot_target import PlotTarget + + +@dataclass(frozen=True) +class DataWorld: + # TODO python 3.15: make frozendict + datas: dict[str, DataWithAttrs] = field(default_factory=dict) + active_key: str | None = None + _: KW_ONLY + plot_targets: list[PlotTarget] = field(default_factory=list) + config: PscPlotConfig = field(default_factory=PscPlotConfig.from_env) + + def __post_init__(self): + assert self.active_key is None or self.active_key in self.datas + + @property + def active_data(self) -> DataWithAttrs | None: + if self.active_key is None: + return None + return self.datas[self.active_key] + + def with_active_data( + self, + active_data: DataWithAttrs | None = None, + active_key: str | None = None, + ) -> DataWorld: + if active_data is None: + return replace(self, active_key=active_key) + + active_key = active_key or self.active_key + assert active_key is not None + + new_datas = self.datas.copy() + new_datas[active_key] = active_data + return replace(self, datas=new_datas, active_key=active_key) diff --git a/src/lib/data/loader.py b/src/lib/data/loader.py index 939ace14..cd5432c9 100644 --- a/src/lib/data/loader.py +++ b/src/lib/data/loader.py @@ -2,12 +2,12 @@ from abc import abstractmethod from pathlib import Path -from lib.data.data_source import DataSource -from lib.file_util import get_available_steps -from lib.has_name_fragments import HasNameFragments +from lib.config import PscPlotConfig +from lib.data.adaptor import WorldAdaptor +from lib.data.data_with_attrs import DataWithAttrs -class Loader(DataSource, HasNameFragments): +class Loader(WorldAdaptor): @classmethod @abstractmethod def discover_prefixes(cls, data_dir: Path) -> list[str]: @@ -20,7 +20,6 @@ def suffix(cls) -> str: def __init__(self, prefix: str, active_key: str | None = None): self.prefix = prefix - self.steps = get_available_steps(prefix + ".", "." + self.suffix()) self.active_key = active_key def get_name_fragments(self) -> list[str]: @@ -29,6 +28,12 @@ def get_name_fragments(self) -> list[str]: fragments.append(self.active_key) return fragments + def apply_world(self, world): + return world.with_active_data(self.get_data(world.config), self.prefix) + + @abstractmethod + def get_data(self, config: PscPlotConfig) -> DataWithAttrs: ... + LOADERS: list[type[Loader]] = [] @@ -55,3 +60,8 @@ def discover_loaders(data_dir: Path) -> dict[str, type[Loader]]: ) result[prefix] = cls return result + + +def get_loader(data_dir: Path, prefix: str, active_key: str | None) -> Loader: + loader_types = discover_loaders(data_dir) + return loader_types[prefix](prefix, active_key) diff --git a/src/lib/data/loaders/field_bp.py b/src/lib/data/loaders/field_bp.py index 30d4d744..c6b6117e 100644 --- a/src/lib/data/loaders/field_bp.py +++ b/src/lib/data/loaders/field_bp.py @@ -4,7 +4,8 @@ import pscpy import xarray as xr -from lib.config import CONFIG +from lib import file_util +from lib.config import PscPlotConfig from lib.data.data_with_attrs import Field, FieldMetadata from lib.data.loader import Loader, loader from lib.derived_field_variables import derive_field_variable @@ -14,8 +15,8 @@ _STEP_BP_RE = re.compile(r"^(.+?)\.\d+\.bp$") -def _get_path(prefix: str, step: int) -> Path: - return CONFIG.data_dir / f"{prefix}.{step:09}.bp" +def _get_path(data_dir: Path, prefix: str, step: int) -> Path: + return data_dir / f"{prefix}.{step:09}.bp" def _decode_psc(ds): @@ -33,20 +34,28 @@ def discover_prefixes(cls, data_dir: Path) -> list[str]: def suffix(cls): return "bp" - def get_data(self) -> Field: + def get_data(self, config: PscPlotConfig) -> Field: ds = xr.open_mfdataset( - paths=[_get_path(self.prefix, step) for step in self.steps], + paths=[_get_path(config.data_dir, self.prefix, step) for step in file_util.get_available_steps(config.data_dir, self.prefix + ".", ".bp")], combine="nested", concat_dim="t", preprocess=_decode_psc, parallel=True, ) - if self.active_key is not None: - derive_field_variable(ds, self.active_key, self.prefix) - var_info = {key: lookup(self.prefix, key) for key in ds.variables} - metadata = FieldMetadata( - active_key=self.active_key, - prefix=self.prefix, - var_infos=var_info, + + data = {key: ds[key] for key in ds.data_vars} + var_infos = {key: lookup(self.prefix, key) for key in ds.variables} + + field = Field( + data, + FieldMetadata( + active_key=self.active_key, + prefix=self.prefix, + var_infos=var_infos, + ), ) - return Field(ds, metadata) + + if self.active_key is not None: + field = derive_field_variable(field, self.active_key, self.prefix) + + return field diff --git a/src/lib/data/loaders/particle_bp.py b/src/lib/data/loaders/particle_bp.py index 85645aed..1908e419 100644 --- a/src/lib/data/loaders/particle_bp.py +++ b/src/lib/data/loaders/particle_bp.py @@ -6,7 +6,8 @@ import pandas as pd import xarray as xr -from lib.config import CONFIG +from lib import file_util +from lib.config import PscPlotConfig from lib.data.data_with_attrs import LazyList, ListMetadata from lib.data.loader import Loader, loader from lib.species import SpeciesInfo, build_species_display @@ -15,8 +16,8 @@ _DISCOVER_PARTICLE_BP_PREFIX_RE = re.compile(r"^prt\.([^.]+)\.\d+\.bp$") -def _get_path(prefix: str, step: int) -> pathlib.Path: - return CONFIG.data_dir / f"{prefix}.{step:09}.bp" +def _get_path(data_dir: pathlib.Path, prefix: str, step: int) -> pathlib.Path: + return data_dir / f"{prefix}.{step:09}.bp" def _read_attrs(path: pathlib.Path) -> dict: @@ -90,8 +91,9 @@ def __init__(self, prefix: str, active_key: str | None = None): super().__init__(prefix, active_key) self.species_key = prefix.split(".", 1)[1] - def get_data(self) -> LazyList: - step_attrs = [_read_attrs(_get_path(self.prefix, step)) for step in self.steps] + def get_data(self, config: PscPlotConfig) -> LazyList: + steps = file_util.get_available_steps(config.data_dir, self.prefix + ".", ".bp") + step_attrs = [_read_attrs(_get_path(config.data_dir, self.prefix, step)) for step in steps] times = np.array([float(a["time"]) for a in step_attrs]) head = step_attrs[0] @@ -116,15 +118,15 @@ def get_data(self) -> LazyList: # along its particle dim. dd.from_map propagates downstream column # projection into `_read_chunk` via its `columns` kwarg, so unused # variables are never read from disk. - chunk_size = CONFIG.dask_chunk_size + chunk_size = config.dask_chunk_size paths: list[pathlib.Path] = [] step_times: list[float] = [] particle_dims: list[str] = [] slices: list[slice] = [] partition_ranges = [] offset = 0 - for step, time in zip(self.steps, times): - path = _get_path(self.prefix, step) + for step, time in zip(steps, times): + path = _get_path(config.data_dir, self.prefix, step) particle_dim, n = _peek_size(path) n_chunks = max(1, (n + chunk_size - 1) // chunk_size) partition_ranges.append((offset, offset + n_chunks)) diff --git a/src/lib/data/loaders/particle_h5.py b/src/lib/data/loaders/particle_h5.py index 77ffa33d..b89b4789 100644 --- a/src/lib/data/loaders/particle_h5.py +++ b/src/lib/data/loaders/particle_h5.py @@ -8,7 +8,8 @@ import h5py import numpy as np -from lib.config import CONFIG +from lib import file_util +from lib.config import PscPlotConfig from lib.data.data_with_attrs import LazyList, ListMetadata from lib.data.loader import Loader, loader from lib.latex import Latex @@ -22,12 +23,12 @@ type Mass = float -def _get_path_at_step(prefix: str, step: int) -> pathlib.Path: - return CONFIG.data_dir / f"{prefix}.{step:09}.h5" +def _get_path_at_step(data_dir: pathlib.Path, prefix: str, step: int) -> pathlib.Path: + return data_dir / f"{prefix}.{step:09}.h5" -def _load_attrs_at_step(prefix: str, step: int) -> dict[str, typing.Any]: - data_path = _get_path_at_step(prefix, step) +def _load_attrs_at_step(data_dir: pathlib.Path, prefix: str, step: int) -> dict[str, typing.Any]: + data_path = _get_path_at_step(data_dir, prefix, step) attrs = {} with h5py.File(data_path) as file: if "time" not in file.keys(): @@ -50,10 +51,10 @@ def _find_first_populated_cell(idx_begin_s: np.ndarray, idx_end_s: np.ndarray) - return int(idx_begin_s[mask].min()) -def _read_species_qm(prefix: str, step: int, missing: set[SpeciesIdx]) -> dict[SpeciesIdx, tuple[Charge, Mass]]: +def _read_species_qm(data_dir: pathlib.Path, prefix: str, step: int, missing: set[SpeciesIdx]) -> dict[SpeciesIdx, tuple[Charge, Mass]]: """Open one step and return {species_index: (q, m)} for any species with particles in that step. Only populates entries for species-indices in `missing`; others are left untouched.""" - with h5py.File(_get_path_at_step(prefix, step)) as f: + with h5py.File(_get_path_at_step(data_dir, prefix, step)) as f: idx_begin = f["particles/idx_begin"][...] idx_end = f["particles/idx_end"][...] particles = f[PRT_PARTICLES_KEY] @@ -67,11 +68,11 @@ def _read_species_qm(prefix: str, step: int, missing: set[SpeciesIdx]) -> dict[S return found -def _discover_species_qm(prefix: str, steps: list[int]) -> dict[SpeciesIdx, tuple[Charge, Mass]]: +def _discover_species_qm(data_dir: pathlib.Path, prefix: str, steps: list[int]) -> dict[SpeciesIdx, tuple[Charge, Mass]]: """For each species index in [0, n_species), find a step where it has particles and read its (q, m). Tries step 0 first, then the last step, then bisects the remaining range. Raises if any species never appears.""" - with h5py.File(_get_path_at_step(prefix, steps[0])) as f: + with h5py.File(_get_path_at_step(data_dir, prefix, steps[0])) as f: n_species = f["particles/idx_begin"].shape[0] qm: dict[SpeciesIdx, tuple[Charge, Mass]] = {} missing = set(range(n_species)) @@ -86,7 +87,7 @@ def _discover_species_qm(prefix: str, steps: list[int]) -> dict[SpeciesIdx, tupl if step in probed: continue probed.add(step) - found = _read_species_qm(prefix, step, missing) + found = _read_species_qm(data_dir, prefix, step, missing) qm.update(found) missing -= set(found.keys()) if not missing: @@ -176,16 +177,17 @@ def discover_prefixes(cls, data_dir: pathlib.Path) -> list[str]: def suffix(cls): return "h5" - def get_data(self) -> LazyList: - species_dict = _build_species_dict(_discover_species_qm(self.prefix, self.steps)) + def get_data(self, config: PscPlotConfig) -> LazyList: + steps = file_util.get_available_steps(config.data_dir, self.prefix + ".", ".h5") + species_dict = _build_species_dict(_discover_species_qm(config.data_dir, self.prefix, steps)) - attrss = [_load_attrs_at_step(self.prefix, step) for step in self.steps] + attrss = [_load_attrs_at_step(config.data_dir, self.prefix, step) for step in steps] times = np.array([attrs["time"] for attrs in attrss]) - data_paths = [_get_path_at_step(self.prefix, step) for step in self.steps] + data_paths = [_get_path_at_step(config.data_dir, self.prefix, step) for step in steps] dfs_of_steps = [] for time, data_path in zip(times, data_paths): - df_of_step: dd.DataFrame = dd.read_hdf(data_path, key=PRT_PARTICLES_KEY, chunksize=CONFIG.dask_chunk_size, lock=True) + df_of_step: dd.DataFrame = dd.read_hdf(data_path, key=PRT_PARTICLES_KEY, chunksize=config.dask_chunk_size, lock=True) df_of_step = df_of_step.assign(t=time) dfs_of_steps.append(df_of_step) diff --git a/src/lib/data/node.py b/src/lib/data/node.py new file mode 100644 index 00000000..16d545d6 --- /dev/null +++ b/src/lib/data/node.py @@ -0,0 +1,152 @@ +import sys +import warnings +from abc import ABC, abstractmethod +from functools import cache +from pathlib import Path + +from lib.config import PscPlotConfig +from lib.data.adaptor import Adaptor +from lib.data.data_world import DataWorld +from lib.plotting.get_plot import get_plot +from lib.plotting.hook import Hook +from lib.plotting.plot import Plot, SaveFormat + + +class DataProcessingNode[D](ABC): + def __init__(self, name_fragments: list[str]): + self.name_fragments = name_fragments + + @abstractmethod + def pull(self) -> D: ... + + def get_save_file_stem(self) -> str: + return "-".join(self.name_fragments) + + +class AdaptorNode(DataProcessingNode[DataWorld]): + def __init__(self, input_node: DataProcessingNode[DataWorld], adaptor: Adaptor): + super().__init__(input_node.name_fragments + adaptor.get_name_fragments()) + self.input_node = input_node + self.adaptor = adaptor + + @cache + def pull(self) -> DataWorld: + return self.adaptor.apply_world(self.input_node.pull()) + + +class RootNode(DataProcessingNode[DataWorld]): + def __init__(self, config: PscPlotConfig): + super().__init__([]) + self.config = config + + def pull(self) -> DataWorld: + return DataWorld(config=self.config) + + +class PlotNode(DataProcessingNode[Plot]): + def __init__(self, input_node: DataProcessingNode[DataWorld], hooks: list[Hook]): + super().__init__(input_node.name_fragments + [frag for hook in hooks for frag in hook.get_name_fragments()]) + self.input_node = input_node + self.hooks = hooks + + @cache + def pull(self) -> Plot: + world = self.input_node.pull() + plot = get_plot(world) + + for hook in self.hooks: + plot.add_hook(hook) + + return plot + + +class ShowPlotNode(DataProcessingNode[None]): + def __init__(self, input_node: DataProcessingNode[Plot]): + super().__init__(input_node.name_fragments) + self.input_node = input_node + + def pull(self) -> None: + self.input_node.pull().show() + + +class SavePlotNode(DataProcessingNode[None]): + def __init__( + self, + input_node: DataProcessingNode[Plot], + *, + save_dir: Path, + save_format: SaveFormat | None, + save_dpi: float | None, + ): + super().__init__(input_node.name_fragments) + self.input_node = input_node + self.save_dir = save_dir + self.save_format = save_format + self.save_dpi = save_dpi + + def pull(self) -> None: + plot = self.input_node.pull() + + save_format = self.save_format + if save_format not in plot.allowed_save_formats(): + if save_format is not None: + message = f"{save_format} is incompatible with the data; reverting to default ({plot.default_save_format()})" + warnings.warn(message) + + save_format = plot.default_save_format() + + self.save_dir.mkdir(exist_ok=True, parents=True) + path = self.save_dir / f"{self.get_save_file_stem()}.{save_format}" + plot.save_to_path(path, dpi=self.save_dpi) + print(f"wrote to {path}") + + +class DaskGraphNode(DataProcessingNode[None]): + def __init__( + self, + input_node: DataProcessingNode[DataWorld], + *, + save_dir: Path | None, + show: bool, + ): + super().__init__(input_node.name_fragments) + self.input_node = input_node + self.save_dir = save_dir or Path.cwd() + self.show = show + + def pull(self) -> None: + data = self.input_node.pull().active_data + + collections = data.dask_collections() + if not collections: + print( + f"error: --dask-graph requires dask-backed data; pipeline produced eager {type(data).__name__}", + file=sys.stderr, + ) + sys.exit(1) + + try: + import graphviz # noqa: F401 + except ImportError: + print( + "error: --dask-graph requires the 'graphviz' package; install with `pip install -e \".[dask-graph]\"`", + file=sys.stderr, + ) + sys.exit(1) + + import dask + + self.save_dir.mkdir(exist_ok=True, parents=True) + path = self.save_dir / f"{self.get_save_file_stem()}.daskgraph.svg" + # dask.visualize's optimize_graph flag only lowers legacy HLG collections + # (e.g. dask Arrays), not new-style Expr ones (dask DataFrames) — without + # pre-optimizing the latter, un-lowered nodes (e.g. Concat from dd.concat) + # fail with NotImplementedError in _layer. + collections = [c.optimize() if hasattr(c, "optimize") else c for c in collections] + dask.visualize(*collections, filename=str(path), optimize_graph=True) + print(f"wrote to {path}") + + if self.show: + import webbrowser + + webbrowser.open(path.absolute().as_uri()) diff --git a/src/lib/data/pipeline.py b/src/lib/data/pipeline.py deleted file mode 100644 index 25439981..00000000 --- a/src/lib/data/pipeline.py +++ /dev/null @@ -1,11 +0,0 @@ -from .adaptor import Adaptor - - -class Pipeline(Adaptor): - def __init__(self, *adaptors: Adaptor): - self.adaptors = list(adaptors) - - def apply(self, data): - for adaptor in self.adaptors: - data = adaptor.apply(data) - return data diff --git a/src/lib/data/plot_target.py b/src/lib/data/plot_target.py new file mode 100644 index 00000000..930de61a --- /dev/null +++ b/src/lib/data/plot_target.py @@ -0,0 +1,43 @@ +from abc import ABC, abstractmethod +from dataclasses import KW_ONLY, dataclass, field + +type DimKey = str + + +@dataclass +class SpatialDims(ABC): + ndims: int + + @abstractmethod + def unpack(self) -> tuple[DimKey, DimKey]: ... + + +@dataclass +class SpatialDimsXY(SpatialDims): + x_dim: DimKey + y_dim: DimKey + ndims: int = field(default=2, init=False) + + def unpack(self): + return (self.x_dim, self.y_dim) + + +@dataclass +class SpatialDimsRTheta(SpatialDims): + r_dim: DimKey + theta_dim: DimKey + ndims: int = field(default=2, init=False) + + def unpack(self): + return (self.r_dim, self.theta_dim) + + +@dataclass +class PlotTarget: + prefix: str + _: KW_ONLY + spatial_dims: SpatialDims + color_dim: DimKey | None = None + time_dim: DimKey | None = None + + axes_index: tuple[int, int] = (1, 1) # 1-based diff --git a/src/lib/derived_field_variables/derived_field_variable.py b/src/lib/derived_field_variables/derived_field_variable.py index 233397af..887a63eb 100644 --- a/src/lib/derived_field_variables/derived_field_variable.py +++ b/src/lib/derived_field_variables/derived_field_variable.py @@ -3,6 +3,9 @@ import xarray as xr +from lib.data.data_with_attrs import Field +from lib.var_info_registry import lookup + __all__ = ["derived_field_variable", "derive_field_variable"] @@ -16,13 +19,18 @@ def __init__( name: str, base_var_names: list[str], derive: DeriveField, + prefix: str, ): self.name = name self.base_var_names = base_var_names self.derive = derive + self.prefix = prefix - def assign_to(self, ds: xr.Dataset): - ds[self.name] = self.derive(*(ds[base_var_name] for base_var_name in self.base_var_names)) + def assign_to(self, field: Field) -> Field: + da = self.derive(*(field.data[base_var_name] for base_var_name in self.base_var_names)) + new_data = field.data | {self.name: da} + new_var_infos = field.metadata.var_infos | {key: lookup(self.prefix, key) for key in (self.name, *da.dims)} + return field.assign_data(new_data).assign_metadata(var_infos=new_var_infos) def __repr__(self) -> str: return f"{self.__class__.__name__}(({', '.join(self.base_var_names)}) -> {self.name}: {self.derive!r})" @@ -39,22 +47,22 @@ def derived_field_variable(prefix: str): def derived_field_variable_inner[F: (function, DeriveField)](derive_func: F) -> F: name = derive_func.__name__ base_var_names = list(inspect.signature(derive_func).parameters) - register_derived_field_variable(prefix, DerivedFieldVariable(name, base_var_names, derive_func)) + register_derived_field_variable(prefix, DerivedFieldVariable(name, base_var_names, derive_func, prefix)) return derive_func return derived_field_variable_inner -def derive_field_variable(ds: xr.Dataset, active_key: str, ds_prefix: str): - if active_key in ds.variables: - return +def derive_field_variable(field: Field, active_key: str, ds_prefix: str) -> Field: + if active_key in field.data: + return field elif active_key in DERIVED_FIELD_VARIABLES[ds_prefix]: derived_var = DERIVED_FIELD_VARIABLES[ds_prefix][active_key] for base_var_name in derived_var.base_var_names: - derive_field_variable(ds, base_var_name, ds_prefix) - derived_var.assign_to(ds) + field = derive_field_variable(field, base_var_name, ds_prefix) + return derived_var.assign_to(field) else: message = f"""No variable named '{active_key}'. -The following variables are defined: {list(ds.variables)}. +The following variables are defined: {list(field.data)}. The following variables can be derived: {list(DERIVED_FIELD_VARIABLES[ds_prefix])}.""" raise ValueError(message) diff --git a/src/lib/file_util.py b/src/lib/file_util.py index 6278c4df..177f0287 100644 --- a/src/lib/file_util.py +++ b/src/lib/file_util.py @@ -1,12 +1,12 @@ -from lib.config import CONFIG +from pathlib import Path -def get_available_steps(before_step: str, after_step: str) -> list[int]: - files = CONFIG.data_dir.glob(f"{before_step}*{after_step}") +def get_available_steps(data_dir: Path, before_step: str, after_step: str) -> list[int]: + files = data_dir.glob(f"{before_step}*{after_step}") steps = [int(file.name.removeprefix(before_step).removesuffix(after_step)) for file in files] if not steps: - raise ValueError(f"No steps found matching {CONFIG.data_dir}/{before_step}*{after_step}") + raise ValueError(f"No steps found matching {data_dir}/{before_step}*{after_step}") steps.sort() return steps diff --git a/src/lib/latex.py b/src/lib/latex.py index f52f2d4a..c578e653 100644 --- a/src/lib/latex.py +++ b/src/lib/latex.py @@ -12,7 +12,7 @@ def strip_latex(latex: str) -> str: return plain -@dataclass(frozen=True) +@dataclass(frozen=True, unsafe_hash=True) class Latex: latex: str plain: str = field(init=False) @@ -41,6 +41,11 @@ def prepend(self, latex: str) -> Latex: def append(self, latex: str) -> Latex: return Latex(self.latex + latex) + def maybe_with_dollars(self) -> str: + if self: + return f"${self}$" + return "" + def __str__(self) -> str: return self.latex diff --git a/src/lib/parsing/__init__.py b/src/lib/parsing/__init__.py deleted file mode 100644 index 6ee7a585..00000000 --- a/src/lib/parsing/__init__.py +++ /dev/null @@ -1,3 +0,0 @@ -from .parse import get_parsed_args - -__all__ = ["get_parsed_args"] diff --git a/src/lib/parsing/args.py b/src/lib/parsing/args.py index 9ab6e49e..9d5bc907 100644 --- a/src/lib/parsing/args.py +++ b/src/lib/parsing/args.py @@ -2,17 +2,11 @@ from pathlib import Path from lib.data.adaptor import Adaptor -from lib.data.compile import compile_source -from lib.data.data_source import DataSource -from lib.data.data_with_attrs import DataWithAttrs -from lib.plotting.get_plot import get_plot from lib.plotting.hook import Hook -from lib.plotting.plot import Plot class Args(argparse.Namespace): prefix: str - loader: DataSource variable: str | None adaptors: list[Adaptor] hooks: list[Hook] @@ -21,22 +15,3 @@ class Args(argparse.Namespace): save_format: str | None save_dpi: float | None dask_graph: bool - - def get_data(self) -> DataWithAttrs: - source = compile_source(self.loader, self.adaptors) - return source.get_data() - - def get_animation(self) -> Plot: - data = self.get_data() - - plot = get_plot(data) - - for hook in self.hooks: - plot.add_hook(hook) - - return plot - - def get_save_file_stem(self) -> str: - sources = [self.loader, *self.adaptors, *self.hooks] - fragments = [frag for src in sources for frag in src.get_name_fragments()] - return "-".join(fragments) diff --git a/src/lib/parsing/parse.py b/src/lib/parsing/parse.py index 95c7f59c..52464b89 100644 --- a/src/lib/parsing/parse.py +++ b/src/lib/parsing/parse.py @@ -1,18 +1,15 @@ import argparse from pathlib import Path -from typing import Iterable -from lib.config import CONFIG -from lib.data.loader import discover_loaders from lib.parsing.args import Args from lib.parsing.args_registry import CUSTOM_ARGS -def _get_parser(prefixes: Iterable[str]) -> argparse.ArgumentParser: +def _get_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(prog="psc-plot") - parser.add_argument("prefix", choices=prefixes, help="data file prefix (auto-discovered from the data directory)") - parser.add_argument("variable", nargs="?", default=None, help="field variable to work with") + parser.add_argument("prefix", help="initial active prefix") + parser.add_argument("variable", nargs="?", default=None, help="initial active variable") parser.add_argument( "-s", "--save", @@ -49,9 +46,6 @@ def _get_parser(prefixes: Iterable[str]) -> argparse.ArgumentParser: return parser -def get_parsed_args(args_list: list[str] | None = None) -> Args: - prefix_to_loader = discover_loaders(CONFIG.data_dir) - parser = _get_parser(prefix_to_loader.keys()) - args = parser.parse_args(args_list, namespace=Args()) - args.loader = prefix_to_loader[args.prefix](args.prefix, active_key=args.variable) - return args +def parse_args(args_list: list[str] | None = None) -> Args: + parser = _get_parser() + return parser.parse_args(args_list, namespace=Args()) diff --git a/src/lib/parsing/parse_util.py b/src/lib/parsing/parse_util.py index 8b871a01..6a241d3b 100644 --- a/src/lib/parsing/parse_util.py +++ b/src/lib/parsing/parse_util.py @@ -4,26 +4,29 @@ def _is_identifier(val: str) -> bool: - return re.match(r"^\w[\d\w]*$", val) + return all(re.match(r"^\w[\d\w]*$", v) for v in val.split(".")) def fail_format(arg: str, format: str): raise argparse.ArgumentTypeError(f"Expected value of form '{format}'; got '{arg}'") -def check_value[T](val: T, val_name: str, valid_options: typing.Container[T]): +def parse_value[T](val: T, val_name: str, valid_options: typing.Container[T]) -> T: if val not in valid_options: raise argparse.ArgumentTypeError(f"Expected {val_name} to be one of {set(valid_options)}; got '{val}'") + return val -def check_identifier(val: str, val_name: str): +def parse_identifier(val: str, val_name: str) -> str: if not _is_identifier(val): raise argparse.ArgumentTypeError(f"Expected {val_name} to be an identifier; got '{val}'") + return val -def check_optional_identifier(val: str | None, val_name: str): +def parse_optional_identifier(val: str | None, val_name: str) -> str | None: if val and not _is_identifier(val): raise argparse.ArgumentTypeError(f"Expected {val_name} to be an identifier or ''; got '{val}'") + return val def check_order[T](lower: T | None, upper: T | None, lower_name: str, upper_name: str): diff --git a/src/lib/plotting/animated_plot.py b/src/lib/plotting/animated_plot.py index 5efda2b0..2e806968 100644 --- a/src/lib/plotting/animated_plot.py +++ b/src/lib/plotting/animated_plot.py @@ -3,11 +3,11 @@ import sys from pathlib import Path -import matplotlib.pyplot as plt from matplotlib.animation import FFMpegWriter, FuncAnimation, PillowWriter -from lib.data.adaptors.idx import Idx +from lib.config import PscPlotConfig from lib.data.data_with_attrs import DataWithAttrs +from lib.plotting.hook import DrawMessage from lib.plotting.plot import Plot, SaveFormat from lib.plotting.renderer import Renderer @@ -18,49 +18,38 @@ def print_progress(current_frame: int, n_frames: int): print(f"frame {current_frame_padded}/{n_frames}", end=end) -class AnimatedPlot[Data: DataWithAttrs](Plot[Data]): - def __init__(self, renderer: Renderer[Data], data: Data): - super().__init__(renderer, data) - self.time_dim: str = self.data.metadata.time_dim +class AnimatedPlot(Plot): + def __init__(self, renderers: list[Renderer[DataWithAttrs]], config: PscPlotConfig, n_frames: int): + super().__init__(renderers, config) + self.n_frames = n_frames - self.fig, self.ax = plt.subplots(subplot_kw=renderer.subplot_kw()) - self._initialized = False + def _initialize(self): + super()._initialize() # FIXME get blitting to work with the title - self.n_frames = len(data.coordss[data.metadata.time_dim]) self.anim = FuncAnimation(self.fig, self._next_frame, frames=self.n_frames, blit=False) - def _get_data_at_frame(self, frame: int) -> Data: - return Idx({self.time_dim: frame}).apply(self.data) - def _next_frame(self, frame: int): - frame_data = self._get_data_at_frame(frame) - update_data = self.renderer.make_update_data(self.ax, frame_data) - self.pre_update_fig(update_data) - self.renderer.draw(self.ax, frame_data, update_data) - self.post_update_fig(update_data) + for renderer in self.renderers: + renderer.update_plot_info(frame) + self.post_update_fig(DrawMessage(plot_info=self.renderers[0].plot_info, axes=self.fig.axes[0], frame_data=self.renderers[0]._get_data_at_frame(frame))) print_progress(frame, self.n_frames) - def _initialize(self): - if self._initialized: - return - self._initialized = True - - frame_0 = self._get_data_at_frame(0) - init_data = self.renderer.make_init_data(self.fig, self.ax, frame_0) - self.pre_init_fig(init_data) - self.renderer.init(self.fig, self.ax, self.data, frame_0, init_data) - self.post_init_fig(init_data) - self.fig.tight_layout() - - def show(self): - self._initialize() - plt.show() - def allowed_save_formats(self) -> list[SaveFormat]: - return ["mp4", "gif"] + if self.config.ffmpeg_bin: + return ["mp4", "gif"] + else: + return ["gif"] def save_to_path(self, path: Path, *, dpi: float | None = None): self._initialize() - writer = PillowWriter() if path.suffix == ".gif" else FFMpegWriter() + + if path.suffix == ".mp4": + from matplotlib import pyplot as plt + + plt.rcParams["animation.ffmpeg_path"] = str(self.config.ffmpeg_bin) + writer = FFMpegWriter() + else: + writer = PillowWriter() + self.anim.save(path, writer=writer, dpi=dpi) diff --git a/src/lib/plotting/frame_data_traits.py b/src/lib/plotting/frame_data_traits.py deleted file mode 100644 index e2e792e4..00000000 --- a/src/lib/plotting/frame_data_traits.py +++ /dev/null @@ -1,77 +0,0 @@ -from dataclasses import dataclass -from typing import Any, TypeGuard - -from matplotlib.axes import Axes - -from lib.data.compatability import isinstance2 -from lib.data.data_with_attrs import Field, FullList, List -from lib.plotting import plt_util -from lib.plotting.hook import Hook - - -def check_impl[T](data: Any, data_type: type[T]) -> TypeGuard[T]: - for superclass in data_type.mro(): - if not hasattr(superclass, "__annotations__"): - continue - - for field_name, field_type in superclass.__annotations__.items(): - if not hasattr(data, field_name): - return False - - if not isinstance2(getattr(data, field_name), field_type): - return False - - return True - - -def assert_impl[T](data: Any, required_type: type[T]) -> T: - if not check_impl(data, required_type): - raise TypeError("TODO better message") - return data - - -@dataclass(kw_only=True) -class HasData: - data: Field | List - - -@dataclass(kw_only=True) -class HasFieldData: - data: Field - - -@dataclass(kw_only=True) -class HasListData: - data: List - - -@dataclass(kw_only=True) -class HasFullListData: - data: FullList - - -@dataclass(kw_only=True) -class HasLineType: - line_type: str - - -@dataclass(kw_only=True) -class HasAxes: - axes: Axes - - -@dataclass(kw_only=True) -class HasSpatialScales: - spatial_scales: list[plt_util.AxisScaleArg] - last_spatial_dim_is_dependent: bool = False - - -@dataclass(kw_only=True) -class HasColorNorm: - color_norm: plt_util.ColorNormArg - color_is_dependent: bool = False - - -@dataclass(kw_only=True) -class HasHookList: - hooks: list[Hook] diff --git a/src/lib/plotting/get_plot.py b/src/lib/plotting/get_plot.py index 27e2e997..09d3bef2 100644 --- a/src/lib/plotting/get_plot.py +++ b/src/lib/plotting/get_plot.py @@ -1,4 +1,6 @@ -from lib.data.data_with_attrs import DataWithAttrs, Field, List +from lib.data.data_with_attrs import Field, List +from lib.data.data_world import DataWorld +from lib.data.plot_target import SpatialDimsRTheta, SpatialDimsXY from lib.plotting.animated_plot import AnimatedPlot from lib.plotting.plot import Plot from lib.plotting.renderer import Renderer @@ -9,32 +11,39 @@ from lib.plotting.static_plot import StaticPlot -def _get_renderer(data: DataWithAttrs) -> Renderer: - spatial_dims = data.metadata.spatial_dims +def get_plot(world: DataWorld) -> Plot: + renderers = get_renderers(world) + n_frames = max(r.get_n_frames() for r in renderers) - if isinstance(data, Field): - if len(spatial_dims) == 1: - return Field1dRenderer() - elif len(spatial_dims) == 2: - if data.metadata.var_infos[spatial_dims[0]].geometry == "polar:r" and data.metadata.var_infos[spatial_dims[1]].geometry == "polar:theta": - return PolarFieldRenderer() - return Field2dRenderer() - else: - raise NotImplementedError("don't have 3D field plots yet") + if n_frames > 1: + return AnimatedPlot(renderers, world.config, n_frames) + else: + return StaticPlot(renderers, world.config) - elif isinstance(data, List): - if len(spatial_dims) == 2: - return ScatterRenderer() - else: - raise NotImplementedError(f"don't have {len(spatial_dims)}D scatter plots yet") - raise TypeError(f"unexpected data type: {type(data)}") +def get_renderers(world: DataWorld) -> list[Renderer]: + renderers = [] + for target in world.plot_targets: + data = world.datas[target.prefix] -def get_plot(data: DataWithAttrs) -> Plot: - renderer = _get_renderer(data) + if isinstance(data, Field): + if not target.color_dim: + renderers.append(Field1dRenderer(data, target)) + elif isinstance(target.spatial_dims, SpatialDimsRTheta): + renderers.append(PolarFieldRenderer(data, target)) + elif isinstance(target.spatial_dims, SpatialDimsXY): + renderers.append(Field2dRenderer(data, target)) + else: + raise NotImplementedError("don't have 3D field plots yet") - if data.metadata.time_dim: - return AnimatedPlot(renderer, data) - else: - return StaticPlot(renderer, data) + elif isinstance(data, List): + if target.spatial_dims.ndims == 2: + renderers.append(ScatterRenderer(data, target)) + else: + raise NotImplementedError(f"don't have {target.spatial_dims.ndims}D scatter plots yet") + + else: + raise TypeError(f"unexpected data type: {type(data)}") + + return renderers diff --git a/src/lib/plotting/hook.py b/src/lib/plotting/hook.py index 9cb47497..5948f6a2 100644 --- a/src/lib/plotting/hook.py +++ b/src/lib/plotting/hook.py @@ -1,20 +1,22 @@ -from typing import Any +from dataclasses import dataclass +from matplotlib.axes import Axes + +from lib.data.data_with_attrs import DataWithAttrs from lib.has_name_fragments import HasNameFragments +from lib.plotting.plot_info import PlotInfo -class Hook(HasNameFragments): - def post_add_hook(self, add_data: Any): - pass +@dataclass(kw_only=True) +class DrawMessage: + plot_info: PlotInfo + axes: Axes + frame_data: DataWithAttrs - def pre_init_fig(self, init_data: Any): - pass - def post_init_fig(self, init_data: Any): - pass - - def pre_update_fig(self, update_data: Any): +class Hook(HasNameFragments): + def post_init_fig(self, message: DrawMessage): pass - def post_update_fig(self, update_data: Any): + def post_update_fig(self, message: DrawMessage): pass diff --git a/src/lib/plotting/hooks/fit.py b/src/lib/plotting/hooks/fit.py index 04aa9883..44c87bf2 100644 --- a/src/lib/plotting/hooks/fit.py +++ b/src/lib/plotting/hooks/fit.py @@ -5,46 +5,35 @@ from lib.data.data_with_attrs import DataWithAttrs, Field, List from lib.parsing import parse_util from lib.parsing.args_registry import arg_parser -from lib.plotting.frame_data_traits import ( - HasAxes, - HasData, - HasLineType, - assert_impl, - check_impl, -) from lib.plotting.hook import Hook +from lib.plotting.plot_info import LineInfo, ScatterInfo class Fit(Hook): - class InitData(HasData, HasAxes): ... - - class UpdateData(HasData, HasAxes): ... - def __init__(self, subdomain: slice): self.subdomain = subdomain - def pre_init_fig(self, init_data): - if check_impl(init_data, HasLineType) and init_data.line_type == "-": - init_data.line_type = "." + def post_init_fig(self, message): + if isinstance(message.plot_info, LineInfo) and message.plot_info.line_style == "-": + message.plot_info.set("line_style", ".") - def post_init_fig(self, init_data): - init_data = assert_impl(init_data, Fit.InitData) + assert isinstance(message.plot_info, (LineInfo, ScatterInfo)) - x_data, y_data = self._get_xy_data(init_data.data) + x_data, y_data = self._get_xy_data(message.frame_data, message.plot_info.x_dim, message.plot_info.y_dim) fit_y_data, label = self._get_fit_y_data(x_data, y_data) - [self.line] = init_data.axes.plot(x_data, fit_y_data, "--", label=label) + [self.line] = message.axes.plot(x_data, fit_y_data, "--", label=label, scalex=False, scaley=False) - init_data.axes.legend() + message.axes.legend() - def post_update_fig(self, update_data): - update_data = assert_impl(update_data, Fit.UpdateData) + def post_update_fig(self, message): + assert isinstance(message.plot_info, (LineInfo, ScatterInfo)) - x_data, y_data = self._get_xy_data(update_data.data) + x_data, y_data = self._get_xy_data(message.frame_data, message.plot_info.x_dim, message.plot_info.y_dim) fit_y_data, label = self._get_fit_y_data(x_data, y_data) self.line.set_data(x_data, fit_y_data) self.line.set_label(label) - update_data.axes.legend() # in case label changed + message.axes.legend() # in case label changed def _get_fit_y_data(self, x_data: np.ndarray, y_data: np.ndarray) -> tuple[np.ndarray, str]: x_log = np.log(x_data) @@ -59,14 +48,13 @@ def _get_fit_y_data(self, x_data: np.ndarray, y_data: np.ndarray) -> tuple[np.nd return y_fit, label - def _get_xy_data(self, data: DataWithAttrs) -> tuple[np.ndarray, np.ndarray]: - spatial_dim = data.metadata.spatial_dims[0] - slicer = Pos({spatial_dim: self.subdomain}) + def _get_xy_data(self, data: DataWithAttrs, x_dim: str, y_dim: str) -> tuple[np.ndarray, np.ndarray]: + slicer = Pos({x_dim: self.subdomain}) data = slicer.apply(data) if isinstance(data, Field): - return (data.coordss[spatial_dim], data.active_data) + return (data.coordss[x_dim], data.data[y_dim]) elif isinstance(data, List): - return (data.data[spatial_dim], data.data[data.metadata.spatial_dims[1]]) + return (data.data[x_dim], data.data[y_dim]) @arg_parser( diff --git a/src/lib/plotting/hooks/grid.py b/src/lib/plotting/hooks/grid.py index c00836ef..6644fbff 100644 --- a/src/lib/plotting/hooks/grid.py +++ b/src/lib/plotting/hooks/grid.py @@ -4,11 +4,10 @@ import numpy.typing as npt from matplotlib.axes import Axes -from lib.data.data_with_attrs import DataWithAttrs from lib.parsing import parse_util from lib.parsing.args_registry import arg_parser -from lib.plotting.frame_data_traits import HasAxes, HasData, assert_impl from lib.plotting.hook import Hook +from lib.plotting.plot_info import PlotInfo2D type DimName = str type MajorGridlineSpacing = float | None @@ -36,46 +35,45 @@ def __init__(self, major: MajorGridParams, minor: MinorGridParams): self.major = major self.minor = minor - class InitData(HasAxes, HasData): ... - - def post_init_fig(self, init_data): - init_data = assert_impl(init_data, Grid.InitData) - + def post_init_fig(self, message): # Ensure major lines are enabled when minor lines are present for dim_name in self.minor: self.major.setdefault(dim_name, None) + assert isinstance(message.plot_info, PlotInfo2D) + # Draw major lines with the given spacing for dim_name, gridline_spacing in self.major.items(): - axis_id = get_axis_id(init_data.data, dim_name) + axis_id = get_axis_id(message.plot_info, dim_name) if gridline_spacing is None: tick_coords = None else: - tick_coords = get_tick_coords(*init_data.data.bounds(dim_name), gridline_spacing) + tick_coords = get_tick_coords(*message.plot_info.dim_bounds[dim_name], gridline_spacing) - set_grid(init_data.axes, "major", axis_id, tick_coords) + set_grid(message.axes, "major", axis_id, tick_coords) # Draw the given number of minor lines between major lines for dim_name, nlines in self.minor.items(): - axis_id = get_axis_id(init_data.data, dim_name) + axis_id = get_axis_id(message.plot_info, dim_name) if nlines is None: tick_coords = None else: - get_ticks = {"x": init_data.axes.get_xticks, "y": init_data.axes.get_yticks}[axis_id] - major_tick_coords = np.concat([get_ticks(minor=False), [init_data.data.upper_bound(dim_name)]]) + get_ticks = {"x": message.axes.get_xticks, "y": message.axes.get_yticks}[axis_id] + major_tick_coords = np.concat([get_ticks(minor=False), [message.plot_info.dim_bounds[dim_name][1]]]) tick_coords = np.concat([np.linspace(left, right, nlines + 1, endpoint=False)[1:] for left, right in zip(major_tick_coords[:-1], major_tick_coords[1:])]) - set_grid(init_data.axes, "minor", axis_id, tick_coords) + set_grid(message.axes, "minor", axis_id, tick_coords) -def get_axis_id(data: DataWithAttrs, dim_name: str) -> Literal["x", "y"]: - if dim_name not in data.metadata.spatial_dims: - message = f"Dimension '{dim_name}' isn't being shown on an axis. Axis dimensions are: {data.metadata.spatial_dims}" +def get_axis_id(info: PlotInfo2D, dim_name: str) -> Literal["x", "y"]: + id_map = {info.x_dim: "x", info.y_dim: "y"} + if dim_name not in id_map: + message = f"Dimension '{dim_name}' isn't being shown on an axis. Axis dimensions are: {id_map.keys()}" raise ValueError(message) - return ["x", "y"][data.metadata.spatial_dims.index(dim_name)] + return id_map[dim_name] MAJOR_MARKER = "major" diff --git a/src/lib/plotting/hooks/scale.py b/src/lib/plotting/hooks/scale.py deleted file mode 100644 index d2ea0802..00000000 --- a/src/lib/plotting/hooks/scale.py +++ /dev/null @@ -1,161 +0,0 @@ -from typing import Literal, Self - -from matplotlib.colors import SymLogNorm -from matplotlib.scale import SymmetricalLogScale - -from lib.data.data_with_attrs import DataWithAttrs -from lib.parsing import parse_util -from lib.parsing.args_registry import arg_parser -from lib.plotting import plt_util -from lib.plotting.frame_data_traits import ( - HasColorNorm, - HasData, - HasSpatialScales, - assert_impl, - check_impl, -) -from lib.plotting.hook import Hook - -type ScaleKey = Literal["linear", "log", "symlog"] -SCALE_KEYS: tuple[ScaleKey, ...] = ScaleKey.__value__.__args__ - - -class Scale: - scale_key: ScaleKey - - def __init_subclass__(cls): - SCALE_TYPES.append(cls) - - def to_axis_scale(self, data: DataWithAttrs) -> plt_util.AxisScaleArg: - return self.scale_key - - def to_color_norm(self, data: DataWithAttrs) -> plt_util.ColorNormArg: - return self.scale_key - - def to_name_fragment_part(self) -> str: - return str(self.scale_key) - - @classmethod - def to_argparse_format(cls) -> str: - return cls.scale_key - - @classmethod - def try_from_argparse_format(cls, arg: str) -> Self | None: - if arg == cls.scale_key: - return cls() - return None - - -SCALE_TYPES: list[Scale] = [] # automatically populated with subclasses - - -class LinearScale(Scale): - scale_key = "linear" - - -class LogScale(Scale): - scale_key = "log" - - -class SymLogScale(Scale): - scale_key = "symlog" - LINEAR_THRESHOLD_ARG_FORMAT = "linear_threshold" - - def __init__(self, linear_threshold: float | None): - self.linear_threshold = linear_threshold - - def _choose_linear_threshold(self, data: DataWithAttrs) -> float: - # TODO pipe it through Magnitude -> Reduce, but reduce via quantile... which requires refactoring Reduce to support subargs - return 0.0001 - - def to_axis_scale(self, data: DataWithAttrs) -> plt_util.AxisScaleArg: - linthresh = self.linear_threshold or self._choose_linear_threshold(data) - return SymmetricalLogScale(None, linthresh=linthresh) - - def to_color_norm(self, data: DataWithAttrs) -> plt_util.ColorNormArg: - linthresh = self.linear_threshold or self._choose_linear_threshold(data) - return SymLogNorm(linthresh) - - def to_name_fragment_part(self) -> str: - if self.linear_threshold is None: - return self.scale_key - return f"{self.scale_key}{parse_util.SUBARG_DELIM}{self.linear_threshold}" - - @classmethod - def to_argparse_format(cls) -> str: - return f"{cls.scale_key}[{parse_util.SUBARG_DELIM}{cls.LINEAR_THRESHOLD_ARG_FORMAT}]" - - @classmethod - def try_from_argparse_format(cls, arg: str) -> Self | None: - scale_key_arg, linear_threshold_arg = parse_util.parse_optional_assignment(arg, cls.to_argparse_format(), delim=parse_util.SUBARG_DELIM) - - if scale_key_arg != cls.scale_key: - return None - - linear_threshold = parse_util.parse_optional_number(linear_threshold_arg, cls.LINEAR_THRESHOLD_ARG_FORMAT, float) - - return cls(linear_threshold) - - -class SetScale(Hook): - def __init__(self, dim_name: str | None, scale: Scale): - self.dim_name = dim_name - self.scale = scale - - def pre_init_fig(self, init_data): - init_data = assert_impl(init_data, HasData) - - data = init_data.data - - if self.dim_name is None: - # find and set the dependent scale/norm - if check_impl(init_data, HasSpatialScales) and init_data.last_spatial_dim_is_dependent: - init_data.spatial_scales[-1] = self.scale.to_axis_scale(data) - elif check_impl(init_data, HasColorNorm) and init_data.color_is_dependent: - init_data.color_norm = self.scale.to_color_norm(data) - else: - message = f"dependent scale not found" - raise Exception(message) - else: - spatial_dims = data.metadata.spatial_dims - color_dim = data.metadata.color_dim - - if self.dim_name in spatial_dims: - init_data = assert_impl(init_data, HasSpatialScales) - init_data.spatial_scales[spatial_dims.index(self.dim_name)] = self.scale.to_axis_scale(data) - elif self.dim_name == color_dim: - init_data = assert_impl(init_data, HasColorNorm) - init_data.color_norm = self.scale.to_color_norm(data) - else: - message = f"'{self.dim_name}' isn't a dimension" - raise Exception(message) - - def get_name_fragments(self) -> list[str]: - maybe_dim_name = f"{self.dim_name}=" if self.dim_name is not None else "" - return [f"scale_{maybe_dim_name}{self.scale.to_name_fragment_part()}"] - - -ANY_SCALE_ARGS_FORMAT = "{" + ",".join(scale_type.to_argparse_format() for scale_type in SCALE_TYPES) + "}" -SCALE_FORMAT = f"[dim_name=]{ANY_SCALE_ARGS_FORMAT}" - - -@arg_parser( - flags="--scale", - metavar=SCALE_FORMAT, - help="set the axis/color scale of the dependent variable or specified dimension (default: linear)", - dest="hooks", -) -def parse_scale(arg: str) -> Scale: - if "=" in arg: - dim_name, scale_arg = parse_util.parse_assignment(arg, SCALE_FORMAT) - parse_util.check_identifier(dim_name, "dim_name") - else: - dim_name = None - scale_arg = arg - - for scale_type in SCALE_TYPES: - maybe_scale = scale_type.try_from_argparse_format(scale_arg) - if maybe_scale: - return SetScale(dim_name, maybe_scale) - - parse_util.fail_format(scale_arg, ANY_SCALE_ARGS_FORMAT) diff --git a/src/lib/plotting/hooks/show_com.py b/src/lib/plotting/hooks/show_com.py index 606be165..3c499f52 100644 --- a/src/lib/plotting/hooks/show_com.py +++ b/src/lib/plotting/hooks/show_com.py @@ -1,7 +1,7 @@ from lib.data.data_with_attrs import Field from lib.parsing.args_registry import const_arg -from lib.plotting.frame_data_traits import HasAxes, HasFieldData, assert_impl from lib.plotting.hook import Hook +from lib.plotting.plot_info import PlotInfo2D def _get_center(field: Field, dim: str) -> float: @@ -15,23 +15,16 @@ def _get_centers(field: Field) -> list[float]: class ShowCom(Hook): - class PostInitData(HasFieldData, HasAxes): ... + def post_init_fig(self, message): + assert isinstance(message.plot_info, PlotInfo2D) - def post_init_fig(self, init_data): - init_data = assert_impl(init_data, ShowCom.PostInitData) + self.scatter = message.axes.scatter(*_get_centers(message.frame_data), marker="x", label="center of mass") + message.axes.legend() - data = init_data.data - assert len(data.dims) == 2 - self.scatter = init_data.axes.scatter(*_get_centers(data), marker="x", label="center of mass") + def post_update_fig(self, message): + assert isinstance(message.plot_info, PlotInfo2D) - init_data.axes.legend() - - class PostUpdateData(HasFieldData): ... - - def post_update_fig(self, update_data): - update_data = assert_impl(update_data, ShowCom.PostUpdateData) - - self.scatter.set_offsets([_get_centers(update_data.data)]) + self.scatter.set_offsets([_get_centers(message.frame_data)]) def get_name_fragments(self) -> list[str]: return ["show_com"] diff --git a/src/lib/plotting/hooks/show_initial.py b/src/lib/plotting/hooks/show_initial.py index c7ae7ca8..3d4b8846 100644 --- a/src/lib/plotting/hooks/show_initial.py +++ b/src/lib/plotting/hooks/show_initial.py @@ -1,28 +1,18 @@ -from lib.data.data_with_attrs import Field from lib.parsing.args_registry import const_arg -from lib.plotting.frame_data_traits import HasAxes, HasData, HasLineType, assert_impl from lib.plotting.hook import Hook +from lib.plotting.plot_info import LineInfo class ShowInitial(Hook): - class PreInitData(HasData, HasAxes, HasLineType): ... + def post_init_fig(self, message): + assert isinstance(message.plot_info, LineInfo) + assert message.plot_info.time_dim - def pre_init_fig(self, init_data): - init_data = assert_impl(init_data, ShowInitial.PreInitData) + message.axes.plot(message.plot_info.x_data, message.plot_info.y_data, "-", label=message.plot_info.get_coord_label(message.plot_info.time_dim)) + message.plot_info.set("line_style", "--") - data = init_data.data - xdata = data.coordss[data.dims[0]] - ydata = data.active_data if isinstance(data, Field) else data.data - time_dim = data.metadata.time_dim - init_data.axes.plot(xdata, ydata, "-", label=data.metadata.var_infos[time_dim].get_coordinate_label(data.coordss[time_dim])) - - init_data.line_type = "--" - - class PostInitData(HasAxes): ... - - def post_init_fig(self, init_data): - init_data = assert_impl(init_data, ShowInitial.PostInitData) - init_data.axes.legend() + def post_init_fig(self, message): + message.axes.legend() @const_arg( diff --git a/src/lib/plotting/hooks/vline.py b/src/lib/plotting/hooks/vline.py index ac06b825..929df623 100644 --- a/src/lib/plotting/hooks/vline.py +++ b/src/lib/plotting/hooks/vline.py @@ -1,6 +1,5 @@ from lib.parsing import parse_util from lib.parsing.args_registry import arg_parser -from lib.plotting.frame_data_traits import HasAxes, assert_impl from lib.plotting.hook import Hook @@ -9,29 +8,21 @@ def __init__(self, pos: float, label: str | None): self.pos = pos self.label = label - class InitData(HasAxes): ... + def post_init_fig(self, message): + if not self.label: + return - def pre_init_fig(self, init_data): - init_data = assert_impl(init_data, VLine.InitData) + lower_spine = message.axes.spines["bottom"] - lower_spine = init_data.axes.spines["bottom"] - init_data.axes.axvline( + message.axes.axvline( self.pos, color="lightgray", linewidth=lower_spine.get_linewidth(), ) - def post_init_fig(self, init_data): - if not self.label: - return - - init_data = assert_impl(init_data, VLine.InitData) - - lower_spine = init_data.axes.spines["bottom"] - - init_data.axes.text( + message.axes.text( self.pos, - init_data.axes.get_ybound()[0], + message.axes.get_ybound()[0], self.label, horizontalalignment="center", verticalalignment="baseline", diff --git a/src/lib/plotting/plot.py b/src/lib/plotting/plot.py index 4b2f4a19..f5d9ae99 100644 --- a/src/lib/plotting/plot.py +++ b/src/lib/plotting/plot.py @@ -2,34 +2,42 @@ from abc import ABC, abstractmethod from pathlib import Path -from typing import Any, Literal +from typing import Literal +from matplotlib import pyplot as plt + +from lib.config import PscPlotConfig from lib.data.data_with_attrs import DataWithAttrs -from lib.plotting.frame_data_traits import HasHookList -from lib.plotting.hook import Hook +from lib.plotting.hook import DrawMessage, Hook from lib.plotting.renderer import Renderer +from lib.plotting.setup_fig import setup_fig type SaveFormat = Literal["mp4", "gif", "png"] -class Plot[Data: DataWithAttrs](ABC): - class AddHookData(HasHookList): ... - - def __init__(self, renderer: Renderer[Data], data: Data): - self.renderer = renderer - self.data = data +class Plot(ABC): + def __init__(self, renderers: list[Renderer[DataWithAttrs]], config: PscPlotConfig): + self.renderers = renderers + self.config = config self.hooks: list[Hook] = [] - def add_hook(self, hook: Hook): - self.hooks.append(hook) + self._initialized = False - post_add_data = Plot.AddHookData(hooks=self.hooks) + def _initialize(self): + if self._initialized: + return + self._initialized = True - for hook in self.hooks.copy(): # hooks might reorder themselves - hook.post_add_hook(post_add_data) + self.fig = setup_fig([r.plot_info for r in self.renderers]) + # TODO hooks should be per-renderer; for now, just apply them to the 1st one + self.post_init_fig(DrawMessage(plot_info=self.renderers[0].plot_info, axes=self.fig.axes[0], frame_data=self.renderers[0]._get_data_at_frame(0))) - @abstractmethod - def show(self): ... + def add_hook(self, hook: Hook): + self.hooks.append(hook) + + def show(self): + self._initialize() + plt.show() @abstractmethod def save_to_path(self, path: Path, *, dpi: float | None = None): ... @@ -40,18 +48,10 @@ def allowed_save_formats(self) -> list[SaveFormat]: ... def default_save_format(self) -> SaveFormat: return self.allowed_save_formats()[0] - def pre_init_fig(self, init_data: Any): - for hook in self.hooks: - hook.pre_init_fig(init_data) - - def post_init_fig(self, init_data: Any): - for hook in self.hooks: - hook.post_init_fig(init_data) - - def pre_update_fig(self, update_data: Any): + def post_init_fig(self, message: DrawMessage): for hook in self.hooks: - hook.pre_update_fig(update_data) + hook.post_init_fig(message) - def post_update_fig(self, update_data: Any): + def post_update_fig(self, message: DrawMessage): for hook in self.hooks: - hook.post_update_fig(update_data) + hook.post_update_fig(message) diff --git a/src/lib/plotting/plot_info.py b/src/lib/plotting/plot_info.py new file mode 100644 index 00000000..4b763480 --- /dev/null +++ b/src/lib/plotting/plot_info.py @@ -0,0 +1,111 @@ +from dataclasses import KW_ONLY, dataclass, field +from typing import Any, Callable, Literal + +import numpy as np +from matplotlib.typing import LineStyleType + +from lib.latex import Latex +from lib.scale import Scale + +type DimKey = str +type AttrKey = str +type VarKey = str +type Projection = Literal["rectilinear", "polar"] + + +@dataclass +class PlotInfo: + _: KW_ONLY + subject: str | None = None + dim_scales: dict[DimKey, Scale] = field(default_factory=dict) + dim_bounds: dict[DimKey, tuple[float | None, float | None]] = field(default_factory=dict) + dim_displays: dict[DimKey, Latex] = field(default_factory=dict) + dim_units: dict[DimKey, Latex] = field(default_factory=dict) + time_dim: DimKey | None = None + scalar_coord_values: dict[DimKey, float] = field(default_factory=dict) + + axes_index: tuple[int, int] = (1, 1) + projection: Projection = field(default="rectilinear", init=False) + + _setter_callbacks: dict[AttrKey | tuple[AttrKey, DimKey], Callable[[Any], None]] = field(default_factory=dict, init=False) + + def set(self, key: AttrKey | tuple[AttrKey, DimKey], value: Any): + if isinstance(key, str): + attr_key = key + setattr(self, key, value) + else: + attr_key, dim_key = key + getattr(self, attr_key)[dim_key] = value + + if key in self._setter_callbacks: + self._setter_callbacks[key](value) + + if attr_key in self._setter_callbacks: + self._setter_callbacks[attr_key](value) + + def get_coord_label(self, dim: DimKey) -> Latex: + display = self.dim_displays.get(dim, f"\\text{{{dim}}}") + coord_val = self.scalar_coord_values[dim] + unit = self.dim_units.get(dim, "") + maybe_space = "\\ " if unit else "" + return Latex(f"{display} = {coord_val:.3f}{maybe_space}{unit}") + + def get_title(self) -> str: + coord_labels_str = ", ".join(f"${self.get_coord_label(dim)}$" for dim in self.scalar_coord_values) + + if self.subject and coord_labels_str: + return f"{self.subject} ({coord_labels_str})" + elif self.subject: + return self.subject + else: + return coord_labels_str + + def get_dim_label(self, dim: DimKey) -> str: + dim_label = f"${self.dim_displays.get(dim, f'\\text{{{dim}}}')}$" + + if unit := self.dim_units.get(dim): + dim_label += f" [${unit}$]" + + return dim_label + + +@dataclass +class PlotInfo2D(PlotInfo): + _: KW_ONLY + x_dim: DimKey + y_dim: DimKey + + +@dataclass +class LineInfo(PlotInfo2D): + _: KW_ONLY + x_data: np.ndarray + y_data: np.ndarray + line_style: LineStyleType = "-" + + +@dataclass +class ImageInfo(PlotInfo2D): + _: KW_ONLY + data: np.ndarray + color_dim: DimKey + + +@dataclass +class ScatterInfo(PlotInfo2D): + _: KW_ONLY + xy_data: np.ndarray + color_data: np.ndarray | None = None + color_dim: DimKey | None = None + + +@dataclass +class PolarMeshInfo(PlotInfo): + _: KW_ONLY + data: np.ndarray + r_vertices: np.ndarray + theta_vertices: np.ndarray + r_dim: DimKey + theta_dim: DimKey + color_dim: DimKey + projection: Projection = field(default="polar", init=False) diff --git a/src/lib/plotting/plt_util.py b/src/lib/plotting/plt_util.py index 81ffbe89..34142c15 100644 --- a/src/lib/plotting/plt_util.py +++ b/src/lib/plotting/plt_util.py @@ -1,21 +1,9 @@ -import typing - import matplotlib.pyplot as plt from matplotlib.axes import Axes from matplotlib.colorizer import _ScalarMappable -from matplotlib.colors import Normalize -from matplotlib.scale import ScaleBase from lib.data.data_with_attrs import ListMetadata, Metadata -type BuiltinAxisScaleKey = typing.Literal["linear", "log"] -SCALES: list[BuiltinAxisScaleKey] = list(BuiltinAxisScaleKey.__value__.__args__) -type AxisScaleArg = BuiltinAxisScaleKey | ScaleBase - -type BuiltinColorNormKey = typing.Literal["linear", "log"] -BUILTIN_COLOR_NORM_KEYS: tuple[BuiltinColorNormKey, ...] = BuiltinColorNormKey.__value__.__args__ -type ColorNormArg = BuiltinColorNormKey | Normalize - def symmetrize_bounds(lower: float, upper: float) -> tuple[float, float]: if lower < 0 < upper: diff --git a/src/lib/plotting/renderer.py b/src/lib/plotting/renderer.py index 8c6c1961..c93ab3e0 100644 --- a/src/lib/plotting/renderer.py +++ b/src/lib/plotting/renderer.py @@ -1,26 +1,37 @@ from __future__ import annotations from abc import ABC, abstractmethod -from typing import Any -from matplotlib.axes import Axes -from matplotlib.figure import Figure - -from lib.data.data_with_attrs import DataWithAttrs +from lib.data.adaptors.idx import Idx +from lib.data.data_with_attrs import DataWithAttrs, Field +from lib.data.plot_target import PlotTarget +from lib.plotting.plot_info import PlotInfo class Renderer[Data: DataWithAttrs](ABC): - def subplot_kw(self) -> dict[str, Any]: - return {} + def __init__(self, full_data: Data, plot_target: PlotTarget): + self.plot_target = plot_target - @abstractmethod - def make_init_data(self, fig: Figure, ax: Axes, frame_data: Data) -> Any: ... + if isinstance(full_data, Field): + self.full_data = full_data.assign_metadata(active_key=plot_target.color_dim or plot_target.spatial_dims.y_dim) + else: + self.full_data = full_data - @abstractmethod - def init(self, fig: Figure, ax: Axes, full_data: Data, frame_data: Data, init_data: Any) -> None: ... + self.plot_info = self.init_plot_info() + + def _get_data_at_frame(self, frame: int) -> Data: + if self.plot_target.time_dim: + frame = min(frame, self.get_n_frames() - 1) + return Idx({self.plot_target.time_dim: frame}).apply(self.full_data) + return self.full_data - def make_update_data(self, ax: Axes, frame_data: Data) -> Any: - return None + def get_n_frames(self) -> int: + if self.plot_target.time_dim: + return len(self.full_data.coordss[self.plot_target.time_dim]) + return 1 - def draw(self, ax: Axes, frame_data: Data, update_data: Any) -> None: - pass + @abstractmethod + def init_plot_info(self) -> PlotInfo: ... + + @abstractmethod + def update_plot_info(self, frame: int): ... diff --git a/src/lib/plotting/renderers/field_1d.py b/src/lib/plotting/renderers/field_1d.py index 921a6e0e..9226b487 100644 --- a/src/lib/plotting/renderers/field_1d.py +++ b/src/lib/plotting/renderers/field_1d.py @@ -1,56 +1,52 @@ -from dataclasses import dataclass - -from matplotlib.axes import Axes -from matplotlib.figure import Figure - from lib.data.data_with_attrs import Field from lib.plotting import plt_util -from lib.plotting.frame_data_traits import ( - HasAxes, - HasFieldData, - HasLineType, - HasSpatialScales, -) +from lib.plotting.plot_info import LineInfo, PlotInfo from lib.plotting.renderer import Renderer class Field1dRenderer(Renderer[Field]): - @dataclass(kw_only=True) - class InitData(HasFieldData, HasLineType, HasAxes, HasSpatialScales): ... - - @dataclass(kw_only=True) - class UpdateData(HasFieldData, HasAxes): ... - - def make_init_data(self, fig: Figure, ax: Axes, frame_data: Field) -> InitData: - return self.InitData( - data=frame_data, - axes=ax, - line_type="-", - spatial_scales=["linear", "linear"], - last_spatial_dim_is_dependent=True, + def init_plot_info(self) -> PlotInfo: + full_data = self.full_data + frame_data = self._get_data_at_frame(0) + + [x_dim, y_dim] = self.plot_target.spatial_dims.unpack() + + plot_info = LineInfo( + x_data=frame_data.coordss[x_dim], + y_data=frame_data.active_data, + x_dim=x_dim, + y_dim=y_dim, + time_dim=self.plot_target.time_dim, + subject=frame_data.metadata.active_var_info.to_axis_label(), + dim_scales={ + x_dim: frame_data.metadata.var_infos[x_dim].scale, + y_dim: frame_data.metadata.var_infos[y_dim].scale, + }, + dim_bounds={ + x_dim: (frame_data.coordss[x_dim][0], frame_data.coordss[x_dim][-1]), + y_dim: plt_util.symmetrize_bounds(*full_data.var_bounds), + }, + dim_displays={ + x_dim: frame_data.metadata.var_infos[x_dim].display, + y_dim: frame_data.metadata.var_infos[y_dim].display, + }, + dim_units={ + x_dim: frame_data.metadata.var_infos[x_dim].unit, + y_dim: frame_data.metadata.var_infos[y_dim].unit, + }, + axes_index=self.plot_target.axes_index, ) - def init(self, fig: Figure, ax: Axes, full_data: Field, frame_data: Field, init_data: InitData) -> None: - [dim_x] = frame_data.metadata.spatial_dims - xdata = frame_data.coordss[dim_x] - ydata = frame_data.active_data - - [self.line] = ax.plot(xdata, ydata, init_data.line_type) - - plt_util.update_title(ax, frame_data.metadata, [frame_data.metadata.var_infos[dim].get_coordinate_label(pos) for dim, pos in frame_data.coordss.items() if pos.shape == ()]) - ax.set_xlabel(frame_data.metadata.var_infos[dim_x].to_axis_label()) - ax.set_ylabel(frame_data.metadata.active_var_info.to_axis_label()) - - ax.set_xscale(init_data.spatial_scales[0]) - ax.set_yscale(init_data.spatial_scales[1]) - - ymin, ymax = plt_util.symmetrize_bounds(*full_data.var_bounds) - ax.set_ybound(ymin, ymax) + for dim, coord in frame_data.coordss.items(): + if coord.shape == (): + plot_info.scalar_coord_values[dim] = coord + plot_info.dim_displays[dim] = frame_data.metadata.var_infos[dim].display + plot_info.dim_units[dim] = frame_data.metadata.var_infos[dim].unit - def make_update_data(self, ax: Axes, frame_data: Field) -> UpdateData: - return self.UpdateData(data=frame_data, axes=ax) + return plot_info - def draw(self, ax: Axes, frame_data: Field, update_data: UpdateData) -> None: - self.line.set_ydata(frame_data.active_data) + def update_plot_info(self, frame: int): + frame_data = self._get_data_at_frame(frame) - plt_util.update_title(ax, frame_data.metadata, [frame_data.metadata.var_infos[dim].get_coordinate_label(pos) for dim, pos in frame_data.coordss.items() if pos.shape == ()]) + self.plot_info.set("y_data", frame_data.active_data) + self.plot_info.set("scalar_coord_values", {dim: coord for dim, coord in frame_data.coordss.items() if coord.shape == ()}) diff --git a/src/lib/plotting/renderers/field_2d.py b/src/lib/plotting/renderers/field_2d.py index b9600f7d..3a01b320 100644 --- a/src/lib/plotting/renderers/field_2d.py +++ b/src/lib/plotting/renderers/field_2d.py @@ -1,12 +1,8 @@ -from dataclasses import dataclass - import xarray as xr -from matplotlib.axes import Axes -from matplotlib.figure import Figure from lib.data.data_with_attrs import Field -from lib.plotting import plt_util -from lib.plotting.frame_data_traits import HasAxes, HasColorNorm, HasFieldData, HasSpatialScales +from lib.data.plot_target import SpatialDimsXY +from lib.plotting.plot_info import ImageInfo, PlotInfo from lib.plotting.renderer import Renderer @@ -17,57 +13,62 @@ def get_extent(da: xr.DataArray, dim: str) -> tuple[float, float]: class Field2dRenderer(Renderer[Field]): - @dataclass(kw_only=True) - class InitData(HasFieldData, HasSpatialScales, HasColorNorm, HasAxes): ... - - @dataclass(kw_only=True) - class UpdateData(HasFieldData, HasAxes): ... - def _transpose(self, data: Field) -> Field: spatial_dims = data.metadata.spatial_dims return data.with_active_data(data.active_data.transpose(*reversed(spatial_dims))) - def make_init_data(self, fig: Figure, ax: Axes, frame_data: Field) -> InitData: - return self.InitData( - data=frame_data, - spatial_scales=["linear", "linear"], - color_norm="linear", - color_is_dependent=True, - axes=ax, - ) - - def init(self, fig: Figure, ax: Axes, full_data: Field, frame_data: Field, init_data: InitData) -> None: - spatial_dims = frame_data.metadata.spatial_dims - frame_data = self._transpose(frame_data) - da = frame_data.active_data - - # must set scale (log, linear) before making image - ax.set_xscale(init_data.spatial_scales[0]) - ax.set_yscale(init_data.spatial_scales[1]) - - self.im = ax.imshow( - da, - origin="lower", - extent=(*get_extent(da, spatial_dims[0]), *get_extent(da, spatial_dims[1])), - norm=init_data.color_norm, - interpolation="nearest", + def init_plot_info(self) -> PlotInfo: + full_data = self.full_data + frame_data = self._get_data_at_frame(0) + + [x_dim, y_dim] = self.plot_target.spatial_dims.unpack() + color_dim = self.plot_target.color_dim + + data = frame_data.active_data.transpose(y_dim, x_dim) + + plot_info = ImageInfo( + data=data, + x_dim=x_dim, + y_dim=y_dim, + color_dim=color_dim, + time_dim=self.plot_target.time_dim, + subject=frame_data.metadata.active_var_info.to_axis_label(), + dim_scales={ + x_dim: frame_data.metadata.var_infos[x_dim].scale, + y_dim: frame_data.metadata.var_infos[y_dim].scale, + color_dim: frame_data.metadata.var_infos[color_dim].scale, + }, + dim_bounds={ + x_dim: get_extent(data, x_dim), + y_dim: get_extent(data, y_dim), + color_dim: full_data.var_bounds, + }, + dim_displays={ + x_dim: frame_data.metadata.var_infos[x_dim].display, + y_dim: frame_data.metadata.var_infos[y_dim].display, + color_dim: None, + }, + dim_units={ + x_dim: frame_data.metadata.var_infos[x_dim].unit, + y_dim: frame_data.metadata.var_infos[y_dim].unit, + color_dim: frame_data.metadata.var_infos[color_dim].unit, + }, + axes_index=self.plot_target.axes_index, ) - fig.colorbar(self.im) - data_lower, data_upper = full_data.var_bounds - plt_util.update_cbar(self.im, data_min_override=data_lower, data_max_override=data_upper) - - plt_util.update_title(ax, frame_data.metadata, [frame_data.metadata.var_infos[dim].get_coordinate_label(pos) for dim, pos in frame_data.coordss.items() if pos.shape == ()]) + for dim, coord in frame_data.coordss.items(): + if coord.shape == (): + plot_info.scalar_coord_values[dim] = coord + plot_info.dim_displays[dim] = frame_data.metadata.var_infos[dim].display + plot_info.dim_units[dim] = frame_data.metadata.var_infos[dim].unit - ax.set_aspect(1 / ax.get_data_ratio()) - ax.set_xlabel(frame_data.metadata.var_infos[spatial_dims[0]].to_axis_label()) - ax.set_ylabel(frame_data.metadata.var_infos[spatial_dims[1]].to_axis_label()) + return plot_info - def make_update_data(self, ax: Axes, frame_data: Field) -> UpdateData: - return self.UpdateData(data=frame_data, axes=ax) + def update_plot_info(self, frame: int): + frame_data = self._get_data_at_frame(frame) - def draw(self, ax: Axes, frame_data: Field, update_data: UpdateData) -> None: - frame_data = self._transpose(frame_data) - self.im.set_data(frame_data.active_data) + [x_dim, y_dim] = self.plot_target.spatial_dims.unpack() + data = frame_data.active_data.transpose(y_dim, x_dim) - plt_util.update_title(ax, frame_data.metadata, [frame_data.metadata.var_infos[dim].get_coordinate_label(pos) for dim, pos in frame_data.coordss.items() if pos.shape == ()]) + self.plot_info.set("data", data) + self.plot_info.set("scalar_coord_values", {dim: coord for dim, coord in frame_data.coordss.items() if coord.shape == ()}) diff --git a/src/lib/plotting/renderers/polar_field.py b/src/lib/plotting/renderers/polar_field.py index ba313176..09425047 100644 --- a/src/lib/plotting/renderers/polar_field.py +++ b/src/lib/plotting/renderers/polar_field.py @@ -1,75 +1,69 @@ -from dataclasses import dataclass -from typing import Any - import numpy as np -from matplotlib.axes import Axes -from matplotlib.figure import Figure from lib.data.data_with_attrs import Field -from lib.plotting import plt_util -from lib.plotting.frame_data_traits import HasAxes, HasColorNorm, HasFieldData, HasSpatialScales +from lib.plotting.plot_info import PlotInfo, PolarMeshInfo from lib.plotting.renderer import Renderer class PolarFieldRenderer(Renderer[Field]): - @dataclass(kw_only=True) - class InitData(HasFieldData, HasSpatialScales, HasColorNorm, HasAxes): ... - - @dataclass(kw_only=True) - class UpdateData(HasFieldData, HasAxes): ... - - def subplot_kw(self) -> dict[str, Any]: - return {"projection": "polar"} - - def make_init_data(self, fig: Figure, ax: Axes, frame_data: Field) -> InitData: - return self.InitData( - data=frame_data, - spatial_scales=["linear", "linear"], - color_norm="linear", - color_is_dependent=True, - axes=ax, - ) - - def init(self, fig: Figure, ax: Axes, full_data: Field, frame_data: Field, init_data: InitData) -> None: - spatial_dims = frame_data.metadata.spatial_dims - - # must set scale before making image - ax.set_rscale(init_data.spatial_scales[0]) - - vertices_theta = frame_data.coordss[spatial_dims[1]] - vertices_theta = np.concat([vertices_theta, [vertices_theta[-1] + vertices_theta[1] - vertices_theta[0]]]) - vertices_r = list(frame_data.coordss[spatial_dims[0]]) - vertices_r = np.concat([vertices_r, [vertices_r[-1] + vertices_r[1] - vertices_r[0]]]) - - if vertices_theta[0] == 0.0: + def init_plot_info(self) -> PlotInfo: + full_data = self.full_data + frame_data = self._get_data_at_frame(0) + + [r_dim, theta_dim] = self.plot_target.spatial_dims.unpack() + color_dim = self.plot_target.color_dim + + theta_vertices = frame_data.coordss[theta_dim] + theta_vertices = np.concat([theta_vertices, [theta_vertices[-1] + theta_vertices[1] - theta_vertices[0]]]) + r_vertices = list(frame_data.coordss[r_dim]) + r_vertices = np.concat([r_vertices, [r_vertices[-1] + r_vertices[1] - r_vertices[0]]]) + if theta_vertices[0] == 0.0: # FIXME hacky check for interpolated values # there are two different ways to go from cartesian to polar: # - interpolating onto a polar grid, in which case theta coords are "cell centered" (and happen to start at 0) # - scattering and binning, in which case theta coords are "node centered" (and happen to start at -pi) # this does a half-cell rotation in the former case to transform to "node centered" coords, which matlotlib expects - vertices_theta -= vertices_theta[1] / 2.0 - - self.im = ax.pcolormesh( - *np.meshgrid(vertices_theta, vertices_r), - frame_data.active_data, - shading="flat", - norm=init_data.color_norm, + theta_vertices -= theta_vertices[1] / 2.0 + + plot_info = PolarMeshInfo( + data=frame_data.active_data, + r_dim=r_dim, + theta_dim=theta_dim, + color_dim=color_dim, + time_dim=self.plot_target.time_dim, + r_vertices=r_vertices, + theta_vertices=theta_vertices, + subject=frame_data.metadata.active_var_info.to_axis_label(), + dim_scales={ + r_dim: frame_data.metadata.var_infos[r_dim].scale, + color_dim: frame_data.metadata.var_infos[color_dim].scale, + }, + dim_bounds={ + color_dim: full_data.var_bounds, + }, + dim_displays={ + r_dim: frame_data.metadata.var_infos[r_dim].display, + theta_dim: frame_data.metadata.var_infos[theta_dim].display, + color_dim: None, + }, + dim_units={ + r_dim: frame_data.metadata.var_infos[r_dim].unit, + theta_dim: frame_data.metadata.var_infos[theta_dim].unit, + color_dim: frame_data.metadata.var_infos[color_dim].unit, + }, + axes_index=self.plot_target.axes_index, ) - fig.colorbar(self.im) - data_lower, data_upper = full_data.var_bounds - plt_util.update_cbar(self.im, data_min_override=data_lower, data_max_override=data_upper) - - plt_util.update_title(ax, frame_data.metadata, [frame_data.metadata.var_infos[dim].get_coordinate_label(pos) for dim, pos in frame_data.coordss.items() if pos.shape == ()]) - - # FIXME make the labels work - # ax.set_xlabel(frame_data.metadata.var_info[spatial_dims[1]].to_axis_label()) - # ax.set_ylabel(frame_data.metadata.var_info[spatial_dims[0]].to_axis_label()) + for dim, coord in frame_data.coordss.items(): + if coord.shape == (): + plot_info.scalar_coord_values[dim] = coord + plot_info.dim_displays[dim] = frame_data.metadata.var_infos[dim].display + plot_info.dim_units[dim] = frame_data.metadata.var_infos[dim].unit - def make_update_data(self, ax: Axes, frame_data: Field) -> UpdateData: - return self.UpdateData(data=frame_data, axes=ax) + return plot_info - def draw(self, ax: Axes, frame_data: Field, update_data: UpdateData) -> None: - self.im.set_array(frame_data.active_data) + def update_plot_info(self, frame: int): + frame_data = self._get_data_at_frame(frame) - plt_util.update_title(ax, frame_data.metadata, [frame_data.metadata.var_infos[dim].get_coordinate_label(pos) for dim, pos in frame_data.coordss.items() if pos.shape == ()]) + self.plot_info.set("data", frame_data.active_data) + self.plot_info.set("scalar_coord_values", {dim: coord for dim, coord in frame_data.coordss.items() if coord.shape == ()}) diff --git a/src/lib/plotting/renderers/scatter.py b/src/lib/plotting/renderers/scatter.py index 29599a6b..4e10a08f 100644 --- a/src/lib/plotting/renderers/scatter.py +++ b/src/lib/plotting/renderers/scatter.py @@ -1,82 +1,65 @@ -from dataclasses import dataclass - import numpy as np -from matplotlib.axes import Axes -from matplotlib.figure import Figure from lib.data.data_with_attrs import FullList -from lib.plotting import plt_util -from lib.plotting.frame_data_traits import ( - HasAxes, - HasColorNorm, - HasFullListData, - HasSpatialScales, -) +from lib.plotting.plot_info import PlotInfo, ScatterInfo from lib.plotting.renderer import Renderer class ScatterRenderer(Renderer[FullList]): - @dataclass(kw_only=True) - class InitData(HasFullListData, HasAxes, HasSpatialScales, HasColorNorm): ... - - @dataclass(kw_only=True) - class UpdateData(HasFullListData, HasAxes): ... - - def make_init_data(self, fig: Figure, ax: Axes, frame_data: FullList) -> InitData: - return self.InitData( - data=frame_data, - axes=ax, - spatial_scales=["linear", "linear"], - last_spatial_dim_is_dependent=True, - color_norm="linear", + def init_plot_info(self) -> PlotInfo: + full_data = self.full_data + frame_data = self._get_data_at_frame(0) + + [x_dim, y_dim] = self.plot_target.spatial_dims.unpack() + + plot_info = ScatterInfo( + xy_data=np.array([frame_data.data[x_dim], frame_data.data[y_dim]]).T, + x_dim=x_dim, + y_dim=y_dim, + time_dim=self.plot_target.time_dim, + subject=f"${frame_data.metadata.subject}$" if frame_data.metadata.subject else None, + dim_scales={ + x_dim: frame_data.metadata.var_infos[x_dim].scale, + y_dim: frame_data.metadata.var_infos[y_dim].scale, + }, + dim_bounds={ + x_dim: full_data.bounds(x_dim), + y_dim: full_data.bounds(y_dim), + }, + dim_displays={ + x_dim: frame_data.metadata.var_infos[x_dim].display, + y_dim: frame_data.metadata.var_infos[y_dim].display, + }, + dim_units={ + x_dim: frame_data.metadata.var_infos[x_dim].unit, + y_dim: frame_data.metadata.var_infos[y_dim].unit, + }, + axes_index=self.plot_target.axes_index, ) - def init(self, fig: Figure, ax: Axes, full_data: FullList, frame_data: FullList, init_data: InitData) -> None: - [dim_x, dim_y] = frame_data.metadata.spatial_dims - df = frame_data.data - - ax.set_xscale(init_data.spatial_scales[0]) - ax.set_yscale(init_data.spatial_scales[1]) - - ax.set_xlabel(frame_data.metadata.var_infos[dim_x].to_axis_label()) - ax.set_ylabel(frame_data.metadata.var_infos[dim_y].to_axis_label()) - - ax.set_xlim(*full_data.bounds(dim_x)) - ax.set_ylim(*full_data.bounds(dim_y)) - - if frame_data.metadata.color_dim: - self.scatter = ax.scatter( - df[dim_x], - df[dim_y], - c=df[frame_data.metadata.color_dim], - norm=init_data.color_norm, - s=1, - ) - - fig.colorbar(self.scatter, label=frame_data.metadata.var_infos[frame_data.metadata.color_dim].to_axis_label()) - data_lower, data_upper = full_data.bounds(frame_data.metadata.color_dim) - plt_util.update_cbar(self.scatter, data_min_override=data_lower, data_max_override=data_upper) - else: - self.scatter = ax.scatter( - df[dim_x], - df[dim_y], - color=ax._get_lines.get_next_color(), - s=0.5, - ) + for dim, coord in frame_data.coordss.items(): + if coord.shape == (): + plot_info.scalar_coord_values[dim] = coord + plot_info.dim_displays[dim] = frame_data.metadata.var_infos[dim].display + plot_info.dim_units[dim] = frame_data.metadata.var_infos[dim].unit - plt_util.update_title(ax, frame_data.metadata, [frame_data.metadata.var_infos[dim].get_coordinate_label(pos) for dim, pos in frame_data.coordss.items() if isinstance(pos, float)]) + if color_dim := self.plot_target.color_dim: + plot_info.color_dim = color_dim + plot_info.color_data = frame_data.data[color_dim] + plot_info.dim_scales[color_dim] = frame_data.metadata.var_infos[color_dim].scale + plot_info.dim_bounds[color_dim] = full_data.bounds(color_dim) + plot_info.dim_displays[color_dim] = frame_data.metadata.var_infos[color_dim].display + plot_info.dim_units[color_dim] = frame_data.metadata.var_infos[color_dim].unit - ax.set_aspect(1 / ax.get_data_ratio()) + return plot_info - def make_update_data(self, ax: Axes, frame_data: FullList) -> UpdateData: - return self.UpdateData(data=frame_data, axes=ax) + def update_plot_info(self, frame: int): + frame_data = self._get_data_at_frame(frame) - def draw(self, ax: Axes, frame_data: FullList, update_data: UpdateData) -> None: - spatial_dims = frame_data.metadata.spatial_dims - df = frame_data.data + [x_dim, y_dim] = self.plot_target.spatial_dims.unpack() - self.scatter.set_offsets(np.array([df[spatial_dims[0]], df[spatial_dims[1]]]).T) - plt_util.update_title(ax, frame_data.metadata, [frame_data.metadata.var_infos[dim].get_coordinate_label(pos) for dim, pos in frame_data.coordss.items() if isinstance(pos, float)]) + self.plot_info.set("xy_data", np.array([frame_data.data[x_dim], frame_data.data[y_dim]]).T) + self.plot_info.set("scalar_coord_values", {dim: coord for dim, coord in frame_data.coordss.items() if coord.shape == ()}) - if frame_data.metadata.color_dim: - self.scatter.set_array(df[frame_data.metadata.color_dim]) + if color_dim := self.plot_info.color_dim: + self.plot_info.set("color_data", frame_data.data[color_dim]) diff --git a/src/lib/plotting/setup_fig.py b/src/lib/plotting/setup_fig.py new file mode 100644 index 00000000..857775b7 --- /dev/null +++ b/src/lib/plotting/setup_fig.py @@ -0,0 +1,504 @@ +from abc import ABC, abstractmethod +from typing import Iterable + +import numpy as np +from matplotlib import pyplot as plt +from matplotlib.axes import Axes +from matplotlib.figure import Figure +from matplotlib.lines import Line2D +from matplotlib.projections import PolarAxes +from matplotlib.text import Text + +from lib.plotting import plt_util +from lib.plotting.plot_info import DimKey, ImageInfo, LineInfo, PlotInfo, PlotInfo2D, PolarMeshInfo, ScatterInfo + +type AxesIdx = tuple[int, int] + + +def _flatten_idx(axes_idx: tuple[int, int], ncols: int) -> int: + return ncols * (axes_idx[1] - 1) + axes_idx[0] + + +def _setup_axes(figure: Figure, plot_infos: list[PlotInfo]) -> dict[AxesIdx, tuple[Axes, list[PlotInfo]]]: + idx_to_infos: dict[AxesIdx, list[PlotInfo]] = {} + for info in plot_infos: + idx_to_infos.setdefault(info.axes_index, []).append(info) + + ncols = max(idx[0] for idx in idx_to_infos) + nrows = max(idx[1] for idx in idx_to_infos) + + ret: dict[AxesIdx, tuple[Axes, list[PlotInfo]]] = {} + for idx, infos in idx_to_infos.items(): + projection = infos[0].projection + for info in infos[1:]: + if info.projection != projection: + raise ValueError("incompatible plots (TODO: better error message)") + ax = figure.add_subplot(nrows, ncols, _flatten_idx(idx, ncols), projection=projection) + ret[idx] = (ax, infos) + + return ret + + +def _one_or_none[T](objs: Iterable[T]) -> T | None: + one = None + for obj in objs: + if one is None: + one = obj + elif obj != one: + return None + return one + + +def find_widest_bounds(boundss: Iterable[tuple[float | None, float | None]]) -> tuple[float | None, float | None]: + lowest_bound = None + highest_bound = None + + for bounds in boundss: + if lowest_bound is None: + lowest_bound = bounds[0] + elif bounds[0] is not None and lowest_bound > bounds[0]: + lowest_bound = bounds[0] + + if highest_bound is None: + highest_bound = bounds[1] + elif bounds[1] is not None and highest_bound < bounds[1]: + highest_bound = bounds[1] + + return (lowest_bound, highest_bound) + + +class UpdateText: + def __init__(self, text: Text, plot_info: PlotInfo): + self.text = text + self.plot_info = plot_info + + def __call__(self, *_): + self.text.set_text(self.plot_info.get_title()) + + +class AxesManager(ABC): + @abstractmethod + def setup(self): ... + + @abstractmethod + def setup_title(self): ... + + @abstractmethod + def setup_labels(self): ... + + @abstractmethod + def setup_scales(self): ... + + @abstractmethod + def setup_bounds(self): ... + + @abstractmethod + def setup_data(self): ... + + +class AxesManagerSingle[A: Axes, PI: PlotInfo](AxesManager): + def __init__(self, ax: A, info: PI): + self.ax = ax + self.info = info + + def setup_title(self): + update_title = UpdateText(self.ax.title, self.info) + self.info._setter_callbacks["subject"] = update_title + self.info._setter_callbacks["dim_displays"] = update_title + self.info._setter_callbacks["dim_units"] = update_title + self.info._setter_callbacks["scalar_coord_values"] = update_title + update_title() + + +class AxesManagerSingle2D[PI2D: PlotInfo2D](AxesManagerSingle[Axes, PI2D]): + def setup(self): + self.setup_title() + self.setup_labels() + self.setup_data() + self.setup_scales() + self.setup_bounds() + + def setup_labels(self): + self.ax.set_xlabel(self.info.get_dim_label(self.info.x_dim)) + self.ax.set_ylabel(self.info.get_dim_label(self.info.y_dim)) + + def setup_scales(self): + self.ax.set_xscale(self.info.dim_scales[self.info.x_dim].to_axis_scale()) + self.ax.set_yscale(self.info.dim_scales[self.info.y_dim].to_axis_scale()) + + def setup_bounds(self): + self.ax.set_xlim(*self.info.dim_bounds[self.info.x_dim]) + self.ax.set_ylim(*self.info.dim_bounds[self.info.y_dim]) + + +class AxesManagerSingleLine(AxesManagerSingle2D[LineInfo]): + def setup_data(self): + [line] = self.ax.plot(self.info.x_data, self.info.y_data, linestyle=self.info.line_style, scalex=False, scaley=False) + self.info._setter_callbacks["x_data"] = line.set_xdata + self.info._setter_callbacks["y_data"] = line.set_ydata + self.info._setter_callbacks["line_style"] = line.set_linestyle + + +class AxesManagerSingleImage(AxesManagerSingle2D[ImageInfo]): + def setup(self): + super().setup() + self.ax.set_aspect(1 / self.ax.get_data_ratio()) + + def setup_data(self): + image = self.ax.imshow( + self.info.data, + origin="lower", + extent=(*self.info.dim_bounds[self.info.x_dim], *self.info.dim_bounds[self.info.y_dim]), + norm=self.info.dim_scales[self.info.color_dim].to_color_norm(), + interpolation="nearest", + ) + self.info._setter_callbacks["data"] = image.set_data + + self.ax.figure.colorbar(image) + data_lower, data_upper = self.info.dim_bounds[self.info.color_dim] + plt_util.update_cbar(image, data_min_override=data_lower, data_max_override=data_upper) + + +class AxesManagerSingleScatter(AxesManagerSingle2D[ScatterInfo]): + def setup(self): + super().setup() + self.ax.set_aspect(1 / self.ax.get_data_ratio()) + + def setup_data(self): + if self.info.color_dim: + scatter = self.ax.scatter( + self.info.xy_data[:, 0], + self.info.xy_data[:, 1], + c=self.info.color_data, + norm=self.info.dim_scales[self.info.color_dim].to_color_norm(), + s=1, + ) + self.info._setter_callbacks["color_data"] = scatter.set_array + + self.ax.figure.colorbar(scatter, label=self.info.get_dim_label(self.info.color_dim)) + data_lower, data_upper = self.info.dim_bounds[self.info.color_dim] + plt_util.update_cbar(scatter, data_min_override=data_lower, data_max_override=data_upper) + else: + scatter = self.ax.scatter( + self.info.xy_data[:, 0], + self.info.xy_data[:, 1], + color=self.ax._get_lines.get_next_color(), + s=0.5, + ) + + update_data = lambda _=None: scatter.set_offsets(self.info.xy_data) + self.info._setter_callbacks["xy_data"] = update_data + + +class AxesManagerSinglePolarMesh(AxesManagerSingle[PolarAxes, PolarMeshInfo]): + def setup(self): + self.setup_title() + self.setup_labels() + self.setup_scales() + self.setup_data() + + def setup_labels(self): + # FIXME make the labels work + pass + + def setup_bounds(self): + pass + + def setup_scales(self): + self.ax.set_rscale(self.info.dim_scales[self.info.r_dim].to_axis_scale()) + + def setup_data(self): + image = self.ax.pcolormesh( + *np.meshgrid(self.info.theta_vertices, self.info.r_vertices), + self.info.data, + shading="flat", + norm=self.info.dim_scales[self.info.color_dim].to_color_norm(), + ) + self.info._setter_callbacks["data"] = image.set_array + + self.ax.figure.colorbar(image) + data_lower, data_upper = self.info.dim_bounds[self.info.color_dim] + plt_util.update_cbar(image, data_min_override=data_lower, data_max_override=data_upper) + + +class AxesManagerMultiLine(AxesManager): + def __init__(self, ax: Axes, infos: list[LineInfo]): + self.ax = ax + self.infos = infos + + self.common_coord_dims: set[DimKey] = set() + self.unique_coord_dimss: list[set[DimKey]] = [set() for _ in self.infos] + self._update_common_scalar_coordinates() + + self.common_subject: str | None = None + self._update_common_subject() + + self.lines: list[Line2D] = [] # populated later + + def _update_common_scalar_coordinates(self): + self.common_coord_dims.clear() + for unique_dims in self.unique_coord_dimss: + unique_dims.clear() + + all_dims: set[DimKey] = {dim for info in self.infos for dim in info.scalar_coord_values} + + for dim in all_dims: + if all(dim in info.scalar_coord_values for info in self.infos) and len({info.get_coord_label(dim) for info in self.infos}) == 1: + self.common_coord_dims.add(dim) + else: + for info, unique_dims in zip(self.infos, self.unique_coord_dimss): + if dim in info.scalar_coord_values: + unique_dims.add(dim) + + def _update_common_subject(self): + self.common_subject = _one_or_none(info.subject for info in self.infos) + + def _get_title(self) -> str: + # by definition, it shouldn't matter which info we use to construct this string + coord_labels_str = ", ".join(self.infos[0].get_coord_label(dim).maybe_with_dollars() for dim in self.common_coord_dims) + + if self.common_subject and coord_labels_str: + return f"{self.common_subject} ({coord_labels_str})" + return self.common_subject or coord_labels_str + + def _get_legend_labels(self) -> list[str]: + legend_labels: list[str] = [] + + for info, unique_dims in zip(self.infos, self.unique_coord_dimss): + coord_labels_str = ", ".join(info.get_coord_label(dim).maybe_with_dollars() for dim in unique_dims) + + if self.common_subject: + legend_labels.append(coord_labels_str) + elif info.subject and coord_labels_str: + legend_labels.append(f"{info.subject} ({coord_labels_str})") + else: + legend_labels.append(info.subject or coord_labels_str) + + return legend_labels + + def _update_title_and_legend(self, *_): + self._update_common_scalar_coordinates() + self._update_common_subject() + self.ax.set_title(self._get_title()) + for line, label in zip(self.lines, self._get_legend_labels()): + line.set_label(label) + + def setup(self): + self.setup_title() + self.setup_labels() + self.setup_data() + self.setup_scales() + self.setup_bounds() + + for info in self.infos: + info._setter_callbacks["scalar_coord_values"] = self._update_title_and_legend + + def setup_title(self): + self.ax.set_title(self._get_title()) + + def setup_labels(self): + x_labels = [info.get_dim_label(info.x_dim) for info in self.infos] + if (x_label := _one_or_none(x_labels)) is not None: + self.ax.set_xlabel(x_label) + else: + raise NotImplementedError(f"x labels must all be the same, but found {x_labels}") + + y_labels = [info.get_dim_label(info.y_dim) for info in self.infos] + y_units = [info.dim_units[info.y_dim] for info in self.infos] + if (y_label := _one_or_none(y_labels)) is not None: + self.ax.set_ylabel(y_label) + elif (y_unit := _one_or_none(y_units)) is not None: + self.ax.set_ylabel(y_unit.maybe_with_dollars()) + else: + raise NotImplementedError(f"y labels must all be the same unit, but found {y_units}") + + def setup_scales(self): + x_scales = [info.dim_scales[info.x_dim] for info in self.infos] + if (x_scale := _one_or_none(x_scales)) is not None: + self.ax.set_xscale(x_scale.to_axis_scale()) + else: + raise NotImplementedError(f"x scales must all be the same, but found {x_scales}") + + y_scales = [info.dim_scales[info.y_dim] for info in self.infos] + if (y_scale := _one_or_none(y_scales)) is not None: + self.ax.set_yscale(y_scale.to_axis_scale()) + else: + raise NotImplementedError(f"y scales must all be the same, but found {y_scales}") + + def setup_bounds(self): + self.ax.set_xbound(*find_widest_bounds(info.dim_bounds[info.x_dim] for info in self.infos)) + self.ax.set_ybound(*find_widest_bounds(info.dim_bounds[info.y_dim] for info in self.infos)) + + def setup_data(self): + for info, label in zip(self.infos, self._get_legend_labels()): + [line] = self.ax.plot(info.x_data, info.y_data, linestyle=info.line_style, scalex=False, scaley=False, label=label) + info._setter_callbacks["x_data"] = line.set_xdata + info._setter_callbacks["y_data"] = line.set_ydata + info._setter_callbacks["line_style"] = line.set_linestyle + self.lines.append(line) + + self.ax.legend() + + +class AxesManagerImageAndLines(AxesManager): + def __init__(self, ax: Axes, image_info: ImageInfo, line_infos: list[LineInfo]): + self.image_ax = ax + self.line_ax = ax.twinx() + self.image_info = image_info + self.line_infos = line_infos + self.infos: list[PlotInfo2D] = [image_info, *line_infos] + + self.common_coord_dims: set[DimKey] = set() + self.unique_coord_dimss: list[set[DimKey]] = [set() for _ in self.line_infos] + self._update_common_scalar_coordinates() + + self.lines: list[Line2D] = [] # populated later + + def _update_common_scalar_coordinates(self): + self.common_coord_dims.clear() + for unique_dims in self.unique_coord_dimss: + unique_dims.clear() + + all_dims: set[DimKey] = {dim for info in self.infos for dim in info.scalar_coord_values} + + for dim in all_dims: + if all(dim in info.scalar_coord_values for info in self.infos) and len({info.get_coord_label(dim) for info in self.infos}) == 1: + self.common_coord_dims.add(dim) + else: + for info, unique_dims in zip(self.infos, self.unique_coord_dimss): + if dim in info.scalar_coord_values: + unique_dims.add(dim) + + def _get_title(self) -> str: + # by definition, it shouldn't matter which info we use to construct this string + coord_labels_str = ", ".join(self.line_infos[0].get_coord_label(dim).maybe_with_dollars() for dim in self.common_coord_dims) + + if self.image_info.subject and coord_labels_str: + return f"{self.image_info.subject} ({coord_labels_str})" + return self.image_info.subject or coord_labels_str + + def _get_legend_labels(self) -> list[str]: + legend_labels: list[str] = [] + + for info, unique_dims in zip(self.line_infos, self.unique_coord_dimss): + coord_labels_str = ", ".join(info.get_coord_label(dim).maybe_with_dollars() for dim in unique_dims) + + if info.subject and coord_labels_str: + legend_labels.append(f"{info.subject} ({coord_labels_str})") + else: + legend_labels.append(info.subject or coord_labels_str) + + return legend_labels + + def _update_title_and_legend(self, *_): + self._update_common_scalar_coordinates() + self.image_ax.set_title(self._get_title()) + for line, label in zip(self.lines, self._get_legend_labels()): + line.set_label(label) + + def setup(self): + self.setup_title() + self.setup_labels() + self.setup_data() + self.setup_scales() + self.setup_bounds() + + for info in self.infos: + info._setter_callbacks["scalar_coord_values"] = self._update_title_and_legend + + def setup_title(self): + self.image_ax.set_title(self._get_title()) + + def setup_labels(self): + x_labels = [info.get_dim_label(info.x_dim) for info in self.infos] + if (x_label := _one_or_none(x_labels)) is not None: + self.image_ax.set_xlabel(x_label) + else: + raise NotImplementedError(f"x labels must all be the same, but found {x_labels}") + + self.image_ax.set_ylabel(self.image_info.get_dim_label(self.image_info.y_dim)) + + y_labels = [info.get_dim_label(info.y_dim) for info in self.line_infos] + y_units = [info.dim_units[info.y_dim] for info in self.line_infos] + if (y_label := _one_or_none(y_labels)) is not None: + self.line_ax.set_ylabel(y_label) + elif (y_unit := _one_or_none(y_units)) is not None: + self.line_ax.set_ylabel(y_unit.maybe_with_dollars()) + else: + raise NotImplementedError(f"line y labels must all be the same unit, but found {y_units}") + + def setup_scales(self): + x_scales = [info.dim_scales[info.x_dim] for info in self.infos] + if (x_scale := _one_or_none(x_scales)) is not None: + self.image_ax.set_xscale(x_scale.to_axis_scale()) + else: + raise NotImplementedError(f"x scales must all be the same, but found {x_scales}") + + self.image_ax.set_yscale(self.image_info.dim_scales[self.image_info.y_dim].to_axis_scale()) + + y_scales = [info.dim_scales[info.y_dim] for info in self.line_infos] + if (y_scale := _one_or_none(y_scales)) is not None: + self.line_ax.set_yscale(y_scale.to_axis_scale()) + else: + raise NotImplementedError(f"y scales must all be the same, but found {y_scales}") + + def setup_bounds(self): + self.image_ax.set_xlim(*find_widest_bounds(info.dim_bounds[info.x_dim] for info in self.infos)) + self.image_ax.set_ylim(*self.image_info.dim_bounds[self.image_info.y_dim]) + self.line_ax.set_ylim(*find_widest_bounds(info.dim_bounds[info.y_dim] for info in self.line_infos)) + + def setup_data(self): + image = self.image_ax.imshow( + self.image_info.data, + origin="lower", + aspect="auto", + extent=(*self.image_info.dim_bounds[self.image_info.x_dim], *self.image_info.dim_bounds[self.image_info.y_dim]), + norm=self.image_info.dim_scales[self.image_info.color_dim].to_color_norm(), + interpolation="nearest", + ) + self.image_info._setter_callbacks["data"] = image.set_data + + self.image_ax.figure.colorbar(image) + data_lower, data_upper = self.image_info.dim_bounds[self.image_info.color_dim] + plt_util.update_cbar(image, data_min_override=data_lower, data_max_override=data_upper) + + for info, label in zip(self.line_infos, self._get_legend_labels()): + [line] = self.line_ax.plot(info.x_data, info.y_data, linestyle=info.line_style, scalex=False, scaley=False, label=label) + info._setter_callbacks["x_data"] = line.set_xdata + info._setter_callbacks["y_data"] = line.set_ydata + info._setter_callbacks["line_style"] = line.set_linestyle + self.lines.append(line) + + self.line_ax.legend() + + +def setup_fig(plot_infos: list[PlotInfo]) -> Figure: + figure = plt.figure(layout="constrained") + + for ax, infos in _setup_axes(figure, plot_infos).values(): + manager: AxesManager + if len(infos) == 1: + info = infos[0] + if isinstance(info, LineInfo): + manager = AxesManagerSingleLine(ax, info) + elif isinstance(info, ImageInfo): + manager = AxesManagerSingleImage(ax, info) + elif isinstance(info, ScatterInfo): + manager = AxesManagerSingleScatter(ax, info) + elif isinstance(info, PolarMeshInfo): + manager = AxesManagerSinglePolarMesh(ax, info) + else: + raise TypeError(f"unknown type: {infos.__class__!r}") + else: + image_infos = [info for info in infos if isinstance(info, ImageInfo)] + line_infos = [info for info in infos if isinstance(info, LineInfo)] + if not image_infos: + manager = AxesManagerMultiLine(ax, line_infos) + elif len(image_infos) == 1: + manager = AxesManagerImageAndLines(ax, image_infos[0], line_infos) + else: + raise NotImplementedError("don't yet support multiple non-line plots per axes") + + manager.setup() + + return figure diff --git a/src/lib/plotting/static_plot.py b/src/lib/plotting/static_plot.py index 4dd9b4af..69644390 100644 --- a/src/lib/plotting/static_plot.py +++ b/src/lib/plotting/static_plot.py @@ -1,34 +1,9 @@ from pathlib import Path -import matplotlib.pyplot as plt - -from lib.data.data_with_attrs import DataWithAttrs from lib.plotting.plot import Plot, SaveFormat -from lib.plotting.renderer import Renderer - - -class StaticPlot[Data: DataWithAttrs](Plot[Data]): - def __init__(self, renderer: Renderer[Data], data: Data): - super().__init__(renderer, data) - - self.fig, self.ax = plt.subplots(subplot_kw=renderer.subplot_kw()) - self._initialized = False - def _initialize(self): - if self._initialized: - return - self._initialized = True - - init_data = self.renderer.make_init_data(self.fig, self.ax, self.data) - self.pre_init_fig(init_data) - self.renderer.init(self.fig, self.ax, self.data, self.data, init_data) - self.post_init_fig(init_data) - self.fig.tight_layout() - - def show(self): - self._initialize() - plt.show() +class StaticPlot(Plot): def allowed_save_formats(self) -> list[SaveFormat]: return ["png"] diff --git a/src/lib/scale.py b/src/lib/scale.py new file mode 100644 index 00000000..7b734834 --- /dev/null +++ b/src/lib/scale.py @@ -0,0 +1,98 @@ +from dataclasses import dataclass, field +from typing import Literal, Self + +from matplotlib.colors import Normalize, SymLogNorm +from matplotlib.scale import ScaleBase, SymmetricalLogScale + +from lib.parsing import parse_util + +type BuiltinAxisScaleKey = Literal["linear", "log"] +SCALES: list[BuiltinAxisScaleKey] = list(BuiltinAxisScaleKey.__value__.__args__) +type AxisScaleArg = BuiltinAxisScaleKey | ScaleBase + +type BuiltinColorNormKey = Literal["linear", "log"] +BUILTIN_COLOR_NORM_KEYS: tuple[BuiltinColorNormKey, ...] = BuiltinColorNormKey.__value__.__args__ +type ColorNormArg = BuiltinColorNormKey | Normalize + +type ScaleKey = Literal["linear", "log", "symlog"] +SCALE_KEYS: tuple[ScaleKey, ...] = ScaleKey.__value__.__args__ + + +@dataclass(frozen=True, unsafe_hash=True) +class Scale: + scale_key: ScaleKey + + def __init_subclass__(cls): + SCALE_TYPES.append(cls) + + def to_axis_scale(self) -> AxisScaleArg: + return self.scale_key + + def to_color_norm(self) -> ColorNormArg: + return self.scale_key + + def to_name_fragment_part(self) -> str: + return str(self.scale_key) + + @classmethod + def to_argparse_format(cls) -> str: + return cls.scale_key + + @classmethod + def try_from_argparse_format(cls, arg: str) -> Self | None: + if arg == cls.scale_key: + return cls() + return None + + +SCALE_TYPES: list[Scale] = [] # automatically populated with subclasses + + +@dataclass(frozen=True, unsafe_hash=True) +class LinearScale(Scale): + scale_key: ScaleKey = field(init=False, default="linear") + + +@dataclass(frozen=True, unsafe_hash=True) +class LogScale(Scale): + scale_key: ScaleKey = field(init=False, default="log") + + +LINEAR_THRESHOLD_ARG_FORMAT = "linear_threshold" + + +@dataclass(frozen=True, unsafe_hash=True) +class SymLogScale(Scale): + linear_threshold: float + scale_key: ScaleKey = field(init=False, default="symlog") + + def to_axis_scale(self) -> AxisScaleArg: + linthresh = self.linear_threshold # or self._choose_linear_threshold(data) + return SymmetricalLogScale(None, linthresh=linthresh) + + def to_color_norm(self) -> ColorNormArg: + linthresh = self.linear_threshold # or self._choose_linear_threshold(data) + return SymLogNorm(linthresh) + + def to_name_fragment_part(self) -> str: + if self.linear_threshold is None: + return self.scale_key + return f"{self.scale_key}{parse_util.SUBARG_DELIM}{self.linear_threshold}" + + @classmethod + def to_argparse_format(cls) -> str: + return f"{cls.scale_key}[{parse_util.SUBARG_DELIM}{LINEAR_THRESHOLD_ARG_FORMAT}]" + + @classmethod + def try_from_argparse_format(cls, arg: str) -> Self | None: + scale_key_arg, linear_threshold_arg = parse_util.parse_optional_assignment(arg, cls.to_argparse_format(), delim=parse_util.SUBARG_DELIM) + + if scale_key_arg != cls.scale_key: + return None + + linear_threshold = parse_util.parse_optional_number(linear_threshold_arg, LINEAR_THRESHOLD_ARG_FORMAT, float) + + return cls(linear_threshold) + + def __eq__(self, value): + return isinstance(value, SymLogScale) and value.scale_key == self.scale_key and self.linear_threshold == value.linear_threshold diff --git a/src/lib/var_info.py b/src/lib/var_info.py index e54931ed..0d36e6c7 100644 --- a/src/lib/var_info.py +++ b/src/lib/var_info.py @@ -1,8 +1,10 @@ from __future__ import annotations -from dataclasses import KW_ONLY, dataclass +from dataclasses import KW_ONLY, dataclass, field from typing import Literal +from lib.scale import LinearScale, Scale + from .latex import Latex INVERSE_ELECTRON_PLASMA_FREQUENCY = Latex("\\omega_\\text{pe}^{-1}") @@ -33,6 +35,7 @@ class VarInfo: geometry: Geometry | None = None _: KW_ONLY key: str = None + scale: Scale = field(default_factory=LinearScale) def __post_init__(self): if self.key is None: @@ -48,11 +51,11 @@ def assign( display = Latex(display) if isinstance(unit, str): unit = Latex(unit) - return VarInfo(display or self.display, unit or self.unit, self.geometry, key=self.key) + return VarInfo(display or self.display, unit or self.unit, self.geometry, key=self.key, scale=self.scale) def to_axis_label(self) -> str: if self.unit: - return f"${self.display}\\ [{self.unit}]$" + return f"${self.display}$ [${self.unit}$]" return f"${self.display}$" def get_coordinate_label(self, coord_val: float) -> str: @@ -70,6 +73,6 @@ def is_fourier(self) -> bool: return self.display.starts_with(FOURIER_KEY_PREFIX) -def check_unit_compatability(dim_1: VarInfo, dim_2: VarInfo, dest_geometry: str): +def check_unit_compatibility(dim_1: VarInfo, dim_2: VarInfo, dest_geometry: str): if dim_1.unit != dim_2.unit: raise ValueError(f"Dimensions {dim_1.display} and {dim_2.display} have incompatible units for transforming to {dest_geometry} coordinates ({dim_1.unit} and {dim_2.unit})") diff --git a/tests/baseline/test_animated_1d.png b/tests/baseline/test_animated_1d.png index 97826359..f78431b4 100644 Binary files a/tests/baseline/test_animated_1d.png and b/tests/baseline/test_animated_1d.png differ diff --git a/tests/baseline/test_animated_1d_bin.png b/tests/baseline/test_animated_1d_bin.png index 69afc9e9..80d10d0b 100644 Binary files a/tests/baseline/test_animated_1d_bin.png and b/tests/baseline/test_animated_1d_bin.png differ diff --git a/tests/baseline/test_animated_1d_downsample.png b/tests/baseline/test_animated_1d_downsample.png index 69afc9e9..80d10d0b 100644 Binary files a/tests/baseline/test_animated_1d_downsample.png and b/tests/baseline/test_animated_1d_downsample.png differ diff --git a/tests/baseline/test_animated_1d_rolling.png b/tests/baseline/test_animated_1d_rolling.png index f0cb0d66..c21c3328 100644 Binary files a/tests/baseline/test_animated_1d_rolling.png and b/tests/baseline/test_animated_1d_rolling.png differ diff --git a/tests/baseline/test_animated_2d.png b/tests/baseline/test_animated_2d.png index 0d7091fa..8982834d 100644 Binary files a/tests/baseline/test_animated_2d.png and b/tests/baseline/test_animated_2d.png differ diff --git a/tests/baseline/test_animated_2d_binned_phase.png b/tests/baseline/test_animated_2d_binned_phase.png index baaa3f60..88d1f328 100644 Binary files a/tests/baseline/test_animated_2d_binned_phase.png and b/tests/baseline/test_animated_2d_binned_phase.png differ diff --git a/tests/baseline/test_animated_2d_derived.png b/tests/baseline/test_animated_2d_derived.png index 24331599..d249c3f2 100644 Binary files a/tests/baseline/test_animated_2d_derived.png and b/tests/baseline/test_animated_2d_derived.png differ diff --git a/tests/baseline/test_animated_2d_idx.png b/tests/baseline/test_animated_2d_idx.png index 5f856732..05817c14 100644 Binary files a/tests/baseline/test_animated_2d_idx.png and b/tests/baseline/test_animated_2d_idx.png differ diff --git a/tests/baseline/test_animated_2d_pos.png b/tests/baseline/test_animated_2d_pos.png index c98a14c9..a22b0fc4 100644 Binary files a/tests/baseline/test_animated_2d_pos.png and b/tests/baseline/test_animated_2d_pos.png differ diff --git a/tests/baseline/test_animated_scatter_electron_positions.png b/tests/baseline/test_animated_scatter_electron_positions.png index 36d4abb2..d3f1aa2e 100644 Binary files a/tests/baseline/test_animated_scatter_electron_positions.png and b/tests/baseline/test_animated_scatter_electron_positions.png differ diff --git a/tests/baseline/test_animated_scatter_ion_phase.png b/tests/baseline/test_animated_scatter_ion_phase.png index a6b4c9da..49e86590 100644 Binary files a/tests/baseline/test_animated_scatter_ion_phase.png and b/tests/baseline/test_animated_scatter_ion_phase.png differ diff --git a/tests/baseline/test_animated_scatter_mul.png b/tests/baseline/test_animated_scatter_mul.png index e1b50571..8b3e217f 100644 Binary files a/tests/baseline/test_animated_scatter_mul.png and b/tests/baseline/test_animated_scatter_mul.png differ diff --git a/tests/baseline/test_animated_scatter_with_variable.png b/tests/baseline/test_animated_scatter_with_variable.png index 14c9fefa..c752eca6 100644 Binary files a/tests/baseline/test_animated_scatter_with_variable.png and b/tests/baseline/test_animated_scatter_with_variable.png differ diff --git a/tests/baseline/test_display_override.png b/tests/baseline/test_display_override.png index b2534055..dd2fd6fb 100644 Binary files a/tests/baseline/test_display_override.png and b/tests/baseline/test_display_override.png differ diff --git a/tests/baseline/test_display_override_dim.png b/tests/baseline/test_display_override_dim.png index 70bf92dd..f2b554f2 100644 Binary files a/tests/baseline/test_display_override_dim.png and b/tests/baseline/test_display_override_dim.png differ diff --git a/tests/baseline/test_gauss_spatial.png b/tests/baseline/test_gauss_spatial.png index 8d464a5b..47f93cdb 100644 Binary files a/tests/baseline/test_gauss_spatial.png and b/tests/baseline/test_gauss_spatial.png differ diff --git a/tests/baseline/test_gauss_temporal.png b/tests/baseline/test_gauss_temporal.png index 165ec749..b4f88d1c 100644 Binary files a/tests/baseline/test_gauss_temporal.png and b/tests/baseline/test_gauss_temporal.png differ diff --git a/tests/baseline/test_hamscan.png b/tests/baseline/test_hamscan.png index 1cd2eee0..6e77f64c 100644 Binary files a/tests/baseline/test_hamscan.png and b/tests/baseline/test_hamscan.png differ diff --git a/tests/baseline/test_image_and_line.png b/tests/baseline/test_image_and_line.png new file mode 100644 index 00000000..7c5bb918 Binary files /dev/null and b/tests/baseline/test_image_and_line.png differ diff --git a/tests/baseline/test_multiline_different_idxs.png b/tests/baseline/test_multiline_different_idxs.png new file mode 100644 index 00000000..dc587697 Binary files /dev/null and b/tests/baseline/test_multiline_different_idxs.png differ diff --git a/tests/baseline/test_multiline_different_vars.png b/tests/baseline/test_multiline_different_vars.png new file mode 100644 index 00000000..fc3026e9 Binary files /dev/null and b/tests/baseline/test_multiline_different_vars.png differ diff --git a/tests/baseline/test_polar_grid.png b/tests/baseline/test_polar_grid.png index 619197cf..cd679411 100644 Binary files a/tests/baseline/test_polar_grid.png and b/tests/baseline/test_polar_grid.png differ diff --git a/tests/baseline/test_scale_symlog.png b/tests/baseline/test_scale_symlog.png index 4f48fa81..a0ed2883 100644 Binary files a/tests/baseline/test_scale_symlog.png and b/tests/baseline/test_scale_symlog.png differ diff --git a/tests/baseline/test_spectrum_1d.png b/tests/baseline/test_spectrum_1d.png index 64d1cbc8..767496d4 100644 Binary files a/tests/baseline/test_spectrum_1d.png and b/tests/baseline/test_spectrum_1d.png differ diff --git a/tests/baseline/test_spectrum_3d.png b/tests/baseline/test_spectrum_3d.png index d81c7cdf..a5ca5ba9 100644 Binary files a/tests/baseline/test_spectrum_3d.png and b/tests/baseline/test_spectrum_3d.png differ diff --git a/tests/baseline/test_static_1d.png b/tests/baseline/test_static_1d.png index 79e3476d..15d02ab2 100644 Binary files a/tests/baseline/test_static_1d.png and b/tests/baseline/test_static_1d.png differ diff --git a/tests/baseline/test_static_1d_rho.png b/tests/baseline/test_static_1d_rho.png index 25a3b051..b15a7005 100644 Binary files a/tests/baseline/test_static_1d_rho.png and b/tests/baseline/test_static_1d_rho.png differ diff --git a/tests/baseline/test_static_1d_rho_derive.png b/tests/baseline/test_static_1d_rho_derive.png index 25a3b051..b15a7005 100644 Binary files a/tests/baseline/test_static_1d_rho_derive.png and b/tests/baseline/test_static_1d_rho_derive.png differ diff --git a/tests/baseline/test_static_1d_rho_i.png b/tests/baseline/test_static_1d_rho_i.png index 5dbbddce..199f0426 100644 Binary files a/tests/baseline/test_static_1d_rho_i.png and b/tests/baseline/test_static_1d_rho_i.png differ diff --git a/tests/baseline/test_static_2d.png b/tests/baseline/test_static_2d.png index df3b7ebd..e2a0add7 100644 Binary files a/tests/baseline/test_static_2d.png and b/tests/baseline/test_static_2d.png differ diff --git a/tests/baseline/test_static_2d_spectogram.png b/tests/baseline/test_static_2d_spectogram.png index 9790dd26..25191585 100644 Binary files a/tests/baseline/test_static_2d_spectogram.png and b/tests/baseline/test_static_2d_spectogram.png differ diff --git a/tests/baseline/test_static_polar.png b/tests/baseline/test_static_polar.png index 8ac90a25..d7586e7f 100644 Binary files a/tests/baseline/test_static_polar.png and b/tests/baseline/test_static_polar.png differ diff --git a/tests/baseline/test_static_scatter.png b/tests/baseline/test_static_scatter.png index b9de26b1..4847406e 100644 Binary files a/tests/baseline/test_static_scatter.png and b/tests/baseline/test_static_scatter.png differ diff --git a/tests/baseline/test_unit_override.png b/tests/baseline/test_unit_override.png index 03ef06f4..238b601e 100644 Binary files a/tests/baseline/test_unit_override.png and b/tests/baseline/test_unit_override.png differ diff --git a/tests/baseline/test_unit_override_dim.png b/tests/baseline/test_unit_override_dim.png index 6a5c494c..ecb785bf 100644 Binary files a/tests/baseline/test_unit_override_dim.png and b/tests/baseline/test_unit_override_dim.png differ diff --git a/tests/baseline/test_vline.png b/tests/baseline/test_vline.png index eb5fbc0a..75902af0 100644 Binary files a/tests/baseline/test_vline.png and b/tests/baseline/test_vline.png differ diff --git a/tests/conftest.py b/tests/conftest.py index ea4b2a5a..3c73e9fd 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -1,56 +1,38 @@ -import os from pathlib import Path -# Must set env var and backend before any lib imports -_TESTS_DIR = Path(__file__).parent -_DATA_DIR = _TESTS_DIR / "data" -os.environ["PSC_PLOT_DATA_DIR"] = str(_DATA_DIR / "test-2d") -os.environ["PSC_PLOT_DASK_NUM_WORKERS"] = "1" - import matplotlib - -matplotlib.use("Agg") - import matplotlib.pyplot as plt import pytest -from lib.config import CONFIG -from lib.parsing.parse import get_parsed_args +from lib.config import PscPlotConfig +from lib.data.compile import compile_plot_node +from lib.parsing.parse import parse_args from lib.plotting.plot import SaveFormat +_TESTS_DIR = Path(__file__).parent +_DATA_DIR = _TESTS_DIR / "data" +CONFIG_2D = PscPlotConfig(data_dir=_DATA_DIR / "test-2d") + +matplotlib.use("Agg") -def make_plot(args_list: list[str], data_dir: str | None = None): - """Parse CLI args, run the full pipeline, and return the initialized figure.""" - if data_dir is not None: - original_dir = CONFIG.data_dir - CONFIG.data_dir = _DATA_DIR / data_dir - try: - args = get_parsed_args(args_list) - plot = args.get_animation() - plot._initialize() - return plot.fig - finally: - if data_dir is not None: - CONFIG.data_dir = original_dir +def make_plot(args_list: list[str], data_dir: str = "test-2d"): + """Parse CLI args, run the full pipeline, and return the initialized figure.""" + args = parse_args(args_list) + plot = compile_plot_node(args, PscPlotConfig(data_dir=_DATA_DIR / data_dir)).pull() + plot._initialize() + return plot.fig -def make_save(args_list: list[str], save_dir: Path, format: SaveFormat, data_dir: str | None = None): +def make_save(args_list: list[str], save_dir: Path, format: SaveFormat, data_dir: str = "test-2d"): """Parse CLI args, run the full pipeline, and save to save_dir. Returns the output file path.""" - if data_dir is not None: - original_dir = CONFIG.data_dir - CONFIG.data_dir = _DATA_DIR / data_dir - - try: - args = get_parsed_args(args_list) - plot = args.get_animation() - save_dir.mkdir(exist_ok=True) - path = save_dir / f"{args.get_save_file_stem()}.{format}" - plot.save_to_path(path) - return path - finally: - if data_dir is not None: - CONFIG.data_dir = original_dir + args = parse_args(args_list) + node = compile_plot_node(args, PscPlotConfig(data_dir=_DATA_DIR / data_dir)) + plot = node.pull() + save_dir.mkdir(exist_ok=True) + path = save_dir / f"{node.get_save_file_stem()}.{format}" + plot.save_to_path(path) + return path @pytest.fixture(autouse=True) diff --git a/tests/test_dask_graph.py b/tests/test_dask_graph.py index 5fe69380..18da57f0 100644 --- a/tests/test_dask_graph.py +++ b/tests/test_dask_graph.py @@ -11,30 +11,28 @@ from conftest import _DATA_DIR -from lib.config import CONFIG -from lib.parsing.parse import get_parsed_args +from lib.config import PscPlotConfig +from lib.data.compile import compile_plot_node +from lib.parsing.parse import parse_args def _read_keys_for_columns(args_list: list[str], data_dir: str = "test-2d") -> list[str]: """Optimize each dask collection produced by `args_list` and return the set of per-column file-read task key strings in the optimized graph.""" - original = CONFIG.data_dir - CONFIG.data_dir = _DATA_DIR / data_dir - try: - args = get_parsed_args(args_list) - data = args.get_data() - collections = data.dask_collections() - assert collections, "expected particle pipeline to be dask-backed" - read_keys: list[str] = [] - for c in collections: - opt = c.optimize() if hasattr(c, "optimize") else c - for k in opt.__dask_graph__(): - key = k[0] if isinstance(k, tuple) else k - if isinstance(key, str) and "open_dataset" in key: - read_keys.append(key) - return read_keys - finally: - CONFIG.data_dir = original + config = PscPlotConfig(data_dir=_DATA_DIR / data_dir) + args = parse_args(args_list) + node = compile_plot_node(args, config) + data = node.input_node.pull().active_data + collections = data.dask_collections() + assert collections, "expected particle pipeline to be dask-backed" + read_keys: list[str] = [] + for c in collections: + opt = c.optimize() if hasattr(c, "optimize") else c + for k in opt.__dask_graph__(): + key = k[0] if isinstance(k, tuple) else k + if isinstance(key, str) and "open_dataset" in key: + read_keys.append(key) + return read_keys def test_particle_load_projects_columns_to_reads(): diff --git a/tests/test_h5_species_discovery.py b/tests/test_h5_species_discovery.py index be8a4342..317b3622 100644 --- a/tests/test_h5_species_discovery.py +++ b/tests/test_h5_species_discovery.py @@ -7,21 +7,14 @@ import pytest from synthetic_particles import write_step -from lib.config import CONFIG +from lib.config import PscPlotConfig from lib.data.loaders.particle_h5 import ParticleLoaderH5 -@pytest.fixture -def isolated_data_dir(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Path: - """Point CONFIG.data_dir at tmp_path for the duration of one test.""" - monkeypatch.setattr(CONFIG, "data_dir", tmp_path) - return tmp_path - - -def test_h5_species_discovery_standard(isolated_data_dir: Path): - write_step(isolated_data_dir / "prt.000000000.h5", time=0.0, species=[(-1.0, 1.0, 10), (1.0, 100.0, 10)], seed=0) +def test_h5_species_discovery_standard(tmp_path: Path): + write_step(tmp_path / "prt.000000000.h5", time=0.0, species=[(-1.0, 1.0, 10), (1.0, 100.0, 10)], seed=0) loader = ParticleLoaderH5("prt", active_key=None) - data = loader.get_data() + data = loader.get_data(PscPlotConfig(data_dir=tmp_path)) assert set(data.metadata.species.keys()) == {"e", "i"} e = data.metadata.species["e"] i = data.metadata.species["i"] @@ -29,43 +22,43 @@ def test_h5_species_discovery_standard(isolated_data_dir: Path): assert i.q == 1.0 and i.m == 100.0 -def test_h5_species_discovery_multiple_ion_masses(isolated_data_dir: Path): +def test_h5_species_discovery_multiple_ion_masses(tmp_path: Path): write_step( - isolated_data_dir / "prt.000000000.h5", + tmp_path / "prt.000000000.h5", time=0.0, species=[(-1.0, 1.0, 10), (1.0, 25.0, 10), (1.0, 100.0, 10)], seed=0, ) loader = ParticleLoaderH5("prt", active_key=None) - data = loader.get_data() + data = loader.get_data(PscPlotConfig(data_dir=tmp_path)) assert set(data.metadata.species.keys()) == {"e", "i25", "i100"} assert data.metadata.species["i25"].m == 25.0 assert data.metadata.species["i100"].m == 100.0 -def test_h5_species_discovery_multiple_ion_charges(isolated_data_dir: Path): +def test_h5_species_discovery_multiple_ion_charges(tmp_path: Path): write_step( - isolated_data_dir / "prt.000000000.h5", + tmp_path / "prt.000000000.h5", time=0.0, species=[(-1.0, 1.0, 10), (1.0, 100.0, 10), (2.0, 100.0, 10)], seed=0, ) loader = ParticleLoaderH5("prt", active_key=None) - data = loader.get_data() + data = loader.get_data(PscPlotConfig(data_dir=tmp_path)) assert set(data.metadata.species.keys()) == {"e", "i+", "i++"} assert data.metadata.species["i+"].q == 1.0 assert data.metadata.species["i++"].q == 2.0 -def test_h5_species_discovery_multiple_ion_everything(isolated_data_dir: Path): +def test_h5_species_discovery_multiple_ion_everything(tmp_path: Path): write_step( - isolated_data_dir / "prt.000000000.h5", + tmp_path / "prt.000000000.h5", time=0.0, species=[(-1.0, 1.0, 10), (1.0, 25.0, 10), (1.0, 100.0, 10), (2.0, 100.0, 10)], seed=0, ) loader = ParticleLoaderH5("prt", active_key=None) - data = loader.get_data() + data = loader.get_data(PscPlotConfig(data_dir=tmp_path)) assert set(data.metadata.species.keys()) == {"e", "i+25", "i+100", "i++100"} assert data.metadata.species["i+25"].q == 1.0 assert data.metadata.species["i+25"].m == 25.0 @@ -75,23 +68,23 @@ def test_h5_species_discovery_multiple_ion_everything(isolated_data_dir: Path): assert data.metadata.species["i++100"].m == 100.0 -def test_h5_species_discovery_electron_merge_warns(isolated_data_dir: Path): +def test_h5_species_discovery_electron_merge_warns(tmp_path: Path): write_step( - isolated_data_dir / "prt.000000000.h5", + tmp_path / "prt.000000000.h5", time=0.0, species=[(-1.0, 1.0, 10), (-1.0, 1.0, 10)], seed=0, ) loader = ParticleLoaderH5("prt", active_key=None) with pytest.warns(UserWarning, match="merging"): - data = loader.get_data() + data = loader.get_data(PscPlotConfig(data_dir=tmp_path)) assert set(data.metadata.species.keys()) == {"e"} -def test_h5_species_discovery_species_at_different_times(isolated_data_dir: Path): +def test_h5_species_discovery_species_at_different_times(tmp_path: Path): # step 0: only species 0 has particles; step 1: only species 1 has particles. - write_step(isolated_data_dir / "prt.000000000.h5", time=0.0, species=[(-1.0, 1.0, 10), (1.0, 1.0, 0)], seed=0) - write_step(isolated_data_dir / "prt.000000001.h5", time=1.0, species=[(-1.0, 1.0, 0), (1.0, 1.0, 10)], seed=1) + write_step(tmp_path / "prt.000000000.h5", time=0.0, species=[(-1.0, 1.0, 10), (1.0, 1.0, 0)], seed=0) + write_step(tmp_path / "prt.000000001.h5", time=1.0, species=[(-1.0, 1.0, 0), (1.0, 1.0, 10)], seed=1) loader = ParticleLoaderH5("prt", active_key=None) - data = loader.get_data() + data = loader.get_data(PscPlotConfig(data_dir=tmp_path)) assert set(data.metadata.species.keys()) == {"e", "i"} diff --git a/tests/test_idx_efficient.py b/tests/test_idx_efficient.py index 365fb76c..2733ca0b 100644 --- a/tests/test_idx_efficient.py +++ b/tests/test_idx_efficient.py @@ -10,8 +10,10 @@ from __future__ import annotations import pytest +from conftest import CONFIG_2D -from lib.parsing.parse import get_parsed_args +from lib.data.compile import compile_plot_node +from lib.parsing.parse import parse_args @pytest.fixture @@ -31,8 +33,8 @@ def counting_read(self: File, var_name: str, index): def test_field_idx_t(files_and_vars): - args = get_parsed_args("pfd ex_ec --idx t=-1 -v y z time= --compute".split()) - args.get_animation()._initialize() + args = parse_args("pfd ex_ec --idx t=-1 -v y z time= --compute".split()) + compile_plot_node(args, CONFIG_2D).pull()._initialize() # 'jeh' is the raw adios2 variable that holds all pfd components. files_read = {f for f, var in files_and_vars if var == "jeh"} @@ -40,8 +42,8 @@ def test_field_idx_t(files_and_vars): def test_particle_bp_idx_t(files_and_vars): - args = get_parsed_args("prt.e --idx t=-1 -v y z time= --compute".split()) - args.get_animation()._initialize() + args = parse_args("prt.e --idx t=-1 -v y z time= --compute".split()) + compile_plot_node(args, CONFIG_2D).pull()._initialize() # Particle position columns; if any of these is read from >1 file, the # loader is scanning steps it shouldn't. @@ -52,16 +54,16 @@ def test_particle_bp_idx_t(files_and_vars): def test_field_pos_t(files_and_vars): # t=999 is past max(t) in test-2d, so "nearest" resolves to the last file. - args = get_parsed_args("pfd ex_ec --pos t=999 -v y z time= --compute".split()) - args.get_animation()._initialize() + args = parse_args("pfd ex_ec --pos t=999 -v y z time= --compute".split()) + compile_plot_node(args, CONFIG_2D).pull()._initialize() files_read = {f for f, var in files_and_vars if var == "jeh"} assert len(files_read) == 1, f"--pos t=999 read 'jeh' from {len(files_read)} files; expected 1. files: {sorted(files_read)}" def test_particle_bp_pos_t(files_and_vars): - args = get_parsed_args("prt.e --pos t=999 -v y z time= --compute".split()) - args.get_animation()._initialize() + args = parse_args("prt.e --pos t=999 -v y z time= --compute".split()) + compile_plot_node(args, CONFIG_2D).pull()._initialize() position_vars = {"y", "z"} files_read = {f for f, var in files_and_vars if var in position_vars} diff --git a/tests/test_memory.py b/tests/test_memory.py index 7360a0b8..9d283339 100644 --- a/tests/test_memory.py +++ b/tests/test_memory.py @@ -14,23 +14,20 @@ import pytest from synthetic_particles import write_steps +from lib.config import PscPlotConfig -def _run_pipeline(data_dir: str, chunksize: int, result_queue: mp.Queue) -> None: - """Child-process entry point: configure env, run the pipeline, report ru_maxrss.""" - import os - - os.environ["PSC_PLOT_DATA_DIR"] = data_dir - os.environ["PSC_PLOT_DASK_NUM_WORKERS"] = "1" - os.environ["PSC_PLOT_DASK_CHUNK_SIZE"] = str(chunksize) +def _run_pipeline(data_dir: pathlib.Path, chunksize: int, result_queue: mp.Queue) -> None: + """Child-process entry point: configure env, run the pipeline, report ru_maxrss.""" import matplotlib matplotlib.use("Agg") - from lib.parsing.parse import get_parsed_args + from lib.data.compile import compile_plot_node + from lib.parsing.parse import parse_args - args = get_parsed_args("prt --species i --bin y py=16 -v y py".split()) - plot = args.get_animation() + args = parse_args("prt --species i --bin y py=16 -v y py".split()) + plot = compile_plot_node(args, PscPlotConfig(data_dir=data_dir, dask_chunk_size=chunksize)).pull() plot._initialize() peak = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss @@ -41,7 +38,7 @@ def _measure(data_dir: pathlib.Path, chunksize: int) -> int: """Run _run_pipeline in a child and return the reported peak ru_maxrss.""" ctx = mp.get_context("spawn") queue = ctx.Queue() - proc = ctx.Process(target=_run_pipeline, args=(str(data_dir), chunksize, queue)) + proc = ctx.Process(target=_run_pipeline, args=(data_dir, chunksize, queue)) proc.start() proc.join(timeout=120) if proc.exitcode != 0: diff --git a/tests/test_particle_bp_perf.py b/tests/test_particle_bp_perf.py index e1ff74fd..aaba829f 100644 --- a/tests/test_particle_bp_perf.py +++ b/tests/test_particle_bp_perf.py @@ -15,37 +15,34 @@ import pytest from synthetic_particles import write_steps, write_steps_bp +from lib.config import PscPlotConfig +from lib.data.compile import compile_plot_node -def _run_h5_pipeline(data_dir: str, result_queue: mp.Queue) -> None: - import os - os.environ["PSC_PLOT_DATA_DIR"] = data_dir - os.environ["PSC_PLOT_DASK_NUM_WORKERS"] = "1" +def _run_h5_pipeline(data_dir: pathlib.Path, result_queue: mp.Queue) -> None: import matplotlib matplotlib.use("Agg") - from lib.parsing.parse import get_parsed_args - args = get_parsed_args("prt --species i --bin y py -v y py".split()) - plot = args.get_animation() + from lib.parsing.parse import parse_args + + args = parse_args("prt --species i --bin y py -v y py".split()) + plot = compile_plot_node(args, PscPlotConfig(data_dir=data_dir)).pull() t0 = time.perf_counter() plot._initialize() elapsed = time.perf_counter() - t0 result_queue.put(elapsed) -def _run_bp_pipeline(data_dir: str, result_queue: mp.Queue) -> None: - import os - - os.environ["PSC_PLOT_DATA_DIR"] = data_dir - os.environ["PSC_PLOT_DASK_NUM_WORKERS"] = "1" +def _run_bp_pipeline(data_dir: pathlib.Path, result_queue: mp.Queue) -> None: import matplotlib matplotlib.use("Agg") - from lib.parsing.parse import get_parsed_args - args = get_parsed_args("prt.i --bin y py -v y py".split()) - plot = args.get_animation() + from lib.parsing.parse import parse_args + + args = parse_args("prt.i --bin y py -v y py".split()) + plot = compile_plot_node(args, PscPlotConfig(data_dir=data_dir)).pull() t0 = time.perf_counter() plot._initialize() elapsed = time.perf_counter() - t0 @@ -55,7 +52,7 @@ def _run_bp_pipeline(data_dir: str, result_queue: mp.Queue) -> None: def _measure(target, data_dir: pathlib.Path) -> float: ctx = mp.get_context("spawn") queue = ctx.Queue() - proc = ctx.Process(target=target, args=(str(data_dir), queue)) + proc = ctx.Process(target=target, args=(data_dir, queue)) proc.start() proc.join(timeout=300) if proc.exitcode != 0: diff --git a/tests/test_particle_bp_vs_h5.py b/tests/test_particle_bp_vs_h5.py index 58c0154d..ccfce164 100644 --- a/tests/test_particle_bp_vs_h5.py +++ b/tests/test_particle_bp_vs_h5.py @@ -3,6 +3,7 @@ import numpy as np import pytest +from conftest import CONFIG_2D from lib.data.adaptors.species_filter import SpeciesFilter from lib.data.loaders.particle_bp import ParticleLoaderBp @@ -11,13 +12,13 @@ def _load_and_filter_h5(species_key: str): loader = ParticleLoaderH5(prefix="prt", active_key=None) - data = loader.get_data() + data = loader.get_data(CONFIG_2D) return SpeciesFilter(species_key).apply_list(data) def _load_bp(species_key: str): loader = ParticleLoaderBp(prefix=f"prt.{species_key}", active_key=None) - return loader.get_data() + return loader.get_data(CONFIG_2D) @pytest.mark.parametrize("species_key", ["i", "e"]) diff --git a/tests/test_plots.py b/tests/test_plots.py index bd8e4782..c8fbcb03 100644 --- a/tests/test_plots.py +++ b/tests/test_plots.py @@ -73,6 +73,27 @@ def test_static_2d_spectogram(): return make_plot("pfd ey_ec -f y --mag --pow 2 --pos k_y=0: -v t k_y time= --nan0 --scale log".split()) +# --- Multiplots --- + + +@pytest.mark.mpl_image_compare(**MPL_KWARGS) +def test_multiline_different_vars(): + """Two line plots of different variables.""" + return make_plot("pfd hx_fc -v y -w hz_fc -v y".split()) + + +@pytest.mark.mpl_image_compare(**MPL_KWARGS) +def test_multiline_different_idxs(): + """Two line plots of the same variable at different slices.""" + return make_plot("pfd --derive hx_fc_copy=hx_fc -i z=0 --display B_x -v y -w hx_fc -i z=1 -v y".split()) + + +@pytest.mark.mpl_image_compare(**MPL_KWARGS) +def test_image_and_line(): + """An image overplotted with a line from a different file.""" + return make_plot("prt.i --bin y py=100 --nan0 --scale log --compute -v y py -w pfd::ey_ec -v y".split()) + + # --- Turbulence power spectrum --- diff --git a/tests/test_save_filename.py b/tests/test_save_filename.py index 4560f792..9b99aeca 100644 --- a/tests/test_save_filename.py +++ b/tests/test_save_filename.py @@ -1,13 +1,8 @@ import pytest +from conftest import CONFIG_2D -from lib.data.compile import compile_source -from lib.parsing.parse import get_parsed_args - - -def _stem(args_list: list[str]) -> str: - args = get_parsed_args(args_list) - compile_source(args.loader, args.adaptors) # mutates args.adaptors: appends default Versus - return args.get_save_file_stem() +from lib.data.compile import compile_plot_node +from lib.parsing.parse import parse_args @pytest.mark.parametrize( @@ -15,9 +10,10 @@ def _stem(args_list: list[str]) -> str: [ (["pfd", "hx_fc"], "pfd-hx_fc-v_y,z"), (["pfd", "hx_fc", "--nan0"], "pfd-hx_fc-nan0-v_y,z"), - (["pfd", "hx_fc", "--scale", "log"], "pfd-hx_fc-v_y,z-scale_log"), + (["pfd", "hx_fc", "--scale", "log"], "pfd-hx_fc-scale_log-v_y,z"), (["pfd", "hx_fc", "-v", "y", "z", "time="], "pfd-hx_fc-v_y,z;time="), ], ) def test_save_file_stem(args_list, expected_stem): - assert _stem(args_list) == expected_stem + actual_stem = compile_plot_node(parse_args(args_list), CONFIG_2D).get_save_file_stem() + assert actual_stem == expected_stem