diff --git a/src/lib/config.py b/src/lib/config.py index ae3e691..8eacccf 100644 --- a/src/lib/config.py +++ b/src/lib/config.py @@ -20,7 +20,7 @@ def parse_optional[T](s: str | None, parser: Callable[[str], T]) -> T | None: @dataclass class PscPlotConfig: _: KW_ONLY - data_dir: Path = field(default_factory=Path.cwd) + data_root: Path = field(default_factory=Path.cwd) ffmpeg_bin: Path | None = None dask_num_workers: int = 1 dask_chunk_size: int = 1_000_000 @@ -30,7 +30,7 @@ class PscPlotConfig: def from_env(cls) -> Self: config = cls() - config.data_dir = parse_optional(os.environ.get(_DATA_DIR_KEY), Path) or config.data_dir + config.data_root = parse_optional(os.environ.get(_DATA_DIR_KEY), Path) or config.data_root 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 diff --git a/src/lib/data/adaptor.py b/src/lib/data/adaptor.py index 801575f..4a7036f 100644 --- a/src/lib/data/adaptor.py +++ b/src/lib/data/adaptor.py @@ -29,7 +29,7 @@ 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)) + return world.with_active(data=self.apply(world.active_data)) def apply(self, data: DataWithAttrs) -> DataWithAttrs: if isinstance(data, List): @@ -59,25 +59,22 @@ def get_modified_unit_latex(self, metadata: Metadata) -> Latex: def apply(self, data: DataWithAttrs) -> DataWithAttrs: data = super().apply(data) - var_infos = data.metadata.var_infos - if data.metadata.active_key is not None and data.metadata.active_key in var_infos: + if info := data.active_info: display_latex = self.get_modified_display_latex(data.metadata) unit_latex = self.get_modified_unit_latex(data.metadata) - old_dim = var_infos[data.metadata.active_key] - new_dim = old_dim.assign(display=display_latex, unit=unit_latex) - var_infos = {**var_infos, data.metadata.active_key: new_dim} + data = data.with_active(info=info.assign(display=display_latex, unit=unit_latex)) - return data.assign_metadata(var_infos=var_infos) + return data class BareAdaptor(MetadataAdaptor): """An adaptor that works with the raw data, no metadata required.""" def apply_field(self, data: Field) -> DataWithAttrs: - return data.with_active_data(self.apply_field_bare(data.active_data)) + return data.with_active(data=self.apply_field_bare(data.require_active_subdata())) def apply_list(self, data: List) -> DataWithAttrs: - return data.with_active_data(self.apply_list_bare(data.active_data)) + return data.with_active(data=self.apply_list_bare(data.require_active_subdata())) def apply_field_bare(self, da: xr.DataArray) -> xr.DataArray: _fail_apply_field(self.__class__) diff --git a/src/lib/data/adaptors/bin.py b/src/lib/data/adaptors/bin.py index a75b969..aafc174 100644 --- a/src/lib/data/adaptors/bin.py +++ b/src/lib/data/adaptors/bin.py @@ -93,7 +93,7 @@ def apply_field(self, data: Field) -> Field: dim_names_to_bin_size[dim_name] = bin_size - return data.with_active_data(data.active_data.coarsen(dim_names_to_bin_size, boundary="pad").mean()) + return data.with_active(data=data.require_active_subdata().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) diff --git a/src/lib/data/adaptors/compute.py b/src/lib/data/adaptors/compute.py index 1fc252d..f9a4b83 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.with_active_data(data.active_data.compute()) + return data.with_active(data=data.require_active_subdata().compute()) diff --git a/src/lib/data/adaptors/derive.py b/src/lib/data/adaptors/derive.py index b576121..fd0399e 100644 --- a/src/lib/data/adaptors/derive.py +++ b/src/lib/data/adaptors/derive.py @@ -3,15 +3,9 @@ from lib import var_info_registry from lib.data.adaptor import WorldAdaptor -from lib.data.data_with_attrs import Field, List -from lib.data.loader import get_loader -from lib.derived_field_variables.derived_field_variable import ( - DERIVED_FIELD_VARIABLES, - derive_field_variable, -) -from lib.derived_particle_variables.derived_particle_variable import ( - derive_particle_variable, -) +from lib.data.data_world import DataWorld +from lib.data.ensure_derived import ensure_derived +from lib.data.loader import load from lib.parsing.args_registry import arg_parser @@ -21,66 +15,15 @@ def __init__(self, expression: str): self.ast = _DERIVE_PARSER.parse(expression) def apply_world(self, world): - active = world.active_data - scoped_prefixes = _collect_scoped_prefixes(self.ast) - - if isinstance(active, List): - if scoped_prefixes: - raise ValueError("--derive: cross-prefix references (prefix::key) are not supported for particle data.") - return world.with_active_data(AssignNewVariable(active).transform(self.ast)) - - if isinstance(active, Field): - siblings = _load_siblings(world, scoped_prefixes) - return world.with_active_data(AssignNewFieldVariable(active, siblings).transform(self.ast)) - - raise ValueError("--derive requires an active variable to derive into; specify one as a positional argument.") + return AssignNewVariable(world).transform(self.ast) def get_name_fragments(self): - if self.ast.data == "assign_default": - return [] return [f'derive_"{self.expression}"'] -def _collect_scoped_prefixes(ast) -> set[str]: - """Distinct prefixes referenced via `prefix::key` anywhere in the expression.""" - prefixes = set() - for tree in ast.find_data("variable"): - if len(tree.children) == 2: - prefixes.add(str(tree.children[0])) - return prefixes - - -def _load_siblings(world, prefixes: set[str]) -> dict[str, Field]: - """Resolve each scoped prefix to a loaded Field, reusing any already in the world - and auto-loading the rest the same way `--with` does.""" - siblings: dict[str, Field] = {} - for prefix in prefixes: - if prefix in world.datas: - siblings[prefix] = world.datas[prefix] - else: - loaded = get_loader(world.config.data_dir, prefix, None).apply_world(world) - siblings[prefix] = loaded.datas[prefix] - return siblings - - -def _resolve_field_variable(field: Field, key: str) -> Field: - """Return a Field guaranteed to contain `key`, deriving it via the prefix's - registry when it isn't already in the dataset.""" - if key in field.data: - return field - prefix = field.metadata.prefix - if prefix is None: - raise ValueError(f"--derive cannot resolve '{key}': field metadata has no prefix.") - if key not in DERIVED_FIELD_VARIABLES.get(prefix, {}): - raise ValueError( - f"--derive: '{key}' is not in the '{prefix}' dataset and not in its derived-variable registry {list(DERIVED_FIELD_VARIABLES.get(prefix, {}))}. Note that earlier adaptors (e.g. --downsample) may have dropped variables that became incompatible with the active grid; consider moving --derive earlier in the pipeline." - ) - return derive_field_variable(field, key, prefix) - - class AssignNewVariable(Transformer_InPlace): - def __init__(self, data: List): - self._data = data + def __init__(self, world: DataWorld): + self.world = world super().__init__(visit_tokens=True) def number(self, toks: list): @@ -89,66 +32,28 @@ def number(self, toks: list): def new_variable(self, toks: list): [tok] = toks - return tok + return str(tok) def variable(self, toks: list): - [tok] = toks - return self._data.data[tok] - - def addition(self, toks: list): - [lhs, rhs] = toks - return lhs + rhs - - def subtraction(self, toks: list): - [lhs, rhs] = toks - return lhs - rhs - - def multiplication(self, toks: list): - [lhs, rhs] = toks - return lhs * rhs - - def division(self, toks: list): - [lhs, rhs] = toks - return lhs / rhs - - def exponentiation(self, toks: list): - [lhs, rhs] = toks - return lhs**rhs - - def assign_default(self, toks: list): - [new_variable] = toks - return derive_particle_variable(self._data, new_variable, "prt") - - def assignment(self, toks: list): - [new_variable, val] = toks - df = self._data.data - df = df.assign(**{new_variable: val}) - return self._data.assign_data(df) - + key = str(toks[0]) + data = self.world.require_active_data() + data = ensure_derived(data, key) + return data[key] -class AssignNewFieldVariable(Transformer_InPlace): - def __init__(self, data: Field, siblings: dict[str, Field]): - self._data = data - self._siblings = siblings - super().__init__(visit_tokens=True) + def prepath(self, toks: list): + return "/".join(str(tok) for tok in toks) - def number(self, toks: list): - [tok] = toks - return float(tok) + def scoped_variable(self, toks: list): + [prepath, key] = [str(tok) for tok in toks] + if prepath not in self.world.datas: + data = load(self.world.config, prepath) + else: + data = self.world.datas[prepath] - def new_variable(self, toks: list): - [tok] = toks - return str(tok) + data = ensure_derived(data, key) + self.world = self.world.with_data(prepath, data) - def variable(self, toks: list): - if len(toks) == 2: - prefix, key = str(toks[0]), str(toks[1]) - sibling = _resolve_field_variable(self._siblings[prefix], key) - self._siblings[prefix] = sibling - return sibling.data[key] - key = str(toks[0]) - self._data = _resolve_field_variable(self._data, key) - return self._data.data[key] + return data[key] def addition(self, toks: list): [lhs, rhs] = toks @@ -170,32 +75,25 @@ def exponentiation(self, toks: list): [lhs, rhs] = toks return lhs**rhs - def assign_default(self, toks: list): - [new_variable] = toks - self._data = _resolve_field_variable(self._data, new_variable) - 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_metadata(active_key=new_variable, var_infos=new_var_infos) - def assignment(self, toks: list): - [new_variable, val] = toks - 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) + [key, subdata] = toks + data = self.world.require_active_data() + info = var_info_registry.lookup(data.metadata.prepath, key) + data = data.with_active(data=subdata, key=key, info=info) + return self.world.with_active(data=data) _DERIVE_GRAMMAR = r""" -?start : assign_default | assignment +?start : assignment -assign_default : new_variable -assignment : new_variable "=" expression -new_variable : CNAME +assignment : new_variable "=" expression +new_variable : CNAME ?expression : _expression_3 _expression_0 : "(" expression ")" | variable + | scoped_variable | number _expression_1 : _expression_0 | exponentiation @@ -212,8 +110,13 @@ def assignment(self, toks: list): addition : _expression_3 "+" _expression_2 subtraction : _expression_3 "-" _expression_2 -variable : (CNAME "::")? CNAME -number : SIGNED_NUMBER +variable : CNAME +scoped_variable : prepath "::" CNAME +number : SIGNED_NUMBER + +prepath : PREFIX +# DIR : /[^\/]+/ TODO figure out a way to make dirs work; as is, the arbitrary chars are incompatible with math symbols, especially / +PREFIX : /[.\w\d]+/ %import common.SIGNED_NUMBER %import common.CNAME @@ -223,16 +126,14 @@ def assignment(self, toks: list): _DERIVE_PARSER = Lark(_DERIVE_GRAMMAR) - -_DERIVE_FORMAT = "new_var_key[=expression]" -_EXPRESSION_DESCRIPTION = "The expression can be any mathematical expression using the standard operators (+, -, *, /, ^), parentheses, signed floating point numbers, and existing variable names. A name may be scoped to another prefix as prefix::key (that prefix is auto-loaded)." +_DERIVE_FORMAT = "new_var_key=expression" @arg_parser( dest="adaptors", flags="--derive", metavar=_DERIVE_FORMAT, - help=f"Create a new variable with the given name. {_EXPRESSION_DESCRIPTION} If the expression is omitted, the variable is derived via the registry of derivable variables.", + help=f"Create a new variable with the given name. The expression can be any mathematical expression using the standard operators (+, -, *, /, ^), parentheses, signed floating point numbers, and other variable names. A name may be scoped to another prefix as `prepath::key` (similar to `--with`).", ) def parse_derive(arg: str) -> Derive: return Derive(arg) diff --git a/src/lib/data/adaptors/display.py b/src/lib/data/adaptors/display.py index e757e08..4a4dfb6 100644 --- a/src/lib/data/adaptors/display.py +++ b/src/lib/data/adaptors/display.py @@ -6,27 +6,25 @@ class Display(Adaptor): """Override the display-LaTeX of the active variable or of a dimension.""" - def __init__(self, target: str | None, value: str): - self.target = target - self.value = value + def __init__(self, key: str | None, display: str): + self.key = key + self.display = display def apply(self, data: DataWithAttrs) -> DataWithAttrs: metadata = data.metadata - target = self.target or metadata.active_key - if target is None: + key = self.key or metadata.active_key + if key is None: raise ValueError("--display requires a target; specify a variable as a positional argument or use --display TARGET=VALUE") - if target not in metadata.var_infos: - raise ValueError(f"--display target {target!r} is not a known key ({sorted(metadata.var_infos)})") + if key not in metadata.var_infos: + raise ValueError(f"--display target {key!r} is not a known key ({sorted(metadata.var_infos)})") - old_dim = metadata.var_infos[target] - new_dim = old_dim.assign(display=self.value) - new_var_infos = {**metadata.var_infos, target: new_dim} - return data.assign_metadata(var_infos=new_var_infos) + info = metadata.var_infos[key].assign(display=self.display) + return data.with_info(key, info) def get_name_fragments(self) -> list[str]: - return [f"display_{self.target or 'active'}={self.value}"] + return [f"display_{self.key or 'active'}={self.display}"] _DISPLAY_FORMAT = "[name=]display_latex" @@ -40,6 +38,6 @@ def get_name_fragments(self) -> list[str]: ) def parse_display(arg: str) -> Display: if "=" in arg: - name, value = arg.split("=", 1) - return Display(target=name, value=value) - return Display(target=None, value=arg) + key, display = arg.split("=", 1) + return Display(key=key, display=display) + return Display(key=None, display=arg) diff --git a/src/lib/data/adaptors/fourier.py b/src/lib/data/adaptors/fourier.py index 60019c9..7121302 100644 --- a/src/lib/data/adaptors/fourier.py +++ b/src/lib/data/adaptors/fourier.py @@ -36,7 +36,7 @@ def __init__(self, dim_keys: str | list[str]): 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 + da = data.require_active_subdata() new_var_infos = data.metadata.var_infos.copy() for key in self.dim_keys: @@ -45,11 +45,11 @@ def apply_field(self, data: Field) -> Field: new_var_infos[f_info.key] = f_info da = toggle_fourier(da, info) - old_active_info = data.metadata.active_var_info + old_active_info = data.active_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) + return data.with_active(data=da).assign(var_infos=new_var_infos) def get_name_fragments(self) -> list[str]: return [f"fourier_{','.join(self.dim_keys)}"] diff --git a/src/lib/data/adaptors/idx.py b/src/lib/data/adaptors/idx.py index 2c08705..3f35d64 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.with_active_data(data.active_data.isel(self.dim_names_to_isel)) + return data.with_active(data=data.require_active_subdata().isel(self.dim_names_to_isel)) def apply_list(self, data: List) -> List: coordss = data.coordss.copy() diff --git a/src/lib/data/adaptors/pos.py b/src/lib/data/adaptors/pos.py index cfdf204..ce0718b 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.with_active_data(data.active_data.sel(dim_names_to_pos, method="nearest").sel(dim_names_to_slice)) + return data.with_active(data=data.require_active_subdata().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. @@ -60,7 +60,7 @@ def apply_list(self, data: List) -> List: df = df[df[dim] >= sel.start] if inc_lo else df[df[dim] > sel.start] if sel.stop is not None: df = df[df[dim] <= sel.stop] if inc_hi else df[df[dim] < sel.stop] - data = data.assign_data(df) + data = data.assign(df) return data diff --git a/src/lib/data/adaptors/scatter.py b/src/lib/data/adaptors/scatter.py index eb93e16..15af561 100644 --- a/src/lib/data/adaptors/scatter.py +++ b/src/lib/data/adaptors/scatter.py @@ -17,14 +17,14 @@ def apply_field(self, data: Field) -> LazyList: ordered_coordss = [coordss[dim] for dim in data.dims] coord_grids = np.meshgrid(*ordered_coordss, indexing="ij") - vars = {data.metadata.active_key: da.ravel(data.active_data)} | dict(zip(data.dims, (da.ravel(coord_grid) for coord_grid in coord_grids))) + vars = {data.active_key: da.ravel(data.require_active_subdata())} | dict(zip(data.dims, (da.ravel(coord_grid) for coord_grid in coord_grids))) df = dd.from_dask_array(da.vstack(vars.values()).T, columns=list(vars)) metadata = ListMetadata.create_from( data.metadata, coordss=coordss, - weight_key=data.metadata.active_key, + weight_key=data.active_key, subject=self.subject, ) diff --git a/src/lib/data/adaptors/set_scale.py b/src/lib/data/adaptors/set_scale.py index fffbdd1..16686c0 100644 --- a/src/lib/data/adaptors/set_scale.py +++ b/src/lib/data/adaptors/set_scale.py @@ -8,18 +8,17 @@ class SetScale(MetadataAdaptor): - def __init__(self, dim_name: str | None, scale: Scale): - self.dim_name = dim_name + def __init__(self, key: str | None, scale: Scale): + self.key = key 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) + key = self.key or data.metadata.active_key + info = replace(data.metadata.var_infos[key], scale=self.scale) + return data.with_info(key, info) def get_name_fragments(self) -> list[str]: - maybe_dim_name = f"{self.dim_name}=" if self.dim_name is not None else "" + maybe_dim_name = f"{self.key}=" if self.key is not None else "" return [f"scale_{maybe_dim_name}{self.scale.to_name_fragment_part()}"] diff --git a/src/lib/data/adaptors/species_filter.py b/src/lib/data/adaptors/species_filter.py index 7df36f8..1bab816 100644 --- a/src/lib/data/adaptors/species_filter.py +++ b/src/lib/data/adaptors/species_filter.py @@ -12,15 +12,18 @@ def apply_list(self, data: List) -> List: if info is None: available = sorted(data.metadata.species.keys()) raise ValueError(f"unknown species {self.species_key!r}; available: {available}") - data = data.assign_metadata(subject=info.display) + + data = data.assign(subject=info.display) + if len(data.metadata.species) == 1: # Dataset is already single-species (e.g. BP per-species file, or H5 # with one species). Row filter is a no-op; skip to avoid requiring # q/m columns that BP data doesn't have. return data + df = data.data df = df[(df["q"] == info.q) & (df["m"] == info.m)] - return data.assign_data(df).assign_metadata(subject=info.display) + return data.assign(df) def get_name_fragments(self) -> list[str]: return [self.species_key] diff --git a/src/lib/data/adaptors/transform_polar.py b/src/lib/data/adaptors/transform_polar.py index 968c3a4..c152e02 100644 --- a/src/lib/data/adaptors/transform_polar.py +++ b/src/lib/data/adaptors/transform_polar.py @@ -68,7 +68,7 @@ def apply_field(self, data: Field) -> Field: xgrid = xr.Variable([key_r, key_theta], xgrid) ygrid = xr.Variable([key_r, key_theta], ygrid) - da = data.active_data + da = data.require_active_subdata() da = da.interp({key_x: xgrid, key_y: ygrid}, assume_sorted=True) da = da.drop_vars([key_x, key_y]) da = da.assign_coords({key_r: rs, key_theta: thetas}) @@ -76,7 +76,7 @@ def apply_field(self, data: Field) -> Field: new_var_infos = {k: v for k, v in data.metadata.var_infos.items() if k not in {key_x, key_y}} new_var_infos[key_r] = dim_r new_var_infos[key_theta] = dim_theta - return data.with_active_data(da).assign_metadata(var_infos=new_var_infos) + return data.with_active(data=da).assign(var_infos=new_var_infos) def apply_list(self, data: List) -> List: dim_x = data.metadata.var_infos[self.dim1_key] @@ -93,7 +93,7 @@ def apply_list(self, data: List) -> List: new_var_infos = dict(data.metadata.var_infos) new_var_infos[key_r] = dim_r new_var_infos[key_theta] = dim_theta - return data.assign_data(df).assign_metadata(var_infos=new_var_infos) + return data.assign(df, var_infos=new_var_infos) def get_name_fragments(self) -> list[str]: return [f"polar_{self.dim1_key},{self.dim2_key}"] diff --git a/src/lib/data/adaptors/transform_spherical.py b/src/lib/data/adaptors/transform_spherical.py index e597f4d..a7e6da7 100644 --- a/src/lib/data/adaptors/transform_spherical.py +++ b/src/lib/data/adaptors/transform_spherical.py @@ -86,7 +86,7 @@ def apply_field(self, data: Field) -> Field: ygrid = xr.Variable([key_r, key_theta, key_phi], ygrid) zgrid = xr.Variable([key_r, key_theta, key_phi], zgrid) - da = data.active_data + da = data.require_active_subdata() da = da.interp({key_x: xgrid, key_y: ygrid, key_z: zgrid}, assume_sorted=True) da = da.drop_vars([key_x, key_y, key_z]) da = da.assign_coords({key_r: rs, key_theta: thetas, key_phi: phis}) @@ -95,7 +95,7 @@ def apply_field(self, data: Field) -> Field: new_var_infos[key_r] = dim_r new_var_infos[key_theta] = dim_theta new_var_infos[key_phi] = dim_phi - return data.with_active_data(da).assign_metadata(var_infos=new_var_infos) + return data.with_active(data=da).assign(var_infos=new_var_infos) def apply_list(self, data: List) -> List: dim_x = data.metadata.var_infos[self.dim1_key] @@ -114,7 +114,7 @@ def apply_list(self, data: List) -> List: new_var_infos[key_r] = dim_r new_var_infos[key_theta] = dim_theta new_var_infos[key_phi] = dim_phi - return data.assign_data(df).assign_metadata(var_infos=new_var_infos) + return data.assign(df, var_infos=new_var_infos) def get_name_fragments(self) -> list[str]: return [f"spherical_{self.dim1_key},{self.dim2_key},{self.dim3_key}"] diff --git a/src/lib/data/adaptors/unit.py b/src/lib/data/adaptors/unit.py index 45d7e8f..6cc5833 100644 --- a/src/lib/data/adaptors/unit.py +++ b/src/lib/data/adaptors/unit.py @@ -6,27 +6,25 @@ class Unit(Adaptor): """Override the unit-LaTeX of the active variable or of a dimension.""" - def __init__(self, target: str | None, value: str): - self.target = target - self.value = value + def __init__(self, key: str | None, unit: str): + self.key = key + self.unit = unit def apply(self, data: DataWithAttrs) -> DataWithAttrs: metadata = data.metadata - target = self.target or metadata.active_key - if target is None: + key = self.key or metadata.active_key + if key is None: raise ValueError("--unit requires a target; specify a variable as a positional argument or use --unit TARGET=VALUE") - if target not in metadata.var_infos: - raise ValueError(f"--unit target {target!r} is not a known key ({sorted(metadata.var_infos)})") + if key not in metadata.var_infos: + raise ValueError(f"--unit target {key!r} is not a known key ({sorted(metadata.var_infos)})") - old_dim = metadata.var_infos[target] - new_dim = old_dim.assign(unit=self.value) - new_var_infos = {**metadata.var_infos, target: new_dim} - return data.assign_metadata(var_infos=new_var_infos) + info = metadata.var_infos[key].assign(unit=self.unit) + return data.with_info(key, info) def get_name_fragments(self) -> list[str]: - return [f"unit_{self.target or 'active'}={self.value}"] + return [f"unit_{self.key or 'active'}={self.unit}"] _UNIT_FORMAT = "[name=]unit_latex" @@ -40,6 +38,6 @@ def get_name_fragments(self) -> list[str]: ) def parse_unit(arg: str) -> Unit: if "=" in arg: - name, value = arg.split("=", 1) - return Unit(target=name, value=value) - return Unit(target=None, value=arg) + key, unit = arg.split("=", 1) + return Unit(key=key, unit=unit) + return Unit(key=None, unit=arg) diff --git a/src/lib/data/adaptors/versus.py b/src/lib/data/adaptors/versus.py index b269031..ffb964d 100644 --- a/src/lib/data/adaptors/versus.py +++ b/src/lib/data/adaptors/versus.py @@ -2,7 +2,6 @@ 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 @@ -24,20 +23,6 @@ def __init__( 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() - - if time_dim := self._get_time_dim(data): - retained_dims.append(time_dim) - - if self.color_dim: - retained_dims.append(self.color_dim) - - if isinstance(data, List) and data.metadata.active_key: - retained_dims.append(data.metadata.active_key) - - return retained_dims - def _get_time_dim(self, data: DataWithAttrs) -> str | None: if self.time_dim_rule != "guess": return self.time_dim_rule @@ -72,7 +57,7 @@ def _get_color_dim(self, data: DataWithAttrs) -> str | None: def apply_world(self, world): data = self.apply(world.active_data) new_plot_target = PlotTarget( - world.active_key, + world.active_prepath, spatial_dims=self._get_spatial_dims(data), color_dim=self._get_color_dim(data), time_dim=self._get_time_dim(data), @@ -81,51 +66,16 @@ def apply_world(self, world): return replace( world, plot_targets=world.plot_targets + [new_plot_target], - datas=world.datas | {world.active_key: data}, + datas=world.datas | {world.active_prepath: data}, ) def apply_field(self, data: Field) -> Field: - # 1. apply implicit coordinate transforms, as necessary - retained_dims = self._get_retained_dim_keys(data) - for dim_name in retained_dims: - # 1a. already have the coordinate; do nothing - if dim_name in data.dims: - continue - - # 1b. need to do a Fourier transform - dim = data.metadata.var_infos[dim_name] - f_dim = dim.toggle_fourier() - if f_dim.key in data.dims: - data = data.assign_metadata( - var_infos={**data.metadata.var_infos, f_dim.key: f_dim}, - ) - fourier = Fourier(f_dim.key) - data = fourier.apply(data) - continue - - # 1c. need to do a coordinate transform - # TODO - - # 2. reduce remaining dimensions via arithmetic mean - reduce_dims = [dim for dim in data.dims if dim not in retained_dims] + used_dims = {*self.spatial_dims, self._get_time_dim(data), self.color_dim} + reduce_dims = [dim for dim in data.dims if dim not in used_dims] # preserve order reduce = Reduce(reduce_dims, "mean") - data = reduce.apply(data) - - return data + return reduce.apply(data) def apply_list(self, data: List) -> List: - # 1. coordinate transform - # TODO - - # 2. drop unused vars - keep_vars = self._get_retained_dim_keys(data) - drop_vars = [active_key for active_key in data.dims if active_key not in keep_vars] - data = data.assign_data(data.data.drop(columns=drop_vars)) - - spatial_dims = self.spatial_dims.copy() - 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 def get_name_fragments(self) -> list[str]: diff --git a/src/lib/data/adaptors/with_.py b/src/lib/data/adaptors/with_.py index 01d5e9b..4d5ddd4 100644 --- a/src/lib/data/adaptors/with_.py +++ b/src/lib/data/adaptors/with_.py @@ -1,55 +1,59 @@ +from lib.config import _DATA_DIR_KEY from lib.data.adaptor import WorldAdaptor -from lib.data.loader import get_loader +from lib.data.ensure_derived import ensure_derived +from lib.data.loader import load +from lib.file_util import Prepath from lib.parsing import parse_util from lib.parsing.args_registry import arg_parser class With(WorldAdaptor): - def __init__(self, prefix_or_key: str, key: str | None = None): - self.prefix_or_key = prefix_or_key + def __init__(self, prepath: Prepath | None, key: str | None = None, *, include_with_in_name_fragment: bool = True): + self.prepath = prepath self.key = key + self.include_with_in_name_fragment = include_with_in_name_fragment 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)) + if not self.prepath: + data = world.require_active_data() + elif self.prepath in world.datas: + data = world.datas[self.prepath] + else: + data = load(world.config, self.prepath) - # case 2: prefix_or_key is a prefix - prefix = self.prefix_or_key - key = self.key + if self.key: + data = ensure_derived(data, self.key) + data = data.with_active(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) + return world.with_active(prepath=self.prepath, data=data) def get_name_fragments(self) -> list[str]: - maybe_prefix = f"{self.prefix_or_key}{SCOPE_OP}" if self.prefix_or_key else "" - return [f"with_{maybe_prefix}{self.key or ''}"] + maybe_with = "with_" if self.include_with_in_name_fragment else "" + maybe_prepath = f"{self.prepath}{SCOPE_OP}" if self.prepath else "" + return [f"{maybe_with}{maybe_prepath}{self.key or ''}"] SCOPE_OP = "::" -WITH_FORMAT = f"prefix[{SCOPE_OP}[key]] | key" +WITH_FORMAT = f"[prepath{SCOPE_OP}][var_key]" @arg_parser( dest="adaptors", flags=["--with", "-w"], metavar=WITH_FORMAT, - help="switch to a different prefix and/or variable", + help=f"Switch to a different prepath (e.g. `run1/pfd`, relative to {_DATA_DIR_KEY}) and/or variable (e.g. `ey_ec`).", nargs="just one", ) def parse_with(arg: str) -> With: 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) + [prepath, key_arg] = split_arg elif len(split_arg) == 1: - prefix_or_key = parse_util.parse_identifier(split_arg[0], "prefix | key") - return With(prefix_or_key) + prepath = None + [key_arg] = split_arg else: parse_util.fail_format(arg, WITH_FORMAT) + + key = parse_util.parse_optional_identifier(key_arg, "key") + return With(prepath, key) diff --git a/src/lib/data/compile.py b/src/lib/data/compile.py index a5cc820..6743e30 100644 --- a/src/lib/data/compile.py +++ b/src/lib/data/compile.py @@ -3,7 +3,7 @@ from lib.config import PscPlotConfig from lib.data.adaptor import Adaptor from lib.data.adaptors.versus import Versus -from lib.data.loader import get_loader +from lib.data.adaptors.with_ import With from lib.data.node import AdaptorNode, DaskGraphNode, DataProcessingNode, PlotNode, RootNode, SavePlotNode, ShowPlotNode from lib.parsing.args import Args @@ -21,7 +21,7 @@ def _with_versus(adaptors: list[Adaptor]) -> list[Adaptor]: def compile_data_node(args: Args, config: PscPlotConfig): node = RootNode(config) - node = AdaptorNode(node, get_loader(config.data_dir, args.prefix, args.variable)) + node = AdaptorNode(node, With(args.prepath, args.variable, include_with_in_name_fragment=False)) for adaptor in _with_versus(args.adaptors): node = AdaptorNode(node, adaptor) diff --git a/src/lib/data/data_with_attrs.py b/src/lib/data/data_with_attrs.py index 7944b09..14569ee 100644 --- a/src/lib/data/data_with_attrs.py +++ b/src/lib/data/data_with_attrs.py @@ -3,7 +3,7 @@ from abc import ABC, abstractmethod from dataclasses import dataclass, field, fields from functools import cached_property -from typing import Any, Callable, Self +from typing import Any, Self import dask.array import dask.dataframe as dd @@ -11,6 +11,7 @@ import pandas as pd import xarray as xr +from lib.file_util import Prepath from lib.latex import Latex from lib.species import SpeciesInfo from lib.var_info import VarInfo @@ -18,6 +19,7 @@ @dataclass(kw_only=True, frozen=True) class Metadata: + prepath: Prepath active_key: str | None = None var_infos: dict[str, VarInfo] = field(default_factory=dict) @@ -57,23 +59,13 @@ def assign(self, **vals: Any) -> Self: return self.__class__(**updated_vals) -@dataclass(frozen=True, init=False) -class DataWithAttrs[D: dict[str, xr.DataArray] | pd.DataFrame | dd.DataFrame, MD: Metadata](ABC): +@dataclass(frozen=True) +class DataWithAttrs[Data, Subdata, MD: Metadata = 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). - # Specifying field data types like this makes their types known in such cases, but doesn't give type hints for the parameters of __init__. - # 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: dict[str, xr.DataArray] | pd.DataFrame | dd.DataFrame - metadata: Metadata - _caches: dict[str, dict[str, Any]] - - def __init__(self, data: D, metadata: MD): - object.__setattr__(self, "data", data) - object.__setattr__(self, "metadata", metadata) - object.__setattr__(self, "_caches", {}) + data: Data + metadata: MD + _caches: dict[str, dict[str, Any]] = field(default_factory=dict, init=False) @property @abstractmethod @@ -83,19 +75,38 @@ def coordss(self) -> dict[str, np.ndarray]: ... @abstractmethod def dims(self) -> list[str]: ... - def assign_data(self, data: D) -> Self: - return self.__class__(data, self.metadata) + @property + def active_key(self) -> str | None: + return self.metadata.active_key + + @property + def active_subdata(self) -> Subdata | None: + if self.active_key is None: + return None + return self[self.active_key] - def assign_metadata(self, metadata: MD | None = None, /, **metadata_vals: Any) -> Self: - if not (metadata or metadata_vals): - return self - return self.__class__(self.data, (metadata or self.metadata).assign(**metadata_vals)) + def require_active_subdata(self) -> Subdata: + if self.active_key is None: + raise ValueError("No active variable.") + return self[self.active_key] - def assign(self, data: D, metadata: MD | None = None, /, **metadata_vals: Any) -> Self: - return self.assign_data(data).assign_metadata(metadata, **metadata_vals) + @property + def active_info(self) -> VarInfo | None: + if self.active_key is None: + return None + return self.metadata.var_infos[self.active_key] - def map_data(self, func: Callable[[D], D]) -> Self: - return self.assign_data(func(self.data)) + @abstractmethod + def __getitem__(self, key: str) -> Subdata: ... + + @abstractmethod + def with_active(self, *, data: Subdata | None = None, key: str | None = None, info: VarInfo | None = None) -> Self: ... + + def with_info(self, key: str, info: VarInfo) -> Self: + return self.assign(var_infos=self.metadata.var_infos | {key: info}) + + def assign(self, data: Data | None = None, /, **metadata_vals: Any) -> Self: + return self.__class__(self.data if data is None else data, self.metadata.assign(**metadata_vals)) @abstractmethod def bounds(self, dim_name: str) -> tuple[float, float]: ... @@ -110,33 +121,37 @@ def upper_bound(self, dim_name: str) -> float: ... def dask_collections(self) -> list: ... -@dataclass(kw_only=True, frozen=True) -class FieldMetadata(Metadata): - prefix: str | None = None +class FieldMetadata(Metadata): ... -class Field(DataWithAttrs[dict[str, xr.DataArray], FieldMetadata]): - data: dict[str, xr.DataArray] - metadata: FieldMetadata +class Field(DataWithAttrs[dict[str, xr.DataArray], xr.DataArray, FieldMetadata]): + def with_active(self, *, data=None, key=None, info=None) -> Self: + if data is None and info is None: + return self.assign(active_key=key) - @property - def active_data(self) -> xr.DataArray: - if self.metadata.active_key is None: - raise ValueError("no active variable; specify one as a positional argument") - return self.data[self.metadata.active_key] + key = key or self.metadata.active_key + assert key is not None + + ret = self + if data is not None: + ret = ret.assign(ret.data | {key: data}, active_key=key) + + if info is not None: + ret = ret.assign(var_infos=ret.metadata.var_infos | {key: info}, active_key=key) - def with_active_data(self, new_da: xr.DataArray) -> Self: - """Returns a shallow copy with the active variable replaced by `new_da`.""" - return self.assign_data(self.data | {self.metadata.active_key: new_da}) + return ret + + def __getitem__(self, key: str): + return self.data[key] @cached_property def coordss(self) -> dict[str, np.ndarray]: - active = self.active_data + active = self.require_active_subdata() return {dim: np.array(active.coords[dim]) for dim in active.coords.keys()} @cached_property def dims(self) -> list[str]: - return list(self.active_data.dims) + return list(self.require_active_subdata().dims) def bounds(self, dim_name): return (self.lower_bound(dim_name), self.upper_bound(dim_name)) @@ -151,7 +166,7 @@ def upper_bound(self, dim_name) -> float: @cached_property def var_bounds(self) -> tuple[float, float]: - active = self.active_data + active = self.require_active_subdata() return dask.compute(np.min(active), np.max(active)) def dask_collections(self) -> list: @@ -177,18 +192,22 @@ class ListMetadata(Metadata): `len(partition_ranges) == len(coordss[partition_dim])`.""" -class List[D: pd.DataFrame | dd.DataFrame](DataWithAttrs[D, ListMetadata]): - data: pd.DataFrame | dd.DataFrame - metadata: ListMetadata +class List[Data: pd.DataFrame | dd.DataFrame = pd.DataFrame | dd.DataFrame, Subdata: pd.Series | dd.Series = pd.Series | dd.Series](DataWithAttrs[Data, Subdata, ListMetadata]): + def with_active(self, *, data=None, key=None, info=None) -> Self: + if data is None and info is None: + return self.assign(active_key=key) - @property - def active_data(self) -> pd.Series | dd.Series: - if self.metadata.active_key is None: - raise ValueError("no active variable; specify one as a positional argument") - return self.data[self.metadata.active_key] + key = key or self.metadata.active_key + assert key is not None + + ret = self + if data is not None: + ret = ret.assign(ret.data.assign(**{key: data}), active_key=key) + + if info is not None: + ret = ret.assign(var_infos=ret.metadata.var_infos | {key: info}, active_key=key) - def with_active_data(self, series: pd.Series | dd.Series) -> Self: - return self.assign_data(self.data.assign(**{self.metadata.active_key: series})) + return ret @abstractmethod def compute(self) -> FullList: ... @@ -203,7 +222,8 @@ def dims(self) -> list[str]: class FullList(List[pd.DataFrame]): - data: pd.DataFrame + def __getitem__(self, key: str): + return self.data[key] def compute(self) -> FullList: return self @@ -236,7 +256,8 @@ def dask_collections(self) -> list: class LazyList(List[dd.DataFrame]): - data: dd.DataFrame + def __getitem__(self, key: str): + return self.data[key] def compute(self) -> FullList: # partition_* describe the dask layout; meaningless after compute. diff --git a/src/lib/data/data_world.py b/src/lib/data/data_world.py index 6850703..ce7f8f2 100644 --- a/src/lib/data/data_world.py +++ b/src/lib/data/data_world.py @@ -5,37 +5,51 @@ from lib.config import PscPlotConfig from lib.data.data_with_attrs import DataWithAttrs from lib.data.plot_target import PlotTarget +from lib.file_util import Prepath @dataclass(frozen=True) class DataWorld: # TODO python 3.15: make frozendict - datas: dict[str, DataWithAttrs] = field(default_factory=dict) - active_key: str | None = None + datas: dict[Prepath, DataWithAttrs] = field(default_factory=dict) + active_prepath: Prepath | 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 + assert self.active_prepath is None or self.active_prepath in self.datas @property def active_data(self) -> DataWithAttrs | None: - if self.active_key is None: + if self.active_prepath is None: return None - return self.datas[self.active_key] + return self.datas[self.active_prepath] - def with_active_data( + def require_active_data(self) -> DataWithAttrs: + if self.active_prepath is None: + raise ValueError("no active dataset; specify one as a positional argument") + return self.datas[self.active_prepath] + + def with_active( self, - active_data: DataWithAttrs | None = None, - active_key: str | None = None, + *, + data: DataWithAttrs | None = None, + prepath: Prepath | None = None, ) -> DataWorld: - if active_data is None: - return replace(self, active_key=active_key) + if data is None: + return replace(self, active_prepath=prepath) - active_key = active_key or self.active_key - assert active_key is not None + prepath = prepath or self.active_prepath + assert prepath is not None new_datas = self.datas.copy() - new_datas[active_key] = active_data - return replace(self, datas=new_datas, active_key=active_key) + new_datas[prepath] = data + return replace(self, datas=new_datas, active_prepath=prepath) + + def with_data( + self, + prepath: Prepath, + data: DataWithAttrs, + ) -> DataWorld: + return replace(self, datas=self.datas | {prepath: data}) diff --git a/src/lib/data/ensure_derived.py b/src/lib/data/ensure_derived.py new file mode 100644 index 0000000..2803928 --- /dev/null +++ b/src/lib/data/ensure_derived.py @@ -0,0 +1,14 @@ +from lib.data.data_with_attrs import DataWithAttrs, Field, List +from lib.derived_field_variables.derived_field_variable import derive_field_variable +from lib.derived_particle_variables.derived_particle_variable import derive_particle_variable +from lib.file_util import split_prepath + + +def ensure_derived[D: DataWithAttrs](data: D, key: str) -> D: + _, prefix = split_prepath(data.metadata.prepath) + if isinstance(data, Field): + return derive_field_variable(data, key, prefix) + elif isinstance(data, List): + return derive_particle_variable(data, key, prefix) + else: + raise TypeError(data.__class__) diff --git a/src/lib/data/loader.py b/src/lib/data/loader.py index cd5432c..9c04f68 100644 --- a/src/lib/data/loader.py +++ b/src/lib/data/loader.py @@ -5,6 +5,7 @@ from lib.config import PscPlotConfig from lib.data.adaptor import WorldAdaptor from lib.data.data_with_attrs import DataWithAttrs +from lib.file_util import Prepath, split_prepath class Loader(WorldAdaptor): @@ -18,18 +19,15 @@ def discover_prefixes(cls, data_dir: Path) -> list[str]: def suffix(cls) -> str: """Return the suffix that this loader supports.""" - def __init__(self, prefix: str, active_key: str | None = None): - self.prefix = prefix - self.active_key = active_key + def __init__(self, prepath: Prepath): + self.prepath = prepath + self.subdir, self.prefix = split_prepath(prepath) def get_name_fragments(self) -> list[str]: - fragments = [self.prefix] - if self.active_key is not None: - fragments.append(self.active_key) - return fragments + return [self.prepath] def apply_world(self, world): - return world.with_active_data(self.get_data(world.config), self.prefix) + return world.with_data(self.prepath, self.get_data(world.config)) @abstractmethod def get_data(self, config: PscPlotConfig) -> DataWithAttrs: ... @@ -62,6 +60,12 @@ def discover_loaders(data_dir: Path) -> dict[str, type[Loader]]: 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) +def get_loader(data_root: Path, prepath: Prepath) -> Loader: + subdir, prefix = split_prepath(prepath) + loader_types = discover_loaders(data_root / subdir) + return loader_types[prefix](prepath) + + +def load(config: PscPlotConfig, prepath: Prepath) -> DataWithAttrs: + loader = get_loader(config.data_root, prepath) + return loader.get_data(config) diff --git a/src/lib/data/loaders/field_bp.py b/src/lib/data/loaders/field_bp.py index c6b6117..f915a65 100644 --- a/src/lib/data/loaders/field_bp.py +++ b/src/lib/data/loaders/field_bp.py @@ -8,7 +8,6 @@ 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 from lib.var_info_registry import lookup _KNOWN_PREFIXES = ("pfd", "pfd_moments", "gauss", "continuity") @@ -36,7 +35,7 @@ def suffix(cls): def get_data(self, config: PscPlotConfig) -> Field: ds = xr.open_mfdataset( - paths=[_get_path(config.data_dir, self.prefix, step) for step in file_util.get_available_steps(config.data_dir, self.prefix + ".", ".bp")], + paths=[_get_path(config.data_root / self.subdir, self.prefix, step) for step in file_util.get_available_steps(config.data_root / self.subdir, self.prefix + ".", ".bp")], combine="nested", concat_dim="t", preprocess=_decode_psc, @@ -46,16 +45,10 @@ def get_data(self, config: PscPlotConfig) -> Field: data = {key: ds[key] for key in ds.data_vars} var_infos = {key: lookup(self.prefix, key) for key in ds.variables} - field = Field( + return Field( data, FieldMetadata( - active_key=self.active_key, - prefix=self.prefix, + prepath=self.prepath, var_infos=var_infos, ), ) - - 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 1908e41..471c696 100644 --- a/src/lib/data/loaders/particle_bp.py +++ b/src/lib/data/loaders/particle_bp.py @@ -87,13 +87,13 @@ def discover_prefixes(cls, data_dir: pathlib.Path) -> list[str]: def suffix(cls): return "bp" - def __init__(self, prefix: str, active_key: str | None = None): - super().__init__(prefix, active_key) - self.species_key = prefix.split(".", 1)[1] + def __init__(self, prepath: file_util.Prepath): + super().__init__(prepath) + self.species_key = prepath.split(".", 1)[1] 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] + steps = file_util.get_available_steps(config.data_root / self.subdir, self.prefix + ".", ".bp") + step_attrs = [_read_attrs(_get_path(config.data_root / self.subdir, self.prefix, step)) for step in steps] times = np.array([float(a["time"]) for a in step_attrs]) head = step_attrs[0] @@ -126,7 +126,7 @@ def get_data(self, config: PscPlotConfig) -> LazyList: partition_ranges = [] offset = 0 for step, time in zip(steps, times): - path = _get_path(config.data_dir, self.prefix, step) + path = _get_path(config.data_root / self.subdir, 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)) @@ -138,7 +138,7 @@ def get_data(self, config: PscPlotConfig) -> LazyList: slices.append(slice(i * chunk_size, (i + 1) * chunk_size)) meta = _build_meta(paths[0]) - df = dd.from_map(_read_chunk, paths, step_times, particle_dims, slices, meta=meta) + df: dd.DataFrame = dd.from_map(_read_chunk, paths, step_times, particle_dims, slices, meta=meta) corners = np.asarray(head["corner"]) lengths = np.asarray(head["length"]) @@ -147,19 +147,13 @@ def get_data(self, config: PscPlotConfig) -> LazyList: coordss["t"] = times metadata = ListMetadata( + prepath=self.prepath, weight_key="w", coordss=coordss, species=species_dict, subject=info.display, partition_dim="t", partition_ranges=partition_ranges, + var_infos={key: lookup("prt", key) for key in df.columns}, ) - data = LazyList(df, metadata) - - # var_info registry is keyed by "prt" (not per-species), so strip the - # species suffix when looking up per-column metadata. - var_infos = {key: lookup("prt", key) for key in data.dims} - return data.assign_metadata( - active_key=self.active_key, - var_infos=var_infos, - ) + return LazyList(df, metadata) diff --git a/src/lib/data/loaders/particle_h5.py b/src/lib/data/loaders/particle_h5.py index b89b478..5cfc3ef 100644 --- a/src/lib/data/loaders/particle_h5.py +++ b/src/lib/data/loaders/particle_h5.py @@ -178,13 +178,13 @@ def suffix(cls): return "h5" 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)) + steps = file_util.get_available_steps(config.data_root / self.subdir, self.prefix + ".", ".h5") + species_dict = _build_species_dict(_discover_species_qm(config.data_root / self.subdir, self.prefix, steps)) - attrss = [_load_attrs_at_step(config.data_dir, self.prefix, step) for step in steps] + attrss = [_load_attrs_at_step(config.data_root / self.subdir, self.prefix, step) for step in steps] times = np.array([attrs["time"] for attrs in attrss]) - data_paths = [_get_path_at_step(config.data_dir, self.prefix, step) for step in steps] + data_paths = [_get_path_at_step(config.data_root / self.subdir, 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) @@ -206,18 +206,14 @@ def get_data(self, config: PscPlotConfig) -> LazyList: coordss["t"] = times metadata = ListMetadata( + prepath=self.prepath, weight_key="w", coordss=coordss, species=species_dict, partition_dim="t", partition_ranges=partition_ranges, - ) - - df_with_metadata = LazyList(df, metadata) - - var_infos = {key: lookup(self.prefix, key) for key in df_with_metadata.dims} - return df_with_metadata.assign_metadata( - active_key=self.active_key, - var_infos=var_infos, + var_infos={key: lookup(self.prepath, key) for key in df.columns}, subject=Latex(r"\text{Particles}"), ) + + return LazyList(df, metadata) diff --git a/src/lib/derived_field_variables/derived_field_variable.py b/src/lib/derived_field_variables/derived_field_variable.py index 887a63e..187e7e6 100644 --- a/src/lib/derived_field_variables/derived_field_variable.py +++ b/src/lib/derived_field_variables/derived_field_variable.py @@ -30,7 +30,7 @@ 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) + return field.assign(new_data, var_infos=new_var_infos) def __repr__(self) -> str: return f"{self.__class__.__name__}(({', '.join(self.base_var_names)}) -> {self.name}: {self.derive!r})" diff --git a/src/lib/derived_particle_variables/derived_particle_variable.py b/src/lib/derived_particle_variables/derived_particle_variable.py index 060b455..27f05bc 100644 --- a/src/lib/derived_particle_variables/derived_particle_variable.py +++ b/src/lib/derived_particle_variables/derived_particle_variable.py @@ -6,7 +6,7 @@ from lib import var_info_registry from lib.data.data_with_attrs import List -__all__ = ["derived_particle_variable", "derive_particle_variable", "DERIVED_PARTICLE_VARIABLES"] +__all__ = ["derived_particle_variable", "derive_particle_variable"] class DeriveParticleVariable(typing.Protocol): @@ -29,7 +29,10 @@ def assign_to(self, data: List) -> List: info = var_info_registry.lookup("prt", self.name) new_var_infos = {**data.metadata.var_infos, self.name: info} - return data.assign_data(df.assign(**{self.name: self.derive(*(df[base_var_name] for base_var_name in self.base_var_names))})).assign_metadata(var_infos=new_var_infos) + return data.assign( + df.assign(**{self.name: self.derive(*(df[base_var_name] for base_var_name in self.base_var_names))}), + var_infos=new_var_infos, + ) def __repr__(self) -> str: return f"{self.__class__.__name__}(({', '.join(self.base_var_names)}) -> {self.name}: {self.derive!r})" @@ -42,6 +45,13 @@ def register_derived_particle_variable(prefix: str, var: DerivedParticleVariable DERIVED_PARTICLE_VARIABLES.setdefault(prefix, {})[var.name] = var +def get_derived_particle_variables(prefix: str) -> dict[str, DerivedParticleVariable]: + if prefix.startswith("prt."): + # FIXME this is a hardcoded hack + prefix = "prt" + return DERIVED_PARTICLE_VARIABLES[prefix] + + def derived_particle_variable(prefix: str): def derived_particle_variable_inner[F: (function, DeriveParticleVariable)](derive_func: F) -> F: name = derive_func.__name__ @@ -55,13 +65,16 @@ def derived_particle_variable_inner[F: (function, DeriveParticleVariable)](deriv def derive_particle_variable(data: List, active_key: str, ds_prefix: str) -> List: if active_key in data.dims: return data - elif active_key in DERIVED_PARTICLE_VARIABLES[ds_prefix]: - derived_var = DERIVED_PARTICLE_VARIABLES[ds_prefix][active_key] + + derived_vars = get_derived_particle_variables(ds_prefix) + + if active_key in derived_vars: + derived_var = derived_vars[active_key] for base_var_name in derived_var.base_var_names: data = derive_particle_variable(data, base_var_name, ds_prefix) return derived_var.assign_to(data) else: message = f"""No variable named '{active_key}'. The following variables are defined: {data.dims}. -The following variables can be derived: {list(DERIVED_PARTICLE_VARIABLES[ds_prefix])}.""" +The following variables can be derived: {list(derived_vars)}.""" raise ValueError(message) diff --git a/src/lib/file_util.py b/src/lib/file_util.py index 177f028..2b330f1 100644 --- a/src/lib/file_util.py +++ b/src/lib/file_util.py @@ -1,5 +1,8 @@ from pathlib import Path +type Prepath = str +"""Path from data root directory to the data-containing directory plus the data prefix, e.g. `"run5/pfd"`.""" + def get_available_steps(data_dir: Path, before_step: str, after_step: str) -> list[int]: files = data_dir.glob(f"{before_step}*{after_step}") @@ -10,3 +13,10 @@ def get_available_steps(data_dir: Path, before_step: str, after_step: str) -> li steps.sort() return steps + + +def split_prepath(prepath: Prepath) -> tuple[Path, str]: + components = prepath.rsplit("/", maxsplit=1) + if len(components) == 1: + return (Path("."), components[0]) + return Path(components[0]), components[1] diff --git a/src/lib/parsing/args.py b/src/lib/parsing/args.py index 9d5bc90..13a9541 100644 --- a/src/lib/parsing/args.py +++ b/src/lib/parsing/args.py @@ -6,7 +6,7 @@ class Args(argparse.Namespace): - prefix: str + prepath: str variable: str | None adaptors: list[Adaptor] hooks: list[Hook] diff --git a/src/lib/parsing/parse.py b/src/lib/parsing/parse.py index 52464b8..e8b8e71 100644 --- a/src/lib/parsing/parse.py +++ b/src/lib/parsing/parse.py @@ -8,7 +8,7 @@ def _get_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(prog="psc-plot") - parser.add_argument("prefix", help="initial active prefix") + parser.add_argument("prepath", help="initial active prepath") parser.add_argument("variable", nargs="?", default=None, help="initial active variable") parser.add_argument( "-s", diff --git a/src/lib/parsing/parse_util.py b/src/lib/parsing/parse_util.py index 6a241d3..b0a7d5f 100644 --- a/src/lib/parsing/parse_util.py +++ b/src/lib/parsing/parse_util.py @@ -7,7 +7,7 @@ def _is_identifier(val: str) -> bool: return all(re.match(r"^\w[\d\w]*$", v) for v in val.split(".")) -def fail_format(arg: str, format: str): +def fail_format(arg: str, format: str) -> typing.NoReturn: raise argparse.ArgumentTypeError(f"Expected value of form '{format}'; got '{arg}'") diff --git a/src/lib/plotting/hooks/show_com.py b/src/lib/plotting/hooks/show_com.py index 3c499f5..74fa848 100644 --- a/src/lib/plotting/hooks/show_com.py +++ b/src/lib/plotting/hooks/show_com.py @@ -6,7 +6,7 @@ def _get_center(field: Field, dim: str) -> float: other_dims = set(field.dims) - {dim} - summed = field.active_data.sum(other_dims) + summed = field.require_active_subdata().sum(other_dims) return (summed * field.coordss[dim]).sum(dim) / summed.sum(dim) diff --git a/src/lib/plotting/renderer.py b/src/lib/plotting/renderer.py index c93ab3e..4db1184 100644 --- a/src/lib/plotting/renderer.py +++ b/src/lib/plotting/renderer.py @@ -13,7 +13,7 @@ def __init__(self, full_data: Data, plot_target: PlotTarget): self.plot_target = plot_target 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) + self.full_data = full_data.with_active(key=plot_target.color_dim or plot_target.spatial_dims.y_dim) else: self.full_data = full_data diff --git a/src/lib/plotting/renderers/field_1d.py b/src/lib/plotting/renderers/field_1d.py index 9226b48..aba21d2 100644 --- a/src/lib/plotting/renderers/field_1d.py +++ b/src/lib/plotting/renderers/field_1d.py @@ -13,7 +13,7 @@ def init_plot_info(self) -> PlotInfo: plot_info = LineInfo( x_data=frame_data.coordss[x_dim], - y_data=frame_data.active_data, + y_data=frame_data.require_active_subdata(), x_dim=x_dim, y_dim=y_dim, time_dim=self.plot_target.time_dim, @@ -48,5 +48,5 @@ def init_plot_info(self) -> PlotInfo: def update_plot_info(self, frame: int): frame_data = self._get_data_at_frame(frame) - self.plot_info.set("y_data", frame_data.active_data) + self.plot_info.set("y_data", frame_data.require_active_subdata()) 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 3a01b32..cfc756e 100644 --- a/src/lib/plotting/renderers/field_2d.py +++ b/src/lib/plotting/renderers/field_2d.py @@ -1,7 +1,6 @@ import xarray as xr from lib.data.data_with_attrs import Field -from lib.data.plot_target import SpatialDimsXY from lib.plotting.plot_info import ImageInfo, PlotInfo from lib.plotting.renderer import Renderer @@ -13,10 +12,6 @@ def get_extent(da: xr.DataArray, dim: str) -> tuple[float, float]: class Field2dRenderer(Renderer[Field]): - 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 init_plot_info(self) -> PlotInfo: full_data = self.full_data frame_data = self._get_data_at_frame(0) @@ -24,7 +19,7 @@ def init_plot_info(self) -> PlotInfo: [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) + data = frame_data.require_active_subdata().transpose(y_dim, x_dim) plot_info = ImageInfo( data=data, @@ -68,7 +63,7 @@ def update_plot_info(self, frame: int): frame_data = self._get_data_at_frame(frame) [x_dim, y_dim] = self.plot_target.spatial_dims.unpack() - data = frame_data.active_data.transpose(y_dim, x_dim) + data = frame_data.require_active_subdata().transpose(y_dim, x_dim) 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 0942504..78c9e7f 100644 --- a/src/lib/plotting/renderers/polar_field.py +++ b/src/lib/plotting/renderers/polar_field.py @@ -26,7 +26,7 @@ def init_plot_info(self) -> PlotInfo: theta_vertices -= theta_vertices[1] / 2.0 plot_info = PolarMeshInfo( - data=frame_data.active_data, + data=frame_data.require_active_subdata(), r_dim=r_dim, theta_dim=theta_dim, color_dim=color_dim, @@ -65,5 +65,5 @@ def init_plot_info(self) -> PlotInfo: def update_plot_info(self, frame: int): frame_data = self._get_data_at_frame(frame) - self.plot_info.set("data", frame_data.active_data) + self.plot_info.set("data", frame_data.require_active_subdata()) self.plot_info.set("scalar_coord_values", {dim: coord for dim, coord in frame_data.coordss.items() if coord.shape == ()}) diff --git a/tests/baseline/test_crossdata_derive.png b/tests/baseline/test_crossdata_derive.png new file mode 100644 index 0000000..a05c0e1 Binary files /dev/null and b/tests/baseline/test_crossdata_derive.png differ diff --git a/tests/baseline/test_crossdir_with.png b/tests/baseline/test_crossdir_with.png new file mode 100644 index 0000000..ef93221 Binary files /dev/null and b/tests/baseline/test_crossdir_with.png differ diff --git a/tests/baseline/test_static_2d_derived_cross_prefix.png b/tests/baseline/test_static_2d_derived_cross_prefix.png deleted file mode 100644 index 9a7913c..0000000 Binary files a/tests/baseline/test_static_2d_derived_cross_prefix.png and /dev/null differ diff --git a/tests/conftest.py b/tests/conftest.py index 3c73e9f..d5f1d53 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -11,7 +11,7 @@ _TESTS_DIR = Path(__file__).parent _DATA_DIR = _TESTS_DIR / "data" -CONFIG_2D = PscPlotConfig(data_dir=_DATA_DIR / "test-2d") +CONFIG_2D = PscPlotConfig(data_root=_DATA_DIR / "test-2d") matplotlib.use("Agg") @@ -19,7 +19,7 @@ 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 = compile_plot_node(args, PscPlotConfig(data_root=_DATA_DIR / data_dir)).pull() plot._initialize() return plot.fig @@ -27,7 +27,7 @@ def make_plot(args_list: list[str], data_dir: str = "test-2d"): 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.""" args = parse_args(args_list) - node = compile_plot_node(args, PscPlotConfig(data_dir=_DATA_DIR / data_dir)) + node = compile_plot_node(args, PscPlotConfig(data_root=_DATA_DIR / data_dir)) plot = node.pull() save_dir.mkdir(exist_ok=True) path = save_dir / f"{node.get_save_file_stem()}.{format}" diff --git a/tests/test_dask_graph.py b/tests/test_dask_graph.py index 18da57f..13d7d5d 100644 --- a/tests/test_dask_graph.py +++ b/tests/test_dask_graph.py @@ -19,7 +19,7 @@ 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.""" - config = PscPlotConfig(data_dir=_DATA_DIR / data_dir) + config = PscPlotConfig(data_root=_DATA_DIR / data_dir) args = parse_args(args_list) node = compile_plot_node(args, config) data = node.input_node.pull().active_data diff --git a/tests/test_h5_species_discovery.py b/tests/test_h5_species_discovery.py index 317b362..0099b66 100644 --- a/tests/test_h5_species_discovery.py +++ b/tests/test_h5_species_discovery.py @@ -8,13 +8,12 @@ from synthetic_particles import write_step from lib.config import PscPlotConfig -from lib.data.loaders.particle_h5 import ParticleLoaderH5 +from lib.data.loader import load 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(PscPlotConfig(data_dir=tmp_path)) + data = load(PscPlotConfig(data_root=tmp_path), "prt") assert set(data.metadata.species.keys()) == {"e", "i"} e = data.metadata.species["e"] i = data.metadata.species["i"] @@ -29,8 +28,7 @@ def test_h5_species_discovery_multiple_ion_masses(tmp_path: Path): 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(PscPlotConfig(data_dir=tmp_path)) + data = load(PscPlotConfig(data_root=tmp_path), "prt") 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 @@ -43,8 +41,7 @@ def test_h5_species_discovery_multiple_ion_charges(tmp_path: Path): 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(PscPlotConfig(data_dir=tmp_path)) + data = load(PscPlotConfig(data_root=tmp_path), "prt") 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 @@ -57,8 +54,7 @@ def test_h5_species_discovery_multiple_ion_everything(tmp_path: Path): 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(PscPlotConfig(data_dir=tmp_path)) + data = load(PscPlotConfig(data_root=tmp_path), "prt") 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,9 +71,8 @@ def test_h5_species_discovery_electron_merge_warns(tmp_path: Path): 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(PscPlotConfig(data_dir=tmp_path)) + data = load(PscPlotConfig(data_root=tmp_path), "prt") assert set(data.metadata.species.keys()) == {"e"} @@ -85,6 +80,5 @@ 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(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(PscPlotConfig(data_dir=tmp_path)) + data = load(PscPlotConfig(data_root=tmp_path), "prt") assert set(data.metadata.species.keys()) == {"e", "i"} diff --git a/tests/test_memory.py b/tests/test_memory.py index 9d28333..ad4e417 100644 --- a/tests/test_memory.py +++ b/tests/test_memory.py @@ -27,7 +27,7 @@ def _run_pipeline(data_dir: pathlib.Path, chunksize: int, result_queue: mp.Queue from lib.parsing.parse import parse_args 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 = compile_plot_node(args, PscPlotConfig(data_root=data_dir, dask_chunk_size=chunksize)).pull() plot._initialize() peak = resource.getrusage(resource.RUSAGE_SELF).ru_maxrss diff --git a/tests/test_particle_bp_perf.py b/tests/test_particle_bp_perf.py index aaba829..6a47f73 100644 --- a/tests/test_particle_bp_perf.py +++ b/tests/test_particle_bp_perf.py @@ -27,7 +27,7 @@ def _run_h5_pipeline(data_dir: pathlib.Path, result_queue: mp.Queue) -> None: 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() + plot = compile_plot_node(args, PscPlotConfig(data_root=data_dir)).pull() t0 = time.perf_counter() plot._initialize() elapsed = time.perf_counter() - t0 @@ -42,7 +42,7 @@ def _run_bp_pipeline(data_dir: pathlib.Path, result_queue: mp.Queue) -> None: 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() + plot = compile_plot_node(args, PscPlotConfig(data_root=data_dir)).pull() t0 = time.perf_counter() plot._initialize() elapsed = time.perf_counter() - t0 diff --git a/tests/test_particle_bp_vs_h5.py b/tests/test_particle_bp_vs_h5.py index ccfce16..888fde5 100644 --- a/tests/test_particle_bp_vs_h5.py +++ b/tests/test_particle_bp_vs_h5.py @@ -11,13 +11,13 @@ def _load_and_filter_h5(species_key: str): - loader = ParticleLoaderH5(prefix="prt", active_key=None) + loader = ParticleLoaderH5(prepath="prt") 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) + loader = ParticleLoaderBp(prepath=f"prt.{species_key}") return loader.get_data(CONFIG_2D) diff --git a/tests/test_plots.py b/tests/test_plots.py index bac4dd3..05849e5 100644 --- a/tests/test_plots.py +++ b/tests/test_plots.py @@ -37,13 +37,6 @@ def test_animated_2d_derived(): return make_plot("pfd h2_cc -v y z".split()) -@pytest.mark.mpl_image_compare(**MPL_KWARGS) -def test_static_2d_derived_cross_prefix(): - """2D view of a variable derived across prefixes: the active `pfd` field `jy_ec` - times `pfd_moments::jy_e`, which auto-loads the `pfd_moments` prefix.""" - return make_plot("pfd --derive electron_power=jy_ec*pfd_moments::jy_e -i t=-1 -v y z time=".split()) - - @pytest.mark.mpl_image_compare(**MPL_KWARGS) def test_animated_2d_idx(): """2D slice of x-component of magnetic field at the x=1 index.""" @@ -101,6 +94,24 @@ def test_image_and_line(): return make_plot("prt.i --bin y py=100 --nan0 --scale log --compute -v y py -w pfd::ey_ec -v y".split()) +# --- Cross-dataset plots --- + + +@pytest.mark.mpl_image_compare(**MPL_KWARGS) +def test_crossdata_derive(): + """Power of electric field acting on ions.""" + return make_plot("pfd --derive ipower=ey_ec*pfd_moments::jy_i".split()) + + +# --- Cross-subdir plots --- + + +@pytest.mark.mpl_image_compare(**MPL_KWARGS) +def test_crossdir_with(): + """2D vs. 3D B_x(y). The "y"s aren't the same length, and that's ok.""" + return make_plot("test-2d/pfd hx_fc -v y --display \\text{2d} --with test-3d/pfd::hx_fc -v y --display \\text{3d}".split(), data_dir=".") + + # --- Turbulence power spectrum --- @@ -182,7 +193,7 @@ def test_static_scatter_bp(): @pytest.mark.mpl_image_compare(**MPL_KWARGS) def test_hamscan(): """Archetypal scan for hammerhead distributions ("hams").""" - return make_plot("prt.e -i t=-1 --derive pzx --bin y py=20 pzx=20 t= -v py pzx time=y --compute".split(), data_dir="test-2d") + return make_plot("prt.e -i t=-1 --with pzx --bin y py=20 pzx=20 t= -v py pzx time=y --compute".split(), data_dir="test-2d") # --- Particle moments --- diff --git a/tests/test_save_filename.py b/tests/test_save_filename.py index 9b99aec..e19e560 100644 --- a/tests/test_save_filename.py +++ b/tests/test_save_filename.py @@ -8,10 +8,10 @@ @pytest.mark.parametrize( "args_list, expected_stem", [ - (["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-scale_log-v_y,z"), - (["pfd", "hx_fc", "-v", "y", "z", "time="], "pfd-hx_fc-v_y,z;time="), + (["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-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):