Source code for escape.storage.storage

from copy import deepcopy
import hashlib
import io
import threading
import time
from traceback import format_exc
import numpy as np
import dask
from dask import array as da
from dask import dataframe as ddf
from dask.diagnostics import ProgressBar
from dask.distributed import get_client, progress, wait
from dask.typing import DaskCollection
import operator

from escape.storage.source import Source
from ..utilities import get_corr, hist_asciicontrast, Hist_ascii, is_local_client_distributed, plot2D, roundto
import logging
from itertools import chain
from numbers import Number
import re
from .. import utilities
import h5py
import hickle
from matplotlib import pyplot as plt
import pandas as pd
from pathlib import Path
import html
import base64
from io import BytesIO
from .storage_tools import ArrayTools, ScanTools


logger = logging.getLogger(__name__)

# Serialise matplotlib pyplot calls that happen inside background repr threads.
# pyplot's global state (current figure, interactive mode) is not thread-safe;
# this lock ensures only one repr computation touches it at a time.
_mpl_repr_lock = threading.Lock()

import escape


class ArraySelector:
    def __init__(self, arrayitem, dims=None):
        """Container object for selecting array subsets in functions mapped on escape Arrays."""
        self.arrayitem = arrayitem
        self.dims = dims

    def __call__(self, sel):
        if max(self.dims) <= (len(sel) - 1):
            return self.arrayitem.__getitem__(tuple(sel[n] for n in self.dims))
        else:
            return self.arrayitem


def _apply_method(
    foo_np,
    foo_da,
    data,
    is_dask_array,
    *args,
    convertesc_axis_kw=False,
    convertOutput2EscData="auto",
    **kwargs,
):
    if convertesc_axis_kw:
        axis = kwargs.get("axis", "noaxis")
        if isinstance(axis, Number):
            axis = [axis]
        if (axis == "noaxis") or (0 in axis):
            convertOutput2EscData = []

    if is_dask_array:
        if not foo_da:
            raise NotImplementedError(
                f"Function {foo_np.__name__} is not defined for dask based arrays!"
            )
        return escaped(foo_da, convertOutput2EscData=convertOutput2EscData)(
            data, *args, **kwargs
        )
    else:
        if not foo_np:
            raise NotImplementedError(
                f"Function {foo_da.__name__} is not defined for numpy based arrays!"
            )
        return escaped(foo_np, convertOutput2EscData=convertOutput2EscData)(
            data, *args, **kwargs
        )


# ---------------------------------------------------------------------------
# Programmatic numpy/dask method injection for Array
# ---------------------------------------------------------------------------
# Table: method_name -> (np_func, da_func_or_None, axis_kw, esc_out)
#   axis_kw=True  → convertesc_axis_kw=True  (reduction: wrapping depends on axis)
#   esc_out=[0]   → convertOutput2EscData=[0] (element-wise: always wrap output)
_ARRAY_DELEGATE_METHODS = {
    # Reductions — wrapping depends on whether event axis is collapsed
    "nansum":        (np.nansum,        da.nansum,        True,  None),
    "nanmean":       (np.nanmean,       da.nanmean,       True,  None),
    "nanstd":        (np.nanstd,        da.nanstd,        True,  None),
    "nanmedian":     (np.nanmedian,     None,             True,  None),
    "nanmin":        (np.nanmin,        da.nanmin,        True,  None),
    "nanmax":        (np.nanmax,        da.nanmax,        True,  None),
    "nanpercentile": (np.nanpercentile, None,             True,  None),
    "nanquantile":   (np.nanquantile,   None,             True,  None),
    "sum":           (np.sum,           da.sum,           True,  None),
    "mean":          (np.mean,          da.mean,          True,  None),
    "average":       (np.average,       da.average,       True,  None),
    "std":           (np.std,           da.std,           True,  None),
    "median":        (np.median,        None,             True,  None),
    "percentile":    (np.percentile,    None,             True,  None),
    "quantile":      (np.quantile,      None,             True,  None),
    "min":           (np.min,           da.min,           True,  None),
    "max":           (np.max,           da.max,           True,  None),
    "all":           (np.all,           da.all,           True,  None),
    "any":           (np.any,           da.any,           True,  None),
    "abs":           (np.abs,           da.abs,           True,  None),
    # Element-wise — always returns an Array with the same shape
    "isnan":         (np.isnan,         da.isnan,         False, [0]),
    "isinf":         (np.isinf,         da.isinf,         False, [0]),
    "isfinite":      (np.isfinite,      da.isfinite,      False, [0]),
}


def _make_array_method(name, np_func, da_func, axis_kw, esc_out):
    """Factory: build an Array method that delegates to a numpy/dask function."""
    np_summary = next(
        (line.strip() for line in (np_func.__doc__ or "").split("\n") if line.strip()), ""
    )
    if esc_out:
        doc = (
            f"Apply :func:`numpy.{name}` element-wise to this Array's data.\n\n"
            f"{np_summary}\n\n"
            "Returns an :class:`Array` with the same index and scan structure."
        )
        kw = {"convertOutput2EscData": esc_out}
    else:
        doc = (
            f"Apply :func:`numpy.{name}` to this Array's data.\n\n"
            f"{np_summary}\n\n"
            "Omitting ``axis`` or passing ``axis=0`` reduces over events and "
            "returns a plain numpy/dask result. Pass ``axis=N`` (N > 0) to "
            "reduce along a non-event axis and receive a new :class:`Array`."
        )
        kw = {"convertesc_axis_kw": True}

    def method(self, *args, **kwargs):
        return _apply_method(np_func, da_func, self, self.is_dask_array(), *args, **kw, **kwargs)

    method.__name__ = name
    method.__qualname__ = f"Array.{name}"
    method.__doc__ = doc
    return method


[docs] class Array: """nd array data wrapper with optional scan metadata and grid support. ``Array`` stores raw measurement data together with an event index and optional scan grouping information. When ``step_lengths`` and ``parameter`` are provided, a lazily constructed ``scan`` property exposes grouped step selection and scan-level operations. Args: data: nd data array or callable returning data. index: Event identifiers aligned with the first dimension of ``data``. step_lengths: List of step sizes for each scan step. parameter: Scan parameter metadata for each step. name: Optional array name. source: Optional source metadata object. grid_specs: Optional metadata used to build ``scan.grid``. """
[docs] def __init__( self, data=None, index=None, step_lengths=None, parameter=None, name=None, source=None, grid_specs=None, ): self.index_dim = 0 if not (callable(data) or callable(index)): assert data.shape[self.index_dim] == len( index ), "lengths of data and event IDs must match!" if not step_lengths is None: if not callable(index): assert sum(step_lengths) == len( index ), "StepsLength need to add up to dataset length!" if parameter is None: logger.debug( "No information about event groups (steps) \ available!" ) else: step_lengths = [len(index)] self._index = index self._data = data self._data_selector = None self._scan = None self._scan_parameter = parameter self._scan_step_lengths = step_lengths self.name = name self.source = source self._touched = False self._tools = None self._grid_specs = grid_specs
@property def scan(self): if self._scan is None: self._scan = Scan(self._scan_parameter, self._scan_step_lengths, self, grid_specs=self._grid_specs if hasattr(self, "_grid_specs") else None) return self._scan @property def grid(self): if hasattr(self.scan,"grid"): return self.scan.grid else: return None @property def tools(self): if self._tools is None: self._tools = ArrayTools(self) return self._tools def _touch(self): if not self._touched: dum = self.index @property def index(self): if isinstance(self._index, da.Array): self._index = self._index.compute() self._index, self._data, self.scan.step_lengths = get_unique_indexes( self._index, self.data, self.scan.step_lengths ) elif callable(self._index): self._index = self._index() self._index, self._data, self.scan.step_lengths = get_unique_indexes( self._index, self.data, self.scan.step_lengths ) return self._index @property def data(self): self._touch() # TODO: try getting the properties outside of storage in the # specific parser section if callable(self._data): # TODO: cludgy solution need fix at some point. op = self._data(data_selector=self._data_selector) if len(op) == 2 and (type(op[1]) is str) and op[1] == "nopersist": return op[0] else: self._data = op return self._data else: return self._data
[docs] def is_dask_array(self): return isinstance(self.data, da.Array)
# numpy/dask reduction and element-wise methods are injected below the class # definition via _ARRAY_DELEGATE_METHODS + _make_array_method.
[docs] def nancount(self): """Return the number of finite (non-NaN) events in this Array.""" return int(np.sum(~np.isnan(self.data.compute() if self.is_dask_array() else self.data)))
[docs] def filter(self, *args, **kwargs): return filter(self, *args, **kwargs)
[docs] def digitize(self, bins, **kwargs): return digitize(self, bins, **kwargs)
[docs] def get_modulo_array(self, mod, offset=0): index = self.index out_bool = np.mod(index, mod) == offset return self[out_bool]
[docs] def update(self, array): """Merge *array* into this Array, adding only events with new pulse IDs. Events already present in ``self`` (matched by pulse ID) are ignored; new events are appended as additional scan steps so that the existing scan structure is preserved and the new events keep their own step grouping. Intended for incremental accumulation during acquisition — call repeatedly with newer snapshots to build up a complete dataset. Parameters ---------- array : escape.Array Source Array whose new events will be added to this one. Returns ------- escape.Array New Array containing all events from ``self`` plus any events in *array* whose pulse ID was absent from ``self``. Returns ``self`` unchanged (same object) if *array* contributes no new events. """ new_mask = ~np.isin(array.index, self.index) if not new_mask.any(): return self new_positions = new_mask.nonzero()[0] _, new_scan = get_scan_step_selections( new_positions, array.scan.step_lengths, scan=array.scan ) new_part = Array( data=array.data[new_mask], index=array.index[new_mask], step_lengths=new_scan.step_lengths, parameter=new_scan.parameter, ) return concatenate([self, new_part])
[docs] def correlation_analysis_to(self, ref, order=2): td, tr = match_arrays(self, ref) std_rel, std_fx_rel = get_corr(td.data, tr.data, order=order) return std_rel, std_fx_rel
[docs] def __len__(self): self._touch() return len(self.index)
[docs] def categorize(self, other_array): """Re-sort and re-group *other_array* to match this Array's index ordering and scan-step boundaries. The returned Array contains *other_array*'s data values at the pulse IDs that are common to both arrays, ordered and grouped exactly as *self*. This is the primary tool for applying a new grouping (obtained e.g. via :meth:`digitize` or :meth:`get_index_array`) to another channel. Parameters ---------- other_array : escape.Array The array to re-sort. Returns ------- escape.Array *other_array* restricted to the common pulse IDs and re-grouped according to *self*'s scan structure. Notes ----- Equivalent to ``escape.match_arrays(self, other_array)[1]``. Examples -------- >>> time_bins = sig.get_index_array(N_index_aggregation=1000) >>> i0_rebinned = time_bins.categorize(i0) """ return match_arrays(self, other_array)[1]
[docs] def __getitem__(self, *args, **kwargs): # this is multi dimensional itemgetting if type(args[0]) is tuple: # expanding ellipses --> TODO: multiple ellipses possible? if Ellipsis in [type(ta) for ta in args[0]]: rargs = list(args[0]) elind = rargs.index(Ellipsis) rargs.pop(elind) eventsel = [ta for ta in rargs if ta] missing_dims = self.ndim - len(eventsel) for n in range(missing_dims): rargs.insert(elind, slice(None, None, None)) args = (tuple(rargs),) # get event selector for the event ID selection eventIx = -1 for n, targ in enumerate(args[0]): if targ: eventIx += 1 if eventIx == self.index_dim: break if type(args[0][eventIx]) is int: rargs = list(args[0]) if rargs[eventIx] == -1: rargs[eventIx] = slice(rargs[eventIx], None) else: rargs[eventIx] = slice(rargs[eventIx], rargs[eventIx] + 1) args = (tuple(rargs),) events = args[0][eventIx] events_type = "eventIx" # Single dimension itemgetting, which is by default along # event dimension, raise error if inconsistent with data shape else: assert self.index_dim == 0, "requesting slice not along event dimension!" # making sure slices are taken in a way the event dimention is not squeezed away. if type(args[0]) is int: rargs = list(args) if rargs[0] == -1: rargs[0] = slice(rargs[0], None) else: rargs[0] = slice(rargs[0], rargs[0] + 1) args = tuple(rargs) events = args[0] events_type = "none" # expand all dimensions for potential use in derived functions if isinstance(events, slice): events = list(range(*events.indices(len(self)))) elif isinstance(events, np.ndarray) and events.dtype == bool: events = events.nonzero()[0] elif isinstance(events, Array): inds_self, [inds_selector], dum = match_indexes(self.index, [events.index]) events = inds_self[events.data[inds_selector].nonzero()[0]] if events_type == "eventIx": args[0][eventIx] = events elif events_type == "none": args = (events,) else: raise Exception( "Issue in escape array getitem using another escape array!" ) stepLengths, scan = get_scan_step_selections( events, self.scan.step_lengths, scan=self.scan ) # Save indices for potential use in derived functions self._data_selector = args return Array( data=self.data.__getitem__(*args), index=self.index.__getitem__(events), step_lengths=stepLengths, parameter=scan.parameter, )
[docs] def get_random_events(self, n, seed=None): np.random.seed(seed) inds = np.random.randint(0, len(self), size=(n,)) return self[list(inds)]
@property def shape(self, *args, **kwargs): self._touch() return self.data.shape @property def ndim(self, *args, **kwargs): return self.data.ndim @property def ndim_nonzero(self, *args, **kwargs): return len(np.asarray(self.shape)[np.nonzero(self.shape)[0]])
[docs] def transpose(self, *args): if not args: axes = tuple(range(self.ndim - 1, -1, -1)) elif len(args) == 1: axes = args[0] else: axes = args return Array( data=self.data.transpose(*args), index=self.index, step_lengths=self.scan.step_lengths, parameter=self.scan.parameter, )
[docs] def ravel_event_data(self): """Flatten all non-event axes into a single dimension per event. Converts an Array of shape ``(N, d1, d2, ...)`` to ``(N, d1*d2*...)``, preserving the event axis and scan structure. Useful for feeding multi-dimensional detector data into functions that expect a 1-D value per event. Returns ------- escape.Array Array with shape ``(N, d1*d2*...)``. Examples -------- >>> imgs.shape # (500, 64, 64) >>> flat = imgs.ravel_event_data() >>> flat.shape # (500, 4096) """ n_events = self.shape[self.index_dim] new_data = self.data.reshape(n_events, -1) return Array( data=new_data, index=self.index, step_lengths=self.scan.step_lengths, parameter=self.scan.parameter, )
@property def T(self): return self.transpose()
[docs] def compute(self, **kwargs): """Evaluate the dask graph and return a new Array backed by a NumPy array. No-op when the data is already a NumPy array (returns *self* with a message). All index and scan metadata are preserved. Parameters ---------- **kwargs Forwarded to :func:`dask.array.Array.compute`. Returns ------- escape.Array Same Array with NumPy data instead of a dask graph. See Also -------- escape.compute : Compute several Arrays in one dask scheduler pass. """ if self.is_dask_array(): with ProgressBar(): return Array( data=self.data.compute(**kwargs), index=self.index, step_lengths=self.scan.step_lengths, parameter=self.scan.parameter, ) else: if self.name: print(f"No `compute` necessary for {self.name}") else: print(f"No `compute` necessary") return self
[docs] def persist(self): self.data.persist()
# def get_progress()
[docs] def map_index_blocks( self, foo, *args, # chunks=None, drop_axis=None, new_axis=None, new_element_size=None, event_dim="same", **kwargs, ): """Apply *foo* block-wise over the event axis using dask's ``map_blocks``. The function ``foo`` receives a **raw NumPy array** (one dask chunk along the event axis) and returns a NumPy array. The result is assembled back into a lazy dask-backed :class:`Array` with the same index and scan metadata. This is the preferred way to apply arbitrary NumPy or SciPy functions (gain correction, thresholding, peak fitting, …) to large detector data without loading everything into memory. Parameters ---------- foo : callable ``f(block, *args, **kwargs) -> ndarray``. *block* has shape ``(n_events_in_chunk, *element_shape)``. *args Extra positional arguments forwarded to ``foo``. drop_axis : int or list of int, optional Axes to remove from the output (forwarded to ``dask.map_blocks``). new_axis : int or list of int, optional New axes to add to the output. new_element_size : list of int, optional Shape of each per-event element in the output (excluding the event axis). Required when ``foo`` changes the per-event shape. **kwargs Extra keyword arguments forwarded to ``foo``. Returns ------- escape.Array Lazy Array with the transformed data. Examples -------- Threshold pixels below 4 keV to NaN:: def threshold(block, thr): out = block.copy() out[out < thr] = np.nan return out imgs_clean = imgs.map_index_blocks(threshold, 4.0) Extract two scalars per event (change per-event shape):: posamp = tt_proj.map_index_blocks( lambda block: np.array([find_edge(row) for row in block]), new_element_size=(2,), dtype=float, ) Notes ----- Formerly called ``map_event_blocks`` in older versions of ``escape``. """ # Test: creating Source instance for origin tracking src = Source( "factory", factory=foo, args=args, kwargs=kwargs, iargout=0, name_dataset=self.name, ) # Getting chunks in the event dimension if event_dim == "same": event_dim = self.index_dim chunks_edim = self.data.chunks[self.index_dim] # making sure that chunks in other dimensions are "flat" shp = self.data.shape if new_element_size: new_size = list(new_element_size) new_size.insert(self.index_dim, None) newchunks = [] rechunk = False for dim, dimchunks in enumerate(self.data.chunks): if dim == self.index_dim: newchunks.append(dimchunks) elif drop_axis and dim in drop_axis: continue else: rechunk = len(dimchunks) > 1 or rechunk if new_element_size: newchunks.append((new_size[dim],)) else: newchunks.append((sum(dimchunks),)) if rechunk: data = self.data.rechunk(tuple(newchunks)) else: data = self.data if new_element_size: chunks = newchunks else: chunks = None # checking if any inputs are to be selected if any([isinstance(x, ArraySelector) for x in chain(args, kwargs.values())]): print("Is arg selector") def get_data(data_selector=None): args_sel = [ ta if not isinstance(ta, ArraySelector) else ta(data_selector) for ta in args ] kwargs_sel = {} for tk, tv in kwargs.items(): if isinstance(tv, ArraySelector): kwargs_sel[tk] = tv(data_selector) else: kwargs_sel[tk] = tv return ( data.map_blocks( foo, *args_sel, chunks=chunks, drop_axis=drop_axis, new_axis=new_axis, **kwargs_sel, ), "nopersist", ) return Array( data=get_data, index=self.index, step_lengths=self.step_lengths, parameter=self.scan.parameter, ) else: return Array( data=data.map_blocks( foo, *args, chunks=chunks, drop_axis=drop_axis, new_axis=new_axis, **kwargs, ), index=self.index, step_lengths=self.scan.step_lengths, parameter=self.scan.parameter, )
[docs] def store(self, parent_h5py=None, name=None, unit=None, lock="auto", **kwargs): """a way to store data, especially expensively computed data, into a new file.""" if lock == "auto": lock = get_lock() if not hasattr(self, "h5"): self.h5 = ArrayH5Dataset(parent_h5py, name) with ProgressBar(): self.h5.append(self.data, self.index, self.scan, lock=lock, **kwargs) self._data = self.h5.get_data_da() self._index = self.h5.index self.scan._save_to_h5(self.h5.grp)
[docs] def set_h5_storage(self, parent_h5py, name=None): if not hasattr(self, "h5"): if not name: name = self.name self.h5 = ArrayH5Dataset(parent_h5py, name) else: try: logger.info( f"h5 storage already set at {name} in {self.h5.file.filename}" ) except: logger.info(f"h5 storage already set for {name}")
[docs] def store_file(self, parent_h5py=None, name=None, unit=None, **kwargs): """a way to store data, especially expensively computed data, into a new file.""" if not hasattr(self, "h5"): self.h5 = ArrayH5Dataset(parent_h5py, name) with ProgressBar(): self.h5.append(self.data, self.index, self.scan) self._data = self.h5.get_data_da() self._index = self.h5.index
[docs] def set_h5_storage_file(self, file_name, parent_group_name, name=None): if not hasattr(self, "h5"): if not name: name = self.name self.h5 = ArrayH5File(file_name, parent_group_name, name) else: logger.info(f"h5 storage already set at {name} in {self.h5.file_name}")
[docs] @classmethod def load_from_h5(cls, parent_h5py, name): h5 = ArrayH5Dataset(parent_h5py, name) try: parameter, step_lengths, grid_specs = Scan._load_from_h5(parent_h5py[name]) except: # print(f"could not read scan metadata of {name}") parameter = None step_lengths = None grid_specs = None data = h5.get_data_da() if data is None: return None else: return cls( index=h5.index, data=data, parameter=parameter, step_lengths=step_lengths, grid_specs=grid_specs, name=name, )
[docs] def ones(self, **kwargs): return Array( data=np.ones(len(self), **kwargs), index=self.index, step_lengths=self.scan.step_lengths, parameter=self.scan.parameter, )
def _get_ana_str(self, perc_limits=[5, 95]): sqaxes = list(range(self.data.ndim)) sqaxes.pop(self.index_dim) try: d = self.data.squeeze(axis=tuple(sqaxes)) except: return "" if d.ndim == 1: ostr = "" if d.dtype == bool: hrange = [0, 1] else: hrange = np.percentile(d[~np.isnan(d)], perc_limits) formnum = lambda num: "{:<9}".format("%0.4g" % (num)) for n, td in enumerate(self.scan): ostr += ( "Step %04d:" % n + hist_asciicontrast( td.data.squeeze(), bins=40, range=hrange, disprange=False ) + "\n" ) ho = Hist_ascii(d, range=hrange, bins=40) ostr += ho.horizontal() return ostr else: return ""
[docs] def get_index_array(self, N_index_aggregation=None): """Return an Array whose data equals its own index (pulse IDs), optionally grouped into contiguous bins. Without aggregation this is a simple 1-D Array where ``data == index``, useful as an "identity" sorter. With *N_index_aggregation* the pulse IDs are binned into groups of width *N_index_aggregation* index units (typically pulse IDs), creating a coarser time-ordered grouping. Parameters ---------- N_index_aggregation : int, optional Width of each pulse-ID bin. If *None* no binning is applied. Returns ------- escape.Array 1-D Array with ``data == index`` (before any binning). Notes ----- The resulting Array can be used with :meth:`categorize` to apply the new grouping to any other channel: Examples -------- >>> # Group into bins of 1000 consecutive pulse IDs >>> time_bins = sig.get_index_array(N_index_aggregation=1000) >>> sig_rebinned = time_bins.categorize(sig) >>> i0_rebinned = time_bins.categorize(i0) """ if N_index_aggregation: tmp = Array(data=self.index, index=self.index) return tmp.digitize( np.arange(min(tmp.data), max(tmp.data), N_index_aggregation) ) else: return Array(data=self.index, index=self.index)
[docs] def correct_for_references( self, isref_bool, N_index_aggregation=None, operation=operator.truediv ): # TODO: incomplete — `indsrt` and `ref` are undefined; implementation needed. refs = self[isref_bool] noref = self[~isref_bool] indxs = self.get_index_array(N_index_aggregation=N_index_aggregation) indxs = indxs[[slice(None), *([None] * (self.ndim - 1))]] return concatenate( operation(tanr, tar) for tanr, tar in zip(indsrt * noref, (indxs * ref).scan.mean(axis=0)) )
[docs] def plot( self, axis=None, linespec=".", *args, **kwargs, ): y = self.data x = self.index if not axis: axis = plt.gca() axis.plot(x, y, linespec, *args, **kwargs) if self.name: axis.set_ylabel(self.name) axis.set_xlabel("index")
[docs] def plot_corr( self, arr, ratio=False, axis=None, linespec=".", polyfit_order=None, *args, **kwargs, ): yarr, xarr = match_arrays(self, arr) y = yarr.data x = xarr.data if not axis: axis = plt.gca() if ratio: axis.plot(x, y / x, linespec, *args, **kwargs) else: axis.plot(x, y, linespec, *args, **kwargs) if arr.name: axis.set_xlabel(arr.name) if self.name and arr.name: axis.set_ylabel(f"{self.name} / {arr.name}") if not polyfit_order is None: pres = np.polyfit(x, y, polyfit_order) xp = np.linspace(np.min(x), np.max(x), 1000) yp = np.polyval(pres, xp) if ratio: plt.plot(xp, yp / xp, "r") else: plt.plot(xp, yp, "r") return pres
[docs] def hist( self, cut_percentage=0, bins="auto", normalize_to=None, scanpar_name=None, plot_results=True, plot_axis=None, ): if self.is_dask_array(): raise Exception( "escape array needs to be numpy type for histogramming, compute first." ) flat = self.data.ravel().astype(float) [hmin, hmax] = np.nanpercentile(flat, [cut_percentage, 100 - cut_percentage]) if not (np.isfinite(hmin) and np.isfinite(hmax) and hmin < hmax): hmin, hmax = float(np.nanmin(flat)), float(np.nanmax(flat)) if hmin == hmax: hmin -= 0.5 hmax += 0.5 hbins = np.histogram_bin_edges(self.data.ravel(), bins, range=[hmin, hmax]) hdat, bin_edges = np.histogram(self.data.ravel(), bins=hbins) if normalize_to == "max": hdat = hdat / hdat.max() elif normalize_to == "sum": hdat = hdat / hdat.sum() if plot_results: if not plot_axis: plot_axis = plt.gca() plt.step(hbins[:-1], hdat, where="post") plt.xlabel(self.name) return hdat, hbins
def __repr__(self, bare=False): s = "<%s.%s object at %s>" % ( self.__class__.__module__, self.__class__.__name__, hex(id(self)), ) s += " {}; shape {}".format(self.name, self.shape) s += "\n" if not bare: if isinstance(self.data, np.ndarray): s += self._get_ana_str() if self.scan: s += self.scan.__repr__() return s @property def dtype(self): return self.data.dtype
[docs] def astype(self, newtype): return Array( data=self.data.astype(newtype), index=self.index, step_lengths=self.scan.step_lengths, parameter=self.scan.parameter, )
def _get_repr_hist_plot(self, fmt="png", figsize=[5, 3]): plt.ioff() f = plt.figure(figsize=figsize) ax = f.add_subplot(111) if self.dtype == bool: flims = [0, 1] else: # Compute a robust filter range from per-step 5–95 percentiles. # Falls back to None (no filtering) when the range is degenerate # (constant data, all-NaN steps, or only a few discrete values). flims = None try: per_step = self.scan.nanpercentile([5, 95]) lows = [float(np.nanmin(s)) for s in per_step] highs = [float(np.nanmax(s)) for s in per_step] lo, hi = float(np.nanmin(lows)), float(np.nanmax(highs)) if np.isfinite(lo) and np.isfinite(hi) and lo < hi: flims = [lo, hi] except Exception: pass if len(self.scan) > 1: arr_to_hist = self.filter(*flims) if flims is not None else self arr_to_hist.scan.hist( plot_axis=ax, cmap=plt.cm.Reds, cut_percentage=0, ) if self.dtype == bool: self.scan.plot(axis=ax, fmt="k.-", use_quantiles=False, label="mean") else: self.scan.plot(axis=ax, fmt="k.-", use_quantiles=False, label="median") plt.ylabel(self.name) ax.legend(fancybox=True, framealpha=0.3, loc="best") else: to_hist = self.data if self.dtype == bool: to_hist = to_hist.astype(int) plt.hist(to_hist, "auto") plt.xlabel(self.name) ax.grid("on") # ax_steps.set_xlim(0, len(self.scan) - 1) # ax_steps.set_xlabel("Step number") f.tight_layout() if fmt == "svg": s = io.StringIO() f.savefig(s, format="svg", bbox_inches="tight") imgobj = s.getvalue() elif fmt == "png": tmpfile = BytesIO() f.savefig(tmpfile, format="png", bbox_inches="tight") imgobj = base64.b64encode(tmpfile.getvalue()).decode("utf-8") plt.ion() return imgobj def _get_repr_map_plot(self, fmt="png", figsize=[5, 3]): plt.ioff() f = plt.figure(figsize=figsize) ax = f.add_subplot(111) # ax_steps = ax.twiny() if len(self.scan) > 1: self.scan.plot(axis=ax, cmap=plt.cm.Greens) else: print("non scan repr not implemented yet for maps/waveforms or images!") ax.grid("on") # ax_steps.set_xlim(0, len(self.scan) - 1) # ax_steps.set_xlabel("Step number") f.tight_layout() if fmt == "svg": s = io.StringIO() f.savefig(s, format="svg", bbox_inches="tight") imgobj = s.getvalue() elif fmt == "png": tmpfile = BytesIO() f.savefig(tmpfile, format="png", bbox_inches="tight") imgobj = base64.b64encode(tmpfile.getvalue()).decode("utf-8") plt.ion() return imgobj def _get_repr_grid_plot(self, fmt="png", figsize=(5, 3.5)): """Repr plot for Arrays with an associated multi-dimensional Grid. * 1-D array, 2-D grid → 2-D heatmap of per-step means, labelled with the grid positions and dimension names. * All other combinations → the standard scan plot with a title that notes the full grid shape, so high dimensionality is visible. """ grid = self.grid is_scalar_array = (self.ndim == 1) or all(ts <= 1 for ts in self.shape[1:]) grid_ndim = len(grid.shape) dim_str = "×".join(str(s) for s in grid.shape) # e.g. "10×8" plt.ioff() f = plt.figure(figsize=figsize) ax = f.add_subplot(111) if is_scalar_array and grid_ndim == 2: step_means = np.asarray(self.scan.nanmean()) grid_data = grid.to_grid(step_means) # shape (nrows, ncols) positions = grid.positions if positions and len(positions) >= 2: x_raw = np.asarray(positions[1]) y_raw = np.asarray(positions[0]) else: x_raw = np.arange(grid.shape[1]) y_raw = np.arange(grid.shape[0]) p = plot2D(x_raw, y_raw, grid_data, axis=ax) plt.colorbar(p, ax=ax, label=self.name or "mean") if grid.dimension_names and len(grid.dimension_names) >= 2: ax.set_xlabel(grid.dimension_names[1]) ax.set_ylabel(grid.dimension_names[0]) ax.set_title(f"Grid mean [{dim_str}]") else: # higher-D grid or per-event arrays: fall back to scan plot try: self.scan.plot(axis=ax) except Exception: pass title = f"Scan plot (grid: {dim_str})" if not is_scalar_array: title += f" | array shape: {'×'.join(str(s) for s in self.shape)}" ax.set_title(title, fontsize=9) ax.grid(True) f.tight_layout() if fmt == "svg": s = io.StringIO() f.savefig(s, format="svg", bbox_inches="tight") imgobj = s.getvalue() else: tmpfile = BytesIO() f.savefig(tmpfile, format="png", bbox_inches="tight") imgobj = base64.b64encode(tmpfile.getvalue()).decode("utf-8") plt.ion() return imgobj # ------------------------------------------------------------------ # Repr helpers: plot dispatch, caching, async display # ------------------------------------------------------------------ def _get_repr_plot_png_b64(self): """Return a base-64 PNG string for the repr plot, or raise on failure. This is the single entry-point for repr plot creation. Change this method (or the private helpers it calls) to alter how the plot looks. """ if self.grid is not None: return self._get_repr_grid_plot(fmt="png") if (self.ndim == 1) or all(ts <= 1 for ts in self.shape[1:]): return self._get_repr_hist_plot(fmt="png") elif self.ndim_nonzero == 2: return self._get_repr_map_plot(fmt="png") else: raise NotImplementedError("No repr plot for this array shape.") def _repr_cache_key(self): """Cheap content hash used as a cache key for the repr PNG.""" try: h = hashlib.sha256() h.update(str(self.shape).encode()) h.update(str(getattr(self, "dtype", "")).encode()) h.update((self.name or "").encode()) data = self.data if hasattr(data, "ravel"): flat = data.ravel() n = min(500, len(flat)) sample = flat[:n] if hasattr(sample, "compute"): sample = sample.compute() h.update(np.asarray(sample).tobytes()) return h.hexdigest()[:24] except Exception: return None def _repr_cache_path(self, key): cache_dir = Path.home() / ".cache" / "escape" / "repr" cache_dir.mkdir(parents=True, exist_ok=True) return cache_dir / f"{key}.b64" def _ipython_display_(self, **kwargs): """Async Jupyter display: show 'Computing…' immediately, then update. * Checks a disk cache keyed by a content hash of the array data before starting any computation — cached results appear instantly. * Runs plot creation in a background thread; if it takes more than 10 s the display is updated with a 'Timed out' notice instead. * Any exception inside the plot thread produces a styled error message rather than propagating to the notebook. """ try: from IPython import display as _idisplay except ImportError: print(self.__repr__()) return header = html.escape(self.__repr__(bare=True)).replace("\n", "<br/>\n") # dask arrays: delegate to dask's own repr, no plot needed if self.is_dask_array(): _idisplay.display( _idisplay.HTML(header + "<br/>\n" + self.data._repr_html_()) ) return # cache look-up cache_key = self._repr_cache_key() if cache_key is not None: cache_path = self._repr_cache_path(cache_key) if cache_path.exists(): try: b64 = cache_path.read_text() _idisplay.display(_idisplay.HTML( header + "<br/>\n" + f"<img src='data:image/png;base64,{b64}'>" )) return except Exception: pass # corrupted cache entry — fall through to recompute # show placeholder immediately so the user sees something right away handle = _idisplay.display( _idisplay.HTML(header + "<br/>\n<em>Computing plot…</em>"), display_id=True, ) png_b64 = [None] err_msg = [None] def _compute(): try: with _mpl_repr_lock: png_b64[0] = self._get_repr_plot_png_b64() except NotImplementedError: pass # shape has no plot — leave png_b64 as None except Exception as exc: err_msg[0] = str(exc) t = threading.Thread(target=_compute, daemon=True) t.start() t.join(timeout=10.0) if t.is_alive(): body = "<em>Plot timed out (computation exceeded 10 s).</em>" elif err_msg[0] is not None: body = ( '<span style="color:#c0392b">Could not create plot: ' + html.escape(err_msg[0]) + "</span>" ) elif png_b64[0] is not None: # persist to cache if cache_key is not None: try: self._repr_cache_path(cache_key).write_text(png_b64[0]) except Exception: pass body = f"<img src='data:image/png;base64,{png_b64[0]}'>" else: body = "<em>No plot available for this array shape.</em>" handle.update(_idisplay.HTML(header + "<br/>\n" + body)) def _repr_html_(self): """HTML repr for static notebook export and non-IPython environments. In a live Jupyter session :meth:`_ipython_display_` is called instead, which provides the async 'Computing…' placeholder and caching. This method is the synchronous fallback used by nbconvert and similar tools. """ header = html.escape(self.__repr__(bare=True)).replace("\n", "<br />\n") if self.is_dask_array(): return header + "<br />\n" + self.data._repr_html_() try: b64 = self._get_repr_plot_png_b64() return ( header + "<br />\n" + f"<img src='data:image/png;base64,{b64}'>" ) except NotImplementedError: return header except Exception as exc: return ( header + "<br />\n" + '<span style="color:#c0392b">Could not create plot: ' + html.escape(str(exc)) + "</span>" )
# Inject numpy/dask delegate methods into Array for _name, (_np, _da, _ax, _esc) in _ARRAY_DELEGATE_METHODS.items(): setattr(Array, _name, _make_array_method(_name, _np, _da, _ax, _esc)) del _name, _np, _da, _ax, _esc # keep module namespace tidy def load_from_h5_file(file_name, name, parent_group_name=""): h5 = ArrayH5File( file_name=file_name, parent_group_name=parent_group_name, name=name ) with h5py.File(file_name, "r") as f: parameter, step_lengths, grid_specs = Scan._load_from_h5(f[h5.group_name]) return Array( index=h5.index, data=h5.get_data_da(), parameter=parameter, step_lengths=step_lengths, grid_specs=grid_specs, name=name, ) def load_from_h5(parent_h5py, name): h5 = ArrayH5Dataset(parent_h5py, name) parameter, step_lengths, grid_specs = Scan._load_from_h5(parent_h5py) return Array( index=h5.index, data=h5.get_data_da(), parameter=parameter, step_lengths=step_lengths, grid_specs=grid_specs, name=name, )
[docs] def escaped(func, convertOutput2EscData="auto"): """Decorator that lifts a NumPy/dask function to operate on escape Arrays. When *any* positional or keyword argument is an :class:`Array`, the decorator automatically: 1. Finds the intersection of all Array indices. 2. Aligns (re-indexes) every Array argument to the common pulse IDs, using the **first** Array found as the ordering reference. 3. Passes the aligned *raw data* arrays to ``func``. 4. Wraps each output whose length matches the number of common events back into a new :class:`Array` carrying the correct index and scan metadata. Non-Array arguments are passed through unchanged. Parameters ---------- func : callable Any function that accepts NumPy or dask arrays. convertOutput2EscData : "auto" or list of int Which output positions to wrap as Arrays. ``"auto"`` (default) wraps every output whose length equals the number of common events. Pass a list of integer indices (e.g. ``[0]``) to wrap only specific outputs. Returns ------- callable Wrapped function with the same signature as ``func``, plus an optional ``escSorter`` keyword to override the master Array (default: first Array found in the argument list). Examples -------- Decorate a function at definition time:: @escape.escaped def normalise(signal, reference): return signal / reference sig_norm = normalise(sig, i0) # sig_norm is an escape.Array Or apply to an existing NumPy function:: my_polyfit = escape.escaped(np.polyfit) """ def wrapped( *args, escSorter="first", convertOutput2EscData=convertOutput2EscData, **kwargs ): args = [ta for ta in args] kwargs = {tk: tv for tk, tv in kwargs.items()} argsIsEsc = [(n, arg) for n, arg in enumerate(args) if isinstance(arg, Array)] kwargsIsEsc = { key: kwarg for key, kwarg in kwargs.items() if isinstance(kwarg, Array) } allEscs = [a for n, a in argsIsEsc] allEscs.extend(kwargsIsEsc.values()) if escSorter == "first": if len(allEscs) > 0: sorter = allEscs[0] else: sorter = None print( "Did not find any Array instance \ in input parameters!" ) else: sorter = escSorter if not sorter is None: ixsorter = allEscs.index(sorter) allEscs.pop(ixsorter) ixmaster, ixslaves, stepLengthsNew = match_indexes( sorter.index, [t.index for t in allEscs] ) ixslaves.insert(ixsorter, ixmaster) ids_res = sorter.index[ixmaster] for n, arg in argsIsEsc: args.pop(n) args.insert(n, arg.data[ixslaves.pop(0)]) for key, kwarg in kwargsIsEsc.items(): kwargs.pop(key) kwargs[key] = kwarg.data[ixslaves.pop(0)] output = func(*args, **kwargs) if not type(output) is tuple: single_output = True output = (output,) else: single_output = False output = list(output) if convertOutput2EscData: stepLengths, scan = get_scan_step_selections( ixmaster, sorter.scan.step_lengths, scan=sorter.scan ) if convertOutput2EscData == "auto": convertOutput2EscData = [] for i, toutput in enumerate(output): try: lentoutput = len(toutput) if len(ids_res) == len(toutput): convertOutput2EscData.append(i) except TypeError: pass for n in convertOutput2EscData: toutput = output.pop(n) output.insert( n, Array( data=toutput, index=ids_res, step_lengths=stepLengths, parameter=scan.parameter, grid_specs=scan.grid.get_grid_specs() if hasattr(scan, "grid") else None, ), ) if len(output) == 1: output = output[0] elif len(output) == 0: output = None return output return wrapped
def scan_escaped(func): def wrapped(*args, **kwargs): args = [ta for ta in args] kwargs = {tk: tv for tk, tv in kwargs.items()} argsIsEsc = [(n, arg) for n, arg in enumerate(args) if isinstance(arg, Scan)] kwargsIsEsc = { key: kwarg for key, kwarg in kwargs.items() if isinstance(kwarg, Scan) } allEscs = [a for n, a in argsIsEsc] allEscs.extend(kwargsIsEsc.values()) scan_lens = [] for tscan in allEscs: scan_lens.append(len(tscan)) if not np.unique(scan_lens): raise Exception( "Scan instances do not have same length, bring them on the same scan parameter set, e.g. using the digitize method!" ) scan_len = scan_lens[0] args_same_len = [] for ta in args: try: tlen = len(ta) except TypeError: continue if tlen == scan_len: args_same_len.append(ta) kwargs_same_len = {} for tn, ta in kwargs.items(): try: tlen = len(ta) except TypeError: continue if tlen == scan_len: kwargs_same_len[tn] = ta res = [] for i_step in range(scan_len): targs = [] for ta in args: if id(ta) in [id(tmp) for tmp in args_same_len]: targs.append(ta[i_step]) else: targs.append(ta) tkwargs = {} for tn, ta in kwargs.items(): if id(ta) in [id(tmp) for tmp in kwargs_same_len.values()]: tkwargs[tn] = ta[i_step] else: tkwargs[tn] = ta tres = func(*targs, **tkwargs) if not isinstance(tres, tuple): tres = (tres,) res.append(tres) fres = [] for tres in zip(*res): if isinstance(tres[0], Array): fres.append(concatenate(list(tres))) else: fres.append(tres) if len(fres) > 1: return tuple(fres) else: return fres[0] return wrapped # if escSorter is "first": # if len(allEscs) > 0: # sorter = allEscs[0] # else: # sorter = None # print( # "Did not find any Array instance \ # in input parameters!" # ) # else: # sorter = escSorter # if not sorter is None: # ixsorter = allEscs.index(sorter) # allEscs.pop(ixsorter) # ixmaster, ixslaves, stepLengthsNew = match_indexes( # sorter.index, [t.index for t in allEscs] # ) # ixslaves.insert(ixsorter, ixmaster) # ids_res = sorter.index[ixmaster] # for n, arg in argsIsEsc: # args.pop(n) # args.insert(n, arg.data[ixslaves.pop(0)]) # for key, kwarg in kwargsIsEsc.items(): # kwargs.pop(key) # kwargs[key] = kwarg.data[ixslaves.pop(0)] # output = func(*args, **kwargs) # if not type(output) is tuple: # output = (output,) # output = list(output) # if convertOutput2EscData: # stepLengths, scan = get_scan_step_selections( # ixmaster, sorter.scan.step_lengths, scan=sorter.scan # ) # if convertOutput2EscData == "auto": # convertOutput2EscData = [] # for i, toutput in enumerate(output): # try: # lentoutput = len(toutput) # if len(ids_res) == len(toutput): # convertOutput2EscData.append(i) # except TypeError: # pass # for n in convertOutput2EscData: # toutput = output.pop(n) # output.insert( # n, # Array( # data=toutput, # index=ids_res, # step_lengths=stepLengths, # parameter=scan.parameter, # ), # ) # if len(output) == 1: # output = output[0] # elif len(output) == 0: # output = None # return output # return wrapped def _scan_wrap(func, **default_kws): def wrapped(scan, **kwargs): default_kws.update(kwargs) return [func(step.data, **default_kws) for step in scan] return wrapped _operatorsJoin = [ (operator.add, "+"), (operator.contains, "in"), (operator.truediv, "/"), (operator.floordiv, "//"), (operator.and_, "&"), (operator.xor, "^"), (operator.or_, "|"), (operator.pow, "**"), (operator.is_, "is"), (operator.is_not, "is not"), (operator.lshift, "<<"), (operator.mod, "%"), (operator.mul, "*"), (operator.rshift, ">>"), (operator.sub, "-"), (operator.lt, "<"), (operator.le, "<="), (operator.eq, "=="), (operator.ne, "!="), (operator.ge, ">="), (operator.gt, ">"), ] _operatorsSingle = [ (operator.invert, "~"), (operator.neg, "-"), (operator.not_, "not"), (operator.pos, "pos"), ] for opJoin, symbol in _operatorsJoin: setattr( Array, "__%s__" % opJoin.__name__.strip("_"), escaped(opJoin, convertOutput2EscData=[0]), ) if ( True ): # any(top in opJoin.__name__ for top in ["add", "sub", "mul", "div", "mod"]): setattr( Array, "__r%s__" % opJoin.__name__.strip("_"), escaped(opJoin, convertOutput2EscData=[0]), ) for opSing, symbol in _operatorsSingle: setattr( Array, "__%s__" % opSing.__name__.strip("_"), escaped(opSing, convertOutput2EscData=[0]), ) def match_scans(a0, a1, parameters=[]): """Match scans of two escape arrays.""" s0 = a0.scan s1 = a1.scan s0arr = np.asarray( [ value["values"] for name, value in s0.parameter.items() if ((not parameters) or (name in parameters)) ] ).T s1arr = np.asarray( [ value["values"] for name, value in s1.parameter.items() if ((not parameters) or (name in parameters)) ] ).T s0_sel = ((np.isin(s0arr, s1arr).all(axis=1))).nonzero()[0] s1_sel = ((np.isin(s1arr, s0arr).all(axis=1))).nonzero()[0] return get_step_indexes(s0, s0_sel), get_step_indexes(s1, s1_sel)
[docs] class Grid:
[docs] def __init__(self, shape, positions, scan=None, grid_dimension_names=None,**kwargs): self.scan = scan self.shape = shape self.positions = positions self.dimension_names = grid_dimension_names
[docs] def get_grid_indices(self): grid_indices = [ tmp['grid_index'] for tmp in self.scan.parameter['scan_step_info']['values'] ] return grid_indices
def _normalize_selector(self, sel, dim_size): if isinstance(sel, slice): return set(range(*sel.indices(dim_size))) if isinstance(sel, (list, tuple, np.ndarray)): normalized = set() for i in sel: if not isinstance(i, (int, np.integer)): raise TypeError(f"Invalid index type in sequence: {type(i)}") if i < 0: i += dim_size if i < 0 or i >= dim_size: raise IndexError( f"index {i} is out of bounds for axis with size {dim_size}" ) normalized.add(int(i)) return normalized if isinstance(sel, (int, np.integer)): if sel < 0: sel = dim_size + sel if sel < 0 or sel >= dim_size: raise IndexError(f"index {sel} is out of bounds for axis with size {dim_size}") return {int(sel)} raise TypeError(f"Invalid index type: {type(sel)}") def _get_subgrid_specs(self, selectors): new_shape = [len(sel) for sel in selectors] if self.positions is None: new_positions = None else: new_positions = [] for pos, selected in zip(self.positions, selectors): if pos is None: new_positions.append(None) continue pos_arr = np.asarray(pos) indices = np.array(sorted(selected), dtype=int) new_positions.append(pos_arr[indices]) new_dimension_names = None if self.dimension_names is not None: new_dimension_names = [ name for name, sel in zip(self.dimension_names, selectors) ] return { "shape": new_shape, "positions": new_positions, "grid_dimension_names": new_dimension_names, }
[docs] def __getitem__(self, sel): if not hasattr(self, "scan") or self.scan is None: raise AttributeError("Grid instance has no associated Scan") ndim = len(self.shape) if not isinstance(sel, tuple): sel = (sel,) if Ellipsis in sel: if sel.count(Ellipsis) > 1: raise IndexError("an index can only have a single ellipsis") ellipsis_index = sel.index(Ellipsis) sel = ( *sel[:ellipsis_index], *[slice(None)] * (ndim - (len(sel) - 1)), *sel[ellipsis_index + 1 :], ) if len(sel) < ndim: sel = tuple(list(sel) + [slice(None)] * (ndim - len(sel))) if len(sel) > ndim: raise IndexError("too many indices for Grid") selectors = [self._normalize_selector(s, dim_size) for s, dim_size in zip(sel, self.shape)] selected_steps = [] for step_idx, grid_index in enumerate(self.get_grid_indices()): if all(grid_index[dim] in selectors[dim] for dim in range(ndim)): selected_steps.append(step_idx) if not selected_steps: raise IndexError("Grid selection returned no matching steps") grid_specs = self._get_subgrid_specs(selectors) if len(selected_steps) == 1: return self.scan.__getitem__(selected_steps[0], grid_specs=grid_specs) return self.scan.__getitem__(selected_steps, grid_specs=grid_specs)
[docs] def to_grid(self, data): cdata = np.asanyarray(data) grid_shape = self.shape grid_data = np.empty(list(grid_shape) + list(np.shape(cdata)[1:]), dtype=cdata.dtype) grid_data[:] = np.nan # or any other fill value for tdat,step_index in zip(data,self.get_grid_indices()): grid_data[tuple(step_index)] = tdat return grid_data
[docs] def get_grid_specs(self): return { "shape": self.shape, "positions": self.positions, "grid_dimension_names": self.dimension_names }
[docs] def fill_count(self): """Return (filled_positions, total_positions, percent_filled). filled_positions is the number of unique grid indices present in the associated Scan. total_positions is the product of `self.shape`. """ if not hasattr(self, "scan") or self.scan is None: return 0, int(np.prod(self.shape)), 0.0 try: grid_indices = self.get_grid_indices() except Exception: grid_indices = [] # Normalize to tuples and deduplicate unique_indices = {tuple(idx) for idx in grid_indices} filled = len(unique_indices) total = int(np.prod(self.shape)) percent = (filled / total * 100.0) if total else 0.0 return filled, total, percent
def __repr__(self): filled, total, percent = self.fill_count() dims = self.dimension_names if self.dimension_names is not None else [] return f"<Grid shape={tuple(self.shape)} dims={dims} filled={filled}/{total} ({percent:0.1f}%)>"
# Names must match methods on Array (injected by _ARRAY_DELEGATE_METHODS) and Scan. _SCAN_STEP_DELEGATE = [ "nansum", "nanmean", "nanstd", "nanmedian", "nanmin", "nanmax", "nanpercentile", "nanquantile", "sum", "mean", "average", "std", "median", "percentile", "quantile", "min", "max", "all", "any", "abs", "isnan", "isinf", "isfinite", ] SCAN_STEP_METHODS = [ *_SCAN_STEP_DELEGATE, # escape-specific aggregators "count", "nancount", "median_and_mad", # "weighted_median_and_mad", "weighted_avg_and_std", "weighted_stat", "correlation_analysis_to", ] # Dynamically attach wrappers for scan step methods that reformat # their per-step outputs into grid-shaped arrays using `to_grid` def _make_grid_scan_wrapper(method_name): def _wrapper(self, *args, **kwargs): """Wraps Scan.{method} with grid reshaping and optional 2-D plotting. Additional keyword arguments (consumed here, not forwarded to the scan): Parameters ---------- plot : bool or matplotlib.axes.Axes or matplotlib.figure.Figure, optional If ``True``, plot the grid-shaped result with :func:`escape.plot2D` using the current axes. Pass a ``matplotlib.axes.Axes`` to target a specific axes, or a ``matplotlib.figure.Figure`` to open a new subplot. No plot is created by default. plot_kws : dict, optional Extra keyword arguments for :func:`escape.plot2D`. The special keys ``'colorbar'`` (bool, default ``True``) and ``'axis'`` (``matplotlib.axes.Axes``) are handled here and not forwarded. """ import matplotlib.axes as _mplaxes # Extract plotting options before forwarding to the underlying scan method plot_opt = kwargs.pop("plot", None) plot_kws = kwargs.pop("plot_kws", {}) or {} if not hasattr(self, "scan") or self.scan is None: raise AttributeError("Grid instance has no associated Scan") if not hasattr(self.scan, method_name): raise AttributeError(f"Scan has no method '{method_name}'") func = getattr(self.scan, method_name) res = func(*args, **kwargs) def _try_to_grid(val): try: arr = np.asarray(val) except Exception: return val try: return self.to_grid(arr) except Exception: return val # Convert results into grid-shaped arrays where appropriate if isinstance(res, tuple): converted = tuple(_try_to_grid(r) for r in res) elif isinstance(res, list): converted = _try_to_grid(np.asarray(res)) else: converted = _try_to_grid(res) # Plotting: if requested and we have a 2D numpy array, call plot2D if plot_opt: try: candidate = converted[0] if isinstance(converted, tuple) else converted if isinstance(candidate, np.ndarray) and candidate.ndim == 2: # copy plot_kws early so we can pop special keys _pkw = dict(plot_kws) add_colorbar = _pkw.pop("colorbar", True) axis_from_kws = _pkw.pop("axis", None) # resolve target axes, in priority order if axis_from_kws is not None: axis = axis_from_kws elif isinstance(plot_opt, _mplaxes.Axes): axis = plot_opt elif hasattr(plot_opt, "add_subplot"): axis = plot_opt.add_subplot(111) elif plot_opt is True: axis = plt.gca() else: axis = plt.gca() # prepare x/y coordinate arrays from grid positions positions = getattr(self, "positions", None) if positions and len(positions) >= 2: x_raw = positions[1] y_raw = positions[0] class _Named: def __init__(self, arr, name=None): self._arr = np.asarray(arr) self.name = name def __array__(self, dtype=None): return self._arr def __len__(self): return len(self._arr) x_named = _Named(x_raw) y_named = _Named(y_raw) try: if self.dimension_names and len(self.dimension_names) > 1: x_named.name = self.dimension_names[1] y_named.name = self.dimension_names[0] except Exception: pass else: x_named = "auto" y_named = "auto" p = plot2D(x_named, y_named, candidate, axis=axis, **_pkw) if add_colorbar: try: plt.colorbar(p, ax=axis) except Exception: pass except Exception: # plotting must not break main functionality pass return converted _wrapper.__doc__ = (_wrapper.__doc__ or "").replace("{method}", method_name) return _wrapper for _method in SCAN_STEP_METHODS: setattr(Grid, _method, _make_grid_scan_wrapper(_method))
[docs] class Scan: """Scan grouping of an Array across defined steps and optional grid metadata. ``Scan`` partitions an ``Array`` into sequential steps defined by ``step_lengths`` and exposes step-level metadata through ``parameter``. When ``grid_specs`` is provided, ``Scan`` also constructs a ``Grid`` that maps scan steps onto an N-D grid layout. Args: parameter: Metadata describing step parameters and values. step_lengths: List of integer lengths for each scan step. array: Underlying ``Array`` instance for this scan. data: Optional raw data used directly by the scan. grid_specs: Optional grid metadata forwarded to ``Grid``. """
[docs] def __init__(self, parameter={}, step_lengths=None, array=None, data=None, grid_specs=None): self.step_lengths = step_lengths self._tools = None self.__step_index_ranges = None if parameter: for par, pardict in parameter.items(): if not len(pardict["values"]) == len(self): raise Exception( f"Parameter array length of {par} does not fit the defined steps." ) else: parameter = {"none": {"values": [1] * len(step_lengths)}} self.parameter = parameter self._array = array # self._add_methods() if data is not None: self._data = data if grid_specs is not None: self.grid = Grid(**grid_specs, scan=self)
@property def _step_index_ranges(self): if self.__step_index_ranges is None: self.__step_index_ranges = np.cumsum( np.hstack([0, np.asarray(self.step_lengths)]) ) return self.__step_index_ranges @property def tools(self): if self._tools is None: self._tools = ScanTools(self) return self._tools
[docs] def append_parameter(self, parameter: {"par_name": {"values": list}}): for par, pardict in parameter.items(): if not len(pardict["values"]) == len(self): lenthis = len(pardict["values"]) raise Exception( f"Parameter array length of {par} ({lenthis}) does not fit the defined steps ({len(self)})." ) self.parameter.update(parameter)
@property def par_steps(self): """pandas.DataFrame with one row per scan step. Columns are the scan parameter names (one column per parameter) plus a ``step_length`` column containing the number of events in each step. The integer row index corresponds to the step number. Returns ------- pandas.DataFrame """ data = {name: value["values"] for name, value in self.parameter.items()} data.update({"step_length": self.step_lengths}) return pd.DataFrame(data, index=list(range(len(self))))
[docs] def steps_where(self, data_condition): """Return an Array containing only the steps where *data_condition* is True. Similar to ``np.where`` but operating at the level of scan steps: each step is tested and only the passing steps are kept, with their scan metadata preserved. Parameters ---------- data_condition : callable ``f(step: Array) -> bool``. Called once per step with the step's :class:`Array`; steps for which it returns a truthy value are included in the result. Returns ------- escape.Array Array whose scan contains only the steps that passed the condition. Examples -------- Keep only steps whose per-step median exceeds 0.5:: result = sig.scan.steps_where(lambda step: np.nanmedian(step.data) > 0.5) Keep only steps that have at least 50 events:: result = sig.scan.steps_where(lambda step: len(step) >= 50) """ step_indices = [n for n, step in enumerate(self) if data_condition(step)] return self[step_indices]
[docs] def count(self): """Return the number of events in each scan step as a list.""" return [len(step) for step in self]
[docs] def nancount(self): """Return the number of non-NaN events in each scan step as a list.""" return [step.nancount() for step in self]
# Remaining step-delegation methods (nansum, nanmean, …, all, any) are # injected below the class definition via _SCAN_STEP_DELEGATE + _make_scan_step_method.
[docs] def median_and_mad(self, axis=None, k_dist=1.4826, norm_samples=False): """Calculate median and median absolute deviation for steps of a scan. Args: axis (int, sequence of int, None, optional): axis argument for median calls. k_dist (float, optional): distribution scale factor, should be 1 for real MAD. Defaults to 1.4826 for gaussian distribution. """ # if self._array.is_dask_array(): # absfoo = da.abs # else: # absfoo = np.abs med = [step.median(axis=axis) for step in self] mad = [ (((step - tmed).abs()) * k_dist).median(axis=axis) for step, tmed in zip(self, med) ] if norm_samples: mad = [tmad / da.sqrt(ct) for tmad, ct in zip(mad, self.count())] return med, mad
# def weighted_median_and_mad(self, weights=None, axis=None, k_dist=1.4826, norm_samples=False): # """Calculate median and median absolute deviation for steps of a scan. # Args: # axis (int, sequence of int, None, optional): axis argument for median calls. # k_dist (float, optional): distribution scale factor, should be # 1 for real MAD. # Defaults to 1.4826 for gaussian distribution. # """ # # if self._array.is_dask_array(): # # absfoo = da.abs # # else: # # absfoo = np.abs # utilities.weighted_quantiles(0.5) # med = [step.median(axis=axis) for step in self] # mad = [ # (((step - tmed).abs()) * k_dist).median(axis=axis) # for step, tmed in zip(self, med) # ] # if norm_samples: # mad = [tmad / da.sqrt(ct) for tmad, ct in zip(mad, self.count())] # return med, mad
[docs] def weighted_avg_and_std(self, weights=None, norm_samples=False, axis=0): avg = [] std = [] for step in self: if weights: (ta, tw) = match_arrays(step, weights) (tavg, tstd) = utilities.weighted_avg_and_std(ta.data, tw.data, axis=axis) else: (tavg, tstd) = utilities.weighted_avg_and_std(step.data, weights, axis=axis) avg.append(tavg) std.append(tstd) if norm_samples: std = [tstd / da.sqrt(ct) for tstd, ct in zip(std, self.count())] return da.asarray(avg), da.asarray(std)
[docs] def weighted_stat(self, weights=None): if weights is None: import warnings warnings.warn("weights not provided, using unweighted median and mad!") return self.median_and_mad() else: array, weightsf = escape.match_arrays(self._array, weights) qsig = 0.682689492 med = [] err = [] # if len(weights.shape) == 3: # weights = weights[:,0,0] for n, (ta, tw) in enumerate(zip(array.scan, weightsf.scan)): if len(ta.shape) == 3: print(f"step {n}/{len(array.scan)}") r = utilities.weighted_quantile( ta.data, [0.5 - qsig / 2, 0.5, 0.5 + qsig / 2], sample_weight=tw.data ) med.append(r[1]) err.append(np.diff(r) / np.sqrt(len(ta.data))) return np.asarray(med), np.asarray(err).T
[docs] def correlation_analysis_to(self, ref, *args, **kwargs): (td, tr) = match_arrays(self._array, ref) return [ step.correlation_analysis_to(tref, *args, **kwargs) for step, tref in zip(td.scan, tr.scan) ]
[docs] def corr_ana_plot(self, referece, scanpar_name=None, axis=None): if not scanpar_name: names = list(self.parameter.keys()) scanpar_name = names[0] x = np.asarray(self.parameter[scanpar_name]["values"]).ravel() corres = self.correlation_analysis_to(referece) if not axis: axis = plt.gca() std = [tc[0] for tc in corres] std_fx = [tc[1] for tc in corres] ordercolors = ["b", "r"] for to, toc in zip([0, 1], ordercolors): axis.plot( x, [tc[to] for tc in std], toc + "--" + ".", label=f"poly. order:{to+1}; zero free", ) for to, toc in zip([0, 1], ordercolors): axis.plot( x, [tc[to] for tc in std_fx], toc + "-" + "o", label=f"poly. order:{to+1}; zero fixed", )
[docs] def plot( self, weights=None, scanpar_name=None, norm_samples=True, axis=None, use_quantiles=True, *args, **kwargs, ): if not scanpar_name: names = list(self.parameter.keys()) scanpar_name = names[0] x = np.asarray(self.parameter[scanpar_name]["values"]).ravel() if self._array.ndim_nonzero == 1: if not weights: if use_quantiles: tmp = np.asarray( self.nanquantile( [0.5, 0.5 - 0.682689492137 / 2, 0.5 + 0.682689492137 / 2], # axis=0, ) ) y = tmp[:, 0] ystd = np.diff(tmp[:, 1:], axis=1)[:, 0] / 2 else: y = np.asarray(self.nanmean(axis=0)).ravel() ystd = np.asarray(self.nanstd(axis=0)).ravel() else: if use_quantiles: print( "Cannot use quantile derivation with weights, using average/std instead!" ) y, ystd = self.weighted_avg_and_std(weights) if norm_samples: yerr = ystd / np.sqrt(np.asarray(self.count())) else: yerr = ystd if not axis: axis = plt.gca() axis.errorbar(x, y, yerr=yerr, *args, **kwargs) axis.set_xlabel(scanpar_name) if self._array.name: axis.set_ylabel(self._array.name) elif self._array.ndim_nonzero == 2: if use_quantiles: tmp = np.asarray( self.nanquantile( [0.5, 0.5 - 0.682689492137 / 2, 0.5 + 0.682689492137 / 2], axis=0, ) ) ic = tmp[ :, 0, :, ] icstd = np.diff(tmp[:, 1:, :], axis=1)[:, 0, :] / 2 # return ic, icstd if not axis: axis = plt.gca() y = np.arange(ic.shape[1]) ih = plot2D(x, y, ic.T, ax=axis, *args, **kwargs) axis.set_xlabel(scanpar_name) plt.colorbar(ih, ax=axis, label=self._array.name) axis.set_ylabel("Median step waveform") return ic, icstd
[docs] def hist( self, cut_percentage=0, bins="auto", normalize_to=None, scanpar_name=None, plot_results=True, plot_axis=None, **kwargs, ): if self._array.is_dask_array(): raise Exception( "escape array needs to be numpy type for histogramming, compute first." ) if not scanpar_name: names = list(self.parameter.keys()) for scanpar_name in names: if not np.isnan( np.asarray(self.parameter[scanpar_name]["values"]).ravel() ).all(): break x_scan = np.asarray(self.parameter[scanpar_name]["values"]).ravel() flat = self._array.data.ravel().astype(float) [hmin, hmax] = np.nanpercentile(flat, [cut_percentage, 100 - cut_percentage]) if not (np.isfinite(hmin) and np.isfinite(hmax) and hmin < hmax): hmin, hmax = float(np.nanmin(flat)), float(np.nanmax(flat)) if hmin == hmax: hmin -= 0.5 hmax += 0.5 hbins = np.histogram_bin_edges( self._array.data.ravel(), bins, range=[hmin, hmax] ) hdat = [np.histogram(td.data.ravel(), bins=hbins)[0] for td in self] if normalize_to == "max": hdat = [td / td.max() for td in hdat] elif normalize_to == "sum": hdat = [td / td.sum() for td in hdat] hdat = np.asarray(hdat) if plot_results: if not plot_axis: plot_axis = plt.gca() # utilities.plot2D(x_scan, utilities.edges_to_center(hbins), hdat.T, **kwargs) plt.pcolormesh(x_scan, utilities.edges_to_center(hbins), hdat.T, **kwargs) plt.xlabel(scanpar_name) return x_scan, hbins, hdat
[docs] def append_step(self, parameter, step_length): self.step_lengths.append(step_length) for par, pardict in parameter: self.parameter[par]["values"].append(pardict["values"])
[docs] def __len__(self): return len(self.step_lengths)
[docs] def __getitem__(self, sel, grid_specs=None): """array getter for scan""" if grid_specs is None and hasattr(self, "grid"): grid_specs = self.grid.get_grid_specs() if isinstance(sel, slice): sel = range(*sel.indices(len(self))) if isinstance(sel, Number): if sel < 0: sel = len(self) + sel return self.get_step_array(sel, grid_specs=grid_specs) else: return concatenate([self.get_step_array(n, grid_specs=grid_specs) for n in sel], grid_specs=grid_specs)
[docs] def get_step_array(self, n, grid_specs=None): """array getter for scan""" assert n >= 0, "Step index needs to be positive" if n == 0 and self.step_lengths is None: data = self._array.data[:] index = self._array.index[:] step_lengths = self._array.step_lengths parameter = self._array.parameter # assert not self.step_lengths is None, "No step sizes defined." elif not n < len(self.step_lengths): raise IndexError(f"Only {len(self.step_lengths)} steps") else: # data = self._array.data[ # sum(self.step_lengths[:n]) : sum(self.step_lengths[: (n + 1)]) # ] # index = self._array.index[ # sum(self.step_lengths[:n]) : sum(self.step_lengths[: (n + 1)]) # ] ix_range = self._step_index_ranges[n : (n + 2)] data = self._array.data[slice(*ix_range)] index = self._array.index[slice(*ix_range)] step_lengths = [self.step_lengths[n]] parameter = {} for par_name, par in self.parameter.items(): parameter[par_name] = {} parameter[par_name]["values"] = [par["values"][n]] if "attributes" in par.keys(): parameter[par_name]["attributes"] = par["attributes"] return Array( data=data, index=index, parameter=parameter, step_lengths=step_lengths, grid_specs=grid_specs )
[docs] def get_step_data(self, n): """data getter for scan""" assert n >= 0, "Step index needs to be positive" if n == 0 and self.step_lengths is None: data = self._array.data[:] # assert not self.step_lengths is None, "No step sizes defined." elif not n < len(self.step_lengths): raise IndexError(f"Only {len(self.step_lengths)} steps") else: ix_range = self._step_index_ranges[n : (n + 2)] data = self._array.data[slice(*ix_range)] return data
[docs] def merge_scans(self, *others, roundto_interval=None, par_name=None): if roundto_interval is None: raise Exception("Please provide roundto value for merging scans.") if par_name is None: par_name = self.par_steps.keys()[0] print(f"Using {par_name} as scan parameter for merging scans.") pars = self.par_steps[par_name] all_pars = roundto(pars, roundto_interval) for other in others: all_pars = np.union1d( all_pars, roundto(other.par_steps[par_name], roundto_interval) ) # all_pars = sorted(all_pars) # return all_pars data = [] index = [] step_lengths = [] for par in all_pars: step_len = 0 for thisscan in [self] + list(others): thisscan_values = roundto( thisscan.par_steps[par_name].values, roundto_interval ) if par in thisscan_values: ind = np.where(thisscan_values == par)[0][0] data.append(thisscan[ind].data) index.append(thisscan[ind].index) step_len += len(thisscan[ind]) step_lengths.append(step_len) if any([isinstance(tmp, da.Array) for tmp in data]): data = da.concatenate(data) else: data = np.concatenate(data) index = np.concatenate(index) return Array( data=data, index=index, step_lengths=step_lengths, parameter={par_name: {"values": list(all_pars)}}, )
[docs] def get_step_indexes(self, ix_step): """ "array getter for multiple steps, more efficient than get_step_array""" ix_to = np.cumsum(self.step_lengths) ix_from = np.hstack([np.asarray([0]), ix_to[:-1]]) index_sel = np.concatenate( [ self._array.index[fr:to] for fr, to in zip(ix_from[ix_step], ix_to[ix_step]) ], axis=0, ) return self._array[np.isin(self._array.index, index_sel).nonzero()[0]]
def _check_consistency(self): for par, pardict in self.parameter.items(): if not len(self) == len(pardict["values"]): raise Exception(f"Scan length does not fit parameter {par}")
[docs] def get_parameter_selection(self, selection): selection = np.atleast_1d(selection) if selection.dtype == bool: selection = selection.nonzero()[0] par_out = {} for par, pardict in self.parameter.items(): par_out[par] = {} par_out[par]["values"] = [pardict["values"][i] for i in selection] if "attributes" in pardict.keys(): par_out[par]["attributes"] = pardict["attributes"] return par_out
[docs] def get_parameter_array(self, key=None): if key: keys = [key] else: keys = self.parameter.keys() dall = {} iall = self._array.index for n, stl in enumerate(self.step_lengths): ons = np.ones(stl) for key in keys: tval = self.parameter[key]["values"][n] if not isinstance(tval, Number): continue if not key in dall.keys(): dall[key] = [] dall[key].append(ons * tval) arrays = [ Array( data=np.hstack(dall[key]), index=iall, name=key, step_lengths=self.step_lengths, parameter=self.parameter, ) for key in dall.keys() ] return arrays if len(arrays) > 1 else arrays[0]
def _save_to_h5(self, group): self._check_consistency() if "scan" in group.keys(): del group["scan"] try: scan_group = group.require_group("scan", track_order=True) except: scan_group = group.require_group("scan") scan_group["step_lengths"] = self.step_lengths par_group = scan_group.require_group("parameter") for parname, pardict in self.parameter.items(): tpg = par_group.require_group(parname) try: tpg["values"] = pardict["values"] except: tpg["values"] = [np.nan] * len(self) if "attributes" in pardict.keys(): tpg.require_group("attributes") for attname, attvalue in pardict["attributes"].items(): if isinstance(attvalue, str): attvalue = str(attvalue) tpg["attributes"][attname] = attvalue if hasattr(self, "grid"): gridspecs = self.grid.get_grid_specs() hickle.dump(gridspecs, scan_group, path="grid_specs") @staticmethod def _load_from_h5(group): if "scan" not in group.keys(): raise Exception("Did not find group scan!") step_lengths = group["scan"]["step_lengths"][()] parameter = {} for parname, pargroup in group["scan"]["parameter"].items(): values = pargroup["values"] if not len(values) == len(step_lengths): raise Exception( f"The length of data array in {parname} parameter in scan does not fit!" ) parameter[parname] = {} parameter[parname]["values"] = list(values[()]) if "attributes" in pargroup.keys(): parameter[parname]["attributes"] = {} for att_name, att_data in pargroup["attributes"].items(): parameter[parname]["attributes"][att_name] = att_data[()] grid_specs = None if "grid_specs" in group["scan"].keys(): try: grid_specs = hickle.load(group["scan"], path="grid_specs") except Exception: grid_specs = None return parameter, step_lengths, grid_specs def __repr__(self): s = "Scan over {} steps".format(len(self)) s += "\n" s += "Parameters {}".format(", ".join(self.parameter.keys())) return s
# def __add__(self,other): # return scan_escaped(operator.add)(self,other) _operatorsJoin = [ (operator.add, "+"), (operator.truediv, "/"), (operator.floordiv, "//"), (operator.and_, "&"), (operator.xor, "^"), (operator.or_, "|"), (operator.pow, "**"), (operator.is_, "is"), (operator.is_not, "is not"), (operator.lshift, "<<"), (operator.mod, "%"), (operator.mul, "*"), (operator.rshift, ">>"), (operator.sub, "-"), (operator.lt, "<"), (operator.le, "<="), (operator.ne, "!="), (operator.ge, ">="), (operator.gt, ">"), # (operator.contains, "in"), # (operator.eq, "=="), ] _operatorsSingle = [ (operator.invert, "~"), (operator.neg, "-"), (operator.not_, "not"), (operator.pos, "pos"), ] for opJoin, symbol in _operatorsJoin: setattr(Scan, "__%s__" % opJoin.__name__.strip("_"), scan_escaped(opJoin)) for opSing, symbol in _operatorsSingle: setattr(Scan, "__%s__" % opSing.__name__.strip("_"), scan_escaped(opSing)) # --------------------------------------------------------------------------- # Programmatic step-delegation method injection for Scan # --------------------------------------------------------------------------- def _make_scan_step_method(name): """Factory: build a Scan method that applies Array.{name} per step.""" np_func = getattr(np, name, None) np_summary = next( (line.strip() for line in (np_func.__doc__ or "").split("\n") if line.strip()), "" ) if np_func else "" doc = ( f"Apply :meth:`Array.{name}` to each scan step.\n\n" + (f"{np_summary}\n\n" if np_summary else "") + "Returns a list with one result per scan step.\n\n" "Parameters\n" "----------\n" "plot : bool or matplotlib.axes.Axes or matplotlib.figure.Figure, optional\n" " If ``True``, plot the per-step results against the scan parameter using\n" " the current axes. Pass a ``matplotlib.axes.Axes`` to target a specific\n" " axes, or a ``matplotlib.figure.Figure`` to open a new subplot. No plot\n" " is created by default.\n" "plot_kws : dict, optional\n" " Extra keyword arguments forwarded to the plot call. Scalar-per-step\n" " results are drawn with ``Axes.plot``; array-per-step results are\n" " rendered as a 2-D heat-map via :func:`escape.plot2D`. The special key\n" " ``'colorbar'`` (bool, default ``True``) toggles the colourbar for 2-D\n" " plots.\n" ) def method(self, *args, **kwargs): plot_opt = kwargs.pop("plot", None) plot_kws = kwargs.pop("plot_kws", {}) or {} result = [getattr(step, name)(*args, **kwargs) for step in self] if plot_opt: try: import matplotlib.axes as _mplaxes if isinstance(plot_opt, _mplaxes.Axes): axis = plot_opt elif hasattr(plot_opt, "add_subplot"): axis = plot_opt.add_subplot(111) elif plot_opt is True: axis = plt.gca() else: axis = plt.gca() par_names = list(self.parameter.keys()) scanpar_name = par_names[0] if par_names else None x = ( np.asarray(self.parameter[scanpar_name]["values"]).ravel() if scanpar_name else np.arange(len(result)) ) try: arr = np.asarray(result, dtype=float) except Exception: arr = None if arr is not None: _pkw = dict(plot_kws) add_colorbar = _pkw.pop("colorbar", True) if arr.ndim == 1: axis.plot(x, arr, **_pkw) if scanpar_name: axis.set_xlabel(scanpar_name) if self._array.name: axis.set_ylabel(f"{name}({self._array.name})") elif arr.ndim == 2: y = np.arange(arr.shape[1]) p = plot2D(x, y, arr.T, axis=axis, **_pkw) if scanpar_name: axis.set_xlabel(scanpar_name) if add_colorbar: try: plt.colorbar( p, ax=axis, label=f"{name}({self._array.name})" if self._array.name else name, ) except Exception: pass except Exception: pass # plotting must not break computation return result method.__name__ = name method.__qualname__ = f"Scan.{name}" method.__doc__ = doc return method for _name in _SCAN_STEP_DELEGATE: setattr(Scan, _name, _make_scan_step_method(_name)) del _name def to_dataframe(*args): """work in progress""" for arg in args: if not np.prod(arg.shape) == len(arg): raise ( NotImplementedError("Only 1D Arrays can be converted to dataframes.") ) dfs = [ ddf.from_dask_array(arg.data.ravel(), columns=[arg.name], index=arg.index) for arg in args ] return ddf.concat(dfs, axis=0, join="outer", interleave_partitions=False)
[docs] @escaped def match_arrays(*args): """Return a tuple of Arrays restricted to their common pulse IDs. All output Arrays share the same index and are ordered by the first argument's pulse-ID sequence. This is the building block of index-aligned arithmetic in ``escape``. Parameters ---------- *args : escape.Array Two or more Arrays to align. Returns ------- tuple of escape.Array One Array per input, all restricted to the common event indices. Examples -------- >>> sig_m, i0_m = escape.match_arrays(sig, i0) >>> assert (sig_m.index == i0_m.index).all() """ return args
weighted_avg_and_std = escaped(utilities.weighted_avg_and_std)
[docs] def compute(*args): """compute multiple escape arrays or dask arrays. Interesting when calculating multiple small arrays from the same ancestor dask based array""" argtypes = [] argcollection = [] for arg in args: targtype = [] if isinstance(arg, Array): targtype.append("esc-array") if arg.is_dask_array(): targtype.append("dask_array") argcollection.append(arg.data) elif isinstance(arg, DaskCollection): targtype.append("daskcollection") if isinstance(arg, da.Array): targtype.append("dask_array") argcollection.append(arg) else: targtype.append("nodask") argtypes.append(targtype) with ProgressBar(): res = da.compute(*argcollection) next_dask_index = 0 out = [] for ta, argtype in zip(args, argtypes): if ("esc-array" in argtype) and ("dask_array" in argtype): out.append( Array( data=res[next_dask_index], index=ta.index, step_lengths=ta.scan.step_lengths, parameter=ta.scan.parameter, ) ) next_dask_index += 1 elif ("daskcollection" in argtype) and ("dask_array" in argtype): out.append(res[next_dask_index]) next_dask_index += 1 else: out.append(ta) return tuple(out)
[docs] def store(arrays, lock="auto", **kwargs): """ Storing of multiple escape arrays (as iterable, list or similar), efficient when they originate from the same ancestor """ if lock == "auto": lock = get_lock() prep = [ array.h5.append(array.data, array.index, prep_run="store_numpy") for array in arrays ] # return prep if not any(prep): print("Nothing to append") # arrays_store = arrays else: arrays_store, ndatas, dsets, n_news = zip( *[(tarray, *tprep) for tarray, tprep in zip(arrays, prep) if tprep] ) if is_local_client_distributed(): client = get_client() print("Found dask distributed client") t_tmp = time.time() storage_return = da.store(ndatas, dsets, lock=lock, compute=False, **kwargs) future = client.persist(storage_return) # print("persisted, length of future: ", len(future)) # from IPython.display import display progress(future,notebook=False) # print('calling wait on future') # time.sleep(0.2) wait(future) print(f'storing data done in {time.time() - t_tmp} s.') else: with ProgressBar(): da.store(ndatas, dsets, lock=lock, **kwargs) for array, n_new in zip(arrays_store, n_news): array.h5._n_i.append(n_new) array.h5._n_d.append(n_new) for array in arrays: array._data = array.h5.get_data_da() array._index = array.h5.index array.scan._save_to_h5(array.h5.grp)
def store_all( arrays, parent_h5py=None, lock="auto", names=None, indexes=None, scans=None, **kwargs ): """ Store a mix of escape arrays and raw dask/numpy arrays efficiently. Handles three types of inputs: 1. Escape arrays with h5 already configured (via set_h5_storage) 2. Escape arrays without h5 configured (will set up h5 storage) 3. Raw dask/numpy arrays (will wrap in escape.Array objects with h5 storage) Args: arrays: List of mixed array types (escape.Array, da.Array, or np.ndarray) parent_h5py: h5py parent group/file for storing data lock: Lock mechanism for dask store operations ("auto" or specific lock) names: Optional list of dataset names (defaults to array.name or auto-generated) indexes: Optional list of event IDs for each array scans: Optional list of Scan objects for each array **kwargs: Additional arguments passed to da.store() Returns: List of escape.Array objects with h5 storage configured Example: >>> # Mix of escape arrays and raw dask arrays >>> esc_array = escape.Array(data=my_data, index=my_index, name="array1") >>> raw_dask = da.from_delayed(...) >>> result = store_all( ... [esc_array, raw_dask], ... parent_h5py=h5file, ... indexes=[esc_array.index, raw_index], ... names=["array1", "array2"] ... ) """ if lock == "auto": lock = get_lock() if parent_h5py is None: raise ValueError("parent_h5py must be provided") # Standardize input to escape.Array objects normalized_arrays = [] for i, arr in enumerate(arrays): if isinstance(arr, Array): # Already an escape array normalized_arrays.append(arr) elif isinstance(arr, (da.Array, np.ndarray)): # Raw dask or numpy array - wrap in escape.Array name = names[i] if names and i < len(names) else f"imported_{i:04d}" index = indexes[i] if indexes and i < len(indexes) else np.arange(arr.shape[0]) scan_obj = scans[i] if scans and i < len(scans) else None normalized_arrays.append( Array(data=arr, index=index, parameter=scan_obj, name=name) ) else: raise TypeError(f"Unsupported array type: {type(arr)}") # Set up h5 storage for arrays that don't have it for arr in normalized_arrays: if not hasattr(arr, "h5"): arr.set_h5_storage(parent_h5py, name=arr.name) # Now use the existing store() logic for all normalized escape arrays prep = [ array.h5.append(array.data, array.index, prep_run="store_numpy") for array in normalized_arrays ] if not any(prep): print("Nothing to append") else: arrays_store, ndatas, dsets, n_news = zip( *[(tarray, *tprep) for tarray, tprep in zip(normalized_arrays, prep) if tprep] ) if is_local_client_distributed(): client = get_client() print(f"Found dask distributed client, storing {len(ndatas)} datasets") storage_return = da.store(ndatas, dsets, lock=lock, compute=False, **kwargs) future = client.persist(storage_return) print("persisted, length of future: ", len(future)) progress(future, notebook=False) print("calling wait on future") time.sleep(0.2) t_tmp = time.time() wait(future) print(f"waited for future ({time.time() - t_tmp:.2f}s)") else: with ProgressBar(): da.store(ndatas, dsets, lock=lock, **kwargs) for array, n_new in zip(arrays_store, n_news): array.h5._n_i.append(n_new) array.h5._n_d.append(n_new) # Update all arrays with stored data for array in normalized_arrays: array._data = array.h5.get_data_da() array._index = array.h5.index array.scan._save_to_h5(array.h5.grp) return normalized_arrays def get_lock(): if escape.STORAGE_LOCK: return escape.STORAGE_LOCK else: return True
[docs] def concatenate(arraylist, grid_specs=None): """Concatenate a list of Arrays along the event axis. Merges data, indices, and scan metadata (step lengths and parameter values) from all input Arrays into a single Array. All input Arrays must share the same scan parameter names. Parameters ---------- arraylist : list of escape.Array Arrays to concatenate. They must have identical scan parameter keys. grid_specs : dict, optional Grid metadata to attach to the resulting Array. Returns ------- escape.Array Combined Array whose scan has ``len(arraylist[0].scan) + …`` steps. Examples -------- >>> combined = escape.concatenate([run1, run2, run3]) >>> print(combined.scan.par_steps) """ if all([ta.is_dask_array() for ta in arraylist]): data = da.concatenate([array.data for array in arraylist], axis=0) else: data = np.concatenate([array.data for array in arraylist], axis=0) index = np.concatenate([array.index for array in arraylist]) parameter = {} step_lengths = [] for array in arraylist: if not parameter: parameter.update(deepcopy(array.scan.parameter)) else: if not all(tk in parameter.keys() for tk in array.scan.parameter.keys()): raise Exception( "Scans can not be concatenated due to mismatch in parameters!" ) for par_name, par_dict in array.scan.parameter.items(): parameter[par_name]["values"].extend(list(deepcopy(par_dict["values"]))) if hasattr(par_dict, "attributes") and ( not parameter[par_name]["attributes"] == par_dict["attributes"] ): raise Exception( f"parameter attributes of {par_name} don't fit toghether in concatenated arrays." ) step_lengths.extend(list(array.scan.step_lengths)) return Array( data=data, index=index, parameter=parameter, step_lengths=step_lengths, grid_specs=grid_specs, )
def match_indexes(ids_master, ids_slaves, stepLengths_master=None): ids_res = ids_master for tid in ids_slaves: ids_res = ids_res[np.isin(ids_res, tid, assume_unique=True)] inds_slaves = [] for tid in ids_slaves: srt = tid.argsort(axis=0) inds_slaves.append(srt[np.searchsorted(tid, ids_res, sorter=srt)]) srt = ids_master.argsort(axis=0) inds_master = srt[np.searchsorted(ids_master, ids_res, sorter=srt)] if not stepLengths_master is None: stepLensNew = np.bincount( np.digitize(inds_master, bins=np.cumsum(stepLengths_master)) ) else: stepLensNew = None return inds_master, inds_slaves, stepLensNew def intersect_indexes(ids_all,stepLengths_all): # main format checks of input if not len(ids_all) == len(stepLengths_all): raise Exception("Length of ids_all and stepLengths_all needs to fit!") if not all(len(tid)==sum(tsl) for tid, tsl in zip(ids_all, stepLengths_all)): raise Exception("Length of ids_all entries needs to fit stepLengths_all entries!") ixgr = [] slgr = [] shape = [len(tsl) for tsl in stepLengths_all] gixl = [np.ravel(t) for t in np.indices(shape)] for n in range(len(gixl[0])): sets = [] for ids, stepLengths,i in zip(ids_all,stepLengths_all, gixl): sets.append(set(ids[sum(stepLengths[:i[n]]):sum(stepLengths[:(i[n]+1)])])) tsec = set.intersection(*sets) ixgr += list(tsec) slgr.append(len(tsec)) return np.asarray(ixgr), slgr, shape def get_unique_indexes(index, array_data, stepLengths=None, delete_Ids=[0]): index, idxs = np.unique(index, return_index=True) good_Ids = np.ones_like(idxs, dtype=bool) for bad_Id in delete_Ids: good_Ids[index == bad_Id] = False index = index[good_Ids] idxs = idxs[good_Ids] if stepLengths: stepLengths = np.bincount(np.digitize(idxs, bins=np.cumsum(stepLengths))) return index, array_data[idxs], stepLengths def get_scan_step_selections(ix, stepLengths, scan=None): ix = np.atleast_1d(ix) stepLengths = np.bincount( np.digitize(ix, bins=np.cumsum(stepLengths)), minlength=len(stepLengths) ) validsteps = ~(stepLengths == 0) stepLengths = stepLengths[validsteps] if scan: scan = Scan( parameter=scan.get_parameter_selection(validsteps), step_lengths=stepLengths, grid_specs=scan.grid.get_grid_specs() if hasattr(scan, "grid") else None, ) return stepLengths, scan def escaped_FuncsOnEscArray(array, inst_funcs, *args, **kwargs): # TODO for inst, func in inst_funcs: if isinstance(array.data, inst): return escaped(func, *args, **kwargs)
[docs] def digitize( array, bins, include_outlier_bins=False, sort_groups_by_index=True, right=False, foo=np.digitize, **kwargs, ): """Digitization function for escape arrays according to numpy.digitize. Works for 1D arrays only. Args: array (escape.Array): the escape array holding data that are supposed to be sorted/digitized. bins (array_kile): array of bins, has to be 1-dimensional and monotonic. include_outlier_bins (bool/'right'/'left' optional): option to include outliers of described bin edges on either or both siges of the bins array. Defaults to False. sort_groups_by_index (bool, optional): sorting escape.Array data within bins according to their index value. Defaults to True. right (bool, optional): Indicating whether the intervals include the right or the left bin edge. Default behavior is (right==False) indicating that the interval does not include the right edge. The left bin end is open in this case, i.e., bins[i-1] <= x < bins[i] is the default behavior for monotonically increasing bins. Defaults to False. foo (function, optional): option to modify the digitisation function, needs still to behave closely to np digitize. Defaults to np.digitize. Raises: NotImplementedError: error if no 1d escape.Array is provided as array argument. Returns: escape.Array: Digitized/ resorted escape.Array """ if not np.prod(np.asarray(array.shape)) == array.shape[array.index_dim]: raise NotImplementedError( "Only 1d escape arrays can be digitized in a sensible way." ) darray = array.data.ravel() if include_outlier_bins: direction = np.sign(bins[-1] - bins[0]) if include_outlier_bins == "right": bins = np.concatenate( [ bins, np.atleast_1d(direction * -np.inf), ] ) elif include_outlier_bins == "left": bins = np.concatenate( [ np.atleast_1d(direction * np.inf), bins, ] ) else: bins = np.concatenate( [ np.atleast_1d(direction * np.inf), bins, np.atleast_1d(direction * -np.inf), ] ) if foo is np.digitize: kwargs["right"] = right inds = foo(darray, bins, **kwargs) ix = inds.argsort()[(0 < inds) & (inds < len(bins))] bin_nos, counts = np.unique( inds[(0 < inds) & (inds < len(bins))] - 1, return_counts=True ) if sort_groups_by_index: for n, bin_no in enumerate(bin_nos): tmn = sum(counts[:n]) tmx = sum(counts[: n + 1]) tix = array.index[ix[tmn:tmx]].argsort() ix[tmn:tmx] = ix[tmn:tmx][tix] bin_left = bins[1:] bin_right = bins[:-1] bin_center = (bin_left + bin_right) / 2 parameter = { f"bin_center_{array.name}": { "values": bin_center[bin_nos], "attributes": kwargs, }, f"bin_left_{array.name}": {"values": bin_left[bin_nos], "attributes": kwargs}, f"bin_right_{array.name}": {"values": bin_right[bin_nos], "attributes": kwargs}, } return Array( data=array.data[ix], index=array.index[ix], parameter=parameter, step_lengths=counts, )
def unravel_scans(*arrays, categorize_target=None): """Create a grid-sorter Array spanning the Cartesian product of multiple scan structures. Takes N Arrays that each carry a 1-D scan (one parameter axis each) and produces a single sorter Array whose scan steps correspond to all combinations of those parameter axes — effectively "unravelling" the individual scans into an N-D grid. The resulting sorter can be used with :meth:`~escape.Array.categorize` to apply the grid grouping to any other channel. Parameters ---------- *arrays : escape.Array Two or more Arrays whose scan structures define the grid axes. The number of steps in each Array becomes the size of one grid dimension. categorize_target : escape.Array, optional If provided, immediately :meth:`~escape.Array.categorize` this Array onto the resulting grid and return the categorized result instead of the raw sorter. Returns ------- escape.Array A sorter Array with ``product(n_steps_per_input)`` scan steps and a :class:`~escape.storage.storage.Grid` attached, or — when *categorize_target* is given — the categorized target Array. Examples -------- Build a 10 × 8 grid from two independent 1-D scans and compute per-cell means:: sorter = escape.unravel_scans(scan_a, scan_b) sig_grid = sorter.categorize(sig) means = sig_grid.grid.nanmean() # shape (10, 8) Notes ----- Previously named ``unravel_arrays``; that name is kept as a deprecated alias. """ ixs = [] sls = [] for array in arrays: ixs.append(array.index) sls.append(array.scan.step_lengths) ixgr, slgr, shape = intersect_indexes(ixs, sls) # Create grid_index mappings for the intersection grid grid_indices = [] grid_coords = list(np.ndindex(*shape)) # Generate all N-D coordinates for coord in grid_coords: grid_indices.append({'grid_index': list(coord)}) # Build parameter dict with scan_step_info parameter = { 'scan_step_info': { 'values': grid_indices } } # Create grid_specs grid_specs = { 'shape': shape, 'positions': None, 'grid_dimension_names': None } index_sort_array = Array( index=ixgr, data=ixgr, step_lengths=slgr, parameter=parameter, grid_specs=grid_specs ) if categorize_target: return index_sort_array.categorize(categorize_target) else: return index_sort_array # Backward-compatible alias unravel_arrays = unravel_scans
[docs] def filter(array, *args, foos_filtering=[operator.ge, operator.le], **kwargs): """general filter function for escape arrays. checking for 1D arrays, applies arbitrary number of filter functions that take one argument as input and""" if not np.prod(np.asarray(array.shape)) == array.shape[array.index_dim]: raise NotImplementedError( "Only 1d escape arrays can be filtered in a sensible way." ) darray = array.data if isinstance(darray, da.Array): print("filtering, i.e. downsizing of arrays requires to convert to numpy.") darray = darray.compute() # darray = array.data.ravel() ix = da.logical_and( *[tfoo(darray, targ) for tfoo, targ in zip(foos_filtering, args)] ).nonzero()[0] stepLengths, scan = get_scan_step_selections( ix, array.scan.step_lengths, scan=array.scan ) return Array( data=array.data[ix], index=array.index[ix], step_lengths=stepLengths, parameter=scan.parameter, )
def broadcast_to(ndarray_list, arraydef): if isinstance(arraydef, Array): index = arraydef.index step_lengths = arraydef.scan.step_lengths parameter = arraydef.scan.parameter if not len(ndarray_list) == len(step_lengths): raise Exception( "Cannot broadcast list of arrays that does not fit length of array scan" ) data = [] for ndarray, step_length in zip(ndarray_list, step_lengths): tbc = da.atleast_1d(ndarray) data.append(da.broadcast_to(tbc, [step_length] + list(tbc.shape))) data = da.concatenate(data, axis=0) return Array(data=data, index=index, parameter=parameter, step_lengths=step_lengths) class ArrayH5Dataset: def __init__(self, parent, name): self.parent = parent try: self.grp = parent[name] except: self.grp = parent.require_group(name) if ( "esc_type" in self.grp.attrs.keys() and self.grp.attrs["esc_type"] == "array_dataset" ): pass else: try: self.grp.attrs["esc_type"] = "array_dataset" except: print("Could not put esc_type metadata.") self._data_finder = re.compile("^data_[0-9]{4}$") self._index_finder = re.compile("^index_[0-9]{4}$") self._check_stored_data() def _check_stored_data(self): self._n_d = [] self._n_i = [] for key in self.grp.keys(): if self._data_finder.match(key): self._n_d.append(int(key[-4:])) if self._index_finder.match(key): self._n_i.append(int(key[-4:])) self._n_d.sort() self._n_i.sort() if not self._n_d == self._n_i: raise Exception( "Corrupt escape ArrayH5Dataset, not equal numbered data and id sub-datasets!" ) def clear_stored_data(self): for key in self.grp.keys(): try: del self.grp[key] except: print(f"Did not succeed to delete key {key}!") @property def index(self): if self._n_i: return np.concatenate( [np.asarray(self.grp[f"index_{n:04d}"][:]) for n in self._n_i], axis=0 ) else: return np.asarray([], dtype=int) def append(self, data, event_ids, scan=None, prep_run=False, lock="auto", **kwargs): """ expects to extend a former dataset, i.e. data includes data already existing, this will likely change in future to also allow real appending of entirely new data. """ if lock == "auto": lock = get_lock() n_new = len(self._n_i) ids_stored = self.index in_previous_indexes = np.isin(event_ids, ids_stored) if ~in_previous_indexes.any(): # real appending data new_event_ids = event_ids new_data = data elif in_previous_indexes.all(): # real extending of data if len(event_ids) < len(ids_stored): raise Exception("fewer event_ids to append than already stored!") if not (event_ids[: len(ids_stored)] == ids_stored).all(): raise Exception("new event_ids don't extend existing ones!") if len(event_ids) == len(ids_stored): print("Nothing new to append.") return new_event_ids = event_ids[len(ids_stored) :] new_data = data[len(ids_stored) :, ...] self.grp[f"index_{n_new:04d}"] = new_event_ids if isinstance(data, np.ndarray): if prep_run: if prep_run == "store_numpy": pass else: raise Exception( "Trying dry_run on numpy array data on {self.grp.name}." ) self.grp[f"data_{n_new:04d}"] = new_data elif isinstance(data, da.Array): # ToDo, smarter chunking when writing small data new_chunks = tuple(c[0] for c in new_data.chunks) try: if "default_dataset_compression" in self.grp.file.attrs: compression = self.grp.file.attrs["default_dataset_compression"] else: compression = None if "default_dataset_compression_opts" in self.grp.file.attrs: compression_opts = self.grp.file.attrs[ "default_dataset_compression_opts" ] else: compression_opts = None dset = self.grp.create_dataset( f"data_{n_new:04d}", shape=new_data.shape, chunks=new_chunks, dtype=new_data.dtype, compression=compression, compression_opts=compression_opts, ) except: compression = None compression_opts = None dset = self.grp.create_dataset( f"data_{n_new:04d}", shape=new_data.shape, chunks=new_chunks, dtype=new_data.dtype, ) if prep_run: return new_data, dset, n_new da.store(new_data, dset, lock=lock, **kwargs) if scan: scan._save_to_h5(self.grp) self._n_i.append(n_new) self._n_d.append(n_new) def get_data_da(self, memlimit_MB=50): allarrays = [] for n in self._n_i: ds = self.grp[f"data_{n:04d}"] if ds.chunks: chunk_size = list(ds.chunks) else: chunk_size = list(ds.shape) if chunk_size[0] == 1: size_element = ( np.dtype(ds.dtype).itemsize * np.prod(ds.shape[1:]) / 1024**2 ) chunk_size[0] = int(memlimit_MB // size_element) allarrays.append(da.from_array(ds, chunks=chunk_size)) # if len(allarrays) < 1: # print(ds, ds.shape, chunk_size) if len(allarrays) == 0: return None else: return da.concatenate(allarrays) def create_array(self): return Array(data=self.get_data_da(), index=self.index) class ArrayH5File: def __init__(self, file_name, parent_group_name="", name=None): self.file_name = Path(file_name) self.parent_group_name = parent_group_name self.name = name self.require_group() self.update_dataset_status() @property def group_name(self): return (Path(self.parent_group_name) / Path(self.name)).as_posix() def require_group(self): with h5py.File(self.file_name, "a") as f: f[self.parent_group_name].require_group(self.name) def update_dataset_status(self): # self.grp = parent.require_group(name) self._data_finder = re.compile("^data_[0-9]{4}$") self._index_finder = re.compile("^index_[0-9]{4}$") self._n_d = [] self._n_i = [] with h5py.File(self.file_name, "r") as f: keys = f[self.group_name].keys() for key in keys: if self._data_finder.match(key): self._n_d.append(int(key[-4:])) if self._index_finder.match(key): self._n_i.append(int(key[-4:])) self._n_d.sort() self._n_i.sort() if not self._n_d == self._n_i: raise Exception( "Corrupt escape ArrayH5Dataset, not equally sized data and index sub-datasets!" ) @property def index(self): if self._n_i: with h5py.File(self.file_name, "r") as f: return np.concatenate( [ np.asarray(f[self.group_name][f"index_{n:04d}"][:]) for n in self._n_i ], axis=0, ) else: return np.asarray([], dtype=int) def append(self, data, event_ids, scan=None, prep_run=False): """ expects to extend a former dataset, i.e. data includes data already existing, this will likely change in future to also allow real appending of entirely new data. """ n_new = len(self._n_i) ids_stored = self.index in_previous_indexes = np.isinb(event_ids, ids_stored) if ~in_previous_indexes.any(): # real appending data new_event_ids = event_ids new_data = data elif in_previous_indexes.all(): # real extending of data if len(event_ids) < len(ids_stored): raise Exception("fewer event_ids to append than already stored!") if not (event_ids[: len(ids_stored)] == ids_stored).all(): raise Exception("new event_ids don't extend existing ones!") if len(event_ids) == len(ids_stored): print("Nothing new to append.") return new_event_ids = event_ids[len(ids_stored) :] new_data = data[len(ids_stored) :, ...] with h5py.File(self.file_name, "a") as f: f[self.group_name + f"/index_{n_new:04d}"] = new_event_ids if isinstance(data, np.ndarray): if prep_run: if prep_run == "store_numpy": pass else: raise Exception( "Trying dry_run on numpy array data on {self.grp.name}." ) with h5py.File(self.file_name, "a") as f: f[self.group_name + f"/data_{n_new:04d}"] = new_data elif isinstance(data, da.Array): # ToDo, smarter chunking when writing small data new_chunks = tuple(c[0] for c in new_data.chunks) location = { "file_name": self.file_name, "dataset_name": self.group_name + f"/data_{n_new:04d}", } if prep_run: return new_data, location, n_new else: data.to_hdf5(self.file_name, location["dataset_name"]) if scan: with h5py.File(self.file_name, "a") as f: scan._save_to_h5(f[self.group_name]) self._n_i.append(n_new) self._n_d.append(n_new) def get_data_da(self, memlimit_MB=50): allarrays = [] for n in self._n_i: ds_name = self.group_name + f"/data_{n:04d}" h5store = self.analyze_h5_dataset(ds_name, memlimit_MB=memlimit_MB) tda = self.h5store_to_da(h5store=h5store) allarrays.append(tda) return da.concatenate(allarrays) def create_array(self): return Array(data=self.get_data_da(), index=self.index) # @dask.delayed def analyze_h5_dataset(self, dataset_path, memlimit_MB=100): """Data parser assuming the standard swissfel h5 format for raw data""" with h5py.File(self.file_name, mode="r") as fh: ds_data = fh[dataset_path] if memlimit_MB: dtype = np.dtype(ds_data.dtype) size_element = ( np.dtype(ds_data.dtype).itemsize * np.prod(ds_data.shape[1:]) / 1024**2 ) chunk_length = int(memlimit_MB // size_element) dset_size = ds_data.shape chunk_shapes = [] slices = [] for chunk_start in range(0, dset_size[0], chunk_length): slice_0dim = [ chunk_start, min(chunk_start + chunk_length, dset_size[0]), ] chunk_shape = list(dset_size) chunk_shape[0] = slice_0dim[1] - slice_0dim[0] slices.append(slice_0dim) chunk_shapes.append(chunk_shape) h5store = { "file_path": self.file_name, "dataset_name": ds_data.name, "dataset_shape": ds_data.shape, "dataset_dtype": dtype, "dataset_chunks": {"slices": slices, "shapes": chunk_shapes}, } return h5store @dask.delayed def read_h5_chunk(self, ds_path, slice_args): with h5py.File(self.file_name, "r") as fh: dat = fh[ds_path][slice(*slice_args)] return dat def h5store_to_da(self, h5store): arrays = [ dask.array.from_delayed( self.read_h5_chunk(h5store["dataset_name"], tslice), tshape, dtype=h5store["dataset_dtype"], ) for tslice, tshape in zip( h5store["dataset_chunks"]["slices"], h5store["dataset_chunks"]["shapes"] ) ] data = dask.array.concatenate(arrays, axis=0) return data # if 'data' in grp.keys(): # print(f'Dataset {name} already exists, data:') # print(str(grp['data'])) # if input('Would you like to delete and overwrite the data ? (y/n)')=='y': # del grp['data'] # del grp['event_ids'] # else: # return # grp['event_ids'] = self.index # if isinstance(self.data,np.array): # grp['data'] = self.data # elif isinstance(self.data,da.array): # dset = grp.create_dataset('data',shape=self.data.shape,chunks=self.data.chunks,dtype=self.data.dtype) # self.data.store(dset)