Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -515,6 +515,7 @@ jobs:
- check
- spelling
- fmt
- citation
- build # includes the dist artifact
- test
- typing
Expand Down
45 changes: 30 additions & 15 deletions packages/xarray-safeguards/src/xarray_safeguards/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -115,6 +115,7 @@
ctx,
)
from compression_safeguards.utils.typing import JSON, S, T
from xarray.namedarray.parallelcompat import get_chunked_array_type

DataValue: TypeAlias = int | float | np.number | xr.DataArray
"""
Expand Down Expand Up @@ -464,8 +465,9 @@ def _compute_independent_chunk_correction(

da_correction = (
data.copy(
data=data.data.map_blocks(
data=dask.array.map_blocks(
_compute_independent_chunk_correction,
data.data,
approximation.data,
*chunked_late_bound.values(),
dtype=correction_dtype,
Expand Down Expand Up @@ -890,23 +892,36 @@ def apply_data_array_correction(
| ctx
)

if approximation.chunks is None:
return approximation.copy(
data=safeguards.apply_correction(approximation.data, correction.data)
).assign_attrs(safeguards=correction.attrs["safeguards"])

with ctx.parameter("correction"):
chunkmanager = get_chunked_array_type(approximation.data, correction.data)

def _apply_independent_chunk_correction(
approximation_chunk: xr.DataArray,
correction_chunk: xr.DataArray,
approximation_chunk: np.ndarray[S, np.dtype[T]],
correction_chunk: np.ndarray[S, np.dtype[np.unsignedinteger]],
safeguards: Safeguards,
) -> xr.DataArray:
return approximation_chunk.copy(
data=safeguards.apply_correction(
approximation_chunk.values, correction_chunk.values
)
)
) -> np.ndarray[S, np.dtype[T]]:
# ensure that we pass np.ndarray's to the compression-safeguards
approximation_chunk = _ensure_array(approximation_chunk)
correction_chunk = _ensure_array(correction_chunk)

return safeguards.apply_correction(approximation_chunk, correction_chunk)

return xr.map_blocks(
_apply_independent_chunk_correction,
approximation,
args=(correction,),
kwargs=dict(safeguards=safeguards),
template=approximation,
return approximation.copy(
data=chunkmanager.map_blocks(
_apply_independent_chunk_correction,
approximation.data,
correction.data,
dtype=approximation.dtype,
chunks=None,
drop_axis=None,
new_axis=None,
safeguards=safeguards,
)
).assign_attrs(safeguards=correction.attrs["safeguards"])


Expand Down
6 changes: 3 additions & 3 deletions src/compression_safeguards/safeguards/_qois/context.py
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
from typing import Generic, Protocol
from typing import Generic, Literal, Protocol

import numpy as np

Expand Down Expand Up @@ -326,7 +326,7 @@ def on_complete_term(
Xs_upper: np_sndarray[Ps, Ns, np.dtype[F]],
*,
term: int,
where: None | np_sndarray[Ps, Ns, np.dtype[np.bool]] = None,
where: Literal[True] | np_sndarray[Ps, Ns, np.dtype[np.bool]] = True,
) -> None:
"""
Callback that can be passed as the `callback` parameter in
Expand All @@ -347,7 +347,7 @@ def on_complete_term(
`Xs`, for the `term`.
term : int
The index of the term for which the data bounds have been computed.
where : None | np_sndarray[Ps, Ns, np.dtype[np.bool]]
where : Literal[True] | np_sndarray[Ps, Ns, np.dtype[np.bool]]
Optional mask to only integrate the data bounds for the `term` for
some elements.
"""
Expand Down
118 changes: 79 additions & 39 deletions src/compression_safeguards/safeguards/stencil/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
import numpy as np
from typing_extensions import override # MSPV 3.12

from ...utils._compat import _sliding_window_view
from ...utils._compat import _reshape, _sliding_window_view
from ...utils.bindings import Parameter
from ...utils.error import TypeCheckError, ctx, lookup_enum_or_raise
from ...utils.typing import JSON, TB, S
Expand Down Expand Up @@ -394,48 +394,88 @@ def _reverse_neighbourhood_indices(
None if axis.constant_boundary is None else np.full((), data_size),
axis.axis,
)
indices_windows = _sliding_window_view(
indices_boundary,
window,
axis=tuple(axis.axis for axis in neighbourhood),
writeable=False,
).reshape((-1, window_size))

indices_windows: np.ndarray[tuple[int, int] | tuple[int], np.dtype[np.int_]]
indices_windows = _reshape(
_sliding_window_view(
indices_boundary,
window,
axis=tuple(axis.axis for axis in neighbourhood),
writeable=False,
),
(-1, window_size),
)

# track the indices of the window indices
indices_windows_indices = np.arange(indices_windows.size).reshape(
indices_windows.shape
)

fill_value = indices_windows.size

# skip back-contributions from data elements where the safety requirements
# are disabled
if where_flat is not True:
indices_windows = indices_windows[where_flat]
indices_windows_indices = indices_windows_indices[where_flat]

# skip window indices that are not used
indices_windows = indices_windows[:, window_used.flatten()]
indices_windows_indices = indices_windows_indices[:, window_used.flatten()]

indices_windows = indices_windows.flatten()
indices_windows_indices = indices_windows_indices.flatten()

# sort the indices, such that windows that read the same data are together
# use a stable sort to ensure consistent results, independent of chunking
argindices = np.argsort(indices_windows, stable=True)
indices_windows_sorted = indices_windows[argindices]

# indices_windows might include fill values, of value data_size, which
# represent constant values that come from no data index
# exclude those, conveniently largest values, from indices_windows_sorted
# to ensure that we only track valid data indices
only_fill_index = np.searchsorted(indices_windows_sorted, data_size)
argindices = argindices[:only_fill_index]
indices_windows_sorted = indices_windows_sorted[:only_fill_index]

# find the starts of the runs of common indices
indices_run_starts = np.r_[
0, np.flatnonzero(indices_windows_sorted[1:] != indices_windows_sorted[:-1]) + 1
]

# find the inverse mapping from sorted indices to their unique indices
_, indices_windows_sorted_inverse = np.unique(
indices_windows_sorted, return_inverse=True, sorted=True
)

# find the offsets inside each index run, e.g. for a sequence
# [a, b, b, b, c, d, d],
# the offsets will be
# [0, 0, 1, 2, 0, 0, 1]
indices_run_offsets = (
np.arange(indices_windows_sorted.size)
- indices_run_starts[indices_windows_sorted_inverse]
)

indices_max_run_length = np.amax(indices_run_offsets, initial=-1) + 1

# compute the reverse: for each data element, which windows is it in
# i.e. for each data element, which derived elements does it contribute to
# and thus which data bounds affect it
reverse_indices_windows = np.full(
(data_size, np.sum(window_used.astype(int))), indices_windows.size
reverse_indices_windows = np.full(data_size * indices_max_run_length, fill_value)
# store the reverse mapping
# - this is complicated since each data element may be referenced by
# multiple windows, and we need to ensure that they don't override
# each other's contributions when run with vectorisation
# - so we precompute a unique run-slot for each back-reference
# - since we sorted the indices earlier to find the runs, we also need
# to apply the same reordering to the back-references
reverse_indices_windows[
indices_windows_sorted * indices_max_run_length + indices_run_offsets
] = indices_windows_indices[argindices]
reverse_indices_windows = reverse_indices_windows.reshape(
data_size, indices_max_run_length
)
reverse_indices_counter = np.zeros(data_size, dtype=np.intp)
for i, u in enumerate(window_used.flat):
# skip window indices that are not used
if not u:
continue
# manual loop to account for potential aliasing:
# with a wrapping boundary, more than one j for the same window
# position j could refer back to the same data element
for j in range(indices_windows.shape[0]):
# skip back-contributions from data elements where the safety
# requirements are disabled
if (where_flat is not True) and (not where_flat[j]):
continue
idx = indices_windows[j, i]
if idx != data_size:
# lazily allocate more to account for all possible edge cases
if reverse_indices_counter[idx] >= reverse_indices_windows.shape[1]:
new_reverse_indices_windows = np.full(
(data_size, reverse_indices_windows.shape[1] * 2),
indices_windows.size,
)
new_reverse_indices_windows[
:, : reverse_indices_windows.shape[1]
] = reverse_indices_windows
reverse_indices_windows = new_reverse_indices_windows
# update the reverse mapping
reverse_indices_windows[idx][reverse_indices_counter[idx]] = (
j * window_used.size
) + i
reverse_indices_counter[idx] += 1

return reverse_indices_windows
1 change: 1 addition & 0 deletions src/compression_safeguards/safeguards/stencil/qoi/eb.py
Original file line number Diff line number Diff line change
Expand Up @@ -1189,6 +1189,7 @@ def compute_safe_intervals(
# since some data elements may have no data bounds that affect them,
# e.g. because of the valid boundary condition, they may have infinite
# bounds
# FIXME: does this need to be zero-sign sensitive?
data_float_lower: np.ndarray[S, np.dtype[np.floating]] = _reshape(
np.amax(data_windows_float_lower_flat[reverse_indices_windows], axis=1),
data.shape,
Expand Down
Loading
Loading