Skip to content

Commit e65398d

Browse files
authored
Fix stale flow-direction caches (#128)
* fix stream_order cache * other cache fixes
1 parent ab8f002 commit e65398d

3 files changed

Lines changed: 132 additions & 6 deletions

File tree

‎pyflwdir/flwdir.py‎

Lines changed: 6 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -291,7 +291,7 @@ def main_upstream(self, uparea: np.ndarray | None = None) -> np.ndarray:
291291
idxs_us_main = core.main_upstream(
292292
idxs_ds=self.idxs_ds, uparea=self._check_data(uparea, "uparea"), mv=self._mv
293293
)
294-
if self.cache:
294+
if self.cache and uparea is None:
295295
self._cached.update(idxs_us_main=idxs_us_main)
296296
return idxs_us_main
297297

@@ -313,10 +313,11 @@ def add_pits(
313313
# add pits
314314
self.idxs_ds[idxs1] = idxs1
315315
self._pit = np.unique(np.concatenate([self.idxs_pit, idxs1]))
316-
# reset order, nnodes and upstream cell indices
316+
# Reset traversal state and all values derived from the flow topology.
317317
self._seq = None
318318
self._nnodes = None
319-
self._idxs_us_main = None
319+
for key in ("rank", "strord", "idxs_us_main", "distnc"):
320+
self._cached.pop(key, None)
320321

321322
def repair_loops(self) -> None:
322323
"""Repair loops by setting a pit at every cell which does not drain to a pit."""
@@ -593,11 +594,11 @@ def stream_order(
593594
"""
594595
mask = self._check_data(mask, "mask", optional=True)
595596
if type.lower() == "strahler":
596-
if "strord" in self._cached:
597+
if mask is None and "strord" in self._cached:
597598
strord = self._cached["strord"]
598599
else:
599600
strord = streams.strahler_order(self.idxs_ds, self.idxs_seq, mask=mask)
600-
if self.cache:
601+
if self.cache and mask is None:
601602
self._cached.update(strord=strord)
602603
elif type.lower() == "classic":
603604
strord = streams.stream_order(

‎pyflwdir/pyflwdir.py‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -343,6 +343,8 @@ def set_transform(self, transform: Affine, latlon: bool = False) -> None:
343343
raise ValueError("Invalid transform.")
344344
self.transform = transform
345345
self.latlon = latlon
346+
for key in ("area", "distnc", "idxs_us_main"):
347+
self._cached.pop(key, None)
346348

347349
### WRITE / EXPORT ###
348350

‎tests/test_pyflwdir.py‎

Lines changed: 124 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@
99
from affine import Affine
1010

1111
import pyflwdir
12-
from pyflwdir import core
12+
from pyflwdir import core, streams
1313
from pyflwdir.pyflwdir import FlwdirRaster, _get_idxs_dtype
1414

1515
pyflwdir_module = importlib.import_module("pyflwdir.pyflwdir")
@@ -381,6 +381,129 @@ def test_streams(flw_real, flwdir_real_rank):
381381
assert np.all(data_smooth1 == data)
382382

383383

384+
@pytest.mark.integration
385+
@pytest.mark.parametrize("raster", [False, True])
386+
def test_stream_order_mask_cache(flw_real, raster):
387+
if raster:
388+
flw = FlwdirRaster(
389+
flw_real.idxs_ds.copy(),
390+
flw_real.shape,
391+
"d8",
392+
idxs_pit=flw_real.idxs_pit.copy(),
393+
cache=True,
394+
)
395+
else:
396+
flw = pyflwdir.Flwdir(
397+
flw_real.idxs_ds.copy(),
398+
idxs_pit=flw_real.idxs_pit.copy(),
399+
cache=True,
400+
)
401+
402+
all_streams = flw.mask.reshape(flw.shape)
403+
pit_streams = np.zeros(flw.shape, dtype=bool)
404+
pit_streams.flat[flw.idxs_pit] = True
405+
406+
all_order = flw.stream_order(mask=all_streams)
407+
pit_order = flw.stream_order(mask=pit_streams)
408+
expected = streams.strahler_order(
409+
flw.idxs_ds, flw.idxs_seq, mask=pit_streams.ravel()
410+
).reshape(flw.shape)
411+
412+
assert not np.array_equal(all_order, pit_order)
413+
assert np.array_equal(pit_order, expected)
414+
assert np.all(pit_order[~pit_streams] == 0)
415+
416+
417+
@pytest.mark.unit
418+
@pytest.mark.parametrize("raster", [False, True])
419+
def test_add_pits_invalidates_cache(raster):
420+
idxs_ds = np.array([0, 0, 1, 1, core._mv], dtype=np.int32)
421+
idxs_pit = np.array([0], dtype=np.int32)
422+
if raster:
423+
flw = FlwdirRaster(
424+
idxs_ds.copy(), (1, 5), "d8", idxs_pit=idxs_pit.copy(), cache=True
425+
)
426+
else:
427+
flw = pyflwdir.Flwdir(idxs_ds.copy(), idxs_pit=idxs_pit.copy(), cache=True)
428+
429+
old_rank = flw.rank.copy()
430+
old_strord = flw.stream_order().copy()
431+
old_idxs_us_main = flw.idxs_us_main.copy()
432+
if raster:
433+
old_distnc = flw.distnc.copy()
434+
435+
flw.add_pits(idxs=np.array([2]))
436+
437+
assert "rank" not in flw._cached
438+
assert "strord" not in flw._cached
439+
assert "idxs_us_main" not in flw._cached
440+
assert np.array_equal(
441+
flw.rank, core.rank(flw.idxs_ds, mv=flw._mv)[0].reshape(flw.shape)
442+
)
443+
assert not np.array_equal(flw.rank, old_rank)
444+
assert not np.array_equal(flw.stream_order(), old_strord)
445+
assert flw.idxs_us_main[1] == 3
446+
assert not np.array_equal(flw.idxs_us_main, old_idxs_us_main)
447+
if raster:
448+
assert "distnc" not in flw._cached
449+
assert not np.array_equal(flw.distnc, old_distnc)
450+
451+
452+
@pytest.mark.unit
453+
def test_repair_loops_invalidates_topology_cache():
454+
idxs_ds = np.array([0, 2, 1], dtype=np.int32)
455+
flw = pyflwdir.Flwdir(idxs_ds, idxs_pit=np.array([0], dtype=np.int32), cache=True)
456+
old_rank = flw.rank.copy()
457+
old_strord = flw.stream_order().copy()
458+
old_idxs_us_main = flw.idxs_us_main.copy()
459+
460+
flw.repair_loops()
461+
462+
assert flw.isvalid
463+
assert "strord" not in flw._cached
464+
assert "idxs_us_main" not in flw._cached
465+
assert np.array_equal(flw.rank, core.rank(flw.idxs_ds, mv=flw._mv)[0])
466+
assert not np.array_equal(flw.rank, old_rank)
467+
assert not np.array_equal(flw.stream_order(), old_strord)
468+
assert not np.array_equal(flw.idxs_us_main, old_idxs_us_main)
469+
470+
471+
@pytest.mark.unit
472+
def test_set_transform_invalidates_geometry_cache():
473+
idxs_ds = np.array([0, 0, 1, 1, core._mv], dtype=np.int32)
474+
flw = FlwdirRaster(
475+
idxs_ds, (1, 5), "d8", idxs_pit=np.array([0], dtype=np.int32), cache=True
476+
)
477+
old_area = flw.area.copy()
478+
old_distnc = flw.distnc.copy()
479+
old_idxs_us_main = flw.idxs_us_main.copy()
480+
481+
flw.set_transform(Affine.scale(2), latlon=False)
482+
483+
assert "area" not in flw._cached
484+
assert "distnc" not in flw._cached
485+
assert "idxs_us_main" not in flw._cached
486+
mask = flw.mask.reshape(flw.shape)
487+
assert np.all(flw.area[mask] == 4)
488+
assert not np.array_equal(flw.area, old_area)
489+
assert np.all(flw.distnc[mask] == 2 * old_distnc[mask])
490+
assert np.array_equal(flw.idxs_us_main, old_idxs_us_main)
491+
492+
493+
@pytest.mark.unit
494+
def test_main_upstream_custom_area_does_not_replace_default_cache():
495+
idxs_ds = np.array([0, 0, 1, 1, core._mv], dtype=np.int32)
496+
flw = pyflwdir.Flwdir(idxs_ds, idxs_pit=np.array([0], dtype=np.int32), cache=True)
497+
default_idxs_us_main = flw.idxs_us_main.copy()
498+
custom_uparea = np.array([1, 1, 1, 2, 0], dtype=np.float32)
499+
500+
custom_idxs_us_main = flw.main_upstream(uparea=custom_uparea)
501+
502+
assert custom_idxs_us_main[1] == 3
503+
assert default_idxs_us_main[1] == 2
504+
assert np.array_equal(flw.idxs_us_main, default_idxs_us_main)
505+
506+
384507
@pytest.mark.integration
385508
def test_upscale(flw_real, nextxy_real):
386509
flw1, idxs_out = flw_real.upscale(5, method="dmm") # single method

0 commit comments

Comments
 (0)