Skip to content

Commit bd14895

Browse files
committed
adds catch that ensures pyramid chunks match encoding chunks
1 parent c0ae882 commit bd14895

2 files changed

Lines changed: 44 additions & 5 deletions

File tree

src/topozarr/coarsen.py

Lines changed: 20 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -57,19 +57,34 @@ def create_pyramid(
5757
full_encoding = {}
5858

5959
for idx, ds_level in level_datasets.items():
60-
if "spatial_ref" in ds_level.coords:
61-
ds_level = ds_level.drop_vars("spatial_ref")
62-
6360
name = str(idx)
6461
path = f"/{idx}"
65-
dt[path] = DataTree(ds_level, name=name)
66-
full_encoding[path] = create_level_encoding(
62+
63+
level_encoding = create_level_encoding(
6764
ds_level,
6865
x_dim,
6966
y_dim,
7067
target_chunk_bytes=target_chunk_bytes,
7168
target_shard_bytes=target_shard_bytes,
7269
)
7370

71+
dim_chunks = {}
72+
for var_name, var_enc in level_encoding.items():
73+
if var_name in ds_level.data_vars and "chunks" in var_enc:
74+
target_chunks = var_enc["chunks"]
75+
da = ds_level[var_name]
76+
77+
for dim, chunk_size in zip(da.dims, target_chunks):
78+
if dim not in dim_chunks:
79+
dim_chunks[dim] = chunk_size
80+
else:
81+
dim_chunks[dim] = min(dim_chunks[dim], chunk_size)
82+
83+
if dim_chunks:
84+
ds_level = ds_level.chunk(dim_chunks)
85+
86+
dt[path] = DataTree(ds_level, name=name)
87+
full_encoding[path] = level_encoding
88+
7489
dt.attrs = create_multiscale_metadata(levels, crs_str, method, spec=spec)
7590
return Pyramid(datatree=dt, encoding=full_encoding)

tests/test_chunking.py

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -50,3 +50,27 @@ def test_shard_size_overrides(create_dataset):
5050

5151
enc = pyramid.encoding["/0"]["elevation"]
5252
assert enc["chunks"] == enc["shards"]
53+
54+
55+
def test_dask_chunks_match_encoding(create_dataset):
56+
from topozarr.coarsen import create_pyramid
57+
58+
ds = create_dataset(nx=1000, ny=1000)
59+
pyramid = create_pyramid(ds, levels=2, target_chunk_bytes=1024)
60+
61+
for level_path, level_encoding in pyramid.encoding.items():
62+
ds_level = pyramid.datatree[level_path].ds
63+
64+
for var_name, var_enc in level_encoding.items():
65+
if var_name not in ds_level.data_vars:
66+
continue
67+
68+
da = ds_level[var_name]
69+
expected_chunks = var_enc["chunks"]
70+
71+
if hasattr(da.data, "chunksize"):
72+
actual_chunks = da.data.chunksize
73+
assert actual_chunks == expected_chunks, (
74+
f"lvl {level_path}, var {var_name}: "
75+
f"dask chunks {actual_chunks} do not match the encoding specified chunks {expected_chunks}"
76+
)

0 commit comments

Comments
 (0)