Skip to content

Commit da09179

Browse files
committed
add code to work with both zarr 2 and 3, add tests
1 parent 6b270c9 commit da09179

5 files changed

Lines changed: 28617 additions & 47 deletions

File tree

.gitignore

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -8,10 +8,10 @@
88
*.orig
99
*.log
1010
*.pot
11-
__pycache__/*
11+
__pycache__/
1212
.cache/*
1313
.*.swp
14-
*/.ipynb_checkpoints/*
14+
.ipynb_checkpoints/
1515
.DS_Store
1616

1717
# Project files

src/scimilarity/zarr_dataset.py

Lines changed: 38 additions & 45 deletions
Original file line numberDiff line numberDiff line change
@@ -1,11 +1,10 @@
1+
from packaging.version import Version
12
from scipy.sparse import csr_matrix, csc_matrix, coo_matrix
2-
from typing import Dict, Optional, Tuple, Union, TYPE_CHECKING, Any
3+
from typing import Any, Dict, Optional, Tuple, Union
34

4-
if TYPE_CHECKING:
5-
import numpy
6-
import pandas
7-
import zarr
5+
import zarr
86

7+
ZARR_V3 = Version(zarr.__version__) >= Version("3.0.0")
98

109
ARRAY_FORMATS = {
1110
"csr_matrix": csr_matrix,
@@ -31,12 +30,14 @@ class ZarrDataset:
3130
"""
3231

3332
def __init__(self, store_path: str, mode: str = "r"):
34-
import zarr
35-
36-
self.store_path = zarr.DirectoryStore(store_path)
37-
self.root = zarr.open_group(
38-
self.store_path, mode=mode, chunk_store=self.store_path
39-
)
33+
if ZARR_V3:
34+
self.store_path = store_path
35+
self.root = zarr.open_group(store_path, mode=mode)
36+
else:
37+
self.store_path = zarr.DirectoryStore(store_path)
38+
self.root = zarr.open_group(
39+
self.store_path, mode=mode, chunk_store=self.store_path
40+
)
4041

4142
@property
4243
def dataset_info(self) -> Dict[str, list]:
@@ -732,18 +733,18 @@ def set_matrix(
732733
group.attrs.setdefault("encoding-version", "0.1.0")
733734
group.attrs.setdefault("shape", list(matrix.shape))
734735

736+
if ZARR_V3:
737+
create = lambda name, data, **_: group.create_array(name, data=data)
738+
else:
739+
create = lambda name, data, **kw: group.create_dataset(name, data=data, **kw)
735740
if encoding_type in ["csr_matrix", "csc_matrix"]:
736-
group.create_dataset("data", data=matrix.data, dtype=matrix.data.dtype)
737-
group.create_dataset(
738-
"indptr", data=matrix.indptr, dtype=matrix.indptr.dtype
739-
)
740-
group.create_dataset(
741-
"indices", data=matrix.indices, dtype=matrix.indices.dtype
742-
)
741+
create("data", data=matrix.data, dtype=matrix.data.dtype)
742+
create("indptr", data=matrix.indptr, dtype=matrix.indptr.dtype)
743+
create("indices", data=matrix.indices, dtype=matrix.indices.dtype)
743744
elif encoding_type in ["coo_matrix"]:
744-
group.create_dataset("data", data=matrix.data, dtype=matrix.data.dtype)
745-
group.create_dataset("row", data=matrix.row, dtype=matrix.row.dtype)
746-
group.create_dataset("col", data=matrix.col, dtype=matrix.col.dtype)
745+
create("data", data=matrix.data, dtype=matrix.data.dtype)
746+
create("row", data=matrix.row, dtype=matrix.row.dtype)
747+
create("col", data=matrix.col, dtype=matrix.col.dtype)
747748

748749
def append_matrix(
749750
self,
@@ -903,7 +904,6 @@ def get_annotation_column(
903904
"""
904905

905906
import pandas as pd
906-
import zarr
907907

908908
if column in group:
909909
series = group[column]
@@ -955,12 +955,18 @@ def set_annotation(self, annotation: str, df: "pandas.DataFrame"):
955955
anno.attrs.setdefault("encoding-type", "dataframe")
956956
anno.attrs.setdefault("encoding-version", "0.2.0")
957957

958-
anno.create_dataset(
959-
"_index",
960-
data=df.index._values,
961-
dtype=df.index._values.dtype,
962-
object_codec=numcodecs.JSON(),
963-
)
958+
def create_array(group, name, data, is_string=False):
959+
if ZARR_V3:
960+
if is_string:
961+
data = data.astype(str) if data.dtype == object else data
962+
group.create_array(name, data=data)
963+
else:
964+
kwargs = {"data": data, "dtype": data.dtype}
965+
if is_string:
966+
kwargs["object_codec"] = numcodecs.JSON()
967+
group.create_dataset(name, **kwargs)
968+
969+
create_array(anno, "_index", df.index._values, is_string=True)
964970
anno["_index"].attrs.setdefault("encoding-type", "string-array")
965971
anno["_index"].attrs.setdefault("encoding-version", "0.2.0")
966972
for k in df.columns:
@@ -971,30 +977,17 @@ def set_annotation(self, annotation: str, df: "pandas.DataFrame"):
971977
anno[k].attrs.setdefault("encoding-version", "0.2.0")
972978
anno[k].attrs.setdefault("ordered", False)
973979

974-
anno[k].create_dataset(
975-
"categories",
976-
data=v.categories._values,
977-
dtype=v.categories._values.dtype,
978-
object_codec=numcodecs.JSON(),
980+
create_array(
981+
anno[k], "categories", v.categories._values, is_string=True
979982
)
980983
anno[k]["categories"].attrs.setdefault("encoding-type", "string-array")
981984
anno[k]["categories"].attrs.setdefault("encoding-version", "0.2.0")
982985

983-
anno[k].create_dataset("codes", data=v.codes)
986+
create_array(anno[k], "codes", v.codes)
984987
anno[k]["codes"].attrs.setdefault("encoding-type", "array")
985988
anno[k]["codes"].attrs.setdefault("encoding-version", "0.2.0")
986989
elif isinstance(df[k], pd.Series):
987-
if df[k].dtype == "O":
988-
anno.create_dataset(
989-
k,
990-
data=df[k]._values,
991-
dtype=df[k]._values.dtype,
992-
object_codec=numcodecs.JSON(),
993-
)
994-
else:
995-
anno.create_dataset(
996-
k, data=df[k]._values, dtype=df[k]._values.dtype
997-
)
990+
create_array(anno, k, df[k]._values, is_string=(df[k].dtype == "O"))
998991
anno[k].attrs.setdefault("encoding-type", "array")
999992
anno[k].attrs.setdefault("encoding-version", "0.2.0")
1000993

0 commit comments

Comments
 (0)