1+ from packaging .version import Version
12from 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
109ARRAY_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