@@ -71,6 +71,39 @@ def _read_dataset_slice(
7171 return data
7272
7373
74+ git def _decode_string_array (values : np .ndarray ) -> np .ndarray :
75+ if values .dtype .kind in {"S" , "O" }:
76+ decoded = np .empty (len (values ), dtype = object )
77+ for idx , value in enumerate (values ):
78+ decoded [idx ] = _decode_scalar (value )
79+ return decoded
80+ if values .dtype .kind == "U" :
81+ return values .astype (object )
82+ return values
83+
84+
85+ def _read_nullable_group_slice (
86+ field_node : h5py .Group ,
87+ field_name : str ,
88+ start_idx : int ,
89+ end_idx : int
90+ ) -> np .ndarray :
91+ if "values" not in field_node or "mask" not in field_node :
92+ msg = f"Nullable obs field '{ field_name } ' is missing required datasets"
93+ raise ValueError (msg )
94+ values = field_node ["values" ][start_idx :end_idx ]
95+ mask = field_node ["mask" ][start_idx :end_idx ]
96+ values_array = _ensure_numpy_array (values )
97+ if values_array .dtype .kind in {"S" , "O" , "U" }:
98+ values_array = _decode_string_array (values_array )
99+ mask_array = np .asarray (mask , dtype = bool )
100+ if mask_array .size == 0 or not mask_array .any ():
101+ return values_array
102+ result = values_array .astype (object , copy = True )
103+ result [mask_array ] = None
104+ return result
105+
106+
74107def _read_obs_field_slice (
75108 field_node : h5py .Dataset | h5py .Group ,
76109 field_name : str ,
@@ -87,6 +120,11 @@ def _read_obs_field_slice(
87120 categorical_cache [field_name ] = _load_categorical_values (field_node )
88121 categories = categorical_cache [field_name ]
89122 return _decode_categorical_codes (field_node , categories , start_idx , end_idx )
123+ if encoding_type in {"nullable-boolean" , "nullable-integer" , "nullable-string-array" }:
124+ if not isinstance (field_node , h5py .Group ):
125+ msg = f"Unexpected nullable node type for field '{ field_name } '"
126+ raise ValueError (msg )
127+ return _read_nullable_group_slice (field_node , field_name , start_idx , end_idx )
90128 if not isinstance (field_node , h5py .Dataset ):
91129 msg = f"Unsupported obs field storage for '{ field_name } '"
92130 raise ValueError (msg )
@@ -244,6 +282,8 @@ def preload_complex_obs_fields(
244282 is_dataset = isinstance (node , h5py .Dataset )
245283 if is_dataset or encoding_type == "categorical" :
246284 continue
285+ if encoding_type in {"nullable-boolean" , "nullable-integer" , "nullable-string-array" }:
286+ continue
247287 try :
248288 values = read_elem (node )
249289 except Exception as exc :
@@ -270,7 +310,7 @@ def infer_obs_schema(obs_group: h5py.Group) -> dict[str, pl.datatypes.DataType]:
270310 if name == "index" :
271311 continue
272312 encoding_type = node .attrs .get ("encoding-type" , "array" )
273- if encoding_type in {"categorical" , "string-array" }:
313+ if encoding_type in {"categorical" , "string-array" , "nullable-string-array" }:
274314 schema [name ] = pl .String
275315 continue
276316 if encoding_type == "nullable-boolean" :
@@ -279,9 +319,6 @@ def infer_obs_schema(obs_group: h5py.Group) -> dict[str, pl.datatypes.DataType]:
279319 if encoding_type == "nullable-integer" :
280320 schema [name ] = pl .Int64
281321 continue
282- if encoding_type == "nullable-string-array" :
283- schema [name ] = pl .String
284- continue
285322 if not isinstance (node , h5py .Dataset ):
286323 continue
287324 kind = node .dtype .kind
0 commit comments