Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 9 additions & 6 deletions src/modelarrayio/storage/tiledb_storage.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,9 @@

logger = logging.getLogger(__name__)

# TileDBArray R package expects dimensions to be R integers.
_DIM_DTYPE = np.int32


def resolve_dtype(storage_dtype):
"""Resolve a storage dtype to a supported NumPy floating type.
Expand Down Expand Up @@ -162,10 +165,10 @@ def create_scalar_matrix_array(

# Domain and schema
dim_subjects = tiledb.Dim(
name='subjects', domain=(0, n_files - 1), tile=tile_shape[0], dtype=np.int64
name='subjects', domain=(0, n_files - 1), tile=tile_shape[0], dtype=_DIM_DTYPE
)
dim_items = tiledb.Dim(
name='items', domain=(0, n_elements - 1), tile=tile_shape[1], dtype=np.int64
name='items', domain=(0, n_elements - 1), tile=tile_shape[1], dtype=_DIM_DTYPE
)
dom = tiledb.Domain(dim_subjects, dim_items)
attr_filters = _build_filter_list(compression, compression_level, shuffle)
Expand Down Expand Up @@ -254,10 +257,10 @@ def create_empty_scalar_matrix_array(
_ensure_parent_group(uri)

dim_subjects = tiledb.Dim(
name='subjects', domain=(0, n_files - 1), tile=tile_shape[0], dtype=np.int64
name='subjects', domain=(0, n_files - 1), tile=tile_shape[0], dtype=_DIM_DTYPE
)
dim_items = tiledb.Dim(
name='items', domain=(0, n_elements - 1), tile=tile_shape[1], dtype=np.int64
name='items', domain=(0, n_elements - 1), tile=tile_shape[1], dtype=_DIM_DTYPE
)
dom = tiledb.Domain(dim_subjects, dim_items)
attr_filters = _build_filter_list(compression, compression_level, shuffle)
Expand Down Expand Up @@ -349,7 +352,7 @@ def write_parcel_names(base_uri: str, array_path: str, names: Sequence[str]):

n = len(names)
dim_idx = tiledb.Dim(
name='idx', domain=(0, max(n - 1, 0)), tile=max(1, min(n, 1024)), dtype=np.int64
name='idx', domain=(0, max(n - 1, 0)), tile=max(1, min(n, 1024)), dtype=_DIM_DTYPE
)
dom = tiledb.Domain(dim_idx)
# np.unicode_ was removed in NumPy 2.0; np.str_ is the compatible string scalar.
Expand Down Expand Up @@ -383,7 +386,7 @@ def write_column_names(base_uri: str, scalar: str, sources: Sequence[str]):

n = len(sources)
dim_idx = tiledb.Dim(
name='idx', domain=(0, max(n - 1, 0)), tile=max(1, min(n, 1024)), dtype=np.int64
name='idx', domain=(0, max(n - 1, 0)), tile=max(1, min(n, 1024)), dtype=_DIM_DTYPE
)
dom = tiledb.Domain(dim_idx)
attr_values = tiledb.Attr(name='values', dtype=np.str_)
Expand Down
9 changes: 9 additions & 0 deletions test/test_tiledb_storage.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,11 @@
from modelarrayio.storage import tiledb_storage


def assert_domain_dtypes(array: tiledb.Array, expected: list[str]) -> None:
for idx, dtype in enumerate(expected):
assert array.schema.domain.dim(idx).dtype == np.dtype(dtype)


def test_build_filter_list_variants() -> None:
no_filters = tiledb_storage._build_filter_list(None, None, shuffle=False)
assert isinstance(no_filters, tiledb.FilterList)
Expand Down Expand Up @@ -41,6 +46,7 @@ def test_create_empty_scalar_matrix_array_writes_metadata_and_overwrites(tmp_pat
)
assert tiledb.object_type(uri) == 'array'
with tiledb.open(uri, 'r') as array:
assert_domain_dtypes(array, ['int32', 'int32'])
assert json.loads(array.meta['column_names']) == ['s1', 's2']

uri_again = tiledb_storage.create_empty_scalar_matrix_array(
Expand Down Expand Up @@ -96,6 +102,7 @@ def test_write_parcel_names_and_column_names(tmp_path: Path) -> None:
parcel_uri = base / 'parcels' / 'parcel_id'
assert tiledb.object_type(str(parcel_uri)) == 'array'
with tiledb.open(str(parcel_uri), 'r') as array:
assert_domain_dtypes(array, ['int32'])
np.testing.assert_array_equal(array[:]['values'], np.array(['P1', 'P2'], dtype=object))

# Some TileDB builds do not implicitly create missing parent directories.
Expand All @@ -104,6 +111,7 @@ def test_write_parcel_names_and_column_names(tmp_path: Path) -> None:
tiledb_storage.write_column_names(str(base), 'FA', ['sub-1', 'sub-2'])
column_uri = base / 'scalars' / 'FA' / 'column_names'
with tiledb.open(str(column_uri), 'r') as array:
assert_domain_dtypes(array, ['int32'])
np.testing.assert_array_equal(
array[:]['values'], np.array(['sub-1', 'sub-2'], dtype=object)
)
Expand All @@ -124,5 +132,6 @@ def test_create_scalar_matrix_array_writes_values_and_metadata(tmp_path: Path) -
compression_level=3,
)
with tiledb.open(uri, 'r') as array:
assert_domain_dtypes(array, ['int32', 'int32'])
np.testing.assert_array_equal(array[:]['values'], values)
assert json.loads(array.meta['column_names']) == ['first', 'second']