Skip to content
Merged
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
9 changes: 8 additions & 1 deletion src/spatialdata/_core/operations/aggregate.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -243,7 +243,14 @@ def _create_sdata_from_table_and_shapes(
) -> SpatialData:
from spatialdata._core._deepcopy import deepcopy as _deepcopy

table.obs[instance_key] = table.obs_names.copy()
shapes_index_dtype = shapes.index.dtype if isinstance(shapes, GeoDataFrame) else shapes.dtype
try:
table.obs[instance_key] = table.obs_names.copy().astype(shapes_index_dtype)
except ValueError as err:
raise TypeError(
f"Instance key column dtype in table resulting from aggregation cannot be cast to the dtype of"
f"element {shapes_name}.index"
) from err
table.obs[region_key] = shapes_name
table = TableModel.parse(table, region=shapes_name, region_key=region_key, instance_key=instance_key)

Expand Down
7 changes: 7 additions & 0 deletions src/spatialdata/_core/spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -199,6 +199,13 @@ def validate_table_in_spatialdata(self, table: AnnData) -> None:
else:
dtype = element.index.dtype
if dtype != table.obs[instance_key].dtype:
if dtype == str or table.obs[instance_key].dtype == str:
raise TypeError(
f"Table instance_key column ({instance_key}) has a dtype "
f"({table.obs[instance_key].dtype}) that does not match the dtype of the indices of "
f"the annotated element ({dtype})."
)

warnings.warn(
(
f"Table instance_key column ({instance_key}) has a dtype "
Expand Down
15 changes: 6 additions & 9 deletions src/spatialdata/models/models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from multiscale_spatial_image.multiscale_spatial_image import MultiscaleSpatialImage
from multiscale_spatial_image.to_multiscale.to_multiscale import Methods
from pandas import CategoricalDtype
from pandas.errors import IntCastingNaNError
from shapely._geometry import GeometryType
from shapely.geometry import MultiPolygon, Point, Polygon
from shapely.geometry.collection import GeometryCollection
Expand DownExpand Up@@ -795,6 +794,11 @@ def _validate_table_annotation_metadata(self, data: AnnData) -> None:
raise ValueError(f"`{attr[self.REGION_KEY_KEY]}` not found in `adata.obs`.")
if attr[self.INSTANCE_KEY] not in data.obs:
raise ValueError(f"`{attr[self.INSTANCE_KEY]}` not found in `adata.obs`.")
if (dtype := data.obs[attr[self.INSTANCE_KEY]].dtype) not in [np.int16, np.int32, np.int64, str]:
raise TypeError(
f"Only np.int16, np.int32, np.int64 or string allowed as dtype for "
f"instance_key column in obs. Dtype found to be {dtype}"
)
expected_regions = attr[self.REGION_KEY] if isinstance(attr[self.REGION_KEY], list) else [attr[self.REGION_KEY]]
found_regions = data.obs[attr[self.REGION_KEY_KEY]].unique().tolist()
if len(set(expected_regions).symmetric_difference(set(found_regions))) > 0:
Expand DownExpand Up@@ -881,14 +885,6 @@ def parse(
adata.obs[region_key] = pd.Categorical(adata.obs[region_key])
if instance_key is None:
raise ValueError("`instance_key` must be provided.")
if adata.obs[instance_key].dtype != int:
try:
warnings.warn(
f"Converting `{cls.INSTANCE_KEY}: {instance_key}` to integer dtype.", UserWarning, stacklevel=2
)
adata.obs[instance_key] = adata.obs[instance_key].astype(int)
except IntCastingNaNError as exc:
raise ValueError("Values within table.obs[] must be able to be coerced to int dtype.") from exc

grouped = adata.obs.groupby(region_key, observed=True)
grouped_size = grouped.size()
Expand All@@ -901,6 +897,7 @@ def parse(

attr = {"region": region, "region_key": region_key, "instance_key": instance_key}
adata.uns[cls.ATTRS_KEY] = attr
cls().validate(adata)
return adata


Expand Down
11 changes: 2 additions & 9 deletions tests/core/operations/test_spatialdata_operations.py
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
from __future__ import annotations

import math
import warnings

import numpy as np
import pytest
Expand DownExpand Up@@ -419,10 +418,7 @@ def test_validate_table_in_spatialdata(full_sdata):
region, region_key, _ = get_table_keys(table)
assert region == "labels2d"

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.labels["labels2d"] = Labels2DModel.parse(full_sdata.labels["labels2d"].astype("int16"))
Expand All@@ -437,10 +433,7 @@ def test_validate_table_in_spatialdata(full_sdata):
table.obs[region_key] = "points_0"
full_sdata.set_table_annotates_spatialelement("table", region="points_0")

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.points["points_0"].index = full_sdata.points["points_0"].index.astype("int16")
Expand Down
31 changes: 31 additions & 0 deletions tests/core/query/test_relational_query.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,6 +22,37 @@ def test_match_table_to_element(sdata_query_aggregation):
# TODO: add tests for labels


def test_join_using_string_instance_id_and_index(sdata_query_aggregation):
sdata_query_aggregation["table"].obs["instance_id"] = [
f"string_{i}" for i in sdata_query_aggregation["table"].obs["instance_id"]
]
sdata_query_aggregation["values_circles"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_circles"].index]
)
sdata_query_aggregation["values_polygons"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_polygons"].index]
)

sdata_query_aggregation["values_polygons"] = sdata_query_aggregation["values_polygons"][:5]
sdata_query_aggregation["values_circles"] = sdata_query_aggregation["values_circles"][:5]

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "inner"
)
# Note that we started with 21 n_obs.
assert table.n_obs == 10

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right_exclusive"
)
assert table.n_obs == 11

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right"
)
assert table.n_obs == 21


def test_left_inner_right_exclusive_join(sdata_query_aggregation):
element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, "values_polygons", "table", "right_exclusive"
Expand Down
18 changes: 8 additions & 10 deletions tests/models/test_models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -318,6 +318,14 @@ def test_table_model(
region: str | np.ndarray,
) -> None:
region_key = "reg"
obs = pd.DataFrame(
RNG.choice(np.arange(0, 100, dtype=float), size=(10, 3), replace=False), columns=["A", "B", "C"]
)
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(TypeError, match="Only np.int16"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs = pd.DataFrame(RNG.choice(np.arange(0, 100), size=(10, 3), replace=False), columns=["A", "B", "C"])
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
Expand All@@ -332,16 +340,6 @@ def test_table_model(
assert TableModel.REGION_KEY_KEY in table.uns[TableModel.ATTRS_KEY]
assert table.uns[TableModel.ATTRS_KEY][TableModel.REGION_KEY] == region

obs["A"] = obs["A"].astype(str)
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.warns(UserWarning, match="Converting"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs["A"] = pd.Series(len([chr(ord("a") + i) for i in range(10)]))
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(ValueError, match="Values within"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

@pytest.mark.parametrize("model", [TableModel])
@pytest.mark.parametrize("region", [["sample_1"] * 5 + ["sample_2"] * 5])
def test_table_instance_key_values_not_unique(self, model: TableModel, region: str | np.ndarray):
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { // Add copy buttons to all
 blocks
(function() {
function addCopyButtons() {
document.querySelectorAll('pre code').forEach(function(codeBlock) {
if (codeBlock.parentElement.hasAttribute('data-copy-added')) return;
codeBlock.parentElement.setAttribute('data-copy-added', 'true');
var btn = document.createElement('button');
btn.textContent = 'Copy';
btn.style.cssText = 'position:absolute;top:4px;right:4px;padding:2px 8px;font-size:11px;background:#4ecdc4;border:none;border-radius:4px;color:#1a1a2e;cursor:pointer;opacity:0.7;transition:opacity 0.2s;';
btn.onmouseover = function() { this.style.opacity = '1'; };
btn.onmouseout = function() { this.style.opacity = '0.7'; };
btn.onclick = function() {
navigator.clipboard.writeText(codeBlock.textContent).then(function() {
btn.textContent = 'Copied!';
setTimeout(function() { btn.textContent = 'Copy'; }, 1500);
});
};
codeBlock.parentElement.style.position = 'relative';
codeBlock.parentElement.appendChild(btn);
});
}
addCopyButtons();
// Re-run on dynamic content
var observer = new MutationObserver(addCopyButtons);
observer.observe(document.body, { childList: true, subtree: true });
})();
}
} catch(__e) { console.warn('[Userscript:Add Copy Buttons to Code Blocks]', __e); }
})();
(function(){
try {
var __m = "github.com";
var __re = new RegExp('^' + "github\\.com" + '
Test joins with string indices and instance id by melonora · Pull Request #485 · scverse/spatialdata · GitHub
Skip to content
Merged
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
9 changes: 8 additions & 1 deletion src/spatialdata/_core/operations/aggregate.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -243,7 +243,14 @@ def _create_sdata_from_table_and_shapes(
) -> SpatialData:
from spatialdata._core._deepcopy import deepcopy as _deepcopy

table.obs[instance_key] = table.obs_names.copy()
shapes_index_dtype = shapes.index.dtype if isinstance(shapes, GeoDataFrame) else shapes.dtype
try:
table.obs[instance_key] = table.obs_names.copy().astype(shapes_index_dtype)
except ValueError as err:
raise TypeError(
f"Instance key column dtype in table resulting from aggregation cannot be cast to the dtype of"
f"element {shapes_name}.index"
) from err
table.obs[region_key] = shapes_name
table = TableModel.parse(table, region=shapes_name, region_key=region_key, instance_key=instance_key)

Expand Down
7 changes: 7 additions & 0 deletions src/spatialdata/_core/spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -199,6 +199,13 @@ def validate_table_in_spatialdata(self, table: AnnData) -> None:
else:
dtype = element.index.dtype
if dtype != table.obs[instance_key].dtype:
if dtype == str or table.obs[instance_key].dtype == str:
raise TypeError(
f"Table instance_key column ({instance_key}) has a dtype "
f"({table.obs[instance_key].dtype}) that does not match the dtype of the indices of "
f"the annotated element ({dtype})."
)

warnings.warn(
(
f"Table instance_key column ({instance_key}) has a dtype "
Expand Down
15 changes: 6 additions & 9 deletions src/spatialdata/models/models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from multiscale_spatial_image.multiscale_spatial_image import MultiscaleSpatialImage
from multiscale_spatial_image.to_multiscale.to_multiscale import Methods
from pandas import CategoricalDtype
from pandas.errors import IntCastingNaNError
from shapely._geometry import GeometryType
from shapely.geometry import MultiPolygon, Point, Polygon
from shapely.geometry.collection import GeometryCollection
Expand DownExpand Up@@ -795,6 +794,11 @@ def _validate_table_annotation_metadata(self, data: AnnData) -> None:
raise ValueError(f"`{attr[self.REGION_KEY_KEY]}` not found in `adata.obs`.")
if attr[self.INSTANCE_KEY] not in data.obs:
raise ValueError(f"`{attr[self.INSTANCE_KEY]}` not found in `adata.obs`.")
if (dtype := data.obs[attr[self.INSTANCE_KEY]].dtype) not in [np.int16, np.int32, np.int64, str]:
raise TypeError(
f"Only np.int16, np.int32, np.int64 or string allowed as dtype for "
f"instance_key column in obs. Dtype found to be {dtype}"
)
expected_regions = attr[self.REGION_KEY] if isinstance(attr[self.REGION_KEY], list) else [attr[self.REGION_KEY]]
found_regions = data.obs[attr[self.REGION_KEY_KEY]].unique().tolist()
if len(set(expected_regions).symmetric_difference(set(found_regions))) > 0:
Expand DownExpand Up@@ -881,14 +885,6 @@ def parse(
adata.obs[region_key] = pd.Categorical(adata.obs[region_key])
if instance_key is None:
raise ValueError("`instance_key` must be provided.")
if adata.obs[instance_key].dtype != int:
try:
warnings.warn(
f"Converting `{cls.INSTANCE_KEY}: {instance_key}` to integer dtype.", UserWarning, stacklevel=2
)
adata.obs[instance_key] = adata.obs[instance_key].astype(int)
except IntCastingNaNError as exc:
raise ValueError("Values within table.obs[] must be able to be coerced to int dtype.") from exc

grouped = adata.obs.groupby(region_key, observed=True)
grouped_size = grouped.size()
Expand All@@ -901,6 +897,7 @@ def parse(

attr = {"region": region, "region_key": region_key, "instance_key": instance_key}
adata.uns[cls.ATTRS_KEY] = attr
cls().validate(adata)
return adata


Expand Down
11 changes: 2 additions & 9 deletions tests/core/operations/test_spatialdata_operations.py
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
from __future__ import annotations

import math
import warnings

import numpy as np
import pytest
Expand DownExpand Up@@ -419,10 +418,7 @@ def test_validate_table_in_spatialdata(full_sdata):
region, region_key, _ = get_table_keys(table)
assert region == "labels2d"

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.labels["labels2d"] = Labels2DModel.parse(full_sdata.labels["labels2d"].astype("int16"))
Expand All@@ -437,10 +433,7 @@ def test_validate_table_in_spatialdata(full_sdata):
table.obs[region_key] = "points_0"
full_sdata.set_table_annotates_spatialelement("table", region="points_0")

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.points["points_0"].index = full_sdata.points["points_0"].index.astype("int16")
Expand Down
31 changes: 31 additions & 0 deletions tests/core/query/test_relational_query.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,6 +22,37 @@ def test_match_table_to_element(sdata_query_aggregation):
# TODO: add tests for labels


def test_join_using_string_instance_id_and_index(sdata_query_aggregation):
sdata_query_aggregation["table"].obs["instance_id"] = [
f"string_{i}" for i in sdata_query_aggregation["table"].obs["instance_id"]
]
sdata_query_aggregation["values_circles"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_circles"].index]
)
sdata_query_aggregation["values_polygons"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_polygons"].index]
)

sdata_query_aggregation["values_polygons"] = sdata_query_aggregation["values_polygons"][:5]
sdata_query_aggregation["values_circles"] = sdata_query_aggregation["values_circles"][:5]

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "inner"
)
# Note that we started with 21 n_obs.
assert table.n_obs == 10

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right_exclusive"
)
assert table.n_obs == 11

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right"
)
assert table.n_obs == 21


def test_left_inner_right_exclusive_join(sdata_query_aggregation):
element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, "values_polygons", "table", "right_exclusive"
Expand Down
18 changes: 8 additions & 10 deletions tests/models/test_models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -318,6 +318,14 @@ def test_table_model(
region: str | np.ndarray,
) -> None:
region_key = "reg"
obs = pd.DataFrame(
RNG.choice(np.arange(0, 100, dtype=float), size=(10, 3), replace=False), columns=["A", "B", "C"]
)
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(TypeError, match="Only np.int16"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs = pd.DataFrame(RNG.choice(np.arange(0, 100), size=(10, 3), replace=False), columns=["A", "B", "C"])
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
Expand All@@ -332,16 +340,6 @@ def test_table_model(
assert TableModel.REGION_KEY_KEY in table.uns[TableModel.ATTRS_KEY]
assert table.uns[TableModel.ATTRS_KEY][TableModel.REGION_KEY] == region

obs["A"] = obs["A"].astype(str)
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.warns(UserWarning, match="Converting"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs["A"] = pd.Series(len([chr(ord("a") + i) for i in range(10)]))
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(ValueError, match="Values within"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

@pytest.mark.parametrize("model", [TableModel])
@pytest.mark.parametrize("region", [["sample_1"] * 5 + ["sample_2"] * 5])
def test_table_instance_key_values_not_unique(self, model: TableModel, region: str | np.ndarray):
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { // Force GitHub README to respect dark mode (function() { var style = document.createElement('style'); style.textContent = ' .markdown-body { color-scheme: dark light; } .markdown-body pre { background: #161b22 !important; } .markdown-body code { background: rgba(110, 118, 129, 0.4) !important; } .markdown-body table th, .markdown-body table td { border-color: #30363d !important; } .markdown-body img { background: #0d1117; } .markdown-body blockquote { border-left-color: #8b949e; } .markdown-body hr { border-color: #30363d; } '; document.head.appendChild(style); })(); } } catch(__e) { console.warn('[Userscript:GitHub Dark Mode README Fix]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Test joins with string indices and instance id by melonora · Pull Request #485 · scverse/spatialdata · GitHub
Skip to content
Merged
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
9 changes: 8 additions & 1 deletion src/spatialdata/_core/operations/aggregate.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -243,7 +243,14 @@ def _create_sdata_from_table_and_shapes(
) -> SpatialData:
from spatialdata._core._deepcopy import deepcopy as _deepcopy

table.obs[instance_key] = table.obs_names.copy()
shapes_index_dtype = shapes.index.dtype if isinstance(shapes, GeoDataFrame) else shapes.dtype
try:
table.obs[instance_key] = table.obs_names.copy().astype(shapes_index_dtype)
except ValueError as err:
raise TypeError(
f"Instance key column dtype in table resulting from aggregation cannot be cast to the dtype of"
f"element {shapes_name}.index"
) from err
table.obs[region_key] = shapes_name
table = TableModel.parse(table, region=shapes_name, region_key=region_key, instance_key=instance_key)

Expand Down
7 changes: 7 additions & 0 deletions src/spatialdata/_core/spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -199,6 +199,13 @@ def validate_table_in_spatialdata(self, table: AnnData) -> None:
else:
dtype = element.index.dtype
if dtype != table.obs[instance_key].dtype:
if dtype == str or table.obs[instance_key].dtype == str:
raise TypeError(
f"Table instance_key column ({instance_key}) has a dtype "
f"({table.obs[instance_key].dtype}) that does not match the dtype of the indices of "
f"the annotated element ({dtype})."
)

warnings.warn(
(
f"Table instance_key column ({instance_key}) has a dtype "
Expand Down
15 changes: 6 additions & 9 deletions src/spatialdata/models/models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from multiscale_spatial_image.multiscale_spatial_image import MultiscaleSpatialImage
from multiscale_spatial_image.to_multiscale.to_multiscale import Methods
from pandas import CategoricalDtype
from pandas.errors import IntCastingNaNError
from shapely._geometry import GeometryType
from shapely.geometry import MultiPolygon, Point, Polygon
from shapely.geometry.collection import GeometryCollection
Expand DownExpand Up@@ -795,6 +794,11 @@ def _validate_table_annotation_metadata(self, data: AnnData) -> None:
raise ValueError(f"`{attr[self.REGION_KEY_KEY]}` not found in `adata.obs`.")
if attr[self.INSTANCE_KEY] not in data.obs:
raise ValueError(f"`{attr[self.INSTANCE_KEY]}` not found in `adata.obs`.")
if (dtype := data.obs[attr[self.INSTANCE_KEY]].dtype) not in [np.int16, np.int32, np.int64, str]:
raise TypeError(
f"Only np.int16, np.int32, np.int64 or string allowed as dtype for "
f"instance_key column in obs. Dtype found to be {dtype}"
)
expected_regions = attr[self.REGION_KEY] if isinstance(attr[self.REGION_KEY], list) else [attr[self.REGION_KEY]]
found_regions = data.obs[attr[self.REGION_KEY_KEY]].unique().tolist()
if len(set(expected_regions).symmetric_difference(set(found_regions))) > 0:
Expand DownExpand Up@@ -881,14 +885,6 @@ def parse(
adata.obs[region_key] = pd.Categorical(adata.obs[region_key])
if instance_key is None:
raise ValueError("`instance_key` must be provided.")
if adata.obs[instance_key].dtype != int:
try:
warnings.warn(
f"Converting `{cls.INSTANCE_KEY}: {instance_key}` to integer dtype.", UserWarning, stacklevel=2
)
adata.obs[instance_key] = adata.obs[instance_key].astype(int)
except IntCastingNaNError as exc:
raise ValueError("Values within table.obs[] must be able to be coerced to int dtype.") from exc

grouped = adata.obs.groupby(region_key, observed=True)
grouped_size = grouped.size()
Expand All@@ -901,6 +897,7 @@ def parse(

attr = {"region": region, "region_key": region_key, "instance_key": instance_key}
adata.uns[cls.ATTRS_KEY] = attr
cls().validate(adata)
return adata


Expand Down
11 changes: 2 additions & 9 deletions tests/core/operations/test_spatialdata_operations.py
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
from __future__ import annotations

import math
import warnings

import numpy as np
import pytest
Expand DownExpand Up@@ -419,10 +418,7 @@ def test_validate_table_in_spatialdata(full_sdata):
region, region_key, _ = get_table_keys(table)
assert region == "labels2d"

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.labels["labels2d"] = Labels2DModel.parse(full_sdata.labels["labels2d"].astype("int16"))
Expand All@@ -437,10 +433,7 @@ def test_validate_table_in_spatialdata(full_sdata):
table.obs[region_key] = "points_0"
full_sdata.set_table_annotates_spatialelement("table", region="points_0")

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.points["points_0"].index = full_sdata.points["points_0"].index.astype("int16")
Expand Down
31 changes: 31 additions & 0 deletions tests/core/query/test_relational_query.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,6 +22,37 @@ def test_match_table_to_element(sdata_query_aggregation):
# TODO: add tests for labels


def test_join_using_string_instance_id_and_index(sdata_query_aggregation):
sdata_query_aggregation["table"].obs["instance_id"] = [
f"string_{i}" for i in sdata_query_aggregation["table"].obs["instance_id"]
]
sdata_query_aggregation["values_circles"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_circles"].index]
)
sdata_query_aggregation["values_polygons"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_polygons"].index]
)

sdata_query_aggregation["values_polygons"] = sdata_query_aggregation["values_polygons"][:5]
sdata_query_aggregation["values_circles"] = sdata_query_aggregation["values_circles"][:5]

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "inner"
)
# Note that we started with 21 n_obs.
assert table.n_obs == 10

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right_exclusive"
)
assert table.n_obs == 11

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right"
)
assert table.n_obs == 21


def test_left_inner_right_exclusive_join(sdata_query_aggregation):
element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, "values_polygons", "table", "right_exclusive"
Expand Down
18 changes: 8 additions & 10 deletions tests/models/test_models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -318,6 +318,14 @@ def test_table_model(
region: str | np.ndarray,
) -> None:
region_key = "reg"
obs = pd.DataFrame(
RNG.choice(np.arange(0, 100, dtype=float), size=(10, 3), replace=False), columns=["A", "B", "C"]
)
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(TypeError, match="Only np.int16"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs = pd.DataFrame(RNG.choice(np.arange(0, 100), size=(10, 3), replace=False), columns=["A", "B", "C"])
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
Expand All@@ -332,16 +340,6 @@ def test_table_model(
assert TableModel.REGION_KEY_KEY in table.uns[TableModel.ATTRS_KEY]
assert table.uns[TableModel.ATTRS_KEY][TableModel.REGION_KEY] == region

obs["A"] = obs["A"].astype(str)
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.warns(UserWarning, match="Converting"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs["A"] = pd.Series(len([chr(ord("a") + i) for i in range(10)]))
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(ValueError, match="Values within"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

@pytest.mark.parametrize("model", [TableModel])
@pytest.mark.parametrize("region", [["sample_1"] * 5 + ["sample_2"] * 5])
def test_table_instance_key_values_not_unique(self, model: TableModel, region: str | np.ndarray):
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { // Highlight search terms from Google/DuckDuckGo/Bing referrer (function() { var ref = document.referrer; var terms = []; if (ref.includes('google.com') || ref.includes('duckduckgo.com') || ref.includes('bing.com')) { var url = new URL(ref); var q = url.searchParams.get('q') || url.searchParams.get('p'); if (q) { terms = q.split(/\s+/).filter(function(t) { return t.length > 2; }); } } if (terms.length === 0) return; var style = document.createElement('style'); style.textContent = '.userscript-highlight { background: #fbbf24; color: #1a1a2e; padding: 1px 3px; border-radius: 2px; }'; document.head.appendChild(style); function highlight(node) { if (node.nodeType === 3) { // text node var text = node.textContent; var found = false; terms.forEach(function(term) { var regex = new RegExp('(' + term.replace(/[.*+?^${}()|[\]\\]/g, '\\') + ')', 'gi'); if (regex.test(text)) { found = true; var frag = document.createDocumentFragment(); var parts = text.split(regex); parts.forEach(function(part, i) { if (i % 2 === 0) { frag.appendChild(document.createTextNode(part)); } else { var span = document.createElement('span'); span.className = 'userscript-highlight'; span.textContent = part; frag.appendChild(span); } }); node.parentNode.replaceChild(frag, node); } }); } else if (node.nodeType === 1 && node.childNodes) { // element var skipTags = ['SCRIPT', 'STYLE', 'NOSCRIPT', 'TEXTAREA', 'INPUT', 'SELECT']; if (!skipTags.includes(node.tagName)) { Array.from(node.childNodes).forEach(highlight); } } } highlight(document.body); // Re-highlight on dynamic content var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1 || node.nodeType === 3) highlight(node); }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Highlight Search Terms]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Test joins with string indices and instance id by melonora · Pull Request #485 · scverse/spatialdata · GitHub
Skip to content
Merged
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
9 changes: 8 additions & 1 deletion src/spatialdata/_core/operations/aggregate.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -243,7 +243,14 @@ def _create_sdata_from_table_and_shapes(
) -> SpatialData:
from spatialdata._core._deepcopy import deepcopy as _deepcopy

table.obs[instance_key] = table.obs_names.copy()
shapes_index_dtype = shapes.index.dtype if isinstance(shapes, GeoDataFrame) else shapes.dtype
try:
table.obs[instance_key] = table.obs_names.copy().astype(shapes_index_dtype)
except ValueError as err:
raise TypeError(
f"Instance key column dtype in table resulting from aggregation cannot be cast to the dtype of"
f"element {shapes_name}.index"
) from err
table.obs[region_key] = shapes_name
table = TableModel.parse(table, region=shapes_name, region_key=region_key, instance_key=instance_key)

Expand Down
7 changes: 7 additions & 0 deletions src/spatialdata/_core/spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -199,6 +199,13 @@ def validate_table_in_spatialdata(self, table: AnnData) -> None:
else:
dtype = element.index.dtype
if dtype != table.obs[instance_key].dtype:
if dtype == str or table.obs[instance_key].dtype == str:
raise TypeError(
f"Table instance_key column ({instance_key}) has a dtype "
f"({table.obs[instance_key].dtype}) that does not match the dtype of the indices of "
f"the annotated element ({dtype})."
)

warnings.warn(
(
f"Table instance_key column ({instance_key}) has a dtype "
Expand Down
15 changes: 6 additions & 9 deletions src/spatialdata/models/models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from multiscale_spatial_image.multiscale_spatial_image import MultiscaleSpatialImage
from multiscale_spatial_image.to_multiscale.to_multiscale import Methods
from pandas import CategoricalDtype
from pandas.errors import IntCastingNaNError
from shapely._geometry import GeometryType
from shapely.geometry import MultiPolygon, Point, Polygon
from shapely.geometry.collection import GeometryCollection
Expand DownExpand Up@@ -795,6 +794,11 @@ def _validate_table_annotation_metadata(self, data: AnnData) -> None:
raise ValueError(f"`{attr[self.REGION_KEY_KEY]}` not found in `adata.obs`.")
if attr[self.INSTANCE_KEY] not in data.obs:
raise ValueError(f"`{attr[self.INSTANCE_KEY]}` not found in `adata.obs`.")
if (dtype := data.obs[attr[self.INSTANCE_KEY]].dtype) not in [np.int16, np.int32, np.int64, str]:
raise TypeError(
f"Only np.int16, np.int32, np.int64 or string allowed as dtype for "
f"instance_key column in obs. Dtype found to be {dtype}"
)
expected_regions = attr[self.REGION_KEY] if isinstance(attr[self.REGION_KEY], list) else [attr[self.REGION_KEY]]
found_regions = data.obs[attr[self.REGION_KEY_KEY]].unique().tolist()
if len(set(expected_regions).symmetric_difference(set(found_regions))) > 0:
Expand DownExpand Up@@ -881,14 +885,6 @@ def parse(
adata.obs[region_key] = pd.Categorical(adata.obs[region_key])
if instance_key is None:
raise ValueError("`instance_key` must be provided.")
if adata.obs[instance_key].dtype != int:
try:
warnings.warn(
f"Converting `{cls.INSTANCE_KEY}: {instance_key}` to integer dtype.", UserWarning, stacklevel=2
)
adata.obs[instance_key] = adata.obs[instance_key].astype(int)
except IntCastingNaNError as exc:
raise ValueError("Values within table.obs[] must be able to be coerced to int dtype.") from exc

grouped = adata.obs.groupby(region_key, observed=True)
grouped_size = grouped.size()
Expand All@@ -901,6 +897,7 @@ def parse(

attr = {"region": region, "region_key": region_key, "instance_key": instance_key}
adata.uns[cls.ATTRS_KEY] = attr
cls().validate(adata)
return adata


Expand Down
11 changes: 2 additions & 9 deletions tests/core/operations/test_spatialdata_operations.py
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
from __future__ import annotations

import math
import warnings

import numpy as np
import pytest
Expand DownExpand Up@@ -419,10 +418,7 @@ def test_validate_table_in_spatialdata(full_sdata):
region, region_key, _ = get_table_keys(table)
assert region == "labels2d"

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.labels["labels2d"] = Labels2DModel.parse(full_sdata.labels["labels2d"].astype("int16"))
Expand All@@ -437,10 +433,7 @@ def test_validate_table_in_spatialdata(full_sdata):
table.obs[region_key] = "points_0"
full_sdata.set_table_annotates_spatialelement("table", region="points_0")

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.points["points_0"].index = full_sdata.points["points_0"].index.astype("int16")
Expand Down
31 changes: 31 additions & 0 deletions tests/core/query/test_relational_query.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,6 +22,37 @@ def test_match_table_to_element(sdata_query_aggregation):
# TODO: add tests for labels


def test_join_using_string_instance_id_and_index(sdata_query_aggregation):
sdata_query_aggregation["table"].obs["instance_id"] = [
f"string_{i}" for i in sdata_query_aggregation["table"].obs["instance_id"]
]
sdata_query_aggregation["values_circles"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_circles"].index]
)
sdata_query_aggregation["values_polygons"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_polygons"].index]
)

sdata_query_aggregation["values_polygons"] = sdata_query_aggregation["values_polygons"][:5]
sdata_query_aggregation["values_circles"] = sdata_query_aggregation["values_circles"][:5]

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "inner"
)
# Note that we started with 21 n_obs.
assert table.n_obs == 10

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right_exclusive"
)
assert table.n_obs == 11

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right"
)
assert table.n_obs == 21


def test_left_inner_right_exclusive_join(sdata_query_aggregation):
element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, "values_polygons", "table", "right_exclusive"
Expand Down
18 changes: 8 additions & 10 deletions tests/models/test_models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -318,6 +318,14 @@ def test_table_model(
region: str | np.ndarray,
) -> None:
region_key = "reg"
obs = pd.DataFrame(
RNG.choice(np.arange(0, 100, dtype=float), size=(10, 3), replace=False), columns=["A", "B", "C"]
)
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(TypeError, match="Only np.int16"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs = pd.DataFrame(RNG.choice(np.arange(0, 100), size=(10, 3), replace=False), columns=["A", "B", "C"])
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
Expand All@@ -332,16 +340,6 @@ def test_table_model(
assert TableModel.REGION_KEY_KEY in table.uns[TableModel.ATTRS_KEY]
assert table.uns[TableModel.ATTRS_KEY][TableModel.REGION_KEY] == region

obs["A"] = obs["A"].astype(str)
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.warns(UserWarning, match="Converting"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs["A"] = pd.Series(len([chr(ord("a") + i) for i in range(10)]))
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(ValueError, match="Values within"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

@pytest.mark.parametrize("model", [TableModel])
@pytest.mark.parametrize("region", [["sample_1"] * 5 + ["sample_2"] * 5])
def test_table_instance_key_values_not_unique(self, model: TableModel, region: str | np.ndarray):
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { // Strip utm_, fbclid, gclid, etc. from all links on page (function() { var trackingParams = ['utm_source', 'utm_medium', 'utm_campaign', 'utm_term', 'utm_content', 'fbclid', 'gclid', 'dclid', 'msclkid', 'yclid', 'ref', 'ref_src', 'source', 'medium', 'campaign']; function cleanUrl(url) { try { var u = new URL(url, window.location.origin); var changed = false; trackingParams.forEach(function(p) { if (u.searchParams.has(p)) { u.searchParams.delete(p); changed = true; } }); return changed ? u.toString() : url; } catch (e) { return url; } } function cleanLinks() { document.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } cleanLinks(); var observer = new MutationObserver(function(mutations) { mutations.forEach(function(m) { m.addedNodes.forEach(function(node) { if (node.nodeType === 1) { if (node.tagName === 'A') cleanLinks(); node.querySelectorAll('a[href]').forEach(function(a) { var clean = cleanUrl(a.href); if (clean !== a.href) a.href = clean; }); } }); }); }); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:Remove Tracking Parameters from Links]', __e); } })(); (function(){ try { var __m = "youtube.com"; var __re = new RegExp('^' + "youtube\\.com" + ' Test joins with string indices and instance id by melonora · Pull Request #485 · scverse/spatialdata · GitHub
Skip to content
Merged
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
9 changes: 8 additions & 1 deletion src/spatialdata/_core/operations/aggregate.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -243,7 +243,14 @@ def _create_sdata_from_table_and_shapes(
) -> SpatialData:
from spatialdata._core._deepcopy import deepcopy as _deepcopy

table.obs[instance_key] = table.obs_names.copy()
shapes_index_dtype = shapes.index.dtype if isinstance(shapes, GeoDataFrame) else shapes.dtype
try:
table.obs[instance_key] = table.obs_names.copy().astype(shapes_index_dtype)
except ValueError as err:
raise TypeError(
f"Instance key column dtype in table resulting from aggregation cannot be cast to the dtype of"
f"element {shapes_name}.index"
) from err
table.obs[region_key] = shapes_name
table = TableModel.parse(table, region=shapes_name, region_key=region_key, instance_key=instance_key)

Expand Down
7 changes: 7 additions & 0 deletions src/spatialdata/_core/spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -199,6 +199,13 @@ def validate_table_in_spatialdata(self, table: AnnData) -> None:
else:
dtype = element.index.dtype
if dtype != table.obs[instance_key].dtype:
if dtype == str or table.obs[instance_key].dtype == str:
raise TypeError(
f"Table instance_key column ({instance_key}) has a dtype "
f"({table.obs[instance_key].dtype}) that does not match the dtype of the indices of "
f"the annotated element ({dtype})."
)

warnings.warn(
(
f"Table instance_key column ({instance_key}) has a dtype "
Expand Down
15 changes: 6 additions & 9 deletions src/spatialdata/models/models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from multiscale_spatial_image.multiscale_spatial_image import MultiscaleSpatialImage
from multiscale_spatial_image.to_multiscale.to_multiscale import Methods
from pandas import CategoricalDtype
from pandas.errors import IntCastingNaNError
from shapely._geometry import GeometryType
from shapely.geometry import MultiPolygon, Point, Polygon
from shapely.geometry.collection import GeometryCollection
Expand DownExpand Up@@ -795,6 +794,11 @@ def _validate_table_annotation_metadata(self, data: AnnData) -> None:
raise ValueError(f"`{attr[self.REGION_KEY_KEY]}` not found in `adata.obs`.")
if attr[self.INSTANCE_KEY] not in data.obs:
raise ValueError(f"`{attr[self.INSTANCE_KEY]}` not found in `adata.obs`.")
if (dtype := data.obs[attr[self.INSTANCE_KEY]].dtype) not in [np.int16, np.int32, np.int64, str]:
raise TypeError(
f"Only np.int16, np.int32, np.int64 or string allowed as dtype for "
f"instance_key column in obs. Dtype found to be {dtype}"
)
expected_regions = attr[self.REGION_KEY] if isinstance(attr[self.REGION_KEY], list) else [attr[self.REGION_KEY]]
found_regions = data.obs[attr[self.REGION_KEY_KEY]].unique().tolist()
if len(set(expected_regions).symmetric_difference(set(found_regions))) > 0:
Expand DownExpand Up@@ -881,14 +885,6 @@ def parse(
adata.obs[region_key] = pd.Categorical(adata.obs[region_key])
if instance_key is None:
raise ValueError("`instance_key` must be provided.")
if adata.obs[instance_key].dtype != int:
try:
warnings.warn(
f"Converting `{cls.INSTANCE_KEY}: {instance_key}` to integer dtype.", UserWarning, stacklevel=2
)
adata.obs[instance_key] = adata.obs[instance_key].astype(int)
except IntCastingNaNError as exc:
raise ValueError("Values within table.obs[] must be able to be coerced to int dtype.") from exc

grouped = adata.obs.groupby(region_key, observed=True)
grouped_size = grouped.size()
Expand All@@ -901,6 +897,7 @@ def parse(

attr = {"region": region, "region_key": region_key, "instance_key": instance_key}
adata.uns[cls.ATTRS_KEY] = attr
cls().validate(adata)
return adata


Expand Down
11 changes: 2 additions & 9 deletions tests/core/operations/test_spatialdata_operations.py
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
from __future__ import annotations

import math
import warnings

import numpy as np
import pytest
Expand DownExpand Up@@ -419,10 +418,7 @@ def test_validate_table_in_spatialdata(full_sdata):
region, region_key, _ = get_table_keys(table)
assert region == "labels2d"

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.labels["labels2d"] = Labels2DModel.parse(full_sdata.labels["labels2d"].astype("int16"))
Expand All@@ -437,10 +433,7 @@ def test_validate_table_in_spatialdata(full_sdata):
table.obs[region_key] = "points_0"
full_sdata.set_table_annotates_spatialelement("table", region="points_0")

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.points["points_0"].index = full_sdata.points["points_0"].index.astype("int16")
Expand Down
31 changes: 31 additions & 0 deletions tests/core/query/test_relational_query.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,6 +22,37 @@ def test_match_table_to_element(sdata_query_aggregation):
# TODO: add tests for labels


def test_join_using_string_instance_id_and_index(sdata_query_aggregation):
sdata_query_aggregation["table"].obs["instance_id"] = [
f"string_{i}" for i in sdata_query_aggregation["table"].obs["instance_id"]
]
sdata_query_aggregation["values_circles"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_circles"].index]
)
sdata_query_aggregation["values_polygons"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_polygons"].index]
)

sdata_query_aggregation["values_polygons"] = sdata_query_aggregation["values_polygons"][:5]
sdata_query_aggregation["values_circles"] = sdata_query_aggregation["values_circles"][:5]

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "inner"
)
# Note that we started with 21 n_obs.
assert table.n_obs == 10

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right_exclusive"
)
assert table.n_obs == 11

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right"
)
assert table.n_obs == 21


def test_left_inner_right_exclusive_join(sdata_query_aggregation):
element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, "values_polygons", "table", "right_exclusive"
Expand Down
18 changes: 8 additions & 10 deletions tests/models/test_models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -318,6 +318,14 @@ def test_table_model(
region: str | np.ndarray,
) -> None:
region_key = "reg"
obs = pd.DataFrame(
RNG.choice(np.arange(0, 100, dtype=float), size=(10, 3), replace=False), columns=["A", "B", "C"]
)
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(TypeError, match="Only np.int16"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs = pd.DataFrame(RNG.choice(np.arange(0, 100), size=(10, 3), replace=False), columns=["A", "B", "C"])
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
Expand All@@ -332,16 +340,6 @@ def test_table_model(
assert TableModel.REGION_KEY_KEY in table.uns[TableModel.ATTRS_KEY]
assert table.uns[TableModel.ATTRS_KEY][TableModel.REGION_KEY] == region

obs["A"] = obs["A"].astype(str)
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.warns(UserWarning, match="Converting"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs["A"] = pd.Series(len([chr(ord("a") + i) for i in range(10)]))
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(ValueError, match="Values within"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

@pytest.mark.parametrize("model", [TableModel])
@pytest.mark.parametrize("region", [["sample_1"] * 5 + ["sample_2"] * 5])
def test_table_instance_key_values_not_unique(self, model: TableModel, region: str | np.ndarray):
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { // Auto-enable theater mode on YouTube (function() { function tryTheater() { var btn = document.querySelector('button[aria-label="Theater mode"], ytd-player #player button[title="Theater mode"]'); if (btn && !btn.classList.contains('activated')) { btn.click(); } } // Try immediately tryTheater(); // Try after navigation (SPA) var lastUrl = location.href; setInterval(function() { if (location.href !== lastUrl) { lastUrl = location.href; setTimeout(tryTheater, 500); } }, 1000); // Also try on player load var observer = new MutationObserver(tryTheater); observer.observe(document.body, { childList: true, subtree: true }); })(); } } catch(__e) { console.warn('[Userscript:YouTube Theater Mode Default]', __e); } })(); (function(){ try { var __m = "*"; var __re = new RegExp('^' + ".*" + ' Test joins with string indices and instance id by melonora · Pull Request #485 · scverse/spatialdata · GitHub
Skip to content
Merged
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
9 changes: 8 additions & 1 deletion src/spatialdata/_core/operations/aggregate.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -243,7 +243,14 @@ def _create_sdata_from_table_and_shapes(
) -> SpatialData:
from spatialdata._core._deepcopy import deepcopy as _deepcopy

table.obs[instance_key] = table.obs_names.copy()
shapes_index_dtype = shapes.index.dtype if isinstance(shapes, GeoDataFrame) else shapes.dtype
try:
table.obs[instance_key] = table.obs_names.copy().astype(shapes_index_dtype)
except ValueError as err:
raise TypeError(
f"Instance key column dtype in table resulting from aggregation cannot be cast to the dtype of"
f"element {shapes_name}.index"
) from err
table.obs[region_key] = shapes_name
table = TableModel.parse(table, region=shapes_name, region_key=region_key, instance_key=instance_key)

Expand Down
7 changes: 7 additions & 0 deletions src/spatialdata/_core/spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -199,6 +199,13 @@ def validate_table_in_spatialdata(self, table: AnnData) -> None:
else:
dtype = element.index.dtype
if dtype != table.obs[instance_key].dtype:
if dtype == str or table.obs[instance_key].dtype == str:
raise TypeError(
f"Table instance_key column ({instance_key}) has a dtype "
f"({table.obs[instance_key].dtype}) that does not match the dtype of the indices of "
f"the annotated element ({dtype})."
)

warnings.warn(
(
f"Table instance_key column ({instance_key}) has a dtype "
Expand Down
15 changes: 6 additions & 9 deletions src/spatialdata/models/models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from multiscale_spatial_image.multiscale_spatial_image import MultiscaleSpatialImage
from multiscale_spatial_image.to_multiscale.to_multiscale import Methods
from pandas import CategoricalDtype
from pandas.errors import IntCastingNaNError
from shapely._geometry import GeometryType
from shapely.geometry import MultiPolygon, Point, Polygon
from shapely.geometry.collection import GeometryCollection
Expand DownExpand Up@@ -795,6 +794,11 @@ def _validate_table_annotation_metadata(self, data: AnnData) -> None:
raise ValueError(f"`{attr[self.REGION_KEY_KEY]}` not found in `adata.obs`.")
if attr[self.INSTANCE_KEY] not in data.obs:
raise ValueError(f"`{attr[self.INSTANCE_KEY]}` not found in `adata.obs`.")
if (dtype := data.obs[attr[self.INSTANCE_KEY]].dtype) not in [np.int16, np.int32, np.int64, str]:
raise TypeError(
f"Only np.int16, np.int32, np.int64 or string allowed as dtype for "
f"instance_key column in obs. Dtype found to be {dtype}"
)
expected_regions = attr[self.REGION_KEY] if isinstance(attr[self.REGION_KEY], list) else [attr[self.REGION_KEY]]
found_regions = data.obs[attr[self.REGION_KEY_KEY]].unique().tolist()
if len(set(expected_regions).symmetric_difference(set(found_regions))) > 0:
Expand DownExpand Up@@ -881,14 +885,6 @@ def parse(
adata.obs[region_key] = pd.Categorical(adata.obs[region_key])
if instance_key is None:
raise ValueError("`instance_key` must be provided.")
if adata.obs[instance_key].dtype != int:
try:
warnings.warn(
f"Converting `{cls.INSTANCE_KEY}: {instance_key}` to integer dtype.", UserWarning, stacklevel=2
)
adata.obs[instance_key] = adata.obs[instance_key].astype(int)
except IntCastingNaNError as exc:
raise ValueError("Values within table.obs[] must be able to be coerced to int dtype.") from exc

grouped = adata.obs.groupby(region_key, observed=True)
grouped_size = grouped.size()
Expand All@@ -901,6 +897,7 @@ def parse(

attr = {"region": region, "region_key": region_key, "instance_key": instance_key}
adata.uns[cls.ATTRS_KEY] = attr
cls().validate(adata)
return adata


Expand Down
11 changes: 2 additions & 9 deletions tests/core/operations/test_spatialdata_operations.py
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
from __future__ import annotations

import math
import warnings

import numpy as np
import pytest
Expand DownExpand Up@@ -419,10 +418,7 @@ def test_validate_table_in_spatialdata(full_sdata):
region, region_key, _ = get_table_keys(table)
assert region == "labels2d"

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.labels["labels2d"] = Labels2DModel.parse(full_sdata.labels["labels2d"].astype("int16"))
Expand All@@ -437,10 +433,7 @@ def test_validate_table_in_spatialdata(full_sdata):
table.obs[region_key] = "points_0"
full_sdata.set_table_annotates_spatialelement("table", region="points_0")

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.points["points_0"].index = full_sdata.points["points_0"].index.astype("int16")
Expand Down
31 changes: 31 additions & 0 deletions tests/core/query/test_relational_query.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,6 +22,37 @@ def test_match_table_to_element(sdata_query_aggregation):
# TODO: add tests for labels


def test_join_using_string_instance_id_and_index(sdata_query_aggregation):
sdata_query_aggregation["table"].obs["instance_id"] = [
f"string_{i}" for i in sdata_query_aggregation["table"].obs["instance_id"]
]
sdata_query_aggregation["values_circles"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_circles"].index]
)
sdata_query_aggregation["values_polygons"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_polygons"].index]
)

sdata_query_aggregation["values_polygons"] = sdata_query_aggregation["values_polygons"][:5]
sdata_query_aggregation["values_circles"] = sdata_query_aggregation["values_circles"][:5]

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "inner"
)
# Note that we started with 21 n_obs.
assert table.n_obs == 10

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right_exclusive"
)
assert table.n_obs == 11

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right"
)
assert table.n_obs == 21


def test_left_inner_right_exclusive_join(sdata_query_aggregation):
element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, "values_polygons", "table", "right_exclusive"
Expand Down
18 changes: 8 additions & 10 deletions tests/models/test_models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -318,6 +318,14 @@ def test_table_model(
region: str | np.ndarray,
) -> None:
region_key = "reg"
obs = pd.DataFrame(
RNG.choice(np.arange(0, 100, dtype=float), size=(10, 3), replace=False), columns=["A", "B", "C"]
)
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(TypeError, match="Only np.int16"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs = pd.DataFrame(RNG.choice(np.arange(0, 100), size=(10, 3), replace=False), columns=["A", "B", "C"])
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
Expand All@@ -332,16 +340,6 @@ def test_table_model(
assert TableModel.REGION_KEY_KEY in table.uns[TableModel.ATTRS_KEY]
assert table.uns[TableModel.ATTRS_KEY][TableModel.REGION_KEY] == region

obs["A"] = obs["A"].astype(str)
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.warns(UserWarning, match="Converting"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs["A"] = pd.Series(len([chr(ord("a") + i) for i in range(10)]))
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(ValueError, match="Values within"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

@pytest.mark.parametrize("model", [TableModel])
@pytest.mark.parametrize("region", [["sample_1"] * 5 + ["sample_2"] * 5])
def test_table_instance_key_values_not_unique(self, model: TableModel, region: str | np.ndarray):
Expand Down
, 'i'); if (__m === '*' || __re.test(location.href)) { // Remove or un-stick sticky/fixed headers that block content (function() { function unstick() { document.querySelectorAll('header, nav, [role="banner"], .header, .navbar, .sticky, .fixed-top, [style*="position: fixed"], [style*="position:sticky"]').forEach(function(el) { if (el.style.position === 'fixed' || el.style.position === 'sticky' || getComputedStyle(el).position === 'fixed' || getComputedStyle(el).position === 'sticky') { el.style.position = 'static'; el.style.top = 'auto'; el.style.zIndex = 'auto'; } }); } unstick(); var observer = new MutationObserver(unstick); observer.observe(document.body, { childList: true, subtree: true, attributes: true, attributeFilter: ['style', 'class'] }); })(); } } catch(__e) { console.warn('[Userscript:Kill Sticky Headers]', __e); } })(); })(); Test joins with string indices and instance id by melonora · Pull Request #485 · scverse/spatialdata · GitHub
Skip to content
Merged
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
9 changes: 8 additions & 1 deletion src/spatialdata/_core/operations/aggregate.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -243,7 +243,14 @@ def _create_sdata_from_table_and_shapes(
) -> SpatialData:
from spatialdata._core._deepcopy import deepcopy as _deepcopy

table.obs[instance_key] = table.obs_names.copy()
shapes_index_dtype = shapes.index.dtype if isinstance(shapes, GeoDataFrame) else shapes.dtype
try:
table.obs[instance_key] = table.obs_names.copy().astype(shapes_index_dtype)
except ValueError as err:
raise TypeError(
f"Instance key column dtype in table resulting from aggregation cannot be cast to the dtype of"
f"element {shapes_name}.index"
) from err
table.obs[region_key] = shapes_name
table = TableModel.parse(table, region=shapes_name, region_key=region_key, instance_key=instance_key)

Expand Down
7 changes: 7 additions & 0 deletions src/spatialdata/_core/spatialdata.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -199,6 +199,13 @@ def validate_table_in_spatialdata(self, table: AnnData) -> None:
else:
dtype = element.index.dtype
if dtype != table.obs[instance_key].dtype:
if dtype == str or table.obs[instance_key].dtype == str:
raise TypeError(
f"Table instance_key column ({instance_key}) has a dtype "
f"({table.obs[instance_key].dtype}) that does not match the dtype of the indices of "
f"the annotated element ({dtype})."
)

warnings.warn(
(
f"Table instance_key column ({instance_key}) has a dtype "
Expand Down
15 changes: 6 additions & 9 deletions src/spatialdata/models/models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -20,7 +20,6 @@
from multiscale_spatial_image.multiscale_spatial_image import MultiscaleSpatialImage
from multiscale_spatial_image.to_multiscale.to_multiscale import Methods
from pandas import CategoricalDtype
from pandas.errors import IntCastingNaNError
from shapely._geometry import GeometryType
from shapely.geometry import MultiPolygon, Point, Polygon
from shapely.geometry.collection import GeometryCollection
Expand DownExpand Up@@ -795,6 +794,11 @@ def _validate_table_annotation_metadata(self, data: AnnData) -> None:
raise ValueError(f"`{attr[self.REGION_KEY_KEY]}` not found in `adata.obs`.")
if attr[self.INSTANCE_KEY] not in data.obs:
raise ValueError(f"`{attr[self.INSTANCE_KEY]}` not found in `adata.obs`.")
if (dtype := data.obs[attr[self.INSTANCE_KEY]].dtype) not in [np.int16, np.int32, np.int64, str]:
raise TypeError(
f"Only np.int16, np.int32, np.int64 or string allowed as dtype for "
f"instance_key column in obs. Dtype found to be {dtype}"
)
expected_regions = attr[self.REGION_KEY] if isinstance(attr[self.REGION_KEY], list) else [attr[self.REGION_KEY]]
found_regions = data.obs[attr[self.REGION_KEY_KEY]].unique().tolist()
if len(set(expected_regions).symmetric_difference(set(found_regions))) > 0:
Expand DownExpand Up@@ -881,14 +885,6 @@ def parse(
adata.obs[region_key] = pd.Categorical(adata.obs[region_key])
if instance_key is None:
raise ValueError("`instance_key` must be provided.")
if adata.obs[instance_key].dtype != int:
try:
warnings.warn(
f"Converting `{cls.INSTANCE_KEY}: {instance_key}` to integer dtype.", UserWarning, stacklevel=2
)
adata.obs[instance_key] = adata.obs[instance_key].astype(int)
except IntCastingNaNError as exc:
raise ValueError("Values within table.obs[] must be able to be coerced to int dtype.") from exc

grouped = adata.obs.groupby(region_key, observed=True)
grouped_size = grouped.size()
Expand All@@ -901,6 +897,7 @@ def parse(

attr = {"region": region, "region_key": region_key, "instance_key": instance_key}
adata.uns[cls.ATTRS_KEY] = attr
cls().validate(adata)
return adata


Expand Down
11 changes: 2 additions & 9 deletions tests/core/operations/test_spatialdata_operations.py
Original file line numberDiff line numberDiff line change
@@ -1,7 +1,6 @@
from __future__ import annotations

import math
import warnings

import numpy as np
import pytest
Expand DownExpand Up@@ -419,10 +418,7 @@ def test_validate_table_in_spatialdata(full_sdata):
region, region_key, _ = get_table_keys(table)
assert region == "labels2d"

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.labels["labels2d"] = Labels2DModel.parse(full_sdata.labels["labels2d"].astype("int16"))
Expand All@@ -437,10 +433,7 @@ def test_validate_table_in_spatialdata(full_sdata):
table.obs[region_key] = "points_0"
full_sdata.set_table_annotates_spatialelement("table", region="points_0")

# no warnings
with warnings.catch_warnings():
warnings.simplefilter("error")
full_sdata.validate_table_in_spatialdata(table)
full_sdata.validate_table_in_spatialdata(table)

# dtype mismatch
full_sdata.points["points_0"].index = full_sdata.points["points_0"].index.astype("int16")
Expand Down
31 changes: 31 additions & 0 deletions tests/core/query/test_relational_query.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -22,6 +22,37 @@ def test_match_table_to_element(sdata_query_aggregation):
# TODO: add tests for labels


def test_join_using_string_instance_id_and_index(sdata_query_aggregation):
sdata_query_aggregation["table"].obs["instance_id"] = [
f"string_{i}" for i in sdata_query_aggregation["table"].obs["instance_id"]
]
sdata_query_aggregation["values_circles"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_circles"].index]
)
sdata_query_aggregation["values_polygons"].index = pd.Index(
[f"string_{i}" for i in sdata_query_aggregation["values_polygons"].index]
)

sdata_query_aggregation["values_polygons"] = sdata_query_aggregation["values_polygons"][:5]
sdata_query_aggregation["values_circles"] = sdata_query_aggregation["values_circles"][:5]

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "inner"
)
# Note that we started with 21 n_obs.
assert table.n_obs == 10

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right_exclusive"
)
assert table.n_obs == 11

element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, ["values_circles", "values_polygons"], "table", "right"
)
assert table.n_obs == 21


def test_left_inner_right_exclusive_join(sdata_query_aggregation):
element_dict, table = join_sdata_spatialelement_table(
sdata_query_aggregation, "values_polygons", "table", "right_exclusive"
Expand Down
18 changes: 8 additions & 10 deletions tests/models/test_models.py
Original file line numberDiff line numberDiff line change
Expand Up@@ -318,6 +318,14 @@ def test_table_model(
region: str | np.ndarray,
) -> None:
region_key = "reg"
obs = pd.DataFrame(
RNG.choice(np.arange(0, 100, dtype=float), size=(10, 3), replace=False), columns=["A", "B", "C"]
)
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(TypeError, match="Only np.int16"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs = pd.DataFrame(RNG.choice(np.arange(0, 100), size=(10, 3), replace=False), columns=["A", "B", "C"])
obs[region_key] = region
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
Expand All@@ -332,16 +340,6 @@ def test_table_model(
assert TableModel.REGION_KEY_KEY in table.uns[TableModel.ATTRS_KEY]
assert table.uns[TableModel.ATTRS_KEY][TableModel.REGION_KEY] == region

obs["A"] = obs["A"].astype(str)
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.warns(UserWarning, match="Converting"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

obs["A"] = pd.Series(len([chr(ord("a") + i) for i in range(10)]))
adata = AnnData(RNG.normal(size=(10, 2)), obs=obs)
with pytest.raises(ValueError, match="Values within"):
model.parse(adata, region=region, region_key=region_key, instance_key="A")

@pytest.mark.parametrize("model", [TableModel])
@pytest.mark.parametrize("region", [["sample_1"] * 5 + ["sample_2"] * 5])
def test_table_instance_key_values_not_unique(self, model: TableModel, region: str | np.ndarray):
Expand Down