diff --git a/cortex/dataset/__init__.py b/cortex/dataset/__init__.py index 4f4f2f34..9dc6b576 100644 --- a/cortex/dataset/__init__.py +++ b/cortex/dataset/__init__.py @@ -2,5 +2,5 @@ """ from __future__ import annotations -from .views import Volume, Vertex, VolumeRGB, VertexRGB, Volume2D, Vertex2D, Dataview, Dataview2D, _from_hdf_data, Colors -from .dataset import Dataset, normalize +from .views import Volume, Vertex, VolumeRGB, VertexRGB, Volume2D, Vertex2D, Dataview, Dataview2D, _from_hdf_data, Colors, DataviewJSON +from .dataset import Dataset, normalize, DatasetLike \ No newline at end of file diff --git a/cortex/dataset/braindata.py b/cortex/dataset/braindata.py index 27a92661..4ee81c89 100644 --- a/cortex/dataset/braindata.py +++ b/cortex/dataset/braindata.py @@ -1,6 +1,7 @@ import hashlib import warnings from copy import deepcopy +import os import sys from typing import Generic, Optional, TypeVar, Union, cast if sys.version_info < (3, 11): @@ -33,15 +34,12 @@ def __init__(self, data: Union[npt.NDArray, str], subject: str, **kwargs): nib = nibabel.load(data) data = cast(npt.NDArray, nib.get_fdata().T) self._data = data - try: - basestring - except NameError: - subject = subject if isinstance(subject, str) else subject.decode('utf-8') + subject = subject if isinstance(subject, str) else subject.decode('utf-8') self.subject = subject super().__init__(**kwargs) @property - def data(self): + def data(self) -> npt.NDArray: if isinstance(self._data, h5py.Dataset): return self._data[()] return self._data @@ -51,7 +49,7 @@ def data(self, data: npt.NDArray): self._data = data @property - def name(self): + def name(self) -> str: """Name of this BrainData, computed from hash of data. TODO:WHAT IS THIS USEFUL FOR """ @@ -70,7 +68,7 @@ def uniques(self, collapse=False): def __hash__(self): return hash(_hash(self.data)) - def _write_hdf(self, h5, name=None): + def _write_hdf(self, h5: Union[h5py.File, h5py.Group], name: Optional[str]=None) -> h5py.Dataset: if name is None: name = self.name dgrp = h5.require_group("/data") @@ -107,7 +105,7 @@ def _add_numpy_methods(cls): "__div__", "__pow__", "__neg__", "__abs__"] def make_opfun(op): # function nesting creates closure containing op - def opfun(self, *args): + def opfun(self: Self, *args): return self.copy(getattr(self.data, op)(*args)) return opfun @@ -165,7 +163,7 @@ def to_json(self, simple: bool=False): return sdict @classmethod - def empty(cls, subject: str, xfmname: str, value: float=0, **kwargs): + def empty(cls, subject: str, xfmname: str, value: float=0, **kwargs) -> Self: """ Create a constant-valued VolumeData for the given subject and xfmname. Often useful for testing purposes. @@ -192,7 +190,7 @@ def empty(cls, subject: str, xfmname: str, value: float=0, **kwargs): return cls(np.ones(shape)*value, subject, xfmname, **kwargs) @classmethod - def random(cls, subject: str, xfmname: str, **kwargs): + def random(cls, subject: str, xfmname: str, **kwargs) -> Self: """ Create a random-valued VolumeData for the given subject and xfmname. Random values are from gaussian distribution with mean 0, s.d. 1. @@ -247,7 +245,7 @@ def _check_size(self, mask: Union[npt.NDArray, str, None]) -> None: raise ValueError("Volumetric data (shape %s) is not the same shape as reference for transform (shape %s)" % (str(shape), str(xfm.shape))) self.shape = shape - def map(self, projection="nearest"): + def map(self, projection: str="nearest") -> Self: """Convert this VolumeData into VertexData using the given projection method. @@ -284,11 +282,12 @@ def __repr__(self): maskstr = maskstr[0].upper()+maskstr[1:] return "<%s data for (%s, %s)>"%(maskstr, self.subject, self.xfmname) - def copy(self, data): + def copy(self, data: npt.NDArray) -> Self: return super().copy(data, self.subject, self.xfmname, mask=self._mask) + # TODO: need to include np.ma.MaskedArra in return type? @property - def volume(self): + def volume(self) -> npt.NDArray: """Returns a 3D or 4D volume for this VolumeData, automatically unmasking masked data. """ @@ -303,10 +302,9 @@ def volume(self): return data - def save(self, filename, name=None): + def save(self, filename: Union[str, h5py.Group], name: Optional[str]=None) -> None: """Save the dataset into the hdf file `filename` with the provided name. """ - import os if isinstance(filename, str): fname, ext = os.path.splitext(filename) if ext in (".hdf", ".h5",".hf5"): @@ -335,12 +333,12 @@ def _write_hdf(self, h5: Union[h5py.File, h5py.Group], name: Optional[str]=None) return node - def save_nii(self, filename): + def save_nii(self, filename: os.PathLike) -> None: """Save as a nifti file at the given filename. Nifti headers are copied from the reference image for this VolumeData's transform. """ xfm = db.get_xfm(self.subject, self.xfmname) - affine = xfm.reference.affine + affine = xfm.reference_nifti.affine import nibabel new_nii = nibabel.Nifti1Image(self.volume.T, affine) nibabel.save(new_nii, filename) @@ -468,7 +466,7 @@ def copy(self, data: npt.NDArray) -> Self: """ return super().copy(data, self.subject) - def volume(self, xfmname, projection='nearest', **kwargs): + def volume(self, xfmname: str, projection: str='nearest', **kwargs) -> VolumeData: """ Map this VertexData back to volume space, creating a VolumeData object. This uses the `mapper.backwards` function, which is not particularly @@ -511,8 +509,9 @@ def __getitem__(self, idx): #return VertexData(self.data[idx], self.subject, **self.attrs) return self.copy(self.data[idx]) - - def to_json(self, simple: bool = False): + + # TODO: simple + def to_json(self, simple: bool = False) -> dict[str, list[str]]: if simple: sdict = dict(split=self.llen, frames=self.vertices.shape[0]) sdict.update(super().to_json(simple=simple)) @@ -530,7 +529,7 @@ def vertices(self): return verts @property - def left(self): + def left(self) -> npt.NDArray: """Data for only the left hemisphere vertices. """ if self.movie: @@ -539,7 +538,7 @@ def left(self): return self.data[:self.llen] @property - def right(self): + def right(self) -> npt.NDArray: """Data for only the right hemisphere vertices. """ if self.movie: @@ -547,8 +546,8 @@ def right(self): else: return self.data[self.llen:] - def blend_curvature(self, alpha, threshold=0, brightness=0.5, - contrast=0.25, smooth=20): + def blend_curvature(self, alpha: npt.NDArray[np.floating], threshold: float=0, brightness: float=0.5, + contrast: float=0.25, smooth: float=20): """Blend the data with a curvature map depending on a transparency map. .. deprecated:: @@ -639,9 +638,8 @@ def blend_curvature(self, alpha, threshold=0, brightness=0.5, return blended -def _find_mask(nvox: int, subject: str, xfmname: str): +def _find_mask(nvox: int, subject: str, xfmname: str) -> tuple[str, npt.NDArray[np.bool_]]: import glob - import os import re import nibabel @@ -652,6 +650,7 @@ def _find_mask(nvox: int, subject: str, xfmname: str): if nvox == np.sum(mask): fname = os.path.split(fname)[1] name = re.compile(r'mask_(.+).nii.gz').search(fname) + assert name is not None, f"Mask filename {fname} does not match expected format" return name.group(1), mask raise ValueError('Cannot find a valid mask') @@ -671,12 +670,12 @@ def __getitem__(self, masktype: str) -> T_masker: mask = db.get_mask(self.dv.subject, self.dv.xfmname, masktype) return self.dv.copy(self.dv.volume[:,mask].squeeze()) -def _hash(array): +def _hash(array: npt.ArrayLike) -> str: '''A simple numpy hash function''' array = np.asarray(array) return hashlib.sha1(array.tobytes()).hexdigest() -def _hdf_write(h5, data, name="data", group="/data"): +def _hdf_write(h5: Union[h5py.File, h5py.Group], data: npt.NDArray, name: str="data", group: str="/data") -> h5py.Dataset: try: node = h5.require_dataset("%s/%s"%(group, name), data.shape, data.dtype, exact=True) except TypeError: diff --git a/cortex/dataset/dataset.py b/cortex/dataset/dataset.py index 29a14d8b..9538c684 100644 --- a/cortex/dataset/dataset.py +++ b/cortex/dataset/dataset.py @@ -1,6 +1,7 @@ import tempfile -from typing import Union, overload +from typing import Iterator, Optional, Union, overload, Literal import numpy as np +import numpy.typing as npt import h5py from ..database import db @@ -9,6 +10,7 @@ from .braindata import _hdf_write from .views import normalize as _vnorm from .views import Dataview, Vertex, Volume, _from_hdf_data +from h5py._hl.files import File class Dataset: """ @@ -19,8 +21,8 @@ class Dataset: # TODO: should be BrainData & Dataview, or just Dataview All kwargs should be `BrainData` or `Dataset` objects. """ - def __init__(self, **kwargs: Union[Dataview, dict, str, tuple, "Dataset"]): - self.h5 = None + def __init__(self, **kwargs: Union[Dataview, dict, str, tuple, "Dataset"]) -> None: + self.h5: Optional[h5py.File] = None self.views: dict[str, Dataview] = {} self.append(**kwargs) @@ -41,7 +43,7 @@ def append(self, **kwargs: Union[Dataview, dict, str, tuple, "Dataset"]) -> "Dat return self - def __getattr__(self, attr): + def __getattr__(self, attr: str): if attr in self.__dict__: return self.__dict__[attr] elif attr in self.views: @@ -49,10 +51,10 @@ def __getattr__(self, attr): raise AttributeError - def __getitem__(self, item): + def __getitem__(self, item: str) -> Dataview: return self.views[item] - def __iter__(self): + def __iter__(self) -> Iterator[tuple[str, Dataview]]: for name, dv in sorted(self.views.items(), key=lambda x: x[1].priority): yield name, dv @@ -67,7 +69,7 @@ def __dir__(self): return list(self.__dict__.keys()) + list(self.views.keys()) @classmethod - def from_file(cls, filename, subject=None): + def from_file(cls, filename: str, subject: Optional[str]=None) -> "Dataset": """Load a pycortex Dataset (cortex.Dataset class) from a file Parameters @@ -114,7 +116,7 @@ def from_file(cls, filename, subject=None): return ds - def uniques(self, collapse=False): + def uniques(self, collapse: bool=False) -> set[Dataview]: """Return the set of unique BrainData objects contained by this dataset""" uniques = set() for name, view in self: @@ -124,7 +126,7 @@ def uniques(self, collapse=False): return uniques - def save(self, filename=None, pack=False): + def save(self, filename: Optional[str]=None, pack: bool=False) -> None: if filename is not None: self.h5 = h5py.File(filename, 'a') elif self.h5 is None: @@ -134,9 +136,9 @@ def save(self, filename=None, pack=False): view._write_hdf(self.h5, name=name) if pack: - subjs = set() - xfms = set() - masks = set() + subjs: set[str] = set() + xfms: set[tuple[str, str]] = set() + masks: set[tuple[str, str, str]] = set() for view in self.views.values(): # .uniques() is provided by BrainData for data in view.uniques(): @@ -154,7 +156,24 @@ def save(self, filename=None, pack=False): self.h5.flush() - def get_surf(self, subject, type, hemi='both', merge=False, nudge=False): + # TODO: forcing '*' WILL cause issues. Look for all instances of merge=True ! + @overload + def get_surf(self, subject: str, type: str, hemi: Literal['both']='both', merge: Literal[False]=False, nudge: bool=False) -> tuple[tuple[npt.NDArray[np.floating], npt.NDArray[np.integer]], tuple[npt.NDArray[np.floating], npt.NDArray[np.integer]]]: ... + + @overload + def get_surf(self, subject: str, type: str, hemi: Literal['both']='both', *, merge: Literal[True], nudge: bool=False) -> tuple[npt.NDArray[np.floating], npt.NDArray[np.integer]]: ... + + @overload + def get_surf(self, subject: str, type: str, hemi: Literal['lh', 'rh'], merge: bool=False, nudge: bool=False) -> tuple[npt.NDArray[np.floating], npt.NDArray[np.integer]]: ... + + # Fallthrough case for callers (e.g. Database.get_surf) forwarding a dynamic + # hemi/merge that isn't statically one of the above. + @overload + def get_surf(self, subject: str, type: str, hemi: Literal['both', 'lh', 'rh']='both', merge: bool=False, nudge: bool=False) -> Union[tuple[tuple[npt.NDArray[np.floating], npt.NDArray[np.integer]], tuple[npt.NDArray[np.floating], npt.NDArray[np.integer]]], tuple[npt.NDArray[np.floating], npt.NDArray[np.integer]]]: ... + + def get_surf(self, subject: str, type: str, hemi: Literal['both', 'lh', 'rh']='both', merge: bool=False, nudge: bool=False) -> Union[tuple[tuple[npt.NDArray[np.floating], npt.NDArray[np.integer]], tuple[npt.NDArray[np.floating], npt.NDArray[np.integer]]], tuple[npt.NDArray[np.floating], npt.NDArray[np.integer]]]: + pts: npt.NDArray[np.floating] + polys: npt.NDArray[np.integer] if hemi == 'both': left = self.get_surf(subject, type, "lh", nudge=nudge) right = self.get_surf(subject, type, "rh", nudge=nudge) @@ -181,23 +200,23 @@ def get_surf(self, subject, type, hemi='both', merge=False, nudge=False): except (KeyError, TypeError): raise IOError('Subject not found in package') - def get_xfm(self, subject, xfmname): + def get_xfm(self, subject: str, xfmname: str) -> Transform: try: - group = self.h5['subjects'][subject]['transforms'][xfmname] + group: h5py.Group = self.h5['subjects'][subject]['transforms'][xfmname] return Transform(group['xfm'][:], tuple(group['xfm'].attrs['shape'])) except (KeyError, TypeError): raise IOError('Transform not found in package') - def get_mask(self, subject, xfmname, maskname): + def get_mask(self, subject: str, xfmname: str, maskname: str): try: - group = self.h5['subjects'][subject]['transforms'][xfmname]['masks'] + group: h5py.Group = self.h5['subjects'][subject]['transforms'][xfmname]['masks'] return group[maskname] except (KeyError, TypeError): raise IOError('Mask not found in package') - def get_overlay(self, subject, type='rois', **kwargs): + def get_overlay(self, subject: str, type: str='rois', **kwargs) -> tempfile._TemporaryFileWrapper: try: - group = self.h5['subjects'][subject] + group: h5py.Group = self.h5['subjects'][subject] if type == "rois": tf = tempfile.NamedTemporaryFile() tf.write(group['rois'][0]) @@ -218,6 +237,8 @@ def prepend(self, prefix): return Dataset(**ds) +DatasetLike = Union[Dataset, dict, str] + @overload def normalize(data: Dataview) -> Dataview: ... @@ -227,7 +248,7 @@ def normalize(data: Union[Dataset, dict, str]) -> Dataset: ... @overload def normalize(data: tuple) -> Union[Vertex, Volume]: ... -def normalize(data: Union[Dataset, Dataview, dict, str, tuple]) -> Union[Dataset, Dataview, Vertex, Volume]: +def normalize(data: Union[DatasetLike, Dataview, tuple]) -> Union[Dataset, Dataview, Vertex, Volume]: if isinstance(data, (Dataset, Dataview)): return data elif isinstance(data, dict): @@ -239,7 +260,7 @@ def normalize(data: Union[Dataset, Dataview, dict, str, tuple]) -> Union[Dataset raise TypeError('Unknown input type') -def _pack_subjs(h5, subjects): +def _pack_subjs(h5: File, subjects: set[str]) -> None: for subject in subjects: rois = db.get_overlay(subject, modify_svg_file=False) rnode = h5.require_dataset("/subjects/%s/rois"%subject, (1,), @@ -254,14 +275,14 @@ def _pack_subjs(h5, subjects): _hdf_write(h5, pts, "pts", group) _hdf_write(h5, polys, "polys", group) -def _pack_xfms(h5, xfms): +def _pack_xfms(h5: File, xfms: set[tuple[str, str]]) -> None: for subj, xfmname in xfms: xfm = db.get_xfm(subj, xfmname, 'coord') group = "/subjects/%s/transforms/%s"%(subj, xfmname) node = _hdf_write(h5, np.array(xfm.xfm), "xfm", group) node.attrs['shape'] = xfm.shape -def _pack_masks(h5, masks): +def _pack_masks(h5: File, masks: set[tuple[str, str, str]]) -> None: for subj, xfm, maskname in masks: mask = db.get_mask(subj, xfm, maskname) group = "/subjects/%s/transforms/%s/masks"%(subj, xfm) diff --git a/cortex/dataset/view2D.py b/cortex/dataset/view2D.py index 922810af..089deff9 100644 --- a/cortex/dataset/view2D.py +++ b/cortex/dataset/view2D.py @@ -183,7 +183,7 @@ def _write_hdf(self, h5, name="data"): return viewnode @property - def raw(self): + def raw(self) -> VolumeRGB: """VolumeRGB object containing the colormapped data from this object. """ if self.dim1.xfmname != self.dim2.xfmname: @@ -275,7 +275,7 @@ def __repr__(self): return "<2D vertex data for (%s)>"%self.dim1.subject @property - def raw(self): + def raw(self) -> VertexRGB: """VertexRGB object containing the colormapped data from this object. """ r, g, b, a = self._to_raw(self.dim1.data, self.dim2.data) diff --git a/cortex/dataset/viewRGB.py b/cortex/dataset/viewRGB.py index 0c1309f9..8d2c9d4e 100644 --- a/cortex/dataset/viewRGB.py +++ b/cortex/dataset/viewRGB.py @@ -1,7 +1,7 @@ from __future__ import annotations import colorsys -from typing import Optional, TypeVar, Union +from typing import Literal, Optional, TypeVar, Union, cast import warnings import numpy as np @@ -112,7 +112,7 @@ def uniques(self, collapse=False): if self.alpha is not None: yield self.alpha - def _apply_nan_mask(self, alpha): + def _apply_nan_mask(self, alpha: BrainData): """Apply stored NaN mask to alpha, enforcing transparency for NaN positions even when the user overrides the alpha channel. uint8 RGB channels cannot hold NaN, so the mask is captured before conversion @@ -162,19 +162,19 @@ def get_cmapdict(self): @staticmethod def color_voxels( - channel1, - channel2, - channel3, - channel1color, - channel2color, - channel3Color, - value_max, - saturation_max, - vmin, - vmax, - autorange, - alpha=None, - ): + channel1: Union[npt.NDArray, VolumeData, VertexData], + channel2: Union[npt.NDArray, VolumeData, VertexData], + channel3: Union[npt.NDArray, VolumeData, VertexData], + channel1color: Color[int], + channel2color: Color[int], + channel3Color: Color[int], + value_max: Optional[float], + saturation_max: float, + vmin: Optional[Union[float, tuple[float, float, float]]], + vmax: Optional[Union[float, tuple[float, float, float]]], + autorange: Literal['shared', 'individual'] = 'individual', + alpha: Optional[Union[npt.NDArray, VolumeData, VertexData]] = None, + ) -> tuple[npt.NDArray[np.uint8], npt.NDArray[np.uint8], npt.NDArray[np.uint8], npt.NDArray[np.uint8]]: """ Colors voxels in 3 color dimensions but not necessarily canonical red, green, and blue Parameters @@ -340,7 +340,7 @@ def color_voxels( saturation = 1.0 if value > 1: value = 1.0 - this_color = HSV2RGB([hue, saturation, value]) + this_color = HSV2RGB((hue, saturation, value)) red.flat[i] = this_color[0] green.flat[i] = this_color[1] blue.flat[i] = this_color[2] @@ -348,7 +348,7 @@ def color_voxels( # Now make an alpha volume if alpha is None: alpha = np.ones_like(red, np.uint8) * 255 - alpha[mask] = 0 + alpha[mask] = 0 # TODO: this seems like an actual issue return red, green, blue, alpha @@ -440,14 +440,14 @@ def __init__( alpha: Optional[Union[npt.NDArray, Volume]] = None, description: str = "", state=None, - channel1color: Color = Colors.Red, - channel2color: Color = Colors.Green, - channel3color: Color = Colors.Blue, + channel1color: Color[int] = Colors.Red, + channel2color: Color[int] = Colors.Green, + channel3color: Color[int] = Colors.Blue, max_color_value: Optional[float] = None, max_color_saturation: float = 1.0, vmin: Optional[Union[float, tuple]] = None, vmax: Optional[Union[float, tuple]] = None, - autorange: str = "individual", + autorange: Literal['shared', 'individual'] = "individual", priority: int = 1, ): channel1color = tuple(channel1color) @@ -563,7 +563,7 @@ def __init__( ) @property - def alpha(self): + def alpha(self) -> Volume: """Compute alpha transparency""" alpha = self._alpha if alpha is None: @@ -606,7 +606,7 @@ def to_json(self, simple=False): return sdict @property - def volume(self): + def volume(self) -> np.ndarray[tuple[int, int, int, int, int], np.dtype[np.uint8]]: """5-dimensional volume (t, z, y, x, rgba) with data that has been mapped into 8-bit unsigned integers that correspond to colors. """ @@ -650,7 +650,7 @@ def _write_hdf(self, h5, name="data"): return super()._write_hdf(h5, name=name, xfmname=[self.xfmname]) @property - def raw(self): + def raw(self) -> VolumeRGB: return self @@ -738,15 +738,15 @@ def __init__( alpha: Optional[Union[npt.NDArray, Vertex]] = None, description: str = "", state=None, - channel1color=Colors.Red, - channel2color=Colors.Green, - channel3color=Colors.Blue, - max_color_value=None, - max_color_saturation=1.0, - vmin=None, - vmax=None, - autorange="individual", - priority=1, + channel1color: Color[int] = Colors.Red, + channel2color: Color[int] = Colors.Green, + channel3color: Color[int] = Colors.Blue, + max_color_value: Optional[float] = None, + max_color_saturation: float = 1.0, + vmin: Optional[Union[float, tuple[float, float, float]]] = None, + vmax: Optional[Union[float, tuple[float, float, float]]] = None, + autorange: Literal['shared', 'individual'] = "individual", + priority: int = 1, ): channel1color = tuple(channel1color) channel2color = tuple(channel2color) @@ -839,7 +839,7 @@ def __init__( ) @property - def alpha(self): + def alpha(self) -> Vertex: """Compute alpha transparency""" alpha = self._alpha if alpha is None: @@ -867,14 +867,14 @@ def alpha(self, alpha: Optional[Union[npt.NDArray, Vertex]]): self._alpha = alpha @property - def vertices(self): + def vertices(self) -> npt.NDArray[np.uint8]: """3-dimensional volume (t, v, rgba) with data that has been mapped into 8-bit unsigned integers that correspond to colors. """ verts = [] for dv in (self.red, self.green, self.blue, self.alpha): if dv.vertices.dtype != np.uint8: - vert = dv.vertices.astype("float32", copy=True) + vert = dv.vertices.astype(np.float32, copy=True) if dv.vmin is None: if vert.min() < 0: vert -= vert.min() @@ -902,11 +902,11 @@ def to_json(self, simple=False): return sdict @property - def left(self): + def left(self) -> npt.NDArray[np.uint8]: return self.vertices[:, : self.red.llen] @property - def right(self): + def right(self) -> npt.NDArray[np.uint8]: return self.vertices[:, self.red.llen :] def __repr__(self): @@ -920,5 +920,5 @@ def name(self): return "__%s" % _hash(self.vertices)[:16] @property - def raw(self): + def raw(self) -> VertexRGB: return self diff --git a/cortex/dataset/views.py b/cortex/dataset/views.py index a42e2e5f..018d7a1d 100644 --- a/cortex/dataset/views.py +++ b/cortex/dataset/views.py @@ -3,9 +3,15 @@ import glob import json import os -from typing import Any, Optional, Union, cast, overload, Literal +import sys +from typing import Any, Optional, TypedDict, Union, cast, overload, Literal +if sys.version_info < (3, 11): + from typing_extensions import NotRequired +else: + from typing import NotRequired import h5py +from matplotlib.colors import Colormap import numpy as np import numpy.typing as npt @@ -24,8 +30,11 @@ def register_cmap(cmap): from matplotlib.cm import register_cmap +JSON = Union[dict[str, "JSON"], list["JSON"], str, int, float, bool, None] + + @overload -def normalize(data: tuple[Any, Any, Any]) -> Volume: ... +def normalize(data: tuple[Any, Any, Any]) -> Union[Volume, VolumeRGB]: ... @overload @@ -36,7 +45,7 @@ def normalize(data: tuple[Any, Any]) -> Vertex: ... def normalize(data: Dataview) -> Dataview: ... -def normalize(data: Union[Dataview, tuple]) -> Union[Volume, Vertex, Dataview]: +def normalize(data: Union[Dataview, tuple]) -> Union[Volume, VolumeRGB, Vertex, Dataview]: if isinstance(data, tuple): if len(data) == 3: if data[0].dtype == np.uint8: @@ -154,7 +163,28 @@ def _from_hdf_view( raise ValueError("Invalid Dataview specification") +class ColormapDict(TypedDict): + cmap: Colormap + vmin: Optional[float] + vmax: Optional[float] + + +class DataviewJSON(TypedDict): + state: Any + attrs: dict[str, Any] + desc: str + cmap: Optional[list[str]] + vmin: Optional[list[float]] + vmax: Optional[list[float]] + name: NotRequired[str] + raw: NotRequired[bool] + mosaic: NotRequired[tuple[int, int]] + subject: NotRequired[str] # is this actually from BrainData? + + class Dataview: + _nan_mask: Optional[npt.NDArray[np.bool_]] + def __init__( self, cmap: Optional[str] = None, @@ -196,12 +226,12 @@ def priority(self): def priority(self, value): self.attrs["priority"] = value - def to_json(self, simple=False): + def to_json(self, simple: bool=False) -> DataviewJSON: if simple: return dict() desc = self.description - if hasattr(desc, "decode"): + if isinstance(desc, bytes): desc = desc.decode() sdict = dict(state=self.state, attrs=self.attrs.copy(), desc=desc) try: @@ -287,7 +317,7 @@ def _write_hdf(self, h5, name="data", data=None, xfmname=None): view[7] = json.dumps(xfmname) return view - def get_cmapdict(self): + def get_cmapdict(self) -> ColormapDict: """Returns a dictionary with cmap information.""" from matplotlib import colors @@ -312,11 +342,10 @@ def get_cmapdict(self): # Register colormap to matplotlib to avoid loading it again register_cmap(cmap) - # TODO: create namedtuple - return dict(cmap=cmap, vmin=self.vmin, vmax=self.vmax) + return ColormapDict(cmap=cmap, vmin=self.vmin, vmax=self.vmax) @property - def raw(self): + def raw(self) -> tuple[npt.NDArray[np.uint8], npt.NDArray[np.bool_]]: from matplotlib import cm, colors cmap = self.get_cmapdict()["cmap"] @@ -324,7 +353,7 @@ def raw(self): norm = colors.Normalize(self.vmin, self.vmax) cmapper = cm.ScalarMappable(norm=norm, cmap=cmap) # Capture NaN mask before uint8 conversion (NaN info is lost after) - nan_mask = np.isnan(self.data) + nan_mask: npt.NDArray[np.bool_] = np.isnan(self.data) # TODO: self.data relies on BrainData. Would need common inheritance for this to work. color_data = cmapper.to_rgba(self.data.flatten()).reshape( self.data.shape + (4,) diff --git a/cortex/tests/test_dataset.py b/cortex/tests/test_dataset.py index d3224bd4..25544019 100644 --- a/cortex/tests/test_dataset.py +++ b/cortex/tests/test_dataset.py @@ -110,11 +110,11 @@ def test_rgb_rejects_unknown_kwargs(): """VolumeRGB and VertexRGB should reject unknown keyword arguments.""" red, green, blue = [np.random.randn(nverts) for _ in range(3)] with pytest.raises(TypeError): - dataset.VertexRGB(red, green, blue, subj, bogus_kwarg=True) + dataset.VertexRGB(red, green, blue, subj, bogus_kwarg=True) # type: ignore red, green, blue = [np.random.randn(*volshape) for _ in range(3)] with pytest.raises(TypeError): - dataset.VolumeRGB(red, green, blue, subj, xfmname, bogus_kwarg=True) + dataset.VolumeRGB(red, green, blue, subj, xfmname, bogus_kwarg=True) # type: ignore def test_volumergb_shared_range(): @@ -314,7 +314,7 @@ def test_blend_curvature(): # blend_curvature is deprecated; the warning should fire on every call. with pytest.warns(DeprecationWarning, match="blend_curvature is deprecated"): - view_rgb = view.blend_curvature(alpha) + view_rgb: cortex.VertexRGB = view.blend_curvature(alpha) with pytest.warns(DeprecationWarning): view_rgb = view.blend_curvature(alpha > 0.3) # test that it returns a VertexRGB