code-loader 1.0.208.dev8__py3-none-any.whl → 1.0.210__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -8,7 +8,7 @@ import numpy.typing as npt
8
8
  from code_loader.contract.enums import DataStateType, DataStateEnum, LeapDataType, ConfusionMatrixValue, \
9
9
  MetricDirection, DatasetMetadataType, LatentSpaceReduction, CustomLatentSpaceComputedAt
10
10
  from code_loader.contract.visualizer_classes import LeapImage, LeapText, LeapGraph, LeapHorizontalBar, \
11
- LeapTextMask, LeapImageMask, LeapImageWithBBox, LeapImageWithHeatmap, LeapVideo, LeapAudio
11
+ LeapTextMask, LeapImageMask, LeapImageWithBBox, LeapImageWithHeatmap, LeapVideo, LeapAudio, LeapVolume
12
12
  from code_loader.contract.sim_config import SimConfig
13
13
 
14
14
  custom_latent_space_attribute = "custom_latent_space"
@@ -213,7 +213,7 @@ VisualizerCallableInterface = Union[
213
213
  ]
214
214
 
215
215
  LeapData = Union[LeapImage, LeapText, LeapGraph, LeapHorizontalBar, LeapImageMask, LeapTextMask, LeapImageWithBBox,
216
- LeapImageWithHeatmap, LeapVideo, LeapAudio]
216
+ LeapImageWithHeatmap, LeapVideo, LeapAudio, LeapVolume]
217
217
 
218
218
  CustomCallableInterface = Callable[..., Any]
219
219
 
@@ -29,6 +29,7 @@ class LeapDataType(Enum):
29
29
  ImageWithHeatmap = 'ImageWithHeatmap'
30
30
  Video = 'Video'
31
31
  Audio = 'Audio'
32
+ Volume = 'Volume'
32
33
 
33
34
 
34
35
  class MetricDirection(Enum):
@@ -377,6 +377,60 @@ class LeapAudio:
377
377
  raise LeapValidationError(f'sample_rate must be a positive int, got {self.sample_rate}')
378
378
 
379
379
 
380
+ @dataclass
381
+ class LeapVolume:
382
+ """
383
+ Visualizer representing a volume (3-D scalar field) for Tensorleap: a CT/MRI scan, a microscopy
384
+ z-stack, or any (D, H, W) grid. The UI slices it along any axis; the engine stores it as uint8.
385
+
386
+ Attributes:
387
+ data (npt.NDArray[np.float32] | npt.NDArray[np.uint8]): The volume, shaped [D, H, W]. Any finite value
388
+ range is accepted: the engine min-max quantizes it to uint8 for display and keeps the original
389
+ [min, max] so the UI can show true values.
390
+ mask (Optional[npt.NDArray[np.uint8]]): Optional label volume shaped like `data`; voxel value v is the
391
+ index into `labels`.
392
+ labels (Optional[List[str]]): Names for mask values; required with `mask`, len(labels) > mask.max().
393
+ spacing (Optional[Tuple[float, float, float]]): Physical voxel size per axis (D, H, W), e.g. mm, so the
394
+ UI renders anisotropic voxels with the right aspect. None means (1, 1, 1).
395
+ type (LeapDataType): The data type, default is LeapDataType.Volume.
396
+
397
+ Example:
398
+ volume = np.random.rand(64, 128, 128).astype(np.float32)
399
+ mask = (volume > 0.8).astype(np.uint8)
400
+ leap_volume = LeapVolume(data=volume, mask=mask, labels=["background", "lesion"], spacing=(2.0, 0.7, 0.7))
401
+ """
402
+ data: Union[npt.NDArray[np.float32], npt.NDArray[np.uint8]]
403
+ mask: Optional[npt.NDArray[np.uint8]] = None
404
+ labels: Optional[List[str]] = None
405
+ spacing: Optional[Tuple[float, float, float]] = None
406
+ type: LeapDataType = LeapDataType.Volume
407
+
408
+ def __post_init__(self) -> None:
409
+ validate_type(self.type, LeapDataType.Volume)
410
+ validate_type(type(self.data), np.ndarray)
411
+ validate_type(self.data.dtype, [np.uint8, np.float32])
412
+ validate_type(len(self.data.shape), 3, 'Volume data must be of shape 3 [D, H, W]')
413
+ if not np.isfinite(self.data).all():
414
+ raise LeapValidationError('Volume data must be finite (no NaN/Inf)')
415
+ if (self.mask is None) != (self.labels is None):
416
+ raise LeapValidationError('Volume mask and labels must be given together')
417
+ if self.mask is not None:
418
+ validate_type(type(self.mask), np.ndarray)
419
+ validate_type(self.mask.dtype, np.uint8)
420
+ if self.mask.shape != self.data.shape:
421
+ raise LeapValidationError(
422
+ f'Volume mask shape {self.mask.shape} must equal data shape {self.data.shape}')
423
+ validate_type(type(self.labels), list)
424
+ for label in self.labels:
425
+ validate_type(type(label), str)
426
+ if self.mask.size and int(self.mask.max()) >= len(self.labels):
427
+ raise LeapValidationError(
428
+ f'Volume mask has value {int(self.mask.max())} but only {len(self.labels)} labels')
429
+ if self.spacing is not None:
430
+ if len(self.spacing) != 3 or any(float(s) <= 0 for s in self.spacing):
431
+ raise LeapValidationError(f'Volume spacing must be 3 positive numbers, got {self.spacing}')
432
+
433
+
380
434
  map_leap_data_type_to_visualizer_class = {
381
435
  LeapDataType.Image.value: LeapImage,
382
436
  LeapDataType.Graph.value: LeapGraph,
@@ -387,5 +441,6 @@ map_leap_data_type_to_visualizer_class = {
387
441
  LeapDataType.TextMask.value: LeapTextMask,
388
442
  LeapDataType.ImageWithBBox.value: LeapImageWithBBox,
389
443
  LeapDataType.ImageWithHeatmap.value: LeapImageWithHeatmap,
390
- LeapDataType.Audio.value: LeapAudio
444
+ LeapDataType.Audio.value: LeapAudio,
445
+ LeapDataType.Volume.value: LeapVolume,
391
446
  }
@@ -111,6 +111,10 @@ class LeapBinder:
111
111
 
112
112
  self.mapping_connections: List[NodeConnection] = []
113
113
  self.integration_test_func: Optional[Callable[[str, PreprocessResponse], Any]] = None
114
+ # Set by LeapLoader.exec_script when running the dataset script failed, because that
115
+ # also resets setup_container: an empty container after a failure means "the script
116
+ # crashed", not "the user forgot to register". See get_preprocess_result.
117
+ self.previous_script_failure: Optional[BaseException] = None
114
118
 
115
119
  self.batch_size_to_validate: Optional[int] = None
116
120
  self.leap_analysis_configuration = LeapAnalysisConfiguration()
@@ -803,6 +807,13 @@ class LeapBinder:
803
807
  def get_preprocess_result(self) -> Dict[DataStateEnum, PreprocessResponse]:
804
808
  preprocess = self.setup_container.preprocess
805
809
  if preprocess is None:
810
+ if self.previous_script_failure is not None:
811
+ # Running the dataset script already failed once, and that reset the
812
+ # registrations (LeapLoader.exec_script). Blaming a missing
813
+ # set_preprocess call here would hide the actual crash.
814
+ raise Exception(
815
+ f"The dataset script failed before registering its handlers: "
816
+ f"{self.previous_script_failure}") from self.previous_script_failure
806
817
  raise Exception("Please make sure you call the leap_binder.set_preprocess method")
807
818
  preprocess_results = preprocess.function()
808
819
  preprocess_result_dict = {}
@@ -34,7 +34,7 @@ from code_loader.contract.enums import MetricDirection, LeapDataType, DatasetMet
34
34
  from code_loader import leap_binder, LeapLoader
35
35
  from code_loader.contract.mapping import NodeMapping, NodeMappingType, NodeConnection
36
36
  from code_loader.contract.visualizer_classes import LeapImage, LeapImageMask, LeapTextMask, LeapText, LeapGraph, \
37
- LeapHorizontalBar, LeapImageWithBBox, LeapImageWithHeatmap, LeapVideo, LeapAudio, LeapValidationError, \
37
+ LeapHorizontalBar, LeapImageWithBBox, LeapImageWithHeatmap, LeapVideo, LeapAudio, LeapVolume, LeapValidationError, \
38
38
  map_leap_data_type_to_visualizer_class
39
39
  from code_loader.inner_leap_binder.leapbinder import mapping_runtime_mode_env_var_mame, \
40
40
  _reject_stringized_sample_preprocess_response
@@ -1593,7 +1593,8 @@ def tensorleap_custom_visualizer(name: str, visualizer_type: LeapDataType,
1593
1593
  LeapDataType.ImageWithBBox: LeapImageWithBBox,
1594
1594
  LeapDataType.ImageWithHeatmap: LeapImageWithHeatmap,
1595
1595
  LeapDataType.Video: LeapVideo,
1596
- LeapDataType.Audio: LeapAudio
1596
+ LeapDataType.Audio: LeapAudio,
1597
+ LeapDataType.Volume: LeapVolume,
1597
1598
  }
1598
1599
  validate_output_structure(result, func_name=user_function.__name__,
1599
1600
  expected_type_name=result_type_map[visualizer_type])
@@ -1762,7 +1763,16 @@ def _classify_custom_latent_space_signature(user_function) -> str:
1762
1763
  """Dataset-computed when a parameter is typed PreprocessResponse (a subclass or an
1763
1764
  Optional[...] of it included), model-computed otherwise."""
1764
1765
  params = list(inspect.signature(user_function).parameters.values())
1765
- hints = typing.get_type_hints(user_function)
1766
+ try:
1767
+ hints = typing.get_type_hints(user_function)
1768
+ except NameError as e:
1769
+ raise Exception(
1770
+ f"tensorleap_custom_latent_space validation failed: could not resolve the type "
1771
+ f"annotations of '{user_function.__name__}' ({e}). Tensorleap reads them to tell a "
1772
+ f"dataset-computed latent space from a model-computed one, so every annotation, the "
1773
+ f"return type included, must be resolvable when the function is decorated. Import the "
1774
+ f"annotated types at module level (not only under TYPE_CHECKING) or remove the "
1775
+ f"annotation.") from e
1766
1776
  preprocess_params = [p.name for p in params if _is_preprocess_response_type(hints.get(p.name))]
1767
1777
  if preprocess_params:
1768
1778
  positional = [p.name for p in params
@@ -1789,18 +1799,29 @@ def _classify_custom_latent_space_signature(user_function) -> str:
1789
1799
  f"If you are upgrading an existing project, this signature used to be accepted "
1790
1800
  f"unannotated as dataset-computed; add ': PreprocessResponse' to '{second}' to keep "
1791
1801
  f"the previous behavior.")
1802
+ if len(params) > 2 and not any(p.name in hints for p in params):
1803
+ names = [p.name for p in params]
1804
+ raise Exception(
1805
+ f"tensorleap_custom_latent_space validation failed: '{user_function.__name__}' has "
1806
+ f"parameters {names} and none of them has a type annotation, so Tensorleap cannot "
1807
+ f"tell whether this is a dataset-computed latent space (one sample at a time) or a "
1808
+ f"model-computed one (a batch of model tensors). Please annotate them:\n"
1809
+ f" dataset-computed: def {user_function.__name__}({names[0]}, {names[1]}: PreprocessResponse, ...) -> (d,)\n"
1810
+ f" model-computed: def {user_function.__name__}({names[0]}: np.ndarray, {names[1]}: np.ndarray, ...) -> (batch, d)")
1792
1811
  return 'model'
1793
1812
 
1794
1813
 
1795
1814
  def _model_latent_space_arg_names(user_function) -> List[str]:
1796
- argspec = inspect.getfullargspec(user_function)
1797
- arg_names = list(argspec.args)
1815
+ # inspect.signature follows functools.wraps' __wrapped__, as the classifier does.
1816
+ params = inspect.signature(user_function).parameters.values()
1817
+ arg_names = [p.name for p in params
1818
+ if p.kind in (inspect.Parameter.POSITIONAL_ONLY, inspect.Parameter.POSITIONAL_OR_KEYWORD)]
1798
1819
  spr_count = 0
1799
- for arg_name, arg_type in argspec.annotations.items():
1800
- if arg_name == 'return':
1820
+ for param in params:
1821
+ if param.annotation is inspect.Parameter.empty:
1801
1822
  continue
1802
- _reject_stringized_sample_preprocess_response(user_function, arg_name, arg_type)
1803
- if arg_type == SamplePreprocessResponse:
1823
+ _reject_stringized_sample_preprocess_response(user_function, param.name, param.annotation)
1824
+ if param.annotation == SamplePreprocessResponse:
1804
1825
  spr_count += 1
1805
1826
  if spr_count > 1:
1806
1827
  raise Exception(
@@ -1882,7 +1903,7 @@ def _check_custom_latent_space_values(result: np.ndarray, ls_name: str, has_batc
1882
1903
  key=("tensorleap_custom_latent_space_nonfinite", ls_name),
1883
1904
  message=(
1884
1905
  f"Custom latent space '{ls_name}' returned NaN or inf. Those samples are left out of "
1885
- f"this latent space; if they exceed 5% of the evaluated samples the latent space is "
1906
+ f"this latent space; if they exceed 15% of the evaluated samples the latent space is "
1886
1907
  f"dropped."))
1887
1908
  peak = 0.0
1888
1909
  if result.size and result.dtype.kind == 'f':
code_loader/leaploader.py CHANGED
@@ -1,4 +1,5 @@
1
1
  # mypy: ignore-errors
2
+ import copy
2
3
  import importlib.util
3
4
  import inspect
4
5
  import io
@@ -7,6 +8,7 @@ import sys
7
8
  from contextlib import redirect_stdout
8
9
  from functools import lru_cache
9
10
  from pathlib import Path
11
+ from types import TracebackType
10
12
  from typing import Dict, List, Iterable, Set, FrozenSet, Union, Any, Type, Optional, Callable, Tuple
11
13
 
12
14
  import numpy as np
@@ -33,6 +35,44 @@ from code_loader.utils import get_root_exception_file_and_line_number, get_metad
33
35
  validate_autoregressive_state_types, autoregressive_nests_equal, is_absent_metadata_value, \
34
36
  sample_preprocess_response_arg_name
35
37
 
38
+ _code_loader_dir = str(Path(__file__).parent)
39
+
40
+
41
+ def _innermost_frame_outside_code_loader(e: BaseException) -> Optional[TracebackType]:
42
+ """The deepest traceback frame that is not code_loader's own code, if any.
43
+
44
+ That is the line the user can go and look at. The deepest frame overall is often
45
+ ours — a validation that raises, or a builtin we wrapped — and "raised at
46
+ leapbinder.py" tells the user nothing about their script.
47
+ """
48
+ frame_in_user_code = None
49
+ _traceback = e.__traceback__
50
+ while _traceback is not None:
51
+ if not _traceback.tb_frame.f_code.co_filename.startswith(_code_loader_dir):
52
+ frame_in_user_code = _traceback
53
+ _traceback = _traceback.tb_next
54
+ return frame_in_user_code
55
+
56
+
57
+ def describe_script_exception(e: BaseException) -> str:
58
+ """One line describing a failure raised by user code: its type, its message and where it happened.
59
+
60
+ ``repr(e)`` on its own throws away what the user needs most: ``repr`` of an OSError
61
+ drops the filename (``FileNotFoundError(2, 'No such file or directory')`` — but not
62
+ *which* file), and no ``repr`` says which line of the integration script raised. Both
63
+ are what turns "the dataset script crashed" into an error the user can act on.
64
+ """
65
+ message = getattr(e, 'message', None)
66
+ if not isinstance(message, str) or not message:
67
+ message = str(e)
68
+ description = f'{type(e).__name__}: {message}' if message else repr(e)
69
+
70
+ frame_in_user_code = _innermost_frame_outside_code_loader(e)
71
+ if frame_in_user_code is not None:
72
+ file_name = Path(frame_in_user_code.tb_frame.f_code.co_filename).name
73
+ description = f'{description} (raised at {file_name}, line {frame_in_user_code.tb_lineno})'
74
+ return description
75
+
36
76
 
37
77
  def _serialize_sim_bounds(bounds) -> dict:
38
78
  if isinstance(bounds, (FloatBounds, IntBounds)):
@@ -55,6 +95,8 @@ class LeapLoader(LeapLoaderBase):
55
95
  self._preprocess_result_cached = None
56
96
  self._synthetic_lookup: Dict[str, Tuple[PreprocessResponse, Any]] = {}
57
97
  self._synthetic_populator: Optional[Callable[[str], None]] = None
98
+ # The first exec_script failure, re-raised by every later call. See exec_script.
99
+ self._exec_script_error: Optional[BaseException] = None
58
100
 
59
101
  try:
60
102
  from code_loader.mixpanel_tracker import track_code_loader_loaded
@@ -66,9 +108,33 @@ class LeapLoader(LeapLoaderBase):
66
108
  except Exception:
67
109
  pass
68
110
 
111
+ def _remember_exec_script_error(self, error: BaseException) -> BaseException:
112
+ # Also recorded on the binder: the failure handlers below wipe its setup_container,
113
+ # so anything that reaches the binder afterwards (another LeapLoader in the same
114
+ # process, say) can report the real failure instead of an empty container.
115
+ global_leap_binder.previous_script_failure = error
116
+ self._exec_script_error = error
117
+ return error
118
+
69
119
  @lru_cache()
70
120
  def exec_script(self) -> None:
71
121
  from code_loader.inner_leap_binder import leapbinder_decorators as _leap_dec
122
+
123
+ # A failed exec_script is sticky. lru_cache does not cache exceptions, so without
124
+ # this every later caller re-ran the script — and the re-run reports the WRONG
125
+ # error: the handlers below reset global_leap_binder.setup_container, while
126
+ # re-importing the entry file does not re-register handlers that live in modules
127
+ # already in sys.modules. The second run therefore raised "Please make sure you
128
+ # call the leap_binder.set_preprocess method", burying the real cause (a missing
129
+ # file, a bad path, a typo'd import). Callers do retry — the engine's samples
130
+ # generator calls this once per loop iteration — so the real error was logged once
131
+ # and then drowned in generic repeats.
132
+ if self._exec_script_error is not None:
133
+ # A fresh copy each time: re-raising the stored object would append this call's
134
+ # frames to its __traceback__, so a caller retrying in a loop grows it (and keeps
135
+ # every retry's frames alive) without bound.
136
+ error = self._exec_script_error
137
+ raise copy.copy(error) from error.__cause__
72
138
  try:
73
139
  os.environ[mapping_runtime_mode_env_var_mame] = 'TRUE'
74
140
  self.evaluate_module()
@@ -98,16 +164,19 @@ class LeapLoader(LeapLoaderBase):
98
164
  if is_grouped else
99
165
  PreprocessResponse(state=DataStateType.training, length=0))
100
166
  global_leap_binder.integration_test_func(None, mapping_preprocess)
167
+ global_leap_binder.previous_script_failure = None
101
168
  except TypeError as e:
102
169
  import traceback
103
170
  global_leap_binder.setup_container = DatasetIntegrationSetup()
104
171
  if "leap_binder.set_metadata(" in traceback.format_exc(5):
105
- raise DeprecationWarning(
106
- "Please remove the metadata_type on leap_binder.set_metadata in your dataset script")
107
- raise DatasetScriptException(getattr(e, 'message', repr(e))) from e
172
+ raise self._remember_exec_script_error(DeprecationWarning(
173
+ "Please remove the metadata_type on leap_binder.set_metadata in your dataset script"))
174
+ raise self._remember_exec_script_error(
175
+ DatasetScriptException(describe_script_exception(e))) from e
108
176
  except Exception as e:
109
177
  global_leap_binder.setup_container = DatasetIntegrationSetup()
110
- raise DatasetScriptException(getattr(e, 'message', repr(e))) from e
178
+ raise self._remember_exec_script_error(
179
+ DatasetScriptException(describe_script_exception(e))) from e
111
180
  finally:
112
181
  # ensure that the environment variable is removed after the script execution
113
182
  _leap_dec._mapping_dataset_is_grouped = False
@@ -1513,8 +1582,9 @@ class LeapLoader(LeapLoaderBase):
1513
1582
  # Preprocess runs only when the function asks for a SamplePreprocessResponse; the metrics
1514
1583
  # pod that calls this has no other reason to pay for it.
1515
1584
  if preprocess_response_arg_name is not None:
1516
- input_tensors_by_arg_name[preprocess_response_arg_name] = SamplePreprocessResponse(
1517
- sample_ids, self._preprocess_result()[state])
1585
+ input_tensors_by_arg_name = {
1586
+ **input_tensors_by_arg_name,
1587
+ preprocess_response_arg_name: SamplePreprocessResponse(sample_ids, self._preprocess_result()[state])}
1518
1588
 
1519
1589
  return handler.function(**input_tensors_by_arg_name)
1520
1590
 
@@ -20,7 +20,7 @@ from textwrap import wrap
20
20
  import math
21
21
 
22
22
  from code_loader.contract.visualizer_classes import LeapImage, LeapImageWithBBox, LeapGraph, LeapText, \
23
- LeapHorizontalBar, LeapImageMask, LeapTextMask, LeapImageWithHeatmap, LeapVideo
23
+ LeapHorizontalBar, LeapImageMask, LeapTextMask, LeapImageWithHeatmap, LeapVideo, LeapVolume
24
24
  from code_loader.utils import rescale_min_max
25
25
 
26
26
 
@@ -420,6 +420,28 @@ def plot_video(leap_data: LeapVideo, title: str) -> None:
420
420
  plt.pause(0.1) # Adjust the pause duration as needed
421
421
  plt.show()
422
422
 
423
+ @run_only_on_non_mapping_mode()
424
+ def plot_volume(leap_data: LeapVolume, title: str) -> None:
425
+ """Middle slice along each axis, grayscale as the engine stores it (uint8), mask overlaid."""
426
+ volume = rescale_min_max(leap_data.data.astype(np.float32))
427
+ spacing = leap_data.spacing or (1.0, 1.0, 1.0)
428
+ fig, axes = plt.subplots(1, 3, figsize=(12, 4))
429
+ fig.patch.set_facecolor('black')
430
+ fig.suptitle(title, color='white')
431
+ for axis, ax in enumerate(axes):
432
+ index = volume.shape[axis] // 2
433
+ ax.imshow(np.take(volume, index, axis=axis), cmap='gray', vmin=0, vmax=255)
434
+ if leap_data.mask is not None:
435
+ mask = np.take(leap_data.mask, index, axis=axis)
436
+ ax.imshow(np.ma.masked_where(mask == 0, mask), cmap='tab10', alpha=0.4,
437
+ vmin=0, vmax=max(9, len(leap_data.labels) - 1))
438
+ rows, cols = [i for i in range(3) if i != axis]
439
+ ax.set_aspect(spacing[rows] / spacing[cols])
440
+ ax.set_title(f'axis {axis} / slice {index}', color='white')
441
+ ax.axis('off')
442
+ plt.show()
443
+
444
+
423
445
  @run_only_on_non_mapping_mode()
424
446
  def plot_image_with_heatmap(leap_data: LeapImageWithHeatmap, title: str) -> None:
425
447
  """
@@ -474,4 +496,5 @@ plot_switch = {
474
496
  LeapDataType.ImageWithHeatmap: plot_image_with_heatmap,
475
497
  LeapDataType.ImageWithBBox: plot_image_with_b_box,
476
498
  LeapDataType.Video: plot_video,
499
+ LeapDataType.Volume: plot_volume,
477
500
  }
code_loader/utils.py CHANGED
@@ -250,8 +250,8 @@ def autoregressive_nests_equal(a: Any, b: Any) -> bool:
250
250
 
251
251
 
252
252
  def sample_preprocess_response_arg_name(func: Callable[..., Any]) -> Optional[str]:
253
- # 'return' lives in annotations too and must never be mistaken for a parameter.
254
- for arg_name, arg_type in inspect.getfullargspec(func).annotations.items():
255
- if arg_name != 'return' and arg_type == SamplePreprocessResponse:
256
- return arg_name
253
+ # inspect.signature follows functools.wraps' __wrapped__; getfullargspec does not.
254
+ for param in inspect.signature(func).parameters.values():
255
+ if param.annotation == SamplePreprocessResponse:
256
+ return param.name
257
257
  return None
@@ -5,7 +5,7 @@ import numpy as np
5
5
  import numpy.typing as npt
6
6
 
7
7
  from code_loader.contract.visualizer_classes import LeapImage, LeapGraph, LeapHorizontalBar, LeapText, \
8
- LeapImageMask, LeapTextMask, LeapVideo, LeapAudio
8
+ LeapImageMask, LeapTextMask, LeapVideo, LeapAudio, LeapVolume
9
9
  from code_loader.utils import rescale_min_max
10
10
 
11
11
 
@@ -92,6 +92,20 @@ def default_audio_visualizer(audio: npt.NDArray[np.float32],
92
92
  raise ValueError(f'visual must be 1-D, 2-D, or 3-D, got shape {v.shape}')
93
93
 
94
94
 
95
+ def default_volume_visualizer(data: npt.NDArray[np.float32]) -> LeapVolume:
96
+ """Strip the batch axis and a single channel axis (first or last) to get a [D, H, W] volume."""
97
+ volume = data[0]
98
+ if volume.ndim == 4:
99
+ if volume.shape[0] == 1:
100
+ volume = volume[0]
101
+ elif volume.shape[-1] == 1:
102
+ volume = volume[..., 0]
103
+ if volume.ndim != 3:
104
+ raise ValueError('default_volume_visualizer expects [1, D, H, W], [1, 1, D, H, W] or [1, D, H, W, 1], '
105
+ f'got {data.shape}')
106
+ return LeapVolume(volume)
107
+
108
+
95
109
  def default_graph_visualizer(data: npt.NDArray[np.float32]) -> LeapGraph:
96
110
  return LeapGraph(data[0])
97
111
 
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: code-loader
3
- Version: 1.0.208.dev8
3
+ Version: 1.0.210
4
4
  Summary:
5
5
  Home-page: https://github.com/tensorleap/code-loader
6
6
  License: MIT
@@ -1,13 +1,13 @@
1
1
  LICENSE,sha256=qIwWjdspQeSMTtnFZBC8MuT-95L02FPvzRUdWFxrwJY,1067
2
2
  code_loader/__init__.py,sha256=outxRQ0M-zMfV0QGVJmAed5qWfRmyD0TV6-goEGAzBw,406
3
3
  code_loader/contract/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
4
- code_loader/contract/datasetclasses.py,sha256=yNEPIyZqCO54fKBgmgbFIKEMVBzvpf4n8U2zmPhXv88,18110
5
- code_loader/contract/enums.py,sha256=R4Ge9-oSO7-GYjZAKMTB8Ws5FNsxhCQAYAJruxViJxk,1886
4
+ code_loader/contract/datasetclasses.py,sha256=zDywsj4N9rH7aUU8vkXvOjjj3R-rr93HiC62m0L4rb0,18134
5
+ code_loader/contract/enums.py,sha256=9pv0TwYW6alqsx2jVPb_eylC6qpgMBbdTlCBgG6ZJn0,1908
6
6
  code_loader/contract/exceptions.py,sha256=jWqu5i7t-0IG0jGRsKF4DjJdrsdpJjIYpUkN1F4RiyQ,51
7
7
  code_loader/contract/mapping.py,sha256=i-yVyFuGUITL5lchPNfYlu7bhwk-yuoHPacAeFhcYr8,1490
8
8
  code_loader/contract/responsedataclasses.py,sha256=2SQCccuIlSeUJT0igyvIJRYtmaWSqqlQwtaAA1iEuSI,5044
9
9
  code_loader/contract/sim_config.py,sha256=le8KMALZiP0WU4UcuKnTOSWBW2rNjpnWYfII502NqDM,3493
10
- code_loader/contract/visualizer_classes.py,sha256=vzX9YcwxKOm3IpYj8OaqsA1odPlRgj2Cfvglwd88Wbw,18213
10
+ code_loader/contract/visualizer_classes.py,sha256=ZXrNHbT4kBvFG2r3og1oQRxxhdVndkToqpa8OVnd9gs,21272
11
11
  code_loader/default_losses.py,sha256=NoOQym1106bDN5dcIk56Elr7ZG5quUHArqfP5-Nyxyo,1139
12
12
  code_loader/default_metrics.py,sha256=2XSlyNw_XLDGSJDoz5W_Evi5wbL0dhwq24pPr15vSPc,5025
13
13
  code_loader/experiment_api/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
@@ -21,18 +21,18 @@ code_loader/experiment_api/types.py,sha256=MY8xFARHwdVA7p4dxyhD60ShmttgTvb4qdp1o
21
21
  code_loader/experiment_api/utils.py,sha256=XZHtxge12TS4H4-8PjV3sKuhp8Ud6ojAiIzTZJEqBqc,3304
22
22
  code_loader/experiment_api/workingspace_config_utils.py,sha256=DLzXQCg4dgTV_YgaSbeTVzq-2ja_SQw4zi7LXwKL9cY,990
23
23
  code_loader/inner_leap_binder/__init__.py,sha256=koOlJyMNYzGbEsoIbXathSmQ-L38N_pEXH_HvL7beXU,99
24
- code_loader/inner_leap_binder/leapbinder.py,sha256=fplNN3TfJHg79YpEPHKK3b6CxrR7W0ugDeHhAZfYty0,62784
25
- code_loader/inner_leap_binder/leapbinder_decorators.py,sha256=paUtnncLn-4iFUJhqEOg3rPtKh9DoS3_HVgGUgOz8c4,210906
26
- code_loader/leaploader.py,sha256=DiVCs076JP21HXBfpKkzs3Qaumg8Cz26q1tIc4C-kcg,97552
24
+ code_loader/inner_leap_binder/leapbinder.py,sha256=EhUD4OLSBuVslKYgiXFDoufnfatwCb0LQe_gH9kg5w0,63615
25
+ code_loader/inner_leap_binder/leapbinder_decorators.py,sha256=yxo3L7GlFGpwDuU1wtopp9kscFB5V6RnetYnvII8K8g,212536
26
+ code_loader/leaploader.py,sha256=e7jNqgurKp8o6zbjeXMwPp6Fm3Eltb9VCwv_5IxLENk,101375
27
27
  code_loader/leaploaderbase.py,sha256=JzgpEfY-kNUWbWUb_cZ1G1zHo6oSci6QhT7hqw3E1Jg,12371
28
28
  code_loader/mixpanel_tracker.py,sha256=rNwRmFifNbdUoqLQvvhhgpKczWpWiEmd8MfyJe27sxw,9131
29
29
  code_loader/plot_functions/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
30
- code_loader/plot_functions/plot_functions.py,sha256=2DC-zlVaN13P4VNx5d8csgs80C6SisaeP1-Kq2LW7iM,16075
30
+ code_loader/plot_functions/plot_functions.py,sha256=m2Gre8fZxypE4Pdf1HrYL7Xv8AVLd_7fZPSa8WgydJc,17190
31
31
  code_loader/plot_functions/visualize.py,sha256=gsBAYYkwMh7jIpJeDMPS8G4CW-pxwx6LznoQIvi4vpo,657
32
- code_loader/utils.py,sha256=v6VraCdbFADoNwNx5TH89BLnZ6yfPoWMkyJrqQ_usFc,12377
32
+ code_loader/utils.py,sha256=Yuv2hPx4G70RV2x0Fc_UCube9BLpEryyq4cuKIcsq7s,12347
33
33
  code_loader/visualizers/__init__.py,sha256=47DEQpj8HBSa-_TImW-5JCeuQeRkm5NMpJWZG3hSuFU,0
34
- code_loader/visualizers/default_visualizers.py,sha256=grTPin_lCE9aci8i8CqA7DqQwAyXRB7_EamA3na_pls,5438
35
- code_loader-1.0.208.dev8.dist-info/LICENSE,sha256=qIwWjdspQeSMTtnFZBC8MuT-95L02FPvzRUdWFxrwJY,1067
36
- code_loader-1.0.208.dev8.dist-info/METADATA,sha256=Zv4Z5dCnspuqPU778nnj7nEOkmHqIhxBg44L_ZwRIUw,1095
37
- code_loader-1.0.208.dev8.dist-info/WHEEL,sha256=sP946D7jFCHeNz5Iq4fL4Lu-PrWrFsgfLXbbkciIZwg,88
38
- code_loader-1.0.208.dev8.dist-info/RECORD,,
34
+ code_loader/visualizers/default_visualizers.py,sha256=bu-4A_Uix5egUJpUOq_vQzN3-vidDRoeeazjcC1l-c4,6023
35
+ code_loader-1.0.210.dist-info/LICENSE,sha256=qIwWjdspQeSMTtnFZBC8MuT-95L02FPvzRUdWFxrwJY,1067
36
+ code_loader-1.0.210.dist-info/METADATA,sha256=ULDVtuTKgMNB37laAG_2o51WbTmox-heCsomRRBYAEc,1090
37
+ code_loader-1.0.210.dist-info/WHEEL,sha256=sP946D7jFCHeNz5Iq4fL4Lu-PrWrFsgfLXbbkciIZwg,88
38
+ code_loader-1.0.210.dist-info/RECORD,,