code-loader 1.0.207.dev0__tar.gz → 1.0.208.dev1__tar.gz

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.
Files changed (37) hide show
  1. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/PKG-INFO +1 -1
  2. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/contract/datasetclasses.py +14 -1
  3. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/contract/enums.py +5 -0
  4. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/contract/mapping.py +1 -0
  5. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/contract/responsedataclasses.py +1 -0
  6. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/inner_leap_binder/leapbinder.py +67 -12
  7. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/inner_leap_binder/leapbinder_decorators.py +191 -7
  8. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/leaploader.py +72 -3
  9. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/leaploaderbase.py +25 -0
  10. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/pyproject.toml +1 -1
  11. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/LICENSE +0 -0
  12. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/README.md +0 -0
  13. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/__init__.py +0 -0
  14. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/contract/__init__.py +0 -0
  15. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/contract/exceptions.py +0 -0
  16. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/contract/sim_config.py +0 -0
  17. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/contract/visualizer_classes.py +0 -0
  18. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/default_losses.py +0 -0
  19. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/default_metrics.py +0 -0
  20. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/experiment_api/__init__.py +0 -0
  21. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/experiment_api/api.py +0 -0
  22. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/experiment_api/cli_config_utils.py +0 -0
  23. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/experiment_api/client.py +0 -0
  24. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/experiment_api/epoch.py +0 -0
  25. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/experiment_api/experiment.py +0 -0
  26. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/experiment_api/experiment_context.py +0 -0
  27. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/experiment_api/types.py +0 -0
  28. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/experiment_api/utils.py +0 -0
  29. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/experiment_api/workingspace_config_utils.py +0 -0
  30. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/inner_leap_binder/__init__.py +0 -0
  31. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/mixpanel_tracker.py +0 -0
  32. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/plot_functions/__init__.py +0 -0
  33. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/plot_functions/plot_functions.py +0 -0
  34. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/plot_functions/visualize.py +0 -0
  35. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/utils.py +0 -0
  36. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/visualizers/__init__.py +0 -0
  37. {code_loader-1.0.207.dev0 → code_loader-1.0.208.dev1}/code_loader/visualizers/default_visualizers.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.1
2
2
  Name: code-loader
3
- Version: 1.0.207.dev0
3
+ Version: 1.0.208.dev1
4
4
  Summary:
5
5
  Home-page: https://github.com/tensorleap/code-loader
6
6
  License: MIT
@@ -6,13 +6,21 @@ import numpy as np
6
6
  import numpy.typing as npt
7
7
 
8
8
  from code_loader.contract.enums import DataStateType, DataStateEnum, LeapDataType, ConfusionMatrixValue, \
9
- MetricDirection, DatasetMetadataType
9
+ MetricDirection, DatasetMetadataType, LatentSpaceReduction
10
10
  from code_loader.contract.visualizer_classes import LeapImage, LeapText, LeapGraph, LeapHorizontalBar, \
11
11
  LeapTextMask, LeapImageMask, LeapImageWithBBox, LeapImageWithHeatmap, LeapVideo, LeapAudio
12
12
  from code_loader.contract.sim_config import SimConfig
13
13
 
14
14
  custom_latent_space_attribute = "custom_latent_space"
15
15
 
16
+ # Hard cap on a custom latent space's per-sample width, applied AFTER any reduction. The
17
+ # soft threshold only warns. MAX_CUSTOM_LATENT_SPACES mirrors the engine's
18
+ # src_tensorleap/common/types.py (source of truth) so an over-budget registration fails
19
+ # locally instead of being silently dropped at eval time; keep the two in sync.
20
+ MAX_CUSTOM_LATENT_SPACE_DIM = 4096
21
+ CUSTOM_LATENT_SPACE_WARN_DIM = 1024
22
+ MAX_CUSTOM_LATENT_SPACES = 10
23
+
16
24
  _simulation_context: Dict[str, bool] = {"active": False}
17
25
 
18
26
  SampleId = Union[int, str]
@@ -317,6 +325,11 @@ class CustomLatentSpaceHandler:
317
325
  name: str = 'custom_latent_space'
318
326
  use_ls_for_analysis: bool = False
319
327
  instance_aware: bool = False
328
+ computed_at: str = 'dataset'
329
+ arg_names: Optional[List[str]] = None
330
+ reduce: Optional[LatentSpaceReduction] = None
331
+ n_components: int = 512
332
+ channel_axis: int = -1
320
333
 
321
334
 
322
335
  # How a chain's latent-space vectors are derived from its per-step forward passes.
@@ -68,3 +68,8 @@ class ConfusionMatrixValue(Enum):
68
68
  class TestingSectionEnum(Enum):
69
69
  Warnings = "Warnings"
70
70
  Errors = "Errors"
71
+
72
+
73
+ class LatentSpaceReduction(Enum):
74
+ MEAN_POOL = 'MEAN_POOL'
75
+ RANDOM_PROJECTION = 'RANDOM_PROJECTION'
@@ -38,6 +38,7 @@ class NodeMappingType(Enum):
38
38
  Input8 = 'Input8'
39
39
  Input9 = 'Input9'
40
40
  PredictionLabels = 'PredictionLabels'
41
+ CustomLatentSpace = 'CustomLatentSpace'
41
42
 
42
43
 
43
44
  @dataclass
@@ -144,6 +144,7 @@ class LeapAnalysisConfiguration:
144
144
  class EngineFileContract:
145
145
  node_connections: Optional[List[NodeConnection]] = None
146
146
  leap_analysis_configuration: Optional[LeapAnalysisConfiguration] = None
147
+ latent_space_connections: Optional[List[NodeConnection]] = None
147
148
 
148
149
 
149
150
  @dataclass
@@ -17,12 +17,12 @@ from code_loader.contract.datasetclasses import SectionCallableInterface, InputH
17
17
  CustomMultipleReturnCallableInterfaceMultiArgs, DatasetBaseHandler, custom_latent_space_attribute, \
18
18
  RawInputsForHeatmap, VisualizerHandlerData, MetricHandlerData, CustomLossHandlerData, SamplePreprocessResponse, \
19
19
  ElementInstanceMasksHandler, InstanceCallableInterface, InstanceSectionCallableInterface, \
20
- CustomLatentSpaceHandler, InstanceMetricHandler, \
20
+ CustomLatentSpaceHandler, InstanceMetricHandler, MAX_CUSTOM_LATENT_SPACES, \
21
21
  SimulationHandler, _simulation_context, AutoregressiveStepHandler, AutoregressiveStepCallableInterface, \
22
22
  AUTOREGRESSIVE_LATENT_SPACE_AGGREGATIONS, AUTOREGRESSIVE_IMPLICIT_ARG_NAMES, \
23
23
  AutoregressiveMetricHandler, AutoregressiveLossHandler, AutoregressiveVisualizerHandler
24
24
  from code_loader.contract.enums import LeapDataType, DataStateEnum, DataStateType, MetricDirection, DatasetMetadataType, \
25
- TestingSectionEnum
25
+ TestingSectionEnum, LatentSpaceReduction
26
26
  from code_loader.contract.mapping import NodeConnection, NodeMapping, NodeMappingType
27
27
  from code_loader.contract.responsedataclasses import DatasetTestResultPayload, LeapAnalysisConfiguration
28
28
  from code_loader.contract.visualizer_classes import map_leap_data_type_to_visualizer_class
@@ -67,16 +67,29 @@ def _stringized_annotation_type_name(annotation: Any) -> Optional[str]:
67
67
  return None
68
68
 
69
69
 
70
- def _reject_stringized_sample_preprocess_response(function: Callable[..., Any], arg_name: str, annotation: Any) -> None:
71
- """Fail fast on a stringized SamplePreprocessResponse arg (otherwise it silently
72
- mis-wires as a regular input, surfacing only on the platform)."""
73
- if _stringized_annotation_type_name(annotation) == SamplePreprocessResponse.__name__:
70
+ def _reject_stringized_annotation(function: Callable[..., Any], arg_name: str, annotation: Any,
71
+ expected_cls: type, why: str) -> None:
72
+ """Fail fast on a stringized `expected_cls` arg (otherwise it silently mis-wires,
73
+ surfacing only on the platform)."""
74
+ if _stringized_annotation_type_name(annotation) == expected_cls.__name__:
74
75
  raise Exception(
75
76
  f"Argument '{arg_name}' of function '{function.__name__}' is annotated with a string "
76
- f"('{annotation}') instead of the SamplePreprocessResponse type. This usually means the "
77
+ f"('{annotation}') instead of the {expected_cls.__name__} type. This usually means the "
77
78
  f"file uses 'from __future__ import annotations' (or a quoted hint), which stringizes "
78
- f"annotations and breaks Tensorleap type detection. Remove that import (or the quotes) so "
79
- f"SamplePreprocessResponse is referenced as a real type.")
79
+ f"annotations. {why} Remove that import (or the quotes).")
80
+
81
+
82
+ def _reject_stringized_sample_preprocess_response(function: Callable[..., Any], arg_name: str, annotation: Any) -> None:
83
+ _reject_stringized_annotation(
84
+ function, arg_name, annotation, SamplePreprocessResponse,
85
+ "This breaks Tensorleap type detection, so it must be referenced as a real type.")
86
+
87
+
88
+ def _reject_stringized_preprocess_response(function: Callable[..., Any], arg_name: str, annotation: Any) -> None:
89
+ _reject_stringized_annotation(
90
+ function, arg_name, annotation, PreprocessResponse,
91
+ "Tensorleap uses that annotation to tell a dataset-computed custom latent space from a "
92
+ "model-computed one, so it must be a real type.")
80
93
 
81
94
 
82
95
 
@@ -102,6 +115,7 @@ class LeapBinder:
102
115
  self._extend_with_default_losses()
103
116
 
104
117
  self.mapping_connections: List[NodeConnection] = []
118
+ self.latent_space_connections: List[NodeConnection] = []
105
119
  self.integration_test_func: Optional[Callable[[str, PreprocessResponse], Any]] = None
106
120
 
107
121
  self.batch_size_to_validate: Optional[int] = None
@@ -546,7 +560,12 @@ class LeapBinder:
546
560
  def set_custom_latent_space(self, function: Union[SectionCallableInterface, InstanceSectionCallableInterface],
547
561
  name: Optional[str] = None,
548
562
  use_ls_for_analysis: bool = False,
549
- instance_aware: bool = False) -> None:
563
+ instance_aware: bool = False,
564
+ computed_at: str = 'dataset',
565
+ arg_names: Optional[List[str]] = None,
566
+ reduce: Optional[LatentSpaceReduction] = None,
567
+ n_components: int = 512,
568
+ channel_axis: int = -1) -> None:
550
569
  """
551
570
  Register a custom latent space function.
552
571
 
@@ -567,7 +586,8 @@ class LeapBinder:
567
586
  space for the Out-Of-Distribution and Domain-Gap insights instead of the
568
587
  built-in defaults. At most one registered custom latent space may set this;
569
588
  registering a second one with the flag raises. Not currently supported when
570
- instance_aware=True — it is ignored (with a warning) and forced to False.
589
+ instance_aware=True, or for a model-computed latent space (computed_at='model')
590
+ — it is ignored (with a warning) and forced to False.
571
591
  instance_aware (bool): When True, `function` takes a third `instance_id` argument
572
592
  and is called once per element-instance row instead of once per sample.
573
593
  """
@@ -579,6 +599,14 @@ class LeapBinder:
579
599
  f"@tensorleap_custom_latent_space must have a unique name "
580
600
  f"(pass name='...' to distinguish them)."
581
601
  )
602
+ if len(self.setup_container.custom_latent_spaces) >= MAX_CUSTOM_LATENT_SPACES:
603
+ raise Exception(
604
+ f"Cannot register custom latent space '{name}': Tensorleap supports at most "
605
+ f"{MAX_CUSTOM_LATENT_SPACES} custom latent spaces, and "
606
+ f"{len(self.setup_container.custom_latent_spaces)} are already registered "
607
+ f"({sorted(self.setup_container.custom_latent_spaces)}). Dataset-computed and "
608
+ f"model-computed latent spaces share this budget."
609
+ )
582
610
  # use_ls_for_analysis is not currently wired for instance-aware latent spaces (OOD /
583
611
  # Domain-Gap analyze the sample-level population, not instance rows) — force it off rather
584
612
  # than silently accepting a flag that has no effect.
@@ -588,6 +616,16 @@ class LeapBinder:
588
616
  f"latent space ('{name}'). Ignoring it; the flag will be set to False."
589
617
  )
590
618
  use_ls_for_analysis = False
619
+ # Model-computed latent spaces are fetched via run_model_latent_space, a separate path
620
+ # that get_sample's per-sample custom_latent_spaces dict never populates — the engine has
621
+ # no consumer for a model-computed analysis latent space yet, so force it off rather than
622
+ # silently accepting a flag that has no effect.
623
+ if computed_at == 'model' and use_ls_for_analysis:
624
+ warnings.warn(
625
+ f"use_ls_for_analysis=True is not currently supported for a model-computed custom "
626
+ f"latent space ('{name}'). Ignoring it; the flag will be set to False."
627
+ )
628
+ use_ls_for_analysis = False
591
629
  if use_ls_for_analysis:
592
630
  already_flagged = [
593
631
  existing_name
@@ -602,8 +640,25 @@ class LeapBinder:
602
640
  f"Out-Of-Distribution and Domain-Gap insights). Set it on '{name}' "
603
641
  f"or '{already_flagged[0]}', not both."
604
642
  )
643
+ if reduce is not None:
644
+ if not isinstance(reduce, LatentSpaceReduction):
645
+ raise Exception(
646
+ f"Custom latent space '{name}': reduce must be a LatentSpaceReduction, got "
647
+ f"{type(reduce).__name__}.")
648
+ if reduce is LatentSpaceReduction.RANDOM_PROJECTION and (
649
+ not isinstance(n_components, int) or isinstance(n_components, bool) or n_components <= 0):
650
+ raise Exception(
651
+ f"Custom latent space '{name}': n_components must be a positive int for "
652
+ f"LatentSpaceReduction.RANDOM_PROJECTION, got {n_components!r}.")
653
+ if reduce is LatentSpaceReduction.MEAN_POOL and (
654
+ not isinstance(channel_axis, int) or isinstance(channel_axis, bool)):
655
+ raise Exception(
656
+ f"Custom latent space '{name}': channel_axis must be an int for "
657
+ f"LatentSpaceReduction.MEAN_POOL, got {channel_axis!r}.")
605
658
  self.setup_container.custom_latent_spaces[name] = CustomLatentSpaceHandler(
606
- function=function, name=name, use_ls_for_analysis=use_ls_for_analysis, instance_aware=instance_aware)
659
+ function=function, name=name, use_ls_for_analysis=use_ls_for_analysis, instance_aware=instance_aware,
660
+ computed_at=computed_at, arg_names=arg_names, reduce=reduce, n_components=n_components,
661
+ channel_axis=channel_axis)
607
662
 
608
663
  def set_autoregressive_step(self, function: AutoregressiveStepCallableInterface,
609
664
  latent_space_aggregation: str = 'last_step') -> None:
@@ -26,15 +26,17 @@ from code_loader.contract.datasetclasses import CustomCallableInterfaceMultiArgs
26
26
  CustomMultipleReturnCallableInterfaceMultiArgs, ConfusionMatrixCallableInterfaceMultiArgs, CustomCallableInterface, \
27
27
  VisualizerCallableInterface, MetadataSectionCallableInterface, PreprocessResponse, SectionCallableInterface, \
28
28
  ConfusionMatrixElement, SamplePreprocessResponse, PredictionTypeHandler, InstanceCallableInterface, ElementInstance, \
29
- InstanceLengthCallableInterface, InstanceSectionCallableInterface, AutoregressiveStepCallableInterface
29
+ InstanceLengthCallableInterface, InstanceSectionCallableInterface, AutoregressiveStepCallableInterface, \
30
+ MAX_CUSTOM_LATENT_SPACE_DIM, CUSTOM_LATENT_SPACE_WARN_DIM
30
31
  from code_loader.contract.enums import MetricDirection, LeapDataType, DatasetMetadataType, DataStateType, \
31
- DataStateEnum
32
+ DataStateEnum, LatentSpaceReduction
32
33
  from code_loader import leap_binder, LeapLoader
33
34
  from code_loader.contract.mapping import NodeMapping, NodeMappingType, NodeConnection
34
35
  from code_loader.contract.visualizer_classes import LeapImage, LeapImageMask, LeapTextMask, LeapText, LeapGraph, \
35
36
  LeapHorizontalBar, LeapImageWithBBox, LeapImageWithHeatmap, LeapVideo, LeapAudio, LeapValidationError, \
36
37
  map_leap_data_type_to_visualizer_class
37
- from code_loader.inner_leap_binder.leapbinder import mapping_runtime_mode_env_var_mame
38
+ from code_loader.inner_leap_binder.leapbinder import mapping_runtime_mode_env_var_mame, \
39
+ _reject_stringized_preprocess_response, _reject_stringized_sample_preprocess_response
38
40
  from code_loader.mixpanel_tracker import clear_integration_events, AnalyticsEvent, emit_integration_event_once
39
41
 
40
42
  _called_from_inside_tl_decorator = 0
@@ -264,7 +266,8 @@ def _validate_grouped_result(result, group_size, func_name, validate_single):
264
266
  f'{group_size} arrays, got {type(result)}.')
265
267
 
266
268
 
267
- def _add_mapping_connection(user_unique_name, connection_destinations, arg_names, name, node_mapping_type):
269
+ def _add_mapping_connection(user_unique_name, connection_destinations, arg_names, name, node_mapping_type,
270
+ target_list=None):
268
271
  connection_destinations = [connection_destination for connection_destination in connection_destinations
269
272
  if not isinstance(connection_destination, SamplePreprocessResponse)]
270
273
 
@@ -274,7 +277,9 @@ def _add_mapping_connection(user_unique_name, connection_destinations, arg_names
274
277
  for arg_name, destination in zip(arg_names, connection_destinations):
275
278
  node_inputs[arg_name] = destination.node_mapping
276
279
 
277
- leap_binder.mapping_connections.append(NodeConnection(main_node_mapping, node_inputs))
280
+ if target_list is None:
281
+ target_list = leap_binder.mapping_connections
282
+ target_list.append(NodeConnection(main_node_mapping, node_inputs))
278
283
 
279
284
 
280
285
  def _add_mapping_connections(connects_to, arg_names, node_mapping_type, name):
@@ -1750,7 +1755,92 @@ def tensorleap_metadata(
1750
1755
  return decorating_function
1751
1756
 
1752
1757
 
1753
- def tensorleap_custom_latent_space(name: Optional[str] = None, use_ls_for_analysis: bool = False):
1758
+ def _classify_custom_latent_space_signature(user_function) -> str:
1759
+ argspec = inspect.getfullargspec(user_function)
1760
+ params = argspec.args
1761
+ if len(params) != 2:
1762
+ return 'model'
1763
+
1764
+ second = params[1]
1765
+ first = params[0]
1766
+ if second not in argspec.annotations:
1767
+ raise Exception(
1768
+ f"tensorleap_custom_latent_space validation failed: '{user_function.__name__}' has "
1769
+ f"exactly two parameters ('{first}', '{second}') and '{second}' has no type "
1770
+ f"annotation, so Tensorleap cannot tell whether this is a dataset-computed latent "
1771
+ f"space (one sample at a time) or a model-computed one (a batch of model tensors). "
1772
+ f"Please annotate the second parameter '{second}' to disambiguate:\n"
1773
+ f" dataset-computed: def {user_function.__name__}({first}, {second}: PreprocessResponse) -> (d,)\n"
1774
+ f" model-computed: def {user_function.__name__}({first}: np.ndarray, {second}: np.ndarray) -> (batch, d)\n"
1775
+ f"If you are upgrading an existing project, this signature used to be accepted "
1776
+ f"unannotated as dataset-computed; add ': PreprocessResponse' to '{second}' to keep "
1777
+ f"the previous behavior.")
1778
+
1779
+ annotation = argspec.annotations[second]
1780
+ _reject_stringized_preprocess_response(user_function, second, annotation)
1781
+ if annotation is PreprocessResponse:
1782
+ return 'dataset'
1783
+ return 'model'
1784
+
1785
+
1786
+ def _model_latent_space_arg_names(user_function) -> List[str]:
1787
+ argspec = inspect.getfullargspec(user_function)
1788
+ arg_names = list(argspec.args)
1789
+ preprocess_response_arg_name = None
1790
+ for arg_name, arg_type in argspec.annotations.items():
1791
+ _reject_stringized_sample_preprocess_response(user_function, arg_name, arg_type)
1792
+ if arg_type == SamplePreprocessResponse:
1793
+ if preprocess_response_arg_name is not None:
1794
+ raise Exception(
1795
+ f"tensorleap_custom_latent_space validation failed: only one argument of "
1796
+ f"'{user_function.__name__}' can be of type SamplePreprocessResponse.")
1797
+ preprocess_response_arg_name = arg_name
1798
+ arg_names.remove(arg_name)
1799
+ return arg_names
1800
+
1801
+
1802
+ def _custom_latent_space_width(result, ls_name, reduce, n_components, channel_axis,
1803
+ has_batch_axis: bool) -> int:
1804
+ if reduce is LatentSpaceReduction.RANDOM_PROJECTION:
1805
+ return n_components
1806
+ if reduce is LatentSpaceReduction.MEAN_POOL:
1807
+ axis = channel_axis if channel_axis >= 0 else result.ndim + channel_axis
1808
+ min_axis = 1 if has_batch_axis else 0
1809
+ if axis < min_axis or axis >= result.ndim:
1810
+ raise Exception(
1811
+ f"tensorleap_custom_latent_space validation failed: '{ls_name}' uses "
1812
+ f"LatentSpaceReduction.MEAN_POOL with channel_axis={channel_axis}, which does not "
1813
+ f"select a non-batch axis of the returned shape {tuple(result.shape)}."
1814
+ + (" Axis 0 is the batch." if has_batch_axis else ""))
1815
+ return int(result.shape[axis])
1816
+ dims = result.shape[1:] if has_batch_axis else result.shape
1817
+ return int(np.prod(dims))
1818
+
1819
+
1820
+ def _check_custom_latent_space_width(result, ls_name, reduce, n_components, channel_axis,
1821
+ has_batch_axis: bool) -> None:
1822
+ width = _custom_latent_space_width(result, ls_name, reduce, n_components, channel_axis,
1823
+ has_batch_axis)
1824
+ if width > MAX_CUSTOM_LATENT_SPACE_DIM:
1825
+ raise Exception(
1826
+ f"tensorleap_custom_latent_space validation failed: '{ls_name}' produces "
1827
+ f"{width} dimensions per sample, above the {MAX_CUSTOM_LATENT_SPACE_DIM} limit. "
1828
+ f"Pass reduce=LatentSpaceReduction.RANDOM_PROJECTION (with n_components) to "
1829
+ f"project it down, or reduce=LatentSpaceReduction.MEAN_POOL (with channel_axis) "
1830
+ f"to average the non-channel axes, or return a smaller array.")
1831
+ if width > CUSTOM_LATENT_SPACE_WARN_DIM:
1832
+ store_general_warning(
1833
+ key=("tensorleap_custom_latent_space_width", ls_name, width),
1834
+ message=(
1835
+ f"Custom latent space '{ls_name}' produces {width} dimensions per sample. "
1836
+ f"Wide latent spaces are slow to transit and store. Consider "
1837
+ f"reduce=LatentSpaceReduction.MEAN_POOL or "
1838
+ f"reduce=LatentSpaceReduction.RANDOM_PROJECTION."))
1839
+
1840
+
1841
+ def tensorleap_custom_latent_space(name: Optional[str] = None, use_ls_for_analysis: bool = False,
1842
+ reduce: Optional[LatentSpaceReduction] = None,
1843
+ n_components: int = 512, channel_axis: int = -1):
1754
1844
  assert isinstance(use_ls_for_analysis, bool), \
1755
1845
  ("tensorleap_custom_latent_space validation failed: use_ls_for_analysis must be a bool. "
1756
1846
  f"Got {type(use_ls_for_analysis)}.")
@@ -1758,6 +1848,10 @@ def tensorleap_custom_latent_space(name: Optional[str] = None, use_ls_for_analys
1758
1848
  def decorating_function(user_function: SectionCallableInterface):
1759
1849
  ls_name = name if name is not None else user_function.__name__
1760
1850
 
1851
+ if _classify_custom_latent_space_signature(user_function) == 'model':
1852
+ return _decorate_model_latent_space(user_function, ls_name, use_ls_for_analysis,
1853
+ reduce, n_components, channel_axis)
1854
+
1761
1855
  def _validate_input_args(sample_id: Union[int, str, list], preprocess_response: PreprocessResponse):
1762
1856
  _validate_id_or_group(sample_id, preprocess_response, 'tensorleap_custom_latent_space')
1763
1857
 
@@ -1777,6 +1871,8 @@ def tensorleap_custom_latent_space(name: Optional[str] = None, use_ls_for_analys
1777
1871
  f"inside your function."
1778
1872
  ),
1779
1873
  )
1874
+ _check_custom_latent_space_width(single_result, ls_name, reduce, n_components, channel_axis,
1875
+ has_batch_axis=False)
1780
1876
 
1781
1877
  def _validate_result(result, grouped=False, group_size=None):
1782
1878
  if not grouped:
@@ -1814,7 +1910,9 @@ def tensorleap_custom_latent_space(name: Optional[str] = None, use_ls_for_analys
1814
1910
  return result
1815
1911
 
1816
1912
  leap_binder.set_custom_latent_space(inner_without_validate, ls_name,
1817
- use_ls_for_analysis=use_ls_for_analysis)
1913
+ use_ls_for_analysis=use_ls_for_analysis,
1914
+ computed_at='dataset', reduce=reduce,
1915
+ n_components=n_components, channel_axis=channel_axis)
1818
1916
 
1819
1917
  def inner(sample_id, preprocess_response):
1820
1918
  if os.environ.get(mapping_runtime_mode_env_var_mame):
@@ -1834,6 +1932,92 @@ def tensorleap_custom_latent_space(name: Optional[str] = None, use_ls_for_analys
1834
1932
  return decorating_function
1835
1933
 
1836
1934
 
1935
+ def _decorate_model_latent_space(user_function, ls_name, use_ls_for_analysis, reduce, n_components,
1936
+ channel_axis):
1937
+ arg_names = _model_latent_space_arg_names(user_function)
1938
+
1939
+ def _validate_input_args(*args, **kwargs):
1940
+ assert len(args) + len(kwargs) > 0, (
1941
+ f"tensorleap_custom_latent_space validation failed: '{user_function.__name__}' is "
1942
+ f"model-computed and expects at least one np.ndarray argument, but received none.")
1943
+ for i, arg in enumerate(args):
1944
+ assert isinstance(arg, (np.ndarray, SamplePreprocessResponse)), (
1945
+ f"tensorleap_custom_latent_space validation failed: Argument #{i} of "
1946
+ f"'{user_function.__name__}' should be a numpy array. Got {type(arg)}.")
1947
+ for arg_name, arg in kwargs.items():
1948
+ assert isinstance(arg, (np.ndarray, SamplePreprocessResponse)), (
1949
+ f"tensorleap_custom_latent_space validation failed: Argument {arg_name} of "
1950
+ f"'{user_function.__name__}' should be a numpy array. Got {type(arg)}.")
1951
+
1952
+ def _validate_result(result):
1953
+ assert isinstance(result, np.ndarray), (
1954
+ f"tensorleap_custom_latent_space validation failed: '{user_function.__name__}' is "
1955
+ f"model-computed and should return a numpy array of shape (batch, d). "
1956
+ f"Got {type(result)}.")
1957
+ assert result.ndim >= 2, (
1958
+ f"tensorleap_custom_latent_space validation failed: '{user_function.__name__}' "
1959
+ f"returned shape {tuple(result.shape)}. A model-computed latent space returns "
1960
+ f"(batch, d), so the result needs a batch axis and at least one feature axis.")
1961
+ if leap_binder.batch_size_to_validate:
1962
+ assert result.shape[0] == leap_binder.batch_size_to_validate, (
1963
+ f"tensorleap_custom_latent_space validation failed: '{user_function.__name__}' "
1964
+ f"returned leading dim {result.shape[0]} instead of the batch size "
1965
+ f"{leap_binder.batch_size_to_validate}.")
1966
+
1967
+ _check_custom_latent_space_width(result, ls_name, reduce, n_components, channel_axis,
1968
+ has_batch_axis=True)
1969
+
1970
+ def inner_without_validate(*args, **kwargs):
1971
+ global _called_from_inside_tl_decorator
1972
+ _called_from_inside_tl_decorator += 1
1973
+
1974
+ try:
1975
+ result = user_function(*args, **kwargs)
1976
+ finally:
1977
+ _called_from_inside_tl_decorator -= 1
1978
+
1979
+ return result
1980
+
1981
+ inner_without_validate.__signature__ = inspect.signature(user_function)
1982
+
1983
+ leap_binder.set_custom_latent_space(inner_without_validate, ls_name,
1984
+ use_ls_for_analysis=use_ls_for_analysis,
1985
+ computed_at='model', arg_names=arg_names, reduce=reduce,
1986
+ n_components=n_components, channel_axis=channel_axis)
1987
+
1988
+ def inner(*args, **kwargs):
1989
+ args, kwargs, _was_downcast = _simulate_engine_tensor_dtype(args, kwargs)
1990
+ _validate_input_args(*args, **kwargs)
1991
+ result = inner_without_validate(*args, **kwargs)
1992
+ _validate_result(result)
1993
+ return result
1994
+
1995
+ def mapping_inner(*args, **kwargs):
1996
+ if _mapping_dataset_is_grouped:
1997
+ raise LeapValidationError(
1998
+ f"tensorleap_custom_latent_space validation failed: '{ls_name}' is model-computed, "
1999
+ f"which is not supported for a grouped preprocess response. Use the "
2000
+ f"(sample_id, preprocess: PreprocessResponse) form instead.")
2001
+
2002
+ ordered_connections = [kwargs[arg_name] for arg_name in arg_names if arg_name in kwargs]
2003
+ ordered_connections = list(args) + ordered_connections
2004
+
2005
+ leap_binder.latent_space_connections[:] = [
2006
+ connection for connection in leap_binder.latent_space_connections
2007
+ if connection.node.name != ls_name]
2008
+ _add_mapping_connection(ls_name, ordered_connections, arg_names, ls_name,
2009
+ NodeMappingType.CustomLatentSpace,
2010
+ target_list=leap_binder.latent_space_connections)
2011
+ return None
2012
+
2013
+ def final_inner(*args, **kwargs):
2014
+ if os.environ.get(mapping_runtime_mode_env_var_mame):
2015
+ return mapping_inner(*args, **kwargs)
2016
+ return inner(*args, **kwargs)
2017
+
2018
+ return final_inner
2019
+
2020
+
1837
2021
  def tensorleap_instance_custom_latent_space(name: Optional[str] = None, use_ls_for_analysis: bool = False):
1838
2022
  assert isinstance(use_ls_for_analysis, bool), \
1839
2023
  ("tensorleap_instance_custom_latent_space validation failed: use_ls_for_analysis must be a bool. "
@@ -352,6 +352,9 @@ class LeapLoader(LeapLoaderBase):
352
352
  instance_ls_test_payload = self._check_instance_custom_latent_spaces()
353
353
  if instance_ls_test_payload is not None:
354
354
  test_payloads.append(instance_ls_test_payload)
355
+ model_ls_test_payload = self._check_model_latent_spaces()
356
+ if model_ls_test_payload is not None:
357
+ test_payloads.append(model_ls_test_payload)
355
358
  handlers_test_payloads = self._check_handlers()
356
359
  test_payloads.extend(handlers_test_payloads)
357
360
  simulation_test_payloads = self._check_simulations()
@@ -386,8 +389,10 @@ class LeapLoader(LeapLoaderBase):
386
389
  is_valid_for_model=is_valid_for_model, setup=setup_response,
387
390
  model_setup=model_setup, general_error=general_error,
388
391
  print_log=print_log,
389
- engine_file_contract=EngineFileContract(global_leap_binder.mapping_connections,
390
- global_leap_binder.leap_analysis_configuration))
392
+ engine_file_contract=EngineFileContract(
393
+ global_leap_binder.mapping_connections,
394
+ global_leap_binder.leap_analysis_configuration,
395
+ global_leap_binder.latent_space_connections))
391
396
 
392
397
  def _check_integration_test_exists(self) -> DatasetTestResultPayload:
393
398
  test_result = DatasetTestResultPayload('integration_test')
@@ -425,6 +430,28 @@ class LeapLoader(LeapLoaderBase):
425
430
  test_result.is_passed = False
426
431
  return test_result
427
432
 
433
+ def _check_model_latent_spaces(self) -> Optional[DatasetTestResultPayload]:
434
+ model_names = [
435
+ name for name, handler in global_leap_binder.setup_container.custom_latent_spaces.items()
436
+ if handler.computed_at == 'model'
437
+ ]
438
+ if not model_names:
439
+ return None
440
+
441
+ test_result = DatasetTestResultPayload('model_custom_latent_space')
442
+ grouped_states = [
443
+ state.name for state, preprocess_response in self._preprocess_result().items()
444
+ if preprocess_response.is_grouped
445
+ ]
446
+ if grouped_states:
447
+ test_result.is_passed = False
448
+ test_result.display[TestingSectionEnum.Errors.name] = (
449
+ f"Model-computed custom latent space(s) {model_names} are not supported with a "
450
+ f"grouped preprocess response (grouped: {grouped_states}). Use the "
451
+ f"(sample_id, preprocess: PreprocessResponse) form instead."
452
+ )
453
+ return test_result
454
+
428
455
  def _check_instance_custom_latent_spaces(self) -> Optional[DatasetTestResultPayload]:
429
456
  instance_aware_names = [
430
457
  name for name, handler in global_leap_binder.setup_container.custom_latent_spaces.items()
@@ -1390,7 +1417,9 @@ class LeapLoader(LeapLoaderBase):
1390
1417
  sample_id: Union[int, str],
1391
1418
  preprocess: "PreprocessResponse",
1392
1419
  instance_id: Optional[int] = None) -> Optional[Dict[str, npt.NDArray[np.float32]]]:
1393
- handlers = global_leap_binder.setup_container.custom_latent_spaces
1420
+ handlers = {handler_name: handler for handler_name, handler
1421
+ in global_leap_binder.setup_container.custom_latent_spaces.items()
1422
+ if handler.computed_at != 'model'}
1394
1423
  if not handlers:
1395
1424
  return None
1396
1425
  if preprocess.is_grouped:
@@ -1428,6 +1457,46 @@ class LeapLoader(LeapLoaderBase):
1428
1457
  self.exec_script()
1429
1458
  return tuple(global_leap_binder.setup_container.custom_latent_spaces.keys())
1430
1459
 
1460
+ @lru_cache()
1461
+ def get_dataset_custom_latent_space_names(self) -> Tuple[str, ...]:
1462
+ self.exec_script()
1463
+ return tuple(name for name, handler in global_leap_binder.setup_container.custom_latent_spaces.items()
1464
+ if handler.computed_at != 'model')
1465
+
1466
+ @lru_cache()
1467
+ def get_model_custom_latent_space_names(self) -> Tuple[str, ...]:
1468
+ self.exec_script()
1469
+ return tuple(name for name, handler in global_leap_binder.setup_container.custom_latent_spaces.items()
1470
+ if handler.computed_at == 'model')
1471
+
1472
+ @lru_cache()
1473
+ def get_model_latent_space_specs(self) -> Dict[str, Dict[str, Any]]:
1474
+ self.exec_script()
1475
+ return {
1476
+ name: {
1477
+ 'arg_names': list(handler.arg_names or []),
1478
+ 'reduce': handler.reduce.value if handler.reduce is not None else None,
1479
+ 'n_components': handler.n_components,
1480
+ 'channel_axis': handler.channel_axis,
1481
+ }
1482
+ for name, handler in global_leap_binder.setup_container.custom_latent_spaces.items()
1483
+ if handler.computed_at == 'model'
1484
+ }
1485
+
1486
+ def run_model_latent_space(self, ls_name: str, sample_ids: np.array, state: DataStateEnum,
1487
+ input_tensors_by_arg_name: Dict[str, npt.NDArray[np.float32]]
1488
+ ) -> npt.NDArray[np.float32]:
1489
+ self._preprocess_result()
1490
+
1491
+ handler = global_leap_binder.setup_container.custom_latent_spaces[ls_name]
1492
+ preprocess_response_arg_name = self._get_preprocess_response_arg_name(handler.function)
1493
+
1494
+ if preprocess_response_arg_name is not None:
1495
+ input_tensors_by_arg_name[preprocess_response_arg_name] = SamplePreprocessResponse(
1496
+ sample_ids, self._preprocess_result()[state])
1497
+
1498
+ return handler.function(**input_tensors_by_arg_name)
1499
+
1431
1500
  @lru_cache()
1432
1501
  def get_instance_custom_latent_space_names(self) -> Tuple[str, ...]:
1433
1502
  """Names of registered custom latent spaces that are instance-aware.
@@ -226,6 +226,31 @@ class LeapLoaderBase:
226
226
  def get_custom_latent_space_for_analysis(self) -> Optional[str]:
227
227
  pass
228
228
 
229
+ # These raise rather than `pass` for the same reason as the autoregressive entry points
230
+ # above: an un-overridden `pass` body would report "no model-computed latent spaces" on a
231
+ # loader that simply predates them, and the engine would silently skip computing them.
232
+ @abstractmethod
233
+ def get_dataset_custom_latent_space_names(self) -> Tuple[str, ...]:
234
+ raise NotImplementedError(f'{type(self).__name__} does not implement '
235
+ 'get_dataset_custom_latent_space_names.')
236
+
237
+ @abstractmethod
238
+ def get_model_custom_latent_space_names(self) -> Tuple[str, ...]:
239
+ raise NotImplementedError(f'{type(self).__name__} does not implement '
240
+ 'get_model_custom_latent_space_names.')
241
+
242
+ @abstractmethod
243
+ def get_model_latent_space_specs(self) -> Dict[str, Dict[str, Any]]:
244
+ raise NotImplementedError(f'{type(self).__name__} does not implement '
245
+ 'get_model_latent_space_specs.')
246
+
247
+ @abstractmethod
248
+ def run_model_latent_space(self, ls_name: str, sample_ids: np.array, state: DataStateEnum,
249
+ input_tensors_by_arg_name: Dict[str, npt.NDArray[np.float32]]
250
+ ) -> npt.NDArray[np.float32]:
251
+ raise NotImplementedError(f'{type(self).__name__} does not implement '
252
+ 'run_model_latent_space.')
253
+
229
254
  @abstractmethod
230
255
  def get_heatmap_visualizer_raw_vis_input_arg_name(self, visualizer_name: str) -> Optional[str]:
231
256
  pass
@@ -1,6 +1,6 @@
1
1
  [tool.poetry]
2
2
  name = "code-loader"
3
- version = "1.0.207.dev0"
3
+ version = "1.0.208.dev1"
4
4
  description = ""
5
5
  authors = ["dorhar <doron.harnoy@tensorleap.ai>"]
6
6
  license = "MIT"