code-loader 1.0.208.dev3__tar.gz → 1.0.208.dev5__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.208.dev3 → code_loader-1.0.208.dev5}/PKG-INFO +1 -1
  2. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/contract/datasetclasses.py +2 -2
  3. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/contract/enums.py +5 -0
  4. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/inner_leap_binder/leapbinder.py +2 -2
  5. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/inner_leap_binder/leapbinder_decorators.py +25 -16
  6. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/leaploader.py +27 -19
  7. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/leaploaderbase.py +2 -7
  8. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/utils.py +10 -1
  9. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/pyproject.toml +1 -1
  10. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/LICENSE +0 -0
  11. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/README.md +0 -0
  12. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/__init__.py +0 -0
  13. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/contract/__init__.py +0 -0
  14. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/contract/exceptions.py +0 -0
  15. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/contract/mapping.py +0 -0
  16. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/contract/responsedataclasses.py +0 -0
  17. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/contract/sim_config.py +0 -0
  18. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/contract/visualizer_classes.py +0 -0
  19. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/default_losses.py +0 -0
  20. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/default_metrics.py +0 -0
  21. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/experiment_api/__init__.py +0 -0
  22. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/experiment_api/api.py +0 -0
  23. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/experiment_api/cli_config_utils.py +0 -0
  24. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/experiment_api/client.py +0 -0
  25. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/experiment_api/epoch.py +0 -0
  26. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/experiment_api/experiment.py +0 -0
  27. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/experiment_api/experiment_context.py +0 -0
  28. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/experiment_api/types.py +0 -0
  29. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/experiment_api/utils.py +0 -0
  30. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/experiment_api/workingspace_config_utils.py +0 -0
  31. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/inner_leap_binder/__init__.py +0 -0
  32. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/mixpanel_tracker.py +0 -0
  33. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/plot_functions/__init__.py +0 -0
  34. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/plot_functions/plot_functions.py +0 -0
  35. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/plot_functions/visualize.py +0 -0
  36. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/code_loader/visualizers/__init__.py +0 -0
  37. {code_loader-1.0.208.dev3 → code_loader-1.0.208.dev5}/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.208.dev3
3
+ Version: 1.0.208.dev5
4
4
  Summary:
5
5
  Home-page: https://github.com/tensorleap/code-loader
6
6
  License: MIT
@@ -6,7 +6,7 @@ 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, LatentSpaceReduction
9
+ MetricDirection, DatasetMetadataType, LatentSpaceReduction, CustomLatentSpaceComputedAt
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
@@ -325,7 +325,7 @@ class CustomLatentSpaceHandler:
325
325
  name: str = 'custom_latent_space'
326
326
  use_ls_for_analysis: bool = False
327
327
  instance_aware: bool = False
328
- computed_at: str = 'dataset'
328
+ computed_at: CustomLatentSpaceComputedAt = CustomLatentSpaceComputedAt.DATASET
329
329
  arg_names: Optional[List[str]] = None
330
330
  reduce: Optional[LatentSpaceReduction] = None
331
331
  n_components: int = 512
@@ -73,3 +73,8 @@ class TestingSectionEnum(Enum):
73
73
  class LatentSpaceReduction(Enum):
74
74
  MEAN_POOL = 'MEAN_POOL'
75
75
  RANDOM_PROJECTION = 'RANDOM_PROJECTION'
76
+
77
+
78
+ class CustomLatentSpaceComputedAt(Enum):
79
+ DATASET = 'dataset'
80
+ MODEL = 'model'
@@ -22,7 +22,7 @@ from code_loader.contract.datasetclasses import SectionCallableInterface, InputH
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, LatentSpaceReduction
25
+ TestingSectionEnum, LatentSpaceReduction, CustomLatentSpaceComputedAt
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
@@ -561,7 +561,7 @@ class LeapBinder:
561
561
  name: Optional[str] = None,
562
562
  use_ls_for_analysis: bool = False,
563
563
  instance_aware: bool = False,
564
- computed_at: str = 'dataset',
564
+ computed_at: CustomLatentSpaceComputedAt = CustomLatentSpaceComputedAt.DATASET,
565
565
  arg_names: Optional[List[str]] = None,
566
566
  reduce: Optional[LatentSpaceReduction] = None,
567
567
  n_components: int = 512,
@@ -18,7 +18,7 @@ import numpy.typing as npt
18
18
  from code_loader.utils import map_dict_to_metadata_types, is_absent_metadata_value, \
19
19
  validate_autoregressive_state_types, autoregressive_nests_equal, \
20
20
  simulate_engine_float16_downcast_on_call_args, ENGINE_STORAGE_DTYPE, \
21
- TL_DISABLE_ENGINE_FLOAT16_SIMULATION_ENV_VAR
21
+ TL_DISABLE_ENGINE_FLOAT16_SIMULATION_ENV_VAR, sample_preprocess_response_arg_name
22
22
 
23
23
  logger = logging.getLogger(__name__)
24
24
 
@@ -29,7 +29,7 @@ from code_loader.contract.datasetclasses import CustomCallableInterfaceMultiArgs
29
29
  InstanceLengthCallableInterface, InstanceSectionCallableInterface, AutoregressiveStepCallableInterface, \
30
30
  MAX_CUSTOM_LATENT_SPACE_DIM, CUSTOM_LATENT_SPACE_WARN_DIM
31
31
  from code_loader.contract.enums import MetricDirection, LeapDataType, DatasetMetadataType, DataStateType, \
32
- DataStateEnum, LatentSpaceReduction
32
+ DataStateEnum, LatentSpaceReduction, CustomLatentSpaceComputedAt
33
33
  from code_loader import leap_binder, LeapLoader
34
34
  from code_loader.contract.mapping import NodeMapping, NodeMappingType, NodeConnection
35
35
  from code_loader.contract.visualizer_classes import LeapImage, LeapImageMask, LeapTextMask, LeapText, LeapGraph, \
@@ -291,11 +291,7 @@ def _require_sample_preprocess_response_supplied(user_function: Callable, args:
291
291
  """A SamplePreprocessResponse argument is auto-injected by the platform / check_dataset
292
292
  but NOT inside integration_test, where the author calls the function directly. Fail fast
293
293
  with an actionable message instead of a raw 'missing argument' TypeError."""
294
- spr_arg_name = None
295
- for arg_name, arg_type in inspect.getfullargspec(user_function).annotations.items():
296
- if arg_type == SamplePreprocessResponse:
297
- spr_arg_name = arg_name
298
- break
294
+ spr_arg_name = sample_preprocess_response_arg_name(user_function)
299
295
  if spr_arg_name is None:
300
296
  return
301
297
  signature = inspect.signature(user_function)
@@ -1786,16 +1782,25 @@ def _classify_custom_latent_space_signature(user_function) -> str:
1786
1782
  def _model_latent_space_arg_names(user_function) -> List[str]:
1787
1783
  argspec = inspect.getfullargspec(user_function)
1788
1784
  arg_names = list(argspec.args)
1789
- preprocess_response_arg_name = None
1785
+ spr_count = 0
1790
1786
  for arg_name, arg_type in argspec.annotations.items():
1787
+ if arg_name == 'return':
1788
+ continue
1791
1789
  _reject_stringized_sample_preprocess_response(user_function, arg_name, arg_type)
1792
1790
  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)
1791
+ spr_count += 1
1792
+ if spr_count > 1:
1793
+ raise Exception(
1794
+ f"tensorleap_custom_latent_space validation failed: only one argument of "
1795
+ f"'{user_function.__name__}' can be of type SamplePreprocessResponse.")
1796
+ spr_arg_name = sample_preprocess_response_arg_name(user_function)
1797
+ if spr_arg_name is not None:
1798
+ arg_names.remove(spr_arg_name)
1799
+ if not arg_names:
1800
+ raise Exception(
1801
+ f"tensorleap_custom_latent_space validation failed: '{user_function.__name__}' is "
1802
+ f"model-computed and expects at least one np.ndarray argument, but its signature "
1803
+ f"declares none.")
1799
1804
  return arg_names
1800
1805
 
1801
1806
 
@@ -1911,7 +1916,7 @@ def tensorleap_custom_latent_space(name: Optional[str] = None, use_ls_for_analys
1911
1916
 
1912
1917
  leap_binder.set_custom_latent_space(inner_without_validate, ls_name,
1913
1918
  use_ls_for_analysis=use_ls_for_analysis,
1914
- computed_at='dataset', reduce=reduce,
1919
+ computed_at=CustomLatentSpaceComputedAt.DATASET, reduce=reduce,
1915
1920
  n_components=n_components, channel_axis=channel_axis)
1916
1921
 
1917
1922
  def inner(sample_id, preprocess_response):
@@ -1937,6 +1942,9 @@ def _decorate_model_latent_space(user_function, ls_name, use_ls_for_analysis, re
1937
1942
  arg_names = _model_latent_space_arg_names(user_function)
1938
1943
 
1939
1944
  def _validate_input_args(*args, **kwargs):
1945
+ # Every bound argument must be an array: a ground truth is never passed as None, so a
1946
+ # model-computed LS cannot fall back from ground truth to predictions on unlabeled rows.
1947
+ # The engine skips the LS for those rows instead (see leaploader._check_model_latent_spaces).
1940
1948
  assert len(args) + len(kwargs) > 0, (
1941
1949
  f"tensorleap_custom_latent_space validation failed: '{user_function.__name__}' is "
1942
1950
  f"model-computed and expects at least one np.ndarray argument, but received none.")
@@ -1972,6 +1980,7 @@ def _decorate_model_latent_space(user_function, ls_name, use_ls_for_analysis, re
1972
1980
  _called_from_inside_tl_decorator += 1
1973
1981
 
1974
1982
  try:
1983
+ _require_sample_preprocess_response_supplied(user_function, args, kwargs)
1975
1984
  result = user_function(*args, **kwargs)
1976
1985
  finally:
1977
1986
  _called_from_inside_tl_decorator -= 1
@@ -1982,7 +1991,7 @@ def _decorate_model_latent_space(user_function, ls_name, use_ls_for_analysis, re
1982
1991
 
1983
1992
  leap_binder.set_custom_latent_space(inner_without_validate, ls_name,
1984
1993
  use_ls_for_analysis=use_ls_for_analysis,
1985
- computed_at='model', arg_names=arg_names, reduce=reduce,
1994
+ computed_at=CustomLatentSpaceComputedAt.MODEL, arg_names=arg_names, reduce=reduce,
1986
1995
  n_components=n_components, channel_axis=channel_axis)
1987
1996
 
1988
1997
  def inner(*args, **kwargs):
@@ -17,7 +17,9 @@ from code_loader.contract.datasetclasses import DatasetSample, DatasetBaseHandle
17
17
  PredictionTypeHandler, MetadataHandler, CustomLayerHandler, MetricHandler, VisualizerHandlerData, MetricHandlerData, \
18
18
  MetricCallableReturnType, CustomLossHandlerData, CustomLossHandler, RawInputsForHeatmap, SamplePreprocessResponse, \
19
19
  ElementInstance, custom_latent_space_attribute, DatasetIntegrationSetup, InstanceMetricHandler, _simulation_context
20
- from code_loader.contract.enums import DataStateEnum, TestingSectionEnum, DataStateType, DatasetMetadataType
20
+ from code_loader.contract.enums import DataStateEnum, TestingSectionEnum, DataStateType, DatasetMetadataType, \
21
+ CustomLatentSpaceComputedAt
22
+ from code_loader.contract.mapping import NodeMappingType
21
23
  from code_loader.contract.exceptions import DatasetScriptException
22
24
  from code_loader.contract.responsedataclasses import DatasetIntegParseResult, DatasetTestResultPayload, \
23
25
  DatasetPreprocess, DatasetSetup, DatasetInputInstance, DatasetOutputInstance, DatasetMetadataInstance, \
@@ -28,7 +30,8 @@ from code_loader.inner_leap_binder import global_leap_binder
28
30
  from code_loader.inner_leap_binder.leapbinder import mapping_runtime_mode_env_var_mame
29
31
  from code_loader.leaploaderbase import LeapLoaderBase
30
32
  from code_loader.utils import get_root_exception_file_and_line_number, get_metadata_type_from_variable, \
31
- validate_autoregressive_state_types, autoregressive_nests_equal, is_absent_metadata_value
33
+ validate_autoregressive_state_types, autoregressive_nests_equal, is_absent_metadata_value, \
34
+ sample_preprocess_response_arg_name
32
35
 
33
36
 
34
37
  def _serialize_sim_bounds(bounds) -> dict:
@@ -433,7 +436,7 @@ class LeapLoader(LeapLoaderBase):
433
436
  def _check_model_latent_spaces(self) -> Optional[DatasetTestResultPayload]:
434
437
  model_names = [
435
438
  name for name, handler in global_leap_binder.setup_container.custom_latent_spaces.items()
436
- if handler.computed_at == 'model'
439
+ if handler.computed_at == CustomLatentSpaceComputedAt.MODEL
437
440
  ]
438
441
  if not model_names:
439
442
  return None
@@ -450,6 +453,19 @@ class LeapLoader(LeapLoaderBase):
450
453
  f"grouped preprocess response (grouped: {grouped_states}). Use the "
451
454
  f"(sample_id, preprocess: PreprocessResponse) form instead."
452
455
  )
456
+ if global_leap_binder.setup_container.unlabeled_data_preprocess is not None:
457
+ gt_bound_names = [
458
+ connection.node.name for connection in global_leap_binder.latent_space_connections
459
+ if any(node_input.type == NodeMappingType.GroundTruth
460
+ for node_input in (connection.node_inputs or {}).values())
461
+ ]
462
+ if gt_bound_names:
463
+ test_result.display[TestingSectionEnum.Warnings.name] = (
464
+ f"Model-computed custom latent space(s) {gt_bound_names} read a ground truth, so "
465
+ f"they are skipped for unlabeled samples, which will have no vector in them. A "
466
+ f"latent space cannot fall back from ground truth to predictions; to cover "
467
+ f"unlabeled samples, add one bound only to model predictions."
468
+ )
453
469
  return test_result
454
470
 
455
471
  def _check_instance_custom_latent_spaces(self) -> Optional[DatasetTestResultPayload]:
@@ -800,10 +816,7 @@ class LeapLoader(LeapLoaderBase):
800
816
  @staticmethod
801
817
  def _get_preprocess_response_arg_name(
802
818
  func: Callable) -> Optional[str]:
803
- for arg_name, arg_type in inspect.getfullargspec(func).annotations.items():
804
- if arg_type == SamplePreprocessResponse:
805
- return arg_name
806
- return None
819
+ return sample_preprocess_response_arg_name(func)
807
820
 
808
821
  def run_custom_loss(self, custom_loss_name: str, sample_ids: np.array, state: DataStateEnum,
809
822
  input_tensors_by_arg_name: Dict[str, npt.NDArray[np.float32]]):
@@ -1419,7 +1432,7 @@ class LeapLoader(LeapLoaderBase):
1419
1432
  instance_id: Optional[int] = None) -> Optional[Dict[str, npt.NDArray[np.float32]]]:
1420
1433
  handlers = {handler_name: handler for handler_name, handler
1421
1434
  in global_leap_binder.setup_container.custom_latent_spaces.items()
1422
- if handler.computed_at != 'model'}
1435
+ if handler.computed_at != CustomLatentSpaceComputedAt.MODEL}
1423
1436
  if not handlers:
1424
1437
  return None
1425
1438
  if preprocess.is_grouped:
@@ -1461,20 +1474,14 @@ class LeapLoader(LeapLoaderBase):
1461
1474
  def get_dataset_custom_latent_space_names(self) -> Tuple[str, ...]:
1462
1475
  self.exec_script()
1463
1476
  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')
1477
+ if handler.computed_at != CustomLatentSpaceComputedAt.MODEL)
1471
1478
 
1472
1479
  @lru_cache()
1473
1480
  def get_custom_latent_space_specs(self) -> Dict[str, Dict[str, Any]]:
1474
1481
  self.exec_script()
1475
1482
  return {
1476
1483
  name: {
1477
- 'computed_at': handler.computed_at,
1484
+ 'computed_at': handler.computed_at.value,
1478
1485
  'arg_names': list(handler.arg_names or []),
1479
1486
  'reduce': handler.reduce.value if handler.reduce is not None else None,
1480
1487
  'n_components': handler.n_components,
@@ -1484,18 +1491,19 @@ class LeapLoader(LeapLoaderBase):
1484
1491
  }
1485
1492
 
1486
1493
  @lru_cache()
1487
- def get_model_latent_space_specs(self) -> Dict[str, Dict[str, Any]]:
1494
+ def get_model_custom_latent_space_specs(self) -> Dict[str, Dict[str, Any]]:
1488
1495
  return {name: spec for name, spec in self.get_custom_latent_space_specs().items()
1489
1496
  if spec['computed_at'] == 'model'}
1490
1497
 
1491
1498
  def run_model_latent_space(self, ls_name: str, sample_ids: np.array, state: DataStateEnum,
1492
1499
  input_tensors_by_arg_name: Dict[str, npt.NDArray[np.float32]]
1493
1500
  ) -> npt.NDArray[np.float32]:
1494
- self._preprocess_result()
1495
-
1501
+ self.exec_script()
1496
1502
  handler = global_leap_binder.setup_container.custom_latent_spaces[ls_name]
1497
1503
  preprocess_response_arg_name = self._get_preprocess_response_arg_name(handler.function)
1498
1504
 
1505
+ # Preprocess runs only when the function asks for a SamplePreprocessResponse; the metrics
1506
+ # pod that calls this has no other reason to pay for it.
1499
1507
  if preprocess_response_arg_name is not None:
1500
1508
  input_tensors_by_arg_name[preprocess_response_arg_name] = SamplePreprocessResponse(
1501
1509
  sample_ids, self._preprocess_result()[state])
@@ -234,20 +234,15 @@ class LeapLoaderBase:
234
234
  raise NotImplementedError(f'{type(self).__name__} does not implement '
235
235
  'get_dataset_custom_latent_space_names.')
236
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
237
  @abstractmethod
243
238
  def get_custom_latent_space_specs(self) -> Dict[str, Dict[str, Any]]:
244
239
  raise NotImplementedError(f'{type(self).__name__} does not implement '
245
240
  'get_custom_latent_space_specs.')
246
241
 
247
242
  @abstractmethod
248
- def get_model_latent_space_specs(self) -> Dict[str, Dict[str, Any]]:
243
+ def get_model_custom_latent_space_specs(self) -> Dict[str, Dict[str, Any]]:
249
244
  raise NotImplementedError(f'{type(self).__name__} does not implement '
250
- 'get_model_latent_space_specs.')
245
+ 'get_model_custom_latent_space_specs.')
251
246
 
252
247
  @abstractmethod
253
248
  def run_model_latent_space(self, ls_name: str, sample_ids: np.array, state: DataStateEnum,
@@ -1,3 +1,4 @@
1
+ import inspect
1
2
  import io
2
3
  import math
3
4
  import os
@@ -11,7 +12,7 @@ import numpy as np
11
12
  import numpy.typing as npt
12
13
 
13
14
  from code_loader.contract.datasetclasses import SectionCallableInterface, PreprocessResponse, \
14
- InstanceCallableInterface, ElementInstance
15
+ InstanceCallableInterface, ElementInstance, SamplePreprocessResponse
15
16
  from code_loader.contract.enums import DatasetMetadataType
16
17
 
17
18
 
@@ -246,3 +247,11 @@ def autoregressive_nests_equal(a: Any, b: Any) -> bool:
246
247
  math.isnan(a) and math.isnan(b):
247
248
  return True
248
249
  return bool(a == b)
250
+
251
+
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
257
+ return None
@@ -1,6 +1,6 @@
1
1
  [tool.poetry]
2
2
  name = "code-loader"
3
- version = "1.0.208.dev3"
3
+ version = "1.0.208.dev5"
4
4
  description = ""
5
5
  authors = ["dorhar <doron.harnoy@tensorleap.ai>"]
6
6
  license = "MIT"