code-loader 1.0.195.dev1__tar.gz → 1.0.196.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.195.dev1 → code_loader-1.0.196.dev1}/PKG-INFO +1 -1
  2. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/contract/datasetclasses.py +45 -10
  3. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/inner_leap_binder/leapbinder.py +146 -5
  4. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/inner_leap_binder/leapbinder_decorators.py +859 -105
  5. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/leaploader.py +247 -64
  6. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/leaploaderbase.py +67 -3
  7. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/utils.py +57 -4
  8. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/pyproject.toml +1 -1
  9. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/LICENSE +0 -0
  10. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/README.md +0 -0
  11. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/__init__.py +0 -0
  12. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/contract/__init__.py +0 -0
  13. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/contract/enums.py +0 -0
  14. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/contract/exceptions.py +0 -0
  15. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/contract/mapping.py +0 -0
  16. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/contract/responsedataclasses.py +0 -0
  17. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/contract/sim_config.py +0 -0
  18. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/contract/visualizer_classes.py +0 -0
  19. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/default_losses.py +0 -0
  20. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/default_metrics.py +0 -0
  21. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/experiment_api/__init__.py +0 -0
  22. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/experiment_api/api.py +0 -0
  23. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/experiment_api/cli_config_utils.py +0 -0
  24. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/experiment_api/client.py +0 -0
  25. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/experiment_api/epoch.py +0 -0
  26. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/experiment_api/experiment.py +0 -0
  27. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/experiment_api/experiment_context.py +0 -0
  28. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/experiment_api/types.py +0 -0
  29. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/experiment_api/utils.py +0 -0
  30. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/experiment_api/workingspace_config_utils.py +0 -0
  31. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/inner_leap_binder/__init__.py +0 -0
  32. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/mixpanel_tracker.py +0 -0
  33. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/plot_functions/__init__.py +0 -0
  34. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/plot_functions/plot_functions.py +0 -0
  35. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/plot_functions/visualize.py +0 -0
  36. {code_loader-1.0.195.dev1 → code_loader-1.0.196.dev1}/code_loader/visualizers/__init__.py +0 -0
  37. {code_loader-1.0.195.dev1 → code_loader-1.0.196.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.195.dev1
3
+ Version: 1.0.196.dev1
4
4
  Summary:
5
5
  Home-page: https://github.com/tensorleap/code-loader
6
6
  License: MIT
@@ -1,6 +1,6 @@
1
1
  import warnings
2
2
  from dataclasses import dataclass, field
3
- from typing import Any, Callable, List, Optional, Dict, Union, Type, Literal, cast
3
+ from typing import Any, Callable, List, Optional, Dict, Union, Type, Literal, Tuple, cast
4
4
  import re
5
5
  import numpy as np
6
6
  import numpy.typing as npt
@@ -48,6 +48,8 @@ class PreprocessResponse:
48
48
  instance_to_sample_ids_mappings: Optional[Dict[str, str]] = None # in use only for element instance
49
49
  tl_generated: bool = False
50
50
  _grouped: bool = field(default=False, init=False, repr=False, compare=False)
51
+ _group_pos_cache: Optional[Dict[SampleId, Tuple[List[SampleId], int]]] = field(
52
+ default=None, init=False, repr=False, compare=False)
51
53
 
52
54
  def __post_init__(self) -> None:
53
55
  assert self.sample_ids_to_instance_mappings is None, f"Keep sample_ids_to_instance_mappings None when initializing PreprocessResponse"
@@ -138,15 +140,19 @@ SectionCallableInterface = Callable[[Union[int, str], PreprocessResponse], npt.N
138
140
  InstanceCallableInterface = Callable[[Union[int, str], PreprocessResponse, int], Optional[ElementInstance]]
139
141
  InstanceLengthCallableInterface = Callable[[Union[int, str], PreprocessResponse], int]
140
142
 
141
- # (sample_id, prev_inputs, prev_outputs, preprocess) -> next model inputs, or None to end the chain.
142
- # First call per chain receives prev_inputs=None, prev_outputs=None and returns the initial inputs.
143
- # Keys starting with '_' are passthrough state: never fed to the model, handed back verbatim in
144
- # prev_inputs on the next call. The hook must be stateless and deterministically seeded (the engine
145
- # may replay a step after crash recovery and relies on identical results).
143
+ # (sample_id, prev_inputs, prev_outputs, state, preprocess) -> (next model inputs | None, state).
144
+ # First call per chain receives prev_inputs=None, prev_outputs=None, state=None and returns the
145
+ # initial model inputs plus the initial state. Returning next_inputs=None ends the chain the
146
+ # terminating call still returns state, so the final step participates in state aggregation.
147
+ # State is an arbitrary nest of dicts/lists holding numpy arrays, numbers, strings or bools; it is
148
+ # never fed to the model and rides the engine's redis queue on every step, so accumulate
149
+ # reductions, not raw per-step tensors. The hook must be stateless and deterministically seeded
150
+ # (the engine may replay a step after crash recovery and relies on identical results).
151
+ AutoregressiveChainState = Any
146
152
  AutoregressiveStepCallableInterface = Callable[
147
153
  [Union[int, str], Optional[Dict[str, npt.NDArray[np.float32]]], Optional[Dict[str, npt.NDArray[np.float32]]],
148
- PreprocessResponse],
149
- Optional[Dict[str, npt.NDArray[np.float32]]]]
154
+ AutoregressiveChainState, PreprocessResponse],
155
+ Tuple[Optional[Dict[str, npt.NDArray[np.float32]]], AutoregressiveChainState]]
150
156
 
151
157
  MetadataSectionCallableInterface = Union[
152
158
  Callable[[Union[int, str], PreprocessResponse], int],
@@ -302,16 +308,42 @@ class CustomLatentSpaceHandler:
302
308
  # dominates the model inputs). 'mean': every latent space is the elementwise mean over all steps.
303
309
  AUTOREGRESSIVE_LATENT_SPACE_AGGREGATIONS = ('last_step', 'mean')
304
310
 
311
+ # Reserved argument names of autoregressive metrics/losses/visualizers, fed implicitly by the
312
+ # platform from the finished chain (per-chain dicts, no batch axis). Any other argument is wired
313
+ # to a ground-truth encoder through the integration test.
314
+ AUTOREGRESSIVE_IMPLICIT_ARG_NAMES = ('inputs', 'outputs', 'state')
315
+
305
316
 
306
317
  @dataclass
307
318
  class AutoregressiveStepHandler:
308
319
  function: AutoregressiveStepCallableInterface
309
320
  name: str = 'autoregressive_step'
310
- # Per-model-input shapes (unbatched, underscore passthrough keys excluded), discovered by running
311
- # the hook's first call at parse time. Fills the role InputHandler.shape plays for input encoders.
321
+ # Per-model-input shapes (unbatched), discovered by running the hook's first call at parse
322
+ # time. Fills the role InputHandler.shape plays for input encoders.
312
323
  input_shapes: Optional[Dict[str, List[int]]] = None
313
324
  latent_space_aggregation: str = 'last_step'
314
325
 
326
+
327
+ # Per-chain, unbatched callables: called once per finished chain with the final step's tensors —
328
+ # inputs/outputs/state are fed implicitly by the platform (dicts, no batch axis; outputs keyed by
329
+ # prediction-type names), remaining args are wired ground-truth encoder results.
330
+ @dataclass
331
+ class AutoregressiveMetricHandler:
332
+ metric_handler_data: MetricHandlerData
333
+ function: CustomCallableInterfaceMultiArgs
334
+
335
+
336
+ @dataclass
337
+ class AutoregressiveLossHandler:
338
+ custom_loss_handler_data: CustomLossHandlerData
339
+ function: CustomCallableInterface
340
+
341
+
342
+ @dataclass
343
+ class AutoregressiveVisualizerHandler:
344
+ visualizer_handler_data: VisualizerHandlerData
345
+ function: VisualizerCallableInterface
346
+
315
347
  @dataclass
316
348
  class PredictionTypeHandler:
317
349
  name: str
@@ -353,6 +385,9 @@ class DatasetIntegrationSetup:
353
385
  custom_latent_space: Optional[CustomLatentSpaceHandler] = None
354
386
  simulations: List[SimulationHandler] = field(default_factory=list)
355
387
  autoregressive_step: Optional[AutoregressiveStepHandler] = None
388
+ autoregressive_metrics: List[AutoregressiveMetricHandler] = field(default_factory=list)
389
+ autoregressive_losses: List[AutoregressiveLossHandler] = field(default_factory=list)
390
+ autoregressive_visualizers: List[AutoregressiveVisualizerHandler] = field(default_factory=list)
356
391
 
357
392
 
358
393
  @dataclass
@@ -1,6 +1,9 @@
1
1
 
2
+ import builtins
2
3
  import inspect
3
- from typing import Callable, List, Optional, Dict, Any, Type, Union, get_args, cast
4
+ import os
5
+ from contextlib import contextmanager
6
+ from typing import Callable, List, Optional, Dict, Any, Type, Union, get_args, cast, Iterator, Set
4
7
 
5
8
  import numpy as np
6
9
  import numpy.typing as npt
@@ -14,8 +17,10 @@ from code_loader.contract.datasetclasses import SectionCallableInterface, InputH
14
17
  RawInputsForHeatmap, VisualizerHandlerData, MetricHandlerData, CustomLossHandlerData, SamplePreprocessResponse, \
15
18
  ElementInstanceMasksHandler, InstanceCallableInterface, CustomLatentSpaceHandler, InstanceMetricHandler, \
16
19
  SimulationHandler, _simulation_context, AutoregressiveStepHandler, AutoregressiveStepCallableInterface, \
17
- AUTOREGRESSIVE_LATENT_SPACE_AGGREGATIONS
18
- from code_loader.contract.enums import LeapDataType, DataStateEnum, DataStateType, MetricDirection, DatasetMetadataType
20
+ AUTOREGRESSIVE_LATENT_SPACE_AGGREGATIONS, AUTOREGRESSIVE_IMPLICIT_ARG_NAMES, \
21
+ AutoregressiveMetricHandler, AutoregressiveLossHandler, AutoregressiveVisualizerHandler
22
+ from code_loader.contract.enums import LeapDataType, DataStateEnum, DataStateType, MetricDirection, DatasetMetadataType, \
23
+ TestingSectionEnum
19
24
  from code_loader.contract.mapping import NodeConnection, NodeMapping, NodeMappingType
20
25
  from code_loader.contract.responsedataclasses import DatasetTestResultPayload, LeapAnalysisConfiguration
21
26
  from code_loader.contract.visualizer_classes import map_leap_data_type_to_visualizer_class
@@ -32,6 +37,25 @@ from code_loader.visualizers.default_visualizers import DefaultVisualizer, \
32
37
  mapping_runtime_mode_env_var_mame = '__MAPPING_RUNTIME_MODE__'
33
38
 
34
39
 
40
+ @contextmanager
41
+ def _track_opened_files() -> Iterator[Set[str]]:
42
+ opened: Set[str] = set()
43
+ original_open = builtins.open
44
+
45
+ def tracking_open(file: Any, *args: Any, **kwargs: Any) -> Any:
46
+ mode = args[0] if args else kwargs.get('mode', 'r')
47
+ is_read_mode = isinstance(mode, str) and not any(c in mode for c in ('a', 'w', 'x'))
48
+ if is_read_mode and isinstance(file, (str, os.PathLike)):
49
+ opened.add(os.fspath(file))
50
+ return original_open(file, *args, **kwargs)
51
+
52
+ builtins.open = tracking_open
53
+ try:
54
+ yield opened
55
+ finally:
56
+ builtins.open = original_open
57
+
58
+
35
59
  def _stringized_annotation_type_name(annotation: Any) -> Optional[str]:
36
60
  """Bare type name if ``annotation`` is a stringized annotation (from ``from __future__
37
61
  import annotations`` or a quoted hint), else ``None``. code_loader inspects raw
@@ -548,11 +572,90 @@ class LeapBinder:
548
572
  # Builtin chain metadata, declared at parse time so it survives the reporter's
549
573
  # metadata type mapping; the placeholder values are overwritten by the engine when a
550
574
  # chain finalizes (realized length, truncated-by-safety-cap flag).
551
- def builtin_chain_metadata(idx: Any, preprocess: PreprocessResponse) -> Dict[str, Any]:
575
+ def builtin_chain_metadata(idx: Any, preprocess: PreprocessResponse) -> Any:
576
+ if isinstance(idx, list):
577
+ return [{'length': -1, 'truncated': False} for _ in idx]
552
578
  return {'length': -1, 'truncated': False}
553
579
 
554
580
  self.set_metadata(builtin_chain_metadata, 'builtin_chain')
555
581
 
582
+ def _autoregressive_arg_names(self, function: Callable[..., Any], decorator_name: str,
583
+ name: str) -> List[str]:
584
+ """Split an autoregressive callable's signature into implicit chain args (validated
585
+ present) and wired arg names (returned) — the wired args connect to ground-truth
586
+ encoders through the integration test."""
587
+ argspec = inspect.getfullargspec(function)
588
+ for arg_name, arg_type in argspec.annotations.items():
589
+ _reject_stringized_sample_preprocess_response(function, arg_name, arg_type)
590
+ if arg_type == SamplePreprocessResponse:
591
+ raise Exception(
592
+ f'{decorator_name} "{name}": SamplePreprocessResponse arguments are not '
593
+ f'supported for autoregressive decorators yet. Derive what you need inside '
594
+ f'the autoregressive step hook and carry it in the state.')
595
+ implicit = [arg_name for arg_name in argspec[0]
596
+ if arg_name in AUTOREGRESSIVE_IMPLICIT_ARG_NAMES]
597
+ if not implicit:
598
+ raise Exception(
599
+ f'{decorator_name} "{name}": the function must declare at least one of the '
600
+ f'implicit chain arguments {list(AUTOREGRESSIVE_IMPLICIT_ARG_NAMES)} — they are '
601
+ f'fed by the platform from the finished chain.')
602
+ return [arg_name for arg_name in argspec[0]
603
+ if arg_name not in AUTOREGRESSIVE_IMPLICIT_ARG_NAMES]
604
+
605
+ def _assert_unique_metric_name(self, name: str) -> None:
606
+ existing = [handler.metric_handler_data.name for handler in self.setup_container.metrics] \
607
+ + [handler.metric_handler_data.name
608
+ for handler in self.setup_container.autoregressive_metrics]
609
+ if name in existing:
610
+ raise Exception(f'Metric with name {name} already exists. Please choose another')
611
+
612
+ def _assert_unique_loss_name(self, name: str) -> None:
613
+ existing = [handler.custom_loss_handler_data.name
614
+ for handler in self.setup_container.custom_loss_handlers] \
615
+ + [handler.custom_loss_handler_data.name
616
+ for handler in self.setup_container.autoregressive_losses]
617
+ if name in existing:
618
+ raise Exception(f'Custom loss with name {name} already exists. Please choose another')
619
+
620
+ def _assert_unique_visualizer_name(self, name: str) -> None:
621
+ existing = [handler.visualizer_handler_data.name
622
+ for handler in self.setup_container.visualizers] \
623
+ + [handler.visualizer_handler_data.name
624
+ for handler in self.setup_container.autoregressive_visualizers]
625
+ if name in existing:
626
+ raise Exception(f'Visualizer with name {name} already exists. Please choose another')
627
+
628
+ def add_autoregressive_metric(self, function: CustomCallableInterfaceMultiArgs, name: str,
629
+ direction: Optional[Union[MetricDirection, Dict[str, MetricDirection]]]
630
+ = MetricDirection.Downward,
631
+ compute_insights: Optional[Union[bool, Dict[str, bool]]] = None) -> None:
632
+ self._assert_unique_metric_name(name)
633
+ arg_names = self._autoregressive_arg_names(function, 'tensorleap_autoregressive_metric',
634
+ name)
635
+ metric_handler_data = MetricHandlerData(name, arg_names, direction, compute_insights)
636
+ self.setup_container.autoregressive_metrics.append(
637
+ AutoregressiveMetricHandler(metric_handler_data, function))
638
+
639
+ def add_autoregressive_loss(self, function: CustomCallableInterface, name: str) -> None:
640
+ self._assert_unique_loss_name(name)
641
+ arg_names = self._autoregressive_arg_names(function, 'tensorleap_autoregressive_loss',
642
+ name)
643
+ self.setup_container.autoregressive_losses.append(
644
+ AutoregressiveLossHandler(CustomLossHandlerData(name, arg_names), function))
645
+
646
+ def add_autoregressive_visualizer(self, function: VisualizerCallableInterface, name: str,
647
+ visualizer_type: LeapDataType) -> None:
648
+ self._assert_unique_visualizer_name(name)
649
+ if visualizer_type.value not in map_leap_data_type_to_visualizer_class:
650
+ raise Exception(
651
+ f'The visualizer_type is invalid. current visualizer_type: {visualizer_type}, '
652
+ f'should be one of : {", ".join([arg.__name__ for arg in get_args(LeapData)])}')
653
+ arg_names = self._autoregressive_arg_names(
654
+ function, 'tensorleap_autoregressive_visualizer', name)
655
+ self.setup_container.autoregressive_visualizers.append(
656
+ AutoregressiveVisualizerHandler(VisualizerHandlerData(name, visualizer_type,
657
+ arg_names), function))
658
+
556
659
  def set_custom_layer(self, custom_layer: Type[Any], name: str, inspect_layer: bool = False,
557
660
  kernel_index: Optional[int] = None, use_custom_latent_space: bool = False) -> None:
558
661
  """
@@ -664,15 +767,34 @@ class LeapBinder:
664
767
  preprocess_response: PreprocessResponse, test_result: List[DatasetTestResultPayload],
665
768
  dataset_base_handler: Union[DatasetBaseHandler, MetadataHandler], state: DataStateEnum) -> List[DatasetTestResultPayload]:
666
769
  assert preprocess_response.sample_ids is not None
770
+ opened_files: Optional[Set[str]] = None
667
771
  if preprocess_response.is_grouped:
668
772
  # Grouped: probe with the first group, then reduce to a single sample's result so the
669
773
  # recorded shape/type is per-sample (matching the flat case), not the (B, *) batch.
670
774
  # (casts: sample_ids[0] is a group list here, and the grouped result is a per-sample list.)
671
775
  group = preprocess_response.sample_ids[0]
672
- raw_result = cast(Any, dataset_base_handler.function(cast(Any, group), preprocess_response))[0]
776
+ with _track_opened_files() as opened_files:
777
+ raw_result = cast(Any, dataset_base_handler.function(cast(Any, group), preprocess_response))[0]
673
778
  else:
674
779
  raw_result = dataset_base_handler.function(
675
780
  cast(Union[int, str], preprocess_response.sample_ids[0]), preprocess_response)
781
+
782
+ def _warn_if_group_spans_multiple_files(payloads: List[DatasetTestResultPayload]) -> None:
783
+ if opened_files is None or len(opened_files) <= 1:
784
+ return
785
+ shown = sorted(opened_files)[:5]
786
+ suffix = f" (+{len(opened_files) - 5} more)" if len(opened_files) > 5 else ""
787
+ warning = (
788
+ f"Encoding declared group 0 (size {len(cast(List[Any], group))}) for '{dataset_base_handler.name}' opened "
789
+ f"{len(opened_files)} distinct files: {shown}{suffix}. If this group is meant to be one "
790
+ f"physical source (e.g. one Parquet file), it may be grouped incorrectly - check your "
791
+ f"preprocess grouping logic. (Only files opened via Python's open() are tracked, so this "
792
+ f"can miss other I/O paths or be a false positive if grouping isn't meant to reflect file "
793
+ f"locality.)")
794
+ for payload in payloads:
795
+ existing = payload.display.get(TestingSectionEnum.Errors.name, '')
796
+ payload.display[TestingSectionEnum.Errors.name] = (existing + '\n' if existing else '') + warning
797
+
676
798
  handler_type = 'metadata' if isinstance(dataset_base_handler, MetadataHandler) else None
677
799
  if isinstance(dataset_base_handler, MetadataHandler):
678
800
  if isinstance(raw_result, dict):
@@ -712,6 +834,7 @@ class LeapBinder:
712
834
  else:
713
835
  if raw_result is None:
714
836
  if state != DataStateEnum.training:
837
+ _warn_if_group_spans_multiple_files(test_result)
715
838
  return test_result
716
839
 
717
840
  if dataset_base_handler.metadata_type is None:
@@ -733,6 +856,7 @@ class LeapBinder:
733
856
  # setting shape in setup for all encoders
734
857
  if isinstance(dataset_base_handler, (InputHandler, GroundTruthHandler)):
735
858
  dataset_base_handler.shape = result_shape
859
+ _warn_if_group_spans_multiple_files(test_result)
736
860
  return test_result
737
861
 
738
862
  def check_handlers(self, preprocess_result: Dict[DataStateEnum, PreprocessResponse]) -> None:
@@ -796,6 +920,23 @@ class LeapBinder:
796
920
  appear in any order in the integration file.
797
921
  """
798
922
  if self.setup_container.autoregressive_step is None:
923
+ declared_autoregressive = [
924
+ ('tensorleap_autoregressive_metric',
925
+ [handler.metric_handler_data.name
926
+ for handler in self.setup_container.autoregressive_metrics]),
927
+ ('tensorleap_autoregressive_loss',
928
+ [handler.custom_loss_handler_data.name
929
+ for handler in self.setup_container.autoregressive_losses]),
930
+ ('tensorleap_autoregressive_visualizer',
931
+ [handler.visualizer_handler_data.name
932
+ for handler in self.setup_container.autoregressive_visualizers]),
933
+ ]
934
+ for decorator_name, names in declared_autoregressive:
935
+ if names:
936
+ raise Exception(
937
+ f'{decorator_name} {names} requires a tensorleap_autoregressive_step '
938
+ f'hook — autoregressive decorators consume the finished chain the hook '
939
+ f'drives.')
799
940
  return
800
941
  if self.setup_container.inputs:
801
942
  raise Exception(