Ensure runtime kwargs reset to baseline defaults before inference.
(self)
| 1008 | ]) |
| 1009 | |
| 1010 | def _reset_runtime_configs(self): |
| 1011 | """Ensure runtime kwargs reset to baseline defaults before inference.""" |
| 1012 | base_map = getattr(self, "_base_kwargs_map", None) |
| 1013 | if not base_map: |
| 1014 | return |
| 1015 | |
| 1016 | for name, base in base_map.items(): |
| 1017 | restored = {} |
| 1018 | for k, v in base.items(): |
| 1019 | if k in self._IMMUTABLE_KWARGS_KEYS or not isinstance(v, (dict, list)): |
| 1020 | restored[k] = v |
| 1021 | else: |
| 1022 | restored[k] = copy.deepcopy(v) |
| 1023 | setattr(self, name, restored) |
| 1024 | |
| 1025 | ncpu = _resolve_ncpu(self.kwargs, 4) |
| 1026 | self.kwargs["ncpu"] = ncpu |
| 1027 | for name, value in base_map.items(): |
| 1028 | if name == "kwargs": |
| 1029 | continue |
| 1030 | config = getattr(self, name, None) |
| 1031 | if isinstance(config, dict): |
| 1032 | config.setdefault("ncpu", ncpu) |
| 1033 | if torch.get_num_threads() != ncpu: |
| 1034 | torch.set_num_threads(ncpu) |
no test coverage detected