(
self,
model: Optional[Union[TorchModel, nn.Module, str]] = None,
cfg_file: Optional[str] = None,
cfg_modify_fn: Optional[Callable] = None,
arg_parse_fn: Optional[Callable] = None,
data_collator: Optional[Union[Callable, Dict[str,
Callable]]] = None,
train_dataset: Optional[Union[MsDataset, Dataset]] = None,
eval_dataset: Optional[Union[MsDataset, Dataset]] = None,
preprocessor: Optional[Union[Preprocessor,
Dict[str, Preprocessor]]] = None,
optimizers: Tuple[torch.optim.Optimizer,
torch.optim.lr_scheduler._LRScheduler] = (None,
None),
model_revision: Optional[str] = DEFAULT_MODEL_REVISION,
seed: int = 42,
callbacks: Optional[List[Hook]] = None,
samplers: Optional[Union[Sampler, Dict[str, Sampler]]] = None,
efficient_tuners: Union[Dict[str, TunerConfig],
TunerConfig] = None,
**kwargs)
| 102 | """ |
| 103 | |
| 104 | def __init__( |
| 105 | self, |
| 106 | model: Optional[Union[TorchModel, nn.Module, str]] = None, |
| 107 | cfg_file: Optional[str] = None, |
| 108 | cfg_modify_fn: Optional[Callable] = None, |
| 109 | arg_parse_fn: Optional[Callable] = None, |
| 110 | data_collator: Optional[Union[Callable, Dict[str, |
| 111 | Callable]]] = None, |
| 112 | train_dataset: Optional[Union[MsDataset, Dataset]] = None, |
| 113 | eval_dataset: Optional[Union[MsDataset, Dataset]] = None, |
| 114 | preprocessor: Optional[Union[Preprocessor, |
| 115 | Dict[str, Preprocessor]]] = None, |
| 116 | optimizers: Tuple[torch.optim.Optimizer, |
| 117 | torch.optim.lr_scheduler._LRScheduler] = (None, |
| 118 | None), |
| 119 | model_revision: Optional[str] = DEFAULT_MODEL_REVISION, |
| 120 | seed: int = 42, |
| 121 | callbacks: Optional[List[Hook]] = None, |
| 122 | samplers: Optional[Union[Sampler, Dict[str, Sampler]]] = None, |
| 123 | efficient_tuners: Union[Dict[str, TunerConfig], |
| 124 | TunerConfig] = None, |
| 125 | **kwargs): |
| 126 | |
| 127 | self._seed = seed |
| 128 | set_random_seed(self._seed) |
| 129 | self._metric_values = None |
| 130 | self.optimizers = optimizers |
| 131 | self._mode = ModeKeys.TRAIN |
| 132 | self._hooks: List[Hook] = [] |
| 133 | self._epoch = 0 |
| 134 | self._iter = 0 |
| 135 | self._inner_iter = 0 |
| 136 | self._stop_training = False |
| 137 | self._compile = kwargs.get('compile', False) |
| 138 | self.trust_remote_code = kwargs.get('trust_remote_code', False) |
| 139 | |
| 140 | self.train_dataloader = None |
| 141 | self.eval_dataloader = None |
| 142 | self.data_loader = None |
| 143 | self._samplers = samplers |
| 144 | |
| 145 | if isinstance(model, str): |
| 146 | self.model_dir = self.get_or_download_model_dir( |
| 147 | model, model_revision, kwargs.pop(ThirdParty.KEY, None)) |
| 148 | if cfg_file is None: |
| 149 | cfg_file = os.path.join(self.model_dir, |
| 150 | ModelFile.CONFIGURATION) |
| 151 | self.input_model_id = model |
| 152 | else: |
| 153 | assert cfg_file is not None, 'Config file should not be None if model is not from pretrained!' |
| 154 | self.model_dir = os.path.dirname(cfg_file) |
| 155 | self.input_model_id = None |
| 156 | if hasattr(model, 'model_dir'): |
| 157 | check_local_model_is_latest( |
| 158 | model.model_dir, |
| 159 | user_agent={ |
| 160 | Invoke.KEY: Invoke.LOCAL_TRAINER, |
| 161 | ThirdParty.KEY: kwargs.pop(ThirdParty.KEY, None) |
nothing calls this directly
no test coverage detected