MCPcopy Create free account
hub / github.com/modelscope/modelscope / __init__

Method __init__

modelscope/trainers/trainer.py:104–284  ·  view source on GitHub ↗
(
            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)

Source from the content-addressed store, hash-verified

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)

Callers

nothing calls this directly

Calls 15

rebuild_configMethod · 0.95
build_modelMethod · 0.95
get_preprocessorsMethod · 0.95
build_datasetMethod · 0.95
get_data_collatorMethod · 0.95
tune_moduleMethod · 0.95
register_hookMethod · 0.95
invoke_hookMethod · 0.95
is_dp_group_availableMethod · 0.95
get_metricsMethod · 0.95
print_cfgMethod · 0.95

Tested by

no test coverage detected