Get the data collator for both training and evaluating. Args: data_collator: The input data_collator param. remove_unused_data: Remove the unused data with 'RemoveColumnsCollator'. Returns: The train_data_collator and eval_data_collator, can be No
(self, data_collator, remove_unused_data=False)
| 301 | self.model = self.to_parallel(self.model) |
| 302 | |
| 303 | def get_data_collator(self, data_collator, remove_unused_data=False): |
| 304 | """Get the data collator for both training and evaluating. |
| 305 | |
| 306 | Args: |
| 307 | data_collator: The input data_collator param. |
| 308 | remove_unused_data: Remove the unused data with 'RemoveColumnsCollator'. |
| 309 | Returns: |
| 310 | The train_data_collator and eval_data_collator, can be None. |
| 311 | """ |
| 312 | |
| 313 | train_data_collator, eval_data_collator = None, None |
| 314 | if isinstance(data_collator, Mapping): |
| 315 | if ConfigKeys.train in data_collator: |
| 316 | assert isinstance(data_collator[ConfigKeys.train], Callable) |
| 317 | train_data_collator = data_collator[ConfigKeys.train] |
| 318 | if ConfigKeys.val in data_collator: |
| 319 | assert isinstance(data_collator[ConfigKeys.val], Callable) |
| 320 | eval_data_collator = data_collator[ConfigKeys.val] |
| 321 | else: |
| 322 | collate_fn = default_collate if data_collator is None else data_collator |
| 323 | train_data_collator = collate_fn |
| 324 | eval_data_collator = collate_fn |
| 325 | |
| 326 | if remove_unused_data: |
| 327 | from modelscope.utils.data_collators import RemoveColumnsCollator |
| 328 | |
| 329 | def _set_signature_columns_if_needed(): |
| 330 | signature = inspect.signature(self.model.forward) |
| 331 | return list(signature.parameters.keys()) |
| 332 | |
| 333 | model_inputs = _set_signature_columns_if_needed() |
| 334 | train_data_collator = RemoveColumnsCollator( |
| 335 | train_data_collator, model_inputs) |
| 336 | eval_data_collator = RemoveColumnsCollator(eval_data_collator, |
| 337 | model_inputs) |
| 338 | return train_data_collator, eval_data_collator |
| 339 | |
| 340 | def init_dist(self, launcher=None): |
| 341 | """Init dist and returns the dist information. |
no test coverage detected