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

Method get_data_collator

modelscope/trainers/trainer.py:303–338  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

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.

Callers 1

__init__Method · 0.95

Calls 1

Tested by

no test coverage detected