MCPcopy Create free account
hub / github.com/OpenSparseLLMs/Linear-MoE / finetune

Function finetune

linear_moe/finetune_utils.py:268–350  ·  view source on GitHub ↗

Main finetune function used across all tasks. Args: model (nn.Module): The model to fine-tune. optimizer (Optimizer): The optimizer to use for gradient updates. opt_param_scheduler (Optional): The optimizer parameter scheduler. forward_step (callable): The fo

(train_valid_datasets_provider,
             model_provider,
             model_type=ModelType.encoder_or_decoder,
             forward_step=_cross_entropy_forward_step,
             end_of_epoch_callback_provider=None,
             task_collate_fn=None)

Source from the content-addressed store, hash-verified

266
267
268def finetune(train_valid_datasets_provider,
269 model_provider,
270 model_type=ModelType.encoder_or_decoder,
271 forward_step=_cross_entropy_forward_step,
272 end_of_epoch_callback_provider=None,
273 task_collate_fn=None):
274 """
275 Main finetune function used across all tasks.
276 Args:
277 model (nn.Module): The model to fine-tune.
278 optimizer (Optimizer): The optimizer to use for gradient updates.
279 opt_param_scheduler (Optional): The optimizer parameter scheduler.
280 forward_step (callable): The forward step function for the model.
281 train_dataloader (DataLoader): The dataloader for training data.
282 valid_dataloader (DataLoader): The dataloader for validation data.
283 end_of_epoch_callback (Optional[callable]): The callback function to call at the end of each epoch.
284 """
285 args = get_args()
286 timers = get_timers()
287 assert args.rampup_batch_size is None, \
288 'batch size scaling is not supported for finetuning'
289
290 # Train and validation data loaders.
291 timers('train/valid/test dataset/dataloder', log_level=0).start()
292 if args.epochs > 0:
293 train_dataset, valid_dataset = train_valid_datasets_provider()
294 train_dataloader, valid_dataloader = _build_train_valid_dataloaders(
295 train_dataset, valid_dataset, task_collate_fn)
296 else:
297 args.train_iters = 0
298 timers('train/valid/test dataset/dataloder').stop()
299
300 # Build calback function.
301 timers('callback function', log_level=0).start()
302 end_of_epoch_callback = None
303 if end_of_epoch_callback_provider is not None:
304 end_of_epoch_callback = end_of_epoch_callback_provider()
305 timers('callback function').stop()
306
307 # Build model, optimizer and learning rate scheduler.
308 timers('model and optimizer', log_level=0).start()
309 model, optimizer, opt_param_scheduler = setup_model_and_optimizer(
310 model_provider, model_type)
311 timers('model and optimizer').stop()
312
313 # If pretrained checkpoint is provided and we have not trained for
314 # any iteration (i.e., iteration is zero), then load the pretrained
315 # checkpoint.
316 timers('pretrained checkpoint', log_level=0).start(barrier=True)
317 if args.iteration == 0 and args.pretrained_checkpoint is not None:
318 original_load = args.load
319 args.load = args.pretrained_checkpoint
320 original_rng = args.no_load_rng
321 args.no_load_rng = True
322 _ = load_checkpoint(model, None, None)
323 args.load = original_load
324 args.no_load_rng = original_rng
325 # This is critical when only model is loaded. We should make sure

Callers

nothing calls this directly

Calls 3

_trainFunction · 0.85
get_argsFunction · 0.50

Tested by

no test coverage detected