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)
| 266 | |
| 267 | |
| 268 | def 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 |
nothing calls this directly
no test coverage detected