(
fabric: lightning.Fabric,
data_dir: Path,
resume: Union[Path, Literal[False]],
training_params: TrainingParams,
ignore_index=-100,
)
| 333 | |
| 334 | |
| 335 | def main( |
| 336 | fabric: lightning.Fabric, |
| 337 | data_dir: Path, |
| 338 | resume: Union[Path, Literal[False]], |
| 339 | training_params: TrainingParams, |
| 340 | ignore_index=-100, |
| 341 | ): |
| 342 | monitor = Monitor( |
| 343 | fabric, |
| 344 | window_size=2, |
| 345 | time_unit='seconds', |
| 346 | log_iter_interval=training_params['log_iter_interval'], |
| 347 | ) |
| 348 | |
| 349 | # Create model out folder only for the first node. |
| 350 | if fabric.global_rank == 0: |
| 351 | training_params['out_dir'].mkdir(parents=True, exist_ok=True) |
| 352 | |
| 353 | config = Config.from_name(training_params['model_name']) |
| 354 | |
| 355 | train_dataloader, val_dataloader = create_dataloaders( |
| 356 | break_into_chunks=training_params['devices'], |
| 357 | batch_size=training_params['micro_batch_size'], |
| 358 | block_size=config.block_size, |
| 359 | fabric=fabric, |
| 360 | data_dir=data_dir, |
| 361 | seed=3407, |
| 362 | ) |
| 363 | |
| 364 | if train_dataloader is not None: |
| 365 | fabric.print('Train dataloader prepared successfully') |
| 366 | |
| 367 | if val_dataloader is None: |
| 368 | fabric.print('Setting up fabric dataloader for training only...') |
| 369 | train_dataloader = fabric.setup_dataloaders(train_dataloader) |
| 370 | else: |
| 371 | fabric.print('Validation dataloader prepared successfully.') |
| 372 | fabric.print('Setting up fabric dataloaders for training and validation...') |
| 373 | train_dataloader, val_dataloader = fabric.setup_dataloaders( |
| 374 | train_dataloader, val_dataloader |
| 375 | ) |
| 376 | |
| 377 | # same seed for every process to init model (FSDP) |
| 378 | fabric.seed_everything(3407) |
| 379 | |
| 380 | fabric.print(f'Loading model with {config.__dict__}') |
| 381 | |
| 382 | t0 = time.perf_counter() |
| 383 | |
| 384 | with fabric.init_module(empty_init=False): |
| 385 | model = GPT(config) |
| 386 | if ( |
| 387 | resume is False |
| 388 | ): # this means we are training from scratch, hence need to initalize the weights for all model layers |
| 389 | model.apply(partial(model._init_weights, n_layer=config.n_layer)) |
| 390 | |
| 391 | instantiation_time = time.perf_counter() - t0 |
| 392 |
no test coverage detected