MCPcopy Create free account
hub / github.com/Pints-AI/1.5-Pints / main

Function main

pretrain/main.py:335–436  ·  view source on GitHub ↗
(
    fabric: lightning.Fabric,
    data_dir: Path,
    resume: Union[Path, Literal[False]],
    training_params: TrainingParams,
    ignore_index=-100,
)

Source from the content-addressed store, hash-verified

333
334
335def 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

Callers 1

setupFunction · 0.70

Calls 8

GPTClass · 0.90
num_parametersFunction · 0.90
create_dataloadersFunction · 0.85
trainFunction · 0.85
from_nameMethod · 0.45
applyMethod · 0.45
setupMethod · 0.45
loadMethod · 0.45

Tested by

no test coverage detected