55In order to be used as model inputs, the training data are further split into
66sequences using a rolling window of a fixed size.
77'''
8- import argparsse
8+ import argparse
99from pathlib import Path
1010from datetime import datetime
1111
1818from pytorch_lightning import loggers as pl_loggers
1919from pytorch_lightning .callbacks import EarlyStopping , ModelCheckpoint , LearningRateMonitor
2020from pytorch_lightning .plugins import DDPPlugin
21- from pytorch_lightning .plugins .precision .native_amp import NativeMixedPrecisionPlugin
2221import torch
2322
2423from utils .PAD_datamodule import PADDataModule
2524from utils .tools import font_colors
26- from utils .settings .config import RANDOM_SEED , CROP_ENCODING , LINEAR_ENCODER , CLASS_WEIGHTS
25+ from utils .settings .config import RANDOM_SEED , CROP_ENCODING , LINEAR_ENCODER , CLASS_WEIGHTS , BANDS
2726
2827# Set seed for everything
2928pl .seed_everything (RANDOM_SEED )
@@ -385,7 +384,7 @@ def main():
385384
386385 if args .train :
387386 # Create Data Modules
388- dm = PatchesDataModule (
387+ dm = PADDataModule (
389388 root_path_coco = root_path_coco ,
390389 path_train = path_train ,
391390 path_val = path_val ,
@@ -421,15 +420,13 @@ def main():
421420 dirpath = run_path / 'checkpoints' ,
422421 monitor = monitor ,
423422 mode = 'min' ,
424- period = 1 ,
425423 save_top_k = - 1
426424 )
427425 )
428426
429427 tb_logger = pl_loggers .TensorBoardLogger (run_path / 'tensorboard' )
430428
431429 my_ddp = DDPPlugin (find_unused_parameters = True )
432- mixed_precision = NativeMixedPrecisionPlugin ()
433430
434431 trainer = pl .Trainer (gpus = args .num_gpus ,
435432 num_nodes = args .num_nodes ,
@@ -445,16 +442,15 @@ def main():
445442 checkpoint_callback = True ,
446443 resume_from_checkpoint = resume_from_checkpoint ,
447444 fast_dev_run = args .devtest ,
448- distributed_backend = 'ddp' if args .num_gpus > 1 else None ,
449- plugins = [my_ddp , mixed_precision ],
450- deterministic = True
445+ strategy = 'ddp' if args .num_gpus > 1 else None ,
446+ plugins = [my_ddp ]
451447 )
452448
453449 # Train model
454450 trainer .fit (model , datamodule = dm )
455451 else :
456452 # Create Data Module
457- dm = PatchesDataModule (
453+ dm = PADDataModule (
458454 root_path_coco = root_path_coco ,
459455 path_test = path_test ,
460456 group_freq = args .group_freq ,
@@ -488,9 +484,8 @@ def main():
488484 min_epochs = 1 ,
489485 max_epochs = 2 ,
490486 precision = 32 ,
491- distributed_backend = 'ddp' if args .num_gpus > 1 else None ,
492- plugins = [my_ddp ],
493- deterministic = True
487+ strategy = 'ddp' if args .num_gpus > 1 else None ,
488+ plugins = [my_ddp ]
494489 )
495490
496491 # Test model
0 commit comments