Skip to content

Commit 88fe2a2

Browse files
committed
Merge remote-tracking branch 'origin/fix_leftovers'
2 parents bd3e7f3 + eba80e4 commit 88fe2a2

5 files changed

Lines changed: 21 additions & 26 deletions

File tree

README.md

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -10,9 +10,9 @@
1010

1111
This repository was tested on:
1212
* Python 3.8
13-
* CUDA 11.2
14-
* PyTorch 1.8.1
15-
* PyTorch Lightning 1.2.1
13+
* CUDA 11.4
14+
* PyTorch 1.11
15+
* PyTorch Lightning 1.6
1616

1717
Check `requirements.txt` for other essential modules.
1818

compute_class_weights.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55

66
from pycocotools.coco import COCO
77

8-
from utils.patches_datamodule import PatchesDataModule
8+
from utils.PAD_datamodule import PADDataModule
99
from utils.settings.config import CROP_ENCODING, LINEAR_ENCODER
1010

1111

@@ -43,7 +43,7 @@
4343
class_pixel_counts = pickle.load(open(pixel_cnts_name, 'rb'))
4444
else:
4545
# Create Data Module
46-
dm = PatchesDataModule(
46+
dm = PADDataModule(
4747
root_path_coco=root_path_coco,
4848
path_train=coco_train,
4949
path_val=coco_val,

pad_experiments.py

Lines changed: 8 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -5,7 +5,7 @@
55
In order to be used as model inputs, the training data are further split into
66
sequences using a rolling window of a fixed size.
77
'''
8-
import argparsse
8+
import argparse
99
from pathlib import Path
1010
from datetime import datetime
1111

@@ -18,12 +18,11 @@
1818
from pytorch_lightning import loggers as pl_loggers
1919
from pytorch_lightning.callbacks import EarlyStopping, ModelCheckpoint, LearningRateMonitor
2020
from pytorch_lightning.plugins import DDPPlugin
21-
from pytorch_lightning.plugins.precision.native_amp import NativeMixedPrecisionPlugin
2221
import torch
2322

2423
from utils.PAD_datamodule import PADDataModule
2524
from 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
2928
pl.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

requirements.txt

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@ pandas==1.3.3
99
Pillow==8.3.2
1010
pyarrow==5.0.0
1111
pycocotools==2.0.2
12-
pytorch-lightning==1.2.1
12+
pytorch-lightning==1.6.1
1313
scikit-learn==0.24.2
1414
scikit-multilearn=0.2.0
1515
seaborn==0.11.2
@@ -23,10 +23,10 @@ tensorboard-data-server==0.6.1
2323
tensorboard-plugin-wit==1.8.0
2424
tensorboardX==2.4
2525
timm==0.4.12
26-
torch==1.8.1+cu111
26+
torch==1.11.0+cu113
2727
torchlars==0.1.2
28-
torchmetrics==0.5.1
29-
torchvision==0.9.1+cu111
28+
torchmetrics==0.7.3
29+
torchvision==0.12.0+cu113
3030
tqdm==4.62.3
3131
transformers==4.6.1
3232
xarray==0.19.0

visualize_predictions.py

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -6,10 +6,10 @@
66
from mpl_toolkits.axes_grid1 import ImageGrid
77
from pycocotools.coco import COCO
88

9-
from model.convLSTM_lightning import ConvLSTM
10-
from model.tempCNN_lightning import TempCNN
11-
from model.convSTAR_lightning import ConvSTAR
12-
from model.unet_lightning import UNet
9+
from model.PAD_convLSTM import ConvLSTM
10+
from model.PAD_tempCNN import TempCNN
11+
from model.PAD_convSTAR import ConvSTAR
12+
from model.PAD_unet import UNet
1313

1414
import pytorch_lightning as pl
1515
import torch

0 commit comments

Comments
 (0)