diff --git a/.gitignore b/.gitignore index f291a5db182ccc51a46fecbff981f3929fc79d0b..7b4961193d9cfc038579e8e5880b8a58862b56fb 100644 --- a/.gitignore +++ b/.gitignore @@ -15,5 +15,4 @@ docs/src/**/* */*.so* */**/*.so* */**/*.dylib* -*~ -models/pretrain \ No newline at end of file +*~ \ No newline at end of file diff --git a/README.md b/README.md index f0f620cf79dfc0665d6ea96829b9fd0577af6cb3..0bf793adc8a65941fd4785de51315f6ba524feba 100644 --- a/README.md +++ b/README.md @@ -2,9 +2,9 @@ Here, we provide the pytorch implementation of the paper: Remote Sensing Image Change Detection with Transformers. -![image-20210228153142126](./images/pipeline.png) +For more ore information, please see our published paper at [IEEE TGRS](https://ieeexplore.ieee.org/document/9491802) or [arxiv](https://arxiv.org/abs/2103.00208). -Code is coming soon~~ +![image-20210228153142126](./images/pipeline.png) ## Requirements @@ -13,7 +13,6 @@ Python 3.6 pytorch 1.6.0 torchvision 0.7.0 einops 0.3.0 -pthflops ``` ## Installation @@ -27,12 +26,87 @@ cd BIT_CD ## Quick Start +We have some samples from the [LEVIR-CD](https://justchenhao.github.io/LEVIR/) dataset in the folder `samples` for a quick start. + +Firstly, you can download our BIT pretrained model——by [baidu drive, code: 2lyz](https://pan.baidu.com/s/1HiXwpspl6odYQKda6pMuZQ) or [google drive](https://drive.google.com/file/d/1IVdF5a3e1_7DiSndtMkhpZuCSgDLLFcg/view?usp=sharing). After downloaded the pretrained model, you can put it in `checkpoints/BIT_LEVIR/`. + +Then, run a demo to get started as follows: + +```python +python demo.py +``` + +After that, you can find the prediction results in `samples/predict`. + ## Train -## Test +You can find the training script `run_cd.sh` in the folder `scripts`. You can run the script file by `sh scripts/run_cd.sh` in the command environment. + +The detailed script file `run_cd.sh` is as follows: + +```cmd +gpus=0 +checkpoint_root=checkpoints +data_name=LEVIR # dataset name + +img_size=256 +batch_size=8 +lr=0.01 +max_epochs=200 #training epochs +net_G=base_transformer_pos_s4_dd8 # model name +#base_resnet18 +#base_transformer_pos_s4_dd8 +#base_transformer_pos_s4_dd8_dedim8 +lr_policy=linear + +split=train # training txt +split_val=val #validation txt +project_name=CD_${net_G}_${data_name}_b${batch_size}_lr${lr}_${split}_${split_val}_${max_epochs}_${lr_policy} + +python main_cd.py --img_size ${img_size} --checkpoint_root ${checkpoint_root} --lr_policy ${lr_policy} --split ${split} --split_val ${split_val} --net_G ${net_G} --gpu_ids ${gpus} --max_epochs ${max_epochs} --project_name ${project_name} --batch_size ${batch_size} --data_name ${data_name} --lr ${lr} +``` + +## Evaluate + +You can find the evaluation script `eval.sh` in the folder `scripts`. You can run the script file by `sh scripts/eval.sh` in the command environment. + +The detailed script file `eval.sh` is as follows: + +```cmd +gpus=0 +data_name=LEVIR # dataset name +net_G=base_transformer_pos_s4_dd8_dedim8 # model name +split=test # test.txt +project_name=BIT_LEVIR # the name of the subfolder in the checkpoints folder +checkpoint_name=best_ckpt.pt # the name of evaluated model file + +python eval_cd.py --split ${split} --net_G ${net_G} --checkpoint_name ${checkpoint_name} --gpu_ids ${gpus} --project_name ${project_name} --data_name ${data_name} +``` ## Dataset Preparation +### Data structure + +``` +""" +Change detection data set with pixel-level binary labels; +├─A +├─B +├─label +└─list +""" +``` + +`A`: images of t1 phase; + +`B`:images of t2 phase; + +`label`: label maps; + +`list`: contains `train.txt, val.txt and test.txt`, each file records the image names (XXX.png) in the change detection dataset. + +### Data Download + LEVIR-CD: https://justchenhao.github.io/LEVIR/ WHU-CD: https://study.rsgis.whu.edu.cn/pages/download/building_dataset.html @@ -56,7 +130,7 @@ If you use this code for your research, please cite our paper: volume={}, number={}, pages={1-14}, - doi={} + doi={10.1109/TGRS.2021.3095166} } ``` diff --git a/checkpoints/BIT_LEVIR/best_ckpt.pt b/checkpoints/BIT_LEVIR/best_ckpt.pt new file mode 100644 index 0000000000000000000000000000000000000000..bebe18139dce5e9465be9bf29017c596f94ea7af --- /dev/null +++ b/checkpoints/BIT_LEVIR/best_ckpt.pt @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c159ba76143447f58c9f367ce8126a0014f2e4ba218cdb97cca173952c38cb3b +size 60048699 diff --git a/data_config.py b/data_config.py new file mode 100644 index 0000000000000000000000000000000000000000..ba0a4f1006b6e0e5966c6ed5c08c16c31c97cc45 --- /dev/null +++ b/data_config.py @@ -0,0 +1,22 @@ + +class DataConfig: + data_name = "" + root_dir = "" + label_transform = "norm" + def get_data_config(self, data_name): + self.data_name = data_name + if data_name == 'LEVIR': + self.root_dir = 'path to the root of LEVIR-CD dataset' + elif data_name == 'quick_start': + self.root_dir = './samples/' + else: + raise TypeError('%s has not defined' % data_name) + return self + + +if __name__ == '__main__': + data = DataConfig().get_data_config(data_name='LEVIR') + print(data.data_name) + print(data.root_dir) + print(data.label_transform) + diff --git a/datasets/CD_dataset.py b/datasets/CD_dataset.py new file mode 100644 index 0000000000000000000000000000000000000000..06463c21f95baf6a6c22da0e5dd8084520e6fe36 --- /dev/null +++ b/datasets/CD_dataset.py @@ -0,0 +1,120 @@ +""" +变化检测数据集 +""" + +import os +from PIL import Image +import numpy as np + +from torch.utils import data + +from datasets.data_utils import CDDataAugmentation + + +""" +CD data set with pixel-level labels; +├─image +├─image_post +├─label +└─list +""" +IMG_FOLDER_NAME = "A" +IMG_POST_FOLDER_NAME = 'B' +LIST_FOLDER_NAME = 'list' +ANNOT_FOLDER_NAME = "label" + +IGNORE = 255 + +label_suffix='.png' # jpg for gan dataset, others : png + +def load_img_name_list(dataset_path): + img_name_list = np.loadtxt(dataset_path, dtype=np.str) + if img_name_list.ndim == 2: + return img_name_list[:, 0] + return img_name_list + + +def load_image_label_list_from_npy(npy_path, img_name_list): + cls_labels_dict = np.load(npy_path, allow_pickle=True).item() + return [cls_labels_dict[img_name] for img_name in img_name_list] + + +def get_img_post_path(root_dir,img_name): + return os.path.join(root_dir, IMG_POST_FOLDER_NAME, img_name) + + +def get_img_path(root_dir, img_name): + return os.path.join(root_dir, IMG_FOLDER_NAME, img_name) + + +def get_label_path(root_dir, img_name): + return os.path.join(root_dir, ANNOT_FOLDER_NAME, img_name.replace('.jpg', label_suffix)) + + +class ImageDataset(data.Dataset): + """VOCdataloder""" + def __init__(self, root_dir, split='train', img_size=256, is_train=True,to_tensor=True): + super(ImageDataset, self).__init__() + self.root_dir = root_dir + self.img_size = img_size + self.split = split # train | train_aug | val + # self.list_path = self.root_dir + '/' + LIST_FOLDER_NAME + '/' + self.list + '.txt' + self.list_path = os.path.join(self.root_dir, LIST_FOLDER_NAME, self.split+'.txt') + self.img_name_list = load_img_name_list(self.list_path) + + self.A_size = len(self.img_name_list) # get the size of dataset A + self.to_tensor = to_tensor + if is_train: + self.augm = CDDataAugmentation( + img_size=self.img_size, + with_random_hflip=True, + with_random_vflip=True, + with_scale_random_crop=True, + with_random_blur=True, + ) + else: + self.augm = CDDataAugmentation( + img_size=self.img_size + ) + def __getitem__(self, index): + name = self.img_name_list[index] + A_path = get_img_path(self.root_dir, self.img_name_list[index % self.A_size]) + B_path = get_img_post_path(self.root_dir, self.img_name_list[index % self.A_size]) + + img = np.asarray(Image.open(A_path).convert('RGB')) + img_B = np.asarray(Image.open(B_path).convert('RGB')) + + [img, img_B], _ = self.augm.transform([img, img_B],[], to_tensor=self.to_tensor) + + return {'A': img, 'B': img_B, 'name': name} + + def __len__(self): + """Return the total number of images in the dataset.""" + return self.A_size + + +class CDDataset(ImageDataset): + + def __init__(self, root_dir, img_size, split='train', is_train=True, label_transform=None, + to_tensor=True): + super(CDDataset, self).__init__(root_dir, img_size=img_size, split=split, is_train=is_train, + to_tensor=to_tensor) + self.label_transform = label_transform + + def __getitem__(self, index): + name = self.img_name_list[index] + A_path = get_img_path(self.root_dir, self.img_name_list[index % self.A_size]) + B_path = get_img_post_path(self.root_dir, self.img_name_list[index % self.A_size]) + img = np.asarray(Image.open(A_path).convert('RGB')) + img_B = np.asarray(Image.open(B_path).convert('RGB')) + L_path = get_label_path(self.root_dir, self.img_name_list[index % self.A_size]) + + label = np.array(Image.open(L_path), dtype=np.uint8) + # 二分类中,前景标注为255 + if self.label_transform == 'norm': + label = label // 255 + + [img, img_B], [label] = self.augm.transform([img, img_B], [label], to_tensor=self.to_tensor) + # print(label.max()) + return {'name': name, 'A': img, 'B': img_B, 'L': label} + diff --git a/datasets/data_utils.py b/datasets/data_utils.py new file mode 100644 index 0000000000000000000000000000000000000000..49c9feaa202808b978ad8a8a9cb637cd7ba47421 --- /dev/null +++ b/datasets/data_utils.py @@ -0,0 +1,185 @@ +import random +import numpy as np + +from PIL import Image +from PIL import ImageFilter + +import torchvision.transforms.functional as TF +from torchvision import transforms +import torch + + +def to_tensor_and_norm(imgs, labels): + # to tensor + imgs = [TF.to_tensor(img) for img in imgs] + labels = [torch.from_numpy(np.array(img, np.uint8)).unsqueeze(dim=0) + for img in labels] + + imgs = [TF.normalize(img, mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5]) + for img in imgs] + return imgs, labels + + +class CDDataAugmentation: + + def __init__( + self, + img_size, + with_random_hflip=False, + with_random_vflip=False, + with_random_rot=False, + with_random_crop=False, + with_scale_random_crop=False, + with_random_blur=False, + ): + self.img_size = img_size + if self.img_size is None: + self.img_size_dynamic = True + else: + self.img_size_dynamic = False + self.with_random_hflip = with_random_hflip + self.with_random_vflip = with_random_vflip + self.with_random_rot = with_random_rot + self.with_random_crop = with_random_crop + self.with_scale_random_crop = with_scale_random_crop + self.with_random_blur = with_random_blur + def transform(self, imgs, labels, to_tensor=True): + """ + :param imgs: [ndarray,] + :param labels: [ndarray,] + :return: [ndarray,],[ndarray,] + """ + # resize image and covert to tensor + imgs = [TF.to_pil_image(img) for img in imgs] + if self.img_size is None: + self.img_size = None + + if not self.img_size_dynamic: + if imgs[0].size != (self.img_size, self.img_size): + imgs = [TF.resize(img, [self.img_size, self.img_size], interpolation=3) + for img in imgs] + else: + self.img_size = imgs[0].size[0] + + labels = [TF.to_pil_image(img) for img in labels] + if len(labels) != 0: + if labels[0].size != (self.img_size, self.img_size): + labels = [TF.resize(img, [self.img_size, self.img_size], interpolation=0) + for img in labels] + + random_base = 0.5 + if self.with_random_hflip and random.random() > 0.5: + imgs = [TF.hflip(img) for img in imgs] + labels = [TF.hflip(img) for img in labels] + + if self.with_random_vflip and random.random() > 0.5: + imgs = [TF.vflip(img) for img in imgs] + labels = [TF.vflip(img) for img in labels] + + if self.with_random_rot and random.random() > random_base: + angles = [90, 180, 270] + index = random.randint(0, 2) + angle = angles[index] + imgs = [TF.rotate(img, angle) for img in imgs] + labels = [TF.rotate(img, angle) for img in labels] + + if self.with_random_crop and random.random() > 0: + i, j, h, w = transforms.RandomResizedCrop(size=self.img_size). \ + get_params(img=imgs[0], scale=(0.8, 1.0), ratio=(1, 1)) + + imgs = [TF.resized_crop(img, i, j, h, w, + size=(self.img_size, self.img_size), + interpolation=Image.CUBIC) + for img in imgs] + + labels = [TF.resized_crop(img, i, j, h, w, + size=(self.img_size, self.img_size), + interpolation=Image.NEAREST) + for img in labels] + + if self.with_scale_random_crop: + # rescale + scale_range = [1, 1.2] + target_scale = scale_range[0] + random.random() * (scale_range[1] - scale_range[0]) + + imgs = [pil_rescale(img, target_scale, order=3) for img in imgs] + labels = [pil_rescale(img, target_scale, order=0) for img in labels] + # crop + imgsize = imgs[0].size # h, w + box = get_random_crop_box(imgsize=imgsize, cropsize=self.img_size) + imgs = [pil_crop(img, box, cropsize=self.img_size, default_value=0) + for img in imgs] + labels = [pil_crop(img, box, cropsize=self.img_size, default_value=255) + for img in labels] + + if self.with_random_blur and random.random() > 0: + radius = random.random() + imgs = [img.filter(ImageFilter.GaussianBlur(radius=radius)) + for img in imgs] + + if to_tensor: + # to tensor + imgs = [TF.to_tensor(img) for img in imgs] + labels = [torch.from_numpy(np.array(img, np.uint8)).unsqueeze(dim=0) + for img in labels] + + imgs = [TF.normalize(img, mean=[0.5, 0.5, 0.5],std=[0.5, 0.5, 0.5]) + for img in imgs] + + return imgs, labels + + +def pil_crop(image, box, cropsize, default_value): + assert isinstance(image, Image.Image) + img = np.array(image) + + if len(img.shape) == 3: + cont = np.ones((cropsize, cropsize, img.shape[2]), img.dtype)*default_value + else: + cont = np.ones((cropsize, cropsize), img.dtype)*default_value + cont[box[0]:box[1], box[2]:box[3]] = img[box[4]:box[5], box[6]:box[7]] + + return Image.fromarray(cont) + + +def get_random_crop_box(imgsize, cropsize): + h, w = imgsize + ch = min(cropsize, h) + cw = min(cropsize, w) + + w_space = w - cropsize + h_space = h - cropsize + + if w_space > 0: + cont_left = 0 + img_left = random.randrange(w_space + 1) + else: + cont_left = random.randrange(-w_space + 1) + img_left = 0 + + if h_space > 0: + cont_top = 0 + img_top = random.randrange(h_space + 1) + else: + cont_top = random.randrange(-h_space + 1) + img_top = 0 + + return cont_top, cont_top+ch, cont_left, cont_left+cw, img_top, img_top+ch, img_left, img_left+cw + + +def pil_rescale(img, scale, order): + assert isinstance(img, Image.Image) + height, width = img.size + target_size = (int(np.round(height*scale)), int(np.round(width*scale))) + return pil_resize(img, target_size, order) + + +def pil_resize(img, size, order): + assert isinstance(img, Image.Image) + if size[0] == img.size[0] and size[1] == img.size[1]: + return img + if order == 3: + resample = Image.BICUBIC + elif order == 0: + resample = Image.NEAREST + return img.resize(size[::-1], resample) diff --git a/demo.py b/demo.py new file mode 100644 index 0000000000000000000000000000000000000000..ce552aaa2c96d5bbd394c8dbaf5b1f611b081ffa --- /dev/null +++ b/demo.py @@ -0,0 +1,79 @@ +from argparse import ArgumentParser + +import utils +import torch +from models.basic_model import CDEvaluator + +import os + +""" +quick start + +sample files in ./samples + +save prediction files in the ./samples/predict + +""" + + +def get_args(): + # ------------ + # args + # ------------ + parser = ArgumentParser() + parser.add_argument('--project_name', default='BIT_LEVIR', type=str) + parser.add_argument('--gpu_ids', type=str, default='0', help='gpu ids: e.g. 0 0,1,2, 0,2. use -1 for CPU') + parser.add_argument('--checkpoint_root', default='checkpoints', type=str) + parser.add_argument('--output_folder', default='samples/predict', type=str) + + # data + parser.add_argument('--num_workers', default=0, type=int) + parser.add_argument('--dataset', default='CDDataset', type=str) + parser.add_argument('--data_name', default='quick_start', type=str) + + parser.add_argument('--batch_size', default=1, type=int) + parser.add_argument('--split', default="demo", type=str) + parser.add_argument('--img_size', default=256, type=int) + + # model + parser.add_argument('--n_class', default=2, type=int) + parser.add_argument('--net_G', default='base_transformer_pos_s4_dd8_dedim8', type=str, + help='base_resnet18 | base_transformer_pos_s4_dd8 | base_transformer_pos_s4_dd8_dedim8|') + parser.add_argument('--checkpoint_name', default='best_ckpt.pt', type=str) + + args = parser.parse_args() + return args + + +if __name__ == '__main__': + + args = get_args() + utils.get_device(args) + device = torch.device("cuda:%s" % args.gpu_ids[0] + if torch.cuda.is_available() and len(args.gpu_ids)>0 + else "cpu") + args.checkpoint_dir = os.path.join(args.checkpoint_root, args.project_name) + os.makedirs(args.output_folder, exist_ok=True) + + log_path = os.path.join(args.output_folder, 'log_vis.txt') + + data_loader = utils.get_loader(args.data_name, img_size=args.img_size, + batch_size=args.batch_size, + split=args.split, is_train=False) + + model = CDEvaluator(args) + model.load_checkpoint(args.checkpoint_name) + model.eval() + + for i, batch in enumerate(data_loader): + name = batch['name'] + print('process: %s' % name) + score_map = model._forward_pass(batch) + model._save_predictions() + + + + + + + diff --git a/eval_cd.py b/eval_cd.py new file mode 100644 index 0000000000000000000000000000000000000000..ec8bc177493bd5eba9db2817db49b760c50ff1b9 --- /dev/null +++ b/eval_cd.py @@ -0,0 +1,59 @@ +from argparse import ArgumentParser +import torch +from models.evaluator import * + +print(torch.cuda.is_available()) + + +""" +eval the CD model +""" + +def main(): + # ------------ + # args + # ------------ + parser = ArgumentParser() + parser.add_argument('--gpu_ids', type=str, default='0', help='gpu ids: e.g. 0 0,1,2, 0,2. use -1 for CPU') + parser.add_argument('--project_name', default='test', type=str) + parser.add_argument('--print_models', default=False, type=bool, help='print models') + + # data + parser.add_argument('--num_workers', default=4, type=int) + parser.add_argument('--dataset', default='CDDataset', type=str) + parser.add_argument('--data_name', default='LEVIR', type=str) + + parser.add_argument('--batch_size', default=8, type=int) + parser.add_argument('--split', default="test", type=str) + + parser.add_argument('--img_size', default=256, type=int) + + # model + parser.add_argument('--n_class', default=2, type=int) + parser.add_argument('--net_G', default='base_transformer_pos_s4_dd8_dedim8', type=str, + help='base_resnet18 | base_transformer_pos_s4_dd8 | base_transformer_pos_s4_dd8_dedim8|') + + parser.add_argument('--checkpoint_name', default='best_ckpt.pt', type=str) + + args = parser.parse_args() + utils.get_device(args) + print(args.gpu_ids) + + # checkpoints dir + args.checkpoint_dir = os.path.join('checkpoints', args.project_name) + os.makedirs(args.checkpoint_dir, exist_ok=True) + # visualize dir + args.vis_dir = os.path.join('vis', args.project_name) + os.makedirs(args.vis_dir, exist_ok=True) + + dataloader = utils.get_loader(args.data_name, img_size=args.img_size, + batch_size=args.batch_size, is_train=False, + split=args.split) + model = CDEvaluator(args=args, dataloader=dataloader) + + model.eval_models(checkpoint_name=args.checkpoint_name) + + +if __name__ == '__main__': + main() + diff --git a/main_cd.py b/main_cd.py new file mode 100644 index 0000000000000000000000000000000000000000..320a5f4a1d9893a621489bcac2fc9caa284589e9 --- /dev/null +++ b/main_cd.py @@ -0,0 +1,77 @@ +from argparse import ArgumentParser +import torch +from models.trainer import * + +print(torch.cuda.is_available()) + +""" +the main function for training the CD networks +""" + + +def train(args): + dataloaders = utils.get_loaders(args) + model = CDTrainer(args=args, dataloaders=dataloaders) + model.train_models() + + +def test(args): + from models.evaluator import CDEvaluator + dataloader = utils.get_loader(args.data_name, img_size=args.img_size, + batch_size=args.batch_size, is_train=False, + split='test') + model = CDEvaluator(args=args, dataloader=dataloader) + + model.eval_models() + + +if __name__ == '__main__': + # ------------ + # args + # ------------ + parser = ArgumentParser() + parser.add_argument('--gpu_ids', type=str, default='0', help='gpu ids: e.g. 0 0,1,2, 0,2. use -1 for CPU') + parser.add_argument('--project_name', default='test', type=str) + parser.add_argument('--checkpoint_root', default='checkpoints', type=str) + + # data + parser.add_argument('--num_workers', default=4, type=int) + parser.add_argument('--dataset', default='CDDataset', type=str) + parser.add_argument('--data_name', default='LEVIR', type=str) + + parser.add_argument('--batch_size', default=8, type=int) + parser.add_argument('--split', default="train", type=str) + parser.add_argument('--split_val', default="val", type=str) + + parser.add_argument('--img_size', default=256, type=int) + + # model + parser.add_argument('--n_class', default=2, type=int) + parser.add_argument('--net_G', default='base_transformer_pos_s4_dd8', type=str, + help='base_resnet18 | base_transformer_pos_s4 | ' + 'base_transformer_pos_s4_dd8 | ' + 'base_transformer_pos_s4_dd8_dedim8|') + parser.add_argument('--loss', default='ce', type=str) + + # optimizer + parser.add_argument('--optimizer', default='sgd', type=str) + parser.add_argument('--lr', default=0.01, type=float) + parser.add_argument('--max_epochs', default=100, type=int) + parser.add_argument('--lr_policy', default='linear', type=str, + help='linear | step') + parser.add_argument('--lr_decay_iters', default=100, type=int) + + args = parser.parse_args() + utils.get_device(args) + print(args.gpu_ids) + + # checkpoints dir + args.checkpoint_dir = os.path.join(args.checkpoint_root, args.project_name) + os.makedirs(args.checkpoint_dir, exist_ok=True) + # visualize dir + args.vis_dir = os.path.join('vis', args.project_name) + os.makedirs(args.vis_dir, exist_ok=True) + + train(args) + + test(args) diff --git a/misc/imutils.py b/misc/imutils.py new file mode 100644 index 0000000000000000000000000000000000000000..d9c9028043f4c03096ec296ad21de085debec363 --- /dev/null +++ b/misc/imutils.py @@ -0,0 +1,401 @@ +import random +import numpy as np +import cv2 +from PIL import Image +from PIL import ImageFilter +import PIL +import tifffile + + +def cv_rotate(image, angle, borderValue): + """ + rot angle, fill with borderValue + """ + # grab the dimensions of the image and then determine the + # center + (h, w) = image.shape[:2] + (cX, cY) = (w // 2, h // 2) + + # grab the rotation matrix (applying the negative of the + # angle to rotate clockwise), then grab the sine and cosine + # (i.e., the rotation components of the matrix) + # -angle位置参数为角度参数负值表示顺时针旋转; 1.0位置参数scale是调整尺寸比例(图像缩放参数),建议0.75 + M = cv2.getRotationMatrix2D((cX, cY), -angle, 1.0) + cos = np.abs(M[0, 0]) + sin = np.abs(M[0, 1]) + + # compute the new bounding dimensions of the image + nW = int((h * sin) + (w * cos)) + nH = int((h * cos) + (w * sin)) + + # adjust the rotation matrix to take into account translation + M[0, 2] += (nW / 2) - cX + M[1, 2] += (nH / 2) - cY + if isinstance(borderValue, int): + values = (borderValue, borderValue, borderValue) + else: + values = borderValue + # perform the actual rotation and return the image + return cv2.warpAffine(image, M, (nW, nH), borderValue=values) + + +def pil_resize(img, size, order): + if size[0] == img.shape[0] and size[1] == img.shape[1]: + return img + + if order == 3: + resample = Image.BICUBIC + elif order == 0: + resample = Image.NEAREST + + return np.asarray(Image.fromarray(img).resize(size[::-1], resample)) + + +def pil_rescale(img, scale, order): + height, width = img.shape[:2] + target_size = (int(np.round(height*scale)), int(np.round(width*scale))) + return pil_resize(img, target_size, order) + + +def pil_rotate(img, degree, default_value): + if isinstance(default_value, tuple): + values = (default_value[0], default_value[1], default_value[2], 0) + else: + values = (default_value, default_value, default_value,0) + img = Image.fromarray(img) + if img.mode =='RGB': + # set img padding == default_value + img2 = img.convert('RGBA') + rot = img2.rotate(degree, expand=1) + fff = Image.new('RGBA', rot.size, values) # 灰色 + out = Image.composite(rot, fff, rot) + img = out.convert(img.mode) + + else: + # set label padding == default_value + img2 = img.convert('RGBA') + rot = img2.rotate(degree, expand=1) + # a white image same size as rotated image + fff = Image.new('RGBA', rot.size, values) + # create a composite image using the alpha layer of rot as a mask + out = Image.composite(rot, fff, rot) + img = out.convert(img.mode) + + return np.asarray(img) + + +def random_resize_long_image_list(img_list, min_long, max_long): + target_long = random.randint(min_long, max_long) + h, w = img_list[0].shape[:2] + if w < h: + scale = target_long / h + else: + scale = target_long / w + out = [] + for img in img_list: + out.append(pil_rescale(img, scale, 3) ) + return out + + +def random_resize_long(img, min_long, max_long): + target_long = random.randint(min_long, max_long) + h, w = img.shape[:2] + + if w < h: + scale = target_long / h + else: + scale = target_long / w + + return pil_rescale(img, scale, 3) + + +def random_scale_list(img_list, scale_range, order): + """ + 输入:图像列表 + """ + target_scale = scale_range[0] + random.random() * (scale_range[1] - scale_range[0]) + + if isinstance(img_list, tuple): + assert img_list.__len__() == 2 + img1 = [] + img2 = [] + for img in img_list[0]: + img1.append(pil_rescale(img, target_scale, order[0])) + for img in img_list[1]: + img2.append(pil_rescale(img, target_scale, order[1])) + return (img1, img2) + else: + out = [] + for img in img_list: + out.append(pil_rescale(img, target_scale, order)) + return out + + +def random_scale(img, scale_range, order): + + target_scale = scale_range[0] + random.random() * (scale_range[1] - scale_range[0]) + + if isinstance(img, tuple): + return (pil_rescale(img[0], target_scale, order[0]), pil_rescale(img[1], target_scale, order[1])) + else: + return pil_rescale(img, target_scale, order) + + +def random_rotate_list(img_list, max_degree, default_values): + degree = random.random() * max_degree + if isinstance(img_list, tuple): + assert img_list.__len__() == 2 + img1 = [] + img2 = [] + for img in img_list[0]: + assert isinstance(img, np.ndarray) + img1.append((pil_rotate(img, degree, default_values[0]))) + for img in img_list[1]: + img2.append((pil_rotate(img, degree, default_values[1]))) + return (img1, img2) + else: + out = [] + for img in img_list: + out.append(pil_rotate(img, degree, default_values)) + return out + + +def random_rotate(img, max_degree, default_values): + degree = random.random() * max_degree + if isinstance(img, tuple): + return (pil_rotate(img[0], degree, default_values[0]), + pil_rotate(img[1], degree, default_values[1])) + else: + return pil_rotate(img, degree, default_values) + + +def random_lr_flip_list(img_list): + + if bool(random.getrandbits(1)): + if isinstance(img_list, tuple): + assert img_list.__len__()==2 + img1=list((np.fliplr(m) for m in img_list[0])) + img2=list((np.fliplr(m) for m in img_list[1])) + + return (img1, img2) + else: + return list([np.fliplr(m) for m in img_list]) + else: + return img_list + + +def random_lr_flip(img): + + if bool(random.getrandbits(1)): + if isinstance(img, tuple): + return tuple([np.fliplr(m) for m in img]) + else: + return np.fliplr(img) + else: + return img + + +def get_random_crop_box(imgsize, cropsize): + h, w = imgsize + + ch = min(cropsize, h) + cw = min(cropsize, w) + + w_space = w - cropsize + h_space = h - cropsize + + if w_space > 0: + cont_left = 0 + img_left = random.randrange(w_space + 1) + else: + cont_left = random.randrange(-w_space + 1) + img_left = 0 + + if h_space > 0: + cont_top = 0 + img_top = random.randrange(h_space + 1) + else: + cont_top = random.randrange(-h_space + 1) + img_top = 0 + + return cont_top, cont_top+ch, cont_left, cont_left+cw, img_top, img_top+ch, img_left, img_left+cw + + +def random_crop_list(images_list, cropsize, default_values): + + if isinstance(images_list, tuple): + imgsize = images_list[0][0].shape[:2] + elif isinstance(images_list, list): + imgsize = images_list[0].shape[:2] + else: + raise RuntimeError('do not support the type of image_list') + if isinstance(default_values, int): default_values = (default_values,) + + box = get_random_crop_box(imgsize, cropsize) + if isinstance(images_list, tuple): + assert images_list.__len__()==2 + img1 = [] + img2 = [] + for img in images_list[0]: + f = default_values[0] + if len(img.shape) == 3: + cont = np.ones((cropsize, cropsize, img.shape[2]), img.dtype)*f + else: + cont = np.ones((cropsize, cropsize), img.dtype)*f + cont[box[0]:box[1], box[2]:box[3]] = img[box[4]:box[5], box[6]:box[7]] + img1.append(cont) + for img in images_list[1]: + f = default_values[1] + if len(img.shape) == 3: + cont = np.ones((cropsize, cropsize, img.shape[2]), img.dtype)*f + else: + cont = np.ones((cropsize, cropsize), img.dtype)*f + cont[box[0]:box[1], box[2]:box[3]] = img[box[4]:box[5], box[6]:box[7]] + img2.append(cont) + return (img1, img2) + else: + out = [] + for img in images_list: + f = default_values + if len(img.shape) == 3: + cont = np.ones((cropsize, cropsize, img.shape[2]), img.dtype) * f + else: + cont = np.ones((cropsize, cropsize), img.dtype) * f + cont[box[0]:box[1], box[2]:box[3]] = img[box[4]:box[5], box[6]:box[7]] + out.append(cont) + return out + + +def random_crop(images, cropsize, default_values): + + if isinstance(images, np.ndarray): images = (images,) + if isinstance(default_values, int): default_values = (default_values,) + + imgsize = images[0].shape[:2] + box = get_random_crop_box(imgsize, cropsize) + + new_images = [] + for img, f in zip(images, default_values): + + if len(img.shape) == 3: + cont = np.ones((cropsize, cropsize, img.shape[2]), img.dtype)*f + else: + cont = np.ones((cropsize, cropsize), img.dtype)*f + cont[box[0]:box[1], box[2]:box[3]] = img[box[4]:box[5], box[6]:box[7]] + new_images.append(cont) + + if len(new_images) == 1: + new_images = new_images[0] + + return new_images + + +def top_left_crop(img, cropsize, default_value): + + h, w = img.shape[:2] + + ch = min(cropsize, h) + cw = min(cropsize, w) + + if len(img.shape) == 2: + container = np.ones((cropsize, cropsize), img.dtype)*default_value + else: + container = np.ones((cropsize, cropsize, img.shape[2]), img.dtype)*default_value + + container[:ch, :cw] = img[:ch, :cw] + + return container + + +def center_crop(img, cropsize, default_value=0): + + h, w = img.shape[:2] + + ch = min(cropsize, h) + cw = min(cropsize, w) + + sh = h - cropsize + sw = w - cropsize + + if sw > 0: + cont_left = 0 + img_left = int(round(sw / 2)) + else: + cont_left = int(round(-sw / 2)) + img_left = 0 + + if sh > 0: + cont_top = 0 + img_top = int(round(sh / 2)) + else: + cont_top = int(round(-sh / 2)) + img_top = 0 + + if len(img.shape) == 2: + container = np.ones((cropsize, cropsize), img.dtype)*default_value + else: + container = np.ones((cropsize, cropsize, img.shape[2]), img.dtype)*default_value + + container[cont_top:cont_top+ch, cont_left:cont_left+cw] = \ + img[img_top:img_top+ch, img_left:img_left+cw] + + return container + + +def HWC_to_CHW(img): + return np.transpose(img, (2, 0, 1)) + + +def pil_blur(img, radius): + return np.array(Image.fromarray(img).filter(ImageFilter.GaussianBlur(radius=radius))) + + +def random_blur(img): + radius = random.random() + # print('add blur: ', radius) + if isinstance(img, list): + out = [] + for im in img: + out.append(pil_blur(im, radius)) + return out + elif isinstance(img, np.ndarray): + return pil_blur(img, radius) + else: + print(img) + raise RuntimeError("do not support the input image type!") + + +def save_image(image_numpy, image_path): + """Save a numpy image to the disk + Parameters: + image_numpy (numpy array) -- input numpy array + image_path (str) -- the path of the image + """ + image_pil = Image.fromarray(np.array(image_numpy,dtype=np.uint8)) + image_pil.save(image_path) + + +def im2arr(img_path, mode=1, dtype=np.uint8): + """ + :param img_path: + :param mode: + :return: numpy.ndarray, shape: H*W*C + """ + if mode==1: + img = PIL.Image.open(img_path) + arr = np.asarray(img, dtype=dtype) + else: + arr = tifffile.imread(img_path) + if arr.ndim == 3: + a, b, c = arr.shape + if a < b and a < c: # 当arr为C*H*W时,需要交换通道顺序 + arr = arr.transpose([1,2,0]) + # print('shape: ', arr.shape, 'dytpe: ',arr.dtype) + return arr + + + + + + + diff --git a/misc/logger_tool.py b/misc/logger_tool.py new file mode 100644 index 0000000000000000000000000000000000000000..31df11d3723184d275d07fcd47bb72aae7ad07b0 --- /dev/null +++ b/misc/logger_tool.py @@ -0,0 +1,73 @@ +import sys +import time + + +class Logger(object): + def __init__(self, outfile): + self.terminal = sys.stdout + self.log_path = outfile + now = time.strftime("%c") + self.write('================ (%s) ================\n' % now) + + def write(self, message): + self.terminal.write(message) + with open(self.log_path, mode='a') as f: + f.write(message) + + def write_dict(self, dict): + message = '' + for k, v in dict.items(): + message += '%s: %.7f ' % (k, v) + self.write(message) + + def write_dict_str(self, dict): + message = '' + for k, v in dict.items(): + message += '%s: %s ' % (k, v) + self.write(message) + + def flush(self): + self.terminal.flush() + + +class Timer: + def __init__(self, starting_msg = None): + self.start = time.time() + self.stage_start = self.start + + if starting_msg is not None: + print(starting_msg, time.ctime(time.time())) + + def __enter__(self): + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + return + + def update_progress(self, progress): + self.elapsed = time.time() - self.start + self.est_total = self.elapsed / progress + self.est_remaining = self.est_total - self.elapsed + self.est_finish = int(self.start + self.est_total) + + + def str_estimated_complete(self): + return str(time.ctime(self.est_finish)) + + def str_estimated_remaining(self): + return str(self.est_remaining/3600) + 'h' + + def estimated_remaining(self): + return self.est_remaining/3600 + + def get_stage_elapsed(self): + return time.time() - self.stage_start + + def reset_stage(self): + self.stage_start = time.time() + + def lapse(self): + out = time.time() - self.stage_start + self.stage_start = time.time() + return out + diff --git a/misc/metric_tool.py b/misc/metric_tool.py new file mode 100644 index 0000000000000000000000000000000000000000..a88c125e63cc2bcb4b35ef325bec4389b9f97e46 --- /dev/null +++ b/misc/metric_tool.py @@ -0,0 +1,164 @@ +import numpy as np + + +################### metrics ################### +class AverageMeter(object): + """Computes and stores the average and current value""" + def __init__(self): + self.initialized = False + self.val = None + self.avg = None + self.sum = None + self.count = None + + def initialize(self, val, weight): + self.val = val + self.avg = val + self.sum = val * weight + self.count = weight + self.initialized = True + + def update(self, val, weight=1): + if not self.initialized: + self.initialize(val, weight) + else: + self.add(val, weight) + + def add(self, val, weight): + self.val = val + self.sum += val * weight + self.count += weight + self.avg = self.sum / self.count + + def value(self): + return self.val + + def average(self): + return self.avg + + def get_scores(self): + scores_dict = cm2score(self.sum) + return scores_dict + + def clear(self): + self.initialized = False + + +################### cm metrics ################### +class ConfuseMatrixMeter(AverageMeter): + """Computes and stores the average and current value""" + def __init__(self, n_class): + super(ConfuseMatrixMeter, self).__init__() + self.n_class = n_class + + def update_cm(self, pr, gt, weight=1): + """获得当前混淆矩阵,并计算当前F1得分,并更新混淆矩阵""" + val = get_confuse_matrix(num_classes=self.n_class, label_gts=gt, label_preds=pr) + self.update(val, weight) + current_score = cm2F1(val) + return current_score + + def get_scores(self): + scores_dict = cm2score(self.sum) + return scores_dict + + + +def harmonic_mean(xs): + harmonic_mean = len(xs) / sum((x+1e-6)**-1 for x in xs) + return harmonic_mean + + +def cm2F1(confusion_matrix): + hist = confusion_matrix + n_class = hist.shape[0] + tp = np.diag(hist) + sum_a1 = hist.sum(axis=1) + sum_a0 = hist.sum(axis=0) + # ---------------------------------------------------------------------- # + # 1. Accuracy & Class Accuracy + # ---------------------------------------------------------------------- # + acc = tp.sum() / (hist.sum() + np.finfo(np.float32).eps) + + # recall + recall = tp / (sum_a1 + np.finfo(np.float32).eps) + # acc_cls = np.nanmean(recall) + + # precision + precision = tp / (sum_a0 + np.finfo(np.float32).eps) + + # F1 score + F1 = 2 * recall * precision / (recall + precision + np.finfo(np.float32).eps) + mean_F1 = np.nanmean(F1) + return mean_F1 + + +def cm2score(confusion_matrix): + hist = confusion_matrix + n_class = hist.shape[0] + tp = np.diag(hist) + sum_a1 = hist.sum(axis=1) + sum_a0 = hist.sum(axis=0) + # ---------------------------------------------------------------------- # + # 1. Accuracy & Class Accuracy + # ---------------------------------------------------------------------- # + acc = tp.sum() / (hist.sum() + np.finfo(np.float32).eps) + + # recall + recall = tp / (sum_a1 + np.finfo(np.float32).eps) + # acc_cls = np.nanmean(recall) + + # precision + precision = tp / (sum_a0 + np.finfo(np.float32).eps) + + # F1 score + F1 = 2*recall * precision / (recall + precision + np.finfo(np.float32).eps) + mean_F1 = np.nanmean(F1) + # ---------------------------------------------------------------------- # + # 2. Frequency weighted Accuracy & Mean IoU + # ---------------------------------------------------------------------- # + iu = tp / (sum_a1 + hist.sum(axis=0) - tp + np.finfo(np.float32).eps) + mean_iu = np.nanmean(iu) + + freq = sum_a1 / (hist.sum() + np.finfo(np.float32).eps) + fwavacc = (freq[freq > 0] * iu[freq > 0]).sum() + + # + cls_iou = dict(zip(['iou_'+str(i) for i in range(n_class)], iu)) + + cls_precision = dict(zip(['precision_'+str(i) for i in range(n_class)], precision)) + cls_recall = dict(zip(['recall_'+str(i) for i in range(n_class)], recall)) + cls_F1 = dict(zip(['F1_'+str(i) for i in range(n_class)], F1)) + + score_dict = {'acc': acc, 'miou': mean_iu, 'mf1':mean_F1} + score_dict.update(cls_iou) + score_dict.update(cls_F1) + score_dict.update(cls_precision) + score_dict.update(cls_recall) + return score_dict + + +def get_confuse_matrix(num_classes, label_gts, label_preds): + """计算一组预测的混淆矩阵""" + def __fast_hist(label_gt, label_pred): + """ + Collect values for Confusion Matrix + For reference, please see: https://en.wikipedia.org/wiki/Confusion_matrix + :param label_gt: ground-truth + :param label_pred: prediction + :return: values for confusion matrix + """ + mask = (label_gt >= 0) & (label_gt < num_classes) + hist = np.bincount(num_classes * label_gt[mask].astype(int) + label_pred[mask], + minlength=num_classes**2).reshape(num_classes, num_classes) + return hist + confusion_matrix = np.zeros((num_classes, num_classes)) + for lt, lp in zip(label_gts, label_preds): + confusion_matrix += __fast_hist(lt.flatten(), lp.flatten()) + return confusion_matrix + + +def get_mIoU(num_classes, label_gts, label_preds): + confusion_matrix = get_confuse_matrix(num_classes, label_gts, label_preds) + score_dict = cm2score(confusion_matrix) + return score_dict['miou'] diff --git a/misc/pyutils.py b/misc/pyutils.py new file mode 100644 index 0000000000000000000000000000000000000000..eeaaaccdf8221e6b0d15ad07d053c212b7571492 --- /dev/null +++ b/misc/pyutils.py @@ -0,0 +1,42 @@ +import numpy as np +import os +import random +import glob + + +def seed_random(seed=2020): + # 加入以下随机种子,数据输入,随机扩充等保持一致 + random.seed(seed) + os.environ['PYTHONHASHSEED'] = str(seed) + np.random.seed(seed) + + +def mkdir(path): + """create a single empty directory if it didn't exist + + Parameters: + path (str) -- a single directory path + """ + if not os.path.exists(path): + os.makedirs(path) + + +def get_paths(image_folder_path, suffix='*.png'): + """从文件夹中返回指定格式的文件 + :param image_folder_path: str + :param suffix: str + :return: list + """ + paths = sorted(glob.glob(os.path.join(image_folder_path, suffix))) + return paths + + +def get_paths_from_list(image_folder_path, list): + """从image folder中找到list中的文件,返回path list""" + out = [] + for item in list: + path = os.path.join(image_folder_path,item) + out.append(path) + return sorted(out) + + diff --git a/misc/torchutils.py b/misc/torchutils.py new file mode 100644 index 0000000000000000000000000000000000000000..8902ef1f84ce5e96d1ac3d98a1a49a6d8d766aa0 --- /dev/null +++ b/misc/torchutils.py @@ -0,0 +1,576 @@ +import torch +from torch.optim import lr_scheduler +from torch.utils.data import Subset +import torch.nn.functional as F +import numpy as np +import math +import random +import os +from torch.nn import MaxPool1d,AvgPool1d +from torch import Tensor +from typing import Iterable, Set, Tuple + + +__all__ = ['cls_accuracy'] + + + +def visualize_imgs(*imgs): + """ + 可视化图像,ndarray格式的图像 + :param imgs: ndarray:H*W*C, C=1/3 + :return: + """ + import matplotlib.pyplot as plt + nums = len(imgs) + if nums > 1: + fig, axs = plt.subplots(1, nums) + for i, image in enumerate(imgs): + axs[i].imshow(image, cmap='jet') + elif nums == 1: + fig, ax = plt.subplots(1, nums) + for i, image in enumerate(imgs): + ax.imshow(image, cmap='jet') + plt.show() + plt.show() + +def minmax(tensor): + assert tensor.ndim >= 2 + shape = tensor.shape + tensor = tensor.view([*shape[:-2], shape[-1]*shape[-2]]) + min_, _ = tensor.min(-1, keepdim=True) + max_, _ = tensor.max(-1, keepdim=True) + return min_, max_ + +def norm_tensor(tensor,min_=None,max_=None, mode='minmax'): + """ + 输入:N*C*H*W / C*H*W / H*W + 输出:在H*W维度的归一化的与原始等大的图 + """ + assert tensor.ndim >= 2 + shape = tensor.shape + tensor = tensor.view([*shape[:-2], shape[-1]*shape[-2]]) + if mode == 'minmax': + if min_ is None: + min_, _ = tensor.min(-1, keepdim=True) + if max_ is None: + max_, _ = tensor.max(-1, keepdim=True) + tensor = (tensor - min_) / (max_ - min_ + 0.00000000001) + elif mode == 'thres': + N = tensor.shape[-1] + thres_a = 0.001 + top_k = round(thres_a*N) + max_ = tensor.topk(top_k, dim=-1, largest=True)[0][..., -1] + max_ = max_.unsqueeze(-1) + min_ = tensor.topk(top_k, dim=-1, largest=False)[0][..., -1] + min_ = min_.unsqueeze(-1) + tensor = (tensor - min_) / (max_ - min_ + 0.00000000001) + + elif mode == 'std': + mean, std = torch.std_mean(tensor, [-1], keepdim=True) + tensor = (tensor - mean)/std + min_, _ = tensor.min(-1, keepdim=True) + max_, _ = tensor.max(-1, keepdim=True) + tensor = (tensor - min_) / (max_ - min_ + 0.00000000001) + elif mode == 'exp': + tai = 1 + tensor = torch.nn.functional.softmax(tensor/tai, dim=-1, ) + min_, _ = tensor.min(-1, keepdim=True) + max_, _ = tensor.max(-1, keepdim=True) + tensor = (tensor - min_) / (max_ - min_ + 0.00000000001) + else: + raise NotImplementedError + tensor = torch.clamp(tensor, 0, 1) + return tensor.view(shape) + + # if tensor.ndim == 4: + # B, C, H, W = tensor.shape + # tensor = tensor.view([B, C, -1]) + # min_, _ = tensor.min(-1, keepdim=True) + # max_, _ = tensor.max(-1, keepdim=True) + # tensor = (tensor - min_) / (max_ - min_ + 0.00000000001) + # return tensor.view(B, C, H, W) + # elif tensor.ndim == 3: + # C, H, W = tensor.shape + # tensor = tensor.view([C, -1]) + # min_, _ = tensor.min(-1, keepdim=True) + # max_, _ = tensor.max(-1, keepdim=True) + # tensor = (tensor - min_) / (max_ - min_ + 0.00000000001) + # return tensor.view(C, H, W) + # elif tensor.ndim == 2: + # H, W = tensor.shape + # tensor = tensor.view([-1]) + # min_, _ = tensor.min(-1, keepdim=True) + # max_, _ = tensor.max(-1, keepdim=True) + # tensor = (tensor - min_) / (max_ - min_ + 0.00000000001) + # return tensor.view(H, W) + # else: + # raise NotImplementedError + +def visulize_features(features, normalize=False): + """ + 可视化特征图,各维度make grid到一起 + """ + from torchvision.utils import make_grid + assert features.ndim == 4 + b,c,h,w = features.shape + features = features.view((b*c, 1, h, w)) + if normalize: + features = norm_tensor(features) + grid = make_grid(features) + visualize_tensors(grid) + +def visualize_tensors(*tensors): + """ + 可视化tensor,支持单通道特征或3通道图像 + :param tensors: tensor: C*H*W, C=1/3 + :return: + """ + import matplotlib.pyplot as plt + # from misc.torchutils import tensor2np + images = [] + for tensor in tensors: + assert tensor.ndim == 3 or tensor.ndim==2 + if tensor.ndim ==3: + assert tensor.shape[0] == 1 or tensor.shape[0] == 3 + images.append(tensor2np(tensor)) + nums = len(images) + if nums>1: + fig, axs = plt.subplots(1, nums) + for i, image in enumerate(images): + axs[i].imshow(image, cmap='jet') + plt.show() + elif nums == 1: + fig, ax = plt.subplots(1, nums) + for i, image in enumerate(images): + ax.imshow(image, cmap='jet') + plt.show() + + +def np_to_tensor(image): + """ + input: nd.array: H*W*C/H*W + """ + if isinstance(image, torch.Tensor): + return image + elif isinstance(image, np.ndarray): + if image.ndim == 3: + if image.shape[2]==3: + image = np.transpose(image,[2,0,1]) + elif image.ndim == 2: + image = np.newaxis(image, 0) + image = torch.from_numpy(image) + return image.unsqueeze(0) + + +def seed_torch(seed=2019): + + # 加入以下随机种子,数据输入,随机扩充等保持一致 + random.seed(seed) + os.environ['PYTHONHASHSEED'] = str(seed) + np.random.seed(seed) + torch.manual_seed(seed) + torch.cuda.manual_seed(seed) + # 加入所有随机种子后,模型更新后,中间结果还是不一样, + # 发现这一的现象:前两轮,的结果还是一样;随着模型更新结果会变; + # torch.backends.cudnn.benchmark = False + # torch.backends.cudnn.deterministic = True + +def simplex(t: Tensor, axis=1) -> bool: + _sum = t.sum(axis).type(torch.float32) + _ones = torch.ones_like(_sum, dtype=torch.float32) + return torch.allclose(_sum, _ones) + + +# Assert utils +def uniq(a: Tensor) -> Set: + return set(torch.unique(a.cpu()).numpy()) + +def sset(a: Tensor, sub: Iterable) -> bool: + return uniq(a).issubset(sub) + +def eq(a: Tensor, b) -> bool: + return torch.eq(a, b).all() + +def one_hot(t: Tensor, axis=1) -> bool: + return simplex(t, axis) and sset(t, [0, 1]) + + +def class2one_hot(seg: Tensor, C: int) -> Tensor: + if len(seg.shape) == 2: # Only w, h, used by the dataloader + seg = seg.unsqueeze(dim=0) + assert sset(seg, list(range(C))) + + b, w, h = seg.shape # type: Tuple[int, int, int] + + res = torch.stack([seg == c for c in range(C)], dim=1).type(torch.int32) + assert res.shape == (b, C, w, h) + assert one_hot(res) + + return res + +class ChannelMaxPool(MaxPool1d): + def forward(self, input): + n, c, w, h = input.size() + input = input.view(n,c,w*h).permute(0,2,1) + pooled = F.max_pool1d(input, self.kernel_size, self.stride, + self.padding, self.dilation, self.ceil_mode, + self.return_indices) + _, _, c = pooled.size() + pooled = pooled.permute(0,2,1) + return pooled.view(n,c,w,h) + +class ChannelAvePool(AvgPool1d): + def forward(self, input): + n, c, w, h = input.size() + input = input.view(n,c,w*h).permute(0,2,1) + pooled = F.avg_pool1d(input, self.kernel_size, self.stride, + self.padding) + _, _, c = pooled.size() + pooled = pooled.permute(0,2,1) + return pooled.view(n,c,w,h) + +def cross_entropy(input, target, weight=None, reduction='mean',ignore_index=255): + """ + logSoftmax_with_loss + :param input: torch.Tensor, N*C*H*W + :param target: torch.Tensor, N*1*H*W,/ N*H*W + :param weight: torch.Tensor, C + :return: torch.Tensor [0] + """ + target = target.long() + if target.dim() == 4: + target = torch.squeeze(target, dim=1) + if input.shape[-1] != target.shape[-1]: + input = F.interpolate(input, size=target.shape[1:], mode='bilinear',align_corners=True) + + return F.cross_entropy(input=input, target=target, weight=weight, + ignore_index=ignore_index, reduction=reduction) + +def balanced_cross_entropy(input, target, weight=None,ignore_index=255): + """ + 类别均衡的交叉熵损失,暂时只支持2类 + TODO: 扩展到多类C>2 + """ + if target.dim() == 4: + target = torch.squeeze(target, dim=1) + if input.shape[-1] != target.shape[-1]: + input = F.interpolate(input, size=target.shape[1:], mode='bilinear',align_corners=True) + + # print('target.sum',target.sum()) + pos = (target==1).float() + neg = (target==0).float() + pos_num = torch.sum(pos) + 0.0000001 + neg_num = torch.sum(neg) + 0.0000001 + # print(pos_num) + # print(neg_num) + target_pos = target.float() + target_pos[target_pos!=1] = ignore_index # 忽略不为正样本的区域 + target_neg = target.float() + target_neg[target_neg!=0] = ignore_index # 忽略不为负样本的区域 + + # print('target.sum',target.sum()) + + loss_pos = cross_entropy(input, target_pos,weight=weight,reduction='sum',ignore_index=ignore_index) + loss_neg = cross_entropy(input, target_neg,weight=weight,reduction='sum',ignore_index=ignore_index) + # print(loss_neg, loss_pos) + loss = 0.5 * loss_pos / pos_num + 0.5 * loss_neg / neg_num + # loss = (loss_pos + loss_neg)/ (pos_num+neg_num) + return loss + +def get_scheduler(optimizer, opt): + """Return a learning rate scheduler + """ + if opt.lr_policy == 'linear': + def lambda_rule(epoch): + lr_l = 1.0 - max(0, epoch + opt.epoch_count - opt.niter) / float(opt.niter_decay + 1) + return lr_l + scheduler = lr_scheduler.LambdaLR(optimizer, lr_lambda=lambda_rule) + elif opt.lr_policy == 'poly': + max_step = opt.niter+opt.niter_decay + power = 0.9 + def lambda_rule(epoch): + current_step = epoch + opt.epoch_count + lr_l = (1.0 - current_step / (max_step+1)) ** float(power) + return lr_l + scheduler = lr_scheduler.LambdaLR(optimizer, lr_lambda=lambda_rule) + elif opt.lr_policy == 'step': + scheduler = lr_scheduler.StepLR(optimizer, step_size=opt.lr_decay_iters, gamma=0.1) + else: + return NotImplementedError('learning rate policy [%s] is not implemented', opt.lr_policy) + return scheduler + + +def mul_cls_acc(preds, targets, topk=(1,)): + """计算multi-label分类的top-k准确率topk-acc,topk-error=1-topk-acc; + 首先计算每张图的的平均准确率,再计算所有图的平均准确率 + :param pred: N * C + :param target: N * C + :param topk: + :return: + """ + with torch.no_grad(): + maxk = max(topk) + bs, C = targets.shape + _, pred = preds.topk(maxk, 1, True, True) + pred += 1 # pred 为类别\in [1,C] + # print('pred: ', pred) + # print('targets: ', targets) + correct = torch.zeros([bs, maxk]).long() # 记录预测正确label数量 + if preds.device != torch.device(type='cpu'): + correct = correct.cuda() + for i in range(C): + label = i + 1 + target = targets[:, i] * label + # print('target.view: ', target.view(-1, 1).expand_as(pred)) + # print('pred: ', pred) + correct = correct + pred.eq(target.view(-1, 1).expand_as(pred)).long() + # print('correct: ', pred.eq(target.view(-1, 1).expand_as(pred)).long()) + n = (targets == 1).long().sum(1) # N*1, 每张图中含有目标的数量 + # print(n) + res = [] + for k in topk: + acc_k = correct[:, :k].sum(1).float() / n.float() # 每张图的平均正确率,预测正确目标数/总目标数 + # print(correct[:, :k].sum(1).float()) + acc_k = acc_k.sum()/bs + res.append(acc_k) + # print(acc_k) + return res + + +def cls_accuracy(output, target, topk=(1,)): + """ + Computes the accuracy over the k top predictions for the specified values of k + https://github.com/pytorch/examples/blob/ee964a2eeb41e1712fe719b83645c79bcbd0ba1a/imagenet/main.py#L407 + """ + + with torch.no_grad(): + maxk = max(topk) + batch_size = target.size(0) + + _, pred = output.topk(maxk, 1, True, True) + pred = pred.t() + correct = pred.eq(target.view(1, -1).expand_as(pred)) + + res = [] + for k in topk: + correct_k = correct[:k].view(-1).float().sum(0, keepdim=True) + res.append(correct_k.mul_(100.0 / batch_size)) + return res + +class PolyOptimizer(torch.optim.SGD): + + def __init__(self, params, lr, weight_decay, max_step, init_step=0, momentum=0.9): + super().__init__(params, lr, weight_decay) + + self.global_step = init_step + print(self.global_step) + self.max_step = max_step + self.momentum = momentum + + self.__initial_lr = [group['lr'] for group in self.param_groups] + + + def step(self, closure=None): + + if self.global_step < self.max_step: + lr_mult = (1 - self.global_step / self.max_step) ** self.momentum + + for i in range(len(self.param_groups)): + self.param_groups[i]['lr'] = self.__initial_lr[i] * lr_mult + + super().step(closure) + + self.global_step += 1 + + +class PolyAdamOptimizer(torch.optim.Adam): + def __init__(self, params, lr, betas, max_step, momentum=0.9): + super().__init__(params, lr, betas) + + self.global_step = 0 + self.max_step = max_step + self.momentum = momentum + + self.__initial_lr = [group['lr'] for group in self.param_groups] + + + def step(self, closure=None): + + if self.global_step < self.max_step: + lr_mult = (1 - self.global_step / self.max_step) ** self.momentum + + for i in range(len(self.param_groups)): + self.param_groups[i]['lr'] = self.__initial_lr[i] * lr_mult + + super().step(closure) + self.global_step += 1 +# +# from ranger import RangerQH,Ranger +# # https://github.com/lessw2020/Ranger-Deep-Learning-Optimizer/blob/master/ranger/rangerqh.py +# +# class PolyRangerOptimizer(RangerQH): +# +# def __init__(self, params, lr, betas, max_step, momentum=0.9): +# super().__init__(params, lr, betas) +# +# self.global_step = 0 +# self.max_step = max_step +# self.momentum = momentum +# +# self.__initial_lr = [group['lr'] for group in self.param_groups] +# +# +# def step(self, closure=None): +# +# if self.global_step < self.max_step: +# lr_mult = (1 - self.global_step / self.max_step) ** self.momentum +# +# for i in range(len(self.param_groups)): +# self.param_groups[i]['lr'] = self.__initial_lr[i] * lr_mult +# +# super().step(closure) +# self.global_step += 1 + +class SGDROptimizer(torch.optim.SGD): + + def __init__(self, params, steps_per_epoch, lr=0, weight_decay=0, epoch_start=1, restart_mult=2): + super().__init__(params, lr, weight_decay) + + self.global_step = 0 + self.local_step = 0 + self.total_restart = 0 + + self.max_step = steps_per_epoch * epoch_start + self.restart_mult = restart_mult + + self.__initial_lr = [group['lr'] for group in self.param_groups] + + + def step(self, closure=None): + + if self.local_step >= self.max_step: + self.local_step = 0 + self.max_step *= self.restart_mult + self.total_restart += 1 + + lr_mult = (1 + math.cos(math.pi * self.local_step / self.max_step))/2 / (self.total_restart + 1) + + for i in range(len(self.param_groups)): + self.param_groups[i]['lr'] = self.__initial_lr[i] * lr_mult + + super().step(closure) + + self.local_step += 1 + self.global_step += 1 + + +def split_dataset(dataset, n_splits): + + return [Subset(dataset, np.arange(i, len(dataset), n_splits)) for i in range(n_splits)] + + +def gap2d(x, keepdims=False): + out = torch.mean(x.view(x.size(0), x.size(1), -1), -1) + if keepdims: + out = out.view(out.size(0), out.size(1), 1, 1) + + return out + + +def decode_seg(label_mask, toTensor=False): + """ + :param label_mask: mask (np.ndarray): (M, N)/ tensor: N*C*H*W + :return: color label: (M, N, 3), + """ + if not isinstance(label_mask, np.ndarray): + if isinstance(label_mask, torch.Tensor): # get the data from a variable + image_tensor = label_mask.data + else: + return label_mask + label_mask = image_tensor[0][0].cpu().numpy() + + rgb = np.zeros((label_mask.shape[0], label_mask.shape[1], 3),dtype=np.float) + r = label_mask % 6 + g = (label_mask % 36) // 6 + b = label_mask // 36 + # 归一化到[0-1] + rgb[:, :, 0] = r / 6 + rgb[:, :, 1] = g / 6 + rgb[:, :, 2] = b / 6 + if toTensor: + rgb = torch.from_numpy(rgb.transpose([2,0,1])).unsqueeze(0) + + return rgb + + +def tensor2im(input_image, imtype=np.uint8, normalize=True): + """"Converts a Tensor array into a numpy image array. + Parameters: + input_image (tensor) -- the input image tensor array + imtype (type) -- the desired type of the converted numpy array + """ + if not isinstance(input_image, np.ndarray): + if isinstance(input_image, torch.Tensor): # get the data from a variable + image_tensor = input_image.data + else: + return input_image + image_numpy = image_tensor[0].cpu().float().numpy() # convert it into a numpy array + # if image_numpy.shape[0] == 1: # grayscale to RGB + # image_numpy = np.tile(image_numpy, (3, 1, 1)) + if image_numpy.shape[0] == 3: # if RGB + image_numpy = np.transpose(image_numpy, (1, 2, 0)) + if normalize: + image_numpy = (image_numpy + 1) / 2.0 * 255.0 # post-processing: tranpose and scaling + else: # if it is a numpy array, do nothing + image_numpy = input_image + return image_numpy.astype(imtype) + + +def tensor2np(input_image, if_normalize=True): + """ + :param input_image: C*H*W / H*W + :return: ndarray, H*W*C / H*W + """ + if isinstance(input_image, torch.Tensor): # get the data from a variable + image_tensor = input_image.data + image_numpy = image_tensor.cpu().float().numpy() # convert it into a numpy array + + else: + image_numpy = input_image + if image_numpy.ndim == 2: + return image_numpy + elif image_numpy.ndim == 3: + C, H, W = image_numpy.shape + image_numpy = np.transpose(image_numpy, (1, 2, 0)) + # 如果输入为灰度图C==1,则输出array,ndim==2; + if C == 1: + image_numpy = image_numpy[:, :, 0] + if if_normalize and C == 3: + image_numpy = (image_numpy + 1) / 2.0 * 255.0 # post-processing: tranpose and scaling + # add to prevent extreme noises in visual images + image_numpy[image_numpy<0]=0 + image_numpy[image_numpy>255]=255 + image_numpy = image_numpy.astype(np.uint8) + return image_numpy + + +import ntpath +from misc.imutils import save_image +def save_visuals(visuals, img_dir, name, save_one=True, iter='0'): + """ + """ + # save images to the disk + for label, image in visuals.items(): + N = image.shape[0] + if save_one: + N = 1 + # 保存各个bz的数据 + for j in range(N): + name_ = ntpath.basename(name[j]) + name_ = name_.split(".")[0] + # print(name_) + image_numpy = tensor2np(image[j], if_normalize=True).astype(np.uint8) + # print(image_numpy) + img_path = os.path.join(img_dir, iter+'_%s_%s.png' % (name_, label)) + save_image(image_numpy, img_path) \ No newline at end of file diff --git a/models/__init__.py b/models/__init__.py new file mode 100644 index 0000000000000000000000000000000000000000..b792ca6ecf7cbe1d51c5c1dd72f1f98328fda8b9 --- /dev/null +++ b/models/__init__.py @@ -0,0 +1 @@ +from .resnet import * diff --git a/models/basic_model.py b/models/basic_model.py new file mode 100644 index 0000000000000000000000000000000000000000..06284b0a85fb6dbfee3fe8708305b98ec73ec038 --- /dev/null +++ b/models/basic_model.py @@ -0,0 +1,75 @@ +import os + +import torch + +from misc.imutils import save_image +from models.networks import * + + +class CDEvaluator(): + + def __init__(self, args): + + self.n_class = args.n_class + # define G + self.net_G = define_G(args=args, gpu_ids=args.gpu_ids) + + self.device = torch.device("cuda:%s" % args.gpu_ids[0] + if torch.cuda.is_available() and len(args.gpu_ids)>0 + else "cpu") + + print(self.device) + + self.checkpoint_dir = args.checkpoint_dir + + self.pred_dir = args.output_folder + os.makedirs(self.pred_dir, exist_ok=True) + + def load_checkpoint(self, checkpoint_name='best_ckpt.pt'): + + if os.path.exists(os.path.join(self.checkpoint_dir, checkpoint_name)): + # load the entire checkpoint + checkpoint = torch.load(os.path.join(self.checkpoint_dir, checkpoint_name), + map_location=self.device) + + self.net_G.load_state_dict(checkpoint['model_G_state_dict']) + self.net_G.to(self.device) + # update some other states + self.best_val_acc = checkpoint['best_val_acc'] + self.best_epoch_id = checkpoint['best_epoch_id'] + + else: + raise FileNotFoundError('no such checkpoint %s' % checkpoint_name) + return self.net_G + + + def _visualize_pred(self): + pred = torch.argmax(self.G_pred, dim=1, keepdim=True) + pred_vis = pred * 255 + return pred_vis + + def _forward_pass(self, batch): + self.batch = batch + img_in1 = batch['A'].to(self.device) + img_in2 = batch['B'].to(self.device) + self.shape_h = img_in1.shape[-2] + self.shape_w = img_in1.shape[-1] + self.G_pred = self.net_G(img_in1, img_in2) + return self._visualize_pred() + + def eval(self): + self.net_G.eval() + + def _save_predictions(self): + """ + 保存模型输出结果,二分类图像 + """ + + preds = self._visualize_pred() + name = self.batch['name'] + for i, pred in enumerate(preds): + file_name = os.path.join( + self.pred_dir, name[i].replace('.jpg', '.png')) + pred = pred[0].cpu().numpy() + save_image(pred, file_name) + diff --git a/models/evaluator.py b/models/evaluator.py new file mode 100644 index 0000000000000000000000000000000000000000..e3d509bc204dd63ca4ad7afd20f556cbbd231bb3 --- /dev/null +++ b/models/evaluator.py @@ -0,0 +1,172 @@ +import os +import numpy as np +import matplotlib.pyplot as plt + +from models.networks import * +from misc.metric_tool import ConfuseMatrixMeter +from misc.logger_tool import Logger +from utils import de_norm +import utils + + +# Decide which device we want to run on +# torch.cuda.current_device() + +# device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu") + + +class CDEvaluator(): + + def __init__(self, args, dataloader): + + self.dataloader = dataloader + + self.n_class = args.n_class + # define G + self.net_G = define_G(args=args, gpu_ids=args.gpu_ids) + self.device = torch.device("cuda:%s" % args.gpu_ids[0] if torch.cuda.is_available() and len(args.gpu_ids)>0 + else "cpu") + print(self.device) + + # define some other vars to record the training states + self.running_metric = ConfuseMatrixMeter(n_class=self.n_class) + + # define logger file + logger_path = os.path.join(args.checkpoint_dir, 'log_test.txt') + self.logger = Logger(logger_path) + self.logger.write_dict_str(args.__dict__) + + + # training log + self.epoch_acc = 0 + self.best_val_acc = 0.0 + self.best_epoch_id = 0 + + self.steps_per_epoch = len(dataloader) + + self.G_pred = None + self.pred_vis = None + self.batch = None + self.is_training = False + self.batch_id = 0 + self.epoch_id = 0 + self.checkpoint_dir = args.checkpoint_dir + self.vis_dir = args.vis_dir + + # check and create model dir + if os.path.exists(self.checkpoint_dir) is False: + os.mkdir(self.checkpoint_dir) + if os.path.exists(self.vis_dir) is False: + os.mkdir(self.vis_dir) + + + def _load_checkpoint(self, checkpoint_name='best_ckpt.pt'): + + if os.path.exists(os.path.join(self.checkpoint_dir, checkpoint_name)): + self.logger.write('loading last checkpoint...\n') + # load the entire checkpoint + checkpoint = torch.load(os.path.join(self.checkpoint_dir, checkpoint_name), map_location=self.device) + + self.net_G.load_state_dict(checkpoint['model_G_state_dict']) + + self.net_G.to(self.device) + + # update some other states + self.best_val_acc = checkpoint['best_val_acc'] + self.best_epoch_id = checkpoint['best_epoch_id'] + + self.logger.write('Eval Historical_best_acc = %.4f (at epoch %d)\n' % + (self.best_val_acc, self.best_epoch_id)) + self.logger.write('\n') + + else: + raise FileNotFoundError('no such checkpoint %s' % checkpoint_name) + + + def _visualize_pred(self): + pred = torch.argmax(self.G_pred, dim=1, keepdim=True) + pred_vis = pred * 255 + return pred_vis + + + def _update_metric(self): + """ + update metric + """ + target = self.batch['L'].to(self.device).detach() + G_pred = self.G_pred.detach() + G_pred = torch.argmax(G_pred, dim=1) + + current_score = self.running_metric.update_cm(pr=G_pred.cpu().numpy(), gt=target.cpu().numpy()) + return current_score + + def _collect_running_batch_states(self): + + running_acc = self._update_metric() + + m = len(self.dataloader) + + if np.mod(self.batch_id, 100) == 1: + message = 'Is_training: %s. [%d,%d], running_mf1: %.5f\n' %\ + (self.is_training, self.batch_id, m, running_acc) + self.logger.write(message) + + if np.mod(self.batch_id, 100) == 1: + vis_input = utils.make_numpy_grid(de_norm(self.batch['A'])) + vis_input2 = utils.make_numpy_grid(de_norm(self.batch['B'])) + + vis_pred = utils.make_numpy_grid(self._visualize_pred()) + + vis_gt = utils.make_numpy_grid(self.batch['L']) + vis = np.concatenate([vis_input, vis_input2, vis_pred, vis_gt], axis=0) + vis = np.clip(vis, a_min=0.0, a_max=1.0) + file_name = os.path.join( + self.vis_dir, 'eval_' + str(self.batch_id)+'.jpg') + plt.imsave(file_name, vis) + + + def _collect_epoch_states(self): + + scores_dict = self.running_metric.get_scores() + + np.save(os.path.join(self.checkpoint_dir, 'scores_dict.npy'), scores_dict) + + self.epoch_acc = scores_dict['mf1'] + + with open(os.path.join(self.checkpoint_dir, '%s.txt' % (self.epoch_acc)), + mode='a') as file: + pass + + message = '' + for k, v in scores_dict.items(): + message += '%s: %.5f ' % (k, v) + self.logger.write('%s\n' % message) # save the message + + self.logger.write('\n') + + def _clear_cache(self): + self.running_metric.clear() + + def _forward_pass(self, batch): + self.batch = batch + img_in1 = batch['A'].to(self.device) + img_in2 = batch['B'].to(self.device) + self.G_pred = self.net_G(img_in1, img_in2) + + def eval_models(self,checkpoint_name='best_ckpt.pt'): + + self._load_checkpoint(checkpoint_name) + + ################## Eval ################## + ########################################## + self.logger.write('Begin evaluation...\n') + self._clear_cache() + self.is_training = False + self.net_G.eval() + + # Iterate over data. + for self.batch_id, batch in enumerate(self.dataloader, 0): + with torch.no_grad(): + self._forward_pass(batch) + self._collect_running_batch_states() + self._collect_epoch_states() diff --git a/models/help_funcs.py b/models/help_funcs.py new file mode 100644 index 0000000000000000000000000000000000000000..4f83e73f865fc49a421991e1beafc36d212d0dc9 --- /dev/null +++ b/models/help_funcs.py @@ -0,0 +1,188 @@ +import torch +import torch.nn.functional as F +from einops import rearrange +from torch import nn + + +class TwoLayerConv2d(nn.Sequential): + def __init__(self, in_channels, out_channels, kernel_size=3): + super().__init__(nn.Conv2d(in_channels, in_channels, kernel_size=kernel_size, + padding=kernel_size // 2, stride=1, bias=False), + nn.BatchNorm2d(in_channels), + nn.ReLU(), + nn.Conv2d(in_channels, out_channels, kernel_size=kernel_size, + padding=kernel_size // 2, stride=1) + ) + + +class Residual(nn.Module): + def __init__(self, fn): + super().__init__() + self.fn = fn + def forward(self, x, **kwargs): + return self.fn(x, **kwargs) + x + + +class Residual2(nn.Module): + def __init__(self, fn): + super().__init__() + self.fn = fn + def forward(self, x, x2, **kwargs): + return self.fn(x, x2, **kwargs) + x + + +class PreNorm(nn.Module): + def __init__(self, dim, fn): + super().__init__() + self.norm = nn.LayerNorm(dim) + self.fn = fn + def forward(self, x, **kwargs): + return self.fn(self.norm(x), **kwargs) + + +class PreNorm2(nn.Module): + def __init__(self, dim, fn): + super().__init__() + self.norm = nn.LayerNorm(dim) + self.fn = fn + def forward(self, x, x2, **kwargs): + return self.fn(self.norm(x), self.norm(x2), **kwargs) + + +class FeedForward(nn.Module): + def __init__(self, dim, hidden_dim, dropout = 0.): + super().__init__() + self.net = nn.Sequential( + nn.Linear(dim, hidden_dim), + nn.GELU(), + nn.Dropout(dropout), + nn.Linear(hidden_dim, dim), + nn.Dropout(dropout) + ) + def forward(self, x): + return self.net(x) + + +class Cross_Attention(nn.Module): + def __init__(self, dim, heads = 8, dim_head = 64, dropout = 0., softmax=True): + super().__init__() + inner_dim = dim_head * heads + self.heads = heads + self.scale = dim ** -0.5 + + self.softmax = softmax + self.to_q = nn.Linear(dim, inner_dim, bias=False) + self.to_k = nn.Linear(dim, inner_dim, bias=False) + self.to_v = nn.Linear(dim, inner_dim, bias=False) + + self.to_out = nn.Sequential( + nn.Linear(inner_dim, dim), + nn.Dropout(dropout) + ) + + def forward(self, x, m, mask = None): + + b, n, _, h = *x.shape, self.heads + q = self.to_q(x) + k = self.to_k(m) + v = self.to_v(m) + + q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h = h), [q,k,v]) + + dots = torch.einsum('bhid,bhjd->bhij', q, k) * self.scale + mask_value = -torch.finfo(dots.dtype).max + + if mask is not None: + mask = F.pad(mask.flatten(1), (1, 0), value = True) + assert mask.shape[-1] == dots.shape[-1], 'mask has incorrect dimensions' + mask = mask[:, None, :] * mask[:, :, None] + dots.masked_fill_(~mask, mask_value) + del mask + + if self.softmax: + attn = dots.softmax(dim=-1) + else: + attn = dots + # attn = dots + # vis_tmp(dots) + + out = torch.einsum('bhij,bhjd->bhid', attn, v) + out = rearrange(out, 'b h n d -> b n (h d)') + out = self.to_out(out) + # vis_tmp2(out) + + return out + + +class Attention(nn.Module): + def __init__(self, dim, heads = 8, dim_head = 64, dropout = 0.): + super().__init__() + inner_dim = dim_head * heads + self.heads = heads + self.scale = dim ** -0.5 + + self.to_qkv = nn.Linear(dim, inner_dim * 3, bias = False) + self.to_out = nn.Sequential( + nn.Linear(inner_dim, dim), + nn.Dropout(dropout) + ) + + def forward(self, x, mask = None): + b, n, _, h = *x.shape, self.heads + qkv = self.to_qkv(x).chunk(3, dim = -1) + q, k, v = map(lambda t: rearrange(t, 'b n (h d) -> b h n d', h = h), qkv) + + dots = torch.einsum('bhid,bhjd->bhij', q, k) * self.scale + mask_value = -torch.finfo(dots.dtype).max + + if mask is not None: + mask = F.pad(mask.flatten(1), (1, 0), value = True) + assert mask.shape[-1] == dots.shape[-1], 'mask has incorrect dimensions' + mask = mask[:, None, :] * mask[:, :, None] + dots.masked_fill_(~mask, mask_value) + del mask + + attn = dots.softmax(dim=-1) + + + out = torch.einsum('bhij,bhjd->bhid', attn, v) + out = rearrange(out, 'b h n d -> b n (h d)') + out = self.to_out(out) + return out + + +class Transformer(nn.Module): + def __init__(self, dim, depth, heads, dim_head, mlp_dim, dropout): + super().__init__() + self.layers = nn.ModuleList([]) + for _ in range(depth): + self.layers.append(nn.ModuleList([ + Residual(PreNorm(dim, Attention(dim, heads = heads, dim_head = dim_head, dropout = dropout))), + Residual(PreNorm(dim, FeedForward(dim, mlp_dim, dropout = dropout))) + ])) + def forward(self, x, mask = None): + for attn, ff in self.layers: + x = attn(x, mask = mask) + x = ff(x) + return x + + +class TransformerDecoder(nn.Module): + def __init__(self, dim, depth, heads, dim_head, mlp_dim, dropout, softmax=True): + super().__init__() + self.layers = nn.ModuleList([]) + for _ in range(depth): + self.layers.append(nn.ModuleList([ + Residual2(PreNorm2(dim, Cross_Attention(dim, heads = heads, + dim_head = dim_head, dropout = dropout, + softmax=softmax))), + Residual(PreNorm(dim, FeedForward(dim, mlp_dim, dropout = dropout))) + ])) + def forward(self, x, m, mask = None): + """target(query), memory""" + for attn, ff in self.layers: + x = attn(x, m, mask = mask) + x = ff(x) + return x + + diff --git a/models/losses.py b/models/losses.py new file mode 100644 index 0000000000000000000000000000000000000000..0fc0381198c18ff38bd915fc345db9db009ab898 --- /dev/null +++ b/models/losses.py @@ -0,0 +1,20 @@ +import torch +import torch.nn.functional as F + + +def cross_entropy(input, target, weight=None, reduction='mean',ignore_index=255): + """ + logSoftmax_with_loss + :param input: torch.Tensor, N*C*H*W + :param target: torch.Tensor, N*1*H*W,/ N*H*W + :param weight: torch.Tensor, C + :return: torch.Tensor [0] + """ + target = target.long() + if target.dim() == 4: + target = torch.squeeze(target, dim=1) + if input.shape[-1] != target.shape[-1]: + input = F.interpolate(input, size=target.shape[1:], mode='bilinear',align_corners=True) + + return F.cross_entropy(input=input, target=target, weight=weight, + ignore_index=ignore_index, reduction=reduction) diff --git a/models/networks.py b/models/networks.py new file mode 100644 index 0000000000000000000000000000000000000000..558f413b169de820aad8174cd33c2b6c5b742195 --- /dev/null +++ b/models/networks.py @@ -0,0 +1,367 @@ +import torch +import torch.nn as nn +from torch.nn import init +import torch.nn.functional as F +from torch.optim import lr_scheduler + +import functools +from einops import rearrange + +import models +from models.help_funcs import Transformer, TransformerDecoder, TwoLayerConv2d + + +############################################################################### +# Helper Functions +############################################################################### + +def get_scheduler(optimizer, args): + """Return a learning rate scheduler + + Parameters: + optimizer -- the optimizer of the network + args (option class) -- stores all the experiment flags; needs to be a subclass of BaseOptions.  + opt.lr_policy is the name of learning rate policy: linear | step | plateau | cosine + + For 'linear', we keep the same learning rate for the first epochs + and linearly decay the rate to zero over the next epochs. + For other schedulers (step, plateau, and cosine), we use the default PyTorch schedulers. + See https://pytorch.org/docs/stable/optim.html for more details. + """ + if args.lr_policy == 'linear': + def lambda_rule(epoch): + lr_l = 1.0 - epoch / float(args.max_epochs + 1) + return lr_l + scheduler = lr_scheduler.LambdaLR(optimizer, lr_lambda=lambda_rule) + elif args.lr_policy == 'step': + step_size = args.max_epochs//3 + # args.lr_decay_iters + scheduler = lr_scheduler.StepLR(optimizer, step_size=step_size, gamma=0.1) + else: + return NotImplementedError('learning rate policy [%s] is not implemented', args.lr_policy) + return scheduler + + +class Identity(nn.Module): + def forward(self, x): + return x + + +def get_norm_layer(norm_type='instance'): + """Return a normalization layer + + Parameters: + norm_type (str) -- the name of the normalization layer: batch | instance | none + + For BatchNorm, we use learnable affine parameters and track running statistics (mean/stddev). + For InstanceNorm, we do not use learnable affine parameters. We do not track running statistics. + """ + if norm_type == 'batch': + norm_layer = functools.partial(nn.BatchNorm2d, affine=True, track_running_stats=True) + elif norm_type == 'instance': + norm_layer = functools.partial(nn.InstanceNorm2d, affine=False, track_running_stats=False) + elif norm_type == 'none': + norm_layer = lambda x: Identity() + else: + raise NotImplementedError('normalization layer [%s] is not found' % norm_type) + return norm_layer + + +def init_weights(net, init_type='normal', init_gain=0.02): + """Initialize network weights. + + Parameters: + net (network) -- network to be initialized + init_type (str) -- the name of an initialization method: normal | xavier | kaiming | orthogonal + init_gain (float) -- scaling factor for normal, xavier and orthogonal. + + We use 'normal' in the original pix2pix and CycleGAN paper. But xavier and kaiming might + work better for some applications. Feel free to try yourself. + """ + def init_func(m): # define the initialization function + classname = m.__class__.__name__ + if hasattr(m, 'weight') and (classname.find('Conv') != -1 or classname.find('Linear') != -1): + if init_type == 'normal': + init.normal_(m.weight.data, 0.0, init_gain) + elif init_type == 'xavier': + init.xavier_normal_(m.weight.data, gain=init_gain) + elif init_type == 'kaiming': + init.kaiming_normal_(m.weight.data, a=0, mode='fan_in') + elif init_type == 'orthogonal': + init.orthogonal_(m.weight.data, gain=init_gain) + else: + raise NotImplementedError('initialization method [%s] is not implemented' % init_type) + if hasattr(m, 'bias') and m.bias is not None: + init.constant_(m.bias.data, 0.0) + elif classname.find('BatchNorm2d') != -1: # BatchNorm Layer's weight is not a matrix; only normal distribution applies. + init.normal_(m.weight.data, 1.0, init_gain) + init.constant_(m.bias.data, 0.0) + + print('initialize network with %s' % init_type) + net.apply(init_func) # apply the initialization function + + +def init_net(net, init_type='normal', init_gain=0.02, gpu_ids=[]): + """Initialize a network: 1. register CPU/GPU device (with multi-GPU support); 2. initialize the network weights + Parameters: + net (network) -- the network to be initialized + init_type (str) -- the name of an initialization method: normal | xavier | kaiming | orthogonal + gain (float) -- scaling factor for normal, xavier and orthogonal. + gpu_ids (int list) -- which GPUs the network runs on: e.g., 0,1,2 + + Return an initialized network. + """ + if len(gpu_ids) > 0: + assert(torch.cuda.is_available()) + net.to(gpu_ids[0]) + if len(gpu_ids) > 1: + net = torch.nn.DataParallel(net, gpu_ids) # multi-GPUs + init_weights(net, init_type, init_gain=init_gain) + return net + + +def define_G(args, init_type='normal', init_gain=0.02, gpu_ids=[]): + if args.net_G == 'base_resnet18': + net = ResNet(input_nc=3, output_nc=2, output_sigmoid=False) + + elif args.net_G == 'base_transformer_pos_s4': + net = BASE_Transformer(input_nc=3, output_nc=2, token_len=4, resnet_stages_num=4, + with_pos='learned') + + elif args.net_G == 'base_transformer_pos_s4_dd8': + net = BASE_Transformer(input_nc=3, output_nc=2, token_len=4, resnet_stages_num=4, + with_pos='learned', enc_depth=1, dec_depth=8) + + elif args.net_G == 'base_transformer_pos_s4_dd8_dedim8': + net = BASE_Transformer(input_nc=3, output_nc=2, token_len=4, resnet_stages_num=4, + with_pos='learned', enc_depth=1, dec_depth=8, decoder_dim_head=8) + + else: + raise NotImplementedError('Generator model name [%s] is not recognized' % args.net_G) + return init_net(net, init_type, init_gain, gpu_ids) + + +############################################################################### +# main Functions +############################################################################### + + +class ResNet(torch.nn.Module): + def __init__(self, input_nc, output_nc, + resnet_stages_num=5, backbone='resnet18', + output_sigmoid=False, if_upsample_2x=True): + """ + In the constructor we instantiate two nn.Linear modules and assign them as + member variables. + """ + super(ResNet, self).__init__() + expand = 1 + if backbone == 'resnet18': + self.resnet = models.resnet18(pretrained=True, + replace_stride_with_dilation=[False,True,True]) + elif backbone == 'resnet34': + self.resnet = models.resnet34(pretrained=True, + replace_stride_with_dilation=[False,True,True]) + elif backbone == 'resnet50': + self.resnet = models.resnet50(pretrained=True, + replace_stride_with_dilation=[False,True,True]) + expand = 4 + else: + raise NotImplementedError + self.relu = nn.ReLU() + self.upsamplex2 = nn.Upsample(scale_factor=2) + self.upsamplex4 = nn.Upsample(scale_factor=4, mode='bilinear') + + self.classifier = TwoLayerConv2d(in_channels=32, out_channels=output_nc) + + self.resnet_stages_num = resnet_stages_num + + self.if_upsample_2x = if_upsample_2x + if self.resnet_stages_num == 5: + layers = 512 * expand + elif self.resnet_stages_num == 4: + layers = 256 * expand + elif self.resnet_stages_num == 3: + layers = 128 * expand + else: + raise NotImplementedError + self.conv_pred = nn.Conv2d(layers, 32, kernel_size=3, padding=1) + + self.output_sigmoid = output_sigmoid + self.sigmoid = nn.Sigmoid() + + def forward(self, x1, x2): + x1 = self.forward_single(x1) + x2 = self.forward_single(x2) + x = torch.abs(x1 - x2) + if not self.if_upsample_2x: + x = self.upsamplex2(x) + x = self.upsamplex4(x) + x = self.classifier(x) + + if self.output_sigmoid: + x = self.sigmoid(x) + return x + + def forward_single(self, x): + # resnet layers + x = self.resnet.conv1(x) + x = self.resnet.bn1(x) + x = self.resnet.relu(x) + x = self.resnet.maxpool(x) + + x_4 = self.resnet.layer1(x) # 1/4, in=64, out=64 + x_8 = self.resnet.layer2(x_4) # 1/8, in=64, out=128 + + if self.resnet_stages_num > 3: + x_8 = self.resnet.layer3(x_8) # 1/8, in=128, out=256 + + if self.resnet_stages_num == 5: + x_8 = self.resnet.layer4(x_8) # 1/32, in=256, out=512 + elif self.resnet_stages_num > 5: + raise NotImplementedError + + if self.if_upsample_2x: + x = self.upsamplex2(x_8) + else: + x = x_8 + # output layers + x = self.conv_pred(x) + return x + + +class BASE_Transformer(ResNet): + """ + Resnet of 8 downsampling + BIT + bitemporal feature Differencing + a small CNN + """ + def __init__(self, input_nc, output_nc, with_pos, resnet_stages_num=5, + token_len=4, token_trans=True, + enc_depth=1, dec_depth=1, + dim_head=64, decoder_dim_head=64, + tokenizer=True, if_upsample_2x=True, + pool_mode='max', pool_size=2, + backbone='resnet18', + decoder_softmax=True, with_decoder_pos=None, + with_decoder=True): + super(BASE_Transformer, self).__init__(input_nc, output_nc,backbone=backbone, + resnet_stages_num=resnet_stages_num, + if_upsample_2x=if_upsample_2x, + ) + self.token_len = token_len + self.conv_a = nn.Conv2d(32, self.token_len, kernel_size=1, + padding=0, bias=False) + self.tokenizer = tokenizer + if not self.tokenizer: + # if not use tokenzier,then downsample the feature map into a certain size + self.pooling_size = pool_size + self.pool_mode = pool_mode + self.token_len = self.pooling_size * self.pooling_size + + self.token_trans = token_trans + self.with_decoder = with_decoder + dim = 32 + mlp_dim = 2*dim + + self.with_pos = with_pos + if with_pos is 'learned': + self.pos_embedding = nn.Parameter(torch.randn(1, self.token_len*2, 32)) + decoder_pos_size = 256//4 + self.with_decoder_pos = with_decoder_pos + if self.with_decoder_pos == 'learned': + self.pos_embedding_decoder =nn.Parameter(torch.randn(1, 32, + decoder_pos_size, + decoder_pos_size)) + self.enc_depth = enc_depth + self.dec_depth = dec_depth + self.dim_head = dim_head + self.decoder_dim_head = decoder_dim_head + self.transformer = Transformer(dim=dim, depth=self.enc_depth, heads=8, + dim_head=self.dim_head, + mlp_dim=mlp_dim, dropout=0) + self.transformer_decoder = TransformerDecoder(dim=dim, depth=self.dec_depth, + heads=8, dim_head=self.decoder_dim_head, mlp_dim=mlp_dim, dropout=0, + softmax=decoder_softmax) + + def _forward_semantic_tokens(self, x): + b, c, h, w = x.shape + spatial_attention = self.conv_a(x) + spatial_attention = spatial_attention.view([b, self.token_len, -1]).contiguous() + spatial_attention = torch.softmax(spatial_attention, dim=-1) + x = x.view([b, c, -1]).contiguous() + tokens = torch.einsum('bln,bcn->blc', spatial_attention, x) + + return tokens + + def _forward_reshape_tokens(self, x): + # b,c,h,w = x.shape + if self.pool_mode is 'max': + x = F.adaptive_max_pool2d(x, [self.pooling_size, self.pooling_size]) + elif self.pool_mode is 'ave': + x = F.adaptive_avg_pool2d(x, [self.pooling_size, self.pooling_size]) + else: + x = x + tokens = rearrange(x, 'b c h w -> b (h w) c') + return tokens + + def _forward_transformer(self, x): + if self.with_pos: + x += self.pos_embedding + x = self.transformer(x) + return x + + def _forward_transformer_decoder(self, x, m): + b, c, h, w = x.shape + if self.with_decoder_pos == 'fix': + x = x + self.pos_embedding_decoder + elif self.with_decoder_pos == 'learned': + x = x + self.pos_embedding_decoder + x = rearrange(x, 'b c h w -> b (h w) c') + x = self.transformer_decoder(x, m) + x = rearrange(x, 'b (h w) c -> b c h w', h=h) + return x + + def _forward_simple_decoder(self, x, m): + b, c, h, w = x.shape + b, l, c = m.shape + m = m.expand([h,w,b,l,c]) + m = rearrange(m, 'h w b l c -> l b c h w') + m = m.sum(0) + x = x + m + return x + + def forward(self, x1, x2): + # forward backbone resnet + x1 = self.forward_single(x1) + x2 = self.forward_single(x2) + + # forward tokenzier + if self.tokenizer: + token1 = self._forward_semantic_tokens(x1) + token2 = self._forward_semantic_tokens(x2) + else: + token1 = self._forward_reshape_tokens(x1) + token2 = self._forward_reshape_tokens(x2) + # forward transformer encoder + if self.token_trans: + self.tokens_ = torch.cat([token1, token2], dim=1) + self.tokens = self._forward_transformer(self.tokens_) + token1, token2 = self.tokens.chunk(2, dim=1) + # forward transformer decoder + if self.with_decoder: + x1 = self._forward_transformer_decoder(x1, token1) + x2 = self._forward_transformer_decoder(x2, token2) + else: + x1 = self._forward_simple_decoder(x1, token1) + x2 = self._forward_simple_decoder(x2, token2) + # feature differencing + x = torch.abs(x1 - x2) + if not self.if_upsample_2x: + x = self.upsamplex2(x) + x = self.upsamplex4(x) + # forward small cnn + x = self.classifier(x) + if self.output_sigmoid: + x = self.sigmoid(x) + return x + + diff --git a/models/resnet.py b/models/resnet.py new file mode 100644 index 0000000000000000000000000000000000000000..62f1c13119948ea74405ef865f70d831edce7da5 --- /dev/null +++ b/models/resnet.py @@ -0,0 +1,358 @@ +import torch +import torch.nn as nn +from torchvision.models.utils import load_state_dict_from_url + + +__all__ = ['ResNet', 'resnet18', 'resnet34', 'resnet50', 'resnet101', + 'resnet152', 'resnext50_32x4d', 'resnext101_32x8d', + 'wide_resnet50_2', 'wide_resnet101_2'] + + +model_urls = { + 'resnet18': 'https://download.pytorch.org/models/resnet18-5c106cde.pth', + 'resnet34': 'https://download.pytorch.org/models/resnet34-333f7ec4.pth', + 'resnet50': 'https://download.pytorch.org/models/resnet50-19c8e357.pth', + 'resnet101': 'https://download.pytorch.org/models/resnet101-5d3b4d8f.pth', + 'resnet152': 'https://download.pytorch.org/models/resnet152-b121ed2d.pth', + 'resnext50_32x4d': 'https://download.pytorch.org/models/resnext50_32x4d-7cdf4587.pth', + 'resnext101_32x8d': 'https://download.pytorch.org/models/resnext101_32x8d-8ba56ff5.pth', + 'wide_resnet50_2': 'https://download.pytorch.org/models/wide_resnet50_2-95faca4d.pth', + 'wide_resnet101_2': 'https://download.pytorch.org/models/wide_resnet101_2-32ee1156.pth', +} + + +def conv3x3(in_planes, out_planes, stride=1, groups=1, dilation=1): + """3x3 convolution with padding""" + return nn.Conv2d(in_planes, out_planes, kernel_size=3, stride=stride, + padding=dilation, groups=groups, bias=False, dilation=dilation) + + +def conv1x1(in_planes, out_planes, stride=1): + """1x1 convolution""" + return nn.Conv2d(in_planes, out_planes, kernel_size=1, stride=stride, bias=False) + + +class BasicBlock(nn.Module): + expansion = 1 + + def __init__(self, inplanes, planes, stride=1, downsample=None, groups=1, + base_width=64, dilation=1, norm_layer=None): + super(BasicBlock, self).__init__() + if norm_layer is None: + norm_layer = nn.BatchNorm2d + if groups != 1 or base_width != 64: + raise ValueError('BasicBlock only supports groups=1 and base_width=64') + if dilation > 1: + dilation = 1 + # raise NotImplementedError("Dilation > 1 not supported in BasicBlock") + # Both self.conv1 and self.downsample layers downsample the input when stride != 1 + self.conv1 = conv3x3(inplanes, planes, stride) + self.bn1 = norm_layer(planes) + self.relu = nn.ReLU(inplace=True) + self.conv2 = conv3x3(planes, planes) + self.bn2 = norm_layer(planes) + self.downsample = downsample + self.stride = stride + + def forward(self, x): + identity = x + + out = self.conv1(x) + out = self.bn1(out) + out = self.relu(out) + + out = self.conv2(out) + out = self.bn2(out) + + if self.downsample is not None: + identity = self.downsample(x) + + out += identity + out = self.relu(out) + + return out + + +class Bottleneck(nn.Module): + # Bottleneck in torchvision places the stride for downsampling at 3x3 convolution(self.conv2) + # while original implementation places the stride at the first 1x1 convolution(self.conv1) + # according to "Deep residual learning for image recognition"https://arxiv.org/abs/1512.03385. + # This variant is also known as ResNet V1.5 and improves accuracy according to + # https://ngc.nvidia.com/catalog/model-scripts/nvidia:resnet_50_v1_5_for_pytorch. + + expansion = 4 + + def __init__(self, inplanes, planes, stride=1, downsample=None, groups=1, + base_width=64, dilation=1, norm_layer=None): + super(Bottleneck, self).__init__() + if norm_layer is None: + norm_layer = nn.BatchNorm2d + width = int(planes * (base_width / 64.)) * groups + # Both self.conv2 and self.downsample layers downsample the input when stride != 1 + self.conv1 = conv1x1(inplanes, width) + self.bn1 = norm_layer(width) + self.conv2 = conv3x3(width, width, stride, groups, dilation) + self.bn2 = norm_layer(width) + self.conv3 = conv1x1(width, planes * self.expansion) + self.bn3 = norm_layer(planes * self.expansion) + self.relu = nn.ReLU(inplace=True) + self.downsample = downsample + self.stride = stride + + def forward(self, x): + identity = x + + out = self.conv1(x) + out = self.bn1(out) + out = self.relu(out) + + out = self.conv2(out) + out = self.bn2(out) + out = self.relu(out) + + out = self.conv3(out) + out = self.bn3(out) + + if self.downsample is not None: + identity = self.downsample(x) + + out += identity + out = self.relu(out) + + return out + + +class ResNet(nn.Module): + + def __init__(self, block, layers, num_classes=1000, zero_init_residual=False, + groups=1, width_per_group=64, replace_stride_with_dilation=None, + norm_layer=None, strides=None): + super(ResNet, self).__init__() + if norm_layer is None: + norm_layer = nn.BatchNorm2d + self._norm_layer = norm_layer + + self.strides = strides + if self.strides is None: + self.strides = [2, 2, 2, 2, 2] + + self.inplanes = 64 + self.dilation = 1 + if replace_stride_with_dilation is None: + # each element in the tuple indicates if we should replace + # the 2x2 stride with a dilated convolution instead + replace_stride_with_dilation = [False, False, False] + if len(replace_stride_with_dilation) != 3: + raise ValueError("replace_stride_with_dilation should be None " + "or a 3-element tuple, got {}".format(replace_stride_with_dilation)) + self.groups = groups + self.base_width = width_per_group + self.conv1 = nn.Conv2d(3, self.inplanes, kernel_size=7, stride=self.strides[0], padding=3, + bias=False) + self.bn1 = norm_layer(self.inplanes) + self.relu = nn.ReLU(inplace=True) + self.maxpool = nn.MaxPool2d(kernel_size=3, stride=self.strides[1], padding=1) + self.layer1 = self._make_layer(block, 64, layers[0]) + self.layer2 = self._make_layer(block, 128, layers[1], stride=self.strides[2], + dilate=replace_stride_with_dilation[0]) + self.layer3 = self._make_layer(block, 256, layers[2], stride=self.strides[3], + dilate=replace_stride_with_dilation[1]) + self.layer4 = self._make_layer(block, 512, layers[3], stride=self.strides[4], + dilate=replace_stride_with_dilation[2]) + self.avgpool = nn.AdaptiveAvgPool2d((1, 1)) + self.fc = nn.Linear(512 * block.expansion, num_classes) + + for m in self.modules(): + if isinstance(m, nn.Conv2d): + nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu') + elif isinstance(m, (nn.BatchNorm2d, nn.GroupNorm)): + nn.init.constant_(m.weight, 1) + nn.init.constant_(m.bias, 0) + + # Zero-initialize the last BN in each residual branch, + # so that the residual branch starts with zeros, and each residual block behaves like an identity. + # This improves the model by 0.2~0.3% according to https://arxiv.org/abs/1706.02677 + if zero_init_residual: + for m in self.modules(): + if isinstance(m, Bottleneck): + nn.init.constant_(m.bn3.weight, 0) + elif isinstance(m, BasicBlock): + nn.init.constant_(m.bn2.weight, 0) + + def _make_layer(self, block, planes, blocks, stride=1, dilate=False): + norm_layer = self._norm_layer + downsample = None + previous_dilation = self.dilation + if dilate: + self.dilation *= stride + stride = 1 + if stride != 1 or self.inplanes != planes * block.expansion: + downsample = nn.Sequential( + conv1x1(self.inplanes, planes * block.expansion, stride), + norm_layer(planes * block.expansion), + ) + + layers = [] + layers.append(block(self.inplanes, planes, stride, downsample, self.groups, + self.base_width, previous_dilation, norm_layer)) + self.inplanes = planes * block.expansion + for _ in range(1, blocks): + layers.append(block(self.inplanes, planes, groups=self.groups, + base_width=self.base_width, dilation=self.dilation, + norm_layer=norm_layer)) + + return nn.Sequential(*layers) + + def _forward_impl(self, x): + # See note [TorchScript super()] + x = self.conv1(x) + x = self.bn1(x) + x = self.relu(x) + x = self.maxpool(x) + + x = self.layer1(x) + x = self.layer2(x) + x = self.layer3(x) + x = self.layer4(x) + + x = self.avgpool(x) + x = torch.flatten(x, 1) + x = self.fc(x) + + return x + + def forward(self, x): + return self._forward_impl(x) + + +def _resnet(arch, block, layers, pretrained, progress, **kwargs): + model = ResNet(block, layers, **kwargs) + if pretrained: + state_dict = load_state_dict_from_url(model_urls[arch], + progress=progress) + model.load_state_dict(state_dict) + return model + + +def resnet18(pretrained=False, progress=True, **kwargs): + r"""ResNet-18 model from + `"Deep Residual Learning for Image Recognition" `_ + + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + return _resnet('resnet18', BasicBlock, [2, 2, 2, 2], pretrained, progress, + **kwargs) + + +def resnet34(pretrained=False, progress=True, **kwargs): + r"""ResNet-34 model from + `"Deep Residual Learning for Image Recognition" `_ + + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + return _resnet('resnet34', BasicBlock, [3, 4, 6, 3], pretrained, progress, + **kwargs) + + +def resnet50(pretrained=False, progress=True, **kwargs): + r"""ResNet-50 model from + `"Deep Residual Learning for Image Recognition" `_ + + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + return _resnet('resnet50', Bottleneck, [3, 4, 6, 3], pretrained, progress, + **kwargs) + + +def resnet101(pretrained=False, progress=True, **kwargs): + r"""ResNet-101 model from + `"Deep Residual Learning for Image Recognition" `_ + + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + return _resnet('resnet101', Bottleneck, [3, 4, 23, 3], pretrained, progress, + **kwargs) + + +def resnet152(pretrained=False, progress=True, **kwargs): + r"""ResNet-152 model from + `"Deep Residual Learning for Image Recognition" `_ + + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + return _resnet('resnet152', Bottleneck, [3, 8, 36, 3], pretrained, progress, + **kwargs) + + +def resnext50_32x4d(pretrained=False, progress=True, **kwargs): + r"""ResNeXt-50 32x4d model from + `"Aggregated Residual Transformation for Deep Neural Networks" `_ + + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + kwargs['groups'] = 32 + kwargs['width_per_group'] = 4 + return _resnet('resnext50_32x4d', Bottleneck, [3, 4, 6, 3], + pretrained, progress, **kwargs) + + +def resnext101_32x8d(pretrained=False, progress=True, **kwargs): + r"""ResNeXt-101 32x8d model from + `"Aggregated Residual Transformation for Deep Neural Networks" `_ + + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + kwargs['groups'] = 32 + kwargs['width_per_group'] = 8 + return _resnet('resnext101_32x8d', Bottleneck, [3, 4, 23, 3], + pretrained, progress, **kwargs) + + +def wide_resnet50_2(pretrained=False, progress=True, **kwargs): + r"""Wide ResNet-50-2 model from + `"Wide Residual Networks" `_ + + The model is the same as ResNet except for the bottleneck number of channels + which is twice larger in every block. The number of channels in outer 1x1 + convolutions is the same, e.g. last block in ResNet-50 has 2048-512-2048 + channels, and in Wide ResNet-50-2 has 2048-1024-2048. + + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + kwargs['width_per_group'] = 64 * 2 + return _resnet('wide_resnet50_2', Bottleneck, [3, 4, 6, 3], + pretrained, progress, **kwargs) + + +def wide_resnet101_2(pretrained=False, progress=True, **kwargs): + r"""Wide ResNet-101-2 model from + `"Wide Residual Networks" `_ + + The model is the same as ResNet except for the bottleneck number of channels + which is twice larger in every block. The number of channels in outer 1x1 + convolutions is the same, e.g. last block in ResNet-50 has 2048-512-2048 + channels, and in Wide ResNet-50-2 has 2048-1024-2048. + + Args: + pretrained (bool): If True, returns a model pre-trained on ImageNet + progress (bool): If True, displays a progress bar of the download to stderr + """ + kwargs['width_per_group'] = 64 * 2 + return _resnet('wide_resnet101_2', Bottleneck, [3, 4, 23, 3], + pretrained, progress, **kwargs) diff --git a/models/trainer.py b/models/trainer.py new file mode 100644 index 0000000000000000000000000000000000000000..6efc13112689ccfb33220ab26970df861ac25bc4 --- /dev/null +++ b/models/trainer.py @@ -0,0 +1,297 @@ +import numpy as np +import matplotlib.pyplot as plt +import os + +import utils +from models.networks import * + +import torch +import torch.optim as optim + +from misc.metric_tool import ConfuseMatrixMeter +from models.losses import cross_entropy +import models.losses as losses + +from misc.logger_tool import Logger, Timer + +from utils import de_norm + + +class CDTrainer(): + + def __init__(self, args, dataloaders): + + self.dataloaders = dataloaders + + self.n_class = args.n_class + # define G + self.net_G = define_G(args=args, gpu_ids=args.gpu_ids) + + self.device = torch.device("cuda:%s" % args.gpu_ids[0] if torch.cuda.is_available() and len(args.gpu_ids)>0 + else "cpu") + print(self.device) + + # Learning rate and Beta1 for Adam optimizers + self.lr = args.lr + + # define optimizers + self.optimizer_G = optim.SGD(self.net_G.parameters(), lr=self.lr, + momentum=0.9, + weight_decay=5e-4) + + # define lr schedulers + self.exp_lr_scheduler_G = get_scheduler(self.optimizer_G, args) + + self.running_metric = ConfuseMatrixMeter(n_class=2) + + # define logger file + logger_path = os.path.join(args.checkpoint_dir, 'log.txt') + self.logger = Logger(logger_path) + self.logger.write_dict_str(args.__dict__) + # define timer + self.timer = Timer() + self.batch_size = args.batch_size + + # training log + self.epoch_acc = 0 + self.best_val_acc = 0.0 + self.best_epoch_id = 0 + self.epoch_to_start = 0 + self.max_num_epochs = args.max_epochs + + self.global_step = 0 + self.steps_per_epoch = len(dataloaders['train']) + self.total_steps = (self.max_num_epochs - self.epoch_to_start)*self.steps_per_epoch + + self.G_pred = None + self.pred_vis = None + self.batch = None + self.G_loss = None + self.is_training = False + self.batch_id = 0 + self.epoch_id = 0 + self.checkpoint_dir = args.checkpoint_dir + self.vis_dir = args.vis_dir + + # define the loss functions + if args.loss == 'ce': + self._pxl_loss = cross_entropy + elif args.loss == 'bce': + self._pxl_loss = losses.binary_ce + else: + raise NotImplemented(args.loss) + + self.VAL_ACC = np.array([], np.float32) + if os.path.exists(os.path.join(self.checkpoint_dir, 'val_acc.npy')): + self.VAL_ACC = np.load(os.path.join(self.checkpoint_dir, 'val_acc.npy')) + self.TRAIN_ACC = np.array([], np.float32) + if os.path.exists(os.path.join(self.checkpoint_dir, 'train_acc.npy')): + self.TRAIN_ACC = np.load(os.path.join(self.checkpoint_dir, 'train_acc.npy')) + + # check and create model dir + if os.path.exists(self.checkpoint_dir) is False: + os.mkdir(self.checkpoint_dir) + if os.path.exists(self.vis_dir) is False: + os.mkdir(self.vis_dir) + + + def _load_checkpoint(self, ckpt_name='last_ckpt.pt'): + + if os.path.exists(os.path.join(self.checkpoint_dir, ckpt_name)): + self.logger.write('loading last checkpoint...\n') + # load the entire checkpoint + checkpoint = torch.load(os.path.join(self.checkpoint_dir, ckpt_name), + map_location=self.device) + # update net_G states + self.net_G.load_state_dict(checkpoint['model_G_state_dict']) + + self.optimizer_G.load_state_dict(checkpoint['optimizer_G_state_dict']) + self.exp_lr_scheduler_G.load_state_dict( + checkpoint['exp_lr_scheduler_G_state_dict']) + + self.net_G.to(self.device) + + # update some other states + self.epoch_to_start = checkpoint['epoch_id'] + 1 + self.best_val_acc = checkpoint['best_val_acc'] + self.best_epoch_id = checkpoint['best_epoch_id'] + + self.total_steps = (self.max_num_epochs - self.epoch_to_start)*self.steps_per_epoch + + self.logger.write('Epoch_to_start = %d, Historical_best_acc = %.4f (at epoch %d)\n' % + (self.epoch_to_start, self.best_val_acc, self.best_epoch_id)) + self.logger.write('\n') + + else: + print('training from scratch...') + + def _timer_update(self): + self.global_step = (self.epoch_id-self.epoch_to_start) * self.steps_per_epoch + self.batch_id + + self.timer.update_progress((self.global_step + 1) / self.total_steps) + est = self.timer.estimated_remaining() + imps = (self.global_step + 1) * self.batch_size / self.timer.get_stage_elapsed() + return imps, est + + def _visualize_pred(self): + pred = torch.argmax(self.G_pred, dim=1, keepdim=True) + pred_vis = pred * 255 + return pred_vis + + def _save_checkpoint(self, ckpt_name): + torch.save({ + 'epoch_id': self.epoch_id, + 'best_val_acc': self.best_val_acc, + 'best_epoch_id': self.best_epoch_id, + 'model_G_state_dict': self.net_G.state_dict(), + 'optimizer_G_state_dict': self.optimizer_G.state_dict(), + 'exp_lr_scheduler_G_state_dict': self.exp_lr_scheduler_G.state_dict(), + }, os.path.join(self.checkpoint_dir, ckpt_name)) + + def _update_lr_schedulers(self): + self.exp_lr_scheduler_G.step() + + def _update_metric(self): + """ + update metric + """ + target = self.batch['L'].to(self.device).detach() + G_pred = self.G_pred.detach() + + G_pred = torch.argmax(G_pred, dim=1) + + current_score = self.running_metric.update_cm(pr=G_pred.cpu().numpy(), gt=target.cpu().numpy()) + return current_score + + def _collect_running_batch_states(self): + + running_acc = self._update_metric() + + m = len(self.dataloaders['train']) + if self.is_training is False: + m = len(self.dataloaders['val']) + + imps, est = self._timer_update() + if np.mod(self.batch_id, 100) == 1: + message = 'Is_training: %s. [%d,%d][%d,%d], imps: %.2f, est: %.2fh, G_loss: %.5f, running_mf1: %.5f\n' %\ + (self.is_training, self.epoch_id, self.max_num_epochs-1, self.batch_id, m, + imps*self.batch_size, est, + self.G_loss.item(), running_acc) + self.logger.write(message) + + + if np.mod(self.batch_id, 500) == 1: + vis_input = utils.make_numpy_grid(de_norm(self.batch['A'])) + vis_input2 = utils.make_numpy_grid(de_norm(self.batch['B'])) + + vis_pred = utils.make_numpy_grid(self._visualize_pred()) + + vis_gt = utils.make_numpy_grid(self.batch['L']) + vis = np.concatenate([vis_input, vis_input2, vis_pred, vis_gt], axis=0) + vis = np.clip(vis, a_min=0.0, a_max=1.0) + file_name = os.path.join( + self.vis_dir, 'istrain_'+str(self.is_training)+'_'+ + str(self.epoch_id)+'_'+str(self.batch_id)+'.jpg') + plt.imsave(file_name, vis) + + def _collect_epoch_states(self): + scores = self.running_metric.get_scores() + self.epoch_acc = scores['mf1'] + self.logger.write('Is_training: %s. Epoch %d / %d, epoch_mF1= %.5f\n' % + (self.is_training, self.epoch_id, self.max_num_epochs-1, self.epoch_acc)) + message = '' + for k, v in scores.items(): + message += '%s: %.5f ' % (k, v) + self.logger.write(message+'\n') + self.logger.write('\n') + + def _update_checkpoints(self): + + # save current model + self._save_checkpoint(ckpt_name='last_ckpt.pt') + self.logger.write('Lastest model updated. Epoch_acc=%.4f, Historical_best_acc=%.4f (at epoch %d)\n' + % (self.epoch_acc, self.best_val_acc, self.best_epoch_id)) + self.logger.write('\n') + + # update the best model (based on eval acc) + if self.epoch_acc > self.best_val_acc: + self.best_val_acc = self.epoch_acc + self.best_epoch_id = self.epoch_id + self._save_checkpoint(ckpt_name='best_ckpt.pt') + self.logger.write('*' * 10 + 'Best model updated!\n') + self.logger.write('\n') + + def _update_training_acc_curve(self): + # update train acc curve + self.TRAIN_ACC = np.append(self.TRAIN_ACC, [self.epoch_acc]) + np.save(os.path.join(self.checkpoint_dir, 'train_acc.npy'), self.TRAIN_ACC) + + def _update_val_acc_curve(self): + # update val acc curve + self.VAL_ACC = np.append(self.VAL_ACC, [self.epoch_acc]) + np.save(os.path.join(self.checkpoint_dir, 'val_acc.npy'), self.VAL_ACC) + + def _clear_cache(self): + self.running_metric.clear() + + + def _forward_pass(self, batch): + self.batch = batch + img_in1 = batch['A'].to(self.device) + img_in2 = batch['B'].to(self.device) + self.G_pred = self.net_G(img_in1, img_in2) + + + def _backward_G(self): + gt = self.batch['L'].to(self.device).long() + self.G_loss = self._pxl_loss(self.G_pred, gt) + self.G_loss.backward() + + + def train_models(self): + + self._load_checkpoint() + + # loop over the dataset multiple times + for self.epoch_id in range(self.epoch_to_start, self.max_num_epochs): + + ################## train ################# + ########################################## + self._clear_cache() + self.is_training = True + self.net_G.train() # Set model to training mode + # Iterate over data. + self.logger.write('lr: %0.7f\n' % self.optimizer_G.param_groups[0]['lr']) + for self.batch_id, batch in enumerate(self.dataloaders['train'], 0): + self._forward_pass(batch) + # update G + self.optimizer_G.zero_grad() + self._backward_G() + self.optimizer_G.step() + self._collect_running_batch_states() + self._timer_update() + + self._collect_epoch_states() + self._update_training_acc_curve() + self._update_lr_schedulers() + + + ################## Eval ################## + ########################################## + self.logger.write('Begin evaluation...\n') + self._clear_cache() + self.is_training = False + self.net_G.eval() + + # Iterate over data. + for self.batch_id, batch in enumerate(self.dataloaders['val'], 0): + with torch.no_grad(): + self._forward_pass(batch) + self._collect_running_batch_states() + self._collect_epoch_states() + + ########### Update_Checkpoints ########### + ########################################## + self._update_val_acc_curve() + self._update_checkpoints() + diff --git a/samples/A/test_102_0512_0000.png b/samples/A/test_102_0512_0000.png new file mode 100644 index 0000000000000000000000000000000000000000..58505396e66c42a6e66408926f05376e695ac4b3 --- /dev/null +++ b/samples/A/test_102_0512_0000.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8e3221cef502953e06538e1ebddfdb609780c0c1cb6366aa3eee9208012aeab3 +size 78722 diff --git a/samples/A/test_113_0256.png b/samples/A/test_113_0256.png new file mode 100644 index 0000000000000000000000000000000000000000..d30ffb06fd3e368face026c0d9d499f47294d2ad --- /dev/null +++ b/samples/A/test_113_0256.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:0910d9eede706b3c315e6369f0ea2963b728bf29bc4e140d03c0902c16c5f642 +size 904786 diff --git a/samples/A/test_121_0768_0256.png b/samples/A/test_121_0768_0256.png new file mode 100644 index 0000000000000000000000000000000000000000..c5853e1f6051dd7fb911c716813162a666f2d8ce --- /dev/null +++ b/samples/A/test_121_0768_0256.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f51526b822696395992565a70f0d2223f6af7c5a1bf1dfef47edc9cd7450e257 +size 102129 diff --git a/samples/A/test_2_0000_0000.png b/samples/A/test_2_0000_0000.png new file mode 100644 index 0000000000000000000000000000000000000000..e6db8e84f08be9c77cf148ecaaa28726736cd586 --- /dev/null +++ b/samples/A/test_2_0000_0000.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:99572d7b22f39f2c3947f817ab857784ab85cdbd1307fa0892d41c564f1966bd +size 131272 diff --git a/samples/A/test_2_0000_0512.png b/samples/A/test_2_0000_0512.png new file mode 100644 index 0000000000000000000000000000000000000000..d931bdba2e50cea4a95e1b6ad014120993a34453 --- /dev/null +++ b/samples/A/test_2_0000_0512.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:2b53d8b252e3aa0d205c0d5af432163d4637cc4e6793480ff74d27de5959de10 +size 140014 diff --git a/samples/A/test_55_0256_0000.png b/samples/A/test_55_0256_0000.png new file mode 100644 index 0000000000000000000000000000000000000000..e29b873c2208ae3faf53ba4af198391c60bbe3a6 --- /dev/null +++ b/samples/A/test_55_0256_0000.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fab174503235b57c43a8cf8b8bf4d4a2fb952990b03ef23aa7b86f3d9eb69cab +size 115205 diff --git a/samples/A/test_77_0512_0256.png b/samples/A/test_77_0512_0256.png new file mode 100644 index 0000000000000000000000000000000000000000..c8cf810ea41d965f8730b6a546de5d9ce808a107 --- /dev/null +++ b/samples/A/test_77_0512_0256.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:fe56308b53815e8b8c81705741236a93d1f75948ce3ad66168a0aed56f51684a +size 158709 diff --git a/samples/A/test_7_0256_0512.png b/samples/A/test_7_0256_0512.png new file mode 100644 index 0000000000000000000000000000000000000000..8a054fb8ba011d3704686b46bd85fcd18652775a --- /dev/null +++ b/samples/A/test_7_0256_0512.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:b6c49095a63807ed29acaadef7153a4e66e7b22f9fa974c07a58f2831fcfe7f3 +size 146740 diff --git a/samples/A/train_36_0512_0512.png b/samples/A/train_36_0512_0512.png new file mode 100644 index 0000000000000000000000000000000000000000..f64907adb31ff5d83d53b2d050e14e5575642906 --- /dev/null +++ b/samples/A/train_36_0512_0512.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:faa96e07dc9c4c22be2f5200555712269d260c27977f68f660cc0f1b4cab50cd +size 107958 diff --git a/samples/A/train_386_0512_0768.png b/samples/A/train_386_0512_0768.png new file mode 100644 index 0000000000000000000000000000000000000000..b606e517984adcfbf8e0fa58b99c9c158fa7c83f --- /dev/null +++ b/samples/A/train_386_0512_0768.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:47e0a39a005e32b91a465564f7f244ddf3cf4b288909b492d041c83640fa19d5 +size 126618 diff --git a/samples/A/train_412_0512_0768.png b/samples/A/train_412_0512_0768.png new file mode 100644 index 0000000000000000000000000000000000000000..2833c271e58f2f66d970cbd78c3d71bf0a62a4df --- /dev/null +++ b/samples/A/train_412_0512_0768.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:72f3dbbb090874de7c84f445c24d103b9b8a4254c2549d6aa55586ea662cbe25 +size 102328 diff --git a/samples/A/val_27_0000_0256.png b/samples/A/val_27_0000_0256.png new file mode 100644 index 0000000000000000000000000000000000000000..78555eef5b7066652fc1cde77c296242d11dd3f9 --- /dev/null +++ b/samples/A/val_27_0000_0256.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f8791b7acca8632d1684e5aabe9d01a5b30e06b4b8000a71074f038c76a8c5fa +size 109623 diff --git a/samples/B/test_102_0512_0000.png b/samples/B/test_102_0512_0000.png new file mode 100644 index 0000000000000000000000000000000000000000..43d71aebf5c8006757cd88602a5c795a5f635997 --- /dev/null +++ b/samples/B/test_102_0512_0000.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c9ec0d2f4f58e3537a1f4685ef16368c5546300d04aef9e8a84d9a6dd3a3cc48 +size 127399 diff --git a/samples/B/test_113_0256.png b/samples/B/test_113_0256.png new file mode 100644 index 0000000000000000000000000000000000000000..cfad149599742dd34a8bd6e11c03df0c9fdadf78 --- /dev/null +++ b/samples/B/test_113_0256.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f701c5918da864e2846d073bc4545e0ef29899fdaccad7736e7a39eef56d0021 +size 792835 diff --git a/samples/B/test_121_0768_0256.png b/samples/B/test_121_0768_0256.png new file mode 100644 index 0000000000000000000000000000000000000000..efeb2440d11104732cf3bfcbd02cdf8982b2e9b1 --- /dev/null +++ b/samples/B/test_121_0768_0256.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:867fdc2fdad0292cc6561e5aab71fc97c8d67bedd6ebd8d1012e094e3b6092b9 +size 128079 diff --git a/samples/B/test_2_0000_0000.png b/samples/B/test_2_0000_0000.png new file mode 100644 index 0000000000000000000000000000000000000000..84b916f43dc7775cd981780923c67003cc9b68dd --- /dev/null +++ b/samples/B/test_2_0000_0000.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c37302e1bde614f032291544d371bbe3631f773e3cf4d7130cee5dabfdbae961 +size 129859 diff --git a/samples/B/test_2_0000_0512.png b/samples/B/test_2_0000_0512.png new file mode 100644 index 0000000000000000000000000000000000000000..d3223e33b00e1a36b01985ab25eda74f4fc0ad45 --- /dev/null +++ b/samples/B/test_2_0000_0512.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a3235f14acf81726e9036f0e4fe16b87e8d59d75975b3e688bf47ee7f532626a +size 132645 diff --git a/samples/B/test_55_0256_0000.png b/samples/B/test_55_0256_0000.png new file mode 100644 index 0000000000000000000000000000000000000000..06c2382cd92a41bc2654f74079e8b2f484848493 --- /dev/null +++ b/samples/B/test_55_0256_0000.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:be7248f746cfdb1c592489be9fe23ae2b689961a668f8d047a7b0c398ec94018 +size 138358 diff --git a/samples/B/test_77_0512_0256.png b/samples/B/test_77_0512_0256.png new file mode 100644 index 0000000000000000000000000000000000000000..bb13a516fe37bb882f6dcfcbb3d6e183d06ff50a --- /dev/null +++ b/samples/B/test_77_0512_0256.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:4f53af9029bce438330555c491bbf901b8252eb0062f27351954b5e898eccdaa +size 134562 diff --git a/samples/B/test_7_0256_0512.png b/samples/B/test_7_0256_0512.png new file mode 100644 index 0000000000000000000000000000000000000000..abe3c4173826865b49735506206c005308c5ce5e --- /dev/null +++ b/samples/B/test_7_0256_0512.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:034282f91bfe65709bb147a31f03b4cf713f0980f2da00a576ed30c535628eb0 +size 135170 diff --git a/samples/B/train_36_0512_0512.png b/samples/B/train_36_0512_0512.png new file mode 100644 index 0000000000000000000000000000000000000000..7defb0f774abf785f678208fd6eef4cf4e8d5ae5 --- /dev/null +++ b/samples/B/train_36_0512_0512.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:17bc7800c8dd2a467c6ca89b785dbf15dd1923fbd386fa7a4ddce4324ff6ed67 +size 145392 diff --git a/samples/B/train_386_0512_0768.png b/samples/B/train_386_0512_0768.png new file mode 100644 index 0000000000000000000000000000000000000000..58759aa6d711914e050f8120d60e25ace4a0b327 --- /dev/null +++ b/samples/B/train_386_0512_0768.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c95df4f0d903c88f3c35515b5577c69f9292f88ee953dd5444a8332680f8a11a +size 107747 diff --git a/samples/B/train_412_0512_0768.png b/samples/B/train_412_0512_0768.png new file mode 100644 index 0000000000000000000000000000000000000000..57c9935a05349202e0b7d440475243694f018c32 --- /dev/null +++ b/samples/B/train_412_0512_0768.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:7b246772966f5fbd6200cd29656aaf62c9868a46afbd813ad5bbc4ea9a6f6669 +size 133916 diff --git a/samples/B/val_27_0000_0256.png b/samples/B/val_27_0000_0256.png new file mode 100644 index 0000000000000000000000000000000000000000..da99908e4f6ee4b65a6384bd31cb063160f0d0a1 --- /dev/null +++ b/samples/B/val_27_0000_0256.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5b90bf93b8cc5a810dcf4447316a87c93173ffe633f670fb5c421d8d2ad39679 +size 134770 diff --git a/samples/label/test_102_0512_0000.png b/samples/label/test_102_0512_0000.png new file mode 100644 index 0000000000000000000000000000000000000000..92d983357b9b3a3489517d23f2fca8b45c49cf3f --- /dev/null +++ b/samples/label/test_102_0512_0000.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:a4223b16cda4fda2bda583e4232603ad91a8700676d1a977401851fb746818ff +size 1186 diff --git a/samples/label/test_121_0768_0256.png b/samples/label/test_121_0768_0256.png new file mode 100644 index 0000000000000000000000000000000000000000..7a7fecc32ddfa584ab22e5d53814bfcdfdb46315 --- /dev/null +++ b/samples/label/test_121_0768_0256.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:ce1fec32f2d6792ea3eddcea7861899c5985ccdfde61a179098694f28d927a3a +size 2098 diff --git a/samples/label/test_2_0000_0000.png b/samples/label/test_2_0000_0000.png new file mode 100644 index 0000000000000000000000000000000000000000..02b4eda13860cd006ec3cbc83a0787701ae66036 --- /dev/null +++ b/samples/label/test_2_0000_0000.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:c2094d81c4dbd1739cf5b028dff2bdad9eb0c0c6a8c8150f863fa08de338e107 +size 1075 diff --git a/samples/label/test_2_0000_0512.png b/samples/label/test_2_0000_0512.png new file mode 100644 index 0000000000000000000000000000000000000000..4e63d30891b4f12c8076a5f3e4a7e40a60a6d0bc --- /dev/null +++ b/samples/label/test_2_0000_0512.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:8d6ce41349c2b149edd5700acc0c42657c2ce553d2f52b39cc795e0924828509 +size 1758 diff --git a/samples/label/test_55_0256_0000.png b/samples/label/test_55_0256_0000.png new file mode 100644 index 0000000000000000000000000000000000000000..25071f28bea15b4dac6e2919761ab18f1634f4c2 --- /dev/null +++ b/samples/label/test_55_0256_0000.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:9a4506913ea0880e2deb487bbf59206ff62ad2a1db3935befb34352302b8037d +size 1762 diff --git a/samples/label/test_77_0512_0256.png b/samples/label/test_77_0512_0256.png new file mode 100644 index 0000000000000000000000000000000000000000..94763af135ddf9cd05ec778ca73fb3ed690d481c --- /dev/null +++ b/samples/label/test_77_0512_0256.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:32618a3acc9c976e30915c4689d93cde1d99cba623de625e17b2037b0cf56c9a +size 922 diff --git a/samples/label/test_7_0256_0512.png b/samples/label/test_7_0256_0512.png new file mode 100644 index 0000000000000000000000000000000000000000..e9ebc5228e1ecfbf64492f7cdfd09646ceda6141 --- /dev/null +++ b/samples/label/test_7_0256_0512.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:bcdb87fed6c57ac3cf5b57439c52824909f888e59322a106b45363b7aae45ef4 +size 1576 diff --git a/samples/label/train_36_0512_0512.png b/samples/label/train_36_0512_0512.png new file mode 100644 index 0000000000000000000000000000000000000000..30f9cbe67dcedda7073849ceca6e589a0677d333 --- /dev/null +++ b/samples/label/train_36_0512_0512.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:d3c12f9ac94c0a5c797fd37303cfd04300757ace1bfbd6e1188ef7117f91cf92 +size 2074 diff --git a/samples/label/train_386_0512_0768.png b/samples/label/train_386_0512_0768.png new file mode 100644 index 0000000000000000000000000000000000000000..c1f0672c8f2366699bf44fea27b3e07f04922485 --- /dev/null +++ b/samples/label/train_386_0512_0768.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:39962cd5bc9f4f0446341d3e6e0c6c37336ddeb2e026a17a3d06bb6cb3266daf +size 141 diff --git a/samples/label/train_412_0512_0768.png b/samples/label/train_412_0512_0768.png new file mode 100644 index 0000000000000000000000000000000000000000..7a378e2b6473b9d587a459498b06492ae5ba2b42 --- /dev/null +++ b/samples/label/train_412_0512_0768.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:f3942ebb6427992354b09c0c5ed58d9e1b7e33691a5c8e51caacf7bc52a56027 +size 1685 diff --git a/samples/label/val_27_0000_0256.png b/samples/label/val_27_0000_0256.png new file mode 100644 index 0000000000000000000000000000000000000000..2bee1cd953793cbe2d1d06ce1f440c7404dbb598 --- /dev/null +++ b/samples/label/val_27_0000_0256.png @@ -0,0 +1,3 @@ +version https://git-lfs.github.com/spec/v1 +oid sha256:5db1abaa74eff8c50d2f728f0af7b69165a8b74464f1ae4c6b2716b16687a2d1 +size 1500 diff --git a/samples/list/demo.txt b/samples/list/demo.txt new file mode 100644 index 0000000000000000000000000000000000000000..2426e867e2478c87314ed9ddf917088392dbc843 --- /dev/null +++ b/samples/list/demo.txt @@ -0,0 +1,7 @@ +test_77_0512_0256.png +test_102_0512_0000.png +test_121_0768_0256.png +test_2_0000_0000.png +test_2_0000_0512.png +test_7_0256_0512.png +test_55_0256_0000.png diff --git a/scripts/eval.sh b/scripts/eval.sh new file mode 100644 index 0000000000000000000000000000000000000000..3bb7aaa79277304ef255f23bd6933139a60db4ed --- /dev/null +++ b/scripts/eval.sh @@ -0,0 +1,13 @@ +#!/usr/bin/env bash + +gpus=0 + +data_name=LEVIR +net_G=base_transformer_pos_s4_dd8_dedim8 +split=test +project_name=BIT_LEVIR +checkpoint_name=best_ckpt.pt + +python eval_cd.py --split ${split} --net_G ${net_G} --checkpoint_name ${checkpoint_name} --gpu_ids ${gpus} --project_name ${project_name} --data_name ${data_name} + + diff --git a/scripts/run_cd.sh b/scripts/run_cd.sh new file mode 100644 index 0000000000000000000000000000000000000000..73721639c45180685c2bcc17aac1349302942bfd --- /dev/null +++ b/scripts/run_cd.sh @@ -0,0 +1,21 @@ +#!/usr/bin/env bash + +gpus=0 +checkpoint_root=checkpoints +data_name=LEVIR + +img_size=256 +batch_size=8 +lr=0.01 +max_epochs=200 +net_G=base_transformer_pos_s4_dd8 +#base_resnet18 +#base_transformer_pos_s4_dd8 +#base_transformer_pos_s4_dd8_dedim8 +lr_policy=linear + +split=trainval +split_val=test +project_name=CD_${net_G}_${data_name}_b${batch_size}_lr${lr}_${split}_${split_val}_${max_epochs}_${lr_policy} + +python main_cd.py --img_size ${img_size} --checkpoint_root ${checkpoint_root} --lr_policy ${lr_policy} --split ${split} --split_val ${split_val} --net_G ${net_G} --gpu_ids ${gpus} --max_epochs ${max_epochs} --project_name ${project_name} --batch_size ${batch_size} --data_name ${data_name} --lr ${lr} \ No newline at end of file diff --git a/utils.py b/utils.py new file mode 100644 index 0000000000000000000000000000000000000000..5216c83c107c11c5c1485a1178392777b4619426 --- /dev/null +++ b/utils.py @@ -0,0 +1,84 @@ +import numpy as np +import torch +from torch.utils.data import DataLoader +from torchvision import utils + +import data_config +from datasets.CD_dataset import CDDataset + + +def get_loader(data_name, img_size=256, batch_size=8, split='test', + is_train=False, dataset='CDDataset'): + dataConfig = data_config.DataConfig().get_data_config(data_name) + root_dir = dataConfig.root_dir + label_transform = dataConfig.label_transform + + if dataset == 'CDDataset': + data_set = CDDataset(root_dir=root_dir, split=split, + img_size=img_size, is_train=is_train, + label_transform=label_transform) + else: + raise NotImplementedError( + 'Wrong dataset name %s (choose one from [CDDataset])' + % dataset) + + shuffle = is_train + dataloader = DataLoader(data_set, batch_size=batch_size, + shuffle=shuffle, num_workers=4) + + return dataloader + + +def get_loaders(args): + + data_name = args.data_name + dataConfig = data_config.DataConfig().get_data_config(data_name) + root_dir = dataConfig.root_dir + label_transform = dataConfig.label_transform + split = args.split + split_val = 'val' + if hasattr(args, 'split_val'): + split_val = args.split_val + if args.dataset == 'CDDataset': + training_set = CDDataset(root_dir=root_dir, split=split, + img_size=args.img_size,is_train=True, + label_transform=label_transform) + val_set = CDDataset(root_dir=root_dir, split=split_val, + img_size=args.img_size,is_train=False, + label_transform=label_transform) + else: + raise NotImplementedError( + 'Wrong dataset name %s (choose one from [CDDataset,])' + % args.dataset) + + datasets = {'train': training_set, 'val': val_set} + dataloaders = {x: DataLoader(datasets[x], batch_size=args.batch_size, + shuffle=True, num_workers=args.num_workers) + for x in ['train', 'val']} + + return dataloaders + + +def make_numpy_grid(tensor_data, pad_value=0,padding=0): + tensor_data = tensor_data.detach() + vis = utils.make_grid(tensor_data, pad_value=pad_value,padding=padding) + vis = np.array(vis.cpu()).transpose((1,2,0)) + if vis.shape[2] == 1: + vis = np.stack([vis, vis, vis], axis=-1) + return vis + + +def de_norm(tensor_data): + return tensor_data * 0.5 + 0.5 + + +def get_device(args): + # set gpu ids + str_ids = args.gpu_ids.split(',') + args.gpu_ids = [] + for str_id in str_ids: + id = int(str_id) + if id >= 0: + args.gpu_ids.append(id) + if len(args.gpu_ids) > 0: + torch.cuda.set_device(args.gpu_ids[0]) \ No newline at end of file