CoolFace
Modelpublic

KyanChen/BuildingExtraction

sourceHugging Faceupdated 4y agoView on Hugging Face
3likes
Train.py186 linesDownload Raw Back to root
1import os2# Change the numbers when you want to train with specific gpus3# os.environ['CUDA_VISIBLE_DEVICES'] = '0, 1, 2, 3'4import torch5from STTNet import STTNet6import torch.nn.functional as F7from Utils.Datasets import get_data_loader8from Utils.Utils import make_numpy_img, inv_normalize_img, encode_onehot_to_mask, get_metrics, Logger9import matplotlib.pyplot as plt10import numpy as np11from collections import OrderedDict12from torch.optim.lr_scheduler import MultiStepLR13 14if __name__ == '__main__':15    model_infos = {16        # vgg16_bn, resnet50, resnet1817        'backbone': 'resnet50',18        'pretrained': True,19        'out_keys': ['block4'],20        'in_channel': 3,21        'n_classes': 2,22        'top_k_s': 64,23        'top_k_c': 16,24        'encoder_pos': True,25        'decoder_pos': True,26        'model_pattern': ['X', 'A', 'S', 'C'],27 28        'BATCH_SIZE': 8,29        'IS_SHUFFLE': True,30        'NUM_WORKERS': 0,31        'DATASET': 'Tools/generate_dep_info/train_data.csv',32        'model_path': 'Checkpoints',33        'log_path': 'Results',34        # if you need the validation process.35        'IS_VAL': True,36        'VAL_BATCH_SIZE': 4,37        'VAL_DATASET': 'Tools/generate_dep_info/val_data.csv',38        # if you need the test process.39        'IS_TEST': True,40        'TEST_DATASET': 'Tools/generate_dep_info/test_data.csv',41        'IMG_SIZE': [512, 512],42        'PHASE': 'seg',43 44        # INRIA Dataset45        'PRIOR_MEAN': [0.40672500537632994, 0.42829032416229895, 0.39331840468605667],46        'PRIOR_STD': [0.029498464618176873, 0.027740088491668233, 0.028246722411879095],47        # # # WHU Dataset48        # 'PRIOR_MEAN': [0.4352682576428411, 0.44523221318154493, 0.41307610541534784],49        # 'PRIOR_STD': [0.026973196780331585, 0.026424642808887323, 0.02791246590291434],50 51        # if you want to load state dict52        'load_checkpoint_path': r'E:\BuildingExtractionDataset\INRIA_ckpt_latest.pt',53        # if you want to resume a checkpoint54        'resume_checkpoint_path': '',55 56    }57    os.makedirs(model_infos['model_path'], exist_ok=True)58    if model_infos['IS_VAL']:59        os.makedirs(model_infos['log_path']+'/val', exist_ok=True)60    if model_infos['IS_TEST']:61        os.makedirs(model_infos['log_path']+'/test', exist_ok=True)62    logger = Logger(model_infos['log_path'] + '/log.log')63 64    data_loaders = get_data_loader(model_infos)65    loss_weight = 0.166    model = STTNet(**model_infos)67 68    epoch_start = 069    if model_infos['load_checkpoint_path'] is not None and os.path.exists(model_infos['load_checkpoint_path']):70        logger.write(f'load checkpoint from {model_infos["load_checkpoint_path"]}\n')71        state_dict = torch.load(model_infos['load_checkpoint_path'], map_location='cpu')72        model_dict = state_dict['model_state_dict']73        try:74            model_dict = OrderedDict({k.replace('module.', ''): v for k, v in model_dict.items()})75            model.load_state_dict(model_dict)76        except Exception as e:77            model.load_state_dict(model_dict)78    if model_infos['resume_checkpoint_path'] is not None and os.path.exists(model_infos['resume_checkpoint_path']):79        logger.write(f'resume checkpoint path from {model_infos["resume_checkpoint_path"]}\n')80        state_dict = torch.load(model_infos['resume_checkpoint_path'], map_location='cpu')81        epoch_start = state_dict['epoch_id']82        model_dict = state_dict['model_state_dict']83        logger.write(f'resume checkpoint from epoch {epoch_start}\n')84        try:85            model_dict = OrderedDict({k.replace('module.', ''): v for k, v in model_dict.items()})86            model.load_state_dict(model_dict)87        except Exception as e:88            model.load_state_dict(model_dict)89    model = model.cuda()90    device_ids = range(torch.cuda.device_count())91    if len(device_ids) > 1:92        model = torch.nn.DataParallel(model, device_ids=device_ids)93        logger.write(f'Use GPUs: {device_ids}\n')94    else:95        logger.write(f'Use GPUs: 1\n')96    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)97    max_epoch = 30098    scheduler = MultiStepLR(optimizer, [int(max_epoch*2/3), int(max_epoch*5/6)], 0.5)99 100    for epoch_id in range(epoch_start, max_epoch):101        pattern = 'train'102        model.train()  # Set model to training mode103        for batch_id, batch in enumerate(data_loaders[pattern]):104            # Get data105            img_batch = batch['img'].cuda()106            label_batch = batch['label'].cuda()107 108            # inference109            optimizer.zero_grad()110            logits, att_branch_output = model(img_batch)111 112            # compute loss113            label_downs = F.interpolate(label_batch, att_branch_output.size()[2:], mode='nearest')114            loss_branch = F.binary_cross_entropy_with_logits(att_branch_output, label_downs)115            loss_master = F.binary_cross_entropy_with_logits(logits, label_batch)116            loss = loss_master + loss_weight * loss_branch117            # loss backward118            loss.backward()119            optimizer.step()120 121            if batch_id % 20 == 1:122                logger.write(123                    f'{pattern}: {epoch_id}/{max_epoch} {batch_id}/{len(data_loaders[pattern])} loss: {loss.item():.4f}\n')124 125        scheduler.step()126        patterns = ['val', 'test']127        for pattern_id, is_pattern in enumerate([model_infos['IS_VAL'], model_infos['IS_TEST']]):128            if is_pattern:129                # pred: logits, tensor, nBatch * nClass * W * H130                # target: labels, tensor, nBatch * nClass * W * H131                # output, batch['label']132                collect_result = {'pred': [], 'target': []}133                pattern = patterns[pattern_id]134                model.eval()135                for batch_id, batch in enumerate(data_loaders[pattern]):136                    # Get data137                    img_batch = batch['img'].cuda()138                    label_batch = batch['label'].cuda()139                    img_names = batch['img_name']140                    collect_result['target'].append(label_batch.data.cpu())141 142                    # inference143                    with torch.no_grad():144                        logits, att_branch_output = model(img_batch)145 146                    collect_result['pred'].append(logits.data.cpu())147                    # get segmentation result, when the phase is test.148                    pred_label = torch.argmax(logits, 1)149                    pred_label *= 255150 151                    if pattern == 'test' or batch_id % 5 == 1:152                        batch_size = pred_label.size(0)153                        # k = np.clip(int(0.3 * batch_size), a_min=1, a_max=batch_size)154                        # ids = np.random.choice(range(batch_size), k, replace=False)155                        ids = range(batch_size)156                        for img_id in ids:157                            img = img_batch[img_id].detach().cpu()158                            target = label_batch[img_id].detach().cpu()159                            pred = pred_label[img_id].detach().cpu()160                            img_name = img_names[img_id]161 162                            img = make_numpy_img(163                                inv_normalize_img(img, model_infos['PRIOR_MEAN'], model_infos['PRIOR_STD']))164                            target = make_numpy_img(encode_onehot_to_mask(target)) * 255165                            pred = make_numpy_img(pred)166 167                            vis = np.concatenate([img / 255., target / 255., pred / 255.], axis=0)168                            vis = np.clip(vis, a_min=0, a_max=1)169                            file_name = os.path.join(model_infos['log_path'], pattern, f'Epoch_{epoch_id}_{img_name.split(".")[0]}.png')170                            plt.imsave(file_name, vis)171 172                collect_result['pred'] = torch.cat(collect_result['pred'], dim=0)173                collect_result['target'] = torch.cat(collect_result['target'], dim=0)174                IoU, OA, F1_score = get_metrics('seg', **collect_result)175                logger.write(f'{pattern}: {epoch_id}/{max_epoch} Iou:{IoU[-1]:.4f} OA:{OA[-1]:.4f} F1:{F1_score[-1]:.4f}\n')176        if epoch_id % 20 == 1:177            torch.save({178                'epoch_id': epoch_id,179                'model_state_dict': model.state_dict()180            }, os.path.join(model_infos['model_path'], f'ckpt_{epoch_id}.pt'))181        torch.save({182            'epoch_id': epoch_id,183            'model_state_dict': model.state_dict()184        }, os.path.join(model_infos['model_path'], f'ckpt_latest.pt'))185 186