{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# segmodel_baseline\n\n# 安装segmentation_models_pytorch","metadata":{}},{"cell_type":"code","source":"!cp -r ../input/pytorch-segmentation-models-lib/ ./\n!pip config set global.disable-pip-version-check true\n!pip install -q ./pytorch-segmentation-models-lib/pretrainedmodels-0.7.4/pretrainedmodels-0.7.4\n!pip install -q ./pytorch-segmentation-models-lib/efficientnet_pytorch-0.6.3/efficientnet_pytorch-0.6.3\n!pip install -q ./pytorch-segmentation-models-lib/timm-0.4.12-py3-none-any.whl\n!pip install -q ./pytorch-segmentation-models-lib/segmentation_models_pytorch-0.2.0-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:02:42.092634Z","iopub.execute_input":"2022-08-17T03:02:42.093137Z","iopub.status.idle":"2022-08-17T03:03:34.561336Z","shell.execute_reply.started":"2022-08-17T03:02:42.093044Z","shell.execute_reply":"2022-08-17T03:03:34.559691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 导入库","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('../input/timm-pytorch-image-models/pytorch-image-models-master')","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:34.564364Z","iopub.execute_input":"2022-08-17T03:03:34.564749Z","iopub.status.idle":"2022-08-17T03:03:34.570753Z","shell.execute_reply.started":"2022-08-17T03:03:34.564713Z","shell.execute_reply":"2022-08-17T03:03:34.569518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import gc\nimport os\nimport random\nimport time\nimport warnings\nimport sys\nimport glob\nfrom torch.optim import lr_scheduler\nimport copy\nfrom collections import defaultdict\nfrom torch.cuda import amp\n\nwarnings.simplefilter(\"ignore\")\nimport timm\nimport albumentations as A\n# from albumentations.pytorch import ToTensor\nimport cv2\nfrom matplotlib import pyplot as plt\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image, ImageFilter\nimport segmentation_models_pytorch as smp\nfrom sklearn.model_selection import KFold\nimport torch\nimport torch.backends.cudnn as cudnn\nimport torch.nn as nn\nfrom torch.optim import AdamW, SGD\nfrom torch.nn import functional as F\nimport torch.optim as optim\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom torch.utils.data import DataLoader, Dataset, sampler\nfrom tqdm import tqdm\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, CosineAnnealingWarmRestarts, StepLR, ReduceLROnPlateau\nfrom torch.cuda.amp import autocast, GradScaler\nfrom torchvision import transforms\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold, KFold","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:34.572603Z","iopub.execute_input":"2022-08-17T03:03:34.572965Z","iopub.status.idle":"2022-08-17T03:03:40.627952Z","shell.execute_reply.started":"2022-08-17T03:03:34.572934Z","shell.execute_reply":"2022-08-17T03:03:40.626484Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 固定随机种子","metadata":{}},{"cell_type":"code","source":"def set_seed(seed=2 ** 3):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:40.631135Z","iopub.execute_input":"2022-08-17T03:03:40.632166Z","iopub.status.idle":"2022-08-17T03:03:40.641012Z","shell.execute_reply.started":"2022-08-17T03:03:40.632123Z","shell.execute_reply":"2022-08-17T03:03:40.639676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CFG:\n    seed          = 101\n    debug         = True # \n    exp_name      = 'Baselinev1'\n    comment       = 'unet-efficientnet_b1-512-512-aug2-split2'\n    model_name    = 'Unet'\n    backbone      = 'efficientnet-b1'\n    train_bs      = 24\n    valid_bs      = train_bs*2\n    img_size      = [512, 512]\n    epochs        = 5\n    lr            = 2e-3\n    scheduler     = 'CosineAnnealingLR'\n    min_lr        = 1e-6\n    T_max         = 10\n    T_0           = 25\n    warmup_epochs = 0\n    wd            = 1e-6\n    n_accumulate  = max(1, 32//train_bs)\n    n_fold        = 5\n    num_classes   = 3\n    device        = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:40.642386Z","iopub.execute_input":"2022-08-17T03:03:40.643424Z","iopub.status.idle":"2022-08-17T03:03:40.655142Z","shell.execute_reply.started":"2022-08-17T03:03:40.643379Z","shell.execute_reply":"2022-08-17T03:03:40.653875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 数据准备\ntrain_path = glob.glob('../input/boolart-cityscapes/data/train/*/*')\ndf_train = pd.DataFrame({\n    \"image_path\":train_path\n})\ndf_train['mask_path'] = df_train['image_path'].str.replace(\"data\",'label')\ndf_train['mask_path'] = df_train['mask_path'].apply(lambda x:x.split('.png')[0][:-11] + 'gtFine_labelTrainIds.png')\nFold = KFold(n_splits=5)\nfor n, (train_index, val_index) in enumerate(Fold.split(df_train)):\n    df_train.loc[val_index, 'fold'] = int(n)\ndf_train['fold'] = df_train['fold'].astype(int)","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:40.657395Z","iopub.execute_input":"2022-08-17T03:03:40.657929Z","iopub.status.idle":"2022-08-17T03:03:41.051394Z","shell.execute_reply.started":"2022-08-17T03:03:40.657884Z","shell.execute_reply":"2022-08-17T03:03:41.050221Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 数据加载","metadata":{}},{"cell_type":"code","source":"data_transforms = {\n    \"train\": A.Compose([\n        A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n        A.HorizontalFlip(p=0.5),\n#         A.VerticalFlip(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.05, rotate_limit=10, p=0.5),\n        A.OneOf([\n            A.GridDistortion(num_steps=5, distort_limit=0.05, p=1.0),\n#             A.OpticalDistortion(distort_limit=0.05, shift_limit=0.05, p=1.0),\n            A.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=1.0)\n        ], p=0.25),\n        A.CoarseDropout(max_holes=8, max_height=CFG.img_size[0]//20, max_width=CFG.img_size[1]//20,\n                         min_holes=5, fill_value=0, mask_fill_value=0, p=0.5),\n        ], p=1.0),\n    \n    \"valid\": A.Compose([\n        A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n        ], p=1.0)\n}\n\nclass BuildDataset(torch.utils.data.Dataset):\n    def __init__(self, df, label=True, transforms=None):\n        self.df         = df\n        self.label      = label\n        self.img_paths  = df['image_path'].tolist()\n        self.msk_paths  = df['mask_path'].tolist()\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_path  = self.img_paths[index]\n        img = cv2.imread(img_path)\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        msk_path = self.msk_paths[index]\n        msk = np.array(Image.open(msk_path))\n        if self.transforms:\n            data = self.transforms(image=img, mask=msk)\n            img  = data['image']\n            msk  = data['mask']\n        img = np.transpose(img, (2, 0, 1))\n        msk = torch.tensor(msk)\n        return torch.tensor(img), msk","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:41.053575Z","iopub.execute_input":"2022-08-17T03:03:41.054129Z","iopub.status.idle":"2022-08-17T03:03:41.071351Z","shell.execute_reply.started":"2022-08-17T03:03:41.054079Z","shell.execute_reply":"2022-08-17T03:03:41.070234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_loaders(fold, debug=False):\n    train_df = df_train.query(\"fold!=@fold\").reset_index(drop=True)\n    valid_df = df_train.query(\"fold==@fold\").reset_index(drop=True)\n    if debug:\n        train_df = train_df.head(32*5).reset_index(drop=True)\n        valid_df = valid_df.head(32*3).reset_index(drop=True)\n    train_dataset = BuildDataset(train_df, transforms=data_transforms['train'])\n    valid_dataset = BuildDataset(valid_df, transforms=data_transforms['valid'])\n\n    train_loader = DataLoader(train_dataset, batch_size=CFG.train_bs if not debug else 20, \n                              num_workers=4, shuffle=True, pin_memory=True, drop_last=False)\n    valid_loader = DataLoader(valid_dataset, batch_size=CFG.valid_bs if not debug else 20, \n                              num_workers=4, shuffle=False, pin_memory=True)\n    \n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:41.072877Z","iopub.execute_input":"2022-08-17T03:03:41.073448Z","iopub.status.idle":"2022-08-17T03:03:41.089359Z","shell.execute_reply.started":"2022-08-17T03:03:41.073388Z","shell.execute_reply":"2022-08-17T03:03:41.088265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"增强后的数据可视化","metadata":{}},{"cell_type":"code","source":"train_loader, valid_loader = prepare_loaders(0)\nimgs,masks = next(iter(train_loader))\n\nplt.figure(figsize=(16,16))\nfor i,(img,mask) in enumerate(zip(imgs,masks)):\n    img = img.permute(1,2,0).numpy().astype(np.uint8)\n    plt.subplot(8,8,i+1)\n    plt.imshow(img,vmin=0,vmax=255)\n    plt.imshow(mask.squeeze().numpy(), alpha=0.2)\n    plt.axis('off')\n    plt.subplots_adjust(wspace=None, hspace=None)\n    \ndel train_loader,valid_loader,imgs,masks","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:41.092226Z","iopub.execute_input":"2022-08-17T03:03:41.093099Z","iopub.status.idle":"2022-08-17T03:03:59.170199Z","shell.execute_reply.started":"2022-08-17T03:03:41.093051Z","shell.execute_reply":"2022-08-17T03:03:59.169311Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 模型\n\nsmp下定义的不同分割模型可以通过smp??进行查看，使用方式和下方相同。","metadata":{}},{"cell_type":"code","source":"import segmentation_models_pytorch as smp\n\ndef build_model():\n    model = smp.create_model(arch='Unet',                    # 可以使用不同的分割模型如UNet++，DeepLabV3\n                         encoder_name='resnet34',        # 可以使用不同的骨架模型如se_resnet50，efficientnet-b1\n                         encoder_weights=\"imagenet\",     # 是否对编码模型加载预训练权重\n                         in_channels=3,                  # 模型的输入通道\n                         classes=8,                      # 模型的输出类别\n                         activation=None,\n                        )\n    \n    model.to(CFG.device)\n    return model\n\ndef load_model(path):\n    model = build_model()\n    model.load_state_dict(torch.load(path))\n    model.eval()\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:59.173712Z","iopub.execute_input":"2022-08-17T03:03:59.174289Z","iopub.status.idle":"2022-08-17T03:03:59.181497Z","shell.execute_reply.started":"2022-08-17T03:03:59.174250Z","shell.execute_reply":"2022-08-17T03:03:59.180222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"maskloss1 = smp.losses.SoftCrossEntropyLoss(smooth_factor=0.05, ignore_index=8)\nmaskloss2 = smp.losses.DiceLoss(mode='multiclass', smooth=0.05, ignore_index=8)","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:59.182716Z","iopub.execute_input":"2022-08-17T03:03:59.183103Z","iopub.status.idle":"2022-08-17T03:03:59.213800Z","shell.execute_reply.started":"2022-08-17T03:03:59.183073Z","shell.execute_reply":"2022-08-17T03:03:59.212941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 训练验证","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, optimizer, scheduler, dataloader, device, epoch):\n    model.train()\n    scaler = amp.GradScaler()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    confuse_mat = np.zeros([8, 8])\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Train ')\n    for step, (images, masks) in pbar:         \n        images = images.to(device, dtype=torch.float)\n        masks  = masks.to(device)\n        \n        batch_size = images.size(0)\n        \n        with amp.autocast(enabled=True):\n            y_pred = model(images)\n            loss = 0.5 * maskloss1(y_pred, masks.long()) + 0.5 * maskloss2(y_pred, masks.long())\n            loss   = loss / CFG.n_accumulate\n            \n        scaler.scale(loss).backward()\n        if (step + 1) % CFG.n_accumulate == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            # zero the parameter gradients\n            optimizer.zero_grad()\n            if scheduler is not None:\n                scheduler.step()\n                \n        running_loss += (loss.item() * batch_size)\n        dataset_size += batch_size\n        \n        epoch_loss = running_loss / dataset_size\n        \n        mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n        current_lr = optimizer.param_groups[0]['lr']\n        pbar.set_postfix(train_loss=f'{epoch_loss:0.4f}',\n                        lr=f'{current_lr:0.5f}',\n                        gpu_mem=f'{mem:0.2f} GB')\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return epoch_loss","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:59.215078Z","iopub.execute_input":"2022-08-17T03:03:59.215961Z","iopub.status.idle":"2022-08-17T03:03:59.229367Z","shell.execute_reply.started":"2022-08-17T03:03:59.215927Z","shell.execute_reply":"2022-08-17T03:03:59.227933Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef valid_one_epoch(model, dataloader, device, epoch):\n    model.eval()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    \n    val_scores = []\n    \n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Valid ')\n    for step, (images, masks) in pbar:        \n        images  = images.to(device, dtype=torch.float)\n        masks   = masks.to(device)\n        batch_size = images.size(0)\n        y_pred  = model(images)\n        loss    =0.5 * maskloss1(y_pred, masks.long()) + 0.5 * maskloss2(y_pred, masks.long())\n        \n        running_loss += (loss.item() * batch_size)\n        dataset_size += batch_size\n        epoch_loss = running_loss / dataset_size\n        y_pred = nn.Sigmoid()(y_pred)\n        \n        mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n        current_lr = optimizer.param_groups[0]['lr']\n        pbar.set_postfix(valid_loss=f'{epoch_loss:0.4f}',\n                        lr=f'{current_lr:0.5f}',\n                        gpu_memory=f'{mem:0.2f} GB')\n    torch.cuda.empty_cache()\n    gc.collect()\n    return epoch_loss","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:59.230991Z","iopub.execute_input":"2022-08-17T03:03:59.231645Z","iopub.status.idle":"2022-08-17T03:03:59.243988Z","shell.execute_reply.started":"2022-08-17T03:03:59.231613Z","shell.execute_reply":"2022-08-17T03:03:59.243008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_training(model, optimizer, scheduler, device, num_epochs): \n    start = time.time()\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_loss      = np.inf\n    best_epoch     = -1\n    history = defaultdict(list)\n    \n    for epoch in range(1, num_epochs + 1): \n        gc.collect()\n        print(f'Epoch {epoch}/{num_epochs}', end='')\n        train_loss = train_one_epoch(model, optimizer, scheduler, \n                                           dataloader=train_loader, \n                                           device=CFG.device, epoch=epoch)\n        \n        val_loss = valid_one_epoch(model, valid_loader, \n                                                 device=CFG.device, \n                                                 epoch=epoch)\n        # deep copy the model\n        if val_loss <= best_loss:\n            print(f\"Valid loss Improved ({best_loss:0.4f} ---> {val_loss:0.4f})\")\n            best_loss    = val_loss\n            best_epoch   = epoch\n            best_model_wts = copy.deepcopy(model.state_dict())\n            PATH = f\"best_epoch-{fold:02d}.pth\"\n            torch.save(model.state_dict(), PATH)\n            print(f\"Model Saved\")\n            \n        last_model_wts = copy.deepcopy(model.state_dict())\n        PATH = f\"last_epoch-{fold:02d}.pth\"\n        torch.save(model.state_dict(), PATH)\n    \n    return model, history","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:59.245359Z","iopub.execute_input":"2022-08-17T03:03:59.246105Z","iopub.status.idle":"2022-08-17T03:03:59.260700Z","shell.execute_reply.started":"2022-08-17T03:03:59.246073Z","shell.execute_reply":"2022-08-17T03:03:59.259501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 学习率策略","metadata":{"execution":{"iopub.status.busy":"2022-08-16T02:24:33.582589Z","iopub.execute_input":"2022-08-16T02:24:33.582959Z","iopub.status.idle":"2022-08-16T02:24:33.587542Z","shell.execute_reply.started":"2022-08-16T02:24:33.582927Z","shell.execute_reply":"2022-08-16T02:24:33.586192Z"}}},{"cell_type":"code","source":"def fetch_scheduler(optimizer):\n    if CFG.scheduler == 'CosineAnnealingLR':\n        scheduler = lr_scheduler.CosineAnnealingLR(optimizer,T_max=CFG.T_max, \n                                                   eta_min=CFG.min_lr)\n    elif CFG.scheduler == 'CosineAnnealingWarmRestarts':\n        scheduler = lr_scheduler.CosineAnnealingWarmRestarts(optimizer,T_0=CFG.T_0, \n                                                             eta_min=CFG.min_lr)\n    elif CFG.scheduler == 'ReduceLROnPlateau':\n        scheduler = lr_scheduler.ReduceLROnPlateau(optimizer,\n                                                   mode='min',\n                                                   factor=0.1,\n                                                   patience=7,\n                                                   threshold=0.0001,\n                                                   min_lr=CFG.min_lr,)\n    elif CFG.scheduer == 'ExponentialLR':\n        scheduler = lr_scheduler.ExponentialLR(optimizer, gamma=0.85)\n    elif CFG.scheduler == None:\n        return None\n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:59.262110Z","iopub.execute_input":"2022-08-17T03:03:59.262680Z","iopub.status.idle":"2022-08-17T03:03:59.275301Z","shell.execute_reply.started":"2022-08-17T03:03:59.262646Z","shell.execute_reply":"2022-08-17T03:03:59.274478Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for fold in range(1):\n    print(f'#'*15)\n    print(f'### Fold: {fold}')\n    print(f'#'*15)\n    train_loader, valid_loader = prepare_loaders(fold=fold, debug=CFG.debug)\n    model     = build_model()\n    optimizer = optim.Adam(model.parameters(), lr=CFG.lr, weight_decay=CFG.wd)\n    scheduler = fetch_scheduler(optimizer)\n    model, history = run_training(model, optimizer, scheduler,\n                                  device=CFG.device,\n                                  num_epochs=CFG.epochs)","metadata":{"execution":{"iopub.status.busy":"2022-08-17T03:03:59.276375Z","iopub.execute_input":"2022-08-17T03:03:59.277145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 测试数据准备","metadata":{}},{"cell_type":"code","source":"test_path = glob.glob('../input/boolart-cityscapes/data/test/*/*')\ntest_df = pd.DataFrame({\n    \"image_path\":test_path\n})","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BuildDataset(torch.utils.data.Dataset):\n    def __init__(self, df, label=True, transforms=None):\n        self.df         = df\n        self.label      = label\n        self.img_paths  = df['image_path'].tolist()\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_path  = self.img_paths[index]\n        img = cv2.imread(img_path)\n        h, w = img.shape[:2]\n        img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n        if self.transforms:\n            data = self.transforms(image=img)\n            img  = data['image']\n        img = np.transpose(img, (2, 0, 1))\n        return torch.tensor(img), os.path.basename(img_path),h,w","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_to_string(runs):\n    return ' '.join(str(x) for x in runs)\n# rle编码\ndef rle_encode(mask):\n    pixels = mask.T.flatten()\n    use_padding = False\n    if pixels[0] or pixels[-1]:\n        use_padding = True\n        pixel_padded = np.zeros([len(pixels) + 2], dtype=pixels.dtype)\n        pixel_padded[1:-1] = pixels\n        pixels = pixel_padded\n    rle = np.where(pixels[1:] != pixels[:-1])[0] + 2\n    if use_padding:\n        rle = rle - 1\n    rle[1::2] = rle[1::2] - rle[:-1:2]\n    return rle_to_string(rle)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 推理","metadata":{}},{"cell_type":"code","source":"@torch.no_grad()\ndef inference(model, dataloader, device):\n    model.eval()\n    pred_strings = []\n    pred_ids = []\n    pred_classes = []\n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Valid ')\n    for step, (images, id_, height,width) in pbar:        \n        images  = images.to(device, dtype=torch.float)\n        batch_size = images.size(0)\n        y_pred  = model(images).squeeze()\n        y_pred = torch.nn.Sigmoid()(y_pred)\n        y_pred = (y_pred.permute((1, 2, 0))>0.45).to(torch.uint8).cpu().detach().numpy()\n        y_pred = cv2.resize(y_pred, \n                        dsize=(int(width),int(height)),\n                        interpolation=cv2.INTER_NEAREST) \n        rle = [None]*8\n        for midx in range(8):\n            rle[midx] = rle_encode(y_pred[...,midx])\n            file_name = id_[0][:-4]+ str(midx)+'.png'\n            pred_ids.extend([file_name])\n        pred_strings.extend(rle)\n        pred_classes.extend(['other', 'road', 'person', 'rider','car', 'truck','bus','bicycle'])\n    return pred_strings, pred_ids, pred_classes","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = BuildDataset(test_df, transforms=data_transforms['valid'])\ntest_loader = DataLoader(test_dataset, batch_size=1, \n                          num_workers=1, shuffle=False, pin_memory=True, drop_last=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(torch.load('./last_epoch-00.pth'))\npred_strings, pred_ids, pred_classes = inference(model, test_loader, device=CFG.device,)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pred_ids = sum(pred_ids,[]) # 将列表展开\ndf = pd.DataFrame({\n    \"id\":pred_ids,\n    \"predict\":pred_strings\n})\ndf.to_csv('./submission.csv',index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df.head(10)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}