{"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":"# ****Directory settings****\n\nกำหนดตัวแปรต่างๆ เพื่อใช้ในการจัดเก็บ output และเรียกใช้โมเดลในการทำนายภาพในโปรเจกต์ที่เกี่ยวกับการจำแนกโรคใบต้นมันสำปะหลัง (Cassava Leaf Disease Classification)\n\nตัวแปร OUTPUT_DIR ใช้สำหรับเก็บผลลัพธ์ของโมเดล โดยกำหนดค่าเป็น ./ หรือค่า path ของโฟลเดอร์ปัจจุบัน (current directory) ซึ่งถ้าไม่มีโฟลเดอร์นี้อยู่แล้วก็จะถูกสร้างขึ้นมา ตัวแปร MODEL_DIR ใช้สำหรับเก็บโมเดลที่ได้เทรนเอาไว้ โดยกำหนดค่า path เป็น \"../input/cassava-model/\" ซึ่งอาจจะเป็น path ที่เก็บโมเดลที่เทรนไว้บน cloud storage เช่น Kaggle หรือ Google Drive หรืออื่นๆ ตัวแปร TRAIN_PATH และ TEST_PATH ใช้สำหรับเก็บ path ของข้อมูลภาพที่ใช้ในการเทรนและทดสอบโมเดล โดยกำหนดค่า path เป็น \"../input/cassava-leaf-disease-classification/train_images\" และ \"../input/cassava-leaf-disease-classification/test_images\" ตามลำดับ","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Directory settings\n# ====================================================\nimport os\n\nOUTPUT_DIR = \"./\"\nMODEL_DIR = \"../input/cassava-model/\"\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)\n\nTRAIN_PATH = \"../input/cassava-leaf-disease-classification/train_images\"\nTEST_PATH = \"../input/cassava-leaf-disease-classification/test_images\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-04-12T07:17:37.372546Z","iopub.execute_input":"2023-04-12T07:17:37.372893Z","iopub.status.idle":"2023-04-12T07:17:37.411142Z","shell.execute_reply.started":"2023-04-12T07:17:37.372858Z","shell.execute_reply":"2023-04-12T07:17:37.409263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ****CFG****\n\nกำหนดค่า (configuration settings) สำหรับโมเดลการเรียนรู้เชิงลึก (deep learning model) ในการแก้ปัญหาการจำแนกภาพโรคใบต้นมันสำปะหลัง (cassava leaf disease classification) โดยใช้ไลบรารี PyTorch และ Hugging Face Transformers โค้ดตัวแปร CFG เก็บค่าต่างๆ ที่เกี่ยวข้องกับการเทรนโมเดล และวิธีการทำคลาสระหว่างการทดสอบ (inference)","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# CFG\n# ====================================================\nclass CFG:\n    debug = False\n    num_workers = 4\n    models = [\n        \"tf_efficientnet_b4_ns\",\n        \"vit_base_patch16_384\",\n        \"seresnext50_32x4d\",\n    ]\n    size = {\n        \"tf_efficientnet_b3_ns\": 512,\n        \"tf_efficientnet_b4_ns\": 512,\n        \"vit_base_patch16_384\": 384,\n        \"deit_base_patch16_384\": 384,\n        \"seresnext50_32x4d\": 512,\n    }\n    batch_size = 64\n    seed = 7097\n    target_size = 5\n    target_col = \"label\"\n    n_fold = 5\n    trn_fold = {  # [0, 1, 2, 3, 4, 5, 6, 7, 8, 9]\n        \"tf_efficientnet_b3_ns\": {\n            \"best\": [0, 1, 2, 3, 4],\n            \"final\": [],\n        },\n        \"tf_efficientnet_b4_ns\": {\n            \"best\": [0, 1, 2, 3, 4],\n            \"final\": [],\n        },\n        \"vit_base_patch16_384\": {\"best\": [0, 1, 2, 3, 4], \"final\": []},\n        \"deit_base_patch16_384\": {\"best\": [0, 1, 2, 3, 4], \"final\": []},\n        \"seresnext50_32x4d\": {\"best\": [5, 6, 7, 8, 9], \"final\": []},\n    }\n    data_parallel = {\n        \"tf_efficientnet_b3_ns\": False,\n        \"tf_efficientnet_b4_ns\": True,  # True,\n        \"vit_base_patch16_384\": False,\n        \"deit_base_patch16_384\": False,\n        \"seresnext50_32x4d\": False,\n    }\n    transform = {\n        \"tf_efficientnet_b4_ns\": \"rotate\",\n        \"vit_base_patch16_384\": \"rotate\",\n        \"seresnext50_32x4d\": \"rotate\",\n    }\n    weight = {\n        \"tf_efficientnet_b4_ns\": 1,\n        \"vit_base_patch16_384\": 1,\n        \"seresnext50_32x4d\": 1,\n    }\n    tta = 10  # 1: no TTA, >1: TTA\n    no_tta_weight = tta - 1\n    train = False\n    inference = True","metadata":{"execution":{"iopub.status.busy":"2023-04-12T07:17:37.414395Z","iopub.execute_input":"2023-04-12T07:17:37.417922Z","iopub.status.idle":"2023-04-12T07:17:37.432735Z","shell.execute_reply.started":"2023-04-12T07:17:37.417885Z","shell.execute_reply":"2023-04-12T07:17:37.431823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"คำนวณ weight_sum ที่จะถูกนำไปใช้ในการคำนวณ loss ของโมเดล โดยจะคูณด้วยตัวแปร tta_weight_sum ซึ่งแสดงถึงจำนวนของการเพิ่มเติมภาพ (Test Time Augmentation) และนำไปบวกกันทั้งหมด โดยจะนับ weight ของโมเดลทั้งหมดโดยรวมกันด้วยการคูณด้วย tta_weight_sum เพื่อให้ได้ weight_sum ที่ถูกต้องตามจำนวนการเพิ่มเติมภาพทั้งหมดของ TTA ที่ต้องการใช้งาน","metadata":{}},{"cell_type":"code","source":"tta_weight_sum = CFG.no_tta_weight + (CFG.tta - 1)\nweight_sum = sum([CFG.weight[model] for model in CFG.models]) * tta_weight_sum","metadata":{"execution":{"iopub.status.busy":"2023-04-12T07:17:37.437571Z","iopub.execute_input":"2023-04-12T07:17:37.440296Z","iopub.status.idle":"2023-04-12T07:17:37.446986Z","shell.execute_reply.started":"2023-04-12T07:17:37.440256Z","shell.execute_reply":"2023-04-12T07:17:37.446063Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ****Library****\n\nการ import libraries ต่างๆที่จำเป็นในการเขียนโปรแกรม เช่น PyTorch, NumPy, Pandas, Scikit-learn, Albumentations, และอื่นๆ และกำหนด device เป็น cuda หรือ cpu ตามที่มี GPU ในเครื่องหรือไม่ โดยใช้คำสั่ง torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\") ในการเลือก device โดยใช้ GPU หากมี แต่ถ้าไม่มีก็ใช้ CPU แทน","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Library\n# ====================================================\nimport sys\n\nsys.path.append(\"../input/pytorch-image-models/pytorch-image-models-master\")\n\nimport math\nimport os\nimport random\nimport shutil\nimport time\nimport warnings\nfrom collections import Counter, defaultdict\nfrom contextlib import contextmanager\nfrom functools import partial\nfrom pathlib import Path\n\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport scipy as sp\nimport timm\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.models as models\nfrom albumentations import (\n    CenterCrop,\n    CoarseDropout,\n    Compose,\n    Cutout,\n    HorizontalFlip,\n    HueSaturationValue,\n    IAAAdditiveGaussianNoise,\n    ImageOnlyTransform,\n    Normalize,\n    OneOf,\n    RandomBrightness,\n    RandomBrightnessContrast,\n    RandomContrast,\n    RandomCrop,\n    RandomResizedCrop,\n    Resize,\n    Rotate,\n    ShiftScaleRotate,\n    Transpose,\n    VerticalFlip,\n)\nfrom albumentations.pytorch import ToTensorV2\nfrom PIL import Image\nfrom sklearn import preprocessing\nfrom sklearn.metrics import accuracy_score\nfrom sklearn.model_selection import StratifiedKFold\nfrom torch.nn.parameter import Parameter\nfrom torch.optim import SGD, Adam\nfrom torch.optim.lr_scheduler import CosineAnnealingLR, CosineAnnealingWarmRestarts, ReduceLROnPlateau\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm.auto import tqdm\n\nwarnings.filterwarnings(\"ignore\")\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"execution":{"iopub.status.busy":"2023-04-12T07:17:37.452481Z","iopub.execute_input":"2023-04-12T07:17:37.455368Z","iopub.status.idle":"2023-04-12T07:17:43.757885Z","shell.execute_reply.started":"2023-04-12T07:17:37.455332Z","shell.execute_reply":"2023-04-12T07:17:43.756529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ****Utils****\n\nฟังก์ชันแรกคือ get_score(y_true, y_pred) ที่ใช้สำหรับคำนวณค่าความแม่นยำ (accuracy score) ของผลการทำนาย y_pred กับผลลัพธ์ที่ถูกต้อง y_true ในชุดข้อมูลเทียบกับเฉลี่ยที่ได้จากการทำนายของโมเดล\n\nฟังก์ชันต่อมาคือ timer(name) ซึ่งเป็นฟังก์ชันเพื่อวัดเวลาการทำงานของโปรแกรมในแต่ละส่วน โดยจะแสดงเวลาเริ่มต้นและเวลาสิ้นสุดของการทำงานในแต่ละส่วน โดย name เป็นชื่อของส่วนงาน\n\nฟังก์ชันต่อมาคือ init_logger(log_file) ซึ่งเป็นฟังก์ชันสำหรับกำหนดการทำงานของ logger ซึ่งจะใช้สำหรับบันทึกข้อมูลการทำงานของโปรแกรมในไฟล์ log_file\n\nฟังก์ชันสุดท้ายคือ seed_torch(seed) ซึ่งจะใช้สำหรับกำหนดค่า random seed ให้กับตัวแปรในโปรแกรม ซึ่งจะช่วยในการทำให้ผลการทำงานของโปรแกรมมีความสม่ำเสมอขึ้นในแต่ละครั้งที่รันโปรแกรม ทำให้ผลการทดสอบและการวัดประสิทธิภาพของโมเดลนั้นมีความเท่าเทียมกันในแต่ละครั้ง","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Utils\n# ====================================================\ndef get_score(y_true, y_pred):\n    return accuracy_score(y_true, y_pred)\n\n\n@contextmanager\ndef timer(name):\n    t0 = time.time()\n    LOGGER.info(f\"[{name}] start\")\n    yield\n    LOGGER.info(f\"[{name}] done in {time.time() - t0:.0f} s.\")\n\n\ndef init_logger(log_file=OUTPUT_DIR + \"inference.log\"):\n    from logging import INFO, FileHandler, Formatter, StreamHandler, getLogger\n\n    logger = getLogger(__name__)\n    logger.setLevel(INFO)\n    handler1 = StreamHandler()\n    handler1.setFormatter(Formatter(\"%(message)s\"))\n    handler2 = FileHandler(filename=log_file)\n    handler2.setFormatter(Formatter(\"%(message)s\"))\n    logger.addHandler(handler1)\n    logger.addHandler(handler2)\n    return logger\n\n\nLOGGER = init_logger()\n\n\ndef seed_torch(seed=42):\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\n    torch.backends.cudnn.benchmark = True\n\n\nseed_torch(seed=CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2023-04-12T07:17:43.759896Z","iopub.execute_input":"2023-04-12T07:17:43.760589Z","iopub.status.idle":"2023-04-12T07:17:43.776937Z","shell.execute_reply.started":"2023-04-12T07:17:43.760543Z","shell.execute_reply":"2023-04-12T07:17:43.775767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ****Data Loading****\n\nอ่านไฟล์ sample_submission.csv และเก็บข้อมูลที่ได้อ่านมาในตัวแปร test และแสดงผลหัวของตารางด้วยคำสั่ง test.head() โดยในไฟล์ sample_submission.csv จะมีข้อมูลเกี่ยวกับรหัสของภาพ (image_id) และคำตอบ (label)","metadata":{}},{"cell_type":"code","source":"test = pd.read_csv(\"../input/cassava-leaf-disease-classification/sample_submission.csv\")\ntest.head()","metadata":{"execution":{"iopub.status.busy":"2023-04-12T07:17:43.780147Z","iopub.execute_input":"2023-04-12T07:17:43.780572Z","iopub.status.idle":"2023-04-12T07:17:43.812670Z","shell.execute_reply.started":"2023-04-12T07:17:43.780528Z","shell.execute_reply":"2023-04-12T07:17:43.811633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ****Dataset****\n\nสร้าง Dataset สำหรับการทดสอบ (Test Dataset) โดยกำหนดว่า Dataset จะมีข้อมูลเป็นแบบไฟล์ภาพ และกำหนด path ไปยังไฟล์ภาพของแต่ละรูปภาพ จากนั้นโหลดภาพขึ้นมาใช้ cv2 ในการอ่านไฟล์ภาพ และเมื่อโหลดภาพเสร็จสิ้น จะทำการทำ augmentation ด้วย albumentations ถ้าหากมีการกำหนด transform มาให้ และส่งภาพออกไปในรูปแบบของ tensor สำหรับการนำไปใช้งานในโมเดล CNN","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Dataset\n# ====================================================\nclass TestDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df[\"image_id\"].values\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_name = self.file_names[idx]\n        file_path = f\"{TEST_PATH}/{file_name}\"\n        image = cv2.imread(file_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented[\"image\"]\n        return image","metadata":{"execution":{"iopub.status.busy":"2023-04-12T07:17:43.816355Z","iopub.execute_input":"2023-04-12T07:17:43.816671Z","iopub.status.idle":"2023-04-12T07:17:43.827172Z","shell.execute_reply.started":"2023-04-12T07:17:43.816641Z","shell.execute_reply":"2023-04-12T07:17:43.825973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ****Transforms****\n\nสร้างฟังก์ชัน get_transforms ซึ่งรับ parameter 2 ตัว คือ data และ size โดยฟังก์ชันนี้จะสร้าง Transform ต่างๆ สำหรับการประมวลผลภาพ โดยเรียกใช้ Compose ซึ่งเป็นฟังก์ชันที่ช่วยในการรวม Transform หลายๆ อันเข้าด้วยกัน โดยสามารถกำหนด Transform แต่ละอันได้ในรูปแบบของ list และใช้ ToTensorV2() เพื่อแปลงภาพเป็น tensor ในรูปแบบของ PyTorch โดยฟังก์ชัน get_transforms นี้จะส่งออก Transform ตาม data ที่กำหนด โดยเมื่อกำหนด data เป็น \"train\" จะส่งออก Transform สำหรับการเทรนโมเดล ส่วนเมื่อกำหนด data เป็น \"valid\" จะส่งออก Transform สำหรับการทดสอบโมเดล และเมื่อกำหนด data เป็น \"simple\" หรือ \"rotate\" จะส่งออก Transform สำหรับการใช้งานง่ายหรือการหมุนภาพตามลำดับ","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Transforms\n# ====================================================\ndef get_transforms(*, data, size):\n\n    if data == \"train\":\n        return Compose(\n            [\n                RandomResizedCrop(size, size),\n                Transpose(p=0.5),\n                HorizontalFlip(p=0.5),\n                VerticalFlip(p=0.5),\n                ShiftScaleRotate(p=0.5),\n                HueSaturationValue(hue_shift_limit=0.2, sat_shift_limit=0.2, val_shift_limit=0.2, p=0.5),\n                RandomBrightnessContrast(brightness_limit=(-0.1, 0.1), contrast_limit=(-0.1, 0.1), p=0.5),\n                CoarseDropout(p=0.5),\n                Cutout(p=0.5),\n                Normalize(\n                    mean=[0.485, 0.456, 0.406],\n                    std=[0.229, 0.224, 0.225],\n                ),\n                ToTensorV2(),\n            ]\n        )\n\n    if data == \"valid\":\n        return Compose(\n            [\n                Resize(size, size),\n                CenterCrop(size, size),\n                Normalize(\n                    mean=[0.485, 0.456, 0.406],\n                    std=[0.229, 0.224, 0.225],\n                ),\n                ToTensorV2(),\n            ]\n        )\n\n    if data == \"simple\":\n        return Compose(\n            [\n                RandomResizedCrop(size, size),\n                \n                Normalize(\n                    mean=[0.485, 0.456, 0.406],\n                    std=[0.229, 0.224, 0.225],\n                ),\n                ToTensorV2(),\n            ]\n        )\n\n    if data == \"rotate\":\n        return Compose(\n            [\n                RandomResizedCrop(size, size),\n                Transpose(p=0.5),\n                HorizontalFlip(p=0.5),\n                VerticalFlip(p=0.5),\n                \n                Normalize(\n                    mean=[0.485, 0.456, 0.406],\n                    std=[0.229, 0.224, 0.225],\n                ),\n                ToTensorV2(),\n            ]\n        )","metadata":{"execution":{"iopub.status.busy":"2023-04-12T07:17:43.829213Z","iopub.execute_input":"2023-04-12T07:17:43.829677Z","iopub.status.idle":"2023-04-12T07:17:43.844099Z","shell.execute_reply.started":"2023-04-12T07:17:43.829599Z","shell.execute_reply":"2023-04-12T07:17:43.843010Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ****MODEL****\n\nสร้างโมเดลประสาทเทียม (neural network model) resnext50_32x4d สำหรับการจำแนกภาพในงาน Cassava Leaf Disease Classification โดยใช้ PyTorch","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# MODEL\n# ====================================================\nclass CassvaImgClassifier(nn.Module):\n    def __init__(self, model_name=\"resnext50_32x4d\", pretrained=False):\n        super().__init__()\n\n        if model_name == \"deit_base_patch16_384\":\n            self.model = torch.hub.load(\"../input/fair-deit\", model_name, pretrained=pretrained, source=\"local\")\n            n_features = self.model.head.in_features\n            self.model.head = nn.Linear(n_features, CFG.target_size)\n\n        else:\n            self.model = timm.create_model(model_name, pretrained=pretrained)\n\n            if \"resnext50_32x4d\" in model_name:\n                n_features = self.model.fc.in_features\n                self.model.fc = nn.Linear(n_features, CFG.target_size)\n\n            elif model_name.startswith(\"tf_efficientnet\"):\n                n_features = self.model.classifier.in_features\n                self.model.classifier = nn.Linear(n_features, CFG.target_size)\n\n            elif model_name.startswith(\"vit_\"):\n                n_features = self.model.head.in_features\n                self.model.head = nn.Linear(n_features, CFG.target_size)\n\n    def forward(self, x):\n        x = self.model(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2023-04-12T07:17:43.845577Z","iopub.execute_input":"2023-04-12T07:17:43.846570Z","iopub.status.idle":"2023-04-12T07:17:43.860043Z","shell.execute_reply.started":"2023-04-12T07:17:43.846523Z","shell.execute_reply":"2023-04-12T07:17:43.858829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ****Helper functions****\n\nฟังก์ชัน inference รับโมเดล (model) และสถานะของโมเดล (states) ที่ถูกบันทึกไว้ก่อนหน้า (ใช้ในการทำนายหลังจากเทรนโมเดลเสร็จแล้ว) รวมถึง DataLoader สำหรับชุดข้อมูลที่จะนำมาทำนาย (test_loader) และอุปกรณ์ที่ใช้งาน (device) ซึ่งสามารถเป็น CPU หรือ GPU ได้\n\nโค้ดจะทำการส่งรูปภาพต่อไปยังโมเดลเพื่อทำนายผลลัพธ์ โดยจะนำ states มาใช้โมเดล เนื่องจาก states เก็บค่า weights ของโมเดลที่ได้จากการทำฝึก (training) โมเดลก่อนหน้านี้ จากนั้นจะใช้โมเดลเหล่านั้นมาทำนายผลลัพธ์บนรูปภาพ และจะทำการหาค่าเฉลี่ยของผลลัพธ์ที่ได้จากทุกๆ states แล้วส่งออกมาในรูปของ numpy array ของค่าความน่าจะเป็น (probs)","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# Helper functions\n# ====================================================\ndef inference(model, states, test_loader, device, data_parallel):\n    model.to(device)\n\n    # Use multi GPU\n    if device == torch.device(\"cuda\") and data_parallel:\n        model = torch.nn.DataParallel(model)  # make parallel\n\n    tk0 = tqdm(enumerate(test_loader), total=len(test_loader))\n    probs = []\n    for i, (images) in tk0:\n        images = images.to(device)\n        avg_preds = []\n        for state in states:\n            model.load_state_dict(state[\"model\"])\n            model.eval()\n            with torch.no_grad():\n                y_preds = model(images)\n            avg_preds.append(y_preds.softmax(1).to(\"cpu\").numpy())\n        avg_preds = np.mean(avg_preds, axis=0)\n        probs.append(avg_preds)\n    probs = np.concatenate(probs)\n    return probs","metadata":{"execution":{"iopub.status.busy":"2023-04-12T07:17:43.862238Z","iopub.execute_input":"2023-04-12T07:17:43.862694Z","iopub.status.idle":"2023-04-12T07:17:43.877230Z","shell.execute_reply.started":"2023-04-12T07:17:43.862608Z","shell.execute_reply":"2023-04-12T07:17:43.874492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ****inference****\n\nในแต่ละรอบการทำนาย (iteration) มีการใช้ TTA (test time augmentation) เพื่อเพิ่มประสิทธิภาพในการทำนายผลลัพธ์ และการแบ่งน้ำหนัก (weight) ของแต่ละโมเดลเพื่อให้มีน้ำหนักเท่าๆ กันก่อนจะนำมา ensemble\n\nหลังจากทำนายผลลัพธ์ด้วยโมเดลทั้งหมดแล้ว จะนำผลลัพธ์ทั้งหมดมาหาค่าเฉลี่ย เพื่อทำการทำนายผลลัพธ์ของตัวอย่างที่ใช้ในการส่งงาน และสุดท้ายจะนำผลลัพธ์ที่ได้มาเขียนเป็นไฟล์ csv เพื่อนำไปส่งงาน","metadata":{}},{"cell_type":"code","source":"# ====================================================\n# inference\n# ====================================================\npredictions = None\nfor model_name in CFG.models:\n    for i in range(CFG.tta):\n        model = CassvaImgClassifier(model_name, pretrained=False)\n        states = []\n        for saved_model in [\"best\", \"final\"]:\n            if CFG.trn_fold[model_name][saved_model] != []:\n                LOGGER.info(\n                    f\"========== Model: {model_name}, TTA: {i}, Saved: {saved_model}, Fold: {CFG.trn_fold[model_name][saved_model]} ==========\"\n                )\n                states += [\n                    torch.load(MODEL_DIR + f\"{model_name}_fold{fold}_{saved_model}.pth\")\n                    for fold in CFG.trn_fold[model_name][saved_model]\n                ]\n\n        if i == 0:  # no TTA\n            test_dataset = TestDataset(test, transform=get_transforms(data=\"valid\", size=CFG.size[model_name]))\n            tta_weight = CFG.no_tta_weight\n        else:\n            test_dataset = TestDataset(\n                test, transform=get_transforms(data=CFG.transform[model_name], size=CFG.size[model_name])\n            )\n            tta_weight = 1\n\n        test_loader = DataLoader(\n            test_dataset, batch_size=CFG.batch_size, shuffle=False, num_workers=CFG.num_workers, pin_memory=True\n        )\n\n        inf = inference(model, states, test_loader, device, CFG.data_parallel[model_name])\n        LOGGER.info(f\"Inference example: {inf[0]}\")\n\n        if predictions is None:\n            predictions = inf[np.newaxis] * CFG.weight[model_name] * tta_weight\n        else:\n            predictions = np.append(predictions, inf[np.newaxis] * CFG.weight[model_name] * tta_weight, axis=0)\n\nsub = np.sum(predictions, axis=0) / weight_sum\nLOGGER.info(f\"========== Overall ==========\")\nLOGGER.info(f\"Submission example: {sub[0]}\")\n\n\n# submission\ntest[\"label\"] = sub.argmax(1)\ntest[[\"image_id\", \"label\"]].to_csv(OUTPUT_DIR + \"submission.csv\", index=False)\ntest.head()","metadata":{"execution":{"iopub.status.busy":"2023-04-12T07:17:43.885359Z","iopub.execute_input":"2023-04-12T07:17:43.886352Z","iopub.status.idle":"2023-04-12T07:19:29.864288Z","shell.execute_reply.started":"2023-04-12T07:17:43.886307Z","shell.execute_reply":"2023-04-12T07:19:29.862972Z"},"trusted":true},"execution_count":null,"outputs":[]}]}