{"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":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-07-01T05:12:00.336463Z","iopub.execute_input":"2023-07-01T05:12:00.337180Z","iopub.status.idle":"2023-07-01T05:12:00.342716Z","shell.execute_reply.started":"2023-07-01T05:12:00.337145Z","shell.execute_reply":"2023-07-01T05:12:00.341797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install py7zr","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:12:00.344553Z","iopub.execute_input":"2023-07-01T05:12:00.345153Z","iopub.status.idle":"2023-07-01T05:12:12.113213Z","shell.execute_reply.started":"2023-07-01T05:12:00.345122Z","shell.execute_reply":"2023-07-01T05:12:12.111327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pip install wandb","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:12:12.120209Z","iopub.execute_input":"2023-07-01T05:12:12.121043Z","iopub.status.idle":"2023-07-01T05:12:23.405118Z","shell.execute_reply.started":"2023-07-01T05:12:12.120998Z","shell.execute_reply":"2023-07-01T05:12:23.403589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport matplotlib.pyplot as plt\nimport torch\nfrom torch.utils.data import DataLoader, Dataset, ConcatDataset, SubsetRandomSampler\nfrom torchvision.datasets import MNIST\nfrom torchvision import datasets, transforms\nfrom torchvision.models import resnet18\nimport torchvision\nimport torch.nn as nn\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.model_selection import KFold\nimport cv2\nimport os\nfrom tqdm.notebook import trange, tqdm\nimport albumentations\nfrom timeit import default_timer as timer\nimport glob\nimport py7zr\nfrom sklearn.model_selection import train_test_split\nimport wandb\nfrom torch.nn import functional as F\nimport math\nNO_LABEL = -1","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:12:23.407634Z","iopub.execute_input":"2023-07-01T05:12:23.408051Z","iopub.status.idle":"2023-07-01T05:12:23.419001Z","shell.execute_reply.started":"2023-07-01T05:12:23.408009Z","shell.execute_reply":"2023-07-01T05:12:23.418085Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp_path = '/kaggle/working/'\nif not os.path.exists(temp_path):\n    os.mkdir(temp_path)\ntrain_file_path = '/kaggle/input/c/cifar-10/train.7z'\narchive = py7zr.SevenZipFile(train_file_path, mode='r')\narchive.extractall(path=temp_path)\narchive.close()","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:12:23.422671Z","iopub.execute_input":"2023-07-01T05:12:23.423441Z","iopub.status.idle":"2023-07-01T05:13:25.550316Z","shell.execute_reply.started":"2023-07-01T05:12:23.423409Z","shell.execute_reply":"2023-07-01T05:13:25.549259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp_path = '/kaggle/working/'\nif not os.path.exists(temp_path):\n    os.mkdir(temp_path)\ntest_label_file = '/kaggle/input/c/cifar-10/test.7z'\narchive = py7zr.SevenZipFile(test_label_file, mode='r')\narchive.extractall(path=temp_path)\narchive.close()","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:13:25.551792Z","iopub.execute_input":"2023-07-01T05:13:25.552149Z","iopub.status.idle":"2023-07-01T05:33:46.867174Z","shell.execute_reply.started":"2023-07-01T05:13:25.552117Z","shell.execute_reply":"2023-07-01T05:33:46.866110Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labeled_images=glob.glob('/kaggle/working/train/*.png')\nlen(labeled_images)","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:46.868597Z","iopub.execute_input":"2023-07-01T05:33:46.868969Z","iopub.status.idle":"2023-07-01T05:33:47.058432Z","shell.execute_reply.started":"2023-07-01T05:33:46.868937Z","shell.execute_reply":"2023-07-01T05:33:47.057457Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_images,val_images= train_test_split(labeled_images,test_size=0.33,random_state=42)\n","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:47.060180Z","iopub.execute_input":"2023-07-01T05:33:47.060876Z","iopub.status.idle":"2023-07-01T05:33:47.087604Z","shell.execute_reply.started":"2023-07-01T05:33:47.060841Z","shell.execute_reply":"2023-07-01T05:33:47.086786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_images= glob.glob('/kaggle/working/test/*.png')\nlen(test_images)\n","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:47.089152Z","iopub.execute_input":"2023-07-01T05:33:47.089519Z","iopub.status.idle":"2023-07-01T05:33:48.119366Z","shell.execute_reply.started":"2023-07-01T05:33:47.089475Z","shell.execute_reply":"2023-07-01T05:33:48.118208Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp_label = ['frog', 'truck', 'deer', 'automobile', 'bird', 'horse', 'ship', 'cat', 'dog', 'airplane']\n\nlabel_dict = {}  # Create an empty dictionary\n\nfor index, label in enumerate(temp_label):\n    label_dict[label] = index\n\nprint(label_dict)","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.121270Z","iopub.execute_input":"2023-07-01T05:33:48.121637Z","iopub.status.idle":"2023-07-01T05:33:48.130589Z","shell.execute_reply.started":"2023-07-01T05:33:48.121605Z","shell.execute_reply":"2023-07-01T05:33:48.129496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class Cifar10(Dataset):\n  def __init__(self,labeled_images,label_files= None,transform=None,label_dict=None):\n    self.transform = transform\n    if label_files !=None:\n        self.labels= pd.read_csv(label_files)\n    self.list_images=labeled_images\n    self.label_dict=label_dict\n\n    \n   \n    print(len(self.list_images))\n  def __len__(self):\n     return len(self.list_images)\n\n  def __getitem__(self, idx):\n    image_path=self.list_images[idx]\n    image = cv2.imread(image_path)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    image = cv2.resize(image, (224, 224))\n    \n    \n    if self.transform:\n        res = self.transform(image=image)\n        image1 = res['image'].astype(np.float32)\n    else:\n        image1 = image.astype(np.float32)\n\n    image1= image.transpose(2, 0, 1)\n    image2= image.astype(np.float32)\n    image2= image.transpose(2,0,1)\n    \n  \n    label= self.labels.iloc[int(image_path[22:-4])-1]['label']\n    index=label_dict[label]\n\n    \n    sample= (image1,image2,index)\n    return sample # image, label\nclass Cifar10test(Dataset):\n  def __init__(self,test_images,transform=None):\n    self.transform = transform\n    self.test_images= test_images\n    \n   \n    print(len(self.test_images))\n  def __len__(self):\n     return len(self.test_images)\n\n  def __getitem__(self, idx):\n    image_path=self.test_images[idx]\n    image = cv2.imread(image_path)\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    image = cv2.resize(image, (224, 224))\n    \n    \n    if self.transform:\n        res = self.transform(image=image)\n        image = res['image'].astype(np.float32)\n    else:\n        image = image.astype(np.float32)\n\n    image = image.transpose(2, 0, 1)\n    \n    sample=image\n    return sample # image, label","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.136467Z","iopub.execute_input":"2023-07-01T05:33:48.137529Z","iopub.status.idle":"2023-07-01T05:33:48.164010Z","shell.execute_reply.started":"2023-07-01T05:33:48.137497Z","shell.execute_reply":"2023-07-01T05:33:48.163006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms(image_size=(224, 224)):\n\n    transforms_student = albumentations.Compose([\n        albumentations.HorizontalFlip(p=0.5),\n        albumentations.ImageCompression(quality_lower=99, quality_upper=100),\n        albumentations.ShiftScaleRotate(shift_limit=0.2, scale_limit=0.2, rotate_limit=10, border_mode=0, p=0.7),\n        # albumentations.Resize(image_size, image_size),\n        # albumentations.Cutout(max_h_size=int(image_size * 0.4), max_w_size=int(image_size * 0.4), num_holes=1, p=0.5),\n        albumentations.Normalize()\n    ])\n\n    transforms_val = albumentations.Compose([\n        albumentations.HorizontalFlip(p=0.5),\n        albumentations.ImageCompression(quality_lower=99, quality_upper=100),\n        albumentations.ShiftScaleRotate(shift_limit=0.2, scale_limit=0.2, rotate_limit=10, border_mode=0, p=0.7),\n        # albumentations.Resize(image_size, image_size),\n        # albumentations.Cutout(max_h_size=int(image_size * 0.4), max_w_size=int(image_size * 0.4), num_holes=1, p=0.5),\n        albumentations.Normalize()\n    ])\n    \n\n    return transforms_student, transforms_val","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.165458Z","iopub.execute_input":"2023-07-01T05:33:48.166074Z","iopub.status.idle":"2023-07-01T05:33:48.181377Z","shell.execute_reply.started":"2023-07-01T05:33:48.166043Z","shell.execute_reply":"2023-07-01T05:33:48.180376Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transforms_train, transforms_val = get_transforms(image_size=(224, 224))","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.182866Z","iopub.execute_input":"2023-07-01T05:33:48.183485Z","iopub.status.idle":"2023-07-01T05:33:48.196481Z","shell.execute_reply.started":"2023-07-01T05:33:48.183446Z","shell.execute_reply":"2023-07-01T05:33:48.195182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_datas=Cifar10(train_images,label_files='/kaggle/input/c/cifar-10/trainLabels.csv',transform= transforms_train,label_dict=label_dict)\nval_datas=Cifar10(val_images,label_files='/kaggle/input/c/cifar-10/trainLabels.csv',transform= transforms_val,label_dict=label_dict)\ntest_data= Cifar10test(test_images,transform= None)","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.197986Z","iopub.execute_input":"2023-07-01T05:33:48.198250Z","iopub.status.idle":"2023-07-01T05:33:48.245208Z","shell.execute_reply.started":"2023-07-01T05:33:48.198229Z","shell.execute_reply":"2023-07-01T05:33:48.242920Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader= DataLoader(train_datas, batch_size= 32, shuffle= True)\n\nval_loader= DataLoader(val_datas,batch_size=32, shuffle= True)\ntest_loader= DataLoader(test_data,batch_size=32, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.247863Z","iopub.execute_input":"2023-07-01T05:33:48.248793Z","iopub.status.idle":"2023-07-01T05:33:48.277032Z","shell.execute_reply.started":"2023-07-01T05:33:48.248732Z","shell.execute_reply":"2023-07-01T05:33:48.275863Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MODEL","metadata":{}},{"cell_type":"code","source":"class ResNet224x224(nn.Module):\n    def __init__(self, block, layers, channels, groups=1, num_classes=1000, downsample='basic'):\n        super().__init__()\n        assert len(layers) == 4\n        self.downsample_mode = downsample\n        self.inplanes = 64\n        self.conv1 = nn.Conv2d(3, self.inplanes, kernel_size=7, stride=2, padding=3,\n                               bias=False)\n        self.bn1 = nn.BatchNorm2d(self.inplanes)\n        self.relu = nn.ReLU(inplace=True)\n        self.maxpool = nn.MaxPool2d(kernel_size=3, stride=2, padding=1)\n        self.layer1 = self._make_layer(block, channels, groups, layers[0])\n        self.layer2 = self._make_layer(\n            block, channels * 2, groups, layers[1], stride=2)\n        self.layer3 = self._make_layer(\n            block, channels * 4, groups, layers[2], stride=2)\n        self.layer4 = self._make_layer(\n            block, channels * 8, groups, layers[3], stride=2)\n        self.avgpool = nn.AvgPool2d(7)\n        self.fc1 = nn.Linear(block.out_channels(\n            channels * 8, groups), num_classes)\n        self.fc2 = nn.Linear(block.out_channels(\n            channels * 8, groups), num_classes)\n\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                n = m.kernel_size[0] * m.kernel_size[1] * m.out_channels\n                m.weight.data.normal_(0, math.sqrt(2. / n))\n            elif isinstance(m, nn.BatchNorm2d):\n                m.weight.data.fill_(1)\n                m.bias.data.zero_()\n\n    def _make_layer(self, block, planes, groups, blocks, stride=1):\n        downsample = None\n        if stride != 1 or self.inplanes != block.out_channels(planes, groups):\n            if self.downsample_mode == 'basic' or stride == 1:\n                downsample = nn.Sequential(\n                    nn.Conv2d(self.inplanes, block.out_channels(planes, groups),\n                              kernel_size=1, stride=stride, bias=False),\n                    nn.BatchNorm2d(block.out_channels(planes, groups)),\n                )\n            elif self.downsample_mode == 'shift_conv':\n                downsample = ShiftConvDownsample(in_channels=self.inplanes,\n                                                 out_channels=block.out_channels(planes, groups))\n            else:\n                assert False\n\n        layers = []\n        layers.append(block(self.inplanes, planes, groups, stride, downsample))\n        self.inplanes = block.out_channels(planes, groups)\n        for i in range(1, blocks):\n            layers.append(block(self.inplanes, planes, groups))\n\n        return nn.Sequential(*layers)\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.relu(x)\n        x = self.maxpool(x)\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n        x = self.avgpool(x)\n        x = x.view(x.size(0), -1)\n        return self.fc1(x),self.fc2(x)\n    \nclass BottleneckBlock(nn.Module):\n    @classmethod\n    def out_channels(cls, planes, groups):\n        if groups > 1:\n            return 2 * planes\n        else:\n            return 4 * planes\n\n    def __init__(self, inplanes, planes, groups, stride=1, downsample=None):\n        super().__init__()\n        self.relu = nn.ReLU(inplace=True)\n\n        self.conv_a1 = nn.Conv2d(inplanes, planes, kernel_size=1, bias=False)\n        self.bn_a1 = nn.BatchNorm2d(planes)\n        self.conv_a2 = nn.Conv2d(\n            planes, planes, kernel_size=3, stride=stride, padding=1, bias=False, groups=groups)\n        self.bn_a2 = nn.BatchNorm2d(planes)\n        self.conv_a3 = nn.Conv2d(planes, self.out_channels(\n            planes, groups), kernel_size=1, bias=False)\n        self.bn_a3 = nn.BatchNorm2d(self.out_channels(planes, groups))\n\n        self.downsample = downsample\n        self.stride = stride\n\n    def forward(self, x):\n        a, residual = x, x\n\n        a = self.conv_a1(a)\n        a = self.bn_a1(a)\n        a = self.relu(a)\n        a = self.conv_a2(a)\n        a = self.bn_a2(a)\n        a = self.relu(a)\n        a = self.conv_a3(a)\n        a = self.bn_a3(a)\n\n        if self.downsample is not None:\n            residual = self.downsample(residual)\n\n        return self.relu(residual + a)\n\ndef get_model(is_mean_teacher=False,device=None):\n    model=ResNet224x224(BottleneckBlock,\n                          layers=[3, 8, 36, 3],\n                          channels=32 * 4,\n                          groups=32,\n                          downsample='basic')\n    model= nn.DataParallel(model)\n    model = model.to(device)\n    \n    # Detach params for Exponential Moving Average Model (aka the Mean Teacher).\n    # We'll manually update these params instead of using backprop.\n    if is_mean_teacher:\n        for param in model.parameters():\n            param.detach_()\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.279704Z","iopub.execute_input":"2023-07-01T05:33:48.281118Z","iopub.status.idle":"2023-07-01T05:33:48.347957Z","shell.execute_reply.started":"2023-07-01T05:33:48.281084Z","shell.execute_reply":"2023-07-01T05:33:48.345335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss function","metadata":{}},{"cell_type":"markdown","source":"        # Costs\n        'consistency_type': 'mse',\n        'consistency_rampup': 5,\n        'consistency': 100.0,\n        'logit_distance_cost': .01,\n        'weight_decay': 2e-4,\n\n        # Optimization\n        'lr_rampup': 0,\n        'base_lr': 0.05,\n        'nesterov': True,","metadata":{}},{"cell_type":"code","source":"def softmax_mse_loss(input_logits, target_logits):\n    \n    assert input_logits.size() == target_logits.size()\n    input_softmax = F.softmax(input_logits, dim=1)\n    target_softmax = F.softmax(target_logits, dim=1)\n    num_classes = input_logits.size()[1]\n    return F.mse_loss(input_softmax, target_softmax, size_average=False) / num_classes\n\n\ndef softmax_kl_loss(input_logits, target_logits):\n    \n    assert input_logits.size() == target_logits.size()\n    input_log_softmax = F.log_softmax(input_logits, dim=1)\n    target_softmax = F.softmax(target_logits, dim=1)\n    return F.kl_div(input_log_softmax, target_softmax, size_average=False)\n\n\ndef symmetric_mse_loss(input1, input2):\n    assert input1.size() == input2.size()\n    num_classes = input1.size()[1]\n    return torch.sum((input1 - input2)**2) / num_classes\n","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.350976Z","iopub.execute_input":"2023-07-01T05:33:48.351939Z","iopub.status.idle":"2023-07-01T05:33:48.367957Z","shell.execute_reply.started":"2023-07-01T05:33:48.351902Z","shell.execute_reply":"2023-07-01T05:33:48.366262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def sigmoid_rampup(current, rampup_length):\n  \n    if rampup_length == 0:\n        return 1.0\n    else:\n        current = np.clip(current, 0.0, rampup_length)\n        phase = 1.0 - current / rampup_length\n        return float(np.exp(-5.0 * phase * phase))\n\n\ndef linear_rampup(current, rampup_length):\n   \n    assert current >= 0 and rampup_length >= 0\n    if current >= rampup_length:\n        return 1.0\n    else:\n        return current / rampup_length\n\n\ndef cosine_rampdown(current, rampdown_length):\n   \n    assert 0 <= current <= rampdown_length\n    return float(.5 * (np.cos(np.pi * current / rampdown_length) + 1))","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.371769Z","iopub.execute_input":"2023-07-01T05:33:48.372989Z","iopub.status.idle":"2023-07-01T05:33:48.382120Z","shell.execute_reply.started":"2023-07-01T05:33:48.372955Z","shell.execute_reply":"2023-07-01T05:33:48.380885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_current_consistency_weight(epoch):\n    return 100.0 * sigmoid_rampup(epoch, 5)","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.384016Z","iopub.execute_input":"2023-07-01T05:33:48.384966Z","iopub.status.idle":"2023-07-01T05:33:48.394265Z","shell.execute_reply.started":"2023-07-01T05:33:48.384933Z","shell.execute_reply":"2023-07-01T05:33:48.393187Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"def update_teacher_params(student, teacher, alpha, global_step):\n    # Use the true average until the exponential average is more correct\n    alpha = min(1 - 1 / (global_step + 1), alpha)\n    for ema_param, param in zip(teacher.parameters(), student.parameters()):\n        ema_param.data.mul_(alpha).add_(1 - alpha, param.data)","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.396098Z","iopub.execute_input":"2023-07-01T05:33:48.396948Z","iopub.status.idle":"2023-07-01T05:33:48.404136Z","shell.execute_reply.started":"2023-07-01T05:33:48.396917Z","shell.execute_reply":"2023-07-01T05:33:48.403135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(train_loader,student, teacher, optimizer, device, epoch):\n    \n    class_criterion = nn.CrossEntropyLoss(size_average=False, ignore_index=NO_LABEL).cuda()\n    consistency_criterion = softmax_mse_loss\n    student.train()\n    teacher.train()\n    running_loss=0\n    global_step = 0\n    for i, (images1, images2,labels) in tqdm(enumerate(train_loader),total=len(train_loader)):\n        images1 = images1.to(device)\n        images2= images2.to(device)\n        \n        images1= images1.float()\n        images2= images2.float()\n        labels= labels.float()\n        labels = labels.type(torch.LongTensor)\n        labels = labels.to(device)\n         #meta_images = images + torch.rand(images.size()).to(device) * 1.0 + 0.1 # augment gaussian\n            \n        # forward pass\n        student_pred,_ = student(images1)\n        teacher_pred,_= teacher(images2)\n        \n        student_classif, student_consistency = student_pred, student_pred\n#       \n        \n        student_class_loss = class_criterion(student_classif, labels)/(len(images1))\n        consistency_weight = get_current_consistency_weight(epoch)\n        \n        consistency_loss = consistency_weight * consistency_criterion(student_consistency, teacher_pred) / (len(images1))\n        loss = student_class_loss + consistency_loss\n        running_loss += loss.item()\n\n        # backward and optimizer\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n        global_step += 1\n        update_teacher_params(student, teacher, 0.995, global_step)\n    epoch_loss = running_loss / len(train_loader)\n\n    return loss, student_class_loss, consistency_loss, consistency_weight, epoch_loss","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.405959Z","iopub.execute_input":"2023-07-01T05:33:48.406881Z","iopub.status.idle":"2023-07-01T05:33:48.419626Z","shell.execute_reply.started":"2023-07-01T05:33:48.406850Z","shell.execute_reply":"2023-07-01T05:33:48.418338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def validate(eval_loader, student, epoch,device):\n    class_criterion = nn.CrossEntropyLoss(size_average=False, ignore_index=NO_LABEL).cuda()\n    \n\n    # switch to evaluate mode\n    student.eval()\n  \n    running_loss=0\n    for i, (images1,images2, labels) in enumerate(eval_loader):\n        images1 = images1.to(device)\n      \n        \n        images1= images1.float()\n       \n        labels= labels.float()\n        labels = labels.type(torch.LongTensor)\n        labels = labels.to(device)\n        student_pred,_ = student(images1)\n        valid_loss = class_criterion(student_pred, labels)/(len(images1))\n        running_loss += valid_loss.item()\n    epoch_loss = running_loss / len(eval_loader)\n    return student,epoch_loss","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.421586Z","iopub.execute_input":"2023-07-01T05:33:48.422524Z","iopub.status.idle":"2023-07-01T05:33:48.433526Z","shell.execute_reply.started":"2023-07-01T05:33:48.422494Z","shell.execute_reply":"2023-07-01T05:33:48.432422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_accuracy(model, data_loader, device):\n    correct = 0\n    total = 0\n\n    with torch.no_grad():\n        model.eval()\n        for images1, images2,labels in data_loader:\n            images1 = images1.to(device)\n         \n            images1= images1.float()\n            labels= labels.float()\n            labels = labels.type(torch.LongTensor)\n            labels = labels.to(device)\n            outputs ,_= model(images1)\n            _, predicted = torch.max(outputs.data, 1)\n\n            total += labels.size(0)\n            correct += (predicted == labels).sum().item()\n\n    return 100*(correct/total)","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.435328Z","iopub.execute_input":"2023-07-01T05:33:48.436245Z","iopub.status.idle":"2023-07-01T05:33:48.447320Z","shell.execute_reply.started":"2023-07-01T05:33:48.436212Z","shell.execute_reply":"2023-07-01T05:33:48.446154Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_losses(train_acc, valid_acc, train_loss, valid_loss):\n    # change the style of the plots to seaborn\n    plt.style.use('seaborn')\n\n    train_acc = np.array(train_acc)\n    valid_acc = np.array(valid_acc)\n\n    fig, (ax1, ax2) = plt.subplots(1, 2)\n\n    ax1.plot(train_acc, color=\"blue\", label=\"Train_acc\")\n    ax1.plot(valid_acc, color=\"red\", label=\"Validation_acc\")\n    ax1.set(title=\"Acc over epochs\",\n            xlabel=\"Epoch\",\n            ylabel=\"Acc\")\n    ax1.legend()\n\n    ax2.plot(train_loss, color=\"blue\", label=\"Train_loss\")\n    ax2.plot(valid_loss, color=\"red\", label=\"Validation_loss\")\n    ax2.set(title=\"loss over epochs\",\n            xlabel=\"Epoch\",\n            ylabel=\"Loss\")\n    ax2.legend()\n\n    fig.show()\n\n    # change the plot style to default\n    plt.style.use('default')","metadata":{"execution":{"iopub.status.busy":"2023-07-01T05:33:48.449119Z","iopub.execute_input":"2023-07-01T05:33:48.450037Z","iopub.status.idle":"2023-07-01T05:33:48.459678Z","shell.execute_reply.started":"2023-07-01T05:33:48.450007Z","shell.execute_reply":"2023-07-01T05:33:48.458488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def training_loop(epochs,device,train_loader,valid_loader):\n\n    if not os.path.exists(\"save_model\"):\n        os.mkdir(\"save_model\")\n    student_model = get_model(device=device)\n    teacher_model = get_model(is_mean_teacher=True,device=device)\n    \n    train_losses = []\n    valid_losses = []\n    list_train_acc = []\n    list_val_acc = []\n    optimizer = torch.optim.SGD(student_model.parameters(), 0.05,\n                                momentum=0.9,\n                                weight_decay=2e-4,\n                                nesterov=True)\n\n    for epoch in tqdm(range(0, epochs), total= epochs):\n        # training\n        loss, student_class_loss, consistency_loss, consistency_weights, train_loss= train(train_loader,student_model, teacher_model, optimizer, device, epoch)\n        \n        with torch.no_grad():\n            student_model, valid_loss = validate(valid_loader, student_model, epoch,device)\n        train_acc = get_accuracy(student_model, train_loader, device=device)\n        valid_acc = get_accuracy(student_model, valid_loader, device=device)\n        \n        print('Epochs: {}, Train_loss: {}, Valid_loss: {}, Train_accuracy: {}, Valid_accuracy: {}'.format(\n                    epoch, train_loss, valid_loss, train_acc, valid_acc\n                    ))\n        list_train_acc.append(train_acc)\n        list_val_acc.append(valid_acc)\n        train_losses.append(train_loss)\n        valid_losses.append(valid_loss)\n        \n        torch.save(student_model.state_dict(), \"save_model/epoch_{}_acc{}.pth\".format(epoch+1, valid_acc))\n    plot_losses(list_train_acc, list_val_acc, train_losses, valid_losses)\n\n    return student_model, optimizer, (train_losses, valid_losses)","metadata":{"execution":{"iopub.status.busy":"2023-07-01T06:06:17.335277Z","iopub.execute_input":"2023-07-01T06:06:17.335659Z","iopub.status.idle":"2023-07-01T06:06:17.347231Z","shell.execute_reply.started":"2023-07-01T06:06:17.335629Z","shell.execute_reply":"2023-07-01T06:06:17.345314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"student_model,optimizer,_ = training_loop(10,device,train_loader,val_loader)","metadata":{"execution":{"iopub.status.busy":"2023-07-01T07:28:39.415693Z","iopub.execute_input":"2023-07-01T07:28:39.416588Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}