{"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":"!pip install timm","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-10-29T12:11:56.196821Z","iopub.execute_input":"2022-10-29T12:11:56.197264Z","iopub.status.idle":"2022-10-29T12:12:05.933423Z","shell.execute_reply.started":"2022-10-29T12:11:56.197168Z","shell.execute_reply":"2022-10-29T12:12:05.932071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## ライブラリ","metadata":{}},{"cell_type":"code","source":"import os\nimport sys\nimport random\nfrom contextlib import contextmanager\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nfrom matplotlib import pyplot as plt\nimport matplotlib.image as mpimg\nfrom tqdm.notebook import tqdm\n\nfrom sklearn.metrics import mean_squared_error,accuracy_score\nfrom sklearn.model_selection import KFold\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder\n\nimport cv2\nfrom IPython.core.display import display\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\nfrom torch.optim import Adam\n\n\nimport timm\n\nimport albumentations as transforms\nfrom albumentations.pytorch import ToTensorV2\n\nimport warnings\nwarnings.filterwarnings('ignore')\ntorch.backends.cudnn.benchmark = True\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nsys.path.append('/kaggle/input/h-and-m-personalized-fashion-recommendations/images')\nOUTPUT_DIR = '/kaggle/working/'\nINPUT_DIR = '/kaggle/input/h-and-m-personalized-fashion-recommendations/images'\n# if not os.path.exists(OUTPUT_DIR):\n#     os.makedirs(OUTPUT_DIR)","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:05.937275Z","iopub.execute_input":"2022-10-29T12:12:05.937663Z","iopub.status.idle":"2022-10-29T12:12:08.723399Z","shell.execute_reply.started":"2022-10-29T12:12:05.937622Z","shell.execute_reply":"2022-10-29T12:12:08.722094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## データ表示&EDA","metadata":{}},{"cell_type":"code","source":"TRAIN_DIR = '/kaggle/input/h-and-m-personalized-fashion-recommendations/images'\narticles = pd.read_csv(\"/kaggle/input/h-and-m-personalized-fashion-recommendations/articles.csv\",dtype={\"article_id\": str})\n\n#画像のパスをカラムとして加える\n#idを指定すると相対パスを得られる関数を定義\ndef get_train_file_path(id):\n    folder_id = id[0:3]\n    return f\"{TRAIN_DIR}/{folder_id}/{id}.jpg\"\narticles['img_path'] = articles['article_id'].apply(get_train_file_path)","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:08.725893Z","iopub.execute_input":"2022-10-29T12:12:08.726622Z","iopub.status.idle":"2022-10-29T12:12:09.296671Z","shell.execute_reply.started":"2022-10-29T12:12:08.726577Z","shell.execute_reply":"2022-10-29T12:12:09.295677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def showpics_by_productgroup(df, num_images,random_state,product_group_name = False):\n    #表示させる枚数\n    num_images = num_images\n    #seedを設定\n    random_state = random_state\n    #種類で指定してランダムに取得\n    if not product_group_name:\n        random_sample = df.sample(num_images, random_state=random_state).reset_index(drop=True)\n    else:\n        random_sample = df[df[\"product_group_name\"] == product_group_name].sample(num_images, random_state=random_state).reset_index(drop=True)\n    \n    #The for loop goes as many loops as specified by the num_images\n    plt.subplots(1, num_images, figsize=(20,10))\n    for x in range(num_images):\n        #start from the id in the dataframe\n        image_path = random_sample.iloc[x]['img_path']\n        #use plt.imread() to read in the image file\n        image_array = plt.imread(image_path)\n        #make a subplot space that is 1 down and num_images across\n        plt.subplot(1, num_images, x+1)\n        #title\n        title = random_sample.iloc[x]['product_group_name']\n        plt.title(title) \n        #turn off gridlines\n        plt.axis('off')\n        #then plt.imshow() can display it for you\n        plt.imshow(image_array)\n    plt.show()\n    plt.close()","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:09.299487Z","iopub.execute_input":"2022-10-29T12:12:09.299879Z","iopub.status.idle":"2022-10-29T12:12:09.310067Z","shell.execute_reply.started":"2022-10-29T12:12:09.299843Z","shell.execute_reply":"2022-10-29T12:12:09.308937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 色を指定してn枚表示\nnum_images = 5\nrandom_state = 123\nproduct_group_name = 'Shoes'\nshowpics_by_productgroup(articles, num_images,random_state,product_group_name)","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:09.311590Z","iopub.execute_input":"2022-10-29T12:12:09.312047Z","iopub.status.idle":"2022-10-29T12:12:11.130098Z","shell.execute_reply.started":"2022-10-29T12:12:09.312010Z","shell.execute_reply":"2022-10-29T12:12:11.129099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(articles['product_group_name'].value_counts())","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:11.131279Z","iopub.execute_input":"2022-10-29T12:12:11.132375Z","iopub.status.idle":"2022-10-29T12:12:11.148715Z","shell.execute_reply.started":"2022-10-29T12:12:11.132325Z","shell.execute_reply":"2022-10-29T12:12:11.147629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#article_idに対応する画像がないものもある\nall_file = set()\n\nfor folder in os.listdir(\"/kaggle/input/h-and-m-personalized-fashion-recommendations/images\"):\n    for file in os.listdir(f\"/kaggle/input/h-and-m-personalized-fashion-recommendations/images/{folder}\"):\n        all_file.add(file)\n\nprint(f\"idの数：{articles.shape[0]}\")\nprint(f\"写真の数：{len(all_file)}\")\nprint(f\"写真がないidの数：{articles.shape[0]-len(all_file)}\")","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:11.150350Z","iopub.execute_input":"2022-10-29T12:12:11.150720Z","iopub.status.idle":"2022-10-29T12:12:11.322805Z","shell.execute_reply.started":"2022-10-29T12:12:11.150683Z","shell.execute_reply":"2022-10-29T12:12:11.321793Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CFG","metadata":{}},{"cell_type":"code","source":"class CFG:\n    apex=False\n    debug=False\n    print_freq=10\n    num_workers=8\n    size=224 ##モデルによって変える。\n    model_name='vit_base_patch16_224' ##モデルによって変える\n    scheduler='CosineAnnealingLR' # ['ReduceLROnPlateau', 'CosineAnnealingLR', 'CosineAnnealingWarmRestarts']\n    epochs=5\n    factor=0.2 # ReduceLROnPlateau\n    patience=4 # ReduceLROnPlateau\n    eps=1e-6 # ReduceLROnPlateau\n    T_max=3 # CosineAnnealingLR\n    T_0=3 # CosineAnnealingWarmRestarts\n    lr=3e-5\n    min_lr=1e-6\n    batch_size=32 # sample_size_maltiple=1の時256, 0.1の時16にしていた\n    weight_decay=1e-8\n    gradient_accumulation_steps=1\n    max_grad_norm=1000\n    seed=0\n    target_size=9 ##モデルによって変える\n    target_col='product_group_name_encoded'\n    n_fold=3\n    trn_fold=[0,1,2]\n    train=True\n    grad_cam=True\n    isTransFormer = False ##モデルによって変える\n    # color_space = \"COLOR_BGR2RGB\"\n    minimum_count = 121\n    sample_size_multiple = 0.1\n    \n    \n    ##メモを書く\"\n    note =str(sample_size_multiple)+\"データ使用,batch\"+str(batch_size)+\"-lr\"+str(lr)+\",重みづけ\"\n\n    MODEL_PATH = \"../input/vit-base-models-pretrained-pytorch/jx_vit_base_p16_224-80ecf9dd.pth\"\n\n# seedの固定\ndef fix_seed(seed):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n\nSEED = 0\nfix_seed(SEED)","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:11.324003Z","iopub.execute_input":"2022-10-29T12:12:11.324682Z","iopub.status.idle":"2022-10-29T12:12:11.341690Z","shell.execute_reply.started":"2022-10-29T12:12:11.324642Z","shell.execute_reply":"2022-10-29T12:12:11.340634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## wandb","metadata":{}},{"cell_type":"code","source":"import wandb\n\ntry:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    api_key = user_secrets.get_secret(\"wandb_api\")\n    wandb.login(key=api_key)\n    anony = None\nexcept:\n    anony = \"must\"\n    print('If you want to use your W&B account, go to Add-ons -> Secrets and provide your W&B access token. Use the Label name as wandb_api. \\nGet your W&B access token from here: https://wandb.ai/authorize')\n\n\ndef class2dict(f):\n    return dict((name, getattr(f, name)) for name in dir(f) if not name.startswith('__'))\n\nrun = wandb.init(project=\"H&M 服の9分類\", \n                 config=class2dict(CFG),\n                 name = CFG.note + \"+\" + CFG.model_name,\n                 job_type=\"train\")","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:11.345069Z","iopub.execute_input":"2022-10-29T12:12:11.345465Z","iopub.status.idle":"2022-10-29T12:12:18.986358Z","shell.execute_reply.started":"2022-10-29T12:12:11.345402Z","shell.execute_reply":"2022-10-29T12:12:18.985057Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## データロード","metadata":{}},{"cell_type":"code","source":"TRAIN_DIR = '/kaggle/input/h-and-m-personalized-fashion-recommendations/images'\narticles = pd.read_csv(\"/kaggle/input/h-and-m-personalized-fashion-recommendations/articles.csv\",dtype={\"article_id\": str})\n\n#画像のパスをカラムとして加える\n#idを指定すると相対パスを得られる関数を定義\ndef get_train_file_path(id):\n    folder_id = id[0:3]\n    return f\"{TRAIN_DIR}/{folder_id}/{id}.jpg\"\narticles['img_path'] = articles['article_id'].apply(get_train_file_path)\n\n#必要なカラムのみ取り出す\ncols = ['article_id',\n        'product_group_name',\n        'img_path']\ndf = articles[cols]\n\n#画像が入っていないものを削除\nall_file = set()\n\nfor folder in os.listdir(\"/kaggle/input/h-and-m-personalized-fashion-recommendations/images\"):\n    for file in os.listdir(f\"/kaggle/input/h-and-m-personalized-fashion-recommendations/images/{folder}\"):\n        all_file.add(file)\n\ndf[\"has_img\"] = df[\"article_id\"].map(lambda x: f\"{x}.jpg\" in all_file)\ndf = df[df['has_img']]\n'''\n#色が不明なものは削除する\ndf = df[df['colour_group_name'] != \"Unknown\"]\n'''\n#データ数を調整する\ndf = df.sample(frac=CFG.sample_size_multiple)\n\n#サンプル数が規定より小さい色は削除する\ndf = df[np.where(df['product_group_name'].value_counts()[df[\"product_group_name\"]]>CFG.minimum_count, True, False)]\n\n#残っている種類の数を確認\nprint('残っている種類の数は：' + str(len(df['product_group_name'].value_counts())))\n\n#種類をラベルエンコーディング\nle = LabelEncoder()\nencoded = le.fit_transform(df['product_group_name'].values)\ndecoded = le.inverse_transform(encoded)\ndf['product_group_name_encoded'] = encoded\n\n#エンコーディング後の数が揃っているか確認\nprint(le.classes_)\nprint(len(le.classes_))","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:18.994326Z","iopub.execute_input":"2022-10-29T12:12:18.995504Z","iopub.status.idle":"2022-10-29T12:12:20.218434Z","shell.execute_reply.started":"2022-10-29T12:12:18.995459Z","shell.execute_reply":"2022-10-29T12:12:20.217345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"weights = []\nfor i,label in enumerate(le.classes_):\n    weights.append(1/(articles['product_group_name'].value_counts()[label]))\nweights = torch.tensor(weights).float().cuda()\nprint(weights)","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:20.223340Z","iopub.execute_input":"2022-10-29T12:12:20.225799Z","iopub.status.idle":"2022-10-29T12:12:23.056198Z","shell.execute_reply.started":"2022-10-29T12:12:20.225754Z","shell.execute_reply":"2022-10-29T12:12:23.048207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# trainデータとtestデータに分ける\ndf, df_test = train_test_split(df, test_size=0.2, stratify=df[CFG.target_col])\ndf = df.reset_index(drop=True)\ndf_test = df_test.reset_index(drop=True)\n\n#層化抽出によるK-Fold作成\n#メモ：indexが揃ってないといけない\nskf = StratifiedKFold(n_splits=CFG.n_fold,shuffle=True, random_state=CFG.seed)\n\nfor n, (train_index, val_index) in enumerate(skf.split(df[CFG.target_col], df[CFG.target_col])):\n    df.loc[val_index,'fold'] = int(n)\n\ndf['fold'] = df['fold'].astype(int)\n\n#残っている色の数を確認\nprint(len(df['product_group_name'].value_counts()))\nprint(len(df_test['product_group_name'].value_counts()))","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:23.057760Z","iopub.execute_input":"2022-10-29T12:12:23.058359Z","iopub.status.idle":"2022-10-29T12:12:23.135402Z","shell.execute_reply.started":"2022-10-29T12:12:23.058316Z","shell.execute_reply":"2022-10-29T12:12:23.134345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TrainDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df['img_path'].values\n        self.labels = df[CFG.target_col].values\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_path = self.file_names[idx]\n        try:\n            image = cv2.imread(file_path)\n            image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        except:\n            print(file_path)\n            print(image)\n        \n        if self.transform:\n            image = self.transform(image=image)['image']\n        label = torch.tensor(self.labels[idx]).long()\n        return image, label\n\nclass TestDataset(Dataset):\n    def __init__(self, df, transform=None):\n        self.df = df\n        self.file_names = df['img_path'].values\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        file_path = self.file_names[idx]\n        image = cv2.imread(file_path)\n        print(file_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        if self.transform:\n            image = self.transform(image=image,data=\"valid\")\n        return image","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:23.137992Z","iopub.execute_input":"2022-10-29T12:12:23.138708Z","iopub.status.idle":"2022-10-29T12:12:23.163315Z","shell.execute_reply.started":"2022-10-29T12:12:23.138668Z","shell.execute_reply":"2022-10-29T12:12:23.160777Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms(*, data):\n    if data == 'train':\n        return transforms.Compose([\n            transforms.Resize(CFG.size,CFG.size),\n            transforms.RandomResizedCrop(CFG.size, CFG.size, scale=(0.85, 1.0)),\n            transforms.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])\n\n    elif data == 'valid':\n        return transforms.Compose([\n            transforms.Resize(CFG.size, CFG.size),\n            transforms.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),\n            ToTensorV2(),\n        ])","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:23.168424Z","iopub.execute_input":"2022-10-29T12:12:23.169938Z","iopub.status.idle":"2022-10-29T12:12:23.198510Z","shell.execute_reply.started":"2022-10-29T12:12:23.169900Z","shell.execute_reply":"2022-10-29T12:12:23.197305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## モデル","metadata":{}},{"cell_type":"code","source":"class VitModel(nn.Module):\n    def __init__(self, cfg, pretrained=True):\n        super().__init__()\n        self.cfg = cfg\n        self.model = timm.create_model(self.cfg.model_name, pretrained=pretrained)\n#         if pretrained:\n#             self.model.load_state_dict(torch.load(MODEL_PATH))\n        self.model.head = nn.Linear(self.model.head.in_features, self.cfg.target_size)\n        \n    def forward(self, image):\n        output = self.model(image)\n        return output","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:23.201275Z","iopub.execute_input":"2022-10-29T12:12:23.203306Z","iopub.status.idle":"2022-10-29T12:12:23.220823Z","shell.execute_reply.started":"2022-10-29T12:12:23.203266Z","shell.execute_reply":"2022-10-29T12:12:23.219993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 学習関数","metadata":{}},{"cell_type":"code","source":"class AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\ndef init_logger(log_file=OUTPUT_DIR+'train.log'):\n    from logging import getLogger, INFO, FileHandler,  Formatter,  StreamHandler\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\nLOGGER = init_logger()\n\ndef get_scheduler(optimizer):\n        if CFG.scheduler=='ReduceLROnPlateau':\n            scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=CFG.factor, patience=CFG.patience, verbose=True, eps=CFG.eps)\n        elif CFG.scheduler=='CosineAnnealingLR':\n            scheduler = CosineAnnealingLR(optimizer, T_max=CFG.T_max, eta_min=CFG.min_lr, last_epoch=-1)\n        elif CFG.scheduler=='CosineAnnealingWarmRestarts':\n            scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=CFG.T_0, T_mult=1, eta_min=CFG.min_lr, last_epoch=-1)\n        return scheduler","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:23.222108Z","iopub.execute_input":"2022-10-29T12:12:23.222517Z","iopub.status.idle":"2022-10-29T12:12:23.253689Z","shell.execute_reply.started":"2022-10-29T12:12:23.222485Z","shell.execute_reply":"2022-10-29T12:12:23.252833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_accuracy(y_true,y_pred):\n    y_true = [int(i) for i in y_true.tolist()]\n    score = accuracy_score(y_true= y_true ,y_pred=y_pred)\n    return score\n\ndef get_result(result_df):\n    preds = result_df['preds'].values\n    labels = result_df[CFG.target_col].values\n    score = get_accuracy(labels, preds)\n    LOGGER.info(f'Score: {score:<.4f}')\n\ndef train_fn(fold,train_loader,model,criterion,optimizer,epoch,scheduler,device):\n    model.train()\n    losses = AverageMeter()\n    global_step = 0\n    for step,(images,labels) in enumerate(train_loader):\n        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n        y_preds = model(images)\n        loss = criterion(y_preds, labels)\n      \n        # record loss\n        losses.update(loss.item(), batch_size)\n\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        loss.backward()\n        grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), CFG.max_grad_norm)\n        if (step + 1) % CFG.gradient_accumulation_steps == 0:\n            optimizer.step()\n            optimizer.zero_grad()\n            global_step += 1\n        if step % CFG.print_freq == 0 or step == (len(train_loader)-1):\n            print('Epoch: [{0}][{1}/{2}] '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  'Grad: {grad_norm:.4f}  '\n                  'LR: {lr:.6f}  '\n                  .format(epoch+1, step, len(train_loader),                         \n                          loss=losses,\n                          grad_norm=grad_norm,\n                          lr=scheduler.get_lr()[0]))\n        wandb.log({f\"[fold{fold}] train_loss\": losses.val,\n                    f\"[fold{fold}] train_lr\": scheduler.get_lr()[0]})\n\n            \ndef valid_fn(fold,valid_loader, model, criterion, device):\n    #推論モードに切り替え\n    model.eval()\n    losses = AverageMeter()\n    preds = []\n    for step, (images, labels) in enumerate(valid_loader):\n        images = images.to(device)\n        labels = labels.to(device)\n        batch_size = labels.size(0)\n       \n        # compute loss\n        with torch.no_grad():\n            y_preds = model(images)\n        loss = criterion(y_preds, labels)\n        losses.update(loss.item(), batch_size)\n        # record accuracy\n        preds.append(y_preds.to('cpu').numpy())\n        if CFG.gradient_accumulation_steps > 1:\n            loss = loss / CFG.gradient_accumulation_steps\n        if step % CFG.print_freq == 0 or step == (len(valid_loader)-1):\n            print('EVAL: [{0}/{1}] '\n                  'Loss: {loss.val:.4f}({loss.avg:.4f}) '\n                  .format(step, len(valid_loader),\n                          loss=losses,\n                          ))\n        wandb.log({f\"[fold{fold}] valid_loss\": losses.val})\n    predictions = np.concatenate(preds)\n    return losses.avg, predictions","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:23.257030Z","iopub.execute_input":"2022-10-29T12:12:23.258361Z","iopub.status.idle":"2022-10-29T12:12:23.300286Z","shell.execute_reply.started":"2022-10-29T12:12:23.258327Z","shell.execute_reply":"2022-10-29T12:12:23.299194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_loop(folds, fold):\n    LOGGER.info(f\"========== fold: {fold} training ==========\")\n\n    #dataset\n    trn_idx = folds[folds['fold'] != fold].index\n    val_idx = folds[folds['fold'] == fold].index\n    train_folds = folds.loc[trn_idx].reset_index(drop=True)\n    valid_folds = folds.loc[val_idx].reset_index(drop=True)\n    valid_labels = valid_folds[CFG.target_col].values\n    train_dataset = TrainDataset(train_folds, transform=get_transforms(data='train'))\n    valid_dataset = TrainDataset(valid_folds, transform=get_transforms(data='valid'))\n\n    #dataloader\n    train_loader = DataLoader(train_dataset,\n                              batch_size=CFG.batch_size, \n                              shuffle=True, \n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=True)\n    valid_loader = DataLoader(valid_dataset, \n                              batch_size=CFG.batch_size * 2, #TODO why\n                              shuffle=False, \n                              num_workers=CFG.num_workers, pin_memory=True, drop_last=False)\n    \n\n    model = VitModel(CFG, pretrained=True)\n    model.to(device)\n    \n    optimizer = Adam(model.parameters(),lr= CFG.lr,weight_decay=CFG.weight_decay,amsgrad=False)\n    scheduler = get_scheduler(optimizer)\n    criterion = nn.CrossEntropyLoss(weight=weights)\n    # criterion = RMSELoss()\n\n    #train loop\n    best_score = -np.inf\n    best_loss = np.inf\n    for epoch in range(CFG.epochs):\n        # train\n        avg_loss = train_fn(fold, train_loader, model, criterion, optimizer, epoch, scheduler, device)\n        \n        # validation\n        avg_val_loss, preds = valid_fn(fold,valid_loader, model, criterion, device)\n\n        y_pred = preds.tolist()\n        y_pred = [pred.index(max(pred)) for pred in y_pred]\n        score = get_accuracy(valid_labels, y_pred)\n        LOGGER.info(f'Epoch {epoch+1} - Score: {score:.4f}')\n        wandb.log({f\"[fold{fold}] epoch\": epoch+1, \n                   f\"[fold{fold}] avg_train_loss\": avg_loss, \n                   f\"[fold{fold}] avg_val_loss\": avg_val_loss,\n                   f\"[fold{fold}] score\": score})\n        if score > best_score:\n            best_score = score\n            LOGGER.info(f'Epoch {epoch+1} - Save Best Score: {best_score:.4f} Model')\n            torch.save({'model': model.state_dict(), \n                        'preds': preds},\n                        OUTPUT_DIR+f'{CFG.model_name}_fold{fold}_best.pth')\n        valid_folds['preds'] = y_pred\n\n    return valid_folds","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:23.305559Z","iopub.execute_input":"2022-10-29T12:12:23.308800Z","iopub.status.idle":"2022-10-29T12:12:23.343107Z","shell.execute_reply.started":"2022-10-29T12:12:23.308761Z","shell.execute_reply":"2022-10-29T12:12:23.342232Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#GPU乗ってるか確認\ndevice","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:23.349221Z","iopub.execute_input":"2022-10-29T12:12:23.351413Z","iopub.status.idle":"2022-10-29T12:12:23.365726Z","shell.execute_reply.started":"2022-10-29T12:12:23.351375Z","shell.execute_reply":"2022-10-29T12:12:23.364538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def main():\n    if CFG.train:\n        # train \n        oof_df = pd.DataFrame()\n        for fold in range(CFG.n_fold):\n            if fold in CFG.trn_fold:\n                _oof_df = train_loop(df, fold)\n                oof_df = pd.concat([oof_df, _oof_df])\n                LOGGER.info(f\"========== fold: {fold} result ==========\")\n                get_result(_oof_df)\n        LOGGER.info(f\"========== CV ==========\")\n        get_result(oof_df)\n        #結果を保存\n        oof_df.to_csv(OUTPUT_DIR+'oof_df.csv', index=False)\n    wandb.finish()\nif __name__ == '__main__':\n    main()","metadata":{"execution":{"iopub.status.busy":"2022-10-29T12:12:23.369303Z","iopub.execute_input":"2022-10-29T12:12:23.372227Z","iopub.status.idle":"2022-10-29T13:20:13.414405Z","shell.execute_reply.started":"2022-10-29T12:12:23.372191Z","shell.execute_reply":"2022-10-29T13:20:13.413310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 検証","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import top_k_accuracy_score\n\ndef get_top_k_accuracies(y_true,y_pred,k):\n    accuracies = []\n    for m in range(k):\n        acc_count = 0\n        for i in range(len(y_true)):\n            if int(y_true[i]) in list(np.argsort(y_pred[i])[::-1][0:m+1]):\n                acc_count += 1\n        accuracies.append(acc_count/len(y_true))\n    \n    return accuracies\n    \ndef get_result_k(df_test,preds_proba_list,k):\n    labels = df_test[CFG.target_col].values\n    k_list = np.arange(1,k+1)\n    accuracies = get_top_k_accuracies(labels, preds_proba_list,k)\n    print('accuracies')\n    print(accuracies)\n    \n    fig = plt.figure(figsize=(10, 30))\n    ax1 = fig.add_subplot(1, 1, 1)\n    ax1.plot(k_list, accuracies, label='accuracies')\n    ax1.legend(loc = 'upper right') #凡例\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-10-29T13:20:13.416566Z","iopub.execute_input":"2022-10-29T13:20:13.416926Z","iopub.status.idle":"2022-10-29T13:20:13.426240Z","shell.execute_reply.started":"2022-10-29T13:20:13.416886Z","shell.execute_reply":"2022-10-29T13:20:13.425163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#testデータの推論→自分で書いたけど上の関数使いまわせたかも\ndef set_image(image):\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    transform = transforms.Compose([\n            transforms.Resize(CFG.size, CFG.size),\n            transforms.Normalize(\n                mean=[0.485, 0.456, 0.406],\n                std=[0.229, 0.224, 0.225],\n            ),ToTensorV2(),\n        ])\n    image_tensor = transform(image=image)\n    image = image_tensor[\"image\"].unsqueeze(0)\n    return image\n\ndef set_model(CFG):\n    model = VitModel(CFG, pretrained=False)\n    state = torch.load(CFG.MODEL_PATH, map_location=torch.device('cpu'))['model']\n    model.load_state_dict(state)\n    model.eval()\n    return model\n\ndf_test['preds']=0\npreds_proba_list =[]\n\nfor i in tqdm(range(df_test.shape[0])):\n    image = cv2.imread(df_test['img_path'][i])\n    image = set_image(image)\n    model = set_model(CFG)\n    result= model(image)\n    preds_proba = nn.functional.softmax(result).detach().numpy()[0].astype(float).tolist()\n    preds_proba_list.append(preds_proba)\n    preds = preds_proba.index(max(preds_proba))\n    df_test['preds'][i] = preds","metadata":{"execution":{"iopub.status.busy":"2022-10-29T13:20:13.427729Z","iopub.execute_input":"2022-10-29T13:20:13.428891Z","iopub.status.idle":"2022-10-29T13:20:15.492841Z","shell.execute_reply.started":"2022-10-29T13:20:13.428852Z","shell.execute_reply":"2022-10-29T13:20:15.490951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#acc@nを可視化\nget_result_k(df_test,preds_proba_list, k=10)","metadata":{"execution":{"iopub.status.busy":"2022-10-29T13:20:15.494267Z","iopub.status.idle":"2022-10-29T13:20:15.494809Z","shell.execute_reply.started":"2022-10-29T13:20:15.494530Z","shell.execute_reply":"2022-10-29T13:20:15.494556Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#混合行列の出力\nfrom sklearn.metrics import confusion_matrix\nimport seaborn as sns\n\ndef show_cm(df_test,preds_proba_list):\n    y_true = df_test[CFG.target_col].values\n    y_pred = df_test['preds'].values\n    cm = confusion_matrix(y_true, y_pred)\n    sns.set(rc = {'figure.figsize':(30,30)})\n    sns.heatmap(cm, annot=True)\n    plt.savefig('/kaggle/working/confusion_matrix_annot.png')\n    return cm","metadata":{"execution":{"iopub.status.busy":"2022-10-29T13:20:15.496370Z","iopub.status.idle":"2022-10-29T13:20:15.497275Z","shell.execute_reply.started":"2022-10-29T13:20:15.496972Z","shell.execute_reply":"2022-10-29T13:20:15.497006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cm = show_cm(df_test,preds_proba_list)","metadata":{"execution":{"iopub.status.busy":"2022-10-29T13:20:15.498966Z","iopub.status.idle":"2022-10-29T13:20:15.499467Z","shell.execute_reply.started":"2022-10-29T13:20:15.499217Z","shell.execute_reply":"2022-10-29T13:20:15.499240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install GPUtil\n\nfrom GPUtil import showUtilization as gpu_usage\ngpu_usage()     ","metadata":{"execution":{"iopub.status.busy":"2022-10-29T13:20:15.500827Z","iopub.status.idle":"2022-10-29T13:20:15.501893Z","shell.execute_reply.started":"2022-10-29T13:20:15.501629Z","shell.execute_reply":"2022-10-29T13:20:15.501657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\ntorch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2022-10-29T13:20:15.510031Z","iopub.status.idle":"2022-10-29T13:20:15.510537Z","shell.execute_reply.started":"2022-10-29T13:20:15.510288Z","shell.execute_reply":"2022-10-29T13:20:15.510312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from numba import cuda\ncuda.select_device(0)\ncuda.close()\ncuda.select_device(0)","metadata":{"execution":{"iopub.status.busy":"2022-10-29T13:20:15.512161Z","iopub.status.idle":"2022-10-29T13:20:15.512642Z","shell.execute_reply.started":"2022-10-29T13:20:15.512392Z","shell.execute_reply":"2022-10-29T13:20:15.512415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}