{"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":"### Memo\n\n- kaggleにはlightningでなくpytorch_lightningが入ってるため、importはpytorch_lightningをしている\n - ただし、書き方はplでなく新しいLでimportする\n- RAdamがPytorch本家に入ったのでtorch_optimizerを入れるコードを削除\n- 不要なimportのチェック\n- https://www.kaggle.com/code/phalanx/train-swin-t-pytorch-lightning\n- https://www.kaggle.com/code/yasufuminakama/fb3-deberta-v3-base-baseline-train#\n- https://www.kaggle.com/code/teyosan1229/petfinder2-nfnet-f3-training/notebook\n- https://www.kaggle.com/code/teyosan1229/paris-1dcnn","metadata":{}},{"cell_type":"markdown","source":"## Get env","metadata":{"papermill":{"duration":0.022464,"end_time":"2021-08-11T02:23:34.652267","exception":false,"start_time":"2021-08-11T02:23:34.629803","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# %env CUBLAS_WORKSPACE_CONFIG=:4096:8","metadata":{"execution":{"iopub.status.busy":"2023-06-23T04:17:38.265279Z","iopub.execute_input":"2023-06-23T04:17:38.265676Z","iopub.status.idle":"2023-06-23T04:17:38.270397Z","shell.execute_reply.started":"2023-06-23T04:17:38.265644Z","shell.execute_reply":"2023-06-23T04:17:38.269091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!nvidia-smi\n# 環境によって処理を変えるためのもの\nimport sys\nIN_COLAB = 'google.colab' in sys.modules\nIN_KAGGLE = 'kaggle_web_client' in sys.modules\nLOCAL = not (IN_KAGGLE or IN_COLAB)\nprint(f'IN_COLAB:{IN_COLAB}, IN_KAGGLE:{IN_KAGGLE}, LOCAL:{LOCAL}')","metadata":{"papermill":{"duration":0.033399,"end_time":"2021-08-11T02:23:35.478019","exception":false,"start_time":"2021-08-11T02:23:35.444620","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:38.276810Z","iopub.execute_input":"2023-06-23T04:17:38.277672Z","iopub.status.idle":"2023-06-23T04:17:39.325943Z","shell.execute_reply.started":"2023-06-23T04:17:38.277629Z","shell.execute_reply":"2023-06-23T04:17:39.324596Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# For Colab Download some datasets\n# ==================\nif IN_COLAB:\n    # mount googledrive\n    from google.colab import drive\n    drive.mount('/content/drive')\n    # copy kaggle.json from googledrive\n    ! pip install --upgrade --force-reinstall --no-deps  kaggle > /dev/null\n    ! mkdir ~/.kaggle\n    ! cp \"/content/drive/MyDrive/kaggle/kaggle.json\" ~/.kaggle/\n    ! chmod 600 ~/.kaggle/kaggle.json\n    \n    # if not os.path.exists(\"/content/input/train_short_audio\"):\n    #     !mkdir input\n    #     !kaggle competitions download -c birdclef-2021\n    #     !unzip /content/birdclef-2021.zip -d input","metadata":{"papermill":{"duration":0.049645,"end_time":"2021-08-11T02:23:35.559175","exception":false,"start_time":"2021-08-11T02:23:35.509530","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:39.329375Z","iopub.execute_input":"2023-06-23T04:17:39.330440Z","iopub.status.idle":"2023-06-23T04:17:39.340747Z","shell.execute_reply.started":"2023-06-23T04:17:39.330400Z","shell.execute_reply":"2023-06-23T04:17:39.339731Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Import Libraries","metadata":{"papermill":{"duration":0.024108,"end_time":"2021-08-11T02:24:01.042079","exception":false,"start_time":"2021-08-11T02:24:01.017971","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# # Image Augmentation Library\n# from albumentations import (\n#     Compose, OneOf, Normalize, Resize, RandomResizedCrop, RandomCrop, HorizontalFlip, VerticalFlip, \n#     RandomBrightness, RandomContrast, RandomBrightnessContrast, Rotate, ShiftScaleRotate, Cutout, \n#     IAAAdditiveGaussianNoise, Transpose\n#     )\n# from albumentations.pytorch import ToTensorV2\n# from albumentations.core.transforms_interface import DualTransform\n# from albumentations.augmentations import functional as AF","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-06-23T04:17:39.343756Z","iopub.execute_input":"2023-06-23T04:17:39.344036Z","iopub.status.idle":"2023-06-23T04:17:39.351040Z","shell.execute_reply.started":"2023-06-23T04:17:39.344007Z","shell.execute_reply":"2023-06-23T04:17:39.350155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Hide Warning\nimport warnings\nwarnings.filterwarnings('ignore', category=DeprecationWarning)\nwarnings.filterwarnings('ignore', category=FutureWarning)\nwarnings.filterwarnings('ignore', category=UserWarning)\n\n# Python Libraries\nimport os\nimport math\nimport random\nimport glob\nimport pickle\nimport gc\nfrom collections import defaultdict\nfrom pathlib import Path\n\n# Third party\nimport numpy as np\nimport pandas as pd\nimport polars as pl\nfrom tqdm.notebook import tqdm\n\n# Visualizations\n# from PIL import Image\n# import cv2\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport plotly.express as px\n%matplotlib inline\nsns.set(style=\"whitegrid\")\n\n# Utilities and Metrics\nfrom sklearn.model_selection import StratifiedKFold\nfrom sklearn.metrics import roc_auc_score, accuracy_score\n\n# Pytorch \nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.optim.lr_scheduler import CosineAnnealingWarmRestarts, CosineAnnealingLR, ReduceLROnPlateau\nfrom torch.optim.optimizer import Optimizer, required\n\nfrom torchvision.io import read_image\nimport torchvision.transforms as T\n\n# Pytorch Lightning 新しいほうは import lightning as L\nimport pytorch_lightning as L\nfrom pytorch_lightning import Callback, seed_everything\nfrom pytorch_lightning.callbacks import LearningRateMonitor, ModelCheckpoint, EarlyStopping\nfrom pytorch_lightning.loggers import WandbLogger, CSVLogger\n\nfrom transformers import get_cosine_schedule_with_warmup\n\n# Pytorch Image Models\nimport timm\n\n# Weights and Biases Tool\nimport wandb\n","metadata":{"execution":{"iopub.status.busy":"2023-06-23T04:17:39.352473Z","iopub.execute_input":"2023-06-23T04:17:39.352901Z","iopub.status.idle":"2023-06-23T04:17:39.371810Z","shell.execute_reply.started":"2023-06-23T04:17:39.352872Z","shell.execute_reply":"2023-06-23T04:17:39.370619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config","metadata":{"papermill":{"duration":0.027595,"end_time":"2021-08-11T02:24:09.304931","exception":false,"start_time":"2021-08-11T02:24:09.277336","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class CFG:\n    debug = False\n    log_wandb = True\n    competition='template'\n    exp_name = \"swin_tiny_224\"\n    seed = [29]\n    # model\n    model_name = 'swin_tiny_patch4_window7_224'\n    pretrained = True\n    img_size = 224\n    in_chans = 3\n    # data\n    target_col = 'label' # 目標値のある列名\n    target_size = 5\n    feature_cols = []\n    # optimizer\n    optimizer_name = 'AdamW'#['RAdam', 'sgd', 'AdamW']\n    lr = 1e-4\n    weight_decay = 1e-5\n    amsgrad = False\n    # scheduler\n    epochs = 10\n#     scheduler = 'CosineAnnealingLR' #['CosineAnnealingLR', 'ReduceLROnPlateau']\n    T_max = 300\n    min_lr = 1e-5\n    scheduler = 'get_cosine_schedule_with_warmup'\n    num_warmup_steps_rate = 0.1 # 総ステップのうち何割をwarm upに使うか\n    num_warmup_steps = 1\n    # criterion\n    criterion_name = 'CrossEntropyLoss'\n    \n    mixup = {'alpha':2}\n    \n    # training\n    train = True\n    save_oof = True\n    save_logits = True\n    n_fold = 5\n    trn_fold = [0]\n    precision = 16 #[16, 32, 64]\n    grad_acc = 2\n    # DataLoader\n    loader = {\n        \"train\": {\n            \"batch_size\": 32,\n            \"num_workers\": 0,\n            \"shuffle\": True,\n            \"pin_memory\": True,\n            \"drop_last\": True\n        },\n        \"valid\": {\n            \"batch_size\": 32,\n            \"num_workers\": 0,\n            \"shuffle\": False,\n            \"pin_memory\": True,\n            \"drop_last\": False\n        }\n    }\n    # pl\n    trainer = {\n#         'gpus': 1,\n        'accelerator': \"auto\",\n        'benchmark': False,\n        'deterministic': True,\n#         'deterministic': False,\n        }\nseed_everything(CFG.seed[0])\nif not LOCAL:\n    CFG.loader[\"train\"][\"num_workers\"] = 4\n    CFG.loader[\"valid\"][\"num_workers\"] = 4","metadata":{"papermill":{"duration":0.034283,"end_time":"2021-08-11T02:24:09.363961","exception":false,"start_time":"2021-08-11T02:24:09.329678","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:39.375763Z","iopub.execute_input":"2023-06-23T04:17:39.376165Z","iopub.status.idle":"2023-06-23T04:17:39.388255Z","shell.execute_reply.started":"2023-06-23T04:17:39.376132Z","shell.execute_reply":"2023-06-23T04:17:39.387327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.log_wandb:\n    if IN_KAGGLE:\n        from kaggle_secrets import UserSecretsClient\n        user_secrets = UserSecretsClient()\n        secret_value_0 = user_secrets.get_secret(\"wandb_api\")\n        wandb.login(key=secret_value_0)\n    elif LOCAL:\n        wandb.login()","metadata":{"execution":{"iopub.status.busy":"2023-06-23T04:17:39.389714Z","iopub.execute_input":"2023-06-23T04:17:39.390096Z","iopub.status.idle":"2023-06-23T04:17:39.918274Z","shell.execute_reply.started":"2023-06-23T04:17:39.390067Z","shell.execute_reply":"2023-06-23T04:17:39.917271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Directory & LoadData","metadata":{"papermill":{"duration":0.026035,"end_time":"2021-08-11T02:24:09.415930","exception":false,"start_time":"2021-08-11T02:24:09.389895","status":"completed"},"tags":[]}},{"cell_type":"code","source":"if IN_KAGGLE:\n    INPUT_DIR = Path('/kaggle/input/cassava-leaf-disease-classification')\n    OUTPUT_DIR = './'\nelif IN_COLAB:\n    INPUT_DIR = Path('/content/input/')\n    OUTPUT_DIR = f'/content/drive/MyDrive/kaggle/BirdClef2021/{CFG.exp_name}/'\nif LOCAL:\n    INPUT_DIR = Path(\"F:/Kaggle/atmaCup11/data/input/\")\n    OUTPUT_DIR = f'F:/Kaggle/atmaCup11/data/output/{CFG.exp_name}/'\n\nTRAIN_DIR = INPUT_DIR / \"train_images\"\nTEST_DIR = INPUT_DIR / \"test_images\"\n\ndf_train = pd.read_csv(INPUT_DIR / \"train.csv\")\ndf_test = pd.read_csv(INPUT_DIR / \"train.csv\")\n\nif not os.path.exists(OUTPUT_DIR):\n    os.makedirs(OUTPUT_DIR)\n\nif CFG.debug:\n    CFG.epochs = 10\n    df_train = df_train.sample(n=1000, random_state=CFG.seed[0]).reset_index(drop=True)","metadata":{"papermill":{"duration":0.114526,"end_time":"2021-08-11T02:24:09.556815","exception":false,"start_time":"2021-08-11T02:24:09.442289","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:39.919633Z","iopub.execute_input":"2023-06-23T04:17:39.919931Z","iopub.status.idle":"2023-06-23T04:17:39.964557Z","shell.execute_reply.started":"2023-06-23T04:17:39.919907Z","shell.execute_reply":"2023-06-23T04:17:39.963619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_filepath(name, folder=TRAIN_DIR):\n#     path = os.path.join(folder, f'{name}.jpg')\n    path = os.path.join(folder, f'{name}')\n    return path\n\ndf_train['image_path'] = df_train['image_id'].apply(lambda x: get_filepath(x, TRAIN_DIR))\ndf_test['image_path'] = df_test['image_id'].apply(lambda x: get_filepath(x, TEST_DIR))\nprint(df_train.shape, df_test.shape)\ndisplay(df_train.head())\nplt.figure(figsize=(4, 4))\nplt.imshow(read_image(df_train.loc[0][\"image_path\"]).permute(1, 2, 0))","metadata":{"papermill":{"duration":0.358407,"end_time":"2021-08-11T02:24:09.939426","exception":false,"start_time":"2021-08-11T02:24:09.581019","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:39.965897Z","iopub.execute_input":"2023-06-23T04:17:39.966260Z","iopub.status.idle":"2023-06-23T04:17:40.501617Z","shell.execute_reply.started":"2023-06-23T04:17:39.966227Z","shell.execute_reply":"2023-06-23T04:17:40.500741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utils","metadata":{"papermill":{"duration":0.026119,"end_time":"2021-08-11T02:24:10.069808","exception":false,"start_time":"2021-08-11T02:24:10.043689","status":"completed"},"tags":[]}},{"cell_type":"code","source":"seed_everything(CFG.seed[0])\n# LINEに通知\nimport requests\ndef send_line_notification(message):\n    env = \"\"\n    if IN_COLAB: env = \"colab\"\n    elif IN_KAGGLE: env = \"kaggle\"\n    elif LOCAL: env = \"local\"\n        \n    line_token = os.getenv('LINE_API_KEY')\n    endpoint = 'https://notify-api.line.me/api/notify'\n    message = f\"[{env}]{message}\"\n    payload = {'message': message}\n    headers = {'Authorization': 'Bearer {}'.format(line_token)}\n    requests.post(endpoint, data=payload, headers=headers)\n\ndef metric(true, pred):\n    \n    \"\"\"コンペの評価指標 CVに使う\"\"\"\n    max_pred = pred\n    \n    p = precision_score(true, max_pred, average='micro')\n    r = recall_score(true, max_pred, average='micro')\n    \n    score = (2 * p * r) / (p + r)\n    return score\n\ndef save_pickle(filename, obj):\n    with open(filename, mode='wb') as f:\n        pickle.dump(obj, f)\n        \ndef load_pickle(filename):\n    with open(filename, mode='rb') as f:\n        p = pickle.load(f)\n    return p ","metadata":{"papermill":{"duration":0.036808,"end_time":"2021-08-11T02:24:10.131566","exception":false,"start_time":"2021-08-11T02:24:10.094758","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:40.502869Z","iopub.execute_input":"2023-06-23T04:17:40.503955Z","iopub.status.idle":"2023-06-23T04:17:40.516886Z","shell.execute_reply.started":"2023-06-23T04:17:40.503922Z","shell.execute_reply":"2023-06-23T04:17:40.515973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CV Split","metadata":{"papermill":{"duration":0.025105,"end_time":"2021-08-11T02:24:10.182188","exception":false,"start_time":"2021-08-11T02:24:10.157083","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# df_train[\"fold\"] = -1\n# Fold = StratifiedKFold(n_splits=CFG.n_fold, shuffle=True, random_state=CFG.seed[0])\n# for n, (train_index, val_index) in enumerate(Fold.split(df_train, df_train[CFG.target_col])):\n#     df_train.loc[val_index, 'fold'] = int(n)\n# df_train['fold'] = df_train['fold'].astype(int)\n# print(df_train.groupby(['fold', CFG.target_col]).size())","metadata":{"papermill":{"duration":0.061644,"end_time":"2021-08-11T02:24:10.269743","exception":false,"start_time":"2021-08-11T02:24:10.208099","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:40.518599Z","iopub.execute_input":"2023-06-23T04:17:40.518947Z","iopub.status.idle":"2023-06-23T04:17:40.529640Z","shell.execute_reply.started":"2023-06-23T04:17:40.518900Z","shell.execute_reply":"2023-06-23T04:17:40.528721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_fold(df, cfg, seed):\n    df_fold = pd.DataFrame([-1]*len(df_train),columns=['fold'])\n    \"\"\"\n    StratifiedKFold\n    https://github.com/teyosan/MLTemplate/blob/master/scripts/fold.py\n    \"\"\"\n    Fold = StratifiedKFold(n_splits=cfg.n_fold, shuffle=True, random_state=seed)\n    for n, (train_index, val_index) in enumerate(Fold.split(df, df[cfg.target_col])):\n        df_fold.loc[val_index, 'fold'] = int(n)\n    return df_fold\nfold = get_fold(df_train, CFG, CFG.seed[0])\nprint(pd.concat([df_train,fold],axis=1).groupby(['fold', CFG.target_col]).size())","metadata":{"execution":{"iopub.status.busy":"2023-06-23T04:17:40.531175Z","iopub.execute_input":"2023-06-23T04:17:40.531684Z","iopub.status.idle":"2023-06-23T04:17:40.550553Z","shell.execute_reply.started":"2023-06-23T04:17:40.531652Z","shell.execute_reply":"2023-06-23T04:17:40.549677Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Transforms","metadata":{"papermill":{"duration":0.042021,"end_time":"2021-08-11T02:24:10.423278","exception":false,"start_time":"2021-08-11T02:24:10.381257","status":"completed"},"tags":[]}},{"cell_type":"code","source":"IMAGENET_MEAN = [0.485, 0.456, 0.406]  # RGB\nIMAGENET_STD = [0.229, 0.224, 0.225]  # RGB\ndef get_transforms():\n    transform = {\n        \"train\": T.Compose(\n            [\n                T.RandomHorizontalFlip(),\n#                 T.RandomVerticalFlip(),\n#                 T.RandomAffine(15, translate=(0.1, 0.1), scale=(0.9, 1.1)),\n#                 T.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.1),\n                T.ConvertImageDtype(torch.float),\n                T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n            ]\n        ),\n        \"val\": T.Compose(\n            [\n                T.ConvertImageDtype(torch.float),\n                T.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n            ]\n        ),\n    }\n    return transform","metadata":{"papermill":{"duration":0.065038,"end_time":"2021-08-11T02:24:10.536528","exception":false,"start_time":"2021-08-11T02:24:10.471490","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:40.552049Z","iopub.execute_input":"2023-06-23T04:17:40.552396Z","iopub.status.idle":"2023-06-23T04:17:40.559289Z","shell.execute_reply.started":"2023-06-23T04:17:40.552366Z","shell.execute_reply":"2023-06-23T04:17:40.557812Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{"papermill":{"duration":0.040811,"end_time":"2021-08-11T02:24:10.618682","exception":false,"start_time":"2021-08-11T02:24:10.577871","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class TrainDataset(Dataset):\n    def __init__(self, df, cfg):\n#         self.df = df\n        self._image_paths = df['image_path'].to_numpy()\n        self._features = df[cfg.feature_cols].to_numpy()\n        self._targets = None\n        if cfg.target_col in df.keys():\n            self._targets = df[cfg.target_col].to_numpy()\n        self._transform = T.Compose([\n                                        T.Resize(cfg.img_size),  # 1\n                                        T.CenterCrop([cfg.img_size, cfg.img_size]),  # 2\n                                    ]\n                                    )\n        \n    def __len__(self):\n        return len(self._image_paths)\n    \n    def __getitem__(self, idx):\n        image_path = self._image_paths[idx]\n        # タスクに合わせたロード方法\n#         image = np.load(file_path).astype(np.float32)# (6, 273, 256)\n#         image = np.vstack(image).transpose((1, 0))# (1638, 256) -> (256, 1638)\n#         image = cv2.imread(file_path)\n#         image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        image = read_image(image_path)\n        image = self._transform(image)\n        features = torch.FloatTensor(self._features[idx, :])\n        \n        if self._targets is not None:\n            label = self._targets[idx]\n            return image, features, label\n        return image, features","metadata":{"papermill":{"duration":0.05568,"end_time":"2021-08-11T02:24:10.715327","exception":false,"start_time":"2021-08-11T02:24:10.659647","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:40.560893Z","iopub.execute_input":"2023-06-23T04:17:40.561247Z","iopub.status.idle":"2023-06-23T04:17:40.572246Z","shell.execute_reply.started":"2023-06-23T04:17:40.561216Z","shell.execute_reply":"2023-06-23T04:17:40.571355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.debug:\n    ds = TrainDataset(df_train, CFG)\n    for i in range(3):\n        print(\"=\"*50)\n        print(ds[0][i])\n    plt.figure(figsize=(3, 3))\n    image, _,_ = ds[0]\n    plt.imshow(image[0])\n    del ds","metadata":{"papermill":{"duration":0.368475,"end_time":"2021-08-11T02:24:11.129097","exception":false,"start_time":"2021-08-11T02:24:10.760622","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:40.577672Z","iopub.execute_input":"2023-06-23T04:17:40.577929Z","iopub.status.idle":"2023-06-23T04:17:40.946077Z","shell.execute_reply.started":"2023-06-23T04:17:40.577907Z","shell.execute_reply":"2023-06-23T04:17:40.945224Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## DataModule","metadata":{"papermill":{"duration":0.032091,"end_time":"2021-08-11T02:24:11.545052","exception":false,"start_time":"2021-08-11T02:24:11.512961","status":"completed"},"tags":[]}},{"cell_type":"code","source":"class DataModule(L.LightningDataModule):\n    def __init__(self, train_data, valid_data, test_data, cfg):\n        super().__init__()\n        self._train_data = train_data\n        self._valid_data = valid_data\n        self._test_data = test_data\n        self._cfg = cfg\n        \n    # 必ず呼び出される関数\n    def setup(self, stage=None):\n        self.train_dataset = TrainDataset(self._train_data, self._cfg)\n        self.valid_dataset = TrainDataset(self._valid_data, self._cfg)\n        self.test_dataset = TrainDataset(self._test_data, self._cfg)\n        \n    # Trainer.fit() 時に呼び出される\n    def train_dataloader(self):\n        return DataLoader(self.train_dataset, **self._cfg.loader['train'])\n\n    # Trainer.fit() 時に呼び出される\n    def val_dataloader(self):\n        return DataLoader(self.valid_dataset, **self._cfg.loader['valid'])\n\n    def test_dataloader(self):\n        return DataLoader(self.test_dataset, **self._cfg.loader['valid'])","metadata":{"papermill":{"duration":0.043501,"end_time":"2021-08-11T02:24:11.619998","exception":false,"start_time":"2021-08-11T02:24:11.576497","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:40.947657Z","iopub.execute_input":"2023-06-23T04:17:40.948357Z","iopub.status.idle":"2023-06-23T04:17:40.957574Z","shell.execute_reply.started":"2023-06-23T04:17:40.948323Z","shell.execute_reply":"2023-06-23T04:17:40.956358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.debug:\n    Data = DataModule(df_train.head(50),df_train.head(50),df_train.head(50), CFG)\n    Data.setup()\n    sample_dataloader = Data.train_dataloader()\n    images, features, labels = next(iter(sample_dataloader))\n    plt.figure(figsize=(6, 6))\n    for it, (image, label) in enumerate(zip(images[:4], labels[:4])):\n        plt.subplot(4, 4, it+1)\n        plt.imshow(image.permute(1, 2, 0))\n        plt.axis('off')\n        plt.title(f'label: {int(label)}')","metadata":{"execution":{"iopub.status.busy":"2023-06-23T04:17:40.960787Z","iopub.execute_input":"2023-06-23T04:17:40.961526Z","iopub.status.idle":"2023-06-23T04:17:42.043245Z","shell.execute_reply.started":"2023-06-23T04:17:40.961495Z","shell.execute_reply":"2023-06-23T04:17:42.042274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Pytorch Lightning Module","metadata":{"papermill":{"duration":0.032014,"end_time":"2021-08-11T02:24:11.683189","exception":false,"start_time":"2021-08-11T02:24:11.651175","status":"completed"},"tags":[]}},{"cell_type":"code","source":"# ====================================================\n# criterion\n# ====================================================\ndef get_criterion(config: dict):\n    if config.criterion_name == 'BCEWithLogitsLoss':\n        criterion = nn.BCEWithLogitsLoss(reduction=\"mean\")\n    if config.criterion_name == 'CrossEntropyLoss':\n        criterion = nn.CrossEntropyLoss()\n    else:\n        raise NotImplementedError\n    return criterion\n# ====================================================\n# optimizer\n# ====================================================\ndef get_optimizer(model: nn.Module, config: dict):\n    \"\"\"\n    input:\n    model:model\n    config:optimizer_nameやlrが入ったものを渡す\n    \n    output:optimizer\n    \"\"\"\n    if 'Adam' == config.optimizer_name:\n        return torch.optim.Adam(model.parameters(),\n                    lr=config.lr,\n                    weight_decay=config.weight_decay,\n                    amsgrad=config.amsgrad)\n    elif 'RAdam' == config.optimizer_name:\n        return torch.optim.RAdam(model.parameters(),\n                           lr=config.lr,\n                           weight_decay=config.weight_decay)\n    elif 'AdamW' == config.optimizer_name:\n        return torch.optim.AdamW(model.parameters(),\n                                 lr=config.lr,\n                                 weight_decay=config.weight_decay)\n#     elif 'Ranger' == config.optimizer_name:\n#         return optim.Ranger(model.parameters(),lr=config.lr)\n    elif 'sgd' == config.optimizer_name:\n        return torch.optim.SGD(model.parameters(),\n                   lr=config.lr,\n                   momentum=0.9,\n                   nesterov=True,\n                   weight_decay=config.weight_decay,)\n    else:\n        raise NotImplementedError\n\n# ====================================================\n# scheduler\n# ====================================================\ndef get_scheduler(cfg, optimizer):\n    if cfg.scheduler=='ReduceLROnPlateau':\n        \"\"\"\n        factor : 学習率の減衰率\n        patience : 何ステップ向上しなければ減衰するかの値\n        eps : nanとかInf回避用の微小数\n        \"\"\"\n        scheduler = ReduceLROnPlateau(optimizer, mode='min', factor=cfg.factor, patience=cfg.patience, verbose=True, eps=cfg.eps, min_lr=cfg.min_lr)\n    elif cfg.scheduler=='CosineAnnealingLR':\n        \"\"\"\n        T_max : 1 半周期のステップサイズ\n        eta_min : 最小学習率(極小値)\n        \"\"\"\n        scheduler = CosineAnnealingLR(optimizer, T_max=cfg.T_max, eta_min=cfg.min_lr, last_epoch=-1)\n    elif cfg.scheduler=='CosineAnnealingWarmRestarts':\n        \"\"\"\n        T_0 : 初期の繰りかえし回数\n        T_mult : サイクルのスケール倍率\n        \"\"\"\n        scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=cfg.T_0, T_mult=1, eta_min=cfg.min_lr, last_epoch=-1)\n    elif cfg.scheduler=='get_cosine_schedule_with_warmup':\n        scheduler = get_cosine_schedule_with_warmup(optimizer,\n                                                    num_warmup_steps=cfg.num_warmup_steps,\n                                                    num_training_steps=cfg.T_max)\n    else:\n        raise NotImplementedError\n    return scheduler\n\ndef get_lightning_scheduler(cfg, optimizer):\n    scheduler = get_scheduler(cfg, optimizer)\n    if cfg.scheduler=='ReduceLROnPlateau':\n        return {'scheduler': scheduler,\n                'monitor': 'val_loss_epoch',\n                'interval': 'epoch',\n                'frequency': 1}\n    else:\n        return {'scheduler': scheduler,\n                'interval': 'step',\n                'frequency': 1}","metadata":{"papermill":{"duration":0.045409,"end_time":"2021-08-11T02:24:11.760269","exception":false,"start_time":"2021-08-11T02:24:11.714860","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:42.045090Z","iopub.execute_input":"2023-06-23T04:17:42.045689Z","iopub.status.idle":"2023-06-23T04:17:42.061944Z","shell.execute_reply.started":"2023-06-23T04:17:42.045650Z","shell.execute_reply":"2023-06-23T04:17:42.060976Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ====================================================\n# model\n# ====================================================\nclass CustomModel(nn.Module):\n    def __init__(self, cfg, pretrained=False):\n        super().__init__()\n        self.cfg = cfg\n        self.backbone = timm.create_model(model_name=self.cfg.model_name,\n                                          pretrained=pretrained,\n                                          in_chans=self.cfg.in_chans,\n                                          num_classes=0)\n        self.dropout1 = nn.Dropout(p=0.5)\n        self.dropout2 = nn.Dropout(p=0.5)\n        self.fc = nn.LazyLinear(self.cfg.target_size)\n        \n    def forward(self, x, features):\n        f = self.backbone(x) # (bs, embedding_size)\n        f = self.dropout1(f)\n        if features.shape[1] != 0:\n            f = torch.cat([f, features],dim=1)\n            f =  self.dropout2(f)\n        out = self.fc(f)\n        return out\n    \ndef get_model(cfg):\n    model = CustomModel(cfg, pretrained=cfg.pretrained)\n    return model","metadata":{"execution":{"iopub.status.busy":"2023-06-23T04:17:42.064988Z","iopub.execute_input":"2023-06-23T04:17:42.065688Z","iopub.status.idle":"2023-06-23T04:17:42.076192Z","shell.execute_reply.started":"2023-06-23T04:17:42.065656Z","shell.execute_reply":"2023-06-23T04:17:42.075159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.debug:\n    # # modelの動作確認\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    model = get_model(CFG).to(device)\n    images = images.to(device)\n    features = features.to(device)\n    labels = labels.to(device)\n    transform = get_transforms()\n    images = transform[\"train\"](images)\n    output = model(images,features)\n    print(output)\n# del Data, model, images, features, labels, output\n# gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-06-23T04:17:42.077715Z","iopub.execute_input":"2023-06-23T04:17:42.078042Z","iopub.status.idle":"2023-06-23T04:17:46.951993Z","shell.execute_reply.started":"2023-06-23T04:17:42.078012Z","shell.execute_reply":"2023-06-23T04:17:46.951010Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if CFG.debug:\n    criterion = get_criterion(CFG)\n    print(criterion(output, labels))","metadata":{"execution":{"iopub.status.busy":"2023-06-23T04:17:46.953412Z","iopub.execute_input":"2023-06-23T04:17:46.953915Z","iopub.status.idle":"2023-06-23T04:17:46.961714Z","shell.execute_reply.started":"2023-06-23T04:17:46.953878Z","shell.execute_reply":"2023-06-23T04:17:46.960598Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#schedulerの確認\nif CFG.debug:\n    CFG.T_max = int(math.ceil(len(df_train)/CFG.grad_acc)*CFG.epochs)\n    CFG.num_warmup_steps = int(CFG.T_max * CFG.num_warmup_steps_rate)\n    model = get_model(CFG)\n    optimizer = get_optimizer(model, CFG)\n    scheduler = get_scheduler(CFG,optimizer)\n    from pylab import rcParams\n    lrs = []\n    for step in range(CFG.T_max):\n        scheduler.step(step)\n        lrs.append(optimizer.param_groups[0][\"lr\"])\n    rcParams['figure.figsize'] = 10,2\n#     print(lrs)\n    plt.plot(lrs)\n    print(max(lrs),min(lrs))","metadata":{"papermill":{"duration":0.468206,"end_time":"2021-08-11T02:24:12.260281","exception":false,"start_time":"2021-08-11T02:24:11.792075","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:46.963508Z","iopub.execute_input":"2023-06-23T04:17:46.963841Z","iopub.status.idle":"2023-06-23T04:17:47.921902Z","shell.execute_reply.started":"2023-06-23T04:17:46.963810Z","shell.execute_reply":"2023-06-23T04:17:47.920994Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Trainer(L.LightningModule):\n    def __init__(self, cfg):\n        super().__init__()\n        self.cfg = cfg\n        print(cfg.model_name)\n        self.model = get_model(cfg)\n        self.criterion = get_criterion(cfg)\n        self.transform = get_transforms()\n        \n        self.training_step_outputs = []\n        self.validation_step_outputs = []\n    \n    def forward(self, x, features):\n        output = self.model(x, features)\n        return output\n    \n    # loss のreturnの必要性が分かってない\n    def training_step(self, batch, batch_idx):\n        loss, pred, labels = self.__share_step(batch, 'train')\n        self.training_step_outputs.append({'loss': loss, 'pred': pred, 'labels': labels})\n        return loss # これ返さないとbackproできない？\n#         return {'loss': loss, 'pred': pred, 'labels': labels}\n    \n    def validation_step(self, batch, batch_idx):\n        loss, pred, labels = self.__share_step(batch, 'val')\n        self.validation_step_outputs.append({'pred': pred, 'labels': labels, 'loss': loss.item()})\n#         return {'pred': pred, 'labels': labels}\n    \n    def predict_step(self, batch, batch_idx):\n        if len(batch) == 3:\n            images, features, _ = batch\n        else:\n            images, features = batch\n        images = self.transform['val'](images)\n        logits = self.forward(images, features)\n        pred = logits\n        return pred\n    \n    def __share_step(self, batch, mode):\n        if self.global_step == 0: # 複数回呼び出されてそうだけどそれでもちゃんと動いてそう・・・\n            wandb.define_metric(f'val_loss_epoch', summary='min')\n        images, features, labels = batch\n        images = self.transform[mode](images)\n        # mixup とかしたい場合はここに差し込む\n        logits = self.forward(images, features)\n        loss = self.criterion(logits, labels)\n        pred = logits\n        return loss, pred, labels\n    \n#     def training_epoch_end(self, outputs):\n    def on_train_epoch_end(self):\n        self.__share_epoch_end(self.training_step_outputs, 'train')\n        self.training_step_outputs.clear()  # free memory\n        # なくてもlrloggerが動くけどちゃんと記録されない・・・\n        self.log(\"lr\", self.optimizer.param_groups[0]['lr'], prog_bar=True, logger=True)\n\n#     def validation_epoch_end(self, outputs):\n    def on_validation_epoch_end(self):\n        self.__share_epoch_end(self.validation_step_outputs, 'val')\n        self.validation_step_outputs.clear()  # free memory\n        \n    def __share_epoch_end(self, outputs, mode):\n        \n        loss = sum([o[\"loss\"] for o in outputs])/len(outputs)\n        self.log(f\"{mode}_loss_epoch\", loss)\n#         preds = torch.cat([o[\"pred\"] for o in outputs])\n#         labels = torch.cat([o[\"labels\"] for o in outputs])\n#         score = metric(labels, preds)\n#         self.log(f\"{mode}_score_epoch\", score)\n    \n    def configure_optimizers(self):\n        self.optimizer = get_optimizer(self, self.cfg)\n        self.scheduler = get_lightning_scheduler(self.cfg, self.optimizer)\n        return {'optimizer': self.optimizer, 'lr_scheduler': self.scheduler}\n# tr = Trainer(CFG)","metadata":{"papermill":{"duration":0.050613,"end_time":"2021-08-11T02:24:12.419225","exception":false,"start_time":"2021-08-11T02:24:12.368612","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:47.923488Z","iopub.execute_input":"2023-06-23T04:17:47.924086Z","iopub.status.idle":"2023-06-23T04:17:47.940281Z","shell.execute_reply.started":"2023-06-23T04:17:47.924051Z","shell.execute_reply":"2023-06-23T04:17:47.939357Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{"papermill":{"duration":0.034187,"end_time":"2021-08-11T02:24:12.486943","exception":false,"start_time":"2021-08-11T02:24:12.452756","status":"completed"},"tags":[]}},{"cell_type":"code","source":"def train(seed:int) -> None:\n    seed_everything(CFG.seed[0])\n    df_F = get_fold(df_train, CFG, seed)\n    df_oof = pd.DataFrame([None]*len(df_train),columns=['pred'])\n    arr_logits = np.full((len(df_train), CFG.target_size), None)\n    for fold in range(CFG.n_fold):\n        if not fold in CFG.trn_fold:\n            continue\n        print(f\"{'='*38} Fold: {fold} {'='*38}\")\n        \n        # callbacks\n        #======================================================\n        lr_monitor = LearningRateMonitor(logging_interval='step')\n        earystopping = EarlyStopping(monitor='val_loss_epoch', mode=\"min\", patience=20)\n        # 学習済重みを保存するために必要\n        loss_checkpoint = ModelCheckpoint(\n            dirpath=OUTPUT_DIR,\n            filename=f\"best_loss_seed{seed}_fold{fold}\",\n            monitor=\"val_loss_epoch\",\n            save_last=True,\n            save_top_k=1,\n            save_weights_only=True,\n            mode=\"min\",\n        )\n        # Logger\n        #======================================================\n        csv_logger = CSVLogger(save_dir=str(OUTPUT_DIR), name=f\"seed{seed}_fold{fold}\")\n        wandb_logger = WandbLogger(\n            project=f'{CFG.competition}',\n            group= f'{CFG.exp_name}',\n            name = f'seed{seed}_fold{fold}',\n            save_dir=OUTPUT_DIR\n        )\n        \n        data_module = DataModule(\n          df_train[df_F['fold']!=fold],\n          df_train[df_F['fold']==fold], \n          df_train[df_F['fold']==fold], \n          CFG\n        )\n        data_module.setup()\n        \n        # setting step param\n        # ===================================================================================\n        CFG.T_max = int(math.ceil(len(data_module.train_dataloader())/CFG.grad_acc)*CFG.epochs)\n        CFG.num_warmup_steps = int(CFG.T_max * CFG.num_warmup_steps_rate)\n        print(f\"set schedular T_max {CFG.T_max}\")\n        \n        trainer = L.Trainer(\n            logger=[wandb_logger,csv_logger],\n            callbacks=[lr_monitor,loss_checkpoint,earystopping],\n            default_root_dir=OUTPUT_DIR,\n            accumulate_grad_batches=CFG.grad_acc,\n            max_epochs=CFG.epochs,\n            precision=CFG.precision,\n            **CFG.trainer\n        )\n        \n        # Train\n        # ================================\n        model = Trainer(CFG)\n        trainer.fit(model, data_module)\n        \n        best_model = Trainer.load_from_checkpoint(cfg=CFG,checkpoint_path=loss_checkpoint.best_model_path)\n        torch.save(best_model.model.state_dict(),OUTPUT_DIR + '/' + f'{CFG.exp_name}_seed{seed}_fold{fold}_best.pth')\n        wandb.finish()\n        \n        # OOF\n        # ================================\n        if CFG.save_oof:\n            logits = inference(data_module, OUTPUT_DIR  + f'{CFG.exp_name}_seed{seed}_fold{fold}_best.pth')\n            df_oof.loc[df_F[\"fold\"] == fold, ['pred']] = logits.argmax(axis=1).reshape(-1, 1)\n            if CFG.save_logits:\n                arr_logits[df_F[\"fold\"] == fold] = logits\n    if CFG.save_oof:\n        if CFG.save_logits:\n            df_logits = pd.DataFrame(arr_logits)\n            df_oof = pd.concat([df_logits,df_oof],axis=1)\n        df_oof.to_csv(OUTPUT_DIR + f'oof_{seed}.csv',index=False)\n        \ndef inference(data_module, weight_pass):\n    trainer = L.Trainer(\n            default_root_dir=OUTPUT_DIR,\n            accumulate_grad_batches=CFG.grad_acc,\n            max_epochs=CFG.epochs,\n            precision=CFG.precision,\n            **CFG.trainer\n        )\n    model = Trainer(CFG)\n    model.model.load_state_dict(torch.load(weight_pass))\n    predictions = trainer.predict(model, data_module.test_dataloader())\n    preds= []\n    for p in predictions:\n        preds += p\n    return torch.stack(preds).to('cpu').detach().numpy()","metadata":{"papermill":{"duration":2677.689596,"end_time":"2021-08-11T03:08:50.209818","exception":false,"start_time":"2021-08-11T02:24:12.520222","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:47.941844Z","iopub.execute_input":"2023-06-23T04:17:47.942235Z","iopub.status.idle":"2023-06-23T04:17:47.962821Z","shell.execute_reply.started":"2023-06-23T04:17:47.942205Z","shell.execute_reply":"2023-06-23T04:17:47.961947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for seed in CFG.seed:\n    train(seed)\n    send_line_notification(f\"seed:{seed}[env]finished\")\nwandb.finish()","metadata":{"papermill":{"duration":0.775172,"end_time":"2021-08-11T03:08:51.027159","exception":false,"start_time":"2021-08-11T03:08:50.251987","status":"completed"},"scrolled":true,"tags":[],"execution":{"iopub.status.busy":"2023-06-23T04:17:47.965715Z","iopub.execute_input":"2023-06-23T04:17:47.966053Z","iopub.status.idle":"2023-06-23T04:19:41.060933Z","shell.execute_reply.started":"2023-06-23T04:17:47.966021Z","shell.execute_reply":"2023-06-23T04:19:41.059805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# df_F = get_fold(df_train, CFG, 29)\n# df_oof = pd.DataFrame([None]*len(df_train),columns=['pred'])\n# fold = 0\n# data_module = DataModule(\n#           df_train[df_F['fold']!=fold],\n#           df_train[df_F['fold']==fold], \n#           df_train[df_F['fold']==fold], \n#           CFG\n#         )\n# data_module.setup()\n# logits = inference(data_module,\"/kaggle/working/kaggle_test_seed29_fold0_best.pth\")","metadata":{"execution":{"iopub.status.busy":"2023-06-23T04:19:41.062847Z","iopub.execute_input":"2023-06-23T04:19:41.063573Z","iopub.status.idle":"2023-06-23T04:19:41.069318Z","shell.execute_reply.started":"2023-06-23T04:19:41.063529Z","shell.execute_reply":"2023-06-23T04:19:41.067446Z"},"trusted":true},"execution_count":null,"outputs":[]}]}