{"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":"# Train Code","metadata":{}},{"cell_type":"markdown","source":"## 実行環境のセットアップ\nコードの実行に必要な環境のセットアップ","metadata":{}},{"cell_type":"markdown","source":"### Settings\n機械学習に必要な各種パラメータを設定","metadata":{}},{"cell_type":"code","source":"from pytorch_lightning.callbacks import ModelCheckpoint","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:06:25.606429Z","iopub.execute_input":"2022-09-21T05:06:25.606911Z","iopub.status.idle":"2022-09-21T05:06:29.246104Z","shell.execute_reply.started":"2022-09-21T05:06:25.606856Z","shell.execute_reply":"2022-09-21T05:06:29.244934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install python-gdcm\n!pip install pylibjpeg pylibjpeg-libjpeg pydicom\n!pip install -qU \"python-gdcm\" pydicom pylibjpeg \"opencv-python-headless\"\n!pip install monai","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:06:29.248272Z","iopub.execute_input":"2022-09-21T05:06:29.249025Z","iopub.status.idle":"2022-09-21T05:07:10.571023Z","shell.execute_reply.started":"2022-09-21T05:06:29.248991Z","shell.execute_reply":"2022-09-21T05:07:10.56969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Common Settings\n\n# 学習者氏名\nTRAINER = \"yamaday\"\n# エポック数\nNUM_EPOCH = 10\nCLASS_NUM = 13\n# GCP のプロジェクトID\nGCP_PROJECT_ID = \"ai-facade\"\n# GCS のバケット名\nGCS_BUCKET_NAME = \"ai-facade-competitions\"\n\n# コンペのプロジェクトID\nCEREBRUM_PROJECT_ID = \"202208_cervical_spine_fracture_detection\"","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:10.573371Z","iopub.execute_input":"2022-09-21T05:07:10.573834Z","iopub.status.idle":"2022-09-21T05:07:10.582465Z","shell.execute_reply.started":"2022-09-21T05:07:10.573784Z","shell.execute_reply":"2022-09-21T05:07:10.581453Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# DEV RUN\n\n# Trainer の開発モード指定\n#   - False: Trainer を本番モードで運用\n#   - True : Trainer を開発モードで運用（1ステップで終了）\n#   - (int): Trainer を開発モードで運用（任意のステップ数を実行して終了）\nFAST_DEV_RUN = False\n\n# 開発用のデータ削減\n#   - False: データ数を削減しない\n#   - (int): 指定した数までデータ数を削減\nDEV_DATA_SIZE = False\n\n# for debug: (https://qiita.com/makopo/items/170c939c79dcc5c89e12)\n# from IPython.core.debugger import Pdb; Pdb().set_trace()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:10.586034Z","iopub.execute_input":"2022-09-21T05:07:10.586479Z","iopub.status.idle":"2022-09-21T05:07:10.593404Z","shell.execute_reply.started":"2022-09-21T05:07:10.586444Z","shell.execute_reply":"2022-09-21T05:07:10.592319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training Settings\n\nMODEL_OPTION = {\n    \"model\": {\n        \"name\": \"baseline/segmentation/Unet\",\n        \"params\": {\n            \"base.model\": \"Unet\",\n            \"fc.model\": \"\"\n        }\n    },\n    \"optimizer\": {\n        \"name\": \"AdamW\",\n        \"params\": {\n            \"model.lr\": 2e-3,\n            \"weight_decay\": 1e-6,\n            \"eps\": 1e-6,\n        }\n    },\n    \"scheduler\": {\n        \"name\": \"transformer_cosine\",\n        \"params\": {\n            \"num_warmup_steps\": 0,\n            \"num_cycles\": 0.5\n        }\n    },\n    \"criterion\": {\n        \"name\": [\"BCEWithLogitsLoss\"],\n        \"params\": {\n            \"reduction\": \"none\"\n        }\n    }\n}\n\nDATASET_OPTION = {\n    \"source\": \"original\",\n    \"train_test_split\": [],\n    \"train_valid_split\": [],\n}\n\nTRAINING_OPTION = {\n    \"batch_size\": 4,\n    \"custom\": {\n        \"dev_run\": FAST_DEV_RUN,\n        \"dev_data_size\": DEV_DATA_SIZE,\n        \"n_split\": 5,\n        \"target_fold\": [0],\n#         \"target_fold\": list(range(0, 5)),\n        \"fold_method\": \"StratifiedGroupKFold\",\n        \"early_stopping\": False,\n        \"save_each_model\": False,\n        \"save_best_loss_model\": True,\n        \"save_best_score_model\": True,\n        \"auto_scale_batch_size\": False,\n        \"auto_lr_find\": True,\n        \"image_size\": 224,\n        \"transforms\": True,\n    }\n}\n\nENVIRONMENT_OPTION = {\n    \"random_seed\": 42,\n    \"custom\": {\n        \"apex\": False\n    }\n}","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:10.595251Z","iopub.execute_input":"2022-09-21T05:07:10.595659Z","iopub.status.idle":"2022-09-21T05:07:10.606294Z","shell.execute_reply.started":"2022-09-21T05:07:10.595625Z","shell.execute_reply":"2022-09-21T05:07:10.605034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Setup\nライブラリのインポート、グローバル変数の定義","metadata":{}},{"cell_type":"code","source":"!pip uninstall -y pytorch-lightning\n!pip install pytorch-lightning==1.6.4\n!pip install segmentation-models-pytorch==0.2.1\n\nimport os\nimport sys\nimport gc\nimport itertools\nimport platform\nimport datetime\nimport logging\nimport ast\nimport zipfile\nimport warnings\nimport pandas as pd\nimport numpy as np\nfrom copy import deepcopy\nfrom pytz import timezone, utc\nfrom glob import glob\nfrom joblib import Parallel, delayed\nfrom tqdm.notebook import tqdm\nfrom PIL import Image\nimport cv2\nimport random\nfrom tqdm import tqdm\ntqdm.pandas()\nimport pydicom\nfrom matplotlib import pyplot as plt\nimport re\n\n# sklearn\nfrom sklearn.model_selection import StratifiedKFold, KFold, StratifiedGroupKFold, GroupKFold\nfrom sklearn.metrics import f1_score\nfrom scipy.spatial.distance import directed_hausdorff\n\n# segmentation-models-pytorch\nimport segmentation_models_pytorch as smp\nimport monai\nfrom monai.data import CSVDataset\nfrom monai.data import DataLoader\nimport nibabel as nib\n\n# torch\nimport torch\nimport torch.nn as nn\nfrom torch.optim import AdamW\nfrom torch.utils.data import Dataset, DataLoader\n\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Rectangle\nfrom sklearn import preprocessing\n\nimport albumentations as A\nimport torch.nn.functional as F\n\nimport segmentation_models_pytorch as smp\n\nfrom transformers import (\n    get_linear_schedule_with_warmup,\n    get_cosine_schedule_with_warmup,\n    get_cosine_with_hard_restarts_schedule_with_warmup\n)\n\n# pytorch-lightning\nimport pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import Callback, ModelCheckpoint\nfrom pytorch_lightning.callbacks.early_stopping import EarlyStopping\n\n# google cloud plotform\nfrom google.cloud import storage\n\nwarnings.simplefilter('ignore')\n\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\n\nprint(\"-\" * 20)\nprint(f\"pytorch_lightning version: {pl.__version__}\")\nprint(f\"smp version: {pl.__version__}\")","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:10.607966Z","iopub.execute_input":"2022-09-21T05:07:10.608383Z","iopub.status.idle":"2022-09-21T05:07:43.30243Z","shell.execute_reply.started":"2022-09-21T05:07:10.60835Z","shell.execute_reply":"2022-09-21T05:07:43.301305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# task\nTASK = \"semantic-segmentation-3D\"\n\n# idf path\nLOCAL_IDF_PATH = \"models/{model_id}/fold{fold}.idf.json\"\nGCS_IDF_PATH = \"models/{task}/{model_name}/{model_id}{dev_mode}/fold{fold}.idf.json\"\n\n# model path\nLOCAL_MODEL_PATH = \"models/{model_id}/fold-{fold}/generation-{generation}.ckpt\"\nGCS_MODEL_PATH = \"models/{task}/{model_name}/{model_id}{dev_mode}/fold-{fold}/generation-{generation}.ckpt\"\n\nLOCAL_BEST_MODEL_PATH = \"models/{model_id}/fold-{fold}/\"\nGCS_BEST_MODEL_PATH = \"models/{task}/{model_name}/{model_id}{dev_mode}/fold-{fold}/{filename}\"","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:43.304234Z","iopub.execute_input":"2022-09-21T05:07:43.304852Z","iopub.status.idle":"2022-09-21T05:07:43.313696Z","shell.execute_reply.started":"2022-09-21T05:07:43.304809Z","shell.execute_reply":"2022-09-21T05:07:43.311918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# logger\n\ndef cvt_jst_time(*args):\n    utc_dt = utc.localize(datetime.datetime.utcnow())\n    jst_tz = timezone(\"Asia/Tokyo\")\n    converted = utc_dt.astimezone(jst_tz)\n    return converted.timetuple()\n\n\ndef get_logger():\n    formatter = logging.Formatter(\n        fmt=\"%(asctime)s [%(levelname)s] %(message)s\",\n        datefmt=\"%Y-%m-%d %H:%M:%S\"\n    )\n    formatter.converter = cvt_jst_time\n    \n    logger = logging.getLogger(__name__)\n    logger.setLevel(logging.INFO)\n    \n    handler = logging.StreamHandler(sys.stdout)\n    handler.setLevel(logging.INFO)\n    handler.setFormatter(formatter)\n    \n    logger.addHandler(handler)\n    \n    return logger\n\n\ntry:\n    logger\nexcept NameError:\n    logger = get_logger()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:43.317392Z","iopub.execute_input":"2022-09-21T05:07:43.318341Z","iopub.status.idle":"2022-09-21T05:07:43.361935Z","shell.execute_reply.started":"2022-09-21T05:07:43.318302Z","shell.execute_reply":"2022-09-21T05:07:43.360892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Import Cerebrum Connector\n\nCC_PACKAGE_NAME = \"CerebrumConnector.zip\"\n\ndef download_cc_package():\n    logger.info(\"start: download cc package\")\n    gcs_client = storage.Client(project=GCP_PROJECT_ID)\n    bucket = gcs_client.bucket(GCS_BUCKET_NAME)\n    blob = bucket.blob(CC_PACKAGE_NAME)\n    blob.download_to_filename(CC_PACKAGE_NAME)\n    logger.info(\"complete: download cc package\")\n\n\ndef unzip_package():\n    logger.info(\"start: unzip cc package\")\n    with zipfile.ZipFile(CC_PACKAGE_NAME) as zf:\n        zf.extractall()\n    logger.info(\"complete: unzip cc package\")\n\n\nif not os.path.exists(CC_PACKAGE_NAME):\n    download_cc_package()\n    unzip_package()\n\nfrom CerebrumConnector import CerebrumConnector as cc\nfrom CerebrumConnector import GoogleCloudPlatform as cc_gcp","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:43.363821Z","iopub.execute_input":"2022-09-21T05:07:43.364228Z","iopub.status.idle":"2022-09-21T05:07:43.581036Z","shell.execute_reply.started":"2022-09-21T05:07:43.36419Z","shell.execute_reply":"2022-09-21T05:07:43.579838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Setup Cerebrum Connector\n\ndef get_identification():\n    idf = cc.Identification(\n        cerebrum_project_id=CEREBRUM_PROJECT_ID,\n        trainer=TRAINER,\n        task=TASK,\n        model_name=MODEL_OPTION[\"model\"][\"name\"]\n    )\n    return idf\n\n\ndef get_gcs_connector():\n    gcs_connector = cc_gcp.GCSConnector(\n        gcp_project_name=GCP_PROJECT_ID,\n        bucket_name=GCS_BUCKET_NAME,\n        sub_dir=CEREBRUM_PROJECT_ID\n    )\n    return gcs_connector\n\n\ndef set_options_to_identification(idf):\n    model_option = cc.ModelOption(\n        model=cc.ModelOptionItem(**MODEL_OPTION[\"model\"]),\n        optimizer=cc.ModelOptionItem(**MODEL_OPTION[\"optimizer\"]),\n        scheduler=cc.ModelOptionItem(**MODEL_OPTION[\"scheduler\"]),\n        criterion=cc.ModelOptionItem(**MODEL_OPTION[\"criterion\"])\n    )\n    dataset_option = cc.DatasetOption(**DATASET_OPTION)\n    training_option = cc.TrainingOption(**TRAINING_OPTION)\n    environment_option = cc.EnvironmentOption(**ENVIRONMENT_OPTION)\n    \n    idf.set_options(\n        model_option,\n        dataset_option,\n        training_option,\n        environment_option\n    )\n\n\nidf = get_identification()\nset_options_to_identification(idf)\n\ngcs_connector = get_gcs_connector()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:43.585537Z","iopub.execute_input":"2022-09-21T05:07:43.5859Z","iopub.status.idle":"2022-09-21T05:07:43.59554Z","shell.execute_reply.started":"2022-09-21T05:07:43.585853Z","shell.execute_reply":"2022-09-21T05:07:43.594334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# set random seed\n\ndef set_seed(seed = 42):\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    os.environ['PYTHONHASHSEED'] = str(seed)\n\nset_seed(idf.environment_option.random_seed)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:43.597959Z","iopub.execute_input":"2022-09-21T05:07:43.598651Z","iopub.status.idle":"2022-09-21T05:07:43.611771Z","shell.execute_reply.started":"2022-09-21T05:07:43.598613Z","shell.execute_reply":"2022-09-21T05:07:43.610918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preprocessing\n機械学習の前処理","metadata":{}},{"cell_type":"code","source":"seg_image_paths = glob(os.path.join('../input/rsna-2022-cervical-spine-fracture-detection/segmentations', \"*\"))\n\nseg_image_uids = pd.DataFrame([f\"{os.path.splitext(os.path.basename(path))[0]}\" for path in seg_image_paths]\n, columns=[\"StudyInstanceUID\"])\n\n_train_df = pd.read_csv('../input/rsna-2022-cervical-spine-fracture-detection/train.csv')\ntrain_bounding_boxes_df = pd.read_csv('../input/rsna-2022-cervical-spine-fracture-detection/train_bounding_boxes.csv')\n_train_df = _train_df.merge(seg_image_uids, on=['StudyInstanceUID'], how='inner')","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:43.615174Z","iopub.execute_input":"2022-09-21T05:07:43.615483Z","iopub.status.idle":"2022-09-21T05:07:43.684672Z","shell.execute_reply.started":"2022-09-21T05:07:43.615454Z","shell.execute_reply":"2022-09-21T05:07:43.683738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_train_df","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:43.686364Z","iopub.execute_input":"2022-09-21T05:07:43.686731Z","iopub.status.idle":"2022-09-21T05:07:43.709524Z","shell.execute_reply.started":"2022-09-21T05:07:43.686696Z","shell.execute_reply":"2022-09-21T05:07:43.708567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_path(row):\n    image_paths = glob(os.path.join(f'../input/rsna-2022-cervical-spine-fracture-detection/train_images/{row[\"StudyInstanceUID\"]}', \"*.dcm\"))\n    p = re.compile(r'(\\d+).dcm')\n    image_paths.sort(key=lambda s: int(p.search(s).groups()[0]))\n    \n    if len(image_paths) < idf.training_option.custom['image_size']:\n        ins = [\"\" for i in range(idf.training_option.custom['image_size'] - len(image_paths))]\n        image_paths.extend(ins)\n    row['image_paths'] = image_paths[:idf.training_option.custom['image_size']]\n    row['mask_path'] = f\"../input/rsna-2022-cervical-spine-fracture-detection/segmentations/{row['StudyInstanceUID']}.nii\"\n    row['image_length'] = len(image_paths)\n    return row\n\n_train_df = _train_df.progress_apply(get_path, axis=1)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:43.711702Z","iopub.execute_input":"2022-09-21T05:07:43.712709Z","iopub.status.idle":"2022-09-21T05:07:45.779016Z","shell.execute_reply.started":"2022-09-21T05:07:43.712668Z","shell.execute_reply":"2022-09-21T05:07:45.778107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_train_df['image_length'].hist()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:45.783023Z","iopub.execute_input":"2022-09-21T05:07:45.78548Z","iopub.status.idle":"2022-09-21T05:07:46.101761Z","shell.execute_reply.started":"2022-09-21T05:07:45.785436Z","shell.execute_reply":"2022-09-21T05:07:46.100294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"_train_df[_train_df[\"StudyInstanceUID\"] == \"1.2.826.0.1.3680043.21321\"][\"image_paths\"]","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:46.10309Z","iopub.execute_input":"2022-09-21T05:07:46.103688Z","iopub.status.idle":"2022-09-21T05:07:46.123019Z","shell.execute_reply.started":"2022-09-21T05:07:46.103645Z","shell.execute_reply":"2022-09-21T05:07:46.121745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mm = preprocessing.MinMaxScaler()\n\ndef read_dcm_file(file_path):\n    image_size = (idf.training_option.custom['image_size'], idf.training_option.custom['image_size'])\n    \n    if not file_path:\n        img = np.zeros(image_size, dtype=np.float32)\n    else:\n        img =  pydicom.read_file(file_path).pixel_array\n        img =  mm.fit_transform(img) * 255.0 \n        img = cv2.resize(img, image_size)\n    img = np.tile(img[...,None], [1, 1, 1]) \n    img = img / 255\n\n    return img\n\ndef read_mask(file_path):\n    image_size = (idf.training_option.custom['image_size'], idf.training_option.custom['image_size'])\n    nii = nib.load(file_path)\n    nii_arrays = nii.get_fdata()  # convert to numpy array\n    \n    mask = nii_arrays[:, :, :idf.training_option.custom['image_size']]\n    if mask.shape[2] < idf.training_option.custom['image_size']:\n        ins = np.zeros((mask.shape[0], mask.shape[1], idf.training_option.custom['image_size'] - mask.shape[2]), dtype=np.float32)\n        mask = np.concatenate([mask, ins], axis=2)\n\n    mask = mask / CLASS_NUM\n    mask = cv2.resize(mask, image_size).transpose(2, 1, 0)\n    mask = np.tile(mask[...,None], [1, 1, 1]) # gray to rgb\n\n    return mask","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:46.124693Z","iopub.execute_input":"2022-09-21T05:07:46.125656Z","iopub.status.idle":"2022-09-21T05:07:46.137906Z","shell.execute_reply.started":"2022-09-21T05:07:46.125609Z","shell.execute_reply":"2022-09-21T05:07:46.136938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"p = read_dcm_file(\"../input/rsna-2022-cervical-spine-fracture-detection/train_images/1.2.826.0.1.3680043.20928/80.dcm\")","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:46.139523Z","iopub.execute_input":"2022-09-21T05:07:46.139942Z","iopub.status.idle":"2022-09-21T05:07:46.193174Z","shell.execute_reply.started":"2022-09-21T05:07:46.139906Z","shell.execute_reply":"2022-09-21T05:07:46.192174Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs_ex = []\nfor img_path in _train_df.iloc[0]['image_paths']:\n    imgs_ex.append(read_dcm_file(img_path))\n\nmask_path = _train_df.iloc[0]['mask_path']\nmasks = read_mask(mask_path)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-09-21T05:07:46.194495Z","iopub.execute_input":"2022-09-21T05:07:46.194817Z","iopub.status.idle":"2022-09-21T05:07:50.487471Z","shell.execute_reply.started":"2022-09-21T05:07:46.194784Z","shell.execute_reply":"2022-09-21T05:07:50.486179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs_ex[1].shape","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:50.489196Z","iopub.execute_input":"2022-09-21T05:07:50.489555Z","iopub.status.idle":"2022-09-21T05:07:50.499062Z","shell.execute_reply.started":"2022-09-21T05:07:50.489517Z","shell.execute_reply":"2022-09-21T05:07:50.497935Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"masks.shape","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:50.500435Z","iopub.execute_input":"2022-09-21T05:07:50.501206Z","iopub.status.idle":"2022-09-21T05:07:50.509838Z","shell.execute_reply.started":"2022-09-21T05:07:50.501167Z","shell.execute_reply":"2022-09-21T05:07:50.508724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training\n機械学習を実行","metadata":{}},{"cell_type":"markdown","source":"### データを分割","metadata":{}},{"cell_type":"code","source":"# split dataset (k-fold)\n\ndef k_fold(idf, raw_df):\n    df = raw_df.copy()\n    n_splits = idf.training_option.custom[\"n_split\"]\n    if idf.training_option.custom[\"fold_method\"] == \"StratifiedGroupKFold\":\n        skf = StratifiedGroupKFold(n_splits=n_splits, shuffle=True)\n        for fold, (_, val_idx) in enumerate(skf.split(X=df, y=df['patient_overall'], groups=df['StudyInstanceUID'])):\n            df.loc[val_idx, 'fold'] = int(fold)\n    df['fold'] = df['fold'].astype(np.uint8)\n    return df\n\n# 開発用のデータ削減\nif DEV_DATA_SIZE:\n    logger.info(f\"DEV_MODE: data size: {len(_train_df)} -> {DEV_DATA_SIZE}\")\n    _train_df = _train_df[:DEV_DATA_SIZE]\n    \n_train_df = k_fold(idf, _train_df)\nlogger.info(\"split dataset. size: {}\".format(_train_df.groupby(\"fold\").size().to_dict()))\n\ngc.collect()\n","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:50.511162Z","iopub.execute_input":"2022-09-21T05:07:50.512249Z","iopub.status.idle":"2022-09-21T05:07:50.874054Z","shell.execute_reply.started":"2022-09-21T05:07:50.512208Z","shell.execute_reply":"2022-09-21T05:07:50.872903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataset","metadata":{}},{"cell_type":"code","source":"from sklearn import preprocessing\nmm = preprocessing.MinMaxScaler()\n\nclass CustomDataset(Dataset):\n    def __init__(self, idf, df, transforms=None):\n        self.idf = idf\n        self.df = df\n        self.img_paths  = df['image_paths'].tolist()\n        self.mask_paths  = df['mask_path'].tolist()\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_paths  = self.img_paths[index]\n        imgs = []\n        for img_path in img_paths:\n            imgs.append(self._read_dcm_file(img_path))\n        \n        mask_path = self.mask_paths[index]\n        masks = self._read_mask(mask_path)\n        \n        if self.transforms is not None:\n            imgs = np.array(imgs, dtype=np.float32)\n            aug = self.transforms(image=imgs, mask=masks)\n            imgs = aug['image']\n            masks = aug['mask']\n        \n        imgs = np.array(imgs, dtype=np.float32)\n        imgs = imgs.transpose(3, 1, 2, 0)\n        masks = masks.transpose(3, 1, 2, 0)\n\n        return torch.tensor(imgs, dtype=torch.float), torch.tensor(masks, dtype=torch.float)\n    \n    def _read_dcm_file(self, file_path):\n        image_size = (self.idf.training_option.custom['image_size'], self.idf.training_option.custom['image_size'])\n        if not file_path:\n            img = np.zeros(image_size, dtype=np.float16)\n        else:\n            img =  pydicom.read_file(file_path).pixel_array\n            img =  mm.fit_transform(img) * 255.0 \n            img = cv2.resize(img, image_size)\n        img = np.tile(img[...,None], [1, 1, 1]) \n        img = img / 255\n\n        return img\n\n    def _read_mask(self, file_path):\n        image_size = (self.idf.training_option.custom['image_size'], self.idf.training_option.custom['image_size'])\n        nii = nib.load(file_path)\n        nii_arrays = nii.get_fdata()  # convert to numpy array\n        mask = nii_arrays[:, :, :self.idf.training_option.custom['image_size']]\n        if mask.shape[2] < idf.training_option.custom['image_size']:\n            ins = np.zeros((mask.shape[0], mask.shape[1], self.idf.training_option.custom['image_size'] - mask.shape[2]), dtype=np.float16)\n            mask = np.concatenate([mask, ins], axis=2)\n    \n        mask = mask / CLASS_NUM\n        mask = cv2.resize(mask, image_size).transpose(2, 1, 0)\n        mask = np.tile(mask[...,None], [1, 1, 1]) # gray to rgb\n\n        return mask","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:50.876026Z","iopub.execute_input":"2022-09-21T05:07:50.876485Z","iopub.status.idle":"2022-09-21T05:07:50.893919Z","shell.execute_reply.started":"2022-09-21T05:07:50.876439Z","shell.execute_reply":"2022-09-21T05:07:50.89294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### DataModule","metadata":{}},{"cell_type":"code","source":"class CustomDataModule(pl.LightningDataModule):\n    def __init__(self, idf, train_df, valid_df, transforms=None):\n        super().__init__()\n        \n        self.idf = idf\n        self.train_df = train_df\n        self.valid_df = valid_df\n        self.train_transforms = transforms['train'] if transforms != None else None\n        self.valid_transforms = transforms['valid'] if transforms != None else None   \n        \n    def train_dataloader(self):\n        dataset = CustomDataset(self.idf, self.train_df, self.train_transforms)\n        data_loader = DataLoader(\n            dataset,\n            batch_size=self.idf.training_option.batch_size,\n            shuffle=True,\n            num_workers=os.cpu_count(),\n            pin_memory=True,\n            drop_last=True\n        )\n        return data_loader\n    \n    def val_dataloader(self):\n        dataset = CustomDataset(self.idf, self.valid_df, self.valid_transforms)\n        data_loader = DataLoader(\n            dataset,\n            batch_size=self.idf.training_option.batch_size,\n            shuffle=False,\n            num_workers=os.cpu_count(),\n            pin_memory=True,\n            drop_last=True\n        )\n        return data_loader","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:50.895676Z","iopub.execute_input":"2022-09-21T05:07:50.896178Z","iopub.status.idle":"2022-09-21T05:07:50.907199Z","shell.execute_reply.started":"2022-09-21T05:07:50.896135Z","shell.execute_reply":"2022-09-21T05:07:50.906106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### transforms","metadata":{}},{"cell_type":"code","source":"data_transforms =  {\n    \"train\": A.Compose([\n        A.HorizontalFlip(),\n        A.VerticalFlip(),\n        A.OneOf([\n            A.RandomContrast(),\n            A.RandomGamma(),\n            A.RandomBrightness(),\n        ], p=0.5),\n    ], p=1.0),\n    \"valid\": None} if idf.training_option.custom[\"transforms\"] else None","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:50.908992Z","iopub.execute_input":"2022-09-21T05:07:50.909371Z","iopub.status.idle":"2022-09-21T05:07:50.922915Z","shell.execute_reply.started":"2022-09-21T05:07:50.909336Z","shell.execute_reply":"2022-09-21T05:07:50.921839Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Scoring functions","metadata":{}},{"cell_type":"code","source":"# Scoring functions\nclass ScoringFn():\n    @classmethod\n    def scoring(cls, pred, label):\n        \"\"\" スコアリング \"\"\"\n        score = F.binary_cross_entropy_with_logits(pred, label)\n        return score\n","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:50.924386Z","iopub.execute_input":"2022-09-21T05:07:50.924905Z","iopub.status.idle":"2022-09-21T05:07:50.932996Z","shell.execute_reply.started":"2022-09-21T05:07:50.924851Z","shell.execute_reply":"2022-09-21T05:07:50.932006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Optimizer & Scheduler","metadata":{}},{"cell_type":"code","source":"# Optimizer Utilities\n\nclass OptimizerUtils():\n    @staticmethod\n    def get_optimizer(idf, model, lr):\n        \"\"\" optimizer を取得 \"\"\"\n        optimizer_name = idf.model_option.optimizer.name\n        optimizer_params = idf.model_option.optimizer.params\n\n        if optimizer_name == \"AdamW\":\n            return AdamW(model.parameters(), \n                         lr=lr,\n                         eps=optimizer_params[\"eps\"], \n                         weight_decay=optimizer_params[\"weight_decay\"])\n        else:\n            raise Exception(f\"optimizer: {optimizer_name} is not supported.\")\n    \n    @staticmethod\n    def get_scheduler(idf, optimizer, num_training_steps, train_dl_size):\n        \"\"\" scheduler を取得 \"\"\"\n        scheduler_name = idf.model_option.scheduler.name\n        scheduler_params = idf.model_option.scheduler.params\n        \n        if scheduler_params[\"num_warmup_steps\"] == \"first_epoch\":\n            scheduler_params[\"num_warmup_steps\"] = train_dl_size\n        \n        if scheduler_name == \"transformer_cosine\":\n            return get_cosine_schedule_with_warmup(\n                optimizer,\n                num_training_steps=num_training_steps,\n                **scheduler_params\n            )\n        \n        elif scheduler_name == \"transformer_cosine_with_hard_restarts\":\n            return get_cosine_with_hard_restarts_schedule_with_warmup(\n                optimizer,\n                num_training_steps=num_training_steps,\n                **scheduler_params\n            )\n        \n        elif scheduler_name == \"transformer_linear\":\n            return get_linear_schedule_with_warmup(\n                optimizer,\n                num_training_steps=num_training_steps,\n                **scheduler_params\n            )\n\n        else:\n            raise Exception(f\"scheduler: {scheduler_name} is not supported.\") ","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:50.93499Z","iopub.execute_input":"2022-09-21T05:07:50.935367Z","iopub.status.idle":"2022-09-21T05:07:50.946531Z","shell.execute_reply.started":"2022-09-21T05:07:50.935331Z","shell.execute_reply":"2022-09-21T05:07:50.945401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Criterion","metadata":{}},{"cell_type":"code","source":"# Criterion Utilities\n\nclass CriterionUtils():\n    @staticmethod\n    def get_criterion(criterion_name, criterion_params):\n        \"\"\" criterion を取得 \"\"\"\n        if criterion_name == \"BCEWithLogitsLoss\":\n            return nn.BCEWithLogitsLoss(**criterion_params)\n        else:\n            raise Exception(f\"criterion: {criterion_name} is not supported.\") ","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:50.952802Z","iopub.execute_input":"2022-09-21T05:07:50.953138Z","iopub.status.idle":"2022-09-21T05:07:50.960021Z","shell.execute_reply.started":"2022-09-21T05:07:50.953112Z","shell.execute_reply":"2022-09-21T05:07:50.95881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model","metadata":{}},{"cell_type":"code","source":"# Inner Model\n\nclass InnerModel(nn.Module):\n    def __init__(self, idf):\n        super().__init__()\n        self.idf = idf\n        # model config\n        base_model_name = idf.model_option.model.params[\"base.model\"]\n        self.model = self.__build_base_model(base_model_name)\n        self.fc = nn.Sequential(nn.ReLU())\n        \n    def forward(self, inputs):\n        output = self.model(inputs)\n        output = self.fc(output)\n        \n        return output\n    \n    def __build_base_model(self, base_model_name):\n        if base_model_name == \"Unet\":\n            return monai.networks.nets.UNet(\n                spatial_dims=3,\n                in_channels=1,\n                out_channels=1,\n                channels=(16, 32, 64, 128, 256),\n                strides=(2, 2, 2, 2),\n                num_res_units=2,\n            )\n            raise Exception(f\"model: {fc_model_name} is not supported.\")\n","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:50.962024Z","iopub.execute_input":"2022-09-21T05:07:50.96239Z","iopub.status.idle":"2022-09-21T05:07:50.974044Z","shell.execute_reply.started":"2022-09-21T05:07:50.962353Z","shell.execute_reply":"2022-09-21T05:07:50.972978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# LightningModule\n\nclass CustomModel(pl.LightningModule):\n    def __init__(self, idf, valid_df, train_dl_size, epoch_size):\n        super().__init__()\n        \n        self.idf = idf\n        # スコア算出に使用する valid データセット\n        self.valid_df = valid_df\n        self.lr = self.idf.model_option.optimizer.params[\"model.lr\"]\n        \n        # scheduler の num_training_steps 引数\n        self.train_dl_size = train_dl_size\n        self.epoch_size = epoch_size\n        self.num_training_steps = train_dl_size * epoch_size\n        \n        # model\n        self.model = InnerModel(self.idf)\n        \n    def forward(self, inputs):\n        \"\"\" 予測を実行 \"\"\"\n        output = self.model(inputs)\n\n        return output\n    \n    def configure_optimizers(self):\n        \"\"\" optimizer, scheduler を取得 \"\"\"\n        optimizer = OptimizerUtils.get_optimizer(self.idf, self.model, self.lr)\n        scheduler = {\n            \"scheduler\": OptimizerUtils.get_scheduler(\n                self.idf,\n                optimizer,\n                self.num_training_steps,\n                self.train_dl_size\n            ),\n            \"interval\": \"step\"\n        }\n        return [optimizer], [scheduler]\n    \n    def training_step(self, batch, batch_idx):\n        \"\"\" Train 処理 \"\"\"\n        loss, _ = self.__common_step(batch, mode=\"train\")\n        return {\"loss\": loss}\n    \n    def validation_step(self, batch, batch_idx):\n        \"\"\" Valid 処理\"\"\"\n        loss, preds = self.__common_step(batch, mode=\"valid\")\n        return {\"loss\": loss, \"pred\": preds, \"label\": batch[1]}\n    \n    def __common_step(self, batch, mode):\n        \"\"\" Train と Valid に共通の処理\"\"\"\n        inputs, labels, = batch\n        y_preds = self.forward(inputs)\n\n        flatten_y_preds = y_preds.view(-1, 1)\n        flatten_labels = labels.view(-1, 1)\n\n        # criterion\n        criterion_names = self.idf.model_option.criterion.name\n        criterion_params = self.idf.model_option.criterion.params\n        losses = []\n        for criterion_name in criterion_names:  \n            criterion = CriterionUtils.get_criterion(criterion_name, criterion_params)\n            loss = criterion(flatten_y_preds, flatten_labels)\n            loss = torch.masked_select(loss, flatten_labels != -1).mean()\n            losses.append(loss)\n        loss = sum(losses) / len(losses)\n                    \n        if torch.isnan(loss):\n            loss = torch.tensor(1.0)\n\n        return loss, y_preds\n    \n    def training_epoch_end(self, outputs):\n        \"\"\" Train の最終処理 \"\"\"\n        # Training ループの最後に呼ばれる\n        avg_losses = self.__common_epoch_end(outputs)\n        self.log(\"train_loss\", avg_losses)\n    \n    def validation_epoch_end(self, outputs):\n        \"\"\" Valid の最終処理 \"\"\"\n        # Valid ループの最後に呼ばれる\n        avg_losses = self.__common_epoch_end(outputs)\n\n        self.log(\"valid_loss\", avg_losses)\n        score = self.__scoring(outputs)\n        self.log(\"score\", score)\n        \n    def __common_epoch_end(self, outputs):\n        \"\"\" Train と Valid に共通の最終処理\"\"\"\n        losses = []\n        for output in outputs:\n            loss = output[\"loss\"].to(\"cpu\").detach().numpy().copy()\n            losses.append(loss)\n        losses = np.array(losses)\n        return np.average(losses)\n    \n    def __scoring(self, outputs):\n        \"\"\" スコア算出 \"\"\"\n        scores = []\n        for output in outputs:\n            pred = output[\"pred\"].to(\"cpu\").detach()\n            label = output[\"label\"].to(\"cpu\").detach()\n            val_score = ScoringFn.scoring(pred, label)\n            scores.append(val_score)\n\n        return np.mean(scores, axis=0)\n","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:09:16.042578Z","iopub.execute_input":"2022-09-21T05:09:16.043186Z","iopub.status.idle":"2022-09-21T05:09:16.074018Z","shell.execute_reply.started":"2022-09-21T05:09:16.043137Z","shell.execute_reply":"2022-09-21T05:09:16.072838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Callback","metadata":{}},{"cell_type":"code","source":"# Custom Callback for Cerebrum\n\nclass CustomCallback(Callback):\n    def __init__(self, idf, fold):\n        self.idf = deepcopy(idf)\n        self.fold = fold\n        self.save_each_model = self.idf.training_option.custom[\"save_each_model\"]\n        \n        self.timer = cc.Timer()\n        \n        # Identification に現在の学習環境を登録\n        self.env_id = len(self.idf.environments)\n        env = cc.Environment(\n            env_id=self.env_id,\n            trainer=self.idf.trainer,\n            platform_name=platform.platform(),\n            cpu_count=os.cpu_count(),\n            gpu_name=torch.cuda.get_device_name(),\n            gpu_count=torch.cuda.device_count()\n        )\n        self.idf.add_environment(env)\n    \n    def on_train_epoch_start(self, trainer, pl_model):\n        \"\"\" train epoch 開始時の処理 \"\"\"\n        self.timer.reset()\n        self.timer.start()\n        \n        generation = self.idf.model_generation\n        current_epoch = trainer.current_epoch\n        logger.info(\n            f\"start epoch-{current_epoch}: \"\n            f\"generation: {generation} -> {generation + 1}\"\n        )\n    \n    def on_train_epoch_end(self, trainer, pl_model):\n        \"\"\" train epoch 終了時の処理 \"\"\"\n        self.timer.stop()\n        elapsed_time = self.timer.get_elapsed_time(format=True)\n        \n        train_loss = trainer.callback_metrics[\"train_loss\"]\n        train_loss = train_loss.to(\"cpu\").detach().numpy().copy()\n        \n        valid_loss = trainer.callback_metrics[\"valid_loss\"]\n        valid_loss = valid_loss.to(\"cpu\").detach().numpy().copy()\n        \n        score = trainer.callback_metrics[\"score\"]\n        score = score.to(\"cpu\").detach().numpy().copy()\n        \n        self.__add_history(train_loss, valid_loss, score)\n        \n        self.__save_idf()\n        \n        if self.save_each_model:\n            self.__save_model(trainer)\n            \n    \n    def __add_history(self, train_loss, valid_loss, score):\n        \"\"\" identification に学習履歴を記録 \"\"\"\n        generation = self.idf.model_generation + 1\n        history_item = cc.Generation(\n            generation=generation,\n            env_id=self.env_id,\n            train_loss=float(train_loss),\n            valid_loss=float(valid_loss),\n            score=float(score),\n            elapsed_sec=self.timer.get_elapsed_time(),\n            recorded_at=self.__get_current_timestamp()\n        )\n        self.idf.add_history(history_item)\n        logger.info(\n            \"save history (generation={}): \"\n            \"train_loss: {:.5f}, valid_loss: {:.5f}, score: {:.5f}\".format(\n                generation, train_loss, valid_loss, score\n            )\n        )\n    \n    def __save_idf(self):\n        \"\"\" identification を GCS に保存 \"\"\"\n        local_path = LOCAL_IDF_PATH.format(\n            model_id=self.idf.model_id, fold=self.fold)\n        self.idf.save(local_path)\n        logger.info(f\"save identification to local: {local_path}\")\n        \n        gcs_path = GCS_IDF_PATH.format(\n            task=TASK,\n            model_name=self.idf.model_name,\n            model_id=self.idf.model_id,\n            dev_mode=\"(Dev)\" if FAST_DEV_RUN or DEV_DATA_SIZE else \"\",\n            fold=self.fold\n        )\n        gcs_connector.upload_to_gcs(local_path, gcs_path)\n        logger.info(f\"save identification to GCS: {gcs_path}\")\n    \n    def __save_model(self, trainer):\n        \"\"\" model を GCS に保存 \"\"\"\n        local_path = LOCAL_MODEL_PATH.format(\n            model_id=self.idf.model_id,\n            fold=self.fold,\n            generation=self.idf.model_generation\n        )\n        trainer.save_checkpoint(local_path)\n        logger.info(f\"save model to local: {local_path}\")\n        \n        gcs_path = GCS_MODEL_PATH.format(\n            task=TASK,\n            model_name=self.idf.model_name,\n            model_id=self.idf.model_id,\n            dev_mode=\"(Dev)\" if FAST_DEV_RUN or DEV_DATA_SIZE else \"\",\n            fold=self.fold,\n            generation=self.idf.model_generation\n        )\n        gcs_connector.upload_to_gcs(local_path, gcs_path)\n        logger.info(f\"save model to GCS: {gcs_path}\")\n    \n    def __get_current_timestamp(self):\n        \"\"\" identification に記録する現在時刻を取得 \"\"\"\n        datetime_now = datetime.datetime.now(\n            datetime.timezone(datetime.timedelta(hours=9))\n        )\n        return datetime_now.strftime(\"%Y%m%d-%H%M%S\")","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:50.999359Z","iopub.execute_input":"2022-09-21T05:07:50.999794Z","iopub.status.idle":"2022-09-21T05:07:51.02113Z","shell.execute_reply.started":"2022-09-21T05:07:50.999757Z","shell.execute_reply":"2022-09-21T05:07:51.020081Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Callbacks\n\nclass Callbacks():\n    def __init__(self, idf, fold):\n        self.idf = idf\n        self.fold = fold\n        \n        self.base_callback = self._get_base_callback()\n        self.early_stopping = self._get_early_stopping()\n        self.min_loss_checkpoint = self._get_min_loss_checkpoint()\n        self.max_score_checkpoint = self._get_max_score_checkpoint()\n        \n        _callbacks = [\n            self.base_callback,\n            self.early_stopping,\n            self.min_loss_checkpoint,\n            self.max_score_checkpoint\n        ]\n        self.callbacks = [c for c in _callbacks if c is not None]\n        \n    @property\n    def has_min_loss_checkpoint(self):\n        return self.min_loss_checkpoint is not None\n    \n    @property\n    def has_max_score_checkpoint(self):\n        return self.max_score_checkpoint is not None\n        \n    @property\n    def min_loss_ckpt_path(self):\n        if self.min_loss_checkpoint is None:\n            return None\n        return self.min_loss_checkpoint.best_model_path\n    \n    @property\n    def max_score_ckpt_path(self):\n        if self.max_score_checkpoint is None:\n            return None\n        return self.max_score_checkpoint.best_model_path\n    \n    def _get_base_callback(self):\n        base_callback = CustomCallback(self.idf, self.fold)\n        return base_callback\n    \n    def _get_early_stopping(self):\n        if not self.idf.training_option.custom[\"early_stopping\"]:\n            return None\n        \n        early_stopping = EarlyStopping(monitor=\"valid_loss\")\n        return early_stopping\n    \n    def _get_min_loss_checkpoint(self):\n        if not self.idf.training_option.custom[\"save_best_loss_model\"]:\n            return None\n        \n        dirpath = LOCAL_BEST_MODEL_PATH.format(model_id=self.idf.model_id, fold=self.fold)\n        min_loss_checkpoint = ModelCheckpoint(\n            dirpath=dirpath,\n            filename=\"min-loss_{epoch}\",\n            monitor=\"valid_loss\",\n            save_top_k=1,\n            mode=\"min\",\n            save_weights_only=True\n        )\n        return min_loss_checkpoint\n    \n    def _get_max_score_checkpoint(self):\n        if not self.idf.training_option.custom[\"save_best_score_model\"]:\n            return None\n        \n        dirpath = LOCAL_BEST_MODEL_PATH.format(model_id=self.idf.model_id, fold=self.fold)\n        max_score_checkpoint = ModelCheckpoint(\n            dirpath=dirpath,\n            filename=\"max-score_{epoch}\",\n            monitor=\"score\",\n            save_top_k=1,\n            mode=\"max\",\n            save_weights_only=True\n        )\n        return max_score_checkpoint","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:51.022857Z","iopub.execute_input":"2022-09-21T05:07:51.023488Z","iopub.status.idle":"2022-09-21T05:07:51.037997Z","shell.execute_reply.started":"2022-09-21T05:07:51.023454Z","shell.execute_reply":"2022-09-21T05:07:51.037015Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model Saving Utils","metadata":{}},{"cell_type":"code","source":"# Model Saving Utilities\n\nclass ModelSavingUtils():\n    @staticmethod\n    def upload_base_model_to_gcs(idf, fold, local_path):\n        \"\"\" best-model を GCS にアップロード \"\"\"\n        ckpt_name = os.path.basename(local_path)\n        gcs_path = GCS_BEST_MODEL_PATH.format(\n            task=TASK,\n            model_name=idf.model_name,\n            model_id=idf.model_id,\n            dev_mode=\"(Dev)\" if FAST_DEV_RUN or DEV_DATA_SIZE else \"\",\n            fold=fold,\n            filename=ckpt_name\n        )\n        gcs_connector.upload_to_gcs(local_path, gcs_path)\n        logger.info(f\"save model to GCS: {gcs_path}\")","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:07:51.0408Z","iopub.execute_input":"2022-09-21T05:07:51.041766Z","iopub.status.idle":"2022-09-21T05:07:51.052561Z","shell.execute_reply.started":"2022-09-21T05:07:51.041734Z","shell.execute_reply":"2022-09-21T05:07:51.051571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Train","metadata":{}},{"cell_type":"code","source":"# Machine Learning\n\nlogger.info(\"start training: model_id={}, n_fold={}, target_fold={}, early_stopping={}\".format(\n    idf.model_id,\n    idf.training_option.custom[\"n_split\"],\n    idf.training_option.custom[\"target_fold\"],\n    idf.training_option.custom[\"early_stopping\"]\n))\n\nfor fold in idf.training_option.custom[\"target_fold\"]:\n    logger.info(f\"fold-{fold} start\")\n    train_df = _train_df[_train_df[\"fold\"] != fold]\n    valid_df = _train_df[_train_df[\"fold\"] == fold]\n    \n    datamodule = CustomDataModule(idf, train_df, valid_df, data_transforms)\n    \n    scheduler_steps = len(datamodule.train_dataloader()) * NUM_EPOCH\n    model = CustomModel(\n        idf,\n        valid_df,\n        len(datamodule.train_dataloader()),\n        NUM_EPOCH\n    )\n\n    callbacks = Callbacks(idf, fold)\n    \n    trainer_params = {\n        \"max_epochs\": NUM_EPOCH,\n        \"accelerator\": \"gpu\",\n        \"gpus\": 1,\n        \"amp_backend\": \"native\",\n        \"fast_dev_run\": FAST_DEV_RUN,\n        \"callbacks\": callbacks.callbacks,\n        \"auto_scale_batch_size\": idf.training_option.custom[\"auto_scale_batch_size\"],\n        # \"auto_lr_find\": idf.training_option.custom[\"auto_lr_find\"],\n    }\n    \n    if idf.environment_option.custom[\"apex\"]:\n        trainer_params[\"amp_backend\"] = \"apex\"\n        \n    trainer = pl.Trainer(**trainer_params)\n    \n    trainer.tune(model, datamodule=datamodule)\n    logger.info(f\"model_lr: {trainer.model.lr}\")\n    \n    trainer.fit(model, datamodule=datamodule)\n    \n    # upload min-loss model to GCS\n    if callbacks.has_min_loss_checkpoint:\n        local_path = callbacks.min_loss_ckpt_path\n        if local_path:\n            logger.info(f\"min-loss-model: {local_path}\")\n            ModelSavingUtils.upload_base_model_to_gcs(idf, fold, local_path)\n    \n    # upload max-score model to GCS\n    if callbacks.has_max_score_checkpoint:\n        local_path = callbacks.max_score_ckpt_path\n        if local_path:\n            logger.info(f\"max-score-model: {local_path}\")\n            ModelSavingUtils.upload_base_model_to_gcs(idf, fold, local_path)\n    \n    del datamodule, model, trainer\n    gc.collect()\n    logger.info(f\"fold-{fold} end\")\n\nlogger.info(\"end training\")","metadata":{"execution":{"iopub.status.busy":"2022-09-21T05:09:19.892538Z","iopub.execute_input":"2022-09-21T05:09:19.893073Z","iopub.status.idle":"2022-09-21T06:32:32.891343Z","shell.execute_reply.started":"2022-09-21T05:09:19.893028Z","shell.execute_reply":"2022-09-21T06:32:32.889121Z"},"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}