{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":14897,"databundleVersionId":1020216,"sourceType":"competition"},{"sourceId":10038,"sourceType":"datasetVersion","datasetId":6978},{"sourceId":853555,"sourceType":"datasetVersion","datasetId":451905},{"sourceId":220530094,"sourceType":"kernelVersion"}],"dockerImageVersionId":30839,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# This Python 3 environment comes with many helpful analytics libraries installed\n# It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# For example, here's several helpful packages to load\n\nimport numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n!pip install albumentations\n!pip install pretrainedmodels\n!pip install iterative-stratification\n!pip install pyarrow\nimport matplotlib.pyplot as plt\nfrom iterstrat.ml_stratifiers import MultilabelStratifiedKFold\nimport joblib\nfrom tqdm import tqdm\nimport torch\nimport warnings\nwarnings.filterwarnings('ignore')\nimport torch.nn as nn\nimport albumentations as A\nimport pretrainedmodels\nimport albumentations.pytorch\n\nfrom torch.utils.data import Dataset\nfrom sklearn.metrics import recall_score\n\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os\nfor dirname, _, filenames in os.walk('/kaggle/input'):\n    for filename in filenames:\n        print(os.path.join(dirname, filename))\n\n# You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:34:20.108070Z","iopub.execute_input":"2025-02-14T05:34:20.108336Z","iopub.status.idle":"2025-02-14T05:34:33.137094Z","shell.execute_reply.started":"2025-02-14T05:34:20.108302Z","shell.execute_reply":"2025-02-14T05:34:33.136187Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!pip install torchvision","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:34:39.798387Z","iopub.execute_input":"2025-02-14T05:34:39.798686Z","iopub.status.idle":"2025-02-14T05:34:43.077526Z","shell.execute_reply.started":"2025-02-14T05:34:39.798659Z","shell.execute_reply":"2025-02-14T05:34:43.076528Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 1. Read Dataset","metadata":{}},{"cell_type":"code","source":"data_dir = \"/kaggle/input/bengaliai-cv19/\"\ndf_train = pd.read_csv(os.path.join(data_dir, \"train.csv\"))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:34:43.606680Z","iopub.execute_input":"2025-02-14T05:34:43.607029Z","iopub.status.idle":"2025-02-14T05:34:43.903718Z","shell.execute_reply.started":"2025-02-14T05:34:43.607000Z","shell.execute_reply":"2025-02-14T05:34:43.903061Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:34:45.386495Z","iopub.execute_input":"2025-02-14T05:34:45.386793Z","iopub.status.idle":"2025-02-14T05:34:45.407236Z","shell.execute_reply.started":"2025-02-14T05:34:45.386772Z","shell.execute_reply":"2025-02-14T05:34:45.406509Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.makedirs('/kaggle/temp')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:34:45.747462Z","iopub.execute_input":"2025-02-14T05:34:45.747756Z","iopub.status.idle":"2025-02-14T05:34:45.751600Z","shell.execute_reply.started":"2025-02-14T05:34:45.747733Z","shell.execute_reply":"2025-02-14T05:34:45.750958Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"files_train = [f'train_image_data_{fid}.feather' for fid in range(4)]\nfor fname in files_train:\n    F = os.path.join(\"/kaggle/input/bengaliaicv19feather\", fname)\n    df_save = pd.read_feather(F)\n    img_ids = df_save['image_id'].values\n    img_array = df_save.iloc[:, 1:].values\n    for idx in tqdm(range(len(df_save))):\n        img_id = img_ids[idx]\n        img = img_array[idx]\n        joblib.dump(img, f\"/kaggle/temp/{img_id}.pkl\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:35:16.541124Z","iopub.execute_input":"2025-02-14T05:35:16.541473Z","iopub.status.idle":"2025-02-14T05:38:47.151317Z","shell.execute_reply.started":"2025-02-14T05:35:16.541445Z","shell.execute_reply":"2025-02-14T05:38:47.150381Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(10, 20))\ndf_train[\"grapheme_root\"].value_counts().sort_index().plot.barh()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:43:58.793804Z","iopub.execute_input":"2025-02-14T05:43:58.794185Z","iopub.status.idle":"2025-02-14T05:44:00.116763Z","shell.execute_reply.started":"2025-02-14T05:43:58.794160Z","shell.execute_reply":"2025-02-14T05:44:00.115846Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train[\"vowel_diacritic\"].value_counts().sort_index().plot.barh()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:00.117891Z","iopub.execute_input":"2025-02-14T05:44:00.118250Z","iopub.status.idle":"2025-02-14T05:44:00.323191Z","shell.execute_reply.started":"2025-02-14T05:44:00.118213Z","shell.execute_reply":"2025-02-14T05:44:00.322366Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train[\"consonant_diacritic\"].value_counts().sort_index().plot.barh()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:00.324568Z","iopub.execute_input":"2025-02-14T05:44:00.324831Z","iopub.status.idle":"2025-02-14T05:44:00.483921Z","shell.execute_reply.started":"2025-02-14T05:44:00.324808Z","shell.execute_reply":"2025-02-14T05:44:00.483126Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train['id']=df_train[\"image_id\"].apply(lambda x : int(x.split('_')[1]))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:04.102333Z","iopub.execute_input":"2025-02-14T05:44:04.102733Z","iopub.status.idle":"2025-02-14T05:44:04.230680Z","shell.execute_reply.started":"2025-02-14T05:44:04.102701Z","shell.execute_reply":"2025-02-14T05:44:04.229924Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"X = df_train[['id', 'grapheme_root', 'vowel_diacritic', 'consonant_diacritic']].values[:, 0]\ny = df_train[['id', 'grapheme_root', 'vowel_diacritic', 'consonant_diacritic']].values[:, 1:]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:04.243476Z","iopub.execute_input":"2025-02-14T05:44:04.243688Z","iopub.status.idle":"2025-02-14T05:44:04.256037Z","shell.execute_reply.started":"2025-02-14T05:44:04.243669Z","shell.execute_reply":"2025-02-14T05:44:04.255192Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Split dataset","metadata":{}},{"cell_type":"code","source":"mskf = MultilabelStratifiedKFold(n_splits=6, random_state=42, shuffle=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:06.338475Z","iopub.execute_input":"2025-02-14T05:44:06.338781Z","iopub.status.idle":"2025-02-14T05:44:06.342473Z","shell.execute_reply.started":"2025-02-14T05:44:06.338757Z","shell.execute_reply":"2025-02-14T05:44:06.341621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train[\"fold\"] = -1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:06.689899Z","iopub.execute_input":"2025-02-14T05:44:06.690277Z","iopub.status.idle":"2025-02-14T05:44:06.695021Z","shell.execute_reply.started":"2025-02-14T05:44:06.690250Z","shell.execute_reply":"2025-02-14T05:44:06.694268Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i, (trn_idx, vld_idx) in enumerate(mskf.split(X, y)):\n    print(i, trn_idx, vld_idx)\n    ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:08.241454Z","iopub.execute_input":"2025-02-14T05:44:08.241829Z","iopub.status.idle":"2025-02-14T05:44:13.007340Z","shell.execute_reply.started":"2025-02-14T05:44:08.241799Z","shell.execute_reply":"2025-02-14T05:44:13.006449Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i, (trn_idx, vld_idx) in enumerate(mskf.split(X, y)):\n    df_train.loc[vld_idx, 'fold'] = i","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:13.008541Z","iopub.execute_input":"2025-02-14T05:44:13.008862Z","iopub.status.idle":"2025-02-14T05:44:17.729049Z","shell.execute_reply.started":"2025-02-14T05:44:13.008827Z","shell.execute_reply":"2025-02-14T05:44:17.728335Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train['fold'].value_counts()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:17.730394Z","iopub.execute_input":"2025-02-14T05:44:17.730758Z","iopub.status.idle":"2025-02-14T05:44:17.737769Z","shell.execute_reply.started":"2025-02-14T05:44:17.730710Z","shell.execute_reply":"2025-02-14T05:44:17.737107Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_train","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:17.738600Z","iopub.execute_input":"2025-02-14T05:44:17.738816Z","iopub.status.idle":"2025-02-14T05:44:17.759192Z","shell.execute_reply.started":"2025-02-14T05:44:17.738796Z","shell.execute_reply":"2025-02-14T05:44:17.758562Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trn_fold = [i for i in range(5) if i not in [5]]\nvld_fold = [5]\ntrn_idx = df_train.loc[df_train['fold'].isin(trn_fold)].index\nvld_idx = df_train.loc[df_train['fold'].isin(vld_fold)].index","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:18.069307Z","iopub.execute_input":"2025-02-14T05:44:18.069578Z","iopub.status.idle":"2025-02-14T05:44:18.088725Z","shell.execute_reply.started":"2025-02-14T05:44:18.069555Z","shell.execute_reply":"2025-02-14T05:44:18.088091Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Define Dataset","metadata":{}},{"cell_type":"code","source":"class BengaliDataset(Dataset):\n    def __init__(self, csv, img_height, img_width, transform):\n        self.csv = csv.reset_index()\n        self.img_ids = csv['image_id'].values\n        self.img_height = img_height\n        self.img_width = img_width\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.csv)\n    def __getitem__(self, index):\n        img_id = self.img_ids[index]\n        img = joblib.load(f'/kaggle/temp/{img_id}.pkl').reshape(self.img_height, self.img_width).astype(np.uint8)\n        img = 255 - img\n        \n        img = img[:, :, np.newaxis]\n        img = np.repeat(img, 3, 2)\n        \n        if self.transform is not None:\n            img = self.transform(image=img)['image']\n\n        label_1 = self.csv.iloc[index]['grapheme_root']\n        label_2 = self.csv.iloc[index]['vowel_diacritic']\n        label_3 = self.csv.iloc[index]['consonant_diacritic']\n\n        return img, np.array([label_1, label_2, label_3])\n        ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:29.522600Z","iopub.execute_input":"2025-02-14T05:44:29.522906Z","iopub.status.idle":"2025-02-14T05:44:29.529074Z","shell.execute_reply.started":"2025-02-14T05:44:29.522882Z","shell.execute_reply":"2025-02-14T05:44:29.528232Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Define augmentations","metadata":{}},{"cell_type":"code","source":"train_augmentation = A.Compose([\n    A.Rotate(20),\n    A.Normalize(normalization=\"min_max\"),\n    A.pytorch.transforms.ToTensorV2()\n])\nvalid_augmentation = A.Compose([\n    A.Normalize(normalization=\"min_max\"),\n    A.pytorch.transforms.ToTensorV2()\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:32.095713Z","iopub.execute_input":"2025-02-14T05:44:32.096078Z","iopub.status.idle":"2025-02-14T05:44:32.103336Z","shell.execute_reply.started":"2025-02-14T05:44:32.096049Z","shell.execute_reply":"2025-02-14T05:44:32.102401Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Make dataloader","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:33.112305Z","iopub.execute_input":"2025-02-14T05:44:33.112594Z","iopub.status.idle":"2025-02-14T05:44:33.116503Z","shell.execute_reply.started":"2025-02-14T05:44:33.112572Z","shell.execute_reply":"2025-02-14T05:44:33.115478Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trn_dataset = BengaliDataset(csv = df_train.loc[trn_idx], \n                            img_height = 137,\n                            img_width = 236,\n                            transform = train_augmentation)\nvld_dataset = BengaliDataset(csv = df_train.loc[vld_idx], \n                            img_height = 137,\n                            img_width = 236,\n                            transform = valid_augmentation)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:42.520391Z","iopub.execute_input":"2025-02-14T05:44:42.520680Z","iopub.status.idle":"2025-02-14T05:44:42.552075Z","shell.execute_reply.started":"2025-02-14T05:44:42.520659Z","shell.execute_reply":"2025-02-14T05:44:42.551359Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trn_dataset[0][0].max()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:45.979964Z","iopub.execute_input":"2025-02-14T05:44:45.980289Z","iopub.status.idle":"2025-02-14T05:44:46.088642Z","shell.execute_reply.started":"2025-02-14T05:44:45.980266Z","shell.execute_reply":"2025-02-14T05:44:46.087953Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trn_loader = DataLoader(trn_dataset, \n                       shuffle=True,\n                       num_workers=4,\n                       batch_size=256)\nvld_loader = DataLoader(vld_dataset, \n                       shuffle=False,\n                       num_workers=4,\n                       batch_size=256)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:47.019455Z","iopub.execute_input":"2025-02-14T05:44:47.019756Z","iopub.status.idle":"2025-02-14T05:44:47.024246Z","shell.execute_reply.started":"2025-02-14T05:44:47.019730Z","shell.execute_reply":"2025-02-14T05:44:47.023540Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for inputs, target in trn_loader:\n    break\ninputs.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:48.860250Z","iopub.execute_input":"2025-02-14T05:44:48.860540Z","iopub.status.idle":"2025-02-14T05:44:50.999908Z","shell.execute_reply.started":"2025-02-14T05:44:48.860519Z","shell.execute_reply":"2025-02-14T05:44:50.998884Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"target.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:44:51.001480Z","iopub.execute_input":"2025-02-14T05:44:51.001807Z","iopub.status.idle":"2025-02-14T05:44:51.007146Z","shell.execute_reply.started":"2025-02-14T05:44:51.001782Z","shell.execute_reply":"2025-02-14T05:44:51.006148Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 6. Model, optimizer, criterion","metadata":{}},{"cell_type":"code","source":"import torchvision\nfrom torchvision import models","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:45:05.661182Z","iopub.execute_input":"2025-02-14T05:45:05.661482Z","iopub.status.idle":"2025-02-14T05:45:05.665242Z","shell.execute_reply.started":"2025-02-14T05:45:05.661461Z","shell.execute_reply":"2025-02-14T05:45:05.664204Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = models.resnet34()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T06:03:06.569148Z","iopub.execute_input":"2025-02-14T06:03:06.569505Z","iopub.status.idle":"2025-02-14T06:03:06.847530Z","shell.execute_reply.started":"2025-02-14T06:03:06.569480Z","shell.execute_reply":"2025-02-14T06:03:06.846791Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T06:03:07.052639Z","iopub.execute_input":"2025-02-14T06:03:07.052973Z","iopub.status.idle":"2025-02-14T06:03:07.059153Z","shell.execute_reply.started":"2025-02-14T06:03:07.052902Z","shell.execute_reply":"2025-02-14T06:03:07.058340Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# model.fc1 = nn.Linear(512, 168)  # For grapheme\n# model.fc2 = nn.Linear(512, 11)   # vowel\n# model.fc3 = nn.Linear(512, 7)    # consonant","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:47:53.462202Z","iopub.execute_input":"2025-02-14T05:47:53.462532Z","iopub.status.idle":"2025-02-14T05:47:53.468107Z","shell.execute_reply.started":"2025-02-14T05:47:53.462510Z","shell.execute_reply":"2025-02-14T05:47:53.467361Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T05:47:57.146443Z","iopub.execute_input":"2025-02-14T05:47:57.146733Z","iopub.status.idle":"2025-02-14T05:47:57.153352Z","shell.execute_reply.started":"2025-02-14T05:47:57.146709Z","shell.execute_reply":"2025-02-14T05:47:57.152648Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"n_grapheme=168\nn_vowel=11\nn_consonant=7\nin_features = model.fc.in_features\nmodel.fc = nn.Linear(in_features, n_grapheme + n_vowel + n_consonant)\nmodel","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T06:04:03.356410Z","iopub.execute_input":"2025-02-14T06:04:03.356732Z","iopub.status.idle":"2025-02-14T06:04:03.364692Z","shell.execute_reply.started":"2025-02-14T06:04:03.356696Z","shell.execute_reply":"2025-02-14T06:04:03.363807Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = model.cuda()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T06:04:12.052754Z","iopub.execute_input":"2025-02-14T06:04:12.053096Z","iopub.status.idle":"2025-02-14T06:04:12.344616Z","shell.execute_reply.started":"2025-02-14T06:04:12.053067Z","shell.execute_reply":"2025-02-14T06:04:12.343942Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"optimizer = torch.optim.Adam(model.parameters(), lr=0.001)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T06:04:13.421950Z","iopub.execute_input":"2025-02-14T06:04:13.422288Z","iopub.status.idle":"2025-02-14T06:04:13.426963Z","shell.execute_reply.started":"2025-02-14T06:04:13.422262Z","shell.execute_reply":"2025-02-14T06:04:13.425978Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"loss_fn = nn.CrossEntropyLoss()\nschedule = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, \n                                                     mode='max',\n                                                     verbose=True,\n                                                     patience = 7,\n                                                     factor=0.5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T06:04:16.757164Z","iopub.execute_input":"2025-02-14T06:04:16.757486Z","iopub.status.idle":"2025-02-14T06:04:16.761629Z","shell.execute_reply.started":"2025-02-14T06:04:16.757458Z","shell.execute_reply":"2025-02-14T06:04:16.760883Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 7. training","metadata":{}},{"cell_type":"code","source":"from tqdm import tqdm_notebook","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T06:04:18.957481Z","iopub.execute_input":"2025-02-14T06:04:18.957845Z","iopub.status.idle":"2025-02-14T06:04:18.961619Z","shell.execute_reply.started":"2025-02-14T06:04:18.957813Z","shell.execute_reply":"2025-02-14T06:04:18.960713Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_score = -1","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T06:04:23.844100Z","iopub.execute_input":"2025-02-14T06:04:23.844419Z","iopub.status.idle":"2025-02-14T06:04:23.848019Z","shell.execute_reply.started":"2025-02-14T06:04:23.844395Z","shell.execute_reply":"2025-02-14T06:04:23.847162Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for e in range(2):\n    train_loss = []\n    model.train()\n\n\n    for inputs, targets in tqdm_notebook(trn_loader):\n        \n    \n        inputs = inputs.cuda()\n        targets = targets.cuda()\n        \n        logits = model(inputs)\n        logits_split = torch.split(logits, [n_grapheme, n_vowel, n_consonant], dim=1)\n        \n        loss = loss_fn(logits_split[0], targets[:, 0]) + loss_fn(logits_split[1], targets[:, 1]) + loss_fn(logits_split[2], targets[:, 2])\n        \n        loss.backward()\n        \n        optimizer.step()\n        optimizer.zero_grad()\n        train_loss.append(loss.item())\n    \n    val_loss = []\n    val_true = []\n    val_pred = []\n    \n    model.eval()\n    \n    with torch.no_grad():\n        for inputs, targets in tqdm_notebook(vld_loader):\n            inputs = inputs.cuda()\n            targets = targets.cuda()\n    \n            logits = model(inputs)\n            \n            logits_split = torch.split(logits, [n_grapheme, n_vowel, n_consonant], dim=1)\n        \n            loss = loss_fn(logits_split[0], targets[:, 0]) + loss_fn(logits_split[1], targets[:, 1]) + loss_fn(logits_split[2], targets[:, 2])\n    \n            val_loss.append(loss.item())\n    \n            grapheme = logits_split[0].cpu().argmax(dim=1).data.numpy()\n            vowel = logits_split[1].cpu().argmax(dim=1).data.numpy()\n            cons = logits_split[2].cpu().argmax(dim=1).data.numpy()\n    \n            val_true.append(targets.cpu().numpy())\n            val_pred.append(np.stack([grapheme, vowel, cons], axis=1))\n    \n    \n    val_true = np.concatenate(val_true)\n    val_pred = np.concatenate(val_pred)\n    \n    print(val_true.shape, val_pred.shape)\n    \n    val_loss = np.mean(val_loss)\n    train_loss = np.mean(train_loss)\n    \n    print(val_loss, train_loss)\n    \n    score_g = recall_score(val_true[:, 0], val_pred[:, 0], average='macro')\n    score_v = recall_score(val_true[:, 1], val_pred[:, 1], average='macro')\n    score_c = recall_score(val_true[:, 2], val_pred[:, 2], average='macro')\n    \n    final_score = np.average([score_g, score_v, score_c], weights=[2, 1, 1])\n    \n    print(f'train_loss: {train_loss: .5f}; val_loss: {val_loss: .5f}; score: {final_score: .5f}')\n    print(f'score_g: {score_g: .5f}; score_c: {score_v: .5f}; score: {score_c: .5f}')\n\n    if final_score > best_score:\n        best_score = final_score\n\n        torch.save(model, \"model.pth\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-14T06:09:22.781205Z","iopub.execute_input":"2025-02-14T06:09:22.781554Z","iopub.status.idle":"2025-02-14T06:15:57.634571Z","shell.execute_reply.started":"2025-02-14T06:09:22.781530Z","shell.execute_reply":"2025-02-14T06:15:57.633365Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}