{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -qU timm\n!pip install -qU wandb","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:20:09.832283Z","iopub.execute_input":"2021-11-27T03:20:09.832592Z","iopub.status.idle":"2021-11-27T03:20:29.938953Z","shell.execute_reply.started":"2021-11-27T03:20:09.832512Z","shell.execute_reply":"2021-11-27T03:20:29.938112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os,glob,warnings,random\nimport cv2\nimport numpy as np\nimport pandas as pd\nfrom tqdm.notebook import tqdm\nfrom pprint import pprint\nfrom sklearn.model_selection import train_test_split\n\n\n# PyTorch \nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda import amp\n\nimport timm\n\n# Albumentations for augmentations\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom albumentations import Normalize,Resize,Compose\nwarnings.filterwarnings(\"ignore\")\n\n# For descriptive error messages\nos.environ['CUDA_LAUNCH_BLOCKING'] = \"1\"","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:20:29.941722Z","iopub.execute_input":"2021-11-27T03:20:29.942024Z","iopub.status.idle":"2021-11-27T03:20:38.292746Z","shell.execute_reply.started":"2021-11-27T03:20:29.941984Z","shell.execute_reply":"2021-11-27T03:20:38.292007Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import wandb\n\ntry:\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    api_key = user_secrets.get_secret(\"WANDB\")\n    wandb.login(key=api_key)\n    anonymous = None\nexcept:\n    anonymous = \"must\"\n    print('To use your W&B account,\\nGo to Add-ons -> Secrets and provide your W&B access token. Use the Label name as WANDB. \\nGet your W&B access token from here: https://wandb.ai/authorize')","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:20:38.294291Z","iopub.execute_input":"2021-11-27T03:20:38.294701Z","iopub.status.idle":"2021-11-27T03:20:38.516384Z","shell.execute_reply.started":"2021-11-27T03:20:38.294659Z","shell.execute_reply":"2021-11-27T03:20:38.515651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def fix_all_seeds(seed):\n    np.random.seed(seed)\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n\nfix_all_seeds(2021)","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:20:38.517760Z","iopub.execute_input":"2021-11-27T03:20:38.518220Z","iopub.status.idle":"2021-11-27T03:20:38.526762Z","shell.execute_reply.started":"2021-11-27T03:20:38.518180Z","shell.execute_reply":"2021-11-27T03:20:38.526045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INPUT_PATH = \"../input/sartorius-cell-instance-segmentation/\"\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint('DEVICE: ', DEVICE)","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:20:38.529772Z","iopub.execute_input":"2021-11-27T03:20:38.530108Z","iopub.status.idle":"2021-11-27T03:20:38.590508Z","shell.execute_reply.started":"2021-11-27T03:20:38.530028Z","shell.execute_reply":"2021-11-27T03:20:38.589678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\ndf_train=pd.read_csv(INPUT_PATH+'train.csv')\ndf_train=df_train.groupby('id')[['cell_type']].first().reset_index()\ndisplay(df_train)","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:20:38.593177Z","iopub.execute_input":"2021-11-27T03:20:38.593453Z","iopub.status.idle":"2021-11-27T03:20:39.139202Z","shell.execute_reply.started":"2021-11-27T03:20:38.593421Z","shell.execute_reply":"2021-11-27T03:20:39.138512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_image_paths=[INPUT_PATH+f'train/{i}.png' for i in df_train['id']]\nsemi_image_paths=glob.glob(INPUT_PATH+'train_semi_supervised/*.png')\ntrain_image_paths.extend(semi_image_paths)\n\ntrain_labels = df_train['cell_type'].to_list()\nsemi_labels=[path.split('/')[-1].split('[')[0] for path in semi_image_paths]\nsemi_labels=['astro' if label=='astros' else label for label in semi_labels]\ntrain_labels.extend(semi_labels)\n\ndf=pd.DataFrame({'image_path':train_image_paths,'cell_type':train_labels})\ndisplay(df)","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:20:39.140609Z","iopub.execute_input":"2021-11-27T03:20:39.140866Z","iopub.status.idle":"2021-11-27T03:20:39.208540Z","shell.execute_reply.started":"2021-11-27T03:20:39.140832Z","shell.execute_reply":"2021-11-27T03:20:39.207867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMAGE_RESIZE=(224,224)\nRESNET_MEAN=(0.485, 0.456, 0.406)\nRESNET_STD = (0.229, 0.224, 0.225)\n\nclass DatasetImageCelltype(Dataset):\n    def __init__(self,df):\n        self.df=df\n        self.images_paths=df['image_path']\n        self.labels=df['cell_type']\n        \n    def __getitem__(self,idx):\n        transforms=Compose([Resize(IMAGE_RESIZE[0],IMAGE_RESIZE[1]),\n                            Normalize(mean=RESNET_MEAN,std=RESNET_STD,p=1),\n                            ToTensorV2()\n                           ])\n        image_path=self.images_paths.iloc[idx]\n        image=cv2.imread(image_path)\n        image=transforms(image=image)['image']\n        label_list=['shsy5y','astro','cort']\n        label=self.labels.iloc[idx]\n        label_id=label_list.index(label)\n        return {'image':image,'label':label_id}\n    \n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:20:39.209840Z","iopub.execute_input":"2021-11-27T03:20:39.210073Z","iopub.status.idle":"2021-11-27T03:20:39.217802Z","shell.execute_reply.started":"2021-11-27T03:20:39.210040Z","shell.execute_reply":"2021-11-27T03:20:39.217002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Split into train and validation\ndf_train, df_valid = train_test_split(df, test_size=0.20)\n\n# Dataset\nds_train = DatasetImageCelltype(df_train)\nds_valid = DatasetImageCelltype(df_valid)\n# Data loader\ndl_train = DataLoader(ds_train, batch_size=64, num_workers=0, pin_memory=True, shuffle=True)\ndl_valid = DataLoader(ds_valid, batch_size=64, num_workers=0, pin_memory=True, shuffle=False)\n\nprint(f'Number of train dataset {len(ds_train)}')\nprint(f'Number of valid dataset {len(ds_valid)}')","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:20:39.219290Z","iopub.execute_input":"2021-11-27T03:20:39.219611Z","iopub.status.idle":"2021-11-27T03:20:39.231679Z","shell.execute_reply.started":"2021-11-27T03:20:39.219551Z","shell.execute_reply":"2021-11-27T03:20:39.230833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import timm\nprint(timm.list_models())","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:20:39.232957Z","iopub.execute_input":"2021-11-27T03:20:39.233259Z","iopub.status.idle":"2021-11-27T03:20:39.245035Z","shell.execute_reply.started":"2021-11-27T03:20:39.233224Z","shell.execute_reply":"2021-11-27T03:20:39.244322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model=timm.create_model('swin_base_patch4_window7_224',pretrained=True)\nmodel","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:20:39.246042Z","iopub.execute_input":"2021-11-27T03:20:39.246241Z","iopub.status.idle":"2021-11-27T03:20:57.817766Z","shell.execute_reply.started":"2021-11-27T03:20:39.246219Z","shell.execute_reply":"2021-11-27T03:20:57.816592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.head=nn.Linear(in_features=1024,out_features=3,bias=True)\nmodel","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:20:57.819558Z","iopub.execute_input":"2021-11-27T03:20:57.820317Z","iopub.status.idle":"2021-11-27T03:20:57.834410Z","shell.execute_reply.started":"2021-11-27T03:20:57.820279Z","shell.execute_reply":"2021-11-27T03:20:57.833662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"LEARNING_RATE = 3e-4\nEPOCHS = 5\nmodel.to(DEVICE)\ncriterion=torch.nn.CrossEntropyLoss()\noptimizer=torch.optim.Adam(model.parameters(),lr=LEARNING_RATE)\nwandb.init()\nwandb.watch(model, log_freq=100)\n\nfor epoch in range(1,EPOCHS+1):\n    print(f'Epoch: {epoch}/{EPOCHS}')\n    model.train()\n    scaler=amp.GradScaler()\n    \n    optimizer.zero_grad()\n    loss_train =0.0\n    correct_train=0.0\n    pbar = tqdm(enumerate(dl_train), total=len(dl_train), desc='Train ')\n    for idx,data in pbar:\n        # Input\n        images, labels = data['image'], data['label']\n        images, labels = images.to(DEVICE), labels.to(DEVICE)\n        \n        with amp.autocast(enabled=True):\n            outputs = model(images) # probabilities\n            loss = criterion(outputs, labels)\n            loss_train += loss\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        optimizer.zero_grad()\n        \n        outputs = outputs.argmax(dim=1) # one hot vector\n        correct_train += (labels==outputs).sum()\n        \n        mem=torch.cuda.memory_reserved()/1E9 if torch.cuda.is_available() else 0\n        pbar.set_postfix(train_loss=f'{loss_train / len(dl_train):0.4f}',\n                        lr=optimizer.param_groups[0]['lr'],\n                        gpu_memory=f'{mem:0.2f} GB')\n    loss_train = loss_train / len(dl_train)\n    acc_train = correct_train / len(ds_train)\n    print(f'Train loss: {loss_train:.4f}, Train accuracy: {acc_train*100:.2f}%') \n    \n    model.eval()\n    loss_valid = 0.0\n    correct_valid = 0.0\n    with torch.no_grad():\n        for data in tqdm(dl_valid, total=len(dl_valid), desc='[valid]'):\n            images, labels = data['image'], data['label']\n            images, labels = images.to(DEVICE), labels.to(DEVICE)\n            outputs = model(images) # probabilities\n            loss_valid += criterion(outputs, labels)\n            outputs = outputs.argmax(dim=1) # one hot vector\n            correct_valid += (labels==outputs).sum()\n            \n    \n    loss_valid = loss_valid / len(dl_valid)\n    acc_valid = correct_valid / len(ds_valid)\n    \n    torch.cuda.empty_cache()\n    print(f'Valid loss: {loss_valid:.4f}, Valid accuracy: {acc_valid*100:.2f}%\\n')   \n    wandb.log({\"Train Loss\": loss_train, \n                   \"Valid Loss\": loss_valid,\n                   \"Train Acc\": acc_train,\n                   \"Valid Acc\": acc_valid,\n                   \"LR\":optimizer.param_groups[0]['lr']})","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:20:57.836391Z","iopub.execute_input":"2021-11-27T03:20:57.836625Z","iopub.status.idle":"2021-11-27T03:28:32.757083Z","shell.execute_reply.started":"2021-11-27T03:20:57.836594Z","shell.execute_reply":"2021-11-27T03:28:32.756329Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model, 'swin_ransformer_crassifier.bin')","metadata":{"execution":{"iopub.status.busy":"2021-11-27T03:28:32.762376Z","iopub.execute_input":"2021-11-27T03:28:32.764644Z","iopub.status.idle":"2021-11-27T03:28:33.311787Z","shell.execute_reply.started":"2021-11-27T03:28:32.764600Z","shell.execute_reply":"2021-11-27T03:28:33.310982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}