{"cells":[{"metadata":{},"cell_type":"markdown","source":"# Use Fastai "},{"metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","trusted":true},"cell_type":"code","source":"%reload_ext autoreload\n%autoreload 2\n%matplotlib inline","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"from fastai import *\nfrom fastai.vision import *\nimport pandas as pd\nfrom fastai.utils.mem import *","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"path = Path('/kaggle/input/iwildcam-2019-fgvc6')\n\ndebug =1\nif debug:\n    train_pct=0.04\nelse:\n    train_pct=0.5","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Load train dataframe\ntrain_df = pd.read_csv(path/'train.csv')\ntrain_df = pd.concat([train_df['id'],train_df['category_id']],axis=1,keys=['id','category_id'])\ntrain_df.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# Load sample submission\ntest_df = pd.read_csv(path/'test.csv')\ntest_df = pd.DataFrame(test_df['id'])\ntest_df['predicted'] = 0\ntest_df.head()\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# 資料增強"},{"metadata":{"trusted":true},"cell_type":"code","source":"free = gpu_mem_get_free_no_cache()\n# the max size of bs depends on the available GPU RAM\nif free > 8200: bs=64\nelse:           bs=32\nprint(f\"using bs={bs}, have {free}MB of GPU RAM free\")\n\ntfms = get_transforms(max_rotate=20, max_zoom=1.3, max_lighting=0.4, max_warp=0.4,\n                      p_affine=1., p_lighting=1.)\n# app_train = app_train.append(app_test).reset_index()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train, test = [ImageList.from_df(df, path=path, cols='id', folder=folder, suffix='.jpg') \n               for df, folder in zip([train_df, test_df], ['train_images', 'test_images'])]\nif debug:\n    src= train.split_subsets(train_size=train_pct, valid_size= train_pct*2)\n#     test=test[:1000]\nelse:\n    src= train.split_subsets(train_size=train_pct, valid_size=0.2, seed=2)\n#     src= train.split_by_rand_pct(0.2, seed=2)\n\nprint(src)\n    \ndef get_data(size, bs, padding_mode='reflection'):\n    return (src.label_from_df(cols='category_id')\n           .add_test(test)\n           .transform(tfms, size=size, padding_mode=padding_mode)\n           .databunch(bs=bs).normalize(imagenet_stats))    \n    \n# data = (train.split_by_rand_pct(0.2, seed=2)\n#         .label_from_df(cols='category_id')\n#         .add_test(test)\n#         .transform(get_transforms(), size=32)\n#         .databunch(path=Path('.'), bs=64).normalize())","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"data = get_data(224, bs, 'zeros')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"def _plot(i,j,ax):\n    x,y = data.train_ds[3]\n    x.show(ax, y=y)\n\nplot_multi(_plot, 3, 3, figsize=(8,8))","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Train model\n"},{"metadata":{"trusted":true},"cell_type":"code","source":"# learn = cnn_learner(data, base_arch=models.densenet121, metrics=[FBeta(),accuracy], wd=1e-5).mixup()\ngc.collect()\n# wd=1e-2\nwd=1e-1\nlearn = cnn_learner(data, models.resnet34, metrics=error_rate, bn_final=True, wd=wd )\nlearn.model_dir= '/kaggle/working/'","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# lr=1e-2\n# learn.fit_one_cycle(3, slice(lr), pct_start=0.8)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# learn.save('223')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# learn.unfreeze()\n# learn.lr_find()\n# learn.recorder.plot(suggestion=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# # lr = 2e-2\n# # learn.fit_one_cycle(2, slice(lr))\n\n# learn.fit_one_cycle(6, max_lr=slice(5.75E-06,lr/5), pct_start=0.8)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# learn.save('224')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"!cp /kaggle/input/fastai-starter-iwildcam-2019-ad561b/224.pth /kaggle/working\nlearn.load('224')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"data = get_data(352,bs)\nlearn.data = data\n","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learn.fit_one_cycle(2, max_lr=slice(1e-6,1e-4))\nlearn.save('352')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# !cp /kaggle/input/352pth/352.pth /kaggle/working\n# learn.load('352')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learn.unfreeze()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learn.lr_find()\nlearn.recorder.plot()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Here, we use discriminative learning rates, where lower learning rates are used for the earlier layers in the model."},{"metadata":{"trusted":true},"cell_type":"code","source":"lr = 1e-3\nlearn.fit_one_cycle(8, slice(lr/100, lr))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"learn.save('stage-2-sz32')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# !cp /kaggle/input/stage-2-sz32.pth /kaggle/working\n# learn.load('stage-2-sz32')","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Interpretation\n\nThe fastai library also provides some functions for interpreting the models, such as displaying images with the top losses, displaying confusion matrices, and [more](https://docs.fast.ai/vision.learner.html#ClassificationInterpretation)."},{"metadata":{"trusted":true},"cell_type":"code","source":"interp = ClassificationInterpretation.from_learner(learn)\n\nlosses,idxs = interp.top_losses()\n\nlen(data.valid_ds)==len(losses)==len(idxs)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"interp.plot_confusion_matrix(figsize=(12,12), dpi=60)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Test predictions"},{"metadata":{"trusted":true},"cell_type":"code","source":"test_preds = learn.get_preds(DatasetType.Test)\ntest_df['predicted'] = test_preds[0].argmax(dim=1)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"test_df.shape","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"csv_path ='/kaggle/working/submission.csv'\ntest_df.to_csv(csv_path, index=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# # import the modules we'll need\n# from IPython.display import HTML\n# import base64\n\n\n# # function that takes in a dataframe and creates a text link to  \n# # download it (will only work for files < 2MB or so)\n# def create_download_link(df, title = \"Download CSV file\", filename = \"data.csv\"):  \n#     csv = df.to_csv()\n#     b64 = base64.b64encode(csv.encode())\n#     payload = b64.decode()\n#     html = '<a download=\"{filename}\" href=\"data:text/csv;base64,{payload}\" target=\"_blank\">{title}</a>'\n#     html = html.format(payload=payload,title=title,filename=filename)\n#     return HTML(html)\n\n\n\n# # create a link to download the dataframe\n# create_download_link(test_df[90000:120000])\n","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.6.4"}},"nbformat":4,"nbformat_minor":1}