{"cells":[{"metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"cell_type":"code","source":"from fastai.vision.all import *","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"set_seed(999)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"To properly use timm we need to get creative with the imports:"},{"metadata":{"trusted":true},"cell_type":"code","source":"%cd ../input/timm030/pytorch-image-models-master/pytorch-image-models-master/\nfrom timm import create_model\n%cd ../../../../","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"%ls","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"We'll want to be able to recreate our model fully, so we'll bring in the `wwf` code:"},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true},"cell_type":"code","source":"from fastai.vision.learner import _update_first_layer\n\n# Cell\ndef create_timm_body(arch:str, pretrained=True, cut=None, n_in=3):\n    \"Creates a body from any model in the `timm` library.\"\n    model = create_model(arch, pretrained=pretrained, num_classes=0, global_pool='')\n    _update_first_layer(model, n_in, pretrained)\n    if cut is None:\n        ll = list(enumerate(model.children()))\n        cut = next(i for i,o in reversed(ll) if has_pool_type(o))\n    if isinstance(cut, int): return nn.Sequential(*list(model.children())[:cut])\n    elif callable(cut): return cut(model)\n    else: raise NamedError(\"cut must be either integer or function\")\n\n# Cell\ndef create_timm_model(arch:str, n_out, cut=None, pretrained=True, n_in=3, init=nn.init.kaiming_normal_, custom_head=None,\n                     concat_pool=True, **kwargs):\n    \"Create custom architecture using `arch`, `n_in` and `n_out` from the `timm` library\"\n    body = create_timm_body(arch, pretrained, None, n_in)\n    if custom_head is None:\n        nf = num_features_model(nn.Sequential(*body.children())) * (2 if concat_pool else 1)\n        head = create_head(nf, n_out, concat_pool=concat_pool, **kwargs)\n    else: head = custom_head\n    model = nn.Sequential(body, head)\n    if init is not None: apply_init(model[1], init)\n    return model\n\n# Cell\nfrom fastai.vision.learner import _add_norm\n\n# Cell\ndef timm_learner(dls, arch:str, loss_func=None, pretrained=True, cut=None, splitter=None,\n                y_range=None, config=None, n_out=None, normalize=True, **kwargs):\n    \"Build a convnet style learner from `dls` and `arch` using the `timm` library\"\n    if config is None: config = {}\n    if n_out is None: n_out = get_c(dls)\n    assert n_out, \"`n_out` is not defined, and could not be inferred from data, set `dls.c` or pass `n_out`\"\n    if y_range is None and 'y_range' in config: y_range = config.pop('y_range')\n    model = create_timm_model(arch, n_out, default_split, pretrained, y_range=y_range, **config)\n    learn = Learner(dls, model, loss_func=loss_func, splitter=default_split, **kwargs)\n    if pretrained: learn.freeze()\n    return learn","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Recreate our data to get access to the `test_dl`:"},{"metadata":{"trusted":true},"cell_type":"code","source":"path = Path(\"input\")\ndata_path = path/'cassava-leaf-disease-classification'","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"data_path.ls()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df = pd.read_csv(data_path/'train.csv')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"df['image_id'] = df['image_id'].apply(lambda x: f'train_images/{x}')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"blocks = (ImageBlock, CategoryBlock)\nsplitter = RandomSplitter(valid_pct=0.2)\ndef get_x(row): return data_path/row['image_id']\n\ndef get_y(row): return row['label']\nitem_tfms = [Resize(512)]\nbatch_tfms = [RandomResizedCropGPU(448), *aug_transforms(), Normalize.from_stats(*imagenet_stats)]\nblock = DataBlock(blocks = blocks,\n                 get_x = get_x,\n                 get_y = get_y,\n                 splitter = splitter,\n                 item_tfms = item_tfms,\n                 batch_tfms = batch_tfms)\ndls = block.dataloaders(df, bs=32)\n","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Build a `Learner`"},{"metadata":{"trusted":true},"cell_type":"code","source":"learn = timm_learner(dls, 'efficientnet_b3', metrics=accuracy, pretrained=False)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"And load in our weights"},{"metadata":{"trusted":true},"cell_type":"code","source":"learn.model_dir = Path('input/b3_example_submission')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"load_model(Path('input/b3-example-submission/b3.pth'), learn.model, learn.opt)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample_df = pd.read_csv(data_path/'sample_submission.csv')\nsample_df.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample_copy = sample_df.copy()\nsample_copy['image_id'] = sample_copy['image_id'].apply(lambda x: f'test_images/{x}')","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"Finally we can grab our predictions using TTA"},{"metadata":{"trusted":true},"cell_type":"code","source":"test_dl = learn.dls.test_dl(sample_copy)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"preds, _ = learn.tta(dl=test_dl)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample_df['label'] = preds.argmax(dim=-1).numpy()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"%cd working","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"sample_df.to_csv('submission.csv',index=False)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"","execution_count":null,"outputs":[]}],"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":4,"nbformat_minor":4}