{"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":"The amazing [Pytorch Image Models](https://rwightman.github.io/pytorch-image-models) library (aka _timm_) also provides inference and training timings for all supported models, so we can compare compute performance against ImageNet1k accuracy for all large computer vision models.\n\nAll credits to [@RossWightman](https://www.kaggle.com/rwightman) many many thanks! (And also thank you for the amazing presentation last year at [NeurIPS](https://slideslive.com/38969332/imagenet-models-from-the-trenches))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\n\ntorch.__version__","metadata":{"execution":{"iopub.status.busy":"2022-08-30T09:01:58.737809Z","iopub.execute_input":"2022-08-30T09:01:58.738213Z","iopub.status.idle":"2022-08-30T09:01:58.747124Z","shell.execute_reply.started":"2022-08-30T09:01:58.73818Z","shell.execute_reply":"2022-08-30T09:01:58.745691Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference time vs. ImageNet1k accuracy","metadata":{}},{"cell_type":"code","source":"imgntval = pd.read_csv('https://github.com/rwightman/pytorch-image-models/raw/master/results/results-imagenet.csv', index_col='model')\n\n# plotting pytorch 1.12.0 nevertheless, as it also covers swinv2\nchannels_first = pd.read_csv('https://github.com/rwightman/pytorch-image-models/raw/master/results/benchmark-infer-amp-nchw-pt112-cu113-rtx3090.csv', index_col='model')\nchannels_first['log10_infer_samples_per_sec'] = np.log10(channels_first.infer_samples_per_sec)","metadata":{"execution":{"iopub.status.busy":"2022-08-30T09:01:58.755733Z","iopub.execute_input":"2022-08-30T09:01:58.756315Z","iopub.status.idle":"2022-08-30T09:01:59.259286Z","shell.execute_reply.started":"2022-08-30T09:01:58.756279Z","shell.execute_reply":"2022-08-30T09:01:59.25807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def label_point(x, y, val, ax):\n    ds = pd.concat({'x': pd.Series(x.values), 'y': pd.Series(y.values), 'val': pd.Series(val.values)}, axis=1)\n    for i, point in ds.iterrows():\n        t = ax.text(point['x']+.02, point['y']-.01, str(point['val']), alpha=0.8)\n        t.set_bbox(dict(facecolor='w', alpha=0.4, edgecolor='w', boxstyle='round4'))\n\ndef plot_benchmark(ds, min_x=80, xcol='top1', ycol='infer_samples_per_sec', benchmark=channels_first, figsize=(22, 14), use_sns=True):\n    ds = ds.join(benchmark[[ycol]])\n    sample = ds[ds[xcol] > min_x]\n\n    if use_sns:\n        plt.figure(figsize=figsize)\n        ax = sns.regplot(x=xcol, y=ycol, data=sample)\n    else:\n        ax = sample.plot(x=xcol, y=ycol, style='o', figsize=figsize)\n\n    label_point(sample[xcol], sample[ycol], sample.index, ax)\n    return ax\n\n\n_ = plot_benchmark(imgntval, min_x=85, xcol='top1', ycol='log10_infer_samples_per_sec', benchmark=channels_first)","metadata":{"execution":{"iopub.status.busy":"2022-08-30T09:01:59.261294Z","iopub.execute_input":"2022-08-30T09:01:59.261652Z","iopub.status.idle":"2022-08-30T09:02:00.308113Z","shell.execute_reply.started":"2022-08-30T09:01:59.26162Z","shell.execute_reply":"2022-08-30T09:02:00.307166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Beit, swin and convnext seem to lead here (as well as the smaller versions of deit3 and efficientnet_ns).","metadata":{}},{"cell_type":"markdown","source":"# Channels last\ntimm also provides inference times for channels last, but that doesn't seem to change much for most of the large models above:","metadata":{}},{"cell_type":"code","source":"channels_last = pd.read_csv('https://github.com/rwightman/pytorch-image-models/raw/master/results/benchmark-infer-amp-nhwc-pt112-cu113-rtx3090.csv', index_col='model')\nchannels_last['log10_infer_samples_per_sec'] = np.log10(channels_last.infer_samples_per_sec)\n\n_ = plot_benchmark(imgntval, min_x=85, xcol='top1', ycol='log10_infer_samples_per_sec', benchmark=channels_last)","metadata":{"execution":{"iopub.status.busy":"2022-08-30T09:02:00.309363Z","iopub.execute_input":"2022-08-30T09:02:00.310037Z","iopub.status.idle":"2022-08-30T09:02:01.611319Z","shell.execute_reply.started":"2022-08-30T09:02:00.309996Z","shell.execute_reply":"2022-08-30T09:02:01.609918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Correlation between channels first and channels last","metadata":{}},{"cell_type":"code","source":"columns = ['log10_infer_samples_per_sec'] \nchnl_comp = channels_first[columns].join(channels_last[columns], lsuffix='_c1st', rsuffix='_clast')\nchnl_comp.dropna(inplace=True)\n\nchnl_comp = chnl_comp.join(imgntval[['top1']])\nsample = chnl_comp[chnl_comp['top1'] > 85]\nplt.figure(figsize=(22, 14))\nax = sns.regplot(x='log10_infer_samples_per_sec_c1st', y='log10_infer_samples_per_sec_clast', data=sample)\nlabel_point(sample['log10_infer_samples_per_sec_c1st'], sample['log10_infer_samples_per_sec_clast'], sample.index, ax)","metadata":{"execution":{"iopub.status.busy":"2022-08-30T09:02:01.613604Z","iopub.execute_input":"2022-08-30T09:02:01.613959Z","iopub.status.idle":"2022-08-30T09:02:02.670471Z","shell.execute_reply.started":"2022-08-30T09:02:01.613928Z","shell.execute_reply":"2022-08-30T09:02:02.669258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training timings\n\nTraining performance is also highly correlated.","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv('https://github.com/rwightman/pytorch-image-models/raw/master/results/benchmark-train-amp-nchw-pt112-cu113-rtx3090.csv', index_col='model')\ntrain['log10_train_samples_per_sec'] = np.log10(train.train_samples_per_sec)\n\n_ = plot_benchmark(imgntval, min_x=85, xcol='top1', ycol='log10_train_samples_per_sec', benchmark=train)","metadata":{"execution":{"iopub.status.busy":"2022-08-30T09:02:02.672023Z","iopub.execute_input":"2022-08-30T09:02:02.673294Z","iopub.status.idle":"2022-08-30T09:02:03.897965Z","shell.execute_reply.started":"2022-08-30T09:02:02.67324Z","shell.execute_reply":"2022-08-30T09:02:03.896669Z"},"trusted":true},"execution_count":null,"outputs":[]}]}