{"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":"%%writefile filepaths.yaml\n\n\nraw_dir: '/kaggle/input/birdclef-2023'\nprocessed_dir: ''\ninterim_dir: ''\nexternal_dir: ''\n\nmodels_dir: '/kaggle/input'","metadata":{"papermill":{"duration":6.667352,"end_time":"2022-04-22T06:00:08.901647","exception":false,"start_time":"2022-04-22T06:00:02.234295","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-05-24T09:22:50.804257Z","iopub.execute_input":"2023-05-24T09:22:50.805051Z","iopub.status.idle":"2023-05-24T09:22:50.844857Z","shell.execute_reply.started":"2023-05-24T09:22:50.805007Z","shell.execute_reply":"2023-05-24T09:22:50.843887Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install onnxruntime","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:22:50.846774Z","iopub.execute_input":"2023-05-24T09:22:50.847162Z","iopub.status.idle":"2023-05-24T09:23:07.704547Z","shell.execute_reply.started":"2023-05-24T09:22:50.847125Z","shell.execute_reply":"2023-05-24T09:23:07.703448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Imports","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport soundfile as sf\n\nimport os\nimport yaml\nfrom types import SimpleNamespace\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.nn.functional import cross_entropy\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.distributions import Beta\nfrom torch.nn.parameter import Parameter\n\nimport torch.multiprocessing as mp\nimport torchaudio\n\nfrom sklearn import model_selection\n\nimport torchmetrics\nimport timm\nfrom pathlib import Path\n\nfrom tqdm.notebook import tqdm\n\nfrom transformers import AutoModel, AutoConfig\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport timm\nimport random\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport concurrent\n\nimport time\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchaudio\nimport librosa\nimport numpy as np\n\ntorch.jit.enable_onednn_fusion(True)\n\n# from onnxconverter_common import float16\nimport onnx\nimport onnxruntime as rt","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:07.705918Z","iopub.execute_input":"2023-05-24T09:23:07.706268Z","iopub.status.idle":"2023-05-24T09:23:23.538098Z","shell.execute_reply.started":"2023-05-24T09:23:07.706230Z","shell.execute_reply":"2023-05-24T09:23:23.536577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.set_num_threads(4)","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:23.541340Z","iopub.execute_input":"2023-05-24T09:23:23.541712Z","iopub.status.idle":"2023-05-24T09:23:23.578562Z","shell.execute_reply.started":"2023-05-24T09:23:23.541673Z","shell.execute_reply":"2023-05-24T09:23:23.577510Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n    \nseed_everything()","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:23.582968Z","iopub.execute_input":"2023-05-24T09:23:23.583765Z","iopub.status.idle":"2023-05-24T09:23:23.617048Z","shell.execute_reply.started":"2023-05-24T09:23:23.583713Z","shell.execute_reply":"2023-05-24T09:23:23.615902Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Utils","metadata":{}},{"cell_type":"code","source":"def load_config(filepath):\n    with open(filepath, 'rb') as file:\n        data = yaml.safe_load(file)\n    data = dictionary_to_namespace(data)\n    return data\n\n\ndef load_filepaths(filepath):\n    with open(filepath, 'rb') as file:\n        data = yaml.safe_load(file)\n\n    path_to_file = Path(filepath).parents[0]\n    for key, value in data.items():\n        data[key] = Path(path_to_file / Path(value)).resolve()\n\n    data = dictionary_to_namespace(data)\n    return data\n\n\ndef dictionary_to_namespace(data):\n    if type(data) is list:\n        return list(map(dictionary_to_namespace, data))\n    elif type(data) is dict:\n        sns = SimpleNamespace()\n        for key, value in data.items():\n            setattr(sns, key, dictionary_to_namespace(value))\n        return sns\n    else:\n        return data","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:23.622090Z","iopub.execute_input":"2023-05-24T09:23:23.622762Z","iopub.status.idle":"2023-05-24T09:23:23.634799Z","shell.execute_reply.started":"2023-05-24T09:23:23.622718Z","shell.execute_reply":"2023-05-24T09:23:23.633900Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Dataset","metadata":{}},{"cell_type":"code","source":"class AudioTransform:\n    def __init__(self, always_apply=False, p=0.5):\n        self.always_apply = always_apply\n        self.p = p\n\n    def __call__(self, wav, label, weight, sr):\n        if self.always_apply:\n            return self.apply(wav, label, weight, sr)\n        else:\n            if np.random.rand() < self.p:\n                return self.apply(wav, label, weight, sr)\n            else:\n                return wav, label, weight, sr\n\n    def apply(self, wav, label, weight, sr=32000):\n        raise NotImplementedError\n        \n\nclass Normalize(AudioTransform):\n    def __init__(self, always_apply=False, p=1):\n        super().__init__(always_apply, p)\n\n    def apply(self, wav, label, weight, sr=32000):\n        max_vol = np.abs(wav).max()\n        y_vol = wav * 1 / max_vol\n        return np.asfortranarray(y_vol), label, weight, sr\n\nclass NormalizeMelSpec(nn.Module):\n    def __init__(self, eps=1e-6):\n        super().__init__()\n        self.eps = eps\n\n    def forward(self, X):\n        mean = X.mean((1, 2), keepdim=True)\n        std = X.std((1, 2), keepdim=True)\n        Xstd = (X - mean) / (std + self.eps)\n        norm_min, norm_max = Xstd.min(-1)[0].min(-1)[0], Xstd.max(-1)[0].max(-1)[0]\n        fix_ind = (norm_max - norm_min) > self.eps * torch.ones_like(\n            (norm_max - norm_min)\n        )\n        V = torch.zeros_like(Xstd)\n        if fix_ind.sum():\n            V_fix = Xstd[fix_ind]\n            norm_max_fix = norm_max[fix_ind, None, None]\n            norm_min_fix = norm_min[fix_ind, None, None]\n            V_fix = torch.max(\n                torch.min(V_fix, norm_max_fix),\n                norm_min_fix,\n            )\n            # print(V_fix.shape, norm_min_fix.shape, norm_max_fix.shape)\n            V_fix = (V_fix - norm_min_fix) / (norm_max_fix - norm_min_fix)\n            V[fix_ind] = V_fix\n        return V\n    \n    \nclass TestDataset(Dataset):\n    def __init__(self, audios, dataframe, config):\n        self.audios = audios\n        \n        dfs = []\n        for path, id in zip(dataframe.path.values, dataframe.filename.values):\n            for i in range(5, 605, 5):\n                start_seconds = i - 5\n                end_seconds = i\n                end_index = int(32000 * end_seconds)\n                start_index = int(32000 * start_seconds)\n                row_id = f'{id}_{end_seconds}'\n                dfs.append([path, start_index, end_index, row_id])\n        self.dataframe = pd.DataFrame(dfs, columns=['path', 'start', 'stop', 'row_id'])\n        \n        self.config = config\n        self.sr = self.config.dataset.sample_rate\n\n        self.mel_spec = torchaudio.transforms.MelSpectrogram(\n            32000,\n            n_mels=self.config.dataset.n_mels,\n            n_fft=self.config.dataset.nfft,\n            hop_length=self.config.dataset.hop_length,\n            f_max=self.config.dataset.fmax,\n            f_min=self.config.dataset.fmin,\n        )\n        self.amplitude_to_db = torchaudio.transforms.AmplitudeToDB(top_db=config.dataset.top_db)\n        self.spec_norm = NormalizeMelSpec()\n        \n        self.wav_norm = Normalize(p=1)\n        \n        self.spec_transforms = nn.Sequential(\n                self.mel_spec,\n                self.amplitude_to_db,\n                self.spec_norm,\n            )\n        \n        self.train_period = config.dataset.train_duration\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, item):\n        row = self.dataframe.iloc[item]\n        fp = row.path\n        start = row.start\n        stop = row.stop\n        row_id = row.row_id\n        \n        audio = self.audios[str(fp)][start:stop]\n        audio = torch.tensor(self.wav_norm(audio, None, None, None)[0])\n        image = self.spec_transforms(audio.view(1, audio.shape[0]))\n        return image, row_id","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:23.636935Z","iopub.execute_input":"2023-05-24T09:23:23.637876Z","iopub.status.idle":"2023-05-24T09:23:23.671981Z","shell.execute_reply.started":"2023-05-24T09:23:23.637796Z","shell.execute_reply":"2023-05-24T09:23:23.670887Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Model","metadata":{}},{"cell_type":"code","source":"def gem_freq(x, p=3, eps=1e-6):\n    return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), 1)).pow(1.0 / p)\n\n\nclass GeMFreq(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super().__init__()\n        self.p = torch.nn.Parameter(torch.ones(1) * p)\n        self.eps = eps\n\n    def forward(self, x):\n        return gem_freq(x, p=self.p, eps=self.eps)\n\n\nclass AttHead(nn.Module):\n    def __init__(\n        self, in_chans, p=0.5, num_class=264, train_period=15.0, infer_period=5.0\n    ):\n        super().__init__()\n        self.train_period = train_period\n        self.infer_period = infer_period\n        self.pooling = GeMFreq()\n\n        self.dense_layers = nn.Sequential(\n            nn.Dropout(p / 2),\n            nn.Linear(in_chans, 512),\n            nn.ReLU(),\n            nn.Dropout(p),\n        )\n        self.attention = nn.Conv1d(\n            in_channels=512,\n            out_channels=num_class,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n            bias=True,\n        )\n        self.fix_scale = nn.Conv1d(\n            in_channels=512,\n            out_channels=num_class,\n            kernel_size=1,\n            stride=1,\n            padding=0,\n            bias=True,\n        )\n\n    def forward(self, feat):\n        feat = self.pooling(feat).squeeze(-2).permute(0, 2, 1)  # (bs, time, ch)\n\n        feat = self.dense_layers(feat).permute(0, 2, 1)  # (bs, 512, time)\n        # print(feat.shape)\n        \n        time_att = torch.tanh(self.attention(feat))\n        \n        assert self.train_period >= self.infer_period\n        \n        clipwise_pred = torch.sum(\n                torch.sigmoid(self.fix_scale(feat)) * torch.softmax(time_att, dim=-1),\n                dim=-1,\n            )\n        return clipwise_pred\n\n\nclass TattakaModel(nn.Module):\n    def __init__(\n        self,\n        config,\n        pretrained,\n    ):\n        super().__init__()\n\n        # self.model = get_timm_backbone(config, pretrained)\n        self.model = timm.create_model(\n            config.model.backbone_type, features_only=True, pretrained=False, in_chans=1\n        )\n        encoder_channels = self.model.feature_info.channels()\n        dense_input = encoder_channels[-1]\n        self.head = AttHead(\n            dense_input,\n            p=config.model.dropout,\n            num_class=len(config.dataset.labels),\n            train_period=config.dataset.train_duration,\n            infer_period=config.dataset.valid_duration,\n        )\n        self.criterion = nn.BCEWithLogitsLoss(reduction='none')\n\n    def forward(self, images):\n        spec = images\n        \n        feats = self.model(spec)\n        output_clip = self.head(feats[-1])\n        return output_clip\n","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:23.673957Z","iopub.execute_input":"2023-05-24T09:23:23.674384Z","iopub.status.idle":"2023-05-24T09:23:23.699000Z","shell.execute_reply.started":"2023-05-24T09:23:23.674344Z","shell.execute_reply":"2023-05-24T09:23:23.697771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Process","metadata":{}},{"cell_type":"code","source":"exp_list = [\n#     ['exp290', [-1]],\n#     ['exp315', [0, 1, 2, 3, -1]],\n#     ['exp283', [-1]],\n#     ['exp287', [-1]],\n    ['exp304', [-1]],\n#     ['exp310', [-1]],\n#     ['exp297', [-1]],\n#     ['exp325', [-1]],\n#     ['exp321', [0, 1]],\n#     ['exp290', [0, 1, 2, 3, -1]],\n# #     ['exp306', [0, 1, 2, 3, -1]],\n    \n#     ['exp303', [0, 1, 2, 3, -1]],\n#     ['exp302', [0, 1, 2, 3, -1]],\n    \n#     ['exp294', [0, 1, 2, 3, -1]],\n#     ['exp295', [0, 1, 2, 3, -1]],\n#     ['exp331', [0, 1, 2, 3, -1]],\n#         ['exp300', [0, 2, 3, -1]],\n#     ['exp304', [0, 1, 2, 3, -1]],\n#     ['exp309', [0, 1, 2, 3, -1]],\n#     ['exp305', [0, 1, -1]],\n#     ['exp288', [-1]],\n    \n#     ['exp294', [0, 1, 2, 3, -1]],\n#     ['exp302', [0, 1, 2, 3, -1]],\n]\n\nmultihead_list = [\n]\n\nn_samples = 1\nif len(os.listdir('/kaggle/input/birdclef-2023/test_soundscapes')) > 1:\n    n_samples = 1\n\nfilepaths = load_filepaths('filepaths.yaml')","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:23.703582Z","iopub.execute_input":"2023-05-24T09:23:23.703997Z","iopub.status.idle":"2023-05-24T09:23:23.722709Z","shell.execute_reply.started":"2023-05-24T09:23:23.703958Z","shell.execute_reply":"2023-05-24T09:23:23.721499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = ['abethr1', 'abhori1', 'abythr1', 'afbfly1', 'afdfly1', 'afecuc1',\n       'affeag1', 'afgfly1', 'afghor1', 'afmdov1', 'afpfly1', 'afpkin1',\n       'afpwag1', 'afrgos1', 'afrgrp1', 'afrjac1', 'afrthr1', 'amesun2',\n       'augbuz1', 'bagwea1', 'barswa', 'bawhor2', 'bawman1', 'bcbeat1',\n       'beasun2', 'bkctch1', 'bkfruw1', 'blacra1', 'blacuc1', 'blakit1',\n       'blaplo1', 'blbpuf2', 'blcapa2', 'blfbus1', 'blhgon1', 'blhher1',\n       'blksaw1', 'blnmou1', 'blnwea1', 'bltapa1', 'bltbar1', 'bltori1',\n       'blwlap1', 'brcale1', 'brcsta1', 'brctch1', 'brcwea1', 'brican1',\n       'brobab1', 'broman1', 'brosun1', 'brrwhe3', 'brtcha1', 'brubru1',\n       'brwwar1', 'bswdov1', 'btweye2', 'bubwar2', 'butapa1', 'cabgre1',\n       'carcha1', 'carwoo1', 'categr', 'ccbeat1', 'chespa1', 'chewea1',\n       'chibat1', 'chtapa3', 'chucis1', 'cibwar1', 'cohmar1', 'colsun2',\n       'combul2', 'combuz1', 'comsan', 'crefra2', 'crheag1', 'crohor1',\n       'darbar1', 'darter3', 'didcuc1', 'dotbar1', 'dutdov1', 'easmog1',\n       'eaywag1', 'edcsun3', 'egygoo', 'equaka1', 'eswdov1', 'eubeat1',\n       'fatrav1', 'fatwid1', 'fislov1', 'fotdro5', 'gabgos2', 'gargan',\n       'gbesta1', 'gnbcam2', 'gnhsun1', 'gobbun1', 'gobsta5', 'gobwea1',\n       'golher1', 'grbcam1', 'grccra1', 'grecor', 'greegr', 'grewoo2',\n       'grwpyt1', 'gryapa1', 'grywrw1', 'gybfis1', 'gycwar3', 'gyhbus1',\n       'gyhkin1', 'gyhneg1', 'gyhspa1', 'gytbar1', 'hadibi1', 'hamerk1',\n       'hartur1', 'helgui', 'hipbab1', 'hoopoe', 'huncis1', 'hunsun2',\n       'joygre1', 'kerspa2', 'klacuc1', 'kvbsun1', 'laudov1', 'lawgol',\n       'lesmaw1', 'lessts1', 'libeat1', 'litegr', 'litswi1', 'litwea1',\n       'loceag1', 'lotcor1', 'lotlap1', 'luebus1', 'mabeat1', 'macshr1',\n       'malkin1', 'marsto1', 'marsun2', 'mcptit1', 'meypar1', 'moccha1',\n       'mouwag1', 'ndcsun2', 'nobfly1', 'norbro1', 'norcro1', 'norfis1',\n       'norpuf1', 'nubwoo1', 'pabspa1', 'palfly2', 'palpri1', 'piecro1',\n       'piekin1', 'pitwhy', 'purgre2', 'pygbat1', 'quailf1', 'ratcis1',\n       'raybar1', 'rbsrob1', 'rebfir2', 'rebhor1', 'reboxp1', 'reccor',\n       'reccuc1', 'reedov1', 'refbar2', 'refcro1', 'reftin1', 'refwar2',\n       'rehblu1', 'rehwea1', 'reisee2', 'rerswa1', 'rewsta1', 'rindov',\n       'rocmar2', 'rostur1', 'ruegls1', 'rufcha2', 'sacibi2', 'sccsun2',\n       'scrcha1', 'scthon1', 'shesta1', 'sichor1', 'sincis1', 'slbgre1',\n       'slcbou1', 'sltnig1', 'sobfly1', 'somgre1', 'somtit4', 'soucit1',\n       'soufis1', 'spemou2', 'spepig1', 'spewea1', 'spfbar1', 'spfwea1',\n       'spmthr1', 'spwlap1', 'squher1', 'strher', 'strsee1', 'stusta1',\n       'subbus1', 'supsta1', 'tacsun1', 'tafpri1', 'tamdov1', 'thrnig1',\n       'trobou1', 'varsun2', 'vibsta2', 'vilwea1', 'vimwea1', 'walsta1',\n       'wbgbir1', 'wbrcha2', 'wbswea1', 'wfbeat1', 'whbcan1', 'whbcou1',\n       'whbcro2', 'whbtit5', 'whbwea1', 'whbwhe3', 'whcpri2', 'whctur2',\n       'wheslf1', 'whhsaw1', 'whihel1', 'whrshr1', 'witswa1', 'wlwwar',\n       'wookin1', 'woosan', 'wtbeat1', 'yebapa1', 'yebbar1', 'yebduc1',\n       'yebere1', 'yebgre1', 'yebsto1', 'yeccan1', 'yefcan', 'yelbis1',\n       'yenspu1', 'yertin1', 'yesbar1', 'yespet1', 'yetgre1', 'yewgre1']\n\nlabels = sorted(labels)","metadata":{"papermill":{"duration":0.035353,"end_time":"2022-04-22T06:01:17.283888","exception":false,"start_time":"2022-04-22T06:01:17.248535","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2023-05-24T09:23:23.724309Z","iopub.execute_input":"2023-05-24T09:23:23.724704Z","iopub.status.idle":"2023-05-24T09:23:23.746796Z","shell.execute_reply.started":"2023-05-24T09:23:23.724652Z","shell.execute_reply":"2023-05-24T09:23:23.745876Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = filepaths.raw_dir / 'test_soundscapes'\n\ntest_df = pd.DataFrame(\n     [(path.stem, *path.stem.split(\"_\"), path) for path in path.glob(\"*.ogg\")]*n_samples,\n    columns = [\"filename\", \"name\" ,\"id\", \"path\"]\n)\n    \nprint(test_df.shape)\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:23.749123Z","iopub.execute_input":"2023-05-24T09:23:23.749851Z","iopub.status.idle":"2023-05-24T09:23:23.793142Z","shell.execute_reply.started":"2023-05-24T09:23:23.749792Z","shell.execute_reply":"2023-05-24T09:23:23.791848Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_time = 0","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:23.794412Z","iopub.execute_input":"2023-05-24T09:23:23.794734Z","iopub.status.idle":"2023-05-24T09:23:23.799208Z","shell.execute_reply.started":"2023-05-24T09:23:23.794702Z","shell.execute_reply":"2023-05-24T09:23:23.798123Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"start_time = time.time()\n\naudios = {}\npaths = test_df.path.unique()\nidxs = [i for i in range(len(paths))]\n\ndef read_audios(idx):\n    fp = paths[idx]\n    wav, _ = sf.read(fp, dtype='float32')\n    audios[str(fp)] = wav\n    return True\n\nwith concurrent.futures.ThreadPoolExecutor(max_workers=4) as executor:\n    _ = executor.map(read_audios, idxs)\n    \nelapsed = (time.time() - start_time)*200\ntotal_time += elapsed\nprint('Elapsed time: ', elapsed)","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:23.800786Z","iopub.execute_input":"2023-05-24T09:23:23.801166Z","iopub.status.idle":"2023-05-24T09:23:24.482560Z","shell.execute_reply.started":"2023-05-24T09:23:23.801117Z","shell.execute_reply":"2023-05-24T09:23:24.481344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = load_config(filepaths.models_dir / exp_list[0][0] / 'config.yaml')\n    \nconfig.dataset.train_duration = 5\nconfig.dataset.valid_duration = 5\nconfig.dataset.labels = labels\n\ntest_ds = TestDataset(audios, test_df, config)\ndataloader = torch.utils.data.DataLoader(\n    test_ds,\n    batch_size=8,\n    num_workers=4,\n    shuffle=False,\n    pin_memory=True,\n    drop_last=False,\n)\n\nbatches = [batch for batch in dataloader]\nidxs = [i for i in range(len(batches))]\n\nimgs = batches[0][0]","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:24.484051Z","iopub.execute_input":"2023-05-24T09:23:24.484926Z","iopub.status.idle":"2023-05-24T09:23:25.219940Z","shell.execute_reply.started":"2023-05-24T09:23:24.484885Z","shell.execute_reply":"2023-05-24T09:23:25.218428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from copy import deepcopy","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:25.222059Z","iopub.execute_input":"2023-05-24T09:23:25.222440Z","iopub.status.idle":"2023-05-24T09:23:25.228683Z","shell.execute_reply.started":"2023-05-24T09:23:25.222397Z","shell.execute_reply":"2023-05-24T09:23:25.227382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_config = deepcopy(config)\nnew_config.model.backbone_type = 'tf_efficientnet_b3_ns'\nencoder_prev = TattakaModel(new_config, pretrained=False)\nstate = torch.load('/kaggle/input/exp295/models/fold_0_best.pth', map_location=torch.device('cpu'))\nencoder_prev.load_state_dict(state['model'])\nencoder_prev.model.blocks = encoder_prev.model.blocks[:3]\n","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:25.230211Z","iopub.execute_input":"2023-05-24T09:23:25.230791Z","iopub.status.idle":"2023-05-24T09:23:26.367472Z","shell.execute_reply.started":"2023-05-24T09:23:25.230754Z","shell.execute_reply":"2023-05-24T09:23:26.366418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from copy import deepcopy\n\nnew_config = deepcopy(config)\nnew_config.model.backbone_type = 'tf_efficientnet_b1_ns'\nencoder1 = TattakaModel(new_config, pretrained=False)\nstate = torch.load('/kaggle/input/exp290/models/fold_0_best.pth', map_location=torch.device('cpu'))\nencoder1.load_state_dict(state['model'])\nencoder1.model.blocks = encoder1.model.blocks[:4]\nencoder1.eval()\n\n\nnew_config = deepcopy(config)\nnew_config.model.backbone_type = 'tf_efficientnet_b2_ns'\nencoder2 = TattakaModel(new_config, pretrained=False)\nstate = torch.load('/kaggle/input/exp302/models/fold_0_best.pth', map_location=torch.device('cpu'))\nencoder2.load_state_dict(state['model'])\nencoder2.model.blocks = encoder2.model.blocks[:4]\nencoder2.eval()\n\n\nnew_config = deepcopy(config)\nnew_config.model.backbone_type = 'tf_efficientnet_b3_ns'\nencoder3 = TattakaModel(new_config, pretrained=False)\nstate = torch.load('/kaggle/input/exp303/models/fold_0_best.pth', map_location=torch.device('cpu'))\nencoder3.load_state_dict(state['model'])\nencoder3.model.blocks = encoder3.model.blocks[:4]\nencoder3.eval()\n\n\nnew_config = deepcopy(config)\nnew_config.model.backbone_type = 'tf_efficientnet_b3_ns'\nencoder4 = TattakaModel(new_config, pretrained=False)\nstate = torch.load('/kaggle/input/exp315/models/fold_0_best.pth', map_location=torch.device('cpu'))\nencoder4.load_state_dict(state['model'])\nencoder4.model.blocks = encoder4.model.blocks[:3]\nencoder4.eval()\n\n1","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:26.369226Z","iopub.execute_input":"2023-05-24T09:23:26.373365Z","iopub.status.idle":"2023-05-24T09:23:29.588872Z","shell.execute_reply.started":"2023-05-24T09:23:26.373319Z","shell.execute_reply":"2023-05-24T09:23:29.587750Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"factor = 10**5\nparams = list(encoder_prev.parameters())\nprint('the length of parameters is', len(params))\nfor i in range(len(params)):\n    params[i].data = torch.round(params[i].data*factor) / factor\n    \n    \nparams = list(encoder1.parameters())\nprint('the length of parameters is', len(params))\nfor i in range(len(params)):\n    params[i].data = torch.round(params[i].data*factor) / factor\n    \n    \nparams = list(encoder2.parameters())\nprint('the length of parameters is', len(params))\nfor i in range(len(params)):\n    params[i].data = torch.round(params[i].data*factor) / factor\n    \n    \nparams = list(encoder3.parameters())\nprint('the length of parameters is', len(params))\nfor i in range(len(params)):\n    params[i].data = torch.round(params[i].data*factor) / factor","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:29.593090Z","iopub.execute_input":"2023-05-24T09:23:29.593514Z","iopub.status.idle":"2023-05-24T09:23:29.643980Z","shell.execute_reply.started":"2023-05-24T09:23:29.593477Z","shell.execute_reply":"2023-05-24T09:23:29.642941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"encoder11 = torch.nn.Sequential(encoder1.model.conv_stem, encoder1.model.bn1, encoder1.model.blocks)\nencoder21 = torch.nn.Sequential(encoder2.model.conv_stem, encoder2.model.bn1, encoder2.model.blocks)\nencoder31 = torch.nn.Sequential(encoder3.model.conv_stem, encoder3.model.bn1, encoder3.model.blocks)\nencoder41 = torch.nn.Sequential(encoder4.model.conv_stem, encoder4.model.bn1, encoder4.model.blocks)\nencoder_prev1 = torch.nn.Sequential(encoder_prev.model.conv_stem, encoder_prev.model.bn1, encoder_prev.model.blocks)\n\nparams = list(encoder_prev.parameters())\nprint('the length of parameters is', len(params))\nfor i in range(len(params)):\n    params[i].data = torch.round(params[i].data*factor) / factor\n    \n    \nparams = list(encoder11.parameters())\nprint('the length of parameters is', len(params))\nfor i in range(len(params)):\n    params[i].data = torch.round(params[i].data*factor) / factor\n    \n    \nparams = list(encoder21.parameters())\nprint('the length of parameters is', len(params))\nfor i in range(len(params)):\n    params[i].data = torch.round(params[i].data*factor) / factor\n    \n    \nparams = list(encoder31.parameters())\nprint('the length of parameters is', len(params))\nfor i in range(len(params)):\n    params[i].data = torch.round(params[i].data*factor) / factor\n    \n    \nparams = list(encoder41.parameters())\nprint('the length of parameters is', len(params))\nfor i in range(len(params)):\n    params[i].data = torch.round(params[i].data*factor) / factor\n\nencoder11.eval()\nencoder21.eval()\nencoder31.eval()\nencoder41.eval()\nencoder_prev1.eval()\n\nbi1 = encoder11(imgs)\nbi2 = encoder21(imgs)\nbi3 = encoder31(imgs)\nbi4 = encoder41(imgs)\nbi_prev = encoder_prev1(imgs)","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:29.648731Z","iopub.execute_input":"2023-05-24T09:23:29.651256Z","iopub.status.idle":"2023-05-24T09:23:32.788202Z","shell.execute_reply.started":"2023-05-24T09:23:29.651212Z","shell.execute_reply":"2023-05-24T09:23:32.787091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"start_time = time.time()\n\nmodels = []\nmodels_names = []\n\nfor i, (exp, folds) in enumerate(exp_list):\n    config = load_config(filepaths.models_dir / exp / 'config.yaml')\n    \n    config.dataset.train_duration = 5\n    config.dataset.valid_duration = 5\n    config.dataset.labels = labels\n    \n    for fold in tqdm(folds):\n        if exp in ['exp122', 'exp180', 'exp181', 'exp183', 'exp184', 'exp185', 'exp186', 'exp190', 'exp205']:\n            folder = 'models' if fold != -1 else 'chkp'\n            fn = f'fold_{fold}_best.pth' if fold != -1 else f'fold_{fold}_chkp.pth'\n        else:\n            folder = 'models'\n            fn = f'fold_{fold}_best.pth'\n        state_path = filepaths.models_dir / exp / folder / fn\n\n        print('Weights from: ', state_path)\n\n        model = TattakaModel(config, pretrained=False)\n        state = torch.load(state_path, map_location=torch.device('cpu'))\n         \n        if exp in ['exp294', 'exp295', 'exp303', 'exp300', 'exp321', 'exp331']:\n            model.model.blocks[0] = encoder_prev.model.blocks[0]\n            model.model.blocks[1] = encoder_prev.model.blocks[1]\n            model.model.blocks[2] = encoder_prev.model.blocks[2]\n            model.model.conv_stem = encoder_prev.model.conv_stem\n            model.model.bn1 = encoder_prev.model.bn1\n            \n        elif exp in ['exp315']:\n            model.model.blocks[0] = encoder4.model.blocks[0]\n            model.model.blocks[1] = encoder4.model.blocks[1]\n            model.model.blocks[2] = encoder4.model.blocks[2]\n            model.model.conv_stem = encoder4.model.conv_stem\n            model.model.bn1 = encoder4.model.bn1\n            \n        else:\n            if config.model.backbone_type in ['tf_efficientnet_b0_ns', 'tf_efficientnet_b1_ns']:\n                model.model.blocks[0] = encoder1.model.blocks[0]\n                model.model.blocks[1] = encoder1.model.blocks[1]\n                model.model.blocks[2] = encoder1.model.blocks[2]\n                model.model.blocks[3] = encoder1.model.blocks[3]\n                model.model.conv_stem = encoder1.model.conv_stem\n                model.model.bn1 = encoder1.model.bn1\n\n            elif config.model.backbone_type in ['tf_efficientnet_b2_ns',]:\n                model.model.blocks[0] = encoder2.model.blocks[0]\n                model.model.blocks[1] = encoder2.model.blocks[1]\n                model.model.blocks[2] = encoder2.model.blocks[2]\n                model.model.blocks[3] = encoder2.model.blocks[3]\n                model.model.conv_stem = encoder2.model.conv_stem\n                model.model.bn1 = encoder2.model.bn1\n\n            elif config.model.backbone_type in ['tf_efficientnet_b3_ns',]:\n                model.model.blocks[0] = encoder3.model.blocks[0]\n                model.model.blocks[1] = encoder3.model.blocks[1]\n                model.model.blocks[2] = encoder3.model.blocks[2]\n                model.model.blocks[3] = encoder3.model.blocks[3]\n                model.model.conv_stem = encoder3.model.conv_stem\n                model.model.bn1 = encoder3.model.bn1\n\n        \n        model.load_state_dict(state['model'])\n        head = model.head\n        \n        params = list(model.parameters())\n        print('the length of parameters is', len(params))\n        for i in range(len(params)):\n            params[i].data = torch.round(params[i].data*factor) / factor\n            \n        model.eval()\n        \n        if exp in ['exp290', 'exp306']:\n            model = model.model.blocks[4:]\n            bi = bi1\n        elif exp in ['exp302', 'exp304', 'exp310']:\n            model = model.model.blocks[4:]\n            bi = bi2\n        elif exp in ['exp294', 'exp295', 'exp303', 'exp300', 'exp331']:\n            model = model.model.blocks[3:]\n            bi = bi_prev \n        elif exp in ['exp315',]:\n            model = model.model.blocks[3:]\n            bi = bi4\n        else:\n            model = model.model.blocks[4:]\n            bi = bi3\n\n        model.eval()\n        head_input = model(bi)\n\n        torch.onnx.export(model,                     # model being run\n                          bi,                            # model input (or a tuple for multiple inputs)\n                          f'{exp}_{fold}.onnx',              # where to save the model (can be a file or file-like object)\n                          export_params=True,           # store the trained parameter weights inside the model file\n                          opset_version=12,             # the ONNX version to export the model to\n                          do_constant_folding=True,     # whether to execute constant folding for optimization\n                          input_names = ['input'],      # the model's input names\n                          output_names = ['output'])    # the model's output names\n\n        torch.onnx.export(head,                     # model being run\n                          head_input,                            # model input (or a tuple for multiple inputs)\n                          f'{exp}_{fold}_head.onnx',              # where to save the model (can be a file or file-like object)\n                          export_params=True,           # store the trained parameter weights inside the model file\n                          opset_version=12,             # the ONNX version to export the model to\n                          do_constant_folding=True,     # whether to execute constant folding for optimization\n                          input_names = ['input'],      # the model's input names\n                          output_names = ['output'])    # the model's output names\n        \n#         model = onnx.load(f'{exp}_{fold}.onnx')\n#         model_fp16 = float16.convert_float_to_float16(model)\n# #         model_fp16 = onnx.optimizer.optimize(model_fp16, ['fuse_bn_into_conv'] )\n#         onnx.save(model_fp16, f'{exp}_{fold}.onnx')\n        \n#         model = onnx.load(f'{exp}_{fold}_head.onnx')\n#         model_fp16 = float16.convert_float_to_float16(model)\n# #         model_fp16 = onnx.optimizer.optimize(model_fp16, ['fuse_bn_into_conv'] )\n#         onnx.save(model_fp16, f'{exp}_{fold}_head.onnx')\n    \n#         sess_options = rt.SessionOptions()\n#         sess_options.graph_optimization_level = rt.GraphOptimizationLevel.ORT_ENABLE_ALL\n#         sess_options.optimized_model_filepath = f'{exp}_{fold}.onnx'\n        \n#         session = rt.InferenceSession(f'{exp}_{fold}.onnx', sess_options, providers=['CPUExecutionProvider'])\n        \n#         sess_options.optimized_model_filepath = f'{exp}_{fold}_head.onnx'\n#         session = rt.InferenceSession(f'{exp}_{fold}_head.onnx', sess_options, providers=['CPUExecutionProvider'])","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:32.790105Z","iopub.execute_input":"2023-05-24T09:23:32.790845Z","iopub.status.idle":"2023-05-24T09:23:36.358377Z","shell.execute_reply.started":"2023-05-24T09:23:32.790780Z","shell.execute_reply":"2023-05-24T09:23:36.357343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.onnx.export(encoder_prev1,                     # model being run\n                  imgs,                            # model input (or a tuple for multiple inputs)\n                  f'exp295_enc.onnx',              # where to save the model (can be a file or file-like object)\n                  export_params=True,           # store the trained parameter weights inside the model file\n                  opset_version=12,             # the ONNX version to export the model to\n                  do_constant_folding=True,     # whether to execute constant folding for optimization\n                  input_names = ['input'],      # the model's input names\n                  output_names = ['output'])    # the model's output names\n\n# model = onnx.load(f'exp295_enc.onnx')\n# model_fp16 = float16.convert_float_to_float16(model)\n# # model_fp16 = onnx.optimizer.optimize(model_fp16, ['fuse_bn_into_conv'] )\n# onnx.save(model_fp16, f'exp295_enc.onnx')\n\n\n# sess_options.optimized_model_filepath = f'exp295_enc.onnx'\n# session = rt.InferenceSession(f'exp295_enc.onnx', sess_options, providers=['CPUExecutionProvider'])","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:36.360402Z","iopub.execute_input":"2023-05-24T09:23:36.361271Z","iopub.status.idle":"2023-05-24T09:23:37.704334Z","shell.execute_reply.started":"2023-05-24T09:23:36.361222Z","shell.execute_reply":"2023-05-24T09:23:37.702921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.onnx.export(encoder11,                     # model being run\n                  imgs,                            # model input (or a tuple for multiple inputs)\n                  f'exp290_enc.onnx',              # where to save the model (can be a file or file-like object)\n                  export_params=True,           # store the trained parameter weights inside the model file\n                  opset_version=12,             # the ONNX version to export the model to\n                  do_constant_folding=True,     # whether to execute constant folding for optimization\n                  input_names = ['input'],      # the model's input names\n                  output_names = ['output'])    # the model's output names\n\n# model = onnx.load(f'exp290_enc.onnx')\n# model_fp16 = float16.convert_float_to_float16(model)\n# # model_fp16 = onnx.optimizer.optimize(model_fp16, ['fuse_bn_into_conv'] )\n# onnx.save(model_fp16, f'exp290_enc.onnx')\n\n# sess_options.optimized_model_filepath = f'exp290_enc.onnx'\n# session = rt.InferenceSession(f'exp290_enc.onnx', sess_options, providers=['CPUExecutionProvider'])","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:37.706636Z","iopub.execute_input":"2023-05-24T09:23:37.707019Z","iopub.status.idle":"2023-05-24T09:23:39.371338Z","shell.execute_reply.started":"2023-05-24T09:23:37.706982Z","shell.execute_reply":"2023-05-24T09:23:39.370051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.onnx.export(encoder21,                     # model being run\n                  imgs,                            # model input (or a tuple for multiple inputs)\n                  f'exp302_enc.onnx',              # where to save the model (can be a file or file-like object)\n                  export_params=True,           # store the trained parameter weights inside the model file\n                  opset_version=12,             # the ONNX version to export the model to\n                  do_constant_folding=True,     # whether to execute constant folding for optimization\n                  input_names = ['input'],      # the model's input names\n                  output_names = ['output'])    # the model's output names\n\n# model = onnx.load(f'exp302_enc.onnx')\n# model_fp16 = float16.convert_float_to_float16(model)\n# # model_fp16 = onnx.optimizer.optimize(model_fp16, ['fuse_bn_into_conv'] )\n# onnx.save(model_fp16, f'exp302_enc.onnx')\n\n# sess_options.optimized_model_filepath = f'exp302_enc.onnx'\n# session = rt.InferenceSession(f'exp302_enc.onnx', sess_options, providers=['CPUExecutionProvider'])","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:39.373754Z","iopub.execute_input":"2023-05-24T09:23:39.374224Z","iopub.status.idle":"2023-05-24T09:23:41.022369Z","shell.execute_reply.started":"2023-05-24T09:23:39.374177Z","shell.execute_reply":"2023-05-24T09:23:41.021233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.onnx.export(encoder31,                     # model being run\n                  imgs,                            # model input (or a tuple for multiple inputs)\n                  f'exp303_enc.onnx',              # where to save the model (can be a file or file-like object)\n                  export_params=True,           # store the trained parameter weights inside the model file\n                  opset_version=12,             # the ONNX version to export the model to\n                  do_constant_folding=True,     # whether to execute constant folding for optimization\n                  input_names = ['input'],      # the model's input names\n                  output_names = ['output'])    # the model's output names\n\n# model = onnx.load(f'exp303_enc.onnx')\n# model_fp16 = float16.convert_float_to_float16(model)\n# # model_fp16 = onnx.optimizer.optimize(model_fp16, ['fuse_bn_into_conv'] )\n# onnx.save(model_fp16, f'exp303_enc.onnx')\n\n# sess_options.optimized_model_filepath = f'exp303_enc.onnx'\n# session = rt.InferenceSession(f'exp303_enc.onnx', sess_options, providers=['CPUExecutionProvider'])","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:41.027958Z","iopub.execute_input":"2023-05-24T09:23:41.028542Z","iopub.status.idle":"2023-05-24T09:23:42.885340Z","shell.execute_reply.started":"2023-05-24T09:23:41.028503Z","shell.execute_reply":"2023-05-24T09:23:42.884093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.onnx.export(encoder41,                     # model being run\n                  imgs,                            # model input (or a tuple for multiple inputs)\n                  f'exp315_enc.onnx',              # where to save the model (can be a file or file-like object)\n                  export_params=True,           # store the trained parameter weights inside the model file\n                  opset_version=12,             # the ONNX version to export the model to\n                  do_constant_folding=True,     # whether to execute constant folding for optimization\n                  input_names = ['input'],      # the model's input names\n                  output_names = ['output'])    # the model's output names\n\n# model = onnx.load(f'exp303_enc.onnx')\n# model_fp16 = float16.convert_float_to_float16(model)\n# # model_fp16 = onnx.optimizer.optimize(model_fp16, ['fuse_bn_into_conv'] )\n# onnx.save(model_fp16, f'exp303_enc.onnx')\n\n# sess_options.optimized_model_filepath = f'exp303_enc.onnx'\n# session = rt.InferenceSession(f'exp303_enc.onnx', sess_options, providers=['CPUExecutionProvider'])","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:42.887449Z","iopub.execute_input":"2023-05-24T09:23:42.887854Z","iopub.status.idle":"2023-05-24T09:23:44.076466Z","shell.execute_reply.started":"2023-05-24T09:23:42.887779Z","shell.execute_reply":"2023-05-24T09:23:44.075135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"exp_list = [\n#     ['exp122', [0, 1]],\n#     ['exp170', [-1]],\n#     ['exp168', [-1]],\n#     ['exp280', [-1]],\n#     ['exp280', [0]],\n#     ['exp280', [1]],\n#     ['exp280', [2]],\n#     ['exp280', [3]],\n    \n#     ['exp315', [-1]],\n#     ['exp315', [0]],\n#     ['exp315', [1]],\n#     ['exp315', [2]],\n#     ['exp315', [3]],\n    \n#     ['exp281', [-1]],\n#     ['exp281', [0]],\n#     ['exp281', [1]],\n#     ['exp281', [2]],\n#     ['exp281', [3]],\n    \n#     ['exp282', [-1]],\n#     ['exp282', [0]],\n#     ['exp282', [1]],\n#     ['exp282', [2]],\n#     ['exp282', [3]],\n#     ['exp229', [-1]]\n#     ['exp150', [1]],\n#     ['exp311', [-1]],\n#     ['exp325', [-1]],\n#     ['exp297', [-1]],\n#     ['exp326', [-1]],\n#     ['exp327', [-1]],\n#     ['exp328', [-1]],\n#     ['exp330', [-1]],\n#     ['exp332', [-1]],\n#     ['exp333', [-1]],\n     ['exp279', [-1]],\n]\n\nfor i, (exp, folds) in tqdm(enumerate(exp_list)):\n    config = load_config(filepaths.models_dir / exp / 'config.yaml')\n    \n    config.dataset.train_duration = 5\n    config.dataset.valid_duration = 5\n    config.dataset.labels = labels\n    \n    for fold in tqdm(folds):\n        if exp in ['exp122', 'exp180', 'exp181', 'exp183', 'exp184', 'exp185', 'exp186', 'exp190', 'exp205']:\n            folder = 'models' if fold != -1 else 'chkp'\n            fn = f'fold_{fold}_best.pth' if fold != -1 else f'fold_{fold}_chkp.pth'\n        else:\n            folder = 'models'\n            fn = f'fold_{fold}_best.pth'\n        state_path = filepaths.models_dir / exp / folder / fn\n\n        print('Weights from: ', state_path)\n\n        model = TattakaModel(config, pretrained=False)\n        state = torch.load(state_path, map_location=torch.device('cpu'))\n        model.load_state_dict(state['model'])\n        model.eval()\n        \n        params = list(model.parameters())\n        print('the length of parameters is', len(params))\n        for i in range(len(params)):\n            params[i].data = torch.round(params[i].data*10**5) / 10**5\n        \n        if exp in ['exp315']:\n            encoder = torch.nn.Sequential(model.model.conv_stem, model.model.bn1, model.model.blocks[:3])\n            encoder.eval()\n            bi = encoder(imgs)\n            head = model.head\n\n            model = model.model.blocks[3:]\n            model.eval()\n            head_input = model(bi)\n            \n        else:\n            encoder = torch.nn.Sequential(model.model.conv_stem, model.model.bn1, model.model.blocks[:4])\n            encoder.eval()\n            bi = encoder(imgs)\n            head = model.head\n\n            model = model.model.blocks[4:]\n            model.eval()\n            head_input = model(bi)\n        \n        \n        torch.onnx.export(model,                     # model being run\n                          bi,                            # model input (or a tuple for multiple inputs)\n                          f'{exp}_{fold}.onnx',              # where to save the model (can be a file or file-like object)\n                          export_params=True,           # store the trained parameter weights inside the model file\n                          opset_version=12,             # the ONNX version to export the model to\n                          do_constant_folding=True,     # whether to execute constant folding for optimization\n                          input_names = ['input'],      # the model's input names\n                          output_names = ['output'])    # the model's output names\n        \n        torch.onnx.export(encoder,                     # model being run\n                          imgs,                            # model input (or a tuple for multiple inputs)\n                          f'{exp}_{fold}_enc.onnx',              # where to save the model (can be a file or file-like object)\n                          export_params=True,           # store the trained parameter weights inside the model file\n                          opset_version=12,             # the ONNX version to export the model to\n                          do_constant_folding=True,     # whether to execute constant folding for optimization\n                          input_names = ['input'],      # the model's input names\n                          output_names = ['output'])    # the model's output names\n        \n\n        torch.onnx.export(head,                     # model being run\n                          head_input,                            # model input (or a tuple for multiple inputs)\n                          f'{exp}_{fold}_head.onnx',              # where to save the model (can be a file or file-like object)\n                          export_params=True,           # store the trained parameter weights inside the model file\n                          opset_version=12,             # the ONNX version to export the model to\n                          do_constant_folding=True,     # whether to execute constant folding for optimization\n                          input_names = ['input'],      # the model's input names\n                          output_names = ['output'])    # the model's output names2","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:44.079345Z","iopub.execute_input":"2023-05-24T09:23:44.080189Z","iopub.status.idle":"2023-05-24T09:23:48.706552Z","shell.execute_reply.started":"2023-05-24T09:23:44.080147Z","shell.execute_reply":"2023-05-24T09:23:48.705172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!zip -r output.zip /kaggle/working/","metadata":{"execution":{"iopub.status.busy":"2023-05-24T09:23:48.707930Z","iopub.execute_input":"2023-05-24T09:23:48.708264Z","iopub.status.idle":"2023-05-24T09:23:54.006168Z","shell.execute_reply.started":"2023-05-24T09:23:48.708229Z","shell.execute_reply":"2023-05-24T09:23:54.004584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}