{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":4735408,"sourceType":"datasetVersion","datasetId":2740254},{"sourceId":7771371,"sourceType":"datasetVersion","datasetId":4467656},{"sourceId":7777059,"sourceType":"datasetVersion","datasetId":4393152},{"sourceId":7977056,"sourceType":"datasetVersion","datasetId":4393171},{"sourceId":7985832,"sourceType":"datasetVersion","datasetId":4546182},{"sourceId":7991909,"sourceType":"datasetVersion","datasetId":4660858},{"sourceId":8023553,"sourceType":"datasetVersion","datasetId":4728246},{"sourceId":8023560,"sourceType":"datasetVersion","datasetId":4728252},{"sourceId":8023569,"sourceType":"datasetVersion","datasetId":4728260},{"sourceId":8029117,"sourceType":"datasetVersion","datasetId":4711945},{"sourceId":8031876,"sourceType":"datasetVersion","datasetId":4711794},{"sourceId":8044747,"sourceType":"datasetVersion","datasetId":4743458},{"sourceId":8044765,"sourceType":"datasetVersion","datasetId":4743466},{"sourceId":8044781,"sourceType":"datasetVersion","datasetId":4743480},{"sourceId":8046700,"sourceType":"datasetVersion","datasetId":4383650},{"sourceId":8046884,"sourceType":"datasetVersion","datasetId":4383640},{"sourceId":8047773,"sourceType":"datasetVersion","datasetId":4745619},{"sourceId":8049346,"sourceType":"datasetVersion","datasetId":4695597},{"sourceId":8053181,"sourceType":"datasetVersion","datasetId":4643651},{"sourceId":8054161,"sourceType":"datasetVersion","datasetId":4691216},{"sourceId":143704443,"sourceType":"kernelVersion"},{"sourceId":143704594,"sourceType":"kernelVersion"},{"sourceId":167542286,"sourceType":"kernelVersion"},{"sourceId":169351599,"sourceType":"kernelVersion"}],"dockerImageVersionId":30648,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"### ver34. LB==\n       「change team weight」\n        Imanishi:Max:UEMU = 0.28:0.28:0.44　→ 0.37:0.15:0.48 rust cv版\n### ver33. LB==\n       「change team weight」\n        Imanishi:Max:UEMU = 0.28:0.28:0.44　→ 0.20:0.36:0.44\n### ver31. LB==\n       「change team weight」\n        Imanishi:Max:UEMU = 0.28:0.28:0.44　→ 0.3:0.3:0.4\n### ver30. LB==xx\n       「UEMU_kaxumax_TH Part」\n         add models(fold change) & update Emsemble part(exp040)\n### ver28. LB==0.22(update score)\n       「change team weight」\n        Imanishi:Max:UEMU = 0.25:0.25:0.50　→ 0.28:0.28:0.44\n### ver25. LB==0.22(update score)\n       「D.Imanishi Part」\n         hms_submission3 ⇒hms_submission4  update model\n### ver21. LB=0.22(update score)\n       「'before_softmax_averaging'」\n### ver20. \n       「add logits output in Imanishi part&Chen part」\n### ver19. LB=0.22(update score)\n       「change team weight」\n        Imanishi:Max:UEMU = 0.34:0.33:0.33　→ 0.25:0.25:0.50\n### ver15. LB=0.22(update score)\n       「UEMU_kaxumax_TH Part」\n         update kazumax part (add 3models cv=0.26～0.28)&Emsemble part\n### ver14. LB=0.22(update score)\n       「D.Imanishi Part」\n         hms_submission2 ⇒hms_submission3  delete spectrogram only model\n### ver10. LB=0.23(First Team Merge Sub)\n       「UEMU_kaxumax_TH Part」\n         update TH result part (add 4models cv=0.24～0.28)&Emsemble part ","metadata":{}},{"cell_type":"markdown","source":"# D.Imanishi Part","metadata":{}},{"cell_type":"code","source":"!pip install --no-index --find-links=/kaggle/input/effnet-whl efficientnet","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:06:20.205368Z","iopub.execute_input":"2024-04-06T17:06:20.206207Z","iopub.status.idle":"2024-04-06T17:06:32.348036Z","shell.execute_reply.started":"2024-04-06T17:06:20.206163Z","shell.execute_reply":"2024-04-06T17:06:32.347124Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\nimport os\nimport shutil","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:06:32.350617Z","iopub.execute_input":"2024-04-06T17:06:32.351276Z","iopub.status.idle":"2024-04-06T17:06:32.709712Z","shell.execute_reply.started":"2024-04-06T17:06:32.351235Z","shell.execute_reply":"2024-04-06T17:06:32.709006Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!python /kaggle/input/hms-pyfile/hms_submission4_logits.py ./submission_DImanishi.csv\n\nshutil.rmtree(\"/imanishi_tmp\")\nprint(os.listdir('./'))","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:06:32.710690Z","iopub.execute_input":"2024-04-06T17:06:32.711049Z","iopub.status.idle":"2024-04-06T17:09:26.853027Z","shell.execute_reply.started":"2024-04-06T17:06:32.711025Z","shell.execute_reply":"2024-04-06T17:09:26.852027Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Max Chen Part","metadata":{}},{"cell_type":"code","source":"import sys\nimport pandas as pd\n\nsys.path.append(\"/kaggle/input/hms-v6-pipeline-packed\")\n\nfrom hms_pipeline.online_inference import hms_inference","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:09:26.854815Z","iopub.execute_input":"2024-04-06T17:09:26.855561Z","iopub.status.idle":"2024-04-06T17:09:36.332771Z","shell.execute_reply.started":"2024-04-06T17:09:26.855514Z","shell.execute_reply":"2024-04-06T17:09:36.332002Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MAXC_TEST_MODE = True  ### For submission, set to True.\n\n\ndataset_path = \"/kaggle/input/hms-harmful-brain-activity-classification\"\nMAXC_TMP_DIR = \"/kaggle/working/tmp/\" ### Processed input data directory\n\n\nif MAXC_TEST_MODE:\n    infer_df_path = \"/kaggle/input/hms-harmful-brain-activity-classification/test.csv\"\nelse:\n    infer_df_path = \"/kaggle/input/hms-v6-pipeline-packed/hms_pipeline/folds/val_fold_0.csv\"\n    \n    \n### The weight directories of my part. They will probably be updated in the last week. \n### Eventually all the directories here should come from a single kaggle dataset for easy sharing.  \nmaxc_model_dir_list = [\n    \"/kaggle/input/hms-v6-datasets/models-inference/hmsv6-convnexts-w025-3imgs-kl1-bandclip-wl200-1e-6\",\n    \"/kaggle/input/hms-v6-datasets/models-inference/hmsv6-convnexts-w025-3imgs-kl1-bandclip-longrawfull\",\n    \"/kaggle/input/hms-v6-datasets/models-inference/hmsv6-maxvitt-w025-3imgs-kl1-bandclip-wl200-april\",\n    \"/kaggle/input/hms-v6-datasets/models-inference/hmsv6-maxvits-w025-3imgs-kl1-bandclip\",\n    \"/kaggle/input/hms-v6-datasets/models-inference/hmsv6-maxvitt-w025-3imgs-kl1-bandclip-avgmax\",\n]\n\n\ninfer_df = pd.read_csv(infer_df_path)\n\n### Another option for my own quick debug. Can be removed for team submission.\nif not MAXC_TEST_MODE:\n    infer_df = infer_df.head(300)\n    \n    \nmaxc_subm_df, maxc_subm_logits_df = hms_inference(\n    infer_df = infer_df,\n    model_dir_list=maxc_model_dir_list,\n    test_mode=MAXC_TEST_MODE,\n    tmp_dir=MAXC_TMP_DIR,\n    input_data_dir=dataset_path,\n    verbose=False, # Set to false if you don't want to much to be printed\n)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:09:36.335264Z","iopub.execute_input":"2024-04-06T17:09:36.335544Z","iopub.status.idle":"2024-04-06T17:16:03.195339Z","shell.execute_reply.started":"2024-04-06T17:09:36.335519Z","shell.execute_reply":"2024-04-06T17:16:03.194400Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### quick check to see if my output is in a reasonalble range\nmaxc_subm_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:16:03.197316Z","iopub.execute_input":"2024-04-06T17:16:03.197650Z","iopub.status.idle":"2024-04-06T17:16:03.216597Z","shell.execute_reply.started":"2024-04-06T17:16:03.197618Z","shell.execute_reply":"2024-04-06T17:16:03.215699Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"maxc_subm_logits_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:16:03.218071Z","iopub.execute_input":"2024-04-06T17:16:03.218459Z","iopub.status.idle":"2024-04-06T17:16:03.229887Z","shell.execute_reply.started":"2024-04-06T17:16:03.218427Z","shell.execute_reply":"2024-04-06T17:16:03.229020Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"### the dataframe should be ready to merge. Can be saved to a csv if you want.\nmaxc_subm_df.to_csv(\"submission_maxc.csv\",index=False)\nmaxc_subm_logits_df.to_csv(\"subm_logits_maxc.csv\",index=False)\n\nshutil.rmtree(MAXC_TMP_DIR)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:16:03.231293Z","iopub.execute_input":"2024-04-06T17:16:03.231526Z","iopub.status.idle":"2024-04-06T17:16:03.242587Z","shell.execute_reply.started":"2024-04-06T17:16:03.231505Z","shell.execute_reply":"2024-04-06T17:16:03.241624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# UEMU_kaxumax_TH Part","metadata":{}},{"cell_type":"markdown","source":"## UEMU Part","metadata":{}},{"cell_type":"markdown","source":"## utils","metadata":{}},{"cell_type":"code","source":"!mkdir /kaggle/logs\n!mkdir output","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:16:03.243873Z","iopub.execute_input":"2024-04-06T17:16:03.244240Z","iopub.status.idle":"2024-04-06T17:16:05.331374Z","shell.execute_reply.started":"2024-04-06T17:16:03.244208Z","shell.execute_reply":"2024-04-06T17:16:05.330194Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Common","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport gc\nimport importlib\nimport pickle\nimport yaml\nimport glob\nimport sys\nimport os \nimport re\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom matplotlib import pyplot as plt\nimport ast\n\nwith open('/kaggle/input/hms-code/yaml/config_HMS_inference.yaml', 'r') as yml:\n    config_common           = yaml.safe_load(yml)\n# ===============\n# import mycode\n# ===============\nsys.path.append('/kaggle/input/hms-code/src')\n#utils\nfrom utils.utils import pickle_dump,pickle_load,seed_everything,AttrDict,replace_placeholders\nfrom utils.logger import setup_logger, LOGGER\nfrom postprocess.utils_tta import calc_tta,calc_tta_all,find_fold_x_paths\nfrom preprocess.preprocess import read_rawdata\nsys.path.remove('/kaggle/input/hms-code/src')\n    \nconfig_common        = AttrDict(config_common)\n\n#\nDEBUG                = False\nDEVICE               = config_common['train']['DEVICE']\nBATCH_SIZE_Test      = 64 #32\nNUM_WORKERS          = 2  #2\ncol_labels           = config_common['dataset']['col_labels']\n\n#\ndata_dir             = '/kaggle/input/hms-harmful-brain-activity-classification'\noutput_dir           = '/kaggle/working/output'\npath_test_csv        = '/kaggle/input/hms-harmful-brain-activity-classification/test.csv'\npath_test_eegs       = '/kaggle/input/hms-harmful-brain-activity-classification/test_eegs'\npath_test_spec       = '/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms'\nphase                = 'test'\nif DEBUG:\n    path_test_csv    = '/kaggle/input/hms-harmful-brain-activity-classification/train.csv'\n    path_test_eegs   = '/kaggle/input/hms-harmful-brain-activity-classification/train_eegs'\n    path_test_spec   = '/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms'\n    phase            = 'train'","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:16:05.333085Z","iopub.execute_input":"2024-04-06T17:16:05.333414Z","iopub.status.idle":"2024-04-06T17:16:05.678551Z","shell.execute_reply.started":"2024-04-06T17:16:05.333385Z","shell.execute_reply":"2024-04-06T17:16:05.677772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n\n#===read_data===#\ntest_meta                      = pd.read_csv(path_test_csv)\nfiles_test_eegs                = glob.glob(f'{path_test_eegs}/*')\nfiles_test_eegs_names          = [int(os.path.basename(file).split('.')[0]) for file in files_test_eegs]\nfiles_test_spectrograms        = glob.glob(f'{path_test_spec}/*')\nfiles_test_spectrograms_names  = [int(os.path.basename(file).split('.')[0]) for file in files_test_spectrograms]\n\nif DEBUG:\n    test_meta                  = test_meta.groupby('spectrogram_id').first().reset_index()\n    test_meta                  = test_meta[:3000]\n#     test_meta                  = test_meta[:100]\n    \n    files_test_eegs_names_test_meta         = list(test_meta['eeg_id'].unique())\n    files_test_spectrograms_names_test_meta = list(test_meta['spectrogram_id'].unique())\n    files_test_eegs_names                   = list(set(files_test_eegs_names)&set(files_test_eegs_names_test_meta))\n    files_test_spectrograms_names           = list(set(files_test_spectrograms_names)&set(files_test_spectrograms_names_test_meta))\n\n#====Read RawData====#\ninput_type                                  = 'ALL'\nconfig_common['dataset']['data_read_mode']  = 'calc'\n\ntest_eegs_data_dict,len_test_eegs_data_dict,col_eegs,\\\ndict_col_index,list_max_eeg,list_min_eeg,list_nan_eeg_id,\\\ntest_spectrograms_data_dict,len_test_spectrograms_data_dict,\\\ncol_spectrograms,dict_col_index,list_nan_spec_id               = read_rawdata(config_common,data_dir,output_dir,input_type,\\\n                                                                              files_test_eegs_names,files_test_spectrograms_names,\n                                                                              phase=phase)\n","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:16:05.679689Z","iopub.execute_input":"2024-04-06T17:16:05.679961Z","iopub.status.idle":"2024-04-06T17:16:05.818386Z","shell.execute_reply.started":"2024-04-06T17:16:05.679937Z","shell.execute_reply":"2024-04-06T17:16:05.817464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def calc_inference(model_dirs,input_type,model_type ,config_common,\n                   DEVICE,BATCH_SIZE_Test,NUM_WORKERS,\n                   use_tta=False,dict_tta=dict(),weight_tta=[]):\n    exp_names           = [model_dir.split('/')[-1] for model_dir in model_dirs]\n    \n    for idx,model_dir in enumerate(model_dirs):\n        exp_name        = exp_names[idx]\n        print(f'========Inference:{exp_name}========')\n        print(f'{model_dir}/src')\n        sys.path.append(f'{model_dir}/src')\n        #import\n        import trainer.datasets\n        import models.model_HMS_CNN \n        import models.model_HMS_WAVE_CNN\n        importlib.reload(trainer.datasets)\n        importlib.reload(models.model_HMS_CNN)\n        importlib.reload(models.model_HMS_WAVE_CNN)\n        from trainer.datasets import HMS_Dataset\n        from models.model_HMS_CNN import HMSModel_CNN\n        from models.model_HMS_WAVE_CNN import HMSModel_WAVE_CNN\n        \n        import trainer.datasets\n        import models.model_HMS_CNN \n        import models.model_HMS_WAVE_CNN\n        importlib.reload(trainer.datasets)\n        importlib.reload(models.model_HMS_CNN)\n        importlib.reload(models.model_HMS_WAVE_CNN)\n        from trainer.datasets import HMS_Dataset\n        from models.model_HMS_CNN import HMSModel_CNN\n        from models.model_HMS_WAVE_CNN import HMSModel_WAVE_CNN\n        try:\n            import models.model_HMS_MULTI\n            importlib.reload(models.model_HMS_MULTI)\n            from models.model_HMS_MULTI import HMSModel_MULTI\n            \n            import models.model_HMS_MULTI\n            importlib.reload(models.model_HMS_MULTI)\n            from models.model_HMS_MULTI import HMSModel_MULTI\n        except:#旧Ver\n            pass\n        sys.path.remove(f'{model_dir}/src')#remove_path\n        \n        #setting\n        with open(f'{model_dir}/yaml/config_HMS.yaml', 'r') as yml:\n            config_model                = yaml.safe_load(yml)\n        config_model                    = AttrDict(config_model)\n        config_model['dataset']['EEG']['dataset_mode']     = 'calc'\n        config_model['dataset']['EEG_IMG']['dataset_mode'] = 'calc'\n        config_model['dataset']['EEG_IMG']['dataset_load_EXP_ID'] = 'None'\n        \n        if input_type=='ALL':\n            MULTI_calc_spec             = config_model['model']['MULTI']['calc_spec']\n            MULTI_calc_eeg_wave         = config_model['model']['MULTI']['calc_eeg_wave']\n            MULTI_calc_eeg_img          = config_model['model']['MULTI']['calc_eeg_img']\n        \n        #dataset\n        test_dataset                    = HMS_Dataset(test_meta,test_eegs_data_dict,test_spectrograms_data_dict,dict_col_index,\n                                                    config_model,phase='test')\n        test_loader                     = DataLoader(test_dataset, batch_size=BATCH_SIZE_Test, shuffle=False,\n                                                    num_workers=NUM_WORKERS,pin_memory=True)\n\n        #===model===#\n        model_fold_dirs                 = find_fold_x_paths(model_dir)\n        models                          = []\n        for model_fold_dir in model_fold_dirs:\n            model_path                  = f'{model_fold_dir}/best.pt'\n            model                       = torch.load(model_path, map_location=torch.device(DEVICE))\n            model.eval()\n            model.to(DEVICE)\n            models.append(model)\n        \n        #===inference===#\n        try:\n            #eeg_wave\n            calc_select_time2_wave      = config_model['dataset']['EEG']['select_time2_wave']['calc_select_time2_wave']\n            use_time_ave_wave           = config_model['dataset']['EEG']['select_time2_wave']['use_time_ave_wave']\n            weights_time_wave           = config_model['dataset']['EEG']['select_time2_wave']['weights_time_wave']\n            #eeg_img\n            calc_select_time2_img       = config_model['dataset']['EEG_IMG']['select_time2_img']['calc_select_time2_img']\n            use_time_ave_img            = config_model['dataset']['EEG_IMG']['select_time2_img']['use_time_ave_img']\n            weights_time_img            = config_model['dataset']['EEG_IMG']['select_time2_img']['weights_time_img']\n        except:\n            #eeg_wave\n            calc_select_time2_wave      = False\n            use_time_ave_wave           = False\n            weights_time_wave           = []\n            #eeg_img\n            calc_select_time2_img       = False\n            use_time_ave_img            = False\n            weights_time_img            = []\n            \n        with torch.no_grad():\n            for batch,data in enumerate(test_loader):\n                patient_id              = data['patient_id']\n                eeg_id                  = data['eeg_id']\n                spectrogram_id          = data['spectrogram_id']\n                #=read=\n                if input_type=='EEG_WAVE':\n                    feature             = data['sub_eeg_data_wave']\n                elif input_type=='EEG_IMG':\n                    feature             = data['sub_eeg_data_img']\n                elif input_type=='SPEC':\n                    feature             = data['sub_spectrogram_4ch']\n                elif input_type=='ALL':\n                    if MULTI_calc_spec:\n                        feature_spec    = data['sub_spectrogram_4ch']\n                    else:\n                        feature_spec    = []\n                    if MULTI_calc_eeg_img:\n                        feature_eeg_img = data['sub_eeg_data_img']\n                    else:\n                        feature_eeg_img = []\n                    if MULTI_calc_eeg_wave:\n                        feature_eeg_wave= data['sub_eeg_data_wave']\n                    else:\n                        feature_eeg_wave= []\n                #=preds=\n                for idy in range(len(models)):\n                    model          = models[idy]\n                    if model_type == 'CNN':\n                        if (calc_select_time2_img==True)&(use_time_ave_img==True):\n                            print('cccccccccc')\n                            length              = feature.shape[-1]//5\n                            for idt in range(5):\n                                feature_tmp     = feature[:,:,:,:,idt*length:(idt+1)*length].float().to(DEVICE)\n                                (_,logits_tmp)  = model(feature_tmp)\n                                if idt ==0:\n                                    logits      = logits_tmp*weights_time_img[idt]\n                                else:\n                                    logits      += logits_tmp*weights_time_img[idt]\n                        else:\n                            (_,logits)          = model(feature.float().to(DEVICE))\n                            \n                    elif model_type == 'WAVE_CNN':\n                        #(_,logits)      = model(feature,1)\n                        if (calc_select_time2_wave==True)&(use_time_ave_wave==True):\n                            print('dddddddddddddddddd')\n                            length              = 2000\n                            for idt in range(5):\n                                feature_tmp     = feature[:,:,idt*length:(idt+1)*length].float().to(DEVICE)\n                                (_,logits_tmp)  = model(feature_tmp,1)\n                                if idt ==0:\n                                    logits      = logits_tmp*weights_time_wave[idt]\n                                else:\n                                    logits      += logits_tmp*weights_time_wave[idt]\n                        else:\n                            (_,logits)          = model(feature.float().to(DEVICE),1)\n                            \n                    elif model_type == 'MULTI':\n                        if (calc_select_time2_wave==True)&(use_time_ave_wave==True)&\\\n                            (calc_select_time2_img==True)&(use_time_ave_img==True):#waveもimgも分割する\n\n                            length_wave         = 2000\n                            length_img          = feature_eeg_img.shape[-1]//5\n\n                            if MULTI_calc_spec:\n                                feature_spec    = feature_spec.float().to(DEVICE)\n\n                            for idt in range(5):\n                                if MULTI_calc_eeg_img:\n                                    feature_eeg_img_tmp     = feature_eeg_img[:,:,:,:,idt*length_img:(idt+1)*length_img].float().to(DEVICE)\n                                if MULTI_calc_eeg_wave:\n                                    feature_eeg_wave_tmp    = feature_eeg_wave[:,:,idt*length_wave:(idt+1)*length_wave].float().to(DEVICE)\n\n                                (_,logits_tmp)  = model(feature_spec,\n                                                            feature_eeg_wave_tmp,\n                                                            feature_eeg_img_tmp) \n                                if idt ==0:\n                                    logits      = logits_tmp*weights_time_wave[idt]\n                                else:\n                                    logits      += logits_tmp*weights_time_wave[idt]\n\n                        elif (calc_select_time2_wave==True)&(use_time_ave_wave==True)&\\\n                            (calc_select_time2_img==False)&(use_time_ave_img==False):#waveは分割するけど、imgはしない\n                            length_wave         = 2000\n                            if MULTI_calc_spec:\n                                feature_spec    = feature_spec.float().to(DEVICE)\n                            if MULTI_calc_eeg_img:\n                                feature_eeg_img = feature_eeg_img.float().to(DEVICE)\n                            for idt in range(5):\n                                if MULTI_calc_eeg_wave:\n                                    feature_eeg_wave_tmp    = feature_eeg_wave[:,:,idt*length_wave:(idt+1)*length_wave].float().to(DEVICE)                                        \n                                (_,logits_tmp)  = model(feature_spec,\n                                                            feature_eeg_wave_tmp,\n                                                            feature_eeg_img) \n                                if idt ==0:\n                                    logits      = logits_tmp*weights_time_wave[idt]\n                                else:\n                                    logits      += logits_tmp*weights_time_wave[idt]\n\n                        elif (calc_select_time2_wave==False)&(use_time_ave_wave==False)&\\\n                            (calc_select_time2_img==True)&(use_time_ave_img==True):#waveもimgも分割する\n                            length_img          = feature_eeg_img.shape[-1]//5\n                            if MULTI_calc_spec:\n                                feature_spec    = feature_spec.float().to(DEVICE)\n                            if MULTI_calc_eeg_wave:\n                                feature_eeg_wave= feature_eeg_wave.float().to(DEVICE)\n                            for idt in range(5):\n                                if MULTI_calc_eeg_img:\n                                    feature_eeg_img_tmp     = feature_eeg_img[:,:,:,:,idt*length_img:(idt+1)*length_img].float().to(DEVICE)\n\n                                (_,logits_tmp)  = model(feature_spec,\n                                                        feature_eeg_wave,\n                                                        feature_eeg_img_tmp) \n                                if idt ==0:\n                                    logits      = logits_tmp*weights_time_img[idt]\n                                else:\n                                    logits      += logits_tmp*weights_time_img[idt]\n\n                        else:\n                            if MULTI_calc_spec:\n                                feature_spec    = feature_spec.float().to(DEVICE)\n                            if MULTI_calc_eeg_img:\n                                feature_eeg_img = feature_eeg_img.float().to(DEVICE)\n                            if MULTI_calc_eeg_wave:\n                                feature_eeg_wave= feature_eeg_wave.float().to(DEVICE)\n                            (_,logits)      = model(feature_spec,\n                                                        feature_eeg_wave,\n                                                        feature_eeg_img) \n                        \n                    #fold_average\n                    if idy ==0:\n                        logits_fold_ave = logits/len(models)\n                    else:\n                        logits_fold_ave += logits/len(models)\n\n                #=matome in batch=\n                if batch==0:\n                    test_patient_ids     = patient_id.tolist()\n                    test_eeg_ids         = eeg_id.tolist()\n                    test_spectrogram_ids = spectrogram_id.tolist()\n                    test_logits          = logits_fold_ave.float().detach().cpu()\n                else:\n                    test_patient_ids     += patient_id.tolist()\n                    test_eeg_ids         += eeg_id.tolist()\n                    test_spectrogram_ids += spectrogram_id.tolist()\n                    test_logits          = torch.cat([test_logits,logits_fold_ave.float().detach().cpu()],dim=0)\n\n        #=matome in model=\n        col_labels_logits               = [s.split('_')[0]+'_logits' for s in config_common['dataset']['col_labels']]\n        df_test_output                  = pd.DataFrame()\n        df_test_output['eeg_id']        = test_eeg_ids\n        df_test_output['spectrogram_id']      = test_spectrogram_ids\n        df_test_output['patient_id']          = test_patient_ids\n        df_test_output[col_labels_logits]     = test_logits.numpy()\n        df_test_output                        = df_test_output.sort_values(by=['patient_id', 'eeg_id'], ascending=[True, True]).reset_index(drop=True)\n        df_test_output.to_csv(f'/kaggle/working/output/df_test_output_{exp_name}.csv',index=False)\n        \n        del models,test_dataset,test_loader\n        \n        #modules確認\n        module1                         = sys.modules[HMS_Dataset.__module__]\n        module2                         = sys.modules[HMSModel_CNN.__module__]\n        module3                         = sys.modules[HMSModel_WAVE_CNN.__module__]\n        print(module1)\n        print(module2)\n        print(module3)\n        \n        #del modules\n        del_modules                     = ['trainer.datasets',\n                                            'audio.audio_aug','trainer.datasets_utils','trainer.datasets_aug','trainer.datasets_feat',\n                                            'models.model_HMS_CNN','models.model_HMS_WAVE_CNN','models.model_HMS_MULTI']\n        list_modules                    = list(sys.modules.keys())\n        for mod in del_modules:\n            if mod in list_modules:\n                del sys.modules[mod]\n        \n        gc.collect()\n        torch.cuda.empty_cache()  ","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:16:05.820058Z","iopub.execute_input":"2024-04-06T17:16:05.820708Z","iopub.status.idle":"2024-04-06T17:16:05.873288Z","shell.execute_reply.started":"2024-04-06T17:16:05.820668Z","shell.execute_reply":"2024-04-06T17:16:05.872401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dict_tta                           = dict()\nweight_tta                         = []","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:16:05.874437Z","iopub.execute_input":"2024-04-06T17:16:05.874707Z","iopub.status.idle":"2024-04-06T17:16:05.886068Z","shell.execute_reply.started":"2024-04-06T17:16:05.874683Z","shell.execute_reply":"2024-04-06T17:16:05.885254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n#=SPEC=#\ninput_type             = 'SPEC'\nmodel_type             = 'CNN'\nmodel_dirs_spec_cnn    = [\n                            '/kaggle/input/hms-models4/exp323_SPEC_effb0ns_maxxvit_Fold_TH_HvoteData_pretrain_exp321',\n                            '/kaggle/input/hms-models4/exp331_SPEC_effb0ns_MLP_Fold_TH_HvoteData_pretrain_exp330',\n                            '/kaggle/input/hms-models6/exp341_SPEC_effb0ns_maxxvit_Fold_THv2_HvoteData_pretrain_exp340',\n                         ]\n\n\nuse_tta                            = [False]*len(model_dirs_spec_cnn)\ncalc_inference(model_dirs_spec_cnn,input_type,model_type,config_common,\n                DEVICE,BATCH_SIZE_Test,NUM_WORKERS,\n                use_tta=use_tta,dict_tta=dict_tta,weight_tta=weight_tta)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:16:05.890837Z","iopub.execute_input":"2024-04-06T17:16:05.891286Z","iopub.status.idle":"2024-04-06T17:16:17.172037Z","shell.execute_reply.started":"2024-04-06T17:16:05.891259Z","shell.execute_reply":"2024-04-06T17:16:17.171010Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n#=EEG IMG=#\ninput_type               = 'EEG_IMG'\nmodel_type               = 'CNN'\n\nmodel_dirs_eeg_img_cnn   = [\n                            '/kaggle/input/hms-models4/exp537_EEG_IMG_stft_512_24_16ch_effb0ns_maxxvit_Fold_TH_add_aug_HvoteData_pretrain_exp536',\n                            '/kaggle/input/hms-models4/exp539_EEG_IMG_stft_512_96_16ch_effb0ns_MLP_Fold_TH_HvoteData_pretrain_exp538',\n                            '/kaggle/input/hms-models6/exp551_EEG_IMG_stft_512_24_16ch_effb0ns_maxxvit_Fold_THv2_add_aug_HvoteData_pretrain_exp550',\n                            ]\n\n\nuse_tta                   = [False]*len(model_dirs_eeg_img_cnn)\ncalc_inference(model_dirs_eeg_img_cnn,input_type,model_type,config_common,\n                DEVICE,BATCH_SIZE_Test,NUM_WORKERS,\n                use_tta=use_tta,dict_tta=dict_tta,weight_tta=weight_tta)","metadata":{"execution":{"iopub.status.busy":"2024-04-06T17:16:17.173230Z","iopub.execute_input":"2024-04-06T17:16:17.173519Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n#=EEG WAVE=#\ninput_type                     = 'EEG_WAVE'\nmodel_type                     = 'WAVE_CNN'\n\nmodel_dirs_eeg_wave_wave_cnn   = [\n                                    '/kaggle/input/hms-models4/exp453_EEG_WAVE_multi1D_deep1D_cnn2D_Fold_TH_HvoteData_pretrain_exp451',\n                                    '/kaggle/input/hms-models6/exp461_EEG_WAVE_multi1D_deep1D_cnn2D_Fold_THv2_HvoteData_pretrain_exp460',\n                                    ]\n\nuse_tta                        = [False]*len(model_dirs_eeg_wave_wave_cnn)\ncalc_inference(model_dirs_eeg_wave_wave_cnn,input_type,model_type,config_common,\n                DEVICE,BATCH_SIZE_Test,NUM_WORKERS,\n                use_tta=use_tta,dict_tta=dict_tta,weight_tta=weight_tta)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%time\n#=MULTI=#\nBATCH_SIZE_Test                = 64\ninput_type                     = 'ALL'\nmodel_type                     = 'MULTI'\n\nmodel_dirs_multi               = [\n                                    '/kaggle/input/hms-models4/exp963_MULTI_exp321_451_536_logits_out_Fold_TH_HvoteData_pretrain_exp962',\n                                    '/kaggle/input/hms-models6/exp971_MULTI_exp340_460_550_logits_out_Fold_THv2_HvoteData_pretrain_exp970'\n                                    ]\n\nuse_tta                        = [False]*len(model_dirs_multi)\ncalc_inference(model_dirs_multi,input_type,model_type,config_common,\n                DEVICE,BATCH_SIZE_Test,NUM_WORKERS,\n                use_tta=use_tta,dict_tta=dict_tta,weight_tta=weight_tta)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if 'models' in sys.modules:\n    del sys.modules['models']\nif 'trainer' in sys.modules:\n    del sys.modules['trainer']","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## T.H Part","metadata":{}},{"cell_type":"code","source":"!python -m pip install --no-index --find-links=/kaggle/input/base-library omegaconf \n!python -m pip install --no-index --find-links=/kaggle/input/audio-library audiomentations colorednoise nnAudio","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport gc\nimport sys\nimport copy\nimport pywt\nimport argparse\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn.functional as F\nfrom typing import List, Dict\nfrom dataclasses import dataclass, field\nfrom pathlib import Path\nfrom tqdm.notebook import tqdm\nfrom omegaconf import OmegaConf\nfrom torch.utils.data import DataLoader\n\n%load_ext autoreload\n%autoreload 2","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"MODE = \"test\" # \"test\" or \"val\"\nOUTPUT_FOR_ENSEMBLE = True\nDEBUG = False #False #True\nDEBUG_SAMPLE_NUM = 100\nKERNEL = Path(\"/kaggle/input/hms-harmful-brain-activity-classification\").exists()\n\nCLASSES = [\n    \"seizure_vote\",\n    \"lpd_vote\",\n    \"gpd_vote\",\n    \"lrda_vote\",\n    \"grda_vote\",\n    \"other_vote\",\n]\nDATA = (\n    Path(\"/kaggle/input/hms-harmful-brain-activity-classification\")\n    if KERNEL\n    else Path(\"../hms-harmful-brain-activity-classification\")\n)\nOUTPUT = Path(\"./\") if KERNEL else Path(\"../submissions\")\nif not KERNEL:\n    OUTPUT.mkdir(parents=True, exist_ok=True)\nTMP = Path(\"./.tmp\") if KERNEL else None\nCKPTS = Path(\"\") if KERNEL else Path(\"../checkpoints/\")\n\n\nprint(\"========================\")\nprint(\"mode :\", MODE)\nprint(\"env  :\", \"KERNEL\" if KERNEL else \"LOCAL\")\nprint(\"debug:\", DEBUG)\nprint(\"========================\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if KERNEL:\n    sys.path.append(os.path.join(os.getcwd(), \"/kaggle/input/hms-hbac-src\"))\nelse:\n    sys.path.append(os.path.join(os.getcwd(), \"../src\"))\nfrom mydatasets import HMSHBACSpecDataset, HMSSEDDataset, HMS1DDataset\nfrom mydatasets.make_eeg_spectrograms import spectrogram_from_eeg\nfrom augmentations import hms_spec_augmentations, hms_1D_augmentations\nfrom metrics.kaggle_kl_div import score as calc_kl_div\nimport models as MODELS","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if MODE == \"val\":\n    csv_path = DATA / \"train.csv\" if KERNEL else DATA / \"train_fold_irr_mark.csv\"\n    # csv_path = DATA / \"train.csv\"\n    spec_dir_path = DATA / \"train_spectrograms\"\n    eeg_dir_path = DATA / \"train_eegs\"\nelse:\n    csv_path = DATA / \"test.csv\"\n    spec_dir_path = DATA / \"test_spectrograms\"\n    eeg_dir_path = DATA / \"test_eegs\"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@dataclass\nclass CFG:\n    gpu: int = 0\n    batch_size: int = 16\n    n_workers: int = 1\n    k_fold: int = 4\n    checkpoints: List[str] = field(\n        default_factory=lambda: [\n            # [\n            #     \"run-20240211_185534-pu295w5p-KSpec2D_V2_effnetb0_fold_0/epoch=09-val_loss=0.608.ckpt\",\n            #     \"run-20240211_192558-ovc53rmi-KSpec2D_V2_effnetb0_fold_1/epoch=09-val_loss=0.600.ckpt\",\n            #     \"run-20240211_195609-ie2mxejz-KSpec2D_V2_effnetb0_fold_2/epoch=09-val_loss=0.659.ckpt\",\n            #     \"run-20240211_202650-289cskok-KSpec2D_V2_effnetb0_fold_3/epoch=08-val_loss=0.640.ckpt\",\n            # ],\n            # [\n            #     \"run-20240211_185541-xw7lf0ob-ESpec2D_V2_effnetb0_fold_0/epoch=09-val_loss=0.610.ckpt\",\n            #     \"run-20240211_192553-rdgt2uon-ESpec2D_V2_effnetb0_fold_1/epoch=09-val_loss=0.613.ckpt\",\n            #     \"run-20240211_195622-djkb3hxq-ESpec2D_V2_effnetb0_fold_2/epoch=09-val_loss=0.655.ckpt\",\n            #     \"run-20240211_202629-kvyxif7d-ESpec2D_V2_effnetb0_fold_3/epoch=09-val_loss=0.618.ckpt\",\n\n            # ],\n            # wavenet_maxxvitv2n\n            [\n                \"/kaggle/input/hms-hbac-weight/run-20240312_102706-0an6kdp8-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_finetune1e-4-nv10_e30_fold_0/epoch=21-val_metric_kldiv_high_votes=0.256.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240312_134457-fq6753ul-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_finetune1e-4-nv10_e30_fold_1/epoch=08-val_metric_kldiv_high_votes=0.243.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240312_134503-uk71e5ex-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_finetune1e-4-nv10_e30_fold_2/epoch=18-val_metric_kldiv_high_votes=0.239.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240312_155203-xcob14sf-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_finetune1e-4-nv10_e30_fold_3/epoch=23-val_metric_kldiv_high_votes=0.249.ckpt\"\n            ],\n            \n            # wavenet_maxxvitv2n_downsample\n            [\n                \"/kaggle/input/hms-hbac-weight/run-20240315_153328-zrrn4qhx-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_finetune_fold_0/epoch=16-val_metric_kldiv_high_votes=0.249.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240315_153438-rqwbpbfb-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_finetune_fold_1/epoch=25-val_metric_kldiv_high_votes=0.238.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240315_153442-0m77mgei-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_finetune_fold_2/epoch=28-val_metric_kldiv_high_votes=0.239.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240315_153445-ujddmobr-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_finetune_fold_3/epoch=22-val_metric_kldiv_high_votes=0.252.ckpt\",\n            ],\n            # wavenet_maxxvitv2n_downsample_seed0\n            [\n                \"/kaggle/input/hms-hbac-weight/run-20240401_111123-vuot3hc1-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_seed0_finetune_fold_0/epoch=10-val_metric_kldiv_high_votes=0.259.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240401_124844-1abnrom1-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_seed0_finetune_fold_1/epoch=18-val_metric_kldiv_high_votes=0.250.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240401_142450-1ouy8vnm-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_seed0_finetune_fold_2/epoch=27-val_metric_kldiv_high_votes=0.248.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240401_160306-ierghoep-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_seed0_finetune_fold_3/epoch=20-val_metric_kldiv_high_votes=0.252.ckpt\",\n            ],\n            # wavenet_maxxvitv2n_downsample_seed123\n            [\n                \"/kaggle/input/hms-hbac-weight/run-20240401_182516-c3nkrdyq-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_seed123_finetune_fold_0/epoch=24-val_metric_kldiv_high_votes=0.258.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240401_123308-ycz11bsh-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_seed123_finetune_fold_1/epoch=26-val_metric_kldiv_high_votes=0.240.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240401_140851-akcr7s7k-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_seed123_finetune_fold_2/epoch=19-val_metric_kldiv_high_votes=0.240.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240401_154720-i5m93pz4-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_seed123_finetune_fold_3/epoch=28-val_metric_kldiv_high_votes=0.251.ckpt\",\n            ],\n            # wavenet_effnetb4_downsample\n            [\n                \"/kaggle/input/hms-hbac-weight/run-20240401_111124-9y3ixw3i-1D_cls_RTpIcS_LS_Wavenet-effnetb4_1e-3_standard_e50_warmup_downwample_finetune_fold_0/epoch=11-val_metric_kldiv_high_votes=0.275.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240401_124721-6uxqz0en-1D_cls_RTpIcS_LS_Wavenet-effnetb4_1e-3_standard_e50_warmup_downwample_finetune_fold_1/epoch=07-val_metric_kldiv_high_votes=0.269.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240401_142250-lsohp5jh-1D_cls_RTpIcS_LS_Wavenet-effnetb4_1e-3_standard_e50_warmup_downwample_finetune_fold_2/epoch=06-val_metric_kldiv_high_votes=0.264.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240401_160029-fdugba01-1D_cls_RTpIcS_LS_Wavenet-effnetb4_1e-3_standard_e50_warmup_downwample_finetune_fold_3/epoch=07-val_metric_kldiv_high_votes=0.268.ckpt\",\n            ],\n            # wavenet_maxxvits_downsample\n            [\n                \"/kaggle/input/hms-hbac-weight/run-20240401_111126-xo5scl7d-1D_cls_RTpIcS_LS_Wavenet-maxxvits_1e-3_standard_e50_warmup_downwample_finetune_fold_0/epoch=16-val_metric_kldiv_high_votes=0.256.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240401_125917-uay4ug48-1D_cls_RTpIcS_LS_Wavenet-maxxvits_1e-3_standard_e50_warmup_downwample_finetune_fold_1/epoch=29-val_metric_kldiv_high_votes=0.232.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240401_144726-czyaecih-1D_cls_RTpIcS_LS_Wavenet-maxxvits_1e-3_standard_e50_warmup_downwample_finetune_fold_2/epoch=21-val_metric_kldiv_high_votes=0.246.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240401_163800-65gz7sk5-1D_cls_RTpIcS_LS_Wavenet-maxxvits_1e-3_standard_e50_warmup_downwample_finetune_fold_3/epoch=08-val_metric_kldiv_high_votes=0.247.ckpt\",\n            ],\n            # wavenet_maxxvitv2n_downsample_foldV2\n            [\n                \"/kaggle/input/hms-hbac-weight/run-20240405_152214-wm5w0ew3-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_foldV2_finetune_fold_0/epoch=28-val_metric_kldiv_high_votes=0.240.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240405_152227-th1rruwk-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_foldV2_finetune_fold_1/epoch=20-val_metric_kldiv_high_votes=0.218.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240405_152228-xhwtkxba-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_foldV2_finetune_fold_2/epoch=25-val_metric_kldiv_high_votes=0.257.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240405_165938-x123rriq-1D_cls_RTpIcS_LS_Wavenet-maxxvitv2n_1e-3_standard_e50_warmup_downwample_foldV2_finetune_fold_3/epoch=21-val_metric_kldiv_high_votes=0.241.ckpt\",\n            ],\n            # wavenet_effnetb4_downsample_foldV2\n            [\n                \"/kaggle/input/hms-hbac-weight/run-20240405_185147-8gjfntzq-1D_cls_RTpIcS_LS_Wavenet-effnetb4_1e-3_standard_e15_warmup_downwample_foldV2_finetune_fold_0/epoch=06-val_metric_kldiv_high_votes=0.272.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240405_202552-4q5x7wi1-1D_cls_RTpIcS_LS_Wavenet-effnetb4_1e-3_standard_e15_warmup_downwample_foldV2_finetune_fold_1/epoch=08-val_metric_kldiv_high_votes=0.240.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240405_203742-j6v7ia0d-1D_cls_RTpIcS_LS_Wavenet-effnetb4_1e-3_standard_e15_warmup_downwample_foldV2_finetune_fold_2/epoch=09-val_metric_kldiv_high_votes=0.276.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240405_195044-wteo0zct-1D_cls_RTpIcS_LS_Wavenet-effnetb4_1e-3_standard_e15_warmup_downwample_foldV2_finetune_fold_3/epoch=04-val_metric_kldiv_high_votes=0.273.ckpt\"\n                \n            ],\n            # wavenet_maxxvits_downsample_foldV2\n            [\n                \"/kaggle/input/hms-hbac-weight/run-20240405_170053-kadak6et-1D_cls_RTpIcS_LS_Wavenet-maxxvits_1e-3_standard_e50_warmup_downwample_foldV2_finetune_fold_0/epoch=27-val_metric_kldiv_high_votes=0.240.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240405_165949-lbjfygkz-1D_cls_RTpIcS_LS_Wavenet-maxxvits_1e-3_standard_e50_warmup_downwample_foldV2_finetune_fold_1/epoch=27-val_metric_kldiv_high_votes=0.216.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240405_183700-wgsxjdxt-1D_cls_RTpIcS_LS_Wavenet-maxxvits_1e-3_standard_e50_warmup_downwample_foldV2_finetune_fold_2/epoch=27-val_metric_kldiv_high_votes=0.255.ckpt\",\n                \"/kaggle/input/hms-hbac-weight/run-20240405_184943-9u133n8x-1D_cls_RTpIcS_LS_Wavenet-maxxvits_1e-3_standard_e50_warmup_downwample_foldV2_finetune_fold_3/epoch=06-val_metric_kldiv_high_votes=0.250.ckpt\"\n            ],\n#             # wavenet_coatnet0_downsample\n#             [\n#                 \"/kaggle/input/hms-hbac-weight/run-20240403_173155-e8o9lftc-1D_cls_RTpIcS_LS_Wavenet-coatnet0_1e-3_standard_e50_warmup_downwample_finetune_fold_0/epoch=28-val_metric_kldiv_high_votes=0.283.ckpt\",\n#                 \"/kaggle/input/hms-hbac-weight/run-20240403_173156-83p7voi3-1D_cls_RTpIcS_LS_Wavenet-coatnet0_1e-3_standard_e50_warmup_downwample_finetune_fold_1/epoch=19-val_metric_kldiv_high_votes=0.270.ckpt\",\n#                 \"/kaggle/input/hms-hbac-weight/run-20240403_173157-7vqt1svv-1D_cls_RTpIcS_LS_Wavenet-coatnet0_1e-3_standard_e50_warmup_downwample_finetune_fold_2/epoch=11-val_metric_kldiv_high_votes=0.266.ckpt\",\n#                 \"/kaggle/input/hms-hbac-weight/run-20240403_173159-tz8637ay-1D_cls_RTpIcS_LS_Wavenet-coatnet0_1e-3_standard_e50_warmup_downwample_finetune_fold_3/epoch=14-val_metric_kldiv_high_votes=0.273.ckpt\",\n#             ],\n        ]\n    )\n    # checkpointsのリストごとにデータセットをどれ使うか指定する。\n    datasets: List = field(\n        default_factory=lambda: [\n            HMS1DDataset,\n            HMS1DDataset,\n            HMS1DDataset,\n            HMS1DDataset,\n            HMS1DDataset,\n            HMS1DDataset,\n            HMS1DDataset,\n            HMS1DDataset,\n            HMS1DDataset,\n\n        ]\n    ) \n    # checkpointsのリストごとのweightを決定\n    model_weight: List[int] = field(default_factory=lambda: [1, 1, 1, 1, 1, 1, 1, 1, 1])\n    # Ensembleようにlogit等を保存する際のモデル名\n    model_name: List[str] = field(\n        default_factory=lambda: [\n            \"wavenet_maxxvitv2n\", \n            \"wavenet_maxxvitv2n_downsample\", \n            \"wavenet_maxxvitv2n_downsample_seed0\", \n            \"wavenet_maxxvitv2n_downsample_seed123\", \n            \"wavenet_effnetb4_downsample\",\n            \"wavenet_maxxvits_downsample\",\n            \"wavenet_maxxvitv2n_downsample_foldV2\",\n            \"wavenet_effnetb4_downsample_foldV2\",\n            \"wavenet_maxxvits_downsample_foldV2\",\n        ]\n    )\n    # ttaをどれ使うか\n    tta_type: List[str] = field(default_factory=lambda: [])\n    # tta_type: List[str] = field(default_factory=lambda: [\"Inversion\"])\n    # tta_type: List[str] = field(default_factory=lambda: [\"Reverse\"])\n    # tta_type: List[str] = field(default_factory=lambda: [\"ChannelSwap\"])\n    # tta_type: List[str] = field(default_factory=lambda: [\"Inversion\", \"Reverse\", \"ChannelSwap\"])\n\n    csv_path: Path = csv_path\n    spec_dir_path: Path = spec_dir_path\n    eeg_dir_path: Path = eeg_dir_path\n\n    # postprocess\n    scale_coeff: float = 0\n\n\nGLOBAL_CFG = CFG()\n\nassert len(GLOBAL_CFG.checkpoints) == len(GLOBAL_CFG.datasets) == len(GLOBAL_CFG.model_name) == len(GLOBAL_CFG.model_weight)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def allkeys(x):\n    for key, value in x.items():\n        yield key\n        if isinstance(value, dict):\n            for child in allkeys(value):\n                yield key + \".\" + child\n\n\ndef check_dotlist(cfg, dotlist):\n    cfg_dict = OmegaConf.to_container(cfg, resolve=True)\n    cfg_keys = list(allkeys(cfg_dict))\n    dotlist_dict = OmegaConf.to_container(dotlist, resolve=True)\n    dotlist_keys = list(allkeys(dotlist_dict))\n\n    for d_key in dotlist_keys:\n        assert d_key in cfg_keys, f\"{d_key} dosen't exist in config file.\"\n\n\ndef load_configs(checkpoints):\n    configs_list = []\n    checkpoints_list = []\n    for ckpts in checkpoints:\n        configs = []\n        ckpt_fold = []\n        for c in ckpts:\n            c = CKPTS / c\n            conf_path = c.parent / \"train_config.yaml\"\n            conf = OmegaConf.load(conf_path)\n            ckpt_fold.append(c)\n            configs.append(conf)\n        configs_list.append(configs)\n        checkpoints_list.append(ckpt_fold)\n\n    return checkpoints_list, configs_list\n\n\nif DEBUG:\n    ckpts, configs = load_configs(GLOBAL_CFG.checkpoints)\n    for i in range(len(ckpts)):\n        print(ckpts[i][0])\n        print(configs[i][0])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms(config):\n    # transform\n    height = config.height\n    width = config.width\n    augment_args = config.transforms\n    transforms = hms_1D_augmentations(augment_args=augment_args)\n    return transforms\n\n\nif DEBUG:\n    for i in range(len(configs)):\n        transforms = get_transforms(configs[i][0])\n        print(\"audio_val\\n\", transforms[\"audio_val\"])\n        print(\"torch_val\\n\", transforms[\"torch_val\"])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## preprocess","metadata":{}},{"cell_type":"code","source":"# spectrogram_from_eegの処理がdataloader内で行うと遅いので、別で行う。\ndef create_eeg_spec(denoise_wavelet=None):\n    paths_eegs = list(GLOBAL_CFG.eeg_dir_path.iterdir())\n    print(f\"There are {len(paths_eegs)} EEG spectrograms\")\n    all_eegs = {}\n    counter = 0\n    save_dir = TMP / f\"EEG_Spectrograms\"\n    save_dir.mkdir(parents=True, exist_ok=True)\n    for file_path in tqdm(paths_eegs):\n        file_path = str(file_path)\n        eeg_id = file_path.split(\"/\")[-1].split(\".\")[0]\n        save_path = save_dir / f\"{eeg_id}.npy\"\n        eeg_spectrogram = spectrogram_from_eeg(file_path, denoise_wavelet=denoise_wavelet, display=counter < 1)\n        # all_eegs[int(eeg_id)] = eeg_spectrogram\n        np.save(save_path, eeg_spectrogram)\n        counter += 1\n    return save_path.parent","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## model loader","metadata":{}},{"cell_type":"code","source":"def load_weight(checkpoint, net):\n    ckpt = torch.load(checkpoint, map_location=f\"cuda:{GLOBAL_CFG.gpu}\")[\"state_dict\"]\n    ckpt = {k[k.find(\".\") + 1 :]: v for k, v in ckpt.items()}\n    missing_keys, unexpected_keys = net.load_state_dict(ckpt, strict=False)\n    print(f\"\\nload checkpoint: {checkpoint}\\n\")\n    if len(missing_keys) != 0 or len(unexpected_keys) != 0:\n        print(\"====================================\")\n        print(\"missing_keys:\", missing_keys)\n        print(\"unexpecte_keys:\", unexpected_keys)\n        print(\"====================================\")\n    return net\n\n\ndef get_models_from_checkkpoints(ckpt_paths, configs):\n    model_list = []\n    for ckpt, config in zip(ckpt_paths, configs):\n        model_args = copy.deepcopy(config[\"model\"])\n        model_args.load_checkpoint = None\n        for key in model_args.args.keys():\n            if \"pretrained\" in key:\n                model_args.args[key] = False\n        model = getattr(MODELS, model_args.name)(**model_args.args)\n        model.to(GLOBAL_CFG.gpu)\n        model = load_weight(ckpt, model)\n        model.eval()\n        model_list.append(model)\n    return model_list","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## data loader","metadata":{}},{"cell_type":"code","source":"def get_dataloader(data_module, config, transforms, batch_size=1, num_workers=4):\n    dataset_args = copy.deepcopy(config[\"dataset\"])\n    dataset_args.csv_path = str(GLOBAL_CFG.csv_path)\n    dataset_args.spec_dir_path = str(GLOBAL_CFG.spec_dir_path)\n    dataset_args.eeg_dir_path = str(GLOBAL_CFG.eeg_dir_path)\n    eeg_spec_dir_path = dataset_args.get(\"eeg_spec_dir_path\", None)\n    # TODO: modelごとに異なるEEG_specを使う場合は修正の必要あり\n    if eeg_spec_dir_path is not None and KERNEL:\n        dataset_args.eeg_spec_dir_path = str(TMP / \"EEG_Spectrograms\")\n\n    # 推論時は事前作成済みのデータは使わないのでオフにする(configから削除したのでコメントアウト)\n    # dataset_args.spec_npy_path = None\n\n    # pseudo_label用に学習した場合は、Falseにすることで、Testデータを推論するようにする。\n    if hasattr(dataset_args, \"for_pseudo_label\"):\n        dataset_args.for_pseudo_label = False\n    if KERNEL or (MODE == \"test\"):\n        dataset_args.fold = None\n    \n\n    # modeによらずtestモードのデータセットを作成。(csvファイルの全行を予測)\n    dataset = data_module(mode=\"test\", transforms=transforms, **dataset_args)\n\n    loader = DataLoader(\n        dataset=dataset,\n        batch_size=batch_size,\n        shuffle=False,\n        num_workers=num_workers,\n        pin_memory=False,\n    )\n    return loader","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## inference","metadata":{}},{"cell_type":"code","source":"def tta_pred(net, input):\n    preds = []\n    if \"Inversion\" in GLOBAL_CFG.tta_type:\n        y = net(-1 * input)\n        preds.append(y[\"pred\"])\n    if \"Reverse\" in GLOBAL_CFG.tta_type:\n        y = net(torch.flip(input, dims=[1]))\n        preds.append(y[\"pred\"])\n    if \"ChannelSwap\" in GLOBAL_CFG.tta_type:\n        b, t, c = input.shape\n        mid = c // 2\n        input = torch.cat((input[..., mid:], input[..., :mid]), dim=-1)\n        y = net(input)\n        preds.append(y[\"pred\"])\n    preds = torch.stack(preds, dim=0).mean(dim=0)\n    return preds","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_inference_loop(model_list, dataloader, transform=None):\n    \"\"\"test時のループ\n    各foldのモデルで同じデータセットを予測し平均を取り、バッチ方向に連結。\n    \"\"\"\n    pred_list = []\n    logit_list = []\n    with torch.no_grad():\n        for i, batch in enumerate(tqdm(dataloader)):\n            eeg = batch[\"eeg\"].to(GLOBAL_CFG.gpu)\n            kspec = batch.get(\"Kspec\")\n            if (transform is not None) and (kspec is not None):\n                kspec = transform(kspec.to(GLOBAL_CFG.gpu))\n\n            preds = []\n            logits = []\n            for net in model_list:\n                if len(GLOBAL_CFG.tta_type) > 0:\n                    pred = tta_pred(net, eeg)\n                else:\n                    y = net(eeg)\n                    pred = y[\"pred\"]\n                logits.append(pred) # n, b, c\n                pred = F.softmax(pred, dim=1)\n                preds.append(pred) # n, b, c\n            preds = torch.stack(preds, dim=0).mean(0) # b, c\n            logits = torch.stack(logits, dim=0).mean(0) # b, c\n            preds = preds.detach().cpu().numpy()\n            logits = logits.detach().cpu().numpy()\n            pred_list.append(preds)\n            logit_list.append(logits)\n            if DEBUG and i >= DEBUG_SAMPLE_NUM:\n                break        \n\n    pred_arr = np.concatenate(pred_list)\n    logit_arr = np.concatenate(logit_list)\n    del pred_list, logit_list\n    return {\"pred_arr\": pred_arr, \"logit_arr\": logit_arr}\n\ndef run_validatioan_loop(model_list, dataloader_list, transform=None):\n    \"\"\"validation時のループ\n    各foldごとに別のデータセットを予測しすべてをバッチ方向に連結。\n    \"\"\"\n    pred_list = []\n    logit_list = []\n    with torch.no_grad():\n        for net, dataloader in zip(model_list, dataloader_list):\n            for i, batch in enumerate(tqdm(dataloader)):\n                eeg = batch[\"eeg\"].to(GLOBAL_CFG.gpu)\n                kspec = batch.get(\"Kspec\")\n                if (transform is not None) and (kspec is not None):\n                    kspec = transform(kspec.to(GLOBAL_CFG.gpu))\n                if len(GLOBAL_CFG.tta_type) > 0:\n                    pred = tta_pred(net, eeg)\n                else:\n                    y = net(eeg)\n                    pred = y[\"pred\"]\n                logit = pred\n                pred = F.softmax(pred, dim=1)\n                logit = logit.detach().cpu().numpy()\n                pred = pred.detach().cpu().numpy()\n                logit_list.append(logit)\n                pred_list.append(pred)\n                if DEBUG and i >= DEBUG_SAMPLE_NUM:\n                    break        \n\n    pred_arr = np.concatenate(pred_list)\n    logit_arr = np.concatenate(logit_list)\n    del pred_list, logit_list\n    return {\"pred_arr\": pred_arr, \"logit_arr\": logit_arr} ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ckpts_list, configs_list = load_configs(GLOBAL_CFG.checkpoints)\nmodels_list = []\nfor ckpts, configs in zip(ckpts_list, configs_list):\n    models_list.append(get_models_from_checkkpoints(ckpts, configs))\n\n# TODO: モデルごとに必要なeeg_specを作れるようにする。(現状は一つのみ)\n# kernelの場合事前にeeg_specを作っておく。\nif KERNEL and hasattr(configs_list[0][0].dataset, \"eeg_spec_dir_path\"):\n    suffix = configs_list[0][0].dataset.eeg_spec_dir_path.split(\"_\")[-1]\n    # パスの最後にデノイズ方法が書いてあればそれを使う。書いてない場合は使わない\n    denoise_wavelet = suffix if suffix in pywt.wavelist() else None\n    create_eeg_spec(denoise_wavelet=denoise_wavelet)\n\npred = []\nlogit = []\nfor ckpts, configs, models, data_module in zip(ckpts_list, configs_list, models_list, GLOBAL_CFG.datasets):\n    # transformはmodel(5fold)ごとに一つ作成\n    transforms = get_transforms(configs[0])\n    if KERNEL or (MODE == \"test\"): # kernel or テスト時は一つのdataloaderで良い。\n        dataloader = get_dataloader(\n            data_module,\n            configs[0],\n            transforms[\"audio_val\"],\n            batch_size=GLOBAL_CFG.batch_size,\n            num_workers=GLOBAL_CFG.n_workers,\n        )\n        result = run_inference_loop(\n            models, dataloader, transform=transforms[\"torch_val\"]\n        )\n        pred_arr = result[\"pred_arr\"]\n        logit_arr = result[\"logit_arr\"]\n    else: # validatiaon時はdataloaderをそれぞれ作る。\n        dataloaders = []\n        for config, model in zip(configs, models):\n            dataloader = get_dataloader(\n                data_module,\n                config,\n                transforms[\"audio_val\"],\n                batch_size=GLOBAL_CFG.batch_size,\n                num_workers=GLOBAL_CFG.n_workers,\n            )\n            dataloaders.append(dataloader)\n        result = run_validatioan_loop(\n            models, dataloaders, transform=transforms[\"torch_val\"]\n        )\n        pred_arr = result[\"pred_arr\"]\n        logit_arr = result[\"logit_arr\"]\n        # validation時はdfの情報を持っておく\n        dfs = []\n        for dataloader in dataloaders:\n            if DEBUG:\n                dfs.append(dataloader.dataset.df.iloc[:(DEBUG_SAMPLE_NUM+1)*GLOBAL_CFG.batch_size])\n            else:\n                dfs.append(dataloader.dataset.df)\n        train_df = pd.concat(dfs).reset_index(drop=True)\n\n    pred.append(pred_arr)\n    logit.append(logit_arr)\n\npred_average = np.average(pred, axis=0, weights=GLOBAL_CFG.model_weight)\n\ndel models_list\ntorch.cuda.empty_cache()\ngc.collect()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def postprocess(pred, scale_coef=0.05):\n    \"\"\"ソフトマックス関数の出力を0.5に近づける調整\"\"\"\n    # 0.5からの距離に基づいて調整\n    adjusted = pred + (0.5 - pred) * scale_coef  # 調整係数\n    # 正規化して総和を1に保つ\n    adjusted_normalized = adjusted / adjusted.sum(axis=1, keepdims=True)\n    return adjusted_normalized\n\nif GLOBAL_CFG.scale_coeff > 0:\n    pred = postprocess(pred, scale_coef=GLOBAL_CFG.scale_coeff)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## submit","metadata":{}},{"cell_type":"code","source":"# test.csvにはeeg_idとspectrogram_id, patient_idが格納されている。\n# sample_submissionにはeeg_id列とターゲットの列しか無いので、testに予測結果をマージしてからsample_subにマージ\n\ndef calc_metric(pred_df, columns):\n    solution = pred_df.loc[:, [\"eeg_id\"] + CLASSES]\n    submission = pred_df.loc[:, [\"eeg_id\"] + columns].rename(\n        columns={c: C for c, C in zip(columns, CLASSES)}\n    )\n    cv = calc_kl_div(\n        solution=solution, submission=submission, row_id_column_name=\"eeg_id\"\n    )\n    return cv\n\ndef save_sub(output_path, pred, df, sample_submission):\n    pred_df = pd.DataFrame(pred, columns=CLASSES)\n    pred_df = pd.concat([df[[\"eeg_id\"]], pred_df], axis=1)\n    if not DEBUG:#sample_submissionの順番に合わせる処理\n        pred_df = pd.merge(sample_submission[[\"eeg_id\"]], pred_df, on=\"eeg_id\", how=\"left\")\n    pred_df.to_csv(output_path, index=False)\n    # pred_df.head()\n    \n\nif KERNEL or MODE == \"test\":\n    df = pd.read_csv(GLOBAL_CFG.csv_path)\n    smpl_sub = pd.read_csv(DATA / \"sample_submission.csv\")\n\n    # submissionの保存\n    save_sub(OUTPUT/\"submission.csv\", pred_average, df, smpl_sub)\n\n    # ensemble用に各モデルのprobとlogitを保存\n    if OUTPUT_FOR_ENSEMBLE:\n        for i, (p, l) in enumerate(zip(pred, logit)):\n            save_sub(OUTPUT / f\"{GLOBAL_CFG.model_name[i]}_prob.csv\", p, df, smpl_sub)\n            save_sub(OUTPUT / f\"{GLOBAL_CFG.model_name[i]}_logit.csv\", l, df, smpl_sub)\n\n    # pred_df = pd.DataFrame(pred_average, columns=CLASSES)\n    # pred_df = pd.concat([df[[\"eeg_id\"]], pred_df], axis=1)\n\n    # if not DEBUG: #sample_submissionの順に合わせる処理\n    #     pred_df = pd.merge(\n    #         smpl_sub[[\"eeg_id\"]], pred_df, on=\"eeg_id\", how=\"left\"\n    #     )\n\n    # pred_df.to_csv(OUTPUT / \"submission.csv\", index=False)\n    # pred_df.head()\n\nelse:\n    columns = [\"pred_\" + c for c in CLASSES]\n    pred_df = pd.DataFrame(pred_average, columns=columns)\n    pred_df = pd.concat([train_df, pred_df], axis=1)\n    cv = calc_metric(pred_df, columns)\n    high_vote_cv = calc_metric(pred_df[pred_df[\"n_votes\"] >= 10], columns)\n    low_vote_cv = calc_metric(pred_df[pred_df[\"n_votes\"] < 10], columns)\n    print(\"CV_kldiv             : \", cv)\n    print(\"CV_kldiv_high_votes  : \", high_vote_cv)\n    print(\"CV_kldiv_low_votes   : \", low_vote_cv)\n    \n    if MODE == \"val\" and not DEBUG:\n        pred_df.to_csv(OUTPUT / \"prediction.csv\", index=False)\n        pred_df.head()\n\n    # モデルごとの確率値とlogitを保存\n    for i, (p, l) in enumerate(zip(pred, logit)):\n        cols_p = [\"prob_\" + c for c in CLASSES]\n        cols_l = [\"logit_\" + c for c in CLASSES]\n        sub_prob_df = pd.DataFrame(p, columns=cols_p)\n        sub_logit_df = pd.DataFrame(l, columns=cols_l)\n        sub_prob_df = pd.concat([train_df, sub_prob_df], axis=1)\n        sub_logit_df = pd.concat([train_df, sub_logit_df], axis=1)\n        if MODE == \"val\" and not DEBUG:\n            sub_prob_df.to_csv(OUTPUT / f\"prob_{GLOBAL_CFG.model_name[i]}_oof.csv\", index=False)\n            sub_logit_df.to_csv(OUTPUT / f\"logit_{GLOBAL_CFG.model_name[i]}_oof.csv\", index=False)\n            pred_df.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -r {TMP}/*\n!ls {TMP}","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"##  Kazumax Part","metadata":{}},{"cell_type":"code","source":"!python -m pip install --no-index --find-links=/kaggle/input/download-omegaconf omegaconf\n!python -m pip install -U --no-index --find-links=/kaggle/input/download-timm timm\n# try:\n#     import omegaconf\n# except:\n#     !python -m pip install --no-index --find-links=/kaggle/input/download-omegaconf omegaconf\n#     !python -m pip install -U --no-index --find-links=/kaggle/input/download-timm timm","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport shutil\n# librosaでは1コアしか利用しないように制限することでバッティングを防ぐ\n# librosaではOMP_NUM_THREADSかOPENBLAS_NUM_THREADSを設定すれば良いようだが\n# ライブラリに寄って異なる様で、後学のために関係ありそうな変数を全て1に設定しておく\n# https://chat.openai.com/share/474b7ab7-3164-47a7-9423-a5f053a21d2b\n# os.environ[\"OMP_NUM_THREADS\"] = \"1\"  # OpenMPで使用するスレッド数を制限\n# os.environ[\"OPENBLAS_NUM_THREADS\"] = \"1\"  # OpenBLASで使用するスレッド数を制限\n# os.environ[\"MKL_NUM_THREADS\"] = \"1\"  # MKLで使用するスレッド数を制限\n\nimport argparse\nimport ast\nimport contextlib\nimport copy\nimport math\nimport random\nfrom collections import defaultdict\nfrom contextlib import contextmanager\nfrom pathlib import Path\nfrom typing import List, Optional, Tuple, Union\n\nimport albumentations as A\nimport cv2\nimport joblib\nimport librosa\nimport numpy as np\nimport pandas\nimport pandas as pd\nimport polars as pl\nimport pywt\nimport sklearn\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torchvision.transforms as T\nimport torchvision.transforms.functional as TF\nimport torchvision.transforms.v2 as Tv2\nimport yaml\nfrom albumentations.core.transforms_interface import ImageOnlyTransform\nfrom albumentations.pytorch.transforms import ToTensorV2\nfrom joblib import Parallel, delayed\nfrom omegaconf import OmegaConf\nfrom timm import models\nfrom torch.nn.parameter import Parameter\nfrom torch.utils.data import DataLoader, Dataset\nfrom torchmetrics.functional import f1_score\nfrom tqdm import tqdm","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"Path(\"/kaggle/kazumax_tmp\").mkdir(exist_ok=True, parents=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"target_columns = [\n    \"seizure_vote\",\n    \"lpd_vote\",\n    \"gpd_vote\",\n    \"lrda_vote\",\n    \"grda_vote\",\n    \"other_vote\",\n]\nprediction_columns = [f\"{c}_prediction\" for c in target_columns]\nlogit_columns = [f\"{c}_logit\" for c in target_columns]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@contextlib.contextmanager\ndef tqdm_joblib(total, **kwargs):\n    \"\"\"\n    https://yururi-do.com/use-joblib-and-tqdm-to-display-progress-bar-at-batch-level/\n    \"\"\"\n    progress_bar = tqdm(total=total, smoothing=0, **kwargs)\n\n    class TqdmBatchCompletionCallBack(joblib.parallel.BatchCompletionCallBack):\n        def __call__(self, *args, **kwargs):\n            progress_bar.update(n=self.batch_size)\n\n            return super().__call__(*args, **kwargs)\n\n    old_batch_callback = joblib.parallel.BatchCompletionCallBack\n    joblib.parallel.BatchCompletionCallBack = TqdmBatchCompletionCallBack\n\n    try:\n        yield progress_bar\n    finally:\n        joblib.parallel.BatchCompletionCallBack = old_batch_callback\n        progress_bar.close()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def make_specs(config):\n    if KAZUMAX_DEBUG and KAZUMAX_DEBUG_DATASET == \"train\":\n        if (Path(config[\"p_kaggle_specs_output_root\"]) / \"specs.npy\").exists():\n            # debugモードでの2回目以降の実行時で、既にファイルを作ってある場合はスキップ\n            print(\"DEBUG MODE. skip make_specs.\")\n        else:\n            # debugモードでの初回実行時は、先頭のN件分作成\n            df = pd.read_csv(config[\"csv_path\"])\n            df = df[:KAZUMAX_N_DEBUG]\n            spectrograms = {}\n            for spectrogram_id in tqdm(df.spectrogram_id, desc=\"make_specs\"):\n                f = Path(config[\"p_kaggle_specs_root\"]) / f\"{spectrogram_id}.parquet\"\n                tmp = pd.read_parquet(f)\n                spectrograms[spectrogram_id] = tmp.iloc[:,1:].values\n            np.save(Path(config[\"p_kaggle_specs_output_root\"]) / \"specs.npy\", spectrograms)\n    else:\n        files = list(Path(config[\"p_kaggle_specs_root\"]).glob(\"*.parquet\"))\n        print(f'There are {len(files)} spectrogram parquets')\n        spectrograms = {}\n        for f in tqdm(files, desc=\"make_specs\"):\n            tmp = pd.read_parquet(f)\n            file_id = f.stem\n            spectrograms[int(file_id)] = tmp.iloc[:,1:].values\n        np.save(Path(config[\"p_kaggle_specs_output_root\"]) / \"specs.npy\", spectrograms)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def maddest(d, axis=None):\n    return np.mean(np.absolute(d - np.mean(d, axis)), axis)\n\n\ndef denoise(x, wavelet=\"haar\", level=1):\n    coeff = pywt.wavedec(x, wavelet, mode=\"per\")\n    sigma = (1 / 0.6745) * maddest(coeff[-level])\n\n    uthresh = sigma * np.sqrt(2 * np.log(len(x)))\n    coeff[1:] = (\n        pywt.threshold(i, value=uthresh, mode=\"hard\") for i in coeff[1:]\n    )\n\n    ret = pywt.waverec(coeff, wavelet, mode=\"per\")\n    return ret\n\n\ndef fill_nans(x):\n    m = np.nanmean(x)\n    if np.isnan(x).mean() < 1:\n        x = np.nan_to_num(x, nan=m)\n    else:\n        x[:] = 0\n    return x\n\n\ndef to_spec(\n    x,\n    n_hop,\n    n_fft,\n    fmax,\n    win_length,\n    use_wavelet,\n    audio_transforms,\n):\n    x = fill_nans(x)\n\n    if use_wavelet:\n        x = denoise(x, wavelet=use_wavelet)\n\n    if audio_transforms is not None:\n        x = audio_transforms(x)\n\n    # RAW SPECTROGRAM\n    spec = librosa.stft(\n        y=x,\n        hop_length=len(x) // n_hop,\n        n_fft=n_fft,\n        win_length=win_length,\n    )\n\n    spec = np.abs(spec) ** 2\n\n    width = (spec.shape[1] // 32) * 32\n    spec = spec.astype(np.float32)[:, :width]\n\n    if fmax == 20:\n        spec = spec[:100, :]\n    elif fmax == 70:\n        spec = spec[:350, :]\n    elif fmax == 100:\n        spec = spec[:500, :]\n    elif fmax is None:\n        pass\n\n    spec = cv2.resize(\n        spec,\n        (256, 128),\n        interpolation=cv2.INTER_LINEAR,\n    )\n\n    return spec\n\n\ndef to_mel_spec(\n    x,\n    n_hop,\n    n_fft,\n    n_mels,\n    fmin,\n    fmax,\n    win_length,\n    use_wavelet,\n    audio_transforms,\n):\n    x = fill_nans(x)\n\n    if use_wavelet:\n        x = denoise(x, wavelet=use_wavelet)\n\n    if audio_transforms is not None:\n        x = audio_transforms(x)\n\n    # RAW SPECTROGRAM\n    mel_spec = librosa.feature.melspectrogram(\n        y=x,\n        sr=200,\n        hop_length=len(x) // n_hop,\n        n_fft=n_fft,\n        n_mels=n_mels,\n        fmin=fmin,\n        fmax=fmax,\n        win_length=win_length,\n    )\n\n    # LOG TRANSFORM\n    width = (mel_spec.shape[1] // 32) * 32\n    mel_spec_db = librosa.power_to_db(mel_spec, ref=np.max).astype(np.float32)[\n        :, :width\n    ]\n\n    # STANDARDIZE TO -1 TO 1\n    mel_spec_db = (mel_spec_db + 40) / 40\n    return mel_spec_db\n\n\ndef make_eeg_specs(\n    df,\n    index,\n    spec_type,\n    window_sec,\n    eeg_electrode_solo,\n    eeg_electrode_combination,\n    freq_dict,\n    n_hop,\n    n_fft,\n    n_mels,\n    win_length,\n    audio_transforms,\n    use_wavelet,\n    p_eeg_specs,\n    p_output_root,\n):\n    row = df.iloc[index]\n    eeg_id = row.eeg_id\n    parquet_path = p_eeg_specs / f\"{eeg_id}.parquet\"\n    eeg_org = pd.read_parquet(parquet_path)\n    offset_sec = row.eeg_label_offset_seconds\n    label_id = row.label_id\n    sr = 200  # [Hz]\n\n    eeg_specs = {}\n    eeg_specs_neighbor = {}\n    middle = int((offset_sec + 25) * sr)\n    start = middle - int(window_sec / 2 * sr)\n    end = middle + int(window_sec / 2 * sr)\n    eeg = eeg_org.iloc[start:end]\n\n    for column in eeg_electrode_solo:\n        for freq_name, freqs in freq_dict.items():\n            fmin, fmax = freqs\n            x = eeg[column].values\n            if spec_type == \"mel_spectrogram\":\n                spec = to_mel_spec(\n                    x,\n                    n_hop,\n                    n_fft,\n                    n_mels,\n                    fmin,\n                    fmax,\n                    win_length,\n                    use_wavelet,\n                    audio_transforms,\n                )\n            elif spec_type == \"stft\":\n                spec = to_spec(\n                    x,\n                    n_hop,\n                    n_fft,\n                    fmax,\n                    win_length,\n                    use_wavelet,\n                    audio_transforms,\n                )\n            spec = spec.astype(\"float32\")\n            key = f\"{column}_{freq_name}\" if freq_name != \"\" else column\n            eeg_specs[key] = spec\n\n    for i_spec, (spec_name, columns) in enumerate(\n        eeg_electrode_combination.items()\n    ):\n        # VARIABLE TO HOLD SPECTROGRAM\n        if spec_type == \"mel_spectrogram\":\n            img = np.zeros((n_mels, n_hop), dtype=\"float32\")\n        elif spec_type == \"stft\":\n            # img = np.zeros((n_fft // 2 + 1, n_hop), dtype=\"float32\")\n            img = np.zeros((128, 256), dtype=\"float32\")\n        for freq_name, freqs in freq_dict.items():\n            fmin, fmax = freqs\n            for i in range(len(columns) - 1):\n                # COMPUTE PAIR DIFFERENCES\n                x = eeg[columns[i]].values - eeg[columns[i + 1]].values\n                if spec_type == \"mel_spectrogram\":\n                    spec = to_mel_spec(\n                        x,\n                        n_hop,\n                        n_fft,\n                        n_mels,\n                        fmin,\n                        fmax,\n                        win_length,\n                        use_wavelet,\n                        audio_transforms,\n                    )\n                elif spec_type == \"stft\":\n                    spec = to_spec(\n                        x,\n                        n_hop,\n                        n_fft,\n                        fmax,\n                        win_length,\n                        use_wavelet,\n                        audio_transforms,\n                    )\n                img += spec\n                spec_name_neighbor = f\"{columns[i]}-{columns[i + 1]}\"\n                key = (\n                    f\"{spec_name_neighbor}_{freq_name}\"\n                    if freq_name != \"\"\n                    else spec_name_neighbor\n                )\n                eeg_specs_neighbor[key] = spec\n            # AVERAGE THE 4 MONTAGE DIFFERENCES\n            img /= 4.0\n            key = f\"{spec_name}_{freq_name}\" if freq_name != \"\" else spec_name\n            eeg_specs[key] = img\n    eeg_specs |= eeg_specs_neighbor\n    np.save(p_output_root / f\"{label_id}.npy\", eeg_specs)\n\ndef make_eeg_specs_main(p_df, p_eeg_specs, p_eeg_specs_output_root):\n    csv_path = Path(p_df)\n    spec_type = \"stft\"\n    n_hop = 256\n    n_fft = 1024\n    n_mels = 128\n    window_sec = 50\n    freq_dict = {\n        \"\": [0, 20],\n    }\n    win_length = None\n    # fmt: off\n    eeg_electrode_solo = []\n    # fmt: on\n    eeg_electrode_combination = {\n        \"LL\": [\"Fp1\", \"F7\", \"T3\", \"T5\", \"O1\"],\n        \"LP\": [\"Fp1\", \"F3\", \"C3\", \"P3\", \"O1\"],\n        \"RP\": [\"Fp2\", \"F8\", \"T4\", \"T6\", \"O2\"],\n        \"RR\": [\"Fp2\", \"F4\", \"C4\", \"P4\", \"O2\"],\n    }\n    p_eeg_specs = Path(p_eeg_specs)\n    audio_transforms = None\n    use_wavelet = None\n    p_output_root = Path(f\"{p_eeg_specs_output_root}\")\n    p_output_root.mkdir(exist_ok=True, parents=True)\n\n    df = pd.read_csv(csv_path)\n    \n    # test時の特別対応\n    if \"eeg_label_offset_seconds\" not in df.columns:\n        df[\"eeg_label_offset_seconds\"] = 0\n    if \"label_id\" not in df.columns:\n        df[\"label_id\"] = np.arange(len(df))\n\n    if KAZUMAX_DEBUG:\n        df = df[:KAZUMAX_N_DEBUG]\n\n    if KAZUMAX_DEBUG and len(list(Path(p_eeg_specs_output_root).glob(\"**/*.npy\"))) == KAZUMAX_N_DEBUG:\n        # debugモードでの2回目以降の実行時で、既にファイルを作ってある場合はスキップ\n        print(\"DEBUG MODE. skip make_eeg_specs.\")\n    else:\n        # debugモードでの初回実行時は、先頭のN件分作成\n        total = len(df)\n    #     with tqdm_joblib(total=total):\n    #         Parallel(n_jobs=-1, backend=\"multiprocessing\", verbose=0)(\n    #             delayed(make_eeg_specs)(\n    #                 df,\n    #                 index,\n    #                 spec_type,\n    #                 window_sec,\n    #                 eeg_electrode_solo,\n    #                 eeg_electrode_combination,\n    #                 freq_dict,\n    #                 n_hop,\n    #                 n_fft,\n    #                 n_mels,\n    #                 win_length,\n    #                 audio_transforms,\n    #                 use_wavelet,\n    #                 p_eeg_specs,\n    #                 p_output_root,\n    #             )\n    #             for index in range(total)\n    #         )\n\n        for index in tqdm(range(total), total=total, desc=\"make_eeg_specs\"):\n            make_eeg_specs(\n                df,\n                index,\n                spec_type,\n                window_sec,\n                eeg_electrode_solo,\n                eeg_electrode_combination,\n                freq_dict,\n                n_hop,\n                n_fft,\n                n_mels,\n                win_length,\n                audio_transforms,\n                use_wavelet,\n                p_eeg_specs,\n                p_output_root,\n            )","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Compose:\n    def __init__(self, transforms: list):\n        self.transforms = transforms\n\n    def __call__(self, y: np.ndarray):\n        for trns in self.transforms:\n            y = trns(y)\n        return y\n\n\nclass TimeFreqMasking(ImageOnlyTransform):\n    def __init__(\n        self,\n        time_drop_width: int,\n        time_stripes_num: int,\n        freq_drop_width: int,\n        freq_stripes_num: int,\n        always_apply=False,\n        p=0.5,\n    ):\n        super().__init__(always_apply, p)\n        self.time_drop_width = time_drop_width\n        self.time_stripes_num = time_stripes_num\n        self.freq_drop_width = freq_drop_width\n        self.freq_stripes_num = freq_stripes_num\n\n    def apply(self, img, **params):\n        img_ = img.copy()\n        img_ = drop_stripes(\n            img_,\n            dim=0,\n            drop_width=self.freq_drop_width,\n            stripes_num=self.freq_stripes_num,\n        )\n        img_ = drop_stripes(\n            img_,\n            dim=1,\n            drop_width=self.time_drop_width,\n            stripes_num=self.time_stripes_num,\n        )\n\n        return img_\n\n\nclass HMSAugmentations:\n    def __init__(\n        self,\n        ver=\"ver_1\",\n    ):\n        self.ver = ver\n        self.spec_transform_val = A.Compose([], p=1.0)\n        self.audio_transform_val = Compose([])\n        self.spec_transform_test = A.Compose([], p=1.0)\n        self.audio_transform_test = Compose([])\n\n        if self.ver == \"ver_1\":\n            self.spec_transform_train = A.Compose(\n                [],\n                p=1.0,\n            )\n            self.audio_transform_train = Compose([])\n        if self.ver == \"ver_2\":\n            self.spec_transform_train = A.Compose(\n                [\n                    TimeFreqMasking(\n                        time_drop_width=32,\n                        time_stripes_num=2,\n                        freq_drop_width=16,\n                        freq_stripes_num=2,\n                        p=0.5,\n                    )\n                ],\n                p=1.0,\n            )\n            self.audio_transform_train = Compose([])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HMSOffsetDataset(Dataset):\n    def __init__(\n        self,\n        csv_path: str,\n        csv_path_tuh_tusz: str,\n        p_eeg_spec_root: str,\n        p_eeg_spec_root_tuh_tusz: str,\n        specs,\n        height: int,\n        width: int,\n        resize_h: int,\n        resize_w: int,\n        n_channel: int,\n        use_channel: list,\n        use_only_kaggle_specs: bool,\n        use_only_eeg_specs: bool,\n        use_tuh_tusz_data: bool,\n        order: str,\n        group: str,\n        label_smoothing_ver: str,\n        label_smoothing_k: int,\n        label_smoothing_epsilon: float,\n        label_smoothing_n_evaluator: int,\n        pseudo_label_n_evaluator: int,\n        mode: str,\n        fold: int,\n        k_fold: int,\n        spec_transforms=None,\n        audio_transforms=None,\n        p_fill_zero_some_channel: float = 0.0,\n        p_swap_some_channel: float = 0.0,\n        p_shift_time: float = 0.0,\n        fill_zero_max_size: int = 1,\n        swap_version: str = \"ver_3\",\n        shift_max_ratio: int = 8,\n        resize_kaggle_spec: bool = True,\n        dry_run: int = None,\n    ):\n        self.csv_path = csv_path\n        self.csv_path_tuh_tusz = csv_path_tuh_tusz\n        self.p_eeg_spec_root = p_eeg_spec_root\n        self.p_eeg_spec_root_tuh_tusz = p_eeg_spec_root_tuh_tusz\n        self.specs = specs\n        self.height = height\n        self.width = width\n        self.resize_h = resize_h\n        self.resize_w = resize_w\n        self.n_channel = n_channel\n        if order == \"tile\":\n            assert self.n_channel % 4 == 0\n            self.w_grid = self.n_channel // 4\n        self.use_channel = use_channel\n        self.use_only_kaggle_specs = use_only_kaggle_specs\n        self.use_only_eeg_specs = use_only_eeg_specs\n        self.use_tuh_tusz_data = use_tuh_tusz_data\n        self.order = order\n        self.group = group\n        self.label_smoothing_ver = label_smoothing_ver\n        self.label_smoothing_k = label_smoothing_k\n        self.label_smoothing_epsilon = label_smoothing_epsilon\n        self.label_smoothing_n_evaluator = label_smoothing_n_evaluator\n        self.pseudo_label_n_evaluator = pseudo_label_n_evaluator\n        self.mode = mode\n        self.spec_transforms = spec_transforms\n        self.audio_transforms = audio_transforms\n        self.p_fill_zero_some_channel = p_fill_zero_some_channel\n        self.p_swap_some_channel = p_swap_some_channel\n        self.p_shift_time = p_shift_time\n        self.fill_zero_max_size = fill_zero_max_size\n        self.swap_version = swap_version\n        self.shift_max_ratio = shift_max_ratio\n        self.resize_kaggle_spec = resize_kaggle_spec\n        self.sr_eeg = 200  # [Hz]\n        self.sr_spec = 0.5  # [Hz]\n        self.dry_run = dry_run\n        self.target_columns = [\n            \"seizure_vote\",\n            \"lpd_vote\",\n            \"gpd_vote\",\n            \"lrda_vote\",\n            \"grda_vote\",\n            \"other_vote\",\n        ]\n\n        df = pd.read_csv(self.csv_path)\n\n        df[\"p_eeg_spec_root\"] = p_eeg_spec_root\n\n        # test時の特別対応\n        if \"eeg_label_offset_seconds\" not in df.columns:\n            df[\"eeg_label_offset_seconds\"] = 0\n        if \"label_id\" not in df.columns:\n            df[\"label_id\"] = np.arange(len(df))\n\n        self.df = df\n        if dry_run is not None:\n            self.df = self.df[:dry_run]\n\n    def __len__(self):\n        if self.mode not in [\"test\", \"pseudo_label\"]:\n            return len(self.unique_eeg_id)\n        else:\n            return len(self.df)\n\n    def __getitem__(self, index):\n        self.index = index\n        row = self._get_row(index)\n        spectrogram_id = row.spectrogram_id\n        label_id = row.label_id\n        p_eeg_spec_root = Path(row.p_eeg_spec_root)\n\n        if self.mode == \"test\":\n            r = 0\n        else:\n            r = int(spectrogram_offset * self.sr_spec)\n\n        data = np.zeros(\n            (self.height, self.width, self.n_channel), dtype=\"float32\"\n        )\n        for spec_channel in range(4):\n            # img.shape = (100, 300)\n            img = self.specs[spectrogram_id][\n                r : r + 300, spec_channel * 100 : (spec_channel + 1) * 100\n            ].T\n            img = self.normalize9(img)\n\n            if self.resize_kaggle_spec:\n                img = cv2.resize(\n                    img,\n                    (self.width, self.height),\n                    interpolation=cv2.INTER_LINEAR,\n                )\n                data[:, :, spec_channel] = img\n            else:\n                # CROP TO 256 TIME STEPS\n                data[14:-14, :, spec_channel] = img[:, 22:-22] / 2.0\n\n        # EEG SPECTROGRAMS\n        eeg_specs = np.load(\n            p_eeg_spec_root / f\"{label_id}.npy\",\n            allow_pickle=True,\n        ).item()\n        cnt = 0\n        for k, img in eeg_specs.items():\n            img = self.normalize9(img)\n            if self.use_channel is None:\n                data[:, :, cnt + 4] = img\n                cnt += 1\n            else:\n                if k in self.use_channel:\n                    data[:, :, cnt + 4] = img\n                    cnt += 1\n\n        # self.dump_data(data, prefix=\"before\")\n        if self.spec_transforms is not None:\n            data = self.spec_transforms(image=data)[\"image\"]\n        # self.dump_data(data, prefix=\"after\", stop_index=7)\n\n        if self.order == \"tile\":\n            # 4 x w_gridマスに(height, width)を並べて\n            # (4 x height, w_grid x width)とする\n            # (height, width) x 4 -> (4 x height, width)\n            x_specs = [data[:, :, i] for i in range(4)]\n            x_specs = np.concatenate(x_specs, axis=0)\n\n            # (height, width) x N -> (w_grid - 1, 4, height, width)\n            x_eeg_specs = (\n                data[:, :, 4:]\n                .transpose(2, 0, 1)\n                .reshape(self.w_grid - 1, 4, self.height, self.width)\n            )\n            # -> (4, height, (w_grid - 1) x width)\n            x_eeg_specs = np.concatenate(\n                [x_eeg_specs[i] for i in range(self.w_grid - 1)], axis=-1\n            )\n            # -> (4 x height, (w_grid - 1) x width)\n            x_eeg_specs = np.concatenate(\n                [x_eeg_specs[i] for i in range(4)], axis=0\n            )\n\n            if self.use_only_kaggle_specs:\n                # (4 x height, width)\n                x = x_specs\n            elif self.use_only_eeg_specs:\n                # (4 x height, (w_grid - 1) x width)\n                x = x_eeg_specs\n            else:\n                # -> (4 x height, w_grid x width)\n                x = np.concatenate([x_eeg_specs, x_specs], axis=1)\n\n            h, w = x.shape\n            if self.resize_h != h or self.resize_w != w:\n                x = cv2.resize(\n                    x, (self.resize_w, self.resize_h), cv2.INTER_LINEAR\n                )\n            # -> (3, 4 x height, w_grid x width)\n            data = np.stack([x, x, x], axis=0)\n        elif self.order == \"channel\":\n            # (h, w, c) -> (c, h, w)\n            x_specs = data[:, :, :4].transpose(2, 0, 1)\n            x_eeg_specs = data[:, :, 4:].transpose(2, 0, 1)\n\n            if self.use_only_kaggle_specs:\n                data = x_specs\n            elif self.use_only_eeg_specs:\n                data = x_eeg_specs\n            else:\n                data = np.concatenate([x_eeg_specs, x_specs], axis=0)\n        return data\n\n    def _get_row(self, index):\n        if self.mode not in [\"test\", \"pseudo_label\"]:\n            if self.mode == \"train\":\n                try:\n                    return (\n                        self.df[self.df[\"eeg_id\"] == self.unique_eeg_id[index]]\n                        .sample(n=1)\n                        .iloc[0]\n                    )\n                except ValueError as e:\n                    print(\"error: eeg_id = \", self.unique_eeg_id[index])\n                    print(e)\n\n            else:\n                tmp = self.df[\n                    self.df[\"eeg_id\"] == self.unique_eeg_id[index]\n                ].reset_index(drop=True)\n                choise = len(tmp) // 2\n                return tmp.iloc[choise]\n                # return tmp.iloc[0]\n\n        else:\n            return self.df.iloc[index]\n\n    def normalize9(self, img):\n        ep = 1e-6\n        img = np.log(img + ep)\n        img = np.nan_to_num(img, nan=0.0)\n        return img\n\n    def dump_data(self, image, label=None, prefix=None, stop_index=None):\n        p_output_root = Path(\"./tmp/aug_albu/\")\n        p_output_root.mkdir(exist_ok=True, parents=True)\n        filename = f\"{self.mode}_{self.index}.npz\"\n        if prefix is not None:\n            filename = f\"{prefix}_\" + filename\n        np.savez(p_output_root / filename, image=image, label=label)\n        if self.index == stop_index:\n            exit()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def mac(x):\n    return F.max_pool2d(x, (x.size(-2), x.size(-1)))\n    # return F.adaptive_max_pool2d(x, (1,1)) # alternative\n\n\ndef gem(x, p=3, eps=1e-6):\n    return F.avg_pool2d(x.clamp(min=eps).pow(p), (x.size(-2), x.size(-1))).pow(\n        1.0 / p\n    )\n    # return F.lp_pool2d(F.threshold(x, eps, eps), p, (x.size(-2), x.size(-1))) # alternative\n\n\nclass MAC(nn.Module):\n    def __init__(self):\n        super(MAC, self).__init__()\n\n    def forward(self, x):\n        return mac(x)\n\n    def __repr__(self):\n        return self.__class__.__name__ + \"()\"\n\n\nclass GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6):\n        super(GeM, self).__init__()\n        self.p = Parameter(torch.ones(1) * p)\n        self.eps = eps\n\n    def forward(self, x):\n        return gem(x, p=self.p, eps=self.eps)\n\n    def __repr__(self):\n        return (\n            self.__class__.__name__\n            + \"(\"\n            + \"p=\"\n            + \"{:.4f}\".format(self.p.data.tolist()[0])\n            + \", \"\n            + \"eps=\"\n            + str(self.eps)\n            + \")\"\n        )\n\n\nclass HMS2DModel(nn.Module):\n    def __init__(\n        self,\n        backbone: str,\n        in_chans: int,\n        pretrained: bool = True,\n        pool_type: str = \"avg\",\n        n_hiddens: int = 512,\n        n_classes: int = 2,\n        drop_path_rate: float = 0.0,\n        drop_rate_backbone: float = 0.0,\n        drop_rate_fc: float = 0.0,\n        manifold_mixup_alpha: float = 0.4,\n    ):\n        super().__init__()\n\n        pretrained_cfg = None\n        if \".\" in backbone:\n            backbone, pretrained_cfg = backbone.split(\".\")\n\n        self.backbone = getattr(models, backbone)(\n            pretrained=pretrained,\n            in_chans=in_chans,\n            drop_rate=drop_rate_backbone,\n            drop_path_rate=drop_path_rate,\n            pretrained_cfg=pretrained_cfg,\n        )\n\n        if pool_type == \"avg\":\n            self.pooling = nn.AdaptiveAvgPool2d(1)\n        elif pool_type == \"max\":\n            self.pooling = nn.AdaptiveMaxPool2d(1)\n        elif pool_type == \"gem\":\n            self.pooling = GeM()\n        elif pool_type == \"mac\":\n            self.pooling = MAC()\n        else:\n            raise KeyError\n\n        self.feat_1 = nn.LazyLinear(n_hiddens)\n        self.feat_2 = nn.LazyLinear(n_classes)\n        self.dropout = nn.Dropout(drop_rate_fc)\n        self.bn = nn.BatchNorm1d(n_hiddens)\n        self.mixup_alpha = manifold_mixup_alpha\n\n    def mixup(self, features):\n        lam = (\n            np.random.beta(self.mixup_alpha, self.mixup_alpha)\n            if self.mixup_alpha > 0\n            else 1\n        )\n\n        index = torch.randperm(features.size()[0]).type_as(features).long()\n\n        features = lam * features + (1 - lam) * features[index]\n        return features, lam, index\n\n    def forward_until_pooling(self, x):\n        x = self.backbone.forward_features(x)\n        if x.dim() == 4:  # CNN family\n            x = self.pooling(x).squeeze(-1).squeeze(-1)\n        if x.dim() == 3:  # ViT family\n            if self.backbone.global_pool == \"avg\":\n                # TODO: The first token is excluded in some models.\n                x = x.mean(dim=1)\n            else:\n                x = x[:, 0, :]\n            if (\n                hasattr(self.backbone, \"fc_norm\")\n                and self.backbone.fc_norm is not None\n            ):\n                x = self.backbone.fc_norm(x)\n\n        return x\n\n    def forward(self, x, manifold_mixup=False):\n        x = self.forward_until_pooling(x)\n\n        if manifold_mixup:\n            x, lam, index = self.mixup(x)\n\n        x = self.feat_2(torch.relu(self.bn(self.feat_1(self.dropout(x)))))\n\n        if manifold_mixup:\n            return x, lam, index\n        else:\n            return x\n\n\n# T.H model\nclass HMS2DModelV3(nn.Module):\n    def __init__(\n        self,\n        backbone: str,\n        neck: str,\n        in_chans: int,\n        ver: str = \"ver_1\",\n        pretrained: bool = True,\n        pool_type: str = \"avg\",\n        n_hiddens: int = 512,\n        n_classes: int = 2,\n        drop_path_rate: float = 0.0,\n        drop_rate_backbone: float = 0.0,\n        drop_rate_fc: float = 0.0,\n        manifold_mixup_alpha: float = 0.4,\n    ):\n        super().__init__()\n\n        self.backbone = self.get_timm_model(\n            backbone,\n            pretrained,\n            in_chans,\n            drop_rate_backbone,\n            drop_path_rate,\n        )\n\n        if pool_type == \"avg\":\n            self.pooling = nn.AdaptiveAvgPool2d(1)\n        elif pool_type == \"max\":\n            self.pooling = nn.AdaptiveMaxPool2d(1)\n        elif pool_type == \"gem\":\n            self.pooling = GeM()\n        elif pool_type == \"mac\":\n            self.pooling = MAC()\n        else:\n            raise KeyError\n\n        self.mixup_alpha = manifold_mixup_alpha\n\n        self.ver = ver\n        if self.ver == \"ver_1\":\n            # 特徴ベクトルを並べるバージョン\n            # (b, c, n_features) -> (b, 1, c, n_features) -> 2D Model\n            self.with_pooling = True\n            self.neck = self.get_timm_model(neck)\n            self.feat_1 = nn.LazyLinear(n_hiddens)\n            self.feat_2 = nn.LazyLinear(n_classes)\n            self.dropout = nn.Dropout(drop_rate_fc)\n            self.bn = nn.BatchNorm1d(n_hiddens)\n        elif self.ver == \"ver_2\":\n            # ver_1にたいして、複数のpooling結果をconcatして作った\n            # 特徴ベクトルを並べるバージョン\n            # (b, c, n_features) -> (b, 1, c, n_features) -> multi pooling\n            # -> (b, 1, c, n_features x 3) -> 2D Model\n            self.with_pooling = True\n            self.neck = self.get_timm_model(neck)\n            self.feat_1 = nn.LazyLinear(n_hiddens)\n            self.feat_2 = nn.LazyLinear(n_classes)\n            self.dropout = nn.Dropout(drop_rate_fc)\n            self.bn = nn.BatchNorm1d(n_hiddens)\n        elif self.ver == \"ver_3\":\n            # 入力解像度に縛りがあるモデルをneckに使うバージョン\n            # backboneとneckによって個別の対応が必要\n            # 例としてefn_b0とmaxxvitv2_nanoの組み合わせ\n            # (b, c, n_features) -> (b, 1, c, n_features) -> multi pooling\n            # -> (b, 1, c, n_features x 3) -> 2D Model\n            self.with_pooling = True\n            assert backbone == \"tf_efficientnet_b0_ns\"\n            assert neck == \"maxxvitv2_nano_rw_256\"\n            self.neck = self.get_timm_model(neck)\n            self.feat_1 = nn.LazyLinear(n_hiddens)\n            self.feat_2 = nn.LazyLinear(n_classes)\n            self.dropout = nn.Dropout(drop_rate_fc)\n            self.bn = nn.BatchNorm1d(n_hiddens)\n        elif self.ver == \"ver_4\":\n            # backboneの途中のfeature mapに対してchannel方向にpoolingした画像を\n            # タイル状に並べるバージョン\n            self.backbone = self.get_timm_model(\n                backbone,\n                pretrained,\n                in_chans,\n                drop_rate_backbone,\n                drop_path_rate,\n                features_only=True,\n            )\n            # backboneをfreezeする場合\n            # for param in self.backbone.parameters():\n            #     param.requires_grad = False\n            self.neck = self.get_timm_model(neck)\n            self.feat_1 = nn.LazyLinear(n_hiddens)\n            self.feat_2 = nn.LazyLinear(n_classes)\n            self.dropout = nn.Dropout(drop_rate_fc)\n            self.bn = nn.BatchNorm1d(n_hiddens)\n        elif self.ver == \"ver_5\":\n            # backboneの途中のfeature mapに対してchannel方向にpoolingした画像を\n            # タイル状に並べるバージョン\n            self.backbone = self.get_timm_model(\n                backbone,\n                pretrained,\n                in_chans,\n                drop_rate_backbone,\n                drop_path_rate,\n                features_only=True,\n            )\n            # backboneをfreezeする場合\n            # for param in self.backbone.parameters():\n            #     param.requires_grad = False\n            self.neck = self.get_timm_model(neck)\n            self.feat_1 = nn.LazyLinear(n_hiddens)\n            self.feat_2 = nn.LazyLinear(n_classes)\n            self.dropout = nn.Dropout(drop_rate_fc)\n            self.bn = nn.BatchNorm1d(n_hiddens)\n        elif self.ver == \"ver_6\":\n            # ver_4をベースにKaggleSpecとEEGSpecを別のbackboneに通すバージョン\n            self.backbone = self.get_timm_model(\n                backbone,\n                pretrained,\n                in_chans,\n                drop_rate_backbone,\n                drop_path_rate,\n                features_only=True,\n            )\n            self.backbone_2 = self.get_timm_model(\n                backbone,\n                pretrained,\n                in_chans,\n                drop_rate_backbone,\n                drop_path_rate,\n                features_only=True,\n            )\n            # backboneをfreezeする場合\n            # for param in self.backbone.parameters():\n            #     param.requires_grad = False\n            self.neck = self.get_timm_model(neck)\n            self.feat_1 = nn.LazyLinear(n_hiddens)\n            self.feat_2 = nn.LazyLinear(n_classes)\n            self.dropout = nn.Dropout(drop_rate_fc)\n            self.bn = nn.BatchNorm1d(n_hiddens)\n        elif self.ver == \"ver_7\":\n            # backboneの途中のfeature mapを全部並べるバージョン\n            self.backbone = self.get_timm_model(\n                backbone,\n                pretrained,\n                in_chans,\n                drop_rate_backbone,\n                drop_path_rate,\n                features_only=True,\n            )\n            # backboneをfreezeする場合\n            # for param in self.backbone.parameters():\n            #     param.requires_grad = False\n            self.neck = self.get_timm_model(neck)\n            self.feat_1 = nn.LazyLinear(n_hiddens)\n            self.feat_2 = nn.LazyLinear(n_classes)\n            self.dropout = nn.Dropout(drop_rate_fc)\n            self.bn = nn.BatchNorm1d(n_hiddens)\n\n    def get_timm_model(\n        self,\n        backbone,\n        pretrained=False,\n        in_chans=1,\n        drop_rate_backbone=0.0,\n        drop_path_rate=0.0,\n        features_only=False,\n    ):\n        pretrained_cfg = None\n        if \".\" in backbone:\n            backbone, pretrained_cfg = backbone.split(\".\")\n\n        model = getattr(models, backbone)(\n            pretrained=pretrained,\n            in_chans=in_chans,\n            drop_rate=drop_rate_backbone,\n            drop_path_rate=drop_path_rate,\n            pretrained_cfg=pretrained_cfg,\n            features_only=features_only,\n        )\n        return model\n\n    def mixup(self, features):\n        lam = (\n            np.random.beta(self.mixup_alpha, self.mixup_alpha)\n            if self.mixup_alpha > 0\n            else 1\n        )\n\n        index = torch.randperm(features.size()[0]).type_as(features).long()\n\n        features = lam * features + (1 - lam) * features[index]\n        return features, lam, index\n\n    def forward_until_pooling(self, model, x, with_pooling):\n        x = model.forward_features(x)\n        if x.dim() == 4:  # CNN family\n            if with_pooling:\n                x = self.pooling(x).squeeze(-1).squeeze(-1)\n        if x.dim() == 3:  # ViT family\n            if with_pooling:\n                if model.global_pool == \"avg\":\n                    # TODO: The first token is excluded in some models.\n                    x = x.mean(dim=1)\n                else:\n                    x = x[:, 0, :]\n                if hasattr(model, \"fc_norm\") and model.fc_norm is not None:\n                    x = model.fc_norm(x)\n            else:\n                x = x.unsqueeze(1)\n\n        return x\n\n    def forward(self, x, manifold_mixup=False):\n        # (b, c, h, w) -> (b * c, 1, h, w)\n        b, c, h, w = x.shape\n        assert c > 3\n        x = x.reshape(-1, 1, h, w)\n\n        if self.ver == \"ver_1\":\n            # 特徴ベクトルを並べるバージョン\n            x = self.forward_until_pooling(self.backbone, x, self.with_pooling)\n            if manifold_mixup:\n                x, lam, index = self.mixup(x)\n            # (b * c, f_backbone) -> (b, 1, c, f_backbone)\n            x = x.reshape(b, 1, c, -1)\n            # (b, 1, 256, 256)\n            x = F.interpolate(x, (256, 256), mode=\"bilinear\")\n            # (b, f_neck)\n            x = self.forward_until_pooling(self.neck, x, True)\n            # (b, n_classes)\n            x = self.feat_2(torch.relu(self.bn(self.feat_1(self.dropout(x)))))\n        elif self.ver == \"ver_2\":\n            # ver_1にたいして、複数のpooling結果をconcatして作った\n            # 特徴ベクトルを並べるバージョン\n            # (b, c, n_features) -> (b, 1, c, n_features) -> 2D Model\n            # (b, c, n_features) -> (b, 1, c, n_features) -> multi pooling\n            # -> (b, 1, c, n_features x 3) -> 2D Model\n            x = self.forward_until_pooling(self.backbone, x, self.with_pooling)\n            if manifold_mixup:\n                x, lam, index = self.mixup(x)\n            # (b * c, f_backbone) -> (b, 1, c, f_backbone)\n            x = x.reshape(b, 1, c, -1)\n            x = self.forward_until_pooling(self.neck, x, False)\n            # (b*c, f_backbone)\n            x_avg = F.adaptive_avg_pool2d(x, (1, 1)).squeeze((-2, -1))\n            x_max = F.adaptive_max_pool2d(x, (1, 1)).squeeze((-2, -1))\n            x_gem = (\n                F.adaptive_avg_pool2d(x.clamp(min=1e-6).pow(3), (1, 1))\n                .pow(1 / 3)\n                .squeeze((-2, -1))\n            )\n            # (b, c, f_backbone)\n            x_avg = x_avg.reshape(b, c, -1)\n            x_max = x_max.reshape(b, c, -1)\n            x_gem = x_gem.reshape(b, c, -1)\n            # (b, c, f_backbone x 3)\n            x = torch.cat([x_avg, x_max, x_gem], dim=-1)\n            # x = x.reshape(b, c * 3, -1)\n            # (b, 1, c, f_backbone x 3)\n            x = x.unsqueeze(1)\n            # (b, f_neck)\n            x = self.forward_until_pooling(self.neck, x, True)\n            # (b, n_classes)\n            x = self.feat_2(torch.relu(self.bn(self.feat_1(self.dropout(x)))))\n        elif self.ver == \"ver_3\":\n            # 入力解像度に縛りがあるモデルをneckに使うバージョン\n            # backboneとneckによって個別の対応が必要\n            # (b x 8, 1280) -> (b, 8, 80 x 16) -> (b, 1, 160, 64)\n            # -> (b, 1, 256, 256) -> maxxvitv2_nano_rw_256\n            x = self.forward_until_pooling(self.backbone, x, self.with_pooling)\n            if manifold_mixup:\n                x, lam, index = self.mixup(x)\n            x = x.reshape(b, 2, 4, 80, 16)\n            x = x.permute(0, 1, 3, 2, 4)\n            x = x.reshape(b, 1, 160, 64)\n            # (b, 1, 256, 256)\n            x = F.interpolate(x, (256, 256), mode=\"bilinear\")\n            # (b, f_neck)\n            x = self.forward_until_pooling(self.neck, x, True)\n            # (b, n_classes)\n            x = self.feat_2(torch.relu(self.bn(self.feat_1(self.dropout(x)))))\n        elif self.ver == \"ver_4\":\n            # backboneの途中のfeature mapに対してchannel方向にpoolingした画像を\n            # タイル状に並べるバージョン\n            # (b * c,  f_backbone, mini_h, mini_w)\n            # x = self.backbone(x)[1]\n            # x = self.backbone(x)[2]\n            # x = self.backbone(x)[3]\n            x = self.backbone(x)[4]\n            # (b * c, 1, mini_h, mini_w)\n            x = torch.mean(x, dim=1, keepdims=True)\n            _, _, mini_h, mini_w = x.shape\n            # (b, c, mini_h, mini_w)\n            x = x.reshape(b, c, mini_h, mini_w)\n            if manifold_mixup:\n                x, lam, index = self.mixup(x)\n            x = self.tiling(x)\n            # (b, 1, 256, 256)\n            x = F.interpolate(x, (256, 256), mode=\"nearest\")\n            # (b, f_neck)\n            x = self.forward_until_pooling(self.neck, x, True)\n            # (b, n_classes)\n            x = self.feat_2(torch.relu(self.bn(self.feat_1(self.dropout(x)))))\n        elif self.ver == \"ver_5\":\n            # backboneの途中のfeature mapに対してchannel方向にpoolingした画像を\n            # LL,LP,RP,RRに集約してからタイル状に並べるバージョン\n            # (b * c,  f_backbone, mini_h, mini_w)\n            # x = self.backbone(x)[1]\n            # x = self.backbone(x)[2]\n            # x = self.backbone(x)[3]\n            x = self.backbone(x)[4]\n            # (b * c, 1, mini_h, mini_w)\n            x = torch.mean(x, dim=1, keepdims=True)\n            _, _, mini_h, mini_w = x.shape\n            # (b, c, mini_h, mini_w)\n            x = x.reshape(b, c, mini_h, mini_w)\n            # LL,LP,RP,RRに集約\n            x_1 = x[:, :8, :, :]\n            channel_list = [\n                [8, 9, 10, 11],\n                [12, 13, 14, 15],\n                [16, 17, 18, 19],\n                [20, 21, 22, 23],\n            ]\n            x_2 = torch.stack(\n                [x[:, channel, :, :].mean(dim=1) for channel in channel_list],\n                dim=1,\n            )\n            x = torch.cat([x_1, x_2], dim=1)\n            if manifold_mixup:\n                x, lam, index = self.mixup(x)\n            x = self.tiling(x)\n            # (b, 1, 256, 256)\n            x = F.interpolate(x, (256, 256), mode=\"bilinear\")\n            # (b, f_neck)\n            x = self.forward_until_pooling(self.neck, x, True)\n            # (b, n_classes)\n            x = self.feat_2(torch.relu(self.bn(self.feat_1(self.dropout(x)))))\n        elif self.ver == \"ver_6\":\n            # ver_4をベースにKaggleSpecとEEGSpecを別のbackboneに通すバージョン\n            x = x.reshape(b, c, h, w)\n            x_kaggle_spec = x[:, :4, :, :]\n            x_kaggle_spec = x_kaggle_spec.reshape(-1, 1, h, w)\n            # (b * c,  f_backbone, mini_h, mini_w)\n            # x = self.backbone(x)[1]\n            # x = self.backbone(x)[2]\n            # x = self.backbone(x)[3]\n            x_kaggle_spec = self.backbone(x_kaggle_spec)[4]\n            # (b * c, 1, mini_h, mini_w)\n            x_kaggle_spec = torch.mean(x_kaggle_spec, dim=1, keepdims=True)\n            _, _, mini_h, mini_w = x_kaggle_spec.shape\n            # (b, 4, mini_h, mini_w)\n            x_kaggle_spec = x_kaggle_spec.reshape(b, 4, mini_h, mini_w)\n\n            x_eeg_spec = x[:, 4:, :, :]\n            x_eeg_spec = x_eeg_spec.reshape(-1, 1, h, w)\n            # (b * c,  f_backbone, mini_h, mini_w)\n            # x = self.backbone(x)[1]\n            # x = self.backbone(x)[2]\n            # x = self.backbone(x)[3]\n            x_eeg_spec = self.backbone_2(x_eeg_spec)[4]\n            # (b * c, 1, mini_h, mini_w)\n            x_eeg_spec = torch.mean(x_eeg_spec, dim=1, keepdims=True)\n            _, _, mini_h, mini_w = x_eeg_spec.shape\n            # (b, c - 4, mini_h, mini_w)\n            x_eeg_spec = x_eeg_spec.reshape(b, c - 4, mini_h, mini_w)\n\n            # (b, c, mini_h, mini_w)\n            x = torch.cat([x_kaggle_spec, x_eeg_spec], dim=1)\n\n            if manifold_mixup:\n                x, lam, index = self.mixup(x)\n            x = self.tiling(x)\n            # (b, 1, 256, 256)\n            x = F.interpolate(x, (256, 256), mode=\"bilinear\")\n            # (b, f_neck)\n            x = self.forward_until_pooling(self.neck, x, True)\n            # (b, n_classes)\n            x = self.feat_2(torch.relu(self.bn(self.feat_1(self.dropout(x)))))\n        elif self.ver == \"ver_7\":\n            # backboneの途中のfeature mapを全部並べるバージョン\n            # (b * c,  f_backbone, mini_h, mini_w)\n            # x = self.backbone(x)[1]\n            # x = self.backbone(x)[2]\n            # x = self.backbone(x)[3]\n            x = self.backbone(x)[4]\n            _, n_features, mini_h, mini_w = x.shape\n            x = x.reshape(b, c, n_features, mini_h, mini_w)\n            # (b, n_features, c, mini_h, mini_w)\n            x = x.permute(0, 2, 1, 3, 4)\n            # (b, 1, n_features, c x mini_h x mini_w)\n            x = x.reshape(b, 1, n_features, -1)\n            # (b, 1, 256, 256)\n            x = F.interpolate(x, (256, 256), mode=\"nearest\")\n            if manifold_mixup:\n                x, lam, index = self.mixup(x)\n            # (b, f_neck)\n            x = self.forward_until_pooling(self.neck, x, True)\n            # (b, n_classes)\n            x = self.feat_2(torch.relu(self.bn(self.feat_1(self.dropout(x)))))\n\n        if manifold_mixup:\n            return x, lam, index\n        else:\n            return x\n\n    def tiling(self, x):\n        b, c, h, w = x.shape\n        assert c % 4 == 0\n        n_rows = 4\n        n_cols = c // 4\n        # 2つの列を表すリストを準備\n        columns = []\n        # 列優先で画像をリストに追加\n        for j in range(n_cols):\n            column_images = []\n            for i in range(n_rows):\n                # 対応する画像を列のリストに追加\n                column_images.append(x[:, i + j * n_rows, :, :])\n            # 一つの列に属する画像を垂直方向（dim=1）に結合\n            columns.append(torch.cat(column_images, dim=1))\n\n        # 最終的に列を水平方向（dim=2）に結合して、全体の画像を形成\n        x = torch.cat(columns, dim=2)\n        # (b, 1, new_h, new_w)\n        x = x.unsqueeze(1)\n        return x","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_net(ckpt_path, device):\n    cfg_path = Path(ckpt_path).parents[1] / \"train_config.yaml\"\n    with open(cfg_path, \"r\", encoding=\"utf-8\") as f:\n        cfg = yaml.load(f, Loader=yaml.SafeLoader)\n\n    args = cfg[\"model\"][\"args\"]\n\n    args[\"pretrained\"] = False\n\n    net = eval(cfg[\"model\"][\"name\"])(**args)\n    net.to(device)\n    ckpt = torch.load(ckpt_path, map_location=f\"cuda:{device}\")[\"state_dict\"]\n    ckpt = {k[k.find(\".\") + 1 :]: v for k, v in ckpt.items()}\n\n    missing_keys, unexpected_keys = net.load_state_dict(ckpt, strict=False)\n    assert not missing_keys, f\"{ckpt_path=}\\n{missing_keys=}\"\n    net.eval()\n    net.config = cfg\n\n    return net\n\n\ndef get_loader(\n    cfg,\n    config,\n    csv_path,\n    transforms,\n    batch_size,\n    num_workers,\n    dry_run,\n    mode=\"test\",\n):\n    tmp_cfg = cfg[\"dataset\"].copy()\n    specs = np.load(cfg[\"specs_path\"], allow_pickle=True).item()\n\n    dataset_args = {\n        **tmp_cfg,\n        \"specs\": specs,\n        \"mode\": mode,\n        \"spec_transforms\": transforms.spec_transform_test,\n        \"audio_transforms\": transforms.audio_transform_test,\n        \"dry_run\": dry_run,\n        \"p_eeg_spec_root\": config[\"p_eeg_specs_output_root\"]\n    }\n    if mode == \"pseudo_label\":\n        dataset_args[\"group\"] = None\n\n    dataset = eval(cfg[\"dataset_name\"])(**dataset_args)\n    loader = DataLoader(\n        dataset=dataset,\n        batch_size=batch_size,\n        shuffle=False,\n        num_workers=num_workers,\n        pin_memory=False,\n    )\n    return loader\n\n\ndef apply_tta(net, image, i_tta):\n    if i_tta == 0:\n        pred = net(image)\n    elif i_tta == 1:  # flip width\n        pred = net(image.flip(-1))\n        raise ValueError\n    return pred\n\n\ndef inference_core(loader, transform, ckpt_list, mode, device, n_tta, use_amp):\n    net_list = []\n    for ckpt in ckpt_list:\n        net_list.append(load_net(ckpt, device))\n\n    results = []\n    results_logit = []\n    for i, batch in enumerate(tqdm(loader, desc=f\"prediction\")):\n        if mode == \"test\":\n            image = batch\n        elif mode == \"val\" or mode == \"pseudo_label\":\n            image, target = batch\n        else:\n            raise ValueError\n\n        image = image.to(device)\n        if transform is not None:\n            # image = image.byte()\n            image = transform(image)\n\n        result_net = []\n        result_net_logit = []\n        for net in net_list:\n            result_tta = []\n            result_tta_logit = []\n            for i_tta in range(n_tta):\n                with torch.autocast(\n                    device_type=\"cuda\",\n                    dtype=torch.float16,\n                    enabled=use_amp,\n                ):\n                    if net.config[\"dataset\"][\"order\"] == \"tile\":\n                        _, _, h, w = image.shape\n                        resize_h = net.config[\"dataset\"][\"resize_h\"]\n                        resize_w = net.config[\"dataset\"][\"resize_w\"]\n                        if h != resize_h or w != resize_w:\n                            image = F.interpolate(\n                                image, (resize_h, resize_w), mode=\"bilinear\"\n                            )\n                    pred = apply_tta(net, image, i_tta)\n                result_tta_logit.append(pred.float().detach().cpu().numpy())\n                pred = F.softmax(pred.float(), dim=1).detach().cpu().numpy()\n                result_tta.append(pred)\n            result_tta = np.array(result_tta).mean(axis=0)\n            result_tta_logit = np.array(result_tta_logit).mean(axis=0)\n            result_net.append(result_tta)\n            result_net_logit.append(result_tta_logit)\n        result_net = np.array(result_net).mean(axis=0)\n        result_net_logit = np.array(result_net_logit).mean(axis=0)\n        # ensemble結果のラベルの合計が1になるようにする\n        result_net = result_net / result_net.sum(axis=1, keepdims=True)\n        results.append(result_net)\n        results_logit.append(result_net_logit)\n    results = np.concatenate(results)\n    results_logit = np.concatenate(results_logit)\n    return results, results_logit\n\n\ndef calc_metric(df, prediction_columns):\n    df[prediction_columns] = df[prediction_columns].astype(\"float32\")\n    df[target_columns] = df[target_columns].astype(\"float32\")\n    submission = df[prediction_columns]\n    submission.columns = target_columns\n    solution = df[target_columns]\n    submission[\"id\"] = np.arange(len(submission))\n    solution[\"id\"] = np.arange(len(solution))\n    metric = score(\n        solution=solution, submission=submission, row_id_column_name=\"id\"\n    )\n    metric_dict = {\"metric\": metric}\n    return metric_dict\n\n\ndef inference(config):\n    mode = config[\"mode\"]\n    assert mode in [\"test\"]\n\n    csv_path = config[\"csv_path\"]\n    device = config[\"gpu\"]\n    batch_size = config[\"batch_size\"]\n    n_tta = config[\"n_tta\"]\n    seed = config[\"seed\"]\n    num_workers = config[\"n_workers\"]\n    use_amp = config[\"use_amp\"]\n    dry_run = config[\"dry_run\"]\n\n    torch.backends.cudnn.deterministic = True\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n\n    cfg_path = Path(config[\"checkpoint\"][0]).parents[1] / \"train_config.yaml\"\n    with open(cfg_path, \"r\", encoding=\"utf-8\") as f:\n        cfg = yaml.load(f, Loader=yaml.SafeLoader)\n    \n    cfg[\"specs_path\"] = Path(config[\"p_kaggle_specs_output_root\"]) / \"specs.npy\"\n    cfg[\"dataset\"][\"csv_path\"] = config[\"csv_path\"]\n\n    transforms = eval(cfg[\"augmentation_name\"])(cfg[\"augmentation_ver\"])\n    loader = get_loader(\n        cfg,\n        config,\n        csv_path,\n        transforms,\n        batch_size,\n        num_workers,\n        dry_run,\n        mode,\n    )\n    df = loader.dataset.df\n\n    # torch_transform = transform.torch_test\n    torch_transform = None\n\n    with torch.no_grad():\n        results, results_logit = inference_core(\n            loader,\n            torch_transform,\n            config[\"checkpoint\"],\n            mode,\n            device,\n            n_tta,\n            use_amp,\n        )\n\n    # 念のためnan除け(基本nanは発生しない)\n    results[np.isnan(results).sum(axis=1) > 0] = [\n        1 / 6,\n        1 / 6,\n        1 / 6,\n        1 / 6,\n        1 / 6,\n        1 / 6,\n    ]\n    results_logit[np.isnan(results_logit).sum(axis=1) > 0] = [\n        1 / 6,\n        1 / 6,\n        1 / 6,\n        1 / 6,\n        1 / 6,\n        1 / 6,\n    ]\n\n    df[prediction_columns] = results\n    df[logit_columns] = results_logit\n    df = df[[\"eeg_id\"] + prediction_columns + logit_columns]\n    df.to_csv(f\"{config['output_csv_name']}\", index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile config.yaml\nmode: test\nbatch_size: 32\nn_workers: 20\nseed: 42\nn_tta: 1\ngpu: 0\nuse_amp: True\n# use_amp: False\ndry_run: ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config = OmegaConf.load(\"/kaggle/working/config.yaml\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 2nd-stageのmodel\ncheckpoint_dict = {\n#=================HiroiFoldV1=================\n    \"logit_Kazumax_CV_0.2639_oof.csv\":\n    [\n        \"/kaggle/input/hms-models-9/run-20240403_150347-sghdwjyy/checkpoints/epoch-08-step-1224-val_metric_low-0.7757-val_metric_high-0.2607-val_metric-0.5956.ckpt\",\n        \"/kaggle/input/hms-models-9/run-20240403_150347-8wgldlq1/checkpoints/epoch-06-step-952-val_metric_low-0.7348-val_metric_high-0.2530-val_metric-0.5663.ckpt\",\n        \"/kaggle/input/hms-models-9/run-20240403_150350-x0t9reju/checkpoints/epoch-08-step-1269-val_metric_low-0.8216-val_metric_high-0.2599-val_metric-0.6449.ckpt\",\n        \"/kaggle/input/hms-models-9/run-20240403_150351-p1gt14uh/checkpoints/epoch-09-step-1430-val_metric_low-0.7243-val_metric_high-0.2821-val_metric-0.5931.ckpt\",\n    ],\n    \"logit_Kazumax_CV_0.2777_oof.csv\":\n    [\n        \"/kaggle/input/hms-models-10/run-20240403_155251-o2vt9kw4/checkpoints/epoch-10-step-1496-val_metric_low-0.8012-val_metric_high-0.2683-val_metric-0.6149.ckpt\",\n        \"/kaggle/input/hms-models-10/run-20240403_155257-nhj0dnsv/checkpoints/epoch-14-step-2040-val_metric_low-0.8269-val_metric_high-0.2853-val_metric-0.6375.ckpt\",\n        \"/kaggle/input/hms-models-10/run-20240403_155354-awvcq66k/checkpoints/epoch-18-step-2679-val_metric_low-0.7807-val_metric_high-0.2621-val_metric-0.6175.ckpt\",\n        \"/kaggle/input/hms-models-10/run-20240403_155356-9od3hbtq/checkpoints/epoch-08-step-1278-val_metric_low-0.7830-val_metric_high-0.2952-val_metric-0.6382.ckpt\",\n    ],\n    \"logit_Kazumax_CV_0.2704_oof.csv\":\n    [\n        \"/kaggle/input/hms-models-11/run-20240403_165921-y8vy02sz/checkpoints/epoch-13-step-1904-val_metric_low-0.7753-val_metric_high-0.2516-val_metric-0.5922.ckpt\",\n        \"/kaggle/input/hms-models-11/run-20240403_165937-o3cdlbcn/checkpoints/epoch-14-step-2040-val_metric_low-0.8242-val_metric_high-0.2597-val_metric-0.6267.ckpt\",\n        \"/kaggle/input/hms-models-11/run-20240403_170053-uczw9hpk/checkpoints/epoch-18-step-2679-val_metric_low-0.8504-val_metric_high-0.2709-val_metric-0.6681.ckpt\",\n        \"/kaggle/input/hms-models-11/run-20240403_170147-3yz0p2g7/checkpoints/epoch-19-step-2840-val_metric_low-0.7307-val_metric_high-0.2995-val_metric-0.6027.ckpt\",\n    ],\n#=================HiroiFoldV2=================\n    \"logit_Kazumax_CV_0.2633_oof.csv\":\n    [\n        \"/kaggle/input/hms-models-12/run-20240405_140632-6y1xsnns/checkpoints/epoch-03-step-552-val_metric_low-0.7567-val_metric_high-0.2829-val_metric-0.6018.ckpt\",\n        \"/kaggle/input/hms-models-12/run-20240405_140633-1r7f68an/checkpoints/epoch-07-step-1152-val_metric_low-0.8498-val_metric_high-0.2483-val_metric-0.6585.ckpt\",\n        \"/kaggle/input/hms-models-12/run-20240405_140634-g64s5v63/checkpoints/epoch-06-step-931-val_metric_low-0.6942-val_metric_high-0.2640-val_metric-0.5539.ckpt\",\n        \"/kaggle/input/hms-models-12/run-20240405_140636-8mz4hezc/checkpoints/epoch-09-step-1420-val_metric_low-0.8193-val_metric_high-0.2579-val_metric-0.6283.ckpt\",\n    ],\n    \"logit_Kazumax_CV_0.2729_oof.csv\":\n    [\n        \"/kaggle/input/hms-models-13/run-20240405_145742-opf5720b/checkpoints/epoch-15-step-2192-val_metric_low-0.7290-val_metric_high-0.2890-val_metric-0.5851.ckpt\",\n        \"/kaggle/input/hms-models-13/run-20240405_145757-yxcm6a5e/checkpoints/epoch-12-step-1872-val_metric_low-0.8845-val_metric_high-0.2541-val_metric-0.6840.ckpt\",\n        \"/kaggle/input/hms-models-13/run-20240405_145720-k4znq322/checkpoints/epoch-15-step-2112-val_metric_low-0.7416-val_metric_high-0.2845-val_metric-0.5925.ckpt\",\n        \"/kaggle/input/hms-models-13/run-20240405_145727-52hd46zi/checkpoints/epoch-18-step-2698-val_metric_low-0.8529-val_metric_high-0.2638-val_metric-0.6525.ckpt\",\n    ],\n    \"logit_Kazumax_CV_0.2698_oof.csv\":\n    [\n        \"/kaggle/input/hms-models-14/run-20240405_160440-fa5b1ycw/checkpoints/epoch-05-step-822-val_metric_low-0.7075-val_metric_high-0.2924-val_metric-0.5718.ckpt\",\n        \"/kaggle/input/hms-models-14/run-20240405_160505-2it6zdri/checkpoints/epoch-15-step-2304-val_metric_low-0.8587-val_metric_high-0.2424-val_metric-0.6627.ckpt\",\n        \"/kaggle/input/hms-models-14/run-20240405_160335-sm0ynnbi/checkpoints/epoch-19-step-2640-val_metric_low-0.7260-val_metric_high-0.2839-val_metric-0.5818.ckpt\",\n        \"/kaggle/input/hms-models-14/run-20240405_160419-mqs93aj7/checkpoints/epoch-14-step-2130-val_metric_low-0.8091-val_metric_high-0.2606-val_metric-0.6224.ckpt\",\n    ],\n}\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# モデルパスのチェック\nfor name, ckpt_list in checkpoint_dict.items():\n    for ckpt in ckpt_list:\n        assert Path(ckpt).exists(), f\"{name}: {ckpt}\"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"KAZUMAX_DEBUG = False\n# KAZUMAX_DEBUG = True\n# KAZUMAX_DEBUG_DATASET = \"train\"\nKAZUMAX_DEBUG_DATASET = \"test\"\nKAZUMAX_N_DEBUG = 100","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if KAZUMAX_DEBUG:\n    if KAZUMAX_DEBUG_DATASET == \"train\":\n        csv_path = \"/kaggle/input/hms-harmful-brain-activity-classification/train.csv\"\n        p_eeg_specs = \"/kaggle/input/hms-harmful-brain-activity-classification/train_eegs\"\n        p_eeg_specs_output_root = \"/kaggle/kazumax_tmp/eeg_specs\"\n        p_kaggle_specs_root = \"/kaggle/input/hms-harmful-brain-activity-classification/train_spectrograms/\"\n        p_kaggle_specs_output_root = \"/kaggle/kazumax_tmp\"\n        dry_run = KAZUMAX_N_DEBUG\n    else:\n        csv_path = \"/kaggle/input/hms-harmful-brain-activity-classification/test.csv\"\n        p_eeg_specs = \"/kaggle/input/hms-harmful-brain-activity-classification/test_eegs\"\n        p_eeg_specs_output_root = \"/kaggle/kazumax_tmp/eeg_specs\"\n        p_kaggle_specs_root = \"/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/\"\n        p_kaggle_specs_output_root = \"/kaggle/kazumax_tmp\"\n        dry_run = None\nelse:\n    csv_path = \"/kaggle/input/hms-harmful-brain-activity-classification/test.csv\"\n    p_eeg_specs = \"/kaggle/input/hms-harmful-brain-activity-classification/test_eegs\"\n    p_eeg_specs_output_root = \"/kaggle/kazumax_tmp/eeg_specs\"\n    p_kaggle_specs_root = \"/kaggle/input/hms-harmful-brain-activity-classification/test_spectrograms/\"\n    p_kaggle_specs_output_root = \"/kaggle/kazumax_tmp\"\n    dry_run = None\n\nconfig[\"csv_path\"] = csv_path\nconfig[\"p_eeg_specs\"] = p_eeg_specs\nconfig[\"p_eeg_specs_output_root\"] = p_eeg_specs_output_root\nconfig[\"p_kaggle_specs_root\"] = p_kaggle_specs_root\nconfig[\"p_kaggle_specs_output_root\"] = p_kaggle_specs_output_root\nconfig[\"dry_run\"] = dry_run","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# For UEMU\nmake_specs(config)\nmake_eeg_specs_main(csv_path, p_eeg_specs, p_eeg_specs_output_root)\nfor name, ckpt_list in checkpoint_dict.items():\n    config[\"checkpoint\"] = ckpt_list\n    config[\"output_csv_name\"] = name\n    inference(config)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# submission.csvが見つからないエラーが発生する原因になるので、不要なファイルを消しておく\nif not KAZUMAX_DEBUG:\n    !rm -r /kaggle/kazumax_tmp/\n    !rm -r /kaggle/working/config.yaml","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if KAZUMAX_DEBUG:\n    df = pd.read_csv(\"./logit_Kazumax_CV_0.2639_oof.csv\")\n    display(df)\n    df = pd.read_csv(\"./logit_Kazumax_CV_0.2777_oof.csv\")\n    display(df)\n    df = pd.read_csv(\"./logit_Kazumax_CV_0.2704_oof.csv\")\n    display(df)\n    df = pd.read_csv(\"./logit_Kazumax_CV_0.2633_oof.csv\")\n    display(df)\n    df = pd.read_csv(\"./logit_Kazumax_CV_0.2729_oof.csv\")\n    display(df)\n    df = pd.read_csv(\"./logit_Kazumax_CV_0.2698_oof.csv\")\n    display(df)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"##  Team Ensemble Part","metadata":{}},{"cell_type":"code","source":"#remove_path\nsys.path.remove(\"/kaggle/input/hms-hbac-src\")\n#remove_modules\nif 'models' in sys.modules:\n    del sys.modules['models']\nif 'trainer' in sys.modules:\n    del sys.modules['trainer']","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def df_rename_merge_TH(df_result,df_sample):\n    df_result                                = df_result.rename(columns={\n                                                                'seizure_vote':'seizure_logits',\n                                                                'lpd_vote':'lpd_logits',\n                                                                'gpd_vote':'gpd_logits',\n                                                                'lrda_vote':'lrda_logits',\n                                                                'grda_vote':'grda_logits',\n                                                                'other_vote':'other_logits',\n                                                               })\n    df_result                                = df_result.merge(df_sample[['eeg_id','spectrogram_id','patient_id',]],on='eeg_id',how='left')\n    return df_result\n\ndef df_rename_merge_kazumax(df_result,df_sample):\n    df_result                                = df_result.rename(columns={\n                                                                'seizure_vote_logit':'seizure_logits',\n                                                                'lpd_vote_logit':'lpd_logits',\n                                                                'gpd_vote_logit':'gpd_logits',\n                                                                'lrda_vote_logit':'lrda_logits',\n                                                                'grda_vote_logit':'grda_logits',\n                                                                'other_vote_logit':'other_logits',\n                                                               })\n    df_result                                = df_result[['eeg_id','seizure_logits','lpd_logits','gpd_logits','lrda_logits','grda_logits','other_logits']]\n    df_result                                = df_result.merge(df_sample[['eeg_id','spectrogram_id','patient_id',]],on='eeg_id',how='left')\n    return df_result\n\n#sample_uemu_result\ndf_sample_UEMU                              = pd.read_csv('/kaggle/working/output/df_test_output_exp453_EEG_WAVE_multi1D_deep1D_cnn2D_Fold_TH_HvoteData_pretrain_exp451.csv')\n\n#==TH==#\n#=read_result\n#foldv1\ndf_result_TH_wavenet_maxxvitv2n                    = pd.read_csv('wavenet_maxxvitv2n_logit.csv')\ndf_result_TH_wavenet_maxxvitv2n_downsample         = pd.read_csv('wavenet_maxxvitv2n_downsample_logit.csv')\ndf_result_TH_wavenet_maxxvitv2n_downsample_seed0   = pd.read_csv('wavenet_maxxvitv2n_downsample_seed0_logit.csv')\ndf_result_TH_wavenet_maxxvitv2n_downsample_seed123 = pd.read_csv('wavenet_maxxvitv2n_downsample_seed123_logit.csv')\ndf_result_TH_wavenet_maxxvits_downsample           = pd.read_csv('wavenet_maxxvits_downsample_logit.csv')\ndf_result_TH_wavenet_effnetb4_downsample           = pd.read_csv('wavenet_effnetb4_downsample_logit.csv')\n#foldv2\ndf_result_TH_wavenet_maxxvitv2n_downsample_foldV2  = pd.read_csv('wavenet_maxxvitv2n_downsample_foldV2_logit.csv')\ndf_result_TH_wavenet_effnetb4_downsample_foldV2    = pd.read_csv('wavenet_effnetb4_downsample_foldV2_logit.csv')\ndf_result_TH_wavenet_maxxvits_downsample_foldV2    = pd.read_csv('wavenet_maxxvits_downsample_foldV2_logit.csv')\n\n#=rename\n#foldv1\ndf_result_TH_wavenet_maxxvitv2n                    = df_rename_merge_TH(df_result_TH_wavenet_maxxvitv2n,df_sample_UEMU)\ndf_result_TH_wavenet_maxxvitv2n_downsample         = df_rename_merge_TH(df_result_TH_wavenet_maxxvitv2n_downsample,df_sample_UEMU)\ndf_result_TH_wavenet_maxxvitv2n_downsample_seed0   = df_rename_merge_TH(df_result_TH_wavenet_maxxvitv2n_downsample_seed0,df_sample_UEMU)\ndf_result_TH_wavenet_maxxvitv2n_downsample_seed123 = df_rename_merge_TH(df_result_TH_wavenet_maxxvitv2n_downsample_seed123,df_sample_UEMU)\ndf_result_TH_wavenet_maxxvits_downsample           = df_rename_merge_TH(df_result_TH_wavenet_maxxvits_downsample,df_sample_UEMU)\ndf_result_TH_wavenet_effnetb4_downsample           = df_rename_merge_TH(df_result_TH_wavenet_effnetb4_downsample,df_sample_UEMU)\n#foldv2\ndf_result_TH_wavenet_maxxvitv2n_downsample_foldV2  = df_rename_merge_TH(df_result_TH_wavenet_maxxvitv2n_downsample_foldV2,df_sample_UEMU)\ndf_result_TH_wavenet_effnetb4_downsample_foldV2    = df_rename_merge_TH(df_result_TH_wavenet_effnetb4_downsample_foldV2,df_sample_UEMU)\ndf_result_TH_wavenet_maxxvits_downsample_foldV2    = df_rename_merge_TH(df_result_TH_wavenet_maxxvits_downsample_foldV2,df_sample_UEMU)\n\n\n#=save\n#foldv1\ndf_result_TH_wavenet_maxxvitv2n.to_csv(f'/kaggle/working/output/df_test_output_result_TH_wavenet_maxxvitv2n_oof.csv',index=False)\ndf_result_TH_wavenet_maxxvitv2n_downsample.to_csv(f'/kaggle/working/output/df_test_output_result_TH_wavenet_maxxvitv2n_downsample_oof.csv',index=False)\ndf_result_TH_wavenet_maxxvitv2n_downsample_seed0.to_csv(f'/kaggle/working/output/df_test_output_result_TH_wavenet_maxxvitv2n_downsample_seed0_oof.csv',index=False)\ndf_result_TH_wavenet_maxxvitv2n_downsample_seed123.to_csv(f'/kaggle/working/output/df_test_output_result_TH_wavenet_maxxvitv2n_downsample_seed123_oof.csv',index=False)\ndf_result_TH_wavenet_maxxvits_downsample.to_csv(f'/kaggle/working/output/df_test_output_result_TH_wavenet_maxxvits_downsample_oof.csv',index=False)\ndf_result_TH_wavenet_effnetb4_downsample.to_csv(f'/kaggle/working/output/df_test_output_result_TH_wavenet_effnetb4_downsample_oof.csv',index=False)\n#foldv2\ndf_result_TH_wavenet_maxxvitv2n_downsample_foldV2.to_csv(f'/kaggle/working/output/df_test_output_result_TH_wavenet_maxxvitv2n_downsample_foldV2_oof.csv',index=False)\ndf_result_TH_wavenet_effnetb4_downsample_foldV2.to_csv(f'/kaggle/working/output/df_test_output_result_TH_wavenet_effnetb4_downsample_foldV2_oof.csv',index=False)\ndf_result_TH_wavenet_maxxvits_downsample_foldV2.to_csv(f'/kaggle/working/output/df_test_output_result_TH_wavenet_maxxvits_downsample_foldV2_oof.csv',index=False)\n\n\n#==kazumax==#\n#=read_result\n#foldv1\ndf_result_Kazumax_CV_02639                        = pd.read_csv('logit_Kazumax_CV_0.2639_oof.csv')\ndf_result_Kazumax_CV_02704                        = pd.read_csv('logit_Kazumax_CV_0.2704_oof.csv')\ndf_result_Kazumax_CV_02777                        = pd.read_csv('logit_Kazumax_CV_0.2777_oof.csv')\n#foldv2\ndf_result_Kazumax_CV_02698                        = pd.read_csv('logit_Kazumax_CV_0.2698_oof.csv')\ndf_result_Kazumax_CV_02729                        = pd.read_csv('logit_Kazumax_CV_0.2729_oof.csv')\ndf_result_Kazumax_CV_02633                        = pd.read_csv('logit_Kazumax_CV_0.2633_oof.csv')\n\n#=rename\n#foldv1\ndf_result_Kazumax_CV_02639                        = df_rename_merge_kazumax(df_result_Kazumax_CV_02639,df_sample_UEMU)\ndf_result_Kazumax_CV_02704                        = df_rename_merge_kazumax(df_result_Kazumax_CV_02704,df_sample_UEMU)\ndf_result_Kazumax_CV_02777                        = df_rename_merge_kazumax(df_result_Kazumax_CV_02777,df_sample_UEMU)\n#foldv2\ndf_result_Kazumax_CV_02698                        = df_rename_merge_kazumax(df_result_Kazumax_CV_02698,df_sample_UEMU)\ndf_result_Kazumax_CV_02729                        = df_rename_merge_kazumax(df_result_Kazumax_CV_02729,df_sample_UEMU)\ndf_result_Kazumax_CV_02633                        = df_rename_merge_kazumax(df_result_Kazumax_CV_02633,df_sample_UEMU)\n\n#=save\n#foldv1\ndf_result_Kazumax_CV_02639.to_csv(f'/kaggle/working/output/df_test_output_result_KM_CV_0.2639_oof.csv',index=False)\ndf_result_Kazumax_CV_02704.to_csv(f'/kaggle/working/output/df_test_output_result_KM_CV_0.2704_oof.csv',index=False)\ndf_result_Kazumax_CV_02777.to_csv(f'/kaggle/working/output/df_test_output_result_KM_CV_0.2777_oof.csv',index=False)\n#foldv2\ndf_result_Kazumax_CV_02698.to_csv(f'/kaggle/working/output/df_test_output_result_KM_CV_0.2698_oof.csv',index=False)\ndf_result_Kazumax_CV_02729.to_csv(f'/kaggle/working/output/df_test_output_result_KM_CV_0.2729_oof.csv',index=False)\ndf_result_Kazumax_CV_02633.to_csv(f'/kaggle/working/output/df_test_output_result_KM_CV_0.2633_oof.csv',index=False)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Stacking ～ PP","metadata":{}},{"cell_type":"code","source":"sys.path.append('/kaggle/input/hms-code/src')\nfrom trainer.datasets_pp import HMS_Dataset_PP\n\ndef calc_inference_pp(df_feat,col_feat,col_labels,\n                      model_dir,exp_name,config_common):\n    \n    test_dataset                    = HMS_Dataset_PP(df_feat,col_feat,col_label=[],phase='test')\n    test_loader                     = DataLoader(test_dataset, batch_size=BATCH_SIZE_Test, shuffle=False,\n                                                num_workers=NUM_WORKERS,pin_memory=True)\n    #===model===#\n    model_fold_dirs                 = find_fold_x_paths(model_dir)\n    model_fold_dirs.sort()\n    models                          = []\n    for model_fold_dir in model_fold_dirs:\n        model_path                  = f'{model_fold_dir}/best.pt'\n        #print(model_path)\n        model                       = torch.load(model_path, map_location=torch.device(DEVICE))\n        model.eval()\n        model.to(DEVICE)\n        models.append(model)\n\n    #===inference===#\n    with torch.no_grad():\n        for batch,data in enumerate(test_loader):\n            patient_id              = data['patient_id']\n            eeg_id                  = data['eeg_id']\n            spectrogram_id          = data['spectrogram_id']\n            #input\n            feature                 = data['feat']\n            if torch.cuda.is_available():\n                feature             = feature.float().to(DEVICE)\n\n            #=preds=\n            for idy in range(len(models)):\n                model               = models[idy]\n                (_,logits)          = model(feature)\n                #fold_average\n                if idy ==0:\n                    logits_fold_ave = logits/len(models)\n                else:\n                    logits_fold_ave += logits/len(models)\n\n            #=matome in batch=\n            if batch==0:\n                test_patient_ids     = patient_id.tolist()\n                test_eeg_ids         = eeg_id.tolist()\n                test_spectrogram_ids = spectrogram_id.tolist()\n                test_logits          = logits_fold_ave.float().detach().cpu()\n            else:\n                test_patient_ids     += patient_id.tolist()\n                test_eeg_ids         += eeg_id.tolist()\n                test_spectrogram_ids += spectrogram_id.tolist()\n                test_logits          = torch.cat([test_logits,logits_fold_ave.float().detach().cpu()],dim=0)\n\n    #=matome in model=\n    col_labels_logits                      = [s.split('_')[0]+'_logits' for s in config_common['dataset']['col_labels']]\n    df_test_output                         = pd.DataFrame()\n    df_test_output['eeg_id']               = test_eeg_ids\n    df_test_output['spectrogram_id']       = test_spectrogram_ids\n    df_test_output['patient_id']           = test_patient_ids\n    df_test_output[col_labels_logits]      = test_logits.numpy()\n    df_test_output                         = df_test_output.sort_values(by=['patient_id', 'eeg_id'], ascending=[True, True]).reset_index(drop=True)\n    df_test_output.to_csv(f'/kaggle/working/output/df_test_output_{exp_name}.csv',index=False)\n\n    del models,test_dataset,test_loader\n\n    gc.collect()\n    torch.cuda.empty_cache()  \n    ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def calc_weight_ave(df_ens1):\n    exp_ids             = df_ens1['exp_id'].values\n    weights             = df_ens1['weights'].values\n    weights             = weights/np.sum(weights)\n\n    #read\n    test_output_data            = []\n    for exp_id in exp_ids:\n        df_test_output          = pd.read_csv(f'/kaggle/working/output/df_test_output_{exp_id}.csv')\n        df_test_output          = df_test_output[base_cols+col_labels_feat]\n        df_test_output          = df_test_output.sort_values(by=['patient_id', 'eeg_id'], ascending=[True, True]).reset_index(drop=True)\n        test_output_data.append(df_test_output[col_labels_feat].values)\n\n    #weight_ave\n    test_output_data            = np.stack(test_output_data)\n    weighted_average            = np.tensordot(weights, test_output_data, axes=([0], [0]))\n\n    #df\n    df_test_output_ens                      = df_test_output.copy()\n    df_test_output_ens[col_labels_feat]     = weighted_average\n    return df_test_output_ens","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def add_feat(df_feat,new_col_labels,setting_pp):\n    add_ave                             = setting_pp['add_ave']\n    add_max                             = setting_pp['add_max']\n    add_min                             = setting_pp['add_min']\n    add_std                             = setting_pp['add_std']\n    add_count                           = setting_pp['add_count']\n    add_delta                           = setting_pp['add_delta']\n    new_col_labels_all                  = new_col_labels.copy()\n    #patient_ave\n    if add_ave:\n        new_col_labels_patient_ave      = [s+'_patient_ave' for s in new_col_labels]\n        patient_ave                     = df_feat.groupby('patient_id')[new_col_labels].mean().reset_index()\n        patient_ave                     = patient_ave.rename(columns={s:s+'_patient_ave' for s in new_col_labels})\n        df_feat                         = pd.merge(df_feat, patient_ave, on='patient_id')\n        new_col_labels_all              += new_col_labels_patient_ave\n    #patient_max\n    if add_max:\n        new_col_labels_patient_max      = [s+'_patient_max' for s in new_col_labels]\n        patient_max                     = df_feat.groupby('patient_id')[new_col_labels].max().reset_index()\n        patient_max                     = patient_max.rename(columns={s:s+'_patient_max' for s in new_col_labels})\n        df_feat                         = pd.merge(df_feat, patient_max, on='patient_id')\n        new_col_labels_all              += new_col_labels_patient_max\n    #patient_min\n    if add_min:\n        new_col_labels_patient_min      = [s+'_patient_min' for s in new_col_labels]\n        patient_min                     = df_feat.groupby('patient_id')[new_col_labels].min().reset_index()\n        patient_min                     = patient_min.rename(columns={s:s+'_patient_min' for s in new_col_labels})\n        df_feat                         = pd.merge(df_feat, patient_min, on='patient_id')\n        new_col_labels_all              += new_col_labels_patient_min\n    #patient_std\n    if add_std:\n        new_col_labels_patient_std      = [s+'_patient_std' for s in new_col_labels]\n        patient_std                     = df_feat.groupby('patient_id')[new_col_labels].std().reset_index()\n        patient_std                     = patient_std.rename(columns={s:s+'_patient_std' for s in new_col_labels})\n        df_feat                         = pd.merge(df_feat, patient_std, on='patient_id')\n        new_col_labels_all              += new_col_labels_patient_std\n    #count_mesure\n    if add_count:\n        new_col_labels_patient_count    = ['count_mesure']\n        patient_count                   = df_feat.groupby('patient_id').count().reset_index()\n        patient_count                   = patient_count[['patient_id','eeg_id']].rename(columns={'eeg_id':'count_mesure'})\n        df_feat                         = pd.merge(df_feat, patient_count, on='patient_id')\n        new_col_labels_all              += new_col_labels_patient_count\n    #patient_delta\n    if add_delta:\n        new_col_labels_patient_delta        = [s+'_patient_delta' for s in new_col_labels]\n        df_feat[new_col_labels_patient_delta] = df_feat[new_col_labels].values - df_feat[new_col_labels_patient_ave].values\n        new_col_labels_all                  += new_col_labels_patient_delta\n    \n    return df_feat,new_col_labels_all\n\ndef preprocess_pp_inference(dict_df_oof,setting_pp,col_labels_feat):\n    use_ens                                 = setting_pp['use_ens']\n    use_spec                                = setting_pp['use_spec']\n    use_eeg_wave                            = setting_pp['use_eeg_wave']\n    use_eeg_img                             = setting_pp['use_eeg_img']\n    use_multi                               = setting_pp['use_multi']\n    use_stacking                            = setting_pp['use_stacking']\n    base_cols                               = ['eeg_id', 'spectrogram_id','patient_id']\n    df_feat                                 = dict_df_oof['ens']\n    df_feat                                 = df_feat[base_cols]\n    new_col_labels                          = []\n    if use_ens:\n        df_oof_output_ens                   = dict_df_oof['ens']   \n        new_col_labels_ens                  = [s+'_ens' for s in col_labels_feat]\n        new_col_labels                      += new_col_labels_ens\n        df_feat.loc[:,new_col_labels_ens]   = df_oof_output_ens.loc[:,col_labels_feat].values\n    if use_spec:\n        df_oof_output_spec                  = dict_df_oof['spec']   \n        new_col_labels_spec                 = [s+'_spec' for s in col_labels_feat]\n        new_col_labels                      += new_col_labels_spec\n        df_feat.loc[:,new_col_labels_spec]  = df_oof_output_spec.loc[:,col_labels_feat].values\n    if use_eeg_wave:\n        df_oof_output_eeg_wave              = dict_df_oof['wave']  \n        new_col_labels_eeg_wave             = [s+'_eeg_wave' for s in col_labels_feat]\n        new_col_labels                      += new_col_labels_eeg_wave\n        df_feat.loc[:,new_col_labels_eeg_wave]  = df_oof_output_eeg_wave.loc[:,col_labels_feat].values\n    if use_eeg_img:\n        df_oof_output_eeg_img               = dict_df_oof['img']\n        new_col_labels_eeg_img              = [s+'_eeg_img' for s in col_labels_feat]\n        new_col_labels                      += new_col_labels_eeg_img\n        df_feat.loc[:,new_col_labels_eeg_img]   = df_oof_output_eeg_img.loc[:,col_labels_feat].values\n    if use_multi:\n        df_oof_output_multi                 = dict_df_oof['multi']\n        new_col_labels_multi                = [s+'_multi' for s in col_labels_feat]\n        new_col_labels                      += new_col_labels_multi\n        df_feat.loc[:,new_col_labels_multi] = df_oof_output_multi.loc[:,col_labels_feat].values\n    if use_stacking:\n        df_oof_output_stacking              = dict_df_oof['stack']\n        new_col_labels_stacking             = [s+'_stacking' for s in col_labels_feat]\n        new_col_labels                      += new_col_labels_stacking\n        df_feat.loc[:,new_col_labels_stacking] = df_oof_output_stacking.loc[:,col_labels_feat].values\n        \n    #===add_feat===#\n    df_feat,new_col_labels_all              = add_feat(df_feat,new_col_labels,setting_pp)\n\n    df_feat                                 = df_feat.fillna(0)\n    return df_feat,new_col_labels_all\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# sys.path.append('/kaggle/input/hms-code/src')\n# from trainer.train_pp import preprocess_pp_inference\n\ncalc_pp                 = [True,True]\nFoldset_path_list       = [\n                             #TH FOLD\n                            '/kaggle/input/hms-models5/exp_rspe30_fold_TH_stack30_comb4_pp8_hidden64_layer3_EP200_trial5000_8000_add_kazumax_v2',\n                            '/kaggle/input/hms-models5/exp_rspe40_fold_THV2_stack24_comb4_pp8_hidden64_layer3_EP200_trial5000_8000',\n                            ]\n\n#\ncol_labels              = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\ncol_labels_logits       = [s.split('_')[0]+'_logits' for s in col_labels]\ncol_labels_feat         = col_labels_logits\nbase_cols               = ['eeg_id', 'spectrogram_id','patient_id']\n\nfor idx_Foldset,Foldset_path in enumerate(Foldset_path_list):\n    print(f'=======Foldset:{Foldset_path}==========')\n    #=========1.stacking=========#\n    try:#stackingフォルダあり\n        df_stack        = pd.read_csv(f'{Foldset_path}/stacking/df_stacking.csv')\n        calc_stack      = True\n    except:\n        calc_stack      = False\n        print('Stacking:OFF')\n    if calc_stack:\n        id_stacking_list= list(df_stack['id_stacking'].values)\n        id_comb_list    = list(df_stack['id_comb'].values)\n        print(f'=======1.Stacking==========')\n        for idx_stacking in range(len(id_stacking_list)):\n            #水準読み取り\n            id_stacking  = id_stacking_list[idx_stacking]\n            id_comb      = ast.literal_eval(id_comb_list[idx_stacking])#文字→list\n            #print(f'==Stacking_Foldset_{idx_Foldset}_{id_stacking}==')\n            #print(id_comb)\n            #=get_feat=#\n            col_labels_allfeat          = []\n            for idx_exp,exp_id in enumerate(id_comb):\n                df_test_output          = pd.read_csv(f'/kaggle/working/output/df_test_output_{exp_id}.csv')\n                df_test_output          = df_test_output[base_cols+col_labels_feat]\n                df_test_output          = df_test_output.sort_values(by=['patient_id', 'eeg_id'], ascending=[True, True]).reset_index(drop=True)\n                #rename\n                column_mapping          = {col: f\"{col}_{idx_exp}\" for col in col_labels_feat if col in df_test_output.columns}\n                new_columns             = list(column_mapping.values())\n                df_test_output          = df_test_output.rename(columns=column_mapping)\n                col_labels_allfeat      +=new_columns\n                if idx_exp==0:\n                    df_test_output_all  = df_test_output\n                else:\n                    df_test_output_all  = df_test_output_all.merge(df_test_output[['patient_id','eeg_id']+new_columns],on=['patient_id','eeg_id'],how='left')\n\n            #stacking\n            model_dir_stacking          = f'{Foldset_path}/stacking/{idx_stacking}/'\n            exp_name                    = f'Foldset_{idx_Foldset}_{id_stacking}'\n            calc_inference_pp(df_test_output_all,col_labels_allfeat,col_labels,\n                              model_dir_stacking,exp_name,config_common)\n    else:\n        pass\n    \n    #=========2.ens1=========#\n    print(f'=======2.ens1==========')\n    #=ens1_spec=#\n    df_ens1_spec                    = pd.read_csv(f'{Foldset_path}/ens1/df_ens1_spec.csv')\n    df_test_output_ens1_spec        = calc_weight_ave(df_ens1_spec)\n    #=ens1_wave=#\n    df_ens1_wave                    = pd.read_csv(f'{Foldset_path}/ens1/df_ens1_wave.csv')\n    df_test_output_ens1_wave        = calc_weight_ave(df_ens1_wave)\n    #=ens1_img=#\n    df_ens1_img                     = pd.read_csv(f'{Foldset_path}/ens1/df_ens1_img.csv')\n    df_test_output_ens1_img         = calc_weight_ave(df_ens1_img)\n    #=ens1_multi=#\n    df_ens1_multi                   = pd.read_csv(f'{Foldset_path}/ens1/df_ens1_multi.csv')\n    df_test_output_ens1_multi       = calc_weight_ave(df_ens1_multi)\n    #=ens1_stack=#\n    if calc_stack:\n        df_ens1_stack               = pd.read_csv(f'{Foldset_path}/ens1/df_ens1_stack.csv')\n        df_ens1_stack.loc[df_ens1_stack['model_type'] == 'STACK', 'exp_id'] = f'Foldset_{idx_Foldset}_' + df_ens1_stack['exp_id']\n        df_test_output_ens1_stack   = calc_weight_ave(df_ens1_stack)\n    else:\n        df_test_output_ens1_stack   = pd.DataFrame()\n    #=ens1_all=#\n    df_ens1_all                     = pd.read_csv(f'{Foldset_path}/ens1/df_ens1_all.csv')\n    df_ens1_all.loc[df_ens1_all['model_type'] == 'STACK', 'exp_id'] = f'Foldset_{idx_Foldset}_' + df_ens1_all['exp_id']\n    df_test_output_ens1_all         = calc_weight_ave(df_ens1_all) \n    \n    if calc_pp[idx_Foldset]:#PP実施する\n        #=========3.pp=========#\n        print(f'=======3.pp==========')\n        df_pp                 = pd.read_csv(f'{Foldset_path}/pp/df_pp.csv')\n        id_pp_list            = list(df_pp['id_pp'].values)\n        id_setting_pp_list    = list(df_pp['id_setting_pp'].values)\n        #=get_feat=#\n        dict_df_test_ens1                = {}\n        dict_df_test_ens1['spec']        = df_test_output_ens1_spec\n        dict_df_test_ens1['wave']        = df_test_output_ens1_wave\n        dict_df_test_ens1['img']         = df_test_output_ens1_img\n        dict_df_test_ens1['multi']       = df_test_output_ens1_multi\n        dict_df_test_ens1['stack']       = df_test_output_ens1_stack\n        dict_df_test_ens1['ens']         = df_test_output_ens1_all\n            \n        for idx_pp in range(len(id_pp_list)):\n            #水準読み取り\n            id_pp                        = id_pp_list[idx_pp]\n            setting_pp                   = ast.literal_eval(id_setting_pp_list[idx_pp])#文字→list\n            #print(f'==PP_Foldset_{idx_Foldset}_{id_pp}==')\n            #print(setting_pp)\n            df_feat_pp,new_col_labels_all   = preprocess_pp_inference(dict_df_test_ens1,setting_pp,col_labels_feat)\n            df_feat_pp                      = df_feat_pp.sort_values(by=['patient_id', 'eeg_id'], ascending=[True, True]).reset_index(drop=True)\n            \n            #print(f'==NUM_FEAT: {len(new_col_labels_all)}==')\n            #pp\n            model_dir_pp                     = f'{Foldset_path}/pp/{idx_pp}/'\n            exp_name                         = f'Foldset_{idx_Foldset}_{id_pp}'\n            calc_inference_pp(df_feat_pp,new_col_labels_all,col_labels,\n                              model_dir_pp,exp_name,config_common)\n\n        #=========4.ens2=========#\n        print(f'=======4.ens2==========')\n        #=ens2_all=#\n        df_ens2_all                         = pd.read_csv(f'{Foldset_path}/ens2/df_ens2_all.csv')\n        df_ens2_all.loc[df_ens2_all['model_type'] == 'STACK', 'exp_id'] = f'Foldset_{idx_Foldset}_' + df_ens2_all['exp_id']\n        df_ens2_all.loc[df_ens2_all['model_type'] == 'PP', 'exp_id'] = f'Foldset_{idx_Foldset}_' + df_ens2_all['exp_id']\n        #calc_ens2\n        df_test_output_ens2_all             = calc_weight_ave(df_ens2_all) \n        df_test_output_ens2_all             = df_test_output_ens2_all.fillna(0)\n        #softmax\n        df_test_output_ens2_all[col_labels] = nn.functional.softmax(torch.tensor(df_test_output_ens2_all[col_labels_feat].values) , dim=1).float().detach().cpu()\n        #save\n        exp_name                            = f'Foldset_{idx_Foldset}'\n        df_test_output_ens2_all.to_csv(f'/kaggle/working/output/df_test_output_{exp_name}.csv',index=False)  \n    else:#PP実施しない→ens1_allが最終結果\n        print('PP:OFF')\n        df_test_output_ens2_all             = df_test_output_ens1_all.copy()\n        df_test_output_ens2_all             = df_test_output_ens2_all.fillna(0)\n        #softmax\n        df_test_output_ens2_all[col_labels] = nn.functional.softmax(torch.tensor(df_test_output_ens2_all[col_labels_feat].values) , dim=1).float().detach().cpu()\n        #save\n        exp_name                            = f'Foldset_{idx_Foldset}'\n        df_test_output_ens2_all.to_csv(f'/kaggle/working/output/df_test_output_{exp_name}.csv',index=False)   \n        ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'=======5.ens3==========')\nweight                             = [0.5,0.5]\n#weight                             = [0.0,1.0]\ndf_test_output_ens2_Foldset_0      = pd.read_csv(f'/kaggle/working/output/df_test_output_Foldset_0.csv')\ndf_test_output_ens2_Foldset_1      = pd.read_csv(f'/kaggle/working/output/df_test_output_Foldset_1.csv')\n\ndf_test_output_ens3                = df_test_output_ens2_Foldset_0.copy()\ndf_test_output_ens3[col_labels]         = weight[0]*df_test_output_ens2_Foldset_0[col_labels].values \\\n                                         + weight[1]*df_test_output_ens2_Foldset_1[col_labels].values\ndf_test_output_ens3[col_labels_feat]    = weight[0]*df_test_output_ens2_Foldset_0[col_labels_feat].values \\\n                                         + weight[1]*df_test_output_ens2_Foldset_1[col_labels_feat].values\n\ndf_test_output_ens3","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"sample_submission = pd.read_csv('/kaggle/input/hms-harmful-brain-activity-classification/sample_submission.csv')\nprint(sample_submission.dtypes)\nsample_submission","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(df_test_output_ens3[['eeg_id']+col_labels].dtypes)\ndf_test_output_ens3[['eeg_id']+col_labels]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(df_test_output_ens3[['eeg_id']+col_labels_logits].dtypes)\ndf_test_output_ens3[['eeg_id']+col_labels_logits]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission        = pd.merge(sample_submission[[\"eeg_id\"]], df_test_output_ens3[['eeg_id']+col_labels], on=\"eeg_id\", how=\"left\")\ndf_submission.to_csv(\"submission_uemu_kazumax_th.csv\", index=False)\ndf_submission","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission_logits        = pd.merge(sample_submission[[\"eeg_id\"]], df_test_output_ens3[['eeg_id']+col_labels_logits], on=\"eeg_id\", how=\"left\")\ndf_submission_logits        = df_submission_logits.rename(columns={\n                                                                'seizure_logits':'seizure_vote',\n                                                                'lpd_logits':'lpd_vote',\n                                                                'gpd_logits':'gpd_vote',\n                                                                'lrda_logits':'lrda_vote',\n                                                                'grda_logits':'grda_vote',\n                                                                'other_logits':'other_vote',\n                                                               })\ndf_submission_logits.to_csv(\"submission_uemu_kazumax_th_logits.csv\", index=False)\ndf_submission_logits","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -r /kaggle/logs\n!rm -r  output","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Team Submission (DImanishi/MaxChen/UEMU_kazumax_TH)","metadata":{}},{"cell_type":"code","source":"mode_weightave           = 'before_softmax_averaging' #'after_softmax_averaging' #'after_softmax_averaging' 'before_softmax_averaging'\nweight                   = [0.37,0.15,0.48]\n#weight                   = [0.2,0.2,0.6] \ncol_labels               = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\n#col_labels_logits        = [s.split('_')[0]+'_logits' for s in col_labels]\n\nif mode_weightave =='after_softmax_averaging':\n    df_submission_DImanishi  = pd.read_csv('./submission_DImanishi.csv').sort_values(by=['eeg_id'], ascending=[True]).reset_index(drop=True)\n    df_submission_MaxChen    = pd.read_csv('./submission_maxc.csv').sort_values(by=['eeg_id'], ascending=[True]).reset_index(drop=True)\n    df_submission_UEMU       = pd.read_csv('./submission_uemu_kazumax_th.csv').sort_values(by=['eeg_id'], ascending=[True]).reset_index(drop=True)\n    #averaging\n    df_submission_Team             = df_submission_DImanishi.copy()\n    df_submission_Team[col_labels] = weight[0]*df_submission_DImanishi[col_labels].values \\\n                                    +weight[1]*df_submission_MaxChen[col_labels].values \\\n                                    +weight[2]*df_submission_UEMU[col_labels].values\n    #save\n    df_submission_Team.to_csv(\"submission.csv\", index=False)\n    \nelif mode_weightave =='before_softmax_averaging':\n    #read logits\n    df_submission_DImanishi_logits  = pd.read_csv('./submission_DImanishi_logits.csv').sort_values(by=['eeg_id'], ascending=[True]).reset_index(drop=True)\n    df_submission_MaxChen_logits    = pd.read_csv('./subm_logits_maxc.csv').sort_values(by=['eeg_id'], ascending=[True]).reset_index(drop=True)\n    df_submission_UEMU_logits       = pd.read_csv('./submission_uemu_kazumax_th_logits.csv').sort_values(by=['eeg_id'], ascending=[True]).reset_index(drop=True)\n    \n    print(df_submission_DImanishi_logits.head())\n    print(df_submission_MaxChen_logits.head())\n    print(df_submission_UEMU_logits.head())\n    \n    #averaging\n    df_submission_Team             = df_submission_DImanishi_logits.copy()\n    df_submission_Team[col_labels] = weight[0]*df_submission_DImanishi_logits[col_labels].values \\\n                                    +weight[1]*df_submission_MaxChen_logits[col_labels].values \\\n                                    +weight[2]*df_submission_UEMU_logits[col_labels].values\n    #softmax\n    def softmax_second_axis(x):\n        e_x = np.exp(x - np.max(x, axis=1, keepdims=True))\n        return e_x / np.sum(e_x, axis=1, keepdims=True)\n    df_submission_Team[col_labels] = softmax_second_axis(df_submission_Team[col_labels].values)\n    #save\n    df_submission_Team.to_csv(\"submission.csv\", index=False)\n\nelse:\n    print('error')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_submission_Team","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}