{"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":"nvidiaTeslaT4","dataSources":[{"sourceId":59093,"databundleVersionId":7469972,"sourceType":"competition"},{"sourceId":2799445,"sourceType":"datasetVersion","datasetId":1710071},{"sourceId":6523931,"sourceType":"datasetVersion","datasetId":3771565},{"sourceId":7669384,"sourceType":"datasetVersion","datasetId":4426767},{"sourceId":7777950,"sourceType":"datasetVersion","datasetId":4551275},{"sourceId":8020917,"sourceType":"datasetVersion","datasetId":4525340},{"sourceId":8020921,"sourceType":"datasetVersion","datasetId":4525342},{"sourceId":8063239,"sourceType":"datasetVersion","datasetId":4433528},{"sourceId":8064887,"sourceType":"datasetVersion","datasetId":4433532},{"sourceId":8064907,"sourceType":"datasetVersion","datasetId":4443602},{"sourceId":8064924,"sourceType":"datasetVersion","datasetId":4443604},{"sourceId":8064946,"sourceType":"datasetVersion","datasetId":4443615},{"sourceId":8064969,"sourceType":"datasetVersion","datasetId":4443616},{"sourceId":8064982,"sourceType":"datasetVersion","datasetId":4499095},{"sourceId":8064993,"sourceType":"datasetVersion","datasetId":4499099},{"sourceId":8064994,"sourceType":"datasetVersion","datasetId":4499139},{"sourceId":8065038,"sourceType":"datasetVersion","datasetId":4525291},{"sourceId":8065040,"sourceType":"datasetVersion","datasetId":4525223},{"sourceId":8065045,"sourceType":"datasetVersion","datasetId":4525301},{"sourceId":8065056,"sourceType":"datasetVersion","datasetId":4525302},{"sourceId":8065653,"sourceType":"datasetVersion","datasetId":4525281},{"sourceId":8065701,"sourceType":"datasetVersion","datasetId":4525343},{"sourceId":8095529,"sourceType":"datasetVersion","datasetId":4392176},{"sourceId":8099282,"sourceType":"datasetVersion","datasetId":4406986},{"sourceId":8100014,"sourceType":"datasetVersion","datasetId":4433517},{"sourceId":8100601,"sourceType":"datasetVersion","datasetId":4433525},{"sourceId":8109591,"sourceType":"datasetVersion","datasetId":4406449},{"sourceId":8109601,"sourceType":"datasetVersion","datasetId":1709845},{"sourceId":8120289,"sourceType":"datasetVersion","datasetId":4406452},{"sourceId":8112767,"sourceType":"datasetVersion","datasetId":4406455},{"sourceId":8114328,"sourceType":"datasetVersion","datasetId":4406459},{"sourceId":8114359,"sourceType":"datasetVersion","datasetId":4392173}],"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"MODEL_NAME = 'harmfull-brain-activity'\nSRC_NAME = f'{MODEL_NAME}-code'\nPYTHONPATH = 'PYTHONPATH=.:/kaggle:/kaggle/input/pikachu/utils:/kaggle/input/pikachu/third:$PYTHONPATH'","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-04-12T08:46:34.927369Z","iopub.execute_input":"2024-04-12T08:46:34.927916Z","iopub.status.idle":"2024-04-12T08:46:34.940042Z","shell.execute_reply.started":"2024-04-12T08:46:34.927886Z","shell.execute_reply":"2024-04-12T08:46:34.939101Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os, sys\nimport socket\n!rm -rf ./src\n!cp -rf ../input/{SRC_NAME} ./src\nif os.path.exists('/kaggle') and socket.gethostname() != 'gezi':\n  sys.path.append('/kaggle/input/pikachu/utils')\n  sys.path.append('/kaggle/input/pikachu/third')\n  sys.path.append('.')\nelse:\n  sys.path.append('..')\n  sys.path.append('../../../../utils')\n  sys.path.append('../../../../third')","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:46:34.942081Z","iopub.execute_input":"2024-04-12T08:46:34.942429Z","iopub.status.idle":"2024-04-12T08:46:38.186428Z","shell.execute_reply.started":"2024-04-12T08:46:34.942390Z","shell.execute_reply":"2024-04-12T08:46:38.185249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q icecream --no-index --find-links=file:///kaggle/input/icecream/     \n!pip install -q pandarallel --no-index --find-links=file:///kaggle/input/pandarallel/ \n#注意本地的absl-py和kaggle上一定版本一致 否则flags加载会失败\n!pip uninstall -y absl-py\n!pip install absl-py --no-index --find-links=file:///kaggle/input/absl-py/ \n!wandb disabled","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:46:38.188054Z","iopub.execute_input":"2024-04-12T08:46:38.188447Z","iopub.status.idle":"2024-04-12T08:47:19.890775Z","shell.execute_reply.started":"2024-04-12T08:46:38.188404Z","shell.execute_reply":"2024-04-12T08:47:19.889811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # remove fcwt-py/pyproject.toml\n# !rm -rf /kaggle/working/fcwt-py \n# !cp -rf /kaggle/input/fcwt-py/  /kaggle/working\n# !pip install /kaggle/working/fcwt-py/     \n# !rm -rf /kaggle/working/fcwt-py ","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:47:19.893400Z","iopub.execute_input":"2024-04-12T08:47:19.893717Z","iopub.status.idle":"2024-04-12T08:47:19.898324Z","shell.execute_reply.started":"2024-04-12T08:47:19.893687Z","shell.execute_reply":"2024-04-12T08:47:19.897349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from datetime import datetime, timedelta, timezone\n\nSHA_TZ = timezone(\n  timedelta(hours=8),\n  name='Asia/Shanghai',\n)\nutc_now = datetime.utcnow().replace(tzinfo=timezone.utc)\nbeijing_now = utc_now.astimezone(SHA_TZ)\nbeijing_now.strftime('%Y-%m-%d %H:%M:%S')","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:47:19.899414Z","iopub.execute_input":"2024-04-12T08:47:19.899672Z","iopub.status.idle":"2024-04-12T08:47:19.912374Z","shell.execute_reply.started":"2024-04-12T08:47:19.899650Z","shell.execute_reply":"2024-04-12T08:47:19.911495Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile ./src/infer.py \n#!/usr/bin/env python \n# -*- coding: utf-8 -*-\n# ==============================================================================\n#          \\file   infer.py\n#        \\author   chenghuige  \n#          \\date   2022-05-15 07:00:15.332073\n#   \\Description  \n# ==============================================================================\n# py ./infer.py --model_dir=../working/offline/4/0/torchbase.2 --batch_size=64\n  \nfrom __future__ import absolute_import\nfrom __future__ import division\nfrom __future__ import print_function\n\nimport os, sys\nimport socket\nif os.path.exists('/kaggle') and socket.gethostname() != 'gezi':\n  sys.path.append('/kaggle/input/pikachu/utils')\n  sys.path.append('/kaggle/input/pikachu/third')\n  sys.path.append('.')\nelse:\n  sys.path.append('..')\n  sys.path.append('../../../../utils')\n  sys.path.append('../../../../third')\n\nfrom gezi.common import *\nimport gc\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\nos.environ[\"NCCL_DEBUG\"] = 'WARNING'\n\nfrom src.config import *\nfrom src.preprocess import *\nfrom src import util\n\nflags.DEFINE_string('ofile', '../working/x.pkl', '')\n\n# ic = print\ndef main(argv):\n  ic('infer.py start')\n  model_dir = FLAGS.model_dir\n  root = model_dir\n  model_dir = f'{model_dir}/0'\n  bs = FLAGS.eval_bs\n  out_file = FLAGS.ofile\n#   infer_norm = FLAGS.infer_norm\n  ic(bs)\n  gezi.restore_configs(model_dir, ignores=gezi.get_commandline_flags())\n  ic(mt.eval_batch_size())\n  # FLAGS.clear_first = False\n  FLAGS.mns = ''\n  FLAGS.distributed = False\n  FLAGS.train_allnew = False\n  FLAGS.grad_acc = 1\n  FLAGS.restore_configs = False\n  FLAGS.bs, FLAGS.eval_bs = bs, bs\n  FLAGS.num_gpus = None\n  \n  mt.set_global('eval_batch_size', bs)\n  ic(mt.eval_batch_size(), gezi.eval_batch_size())\n## TODO FIXME notice if remove mt.init test/eval dataset using bs mt.eval_batch_size() here will always be 64 even if you set --eval_bs=128\n## but with mt.init might cause 8min model running 13min?\n#   mt.init()\n  from pandarallel import pandarallel as pdl\n  pdl.initialize(nb_workers=FLAGS.pdl_workers, progress_bar=True)\n  FLAGS.model_dir = model_dir\n  FLAGS.pymp = False\n  FLAGS.num_workers = 2\n  FLAGS.pin_memory = True\n  FLAGS.persistent_workers = True\n  FLAGS.workers = 1\n  \n  FLAGS.convert = True\n  FLAGS.onfly = True\n  if FLAGS.hack_infer:\n    FLAGS.dynamic_spec = False\n  \n  show()\n  FLAGS.mode = 'test'\n  FLAGS.work_mode = 'test'\n  FLAGS.model_dir = model_dir\n  ic(FLAGS.model_dir, os.path.exists(f'{FLAGS.model_dir}/model.pt'))\n\n  model = util.get_model()\n  \n  ensembler = gezi.Ensembler()\n  num_models = int(open(f'{root}/num_models.txt').readline().strip())\n  ic(num_models)\n  for fold in tqdm(range(num_models), desc='folds'):\n    model_dir = f'{root}/{fold}'\n    try:\n      display(pd.read_csv(f'{model_dir}/metrics.csv'))\n      ic(open(f'{root}/path.txt').readline())\n    except Exception as e:\n      logger.warning(e)\n    logger.info('before load weights')\n    ic(gezi.get_mem_gb())\n    gezi.load_weights(model, model_dir, strict=False)\n    ic(gezi.get_mem_gb())\n    logger.info('load weights done')\n    \n    # 每个模型重新生成 test dataset 不复用 目的是 FLAGS.rand_infer = True, test dataset具有多样性 \n    logger.info('before get test_ds')\n    from src.dataset import get_datasets\n    test_ds = get_datasets(mode='test')\n    logger.info('after get test_ds')\n\n    logger.info('before predict')\n    out_keys = [] if not FLAGS.hack_infer else ['label']\n    x = lele.predict(model, test_ds, out_keys=out_keys, amp=True, fp16=True)\n    ic(x['pred'][0])\n    \n    pred = np.asarray(x['pred'], dtype=np.float32)\n    x['pred'] = list(gezi.torch_softmax(pred))\n    ic(x['pred'][0])\n    \n    logger.info('predict done')\n    ensembler.add(x)\n  x = ensembler.finalize()\n\n  gezi.save(x, out_file)\n  \n  gc.collect()\n  torch.cuda.empty_cache()\n\nif __name__ == '__main__':\n  app.run(main)  ","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:47:19.914870Z","iopub.execute_input":"2024-04-12T08:47:19.915218Z","iopub.status.idle":"2024-04-12T08:47:19.924554Z","shell.execute_reply.started":"2024-04-12T08:47:19.915187Z","shell.execute_reply":"2024-04-12T08:47:19.923639Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%writefile ./src/infer.py \n#!/usr/bin/env python \n# -*- coding: utf-8 -*-\n# ==============================================================================\n#          \\file   infer.py\n#        \\author   chenghuige  \n#          \\date   2022-05-15 07:00:15.332073\n#   \\Description  \n# ==============================================================================\n# py ./infer.py --model_dir=../working/offline/4/0/torchbase.2 --batch_size=64\n  \nfrom __future__ import absolute_import\nfrom __future__ import division\nfrom __future__ import print_function\n\nimport os, sys\nimport socket\nif os.path.exists('/kaggle') and socket.gethostname() != 'gezi':\n  sys.path.append('/kaggle/input/pikachu/utils')\n  sys.path.append('/kaggle/input/pikachu/third')\n  sys.path.append('.')\nelse:\n  sys.path.append('..')\n  sys.path.append('../../../../utils')\n  sys.path.append('../../../../third')\n\nfrom gezi.common import *\nimport gc\nos.environ[\"TOKENIZERS_PARALLELISM\"] = \"false\"\nos.environ[\"NCCL_DEBUG\"] = 'WARNING'\n\nfrom src.config import *\nfrom src.preprocess import *\nfrom src import util\n\nflags.DEFINE_string('ofile', '../working/x.pkl', '')\n\n# ic = print\ndef main(argv):\n  ic('infer.py start')\n  model_dir = FLAGS.model_dir\n  root = model_dir\n  model_dir = f'{model_dir}/0'\n  bs = FLAGS.eval_bs\n  out_file = FLAGS.ofile\n#   infer_norm = FLAGS.infer_norm\n  ic(bs)\n  gezi.restore_configs(model_dir, ignores=gezi.get_commandline_flags())\n  ic(mt.eval_batch_size())\n  # FLAGS.clear_first = False\n  FLAGS.mns = ''\n  FLAGS.distributed = False\n  FLAGS.train_allnew = False\n  FLAGS.grad_acc = 1\n  FLAGS.restore_configs = False\n  FLAGS.bs, FLAGS.eval_bs = bs, bs\n  FLAGS.num_gpus = None\n  \n  mt.set_global('eval_batch_size', bs)\n  ic(mt.eval_batch_size(), gezi.eval_batch_size())\n## TODO FIXME notice if remove mt.init test/eval dataset using bs mt.eval_batch_size() here will always be 64 even if you set --eval_bs=128\n## but with mt.init might cause 8min model running 13min?\n#   mt.init()\n  from pandarallel import pandarallel as pdl\n  pdl.initialize(nb_workers=FLAGS.pdl_workers, progress_bar=True)\n  FLAGS.model_dir = model_dir\n  FLAGS.pymp = False\n  FLAGS.num_workers = 2\n  FLAGS.pin_memory = True\n  FLAGS.persistent_workers = True\n  FLAGS.workers = 1\n\n  FLAGS.convert = True\n  FLAGS.onfly = True\n  FLAGS.streaming = True\n  FLAGS.tta = True\n#   FLAGS.tta = False\n  if FLAGS.hack_infer:\n    FLAGS.dynamic_spec = False\n  \n  show()\n  FLAGS.mode = 'test'\n  FLAGS.work_mode = 'test'\n  FLAGS.model_dir = model_dir\n  ic(FLAGS.model_dir, os.path.exists(f'{FLAGS.model_dir}/model.pt'))\n\n  logger.info('before get test_ds')\n  from src.dataset import get_datasets\n  test_ds = get_datasets(mode='test')\n  logger.info('after get test_ds')\n\n  num_models = int(open(f'{root}/num_models.txt').readline().strip())\n  # remove online model to see perfromance of 5 folds only\n  #   num_models = 5 \n  ic(num_models)\n  models = []\n  for fold in tqdm(range(num_models), desc='folds'):\n#     if fold != 5:\n#       continue\n    model = util.get_model()\n    model_dir = f'{root}/{fold}'\n    try:\n      display(pd.read_csv(f'{model_dir}/metrics.csv'))\n      ic(open(f'{root}/path.txt').readline())\n    except Exception as e:\n      logger.warning(e)\n    logger.info('before load weights')\n    ic(gezi.get_mem_gb())\n    gezi.load_weights(model, model_dir, strict=False)\n    ic(gezi.get_mem_gb())\n    logger.info('load weights done')\n    models.append(model)\n\n  logger.info('before predict')\n  out_keys = [] if not FLAGS.hack_infer else ['label']\n  def post_deal(x):\n    x['pred'] = list(gezi.torch_softmax(np.array(list(x['pred'])).astype(np.float32)))\n    return x\n#   x = lele.infer(models, test_ds, post_fn=post_deal, amp=True, fp16=True, verbose=1)\n  x = lele.infer(models, test_ds, post_fn=post_deal, amp=False, fp16=False, verbose=1)\n  ic(x['pred'][0])\n  ic(list(x.keys()))\n#   pred = np.asarray(x['pred'], dtype=np.float32)\n#   x['pred'] = list(gezi.torch_softmax(pred))\n\n  logger.info('predict done')\n\n  gezi.save(x, out_file)\n  \n  gc.collect()\n  torch.cuda.empty_cache()\n\nif __name__ == '__main__':\n  app.run(main)  ","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:47:19.925964Z","iopub.execute_input":"2024-04-12T08:47:19.926292Z","iopub.status.idle":"2024-04-12T08:47:19.937693Z","shell.execute_reply.started":"2024-04-12T08:47:19.926270Z","shell.execute_reply":"2024-04-12T08:47:19.936711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TEST_MODE = 0\n# TEST_MODE = 1\nTEST_MODE","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:47:19.938810Z","iopub.execute_input":"2024-04-12T08:47:19.939464Z","iopub.status.idle":"2024-04-12T08:47:19.948187Z","shell.execute_reply.started":"2024-04-12T08:47:19.939433Z","shell.execute_reply":"2024-04-12T08:47:19.947120Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TEST_ENTROPY = 0\n# TEST_ENTROPY = 1\nTEST_ENTROPY","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:47:19.949350Z","iopub.execute_input":"2024-04-12T08:47:19.949759Z","iopub.status.idle":"2024-04-12T08:47:19.956381Z","shell.execute_reply.started":"2024-04-12T08:47:19.949729Z","shell.execute_reply":"2024-04-12T08:47:19.955383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TEST_OTHER = 0\n# TEST_OTHER = 1\nTEST_OTHER","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:47:19.960705Z","iopub.execute_input":"2024-04-12T08:47:19.961033Z","iopub.status.idle":"2024-04-12T08:47:19.966737Z","shell.execute_reply.started":"2024-04-12T08:47:19.961009Z","shell.execute_reply":"2024-04-12T08:47:19.965773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport gezi\ngezi.init_flags()\nfrom tqdm import tqdm","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:47:19.967865Z","iopub.execute_input":"2024-04-12T08:47:19.968194Z","iopub.status.idle":"2024-04-12T08:47:56.831915Z","shell.execute_reply.started":"2024-04-12T08:47:19.968164Z","shell.execute_reply":"2024-04-12T08:47:56.830968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## efficientnet b0 256 will use 13.6G for safe could use 128\nmodels = [1]\n# models = [0,1,2,3]\n# models = [0, 1, 2, 3]\n# models = [0, 1, 3, 4, 5]\n# models = [0, 1, 3, 4, 5, 6]\n# models = [0, 1, 2, 3, 4, 5, 6]\n# models = [6]\n# models = [0,1,2,3,4,5]\n# models = [19]\n# models = list(range(20))\n# models = list(range(22))\nweights = [1] * len(models)\n# weights = [0.2576146 , 0.34746932, 0.89300917, 0.82881567, 0.02542841,\n#        0.92921973, 0.9580988 , 0.67406156, 0.35969988, 0.45420358,\n#        0.90141364, 0.96586689, 0.33583702, 0.56482746, 0.02930522,\n#        0.91234032, 0.38841254, 0.55903399, 0.28086246, 0.98014346,\n#        0.64953999, 0.50522002]\n## mixnet_xl need 32, 64 oom and seems large batch size not improve speed so just set all 32\nbss = [32] * len(models)\nassert len(models) == len(weights)\n\nmodel_dir_ = f'../input/{MODEL_NAME}-model'\nofile = '../working/x.pkl'\npdl_workers = 2\nhack_count = 2700\n# hack_count = 270\n\ni = 0\nfor model, weight, bs in tqdm(zip(models, weights, bss)):\n  model_dir = f'{model_dir_}{model}'\n  model_path = open(f'{model_dir}/path.txt').readline().strip()\n  ic(i, model_dir, model_path, bs, weight)\n  i += 1\n    \nensembler = gezi.Ensembler()\nfor model, weight, bs in tqdm(zip(models, weights, bss)):\n  model_dir = f'{model_dir_}{model}'\n#   command = f'python ./src/infer.py --model_dir={model_dir} --eval_bs={bs} --ofile={ofile} --hack_infer={TEST_MODE} --pdl_workers={pdl_workers} --folds={folds}'\n  command = f'python ./src/infer.py --model_dir={model_dir} --eval_bs={bs} --ofile={ofile} --hack_infer={TEST_MODE} --pdl_workers={pdl_workers} --hack_count={hack_count}'\n  ic(model_dir, bs, command, weight)\n  os.system(command)\n  x = gezi.load(ofile)\n  ic(gezi.get_mem_gb())\n#   ic(x)\n  ensembler.add(x, weight)\nx = ensembler.finalize()","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:47:56.833090Z","iopub.execute_input":"2024-04-12T08:47:56.833626Z","iopub.status.idle":"2024-04-12T08:51:38.465973Z","shell.execute_reply.started":"2024-04-12T08:47:56.833599Z","shell.execute_reply":"2024-04-12T08:51:38.464886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# if TEST_MODE:\n#     from src.eval import calc_metrics\n#     ic(calc_metrics(x['label'], x['pred']))","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:51:38.467394Z","iopub.execute_input":"2024-04-12T08:51:38.468022Z","iopub.status.idle":"2024-04-12T08:51:38.472294Z","shell.execute_reply.started":"2024-04-12T08:51:38.467986Z","shell.execute_reply":"2024-04-12T08:51:38.471319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nentropy = np.average([gezi.logits_entropy(a) for a in x['pred']])\nother = np.average([a[-1] for a in x['pred']])\nic(entropy, other)","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:51:38.473529Z","iopub.execute_input":"2024-04-12T08:51:38.473827Z","iopub.status.idle":"2024-04-12T08:51:38.540152Z","shell.execute_reply.started":"2024-04-12T08:51:38.473802Z","shell.execute_reply":"2024-04-12T08:51:38.539287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(x['pred'])","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:51:38.541197Z","iopub.execute_input":"2024-04-12T08:51:38.541452Z","iopub.status.idle":"2024-04-12T08:51:38.547257Z","shell.execute_reply.started":"2024-04-12T08:51:38.541430Z","shell.execute_reply":"2024-04-12T08:51:38.546287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.array(list(x['pred'])).sum(-1).mean()","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:51:38.548438Z","iopub.execute_input":"2024-04-12T08:51:38.548715Z","iopub.status.idle":"2024-04-12T08:51:38.555517Z","shell.execute_reply.started":"2024-04-12T08:51:38.548692Z","shell.execute_reply":"2024-04-12T08:51:38.554629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x['pred'][-1]","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:51:38.556565Z","iopub.execute_input":"2024-04-12T08:51:38.556861Z","iopub.status.idle":"2024-04-12T08:51:38.563855Z","shell.execute_reply.started":"2024-04-12T08:51:38.556817Z","shell.execute_reply":"2024-04-12T08:51:38.562809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TEST_ENTROPY and len(x['pred']) > 1:\n    from time import sleep\n    if entropy > 1:\n      print('large entropy no sleep')\n    elif entropy > 0.98:\n      dur = 5 * 60\n      print(f'sleep {dur}')\n      sleep(dur)\n    elif entropy > 0.95:\n      dur = 10 * 60\n      print(f'sleep {dur}')\n      sleep(dur)\n    elif entropy > 0.9:\n      dur = 15 * 60\n      print(f'sleep {dur}')\n      sleep(dur)\n    elif entropy > 0.85:\n      dur = 20 * 60\n      print(f'sleep {dur}')\n      sleep(dur)\n    elif entropy > 0.8:\n      dur = 25 * 60\n      print(f'sleep {dur}')\n      sleep(dur)\n    elif entropy > 0.75:\n      dur = 30 * 60\n      print(f'sleep {dur}')\n      sleep(dur)\n    elif entropy > 0.7:\n      dur = 35 * 60\n      print(f'sleep {dur}')\n      sleep(dur)\n    elif entropy > 0.65:\n      dur = 40 * 60\n      print(f'sleep {dur}')\n      sleep(dur)\n    else:\n      dur = 45 * 60\n      print(f'sleep {dur}')\n      sleep(dur)","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:51:38.565125Z","iopub.execute_input":"2024-04-12T08:51:38.565440Z","iopub.status.idle":"2024-04-12T08:51:38.576020Z","shell.execute_reply.started":"2024-04-12T08:51:38.565396Z","shell.execute_reply":"2024-04-12T08:51:38.574971Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if TEST_OTHER and len(x['pred']) > 1:\n    from time import sleep\n    if other > 0.43:\n      print('large entropy no sleep')\n    elif other > 0.38:\n      dur = 5 * 60\n      print(f'sleep {dur}')\n      sleep(dur)\n    elif other > 0.375:\n      dur = 10 * 60\n      print(f'sleep {dur}')\n      sleep(dur)\n    elif other > 0.3713:\n      dur = 15 * 60\n      print(f'sleep {dur}')\n      sleep(dur)\n    elif other > 0.37:\n      dur = 20 * 60\n      print(f'sleep {dur}')\n      sleep(dur)\n    elif other > 0.369:\n      dur = 25 * 60\n      print(f'sleep {dur}')\n      sleep(dur)\n    else:\n      dur = 30 * 60\n      print(f'sleep {dur}')\n      sleep(dur)","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:51:38.577283Z","iopub.execute_input":"2024-04-12T08:51:38.577604Z","iopub.status.idle":"2024-04-12T08:51:38.588152Z","shell.execute_reply.started":"2024-04-12T08:51:38.577578Z","shell.execute_reply":"2024-04-12T08:51:38.587179Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv('../input/hms-harmful-brain-activity-classification/test.csv')\nprint(f\"Test dataframe shape is: {test_df.shape}\")\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:51:38.589546Z","iopub.execute_input":"2024-04-12T08:51:38.589906Z","iopub.status.idle":"2024-04-12T08:51:38.607090Z","shell.execute_reply.started":"2024-04-12T08:51:38.589870Z","shell.execute_reply":"2024-04-12T08:51:38.606170Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 如果是logits平均集成 最后取softmax才需要 但是似乎效果比prob集成稍差 \n# pred = np.asarray(x['pred'], dtype=np.float32)\n# x['pred'] = list(gezi.torch_softmax(pred))","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:51:38.608212Z","iopub.execute_input":"2024-04-12T08:51:38.608469Z","iopub.status.idle":"2024-04-12T08:51:38.612385Z","shell.execute_reply.started":"2024-04-12T08:51:38.608447Z","shell.execute_reply":"2024-04-12T08:51:38.611361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TARGETS = ['seizure_vote', 'lpd_vote', 'gpd_vote', 'lrda_vote', 'grda_vote', 'other_vote']\nsub = pd.DataFrame({'eeg_id': test_df.eeg_id.values})\nsub[TARGETS] = x['pred'][:len(test_df)]\nsub.to_csv('submission.csv',index=False)\nprint(f'Submissionn shape: {sub.shape}')\nsub.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-12T08:51:38.613603Z","iopub.execute_input":"2024-04-12T08:51:38.613895Z","iopub.status.idle":"2024-04-12T08:51:38.635974Z","shell.execute_reply.started":"2024-04-12T08:51:38.613869Z","shell.execute_reply":"2024-04-12T08:51:38.635022Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}