{"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":8003751,"sourceType":"datasetVersion","datasetId":4392176},{"sourceId":8020917,"sourceType":"datasetVersion","datasetId":4525340},{"sourceId":8020921,"sourceType":"datasetVersion","datasetId":4525342},{"sourceId":8059460,"sourceType":"datasetVersion","datasetId":4406455},{"sourceId":8096894,"sourceType":"datasetVersion","datasetId":4406459},{"sourceId":8059478,"sourceType":"datasetVersion","datasetId":4406986},{"sourceId":8059486,"sourceType":"datasetVersion","datasetId":4433517},{"sourceId":8109601,"sourceType":"datasetVersion","datasetId":1709845},{"sourceId":8062956,"sourceType":"datasetVersion","datasetId":4433525},{"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":8109608,"sourceType":"datasetVersion","datasetId":4392173},{"sourceId":8094967,"sourceType":"datasetVersion","datasetId":4406452},{"sourceId":8109591,"sourceType":"datasetVersion","datasetId":4406449}],"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-08T17:02:04.387938Z","iopub.execute_input":"2024-04-08T17:02:04.388835Z","iopub.status.idle":"2024-04-08T17:02:04.393804Z","shell.execute_reply.started":"2024-04-08T17:02:04.388805Z","shell.execute_reply":"2024-04-08T17:02:04.392906Z"},"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-08T17:02:04.396055Z","iopub.execute_input":"2024-04-08T17:02:04.396697Z","iopub.status.idle":"2024-04-08T17:02:07.100559Z","shell.execute_reply.started":"2024-04-08T17:02:04.396628Z","shell.execute_reply":"2024-04-08T17:02:07.099455Z"},"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-08T17:02:07.102196Z","iopub.execute_input":"2024-04-08T17:02:07.102586Z","iopub.status.idle":"2024-04-08T17:02:49.464981Z","shell.execute_reply.started":"2024-04-08T17:02:07.102547Z","shell.execute_reply":"2024-04-08T17:02:49.463821Z"},"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-08T17:02:49.466742Z","iopub.execute_input":"2024-04-08T17:02:49.467135Z","iopub.status.idle":"2024-04-08T17:02:49.472266Z","shell.execute_reply.started":"2024-04-08T17:02:49.467089Z","shell.execute_reply":"2024-04-08T17:02:49.471329Z"},"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-08T17:02:49.474955Z","iopub.execute_input":"2024-04-08T17:02:49.475257Z","iopub.status.idle":"2024-04-08T17:02:49.493911Z","shell.execute_reply.started":"2024-04-08T17:02:49.475232Z","shell.execute_reply":"2024-04-08T17:02:49.493009Z"},"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-08T17:02:49.495043Z","iopub.execute_input":"2024-04-08T17:02:49.495314Z","iopub.status.idle":"2024-04-08T17:02:49.508274Z","shell.execute_reply.started":"2024-04-08T17:02:49.495291Z","shell.execute_reply":"2024-04-08T17:02:49.507218Z"},"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-10T03:43:23.124308Z","iopub.execute_input":"2024-04-10T03:43:23.125175Z","iopub.status.idle":"2024-04-10T03:43:23.684275Z","shell.execute_reply.started":"2024-04-10T03:43:23.125139Z","shell.execute_reply":"2024-04-10T03:43:23.682871Z"},"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-08T17:02:49.525332Z","iopub.execute_input":"2024-04-08T17:02:49.525868Z","iopub.status.idle":"2024-04-08T17:02:49.538950Z","shell.execute_reply.started":"2024-04-08T17:02:49.525842Z","shell.execute_reply":"2024-04-08T17:02:49.537965Z"},"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-08T17:02:49.565303Z","iopub.execute_input":"2024-04-08T17:02:49.565578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"## efficientnet b0 256 will use 13.6G for safe could use 128\nmodels = [0]\n# models = [3]\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\nensembler = gezi.Ensembler()\n# for model, weight, bs, folds in tqdm(zip(models, weights, bss, foldss)):\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":{"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":{"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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(x['pred'])","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.array(list(x['pred'])).sum(-1).mean()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x['pred'][-1]","metadata":{"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":{"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":{"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":{"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":[]}]}