{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import gc\nimport os\nimport pickle\nimport glob\n\nfrom text_unidecode import unidecode\nfrom typing import Dict, List, Tuple\nimport codecs\n\nimport numpy as np\nimport pandas as pd\n\nfrom tqdm import tqdm\n\nimport seaborn as sns\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom torch.nn import Parameter\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom transformers import AutoModel, AutoTokenizer, AutoConfig\n\nimport warnings\nwarnings.simplefilter('ignore')","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:12:53.299780Z","iopub.execute_input":"2022-07-28T03:12:53.300168Z","iopub.status.idle":"2022-07-28T03:12:57.725356Z","shell.execute_reply.started":"2022-07-28T03:12:53.300057Z","shell.execute_reply":"2022-07-28T03:12:57.724347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def replace_encoding_with_utf8(error: UnicodeError) -> Tuple[bytes, int]:\n    return error.object[error.start : error.end].encode(\"utf-8\"), error.end\n\n\ndef replace_decoding_with_cp1252(error: UnicodeError) -> Tuple[str, int]:\n    return error.object[error.start : error.end].decode(\"cp1252\"), error.end\n\ncodecs.register_error(\"replace_encoding_with_utf8\", replace_encoding_with_utf8)\ncodecs.register_error(\"replace_decoding_with_cp1252\", replace_decoding_with_cp1252)\n\n\ndef resolve_encodings_and_normalize(text: str) -> str:\n    text = (\n        text.encode(\"raw_unicode_escape\")\n        .decode(\"utf-8\", errors=\"replace_decoding_with_cp1252\")\n        .encode(\"cp1252\", errors=\"replace_encoding_with_utf8\")\n        .decode(\"utf-8\", errors=\"replace_decoding_with_cp1252\")\n    )\n    \n    text = unidecode(text)\n    \n    return text\n\n\ndef fetch_essay(essay_id: str, txt_dir: str):\n    essay_path = os.path.join(COMP_DIR + txt_dir, essay_id + '.txt')\n    essay_text = open(essay_path, 'r').read()\n    \n    return essay_text","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:12:57.727107Z","iopub.execute_input":"2022-07-28T03:12:57.727465Z","iopub.status.idle":"2022-07-28T03:12:57.738123Z","shell.execute_reply.started":"2022-07-28T03:12:57.727428Z","shell.execute_reply":"2022-07-28T03:12:57.737055Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.set_option('display.precision', 4)\ncm = sns.light_palette('green', as_cmap=True)\nprops_param = \"color:white; font-weight:bold; background-color:green;\"\n\nN_ROW = 10\n\nCOMP_DIR = \"../input/feedback-prize-effectiveness/\"\nDEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:12:57.739397Z","iopub.execute_input":"2022-07-28T03:12:57.739642Z","iopub.status.idle":"2022-07-28T03:12:57.754764Z","shell.execute_reply.started":"2022-07-28T03:12:57.739615Z","shell.execute_reply":"2022-07-28T03:12:57.753656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_path = COMP_DIR + \"test.csv\"\nsubmission_path = COMP_DIR + \"sample_submission.csv\"\n\ntest_origin = pd.read_csv(test_path)\nsubmission_origin = pd.read_csv(submission_path)","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:12:57.757059Z","iopub.execute_input":"2022-07-28T03:12:57.757390Z","iopub.status.idle":"2022-07-28T03:12:57.796867Z","shell.execute_reply.started":"2022-07-28T03:12:57.757346Z","shell.execute_reply":"2022-07-28T03:12:57.795960Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp=pd.read_csv('../input/feedback-prize-effectiveness/train.csv')","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:12:57.798406Z","iopub.execute_input":"2022-07-28T03:12:57.798974Z","iopub.status.idle":"2022-07-28T03:12:58.303898Z","shell.execute_reply.started":"2022-07-28T03:12:57.798922Z","shell.execute_reply":"2022-07-28T03:12:58.302604Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp['discourse_text_UPD'] = temp['discourse_text'].apply(resolve_encodings_and_normalize)\n\ntemp['essay_text'] = temp['essay_id'].transform(fetch_essay, txt_dir='train')\ntemp['essay_text_UPD'] = temp['essay_text'].apply(resolve_encodings_and_normalize)","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:12:58.305104Z","iopub.execute_input":"2022-07-28T03:12:58.305344Z","iopub.status.idle":"2022-07-28T03:14:16.645112Z","shell.execute_reply.started":"2022-07-28T03:12:58.305316Z","shell.execute_reply":"2022-07-28T03:14:16.644085Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp2=pd.read_csv('../input/feedback-prize-effectiveness/test.csv')\ntemp2['discourse_text_UPD'] = temp2['discourse_text'].apply(resolve_encodings_and_normalize)\n\ntemp2['essay_text'] = temp2['essay_id'].transform(fetch_essay, txt_dir='test')\ntemp2['essay_text_UPD'] = temp2['essay_text'].apply(resolve_encodings_and_normalize)","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:14:16.646804Z","iopub.execute_input":"2022-07-28T03:14:16.647057Z","iopub.status.idle":"2022-07-28T03:14:16.688672Z","shell.execute_reply.started":"2022-07-28T03:14:16.647027Z","shell.execute_reply":"2022-07-28T03:14:16.687445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:14:16.690259Z","iopub.execute_input":"2022-07-28T03:14:16.690609Z","iopub.status.idle":"2022-07-28T03:14:16.722424Z","shell.execute_reply.started":"2022-07-28T03:14:16.690574Z","shell.execute_reply":"2022-07-28T03:14:16.721532Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"temp2.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:14:16.723558Z","iopub.execute_input":"2022-07-28T03:14:16.723987Z","iopub.status.idle":"2022-07-28T03:14:16.741230Z","shell.execute_reply.started":"2022-07-28T03:14:16.723953Z","shell.execute_reply":"2022-07-28T03:14:16.740178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport tensorflow as tf\nimport transformers\nimport pathlib","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:14:16.744908Z","iopub.execute_input":"2022-07-28T03:14:16.745199Z","iopub.status.idle":"2022-07-28T03:14:23.453720Z","shell.execute_reply.started":"2022-07-28T03:14:16.745168Z","shell.execute_reply":"2022-07-28T03:14:23.452727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"try:\n    tpu=tf.distribute.cluster_resolver.TPUClusterResolver()\n    print('Device:',tpu.master())\n    tf.config.experimental_connect_to_cluster(tpu)\n    tf.tpu.experimental.initialize_tpu_system(tpu)\n    strategy=tf.distribute.experimental.TPUStrategy(tpu)\nexcept:\n    print(\"TPU Failed\")\n    tpu=None\n    strategy=tf.distribute.get_strategy()\nprint(\"number of replicas\", strategy.num_replicas_in_sync)","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:14:23.455063Z","iopub.execute_input":"2022-07-28T03:14:23.455994Z","iopub.status.idle":"2022-07-28T03:14:29.792861Z","shell.execute_reply.started":"2022-07-28T03:14:23.455945Z","shell.execute_reply":"2022-07-28T03:14:29.791613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Config:\n    seed = 887\n    model_name = \"tpu_bert_large_v4\"\n    n_fold = 5\n    one_fold = False\n    # inputs\n    input_dir = pathlib.Path(\"/kaggle/input/feedback-prize-effectiveness/\")\n    path_train = input_dir / \"train.csv\"\n    train_dir = input_dir / \"train\"\n    path_test = input_dir / \"test.csv\"\n    test_dir = input_dir / \"test\"\n    path_submission = input_dir / \"sample_submission.csv\"\n    labels = [\"Ineffective\", \"Adequate\", \"Effective\"]\n    label_dict = {v: i for i, v in enumerate(labels)}\n    num_classes = len(labels)\n    id_col = \"discourse_id\"\n    # model\n    pretrained = \"bert-large-cased\"\n    pretrained_dir = pathlib.Path(\"/kaggle/working/pretrained\")\n    max_len = 512\n    dropout = 0.4\n    # train\n    label_smoothing = 0.1\n    learning_rate = 3e-6\n    batch_size = 128\n    steps_per_epoch = 100\n    epochs = 100  # 50\n    patience = 10\n    verbose = 2\n    \ncfg = Config()","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:14:29.794538Z","iopub.execute_input":"2022-07-28T03:14:29.795262Z","iopub.status.idle":"2022-07-28T03:14:29.803766Z","shell.execute_reply.started":"2022-07-28T03:14:29.795221Z","shell.execute_reply":"2022-07-28T03:14:29.803046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tokenizer=transformers.AutoTokenizer.from_pretrained(cfg.pretrained)\ntokenizer.save_pretrained(cfg.pretrained_dir)\nconfig=transformers.AutoConfig.from_pretrained(cfg.pretrained)\nconfig.save_pretrained(cfg.pretrained_dir)\nbase_model=transformers.TFAutoModel.from_pretrained(cfg.pretrained,config=config,from_pt=True)\nbase_model.save_pretrained(cfg.pretrained_dir)\nos.listdir(cfg.pretrained_dir)","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:14:29.804750Z","iopub.execute_input":"2022-07-28T03:14:29.804990Z","iopub.status.idle":"2022-07-28T03:15:53.251480Z","shell.execute_reply.started":"2022-07-28T03:14:29.804963Z","shell.execute_reply":"2022-07-28T03:15:53.250472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Preparation","metadata":{}},{"cell_type":"markdown","source":"Load data and prepare the text discourse_text_UPD +' '+ essay_text_UPD+'[SEP]'+  discourse_type\n\nThe Encoder will encode the above trail data in the form of [CLS]+discoure_text_UPD+' '+essay_text_UPD+[SEP]+discoures_type+[SEP]","metadata":{}},{"cell_type":"code","source":"data=pd.read_csv(cfg.path_train)\ndata['label']=data['discourse_effectiveness'].map(cfg.label_dict)\ndata","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:15:53.252879Z","iopub.execute_input":"2022-07-28T03:15:53.253200Z","iopub.status.idle":"2022-07-28T03:15:53.763781Z","shell.execute_reply.started":"2022-07-28T03:15:53.253158Z","shell.execute_reply":"2022-07-28T03:15:53.762820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tokenizer=transformers.AutoTokenizer.from_pretrained(cfg.pretrained_dir)\ntokenizer","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:15:53.765164Z","iopub.execute_input":"2022-07-28T03:15:53.765404Z","iopub.status.idle":"2022-07-28T03:15:53.844433Z","shell.execute_reply.started":"2022-07-28T03:15:53.765375Z","shell.execute_reply":"2022-07-28T03:15:53.843583Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data['text']=temp['discourse_type']+tokenizer.sep_token+temp['discourse_text_UPD']+tokenizer.sep_token+temp['essay_text_UPD']\ndata","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:15:53.845725Z","iopub.execute_input":"2022-07-28T03:15:53.846154Z","iopub.status.idle":"2022-07-28T03:15:53.940524Z","shell.execute_reply.started":"2022-07-28T03:15:53.846125Z","shell.execute_reply":"2022-07-28T03:15:53.939546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rec = data.sample(n=1).iloc[0].to_dict()\nrec","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:15:53.942219Z","iopub.execute_input":"2022-07-28T03:15:53.942589Z","iopub.status.idle":"2022-07-28T03:15:53.967591Z","shell.execute_reply.started":"2022-07-28T03:15:53.942547Z","shell.execute_reply":"2022-07-28T03:15:53.966279Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"original\",rec['text'])\nprint(\"tokenized:\",tokenizer.tokenize(rec['text']))\nprint(\"encoded:\",tokenizer.encode_plus(rec['text']))\nprint(\"decoded\",tokenizer.decode(tokenizer.encode_plus(rec['text'])['input_ids']))","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:15:53.969261Z","iopub.execute_input":"2022-07-28T03:15:53.969568Z","iopub.status.idle":"2022-07-28T03:15:54.009319Z","shell.execute_reply.started":"2022-07-28T03:15:53.969533Z","shell.execute_reply":"2022-07-28T03:15:54.008158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"markdown","source":"We will apply the tokenizer.encode_plus to all the text in the dataset","metadata":{}},{"cell_type":"code","source":"options=tf.data.Options()\noptions.experimental_distribute.auto_shard_policy=tf.data.experimental.AutoShardPolicy.OFF\n\ndef encode_text(text):\n    \"Encode text with tokenizer and return numpy mode text\"\n    encoded=tokenizer.batch_encode_plus(text,\n                                      max_length=cfg.max_len,\n                                      padding='max_length',\n                                      truncation=True,\n                                      return_attention_mask=True,\n                                      return_token_type_ids=True,\n                                      return_tensors=\"tf\")\n    return {\n        \"input_ids\":encoded['input_ids'].numpy(),\n        \"attention_masks\":encoded['attention_mask'].numpy(),\n        \"token_type_ids\":encoded['token_type_ids'].numpy(),\n    }\n\n\ndef get_dataset(data, batch_size=cfg.batch_size, shuffle=False, repeat=False, include_label=True):\n    \"\"\"Get dataset\"\"\"\n    encoded_text = encode_text(data['text'].to_list())\n    tensor_slices = encoded_text\n    if include_label:\n        label = tf.one_hot(data[\"label\"].to_list(), cfg.num_classes)\n        tensor_slices = (encoded_text, label)\n    ds = tf.data.Dataset.from_tensor_slices(tensor_slices)\n    ds = ds.with_options(options)\n    if repeat:\n        ds = ds.repeat()\n    if shuffle:\n        ds = ds.shuffle(2048)\n    ds = ds.batch(batch_size)\n    ds = ds.prefetch(tf.data.AUTOTUNE)\n    return ds","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:15:54.011096Z","iopub.execute_input":"2022-07-28T03:15:54.011698Z","iopub.status.idle":"2022-07-28T03:15:54.025749Z","shell.execute_reply.started":"2022-07-28T03:15:54.011646Z","shell.execute_reply":"2022-07-28T03:15:54.024871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ds = get_dataset(data.iloc[:5], batch_size=2)\nds.element_spec","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:15:54.027620Z","iopub.execute_input":"2022-07-28T03:15:54.028087Z","iopub.status.idle":"2022-07-28T03:15:54.120643Z","shell.execute_reply.started":"2022-07-28T03:15:54.028034Z","shell.execute_reply":"2022-07-28T03:15:54.119727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"elem = next(iter(ds))\nelem","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:15:54.121763Z","iopub.execute_input":"2022-07-28T03:15:54.121986Z","iopub.status.idle":"2022-07-28T03:15:54.144487Z","shell.execute_reply.started":"2022-07-28T03:15:54.121960Z","shell.execute_reply":"2022-07-28T03:15:54.143833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tensorflow.keras import Model, layers, losses, optimizers, metrics, callbacks, backend","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:15:54.145529Z","iopub.execute_input":"2022-07-28T03:15:54.146332Z","iopub.status.idle":"2022-07-28T03:15:54.151640Z","shell.execute_reply.started":"2022-07-28T03:15:54.146293Z","shell.execute_reply":"2022-07-28T03:15:54.150246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class MeanPooler(layers.Layer):\n    def call(self, inputs, mask=None):\n        broadcast_float_mask = tf.expand_dims(tf.cast(mask, \"float32\"), -1)\n        masked_inputs = inputs * broadcast_float_mask\n        inputs_sum = tf.reduce_sum(masked_inputs, axis=1)\n        mask_sum = tf.reduce_sum(broadcast_float_mask, axis=1)\n        mask_sum = tf.math.maximum(mask_sum, 1e-9)\n        return inputs_sum / mask_sum","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:15:54.152975Z","iopub.execute_input":"2022-07-28T03:15:54.153211Z","iopub.status.idle":"2022-07-28T03:15:54.164305Z","shell.execute_reply.started":"2022-07-28T03:15:54.153183Z","shell.execute_reply":"2022-07-28T03:15:54.163198Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_model():\n    # inputs\n    input_ids = layers.Input(shape=(cfg.max_len,), dtype=\"int32\", name=\"input_ids\")\n    attention_masks = layers.Input(shape=(cfg.max_len,), dtype=\"int32\", name=\"attention_masks\")\n    token_type_ids = layers.Input(shape=(cfg.max_len,), dtype=\"int32\", name=\"token_type_ids\")\n    # base_model\n    base_model_config = transformers.AutoConfig.from_pretrained(\n        cfg.pretrained_dir / \"config.json\"\n    )\n    base_model = transformers.TFAutoModel.from_pretrained(\n        cfg.pretrained_dir / \"tf_model.h5\", config=base_model_config\n    )\n    # base_model.trainable = False\n    base_model_output = base_model(\n        input_ids, attention_mask=attention_masks, token_type_ids=token_type_ids\n    )\n    # x = base_model_output.last_hidden_state[:, 0, :]\n    x = MeanPooler()(base_model_output.last_hidden_state, mask=attention_masks)\n    # head\n    x = layers.Dropout(cfg.dropout)(x)\n    output = layers.Dense(cfg.num_classes, activation='softmax')(x)\n    model = Model(\n        inputs=[input_ids, attention_masks, token_type_ids],\n        outputs=output,\n        name=cfg.model_name,\n    )\n    # compile\n    model.compile(\n        optimizer=optimizers.Adam(cfg.learning_rate),\n        loss=losses.CategoricalCrossentropy(label_smoothing=cfg.label_smoothing),\n        metrics=[\"acc\", metrics.CategoricalCrossentropy(name='xentropy')],\n    )\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:15:54.165742Z","iopub.execute_input":"2022-07-28T03:15:54.166010Z","iopub.status.idle":"2022-07-28T03:15:54.178247Z","shell.execute_reply.started":"2022-07-28T03:15:54.165982Z","shell.execute_reply":"2022-07-28T03:15:54.177139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"backend.clear_session()\nwith strategy.scope():\n    model = create_model()\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:15:54.179844Z","iopub.execute_input":"2022-07-28T03:15:54.180221Z","iopub.status.idle":"2022-07-28T03:16:29.392745Z","shell.execute_reply.started":"2022-07-28T03:15:54.180178Z","shell.execute_reply":"2022-07-28T03:16:29.391721Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.predict(elem[0]), elem[1]","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:16:29.394862Z","iopub.execute_input":"2022-07-28T03:16:29.395234Z","iopub.status.idle":"2022-07-28T03:16:45.979925Z","shell.execute_reply.started":"2022-07-28T03:16:29.395190Z","shell.execute_reply":"2022-07-28T03:16:45.978904Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sklearn.metrics as sk_metrics","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:16:46.006790Z","iopub.execute_input":"2022-07-28T03:16:46.008885Z","iopub.status.idle":"2022-07-28T03:16:46.126309Z","shell.execute_reply.started":"2022-07-28T03:16:46.008827Z","shell.execute_reply":"2022-07-28T03:16:46.125303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_callbacks(filepath):\n    \"\"\"Create callbacks for training\"\"\"\n    return [\n        callbacks.ModelCheckpoint(\n            filepath=filepath, monitor=\"val_xentropy\", save_best_only=True, save_weights_only=True, verbose=1\n        ),\n        callbacks.EarlyStopping(\n            patience=cfg.patience, monitor=\"val_xentropy\", restore_best_weights=False, verbose=1\n        ),\n    ]\n\n\ndef show_history(history):\n    \"\"\"Show history\"\"\"\n    history_df = pd.DataFrame(history.history)\n    history_df.index = pd.Index(history.epoch, name=\"epoch\")\n    display(\n        history_df.style.highlight_min(\n            color=\"green\", subset=[\"val_loss\", \"val_xentropy\"]\n        ).highlight_max(color=\"green\", subset=[\"val_acc\"])\n    )\n    fig, ax = plt.subplots(1, 2, figsize=(16, 8))\n    history_df[[\"loss\", \"val_loss\", \"xentropy\", \"val_xentropy\"]].plot(ax=ax[0], title=\"loss/xentropy\")\n    history_df[[\"acc\", \"val_acc\"]].plot(ax=ax[1], title=\"acc\")\n    plt.tight_layout()\n    plt.show()\n    \n    \ndef compute_oof(model, valid):\n    \"\"\"Compute OOF\"\"\"\n    valid_ds = get_dataset(valid)\n    pred = model.predict(valid_ds, verbose=0)\n    oof = pd.DataFrame(pred, columns=cfg.labels, index=valid[cfg.id_col])\n    oof[\"label\"] = valid.set_index(cfg.id_col)[\"label\"]\n    return oof    \n    \n\ndef compute_score(x):\n    \"\"\"Compute score\"\"\"\n    return sk_metrics.log_loss(y_true=x[\"label\"], y_pred=x[cfg.labels])\n\n\ndef create_lr_scheduler(train_ds):\n    \"\"\"Create learning rate scheduler\"\"\"\n    lr_scheduler = optimizers.schedules.PolynomialDecay(\n        initial_learning_rate=cfg.init_learning_rate,\n        decay_steps=(len(train_ds) * cfg.decay_epochs),\n        end_learning_rate=cfg.end_learning_rate\n    )\n    print(\"lr_scheduler.get_config():\", lr_scheduler.get_config())\n    return lr_scheduler","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:16:46.129840Z","iopub.execute_input":"2022-07-28T03:16:46.130184Z","iopub.status.idle":"2022-07-28T03:16:46.145619Z","shell.execute_reply.started":"2022-07-28T03:16:46.130150Z","shell.execute_reply":"2022-07-28T03:16:46.144465Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_training(train, valid, filename):\n    \"\"\"Run training\"\"\"\n    # https://www.kaggle.com/code/cdeotte/tfrecord-experiments-upsample-and-coarse-dropout\n    if tpu:\n        tf.tpu.experimental.initialize_tpu_system()\n    # create datasets\n    train_ds = get_dataset(train, repeat=True, shuffle=True)\n    valid_ds = get_dataset(valid)\n    # create model\n    backend.clear_session()\n    with strategy.scope():\n        model = create_model()\n    # fit\n    hist = model.fit(\n        train_ds,\n        epochs=cfg.epochs,\n        steps_per_epoch=cfg.steps_per_epoch,\n        validation_data=valid_ds,\n        callbacks=create_callbacks(filename),\n        verbose=cfg.verbose,\n    )\n    model.load_weights(filename)\n    # oof\n    oof = compute_oof(model, valid)\n    return hist, oof","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:16:46.146871Z","iopub.execute_input":"2022-07-28T03:16:46.147355Z","iopub.status.idle":"2022-07-28T03:16:46.162289Z","shell.execute_reply.started":"2022-07-28T03:16:46.147325Z","shell.execute_reply":"2022-07-28T03:16:46.161339Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedKFold","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:16:46.163794Z","iopub.execute_input":"2022-07-28T03:16:46.164070Z","iopub.status.idle":"2022-07-28T03:16:46.197997Z","shell.execute_reply.started":"2022-07-28T03:16:46.164032Z","shell.execute_reply":"2022-07-28T03:16:46.197243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"skf = StratifiedKFold(n_splits=cfg.n_fold, shuffle=True, random_state=cfg.seed)\nskf","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:16:46.199191Z","iopub.execute_input":"2022-07-28T03:16:46.199534Z","iopub.status.idle":"2022-07-28T03:16:46.206009Z","shell.execute_reply.started":"2022-07-28T03:16:46.199506Z","shell.execute_reply":"2022-07-28T03:16:46.204846Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"StratifiedKFold(n_splits=5, random_state=887, shuffle=True)\nd_oof = {}\nfor fold, (iloc_train, iloc_valid) in enumerate(skf.split(data, data['label'])):\n    print(f\"fold: {fold}\")\n    train = data.iloc[iloc_train]\n    valid = data.iloc[iloc_valid]\n    model_filepath = f\"weights__{cfg.model_name}__fold-{fold}.h5\"\n    print(f\"#train: {len(train)},  #valid: {len(valid)} \")\n    print(f\"model_filepath: {model_filepath}\")\n    hist, oof = run_training(train, valid, model_filepath)\n    print(\"OOF score:\", compute_score(oof))\n    show_history(hist)\n    d_oof[fold] = oof\n    if cfg.one_fold:\n        break    ","metadata":{"execution":{"iopub.status.busy":"2022-07-28T03:16:46.207142Z","iopub.execute_input":"2022-07-28T03:16:46.207469Z","iopub.status.idle":"2022-07-28T08:01:54.143559Z","shell.execute_reply.started":"2022-07-28T03:16:46.207443Z","shell.execute_reply":"2022-07-28T08:01:54.140554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"oof = pd.concat(d_oof, names=['fold']).reset_index('fold')\noof.to_csv(\"oof.csv\")\nscore_by_fold = oof.groupby('fold').apply(compute_score)\ndisplay(score_by_fold)\nscore = compute_score(oof)\nprint(f\"\\nOOF score: {score:.6f}\")","metadata":{"execution":{"iopub.status.busy":"2022-07-28T08:01:54.148639Z","iopub.execute_input":"2022-07-28T08:01:54.149862Z","iopub.status.idle":"2022-07-28T08:01:54.557451Z","shell.execute_reply.started":"2022-07-28T08:01:54.149821Z","shell.execute_reply":"2022-07-28T08:01:54.556522Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}