{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.6.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":13451,"datasetId":654585,"databundleVersionId":1188070},{"sourceType":"datasetVersion","sourceId":10288116,"datasetId":6366959,"databundleVersionId":10589853},{"sourceType":"datasetVersion","sourceId":10288130,"datasetId":6366971,"databundleVersionId":10589867},{"sourceType":"datasetVersion","sourceId":707009,"datasetId":361580,"databundleVersionId":727143},{"sourceType":"datasetVersion","sourceId":10288123,"datasetId":6366965,"databundleVersionId":10589860},{"sourceType":"datasetVersion","sourceId":10288122,"datasetId":6366964,"databundleVersionId":10589859},{"sourceType":"datasetVersion","sourceId":10218414,"datasetId":6316200,"databundleVersionId":10511258}],"dockerImageVersionId":29662,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install bert4keras","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:51:16.59063Z","iopub.execute_input":"2024-12-25T00:51:16.591056Z","iopub.status.idle":"2024-12-25T00:51:23.71056Z","shell.execute_reply.started":"2024-12-25T00:51:16.590983Z","shell.execute_reply":"2024-12-25T00:51:23.709743Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport pydicom\nimport os\nimport matplotlib.pyplot as plt\nimport collections\nfrom tqdm import tqdm_notebook as tqdm\nfrom datetime import datetime\n\nfrom math import ceil, floor, log\nimport cv2\n\nimport tensorflow as tf\nimport keras\n\nimport sys\n\n# from keras_applications.resnet import ResNet50\nfrom keras_applications.inception_v3 import InceptionV3\n\nfrom sklearn.model_selection import ShuffleSplit\n\nimport numpy as np\nfrom bert4keras.backend import keras, set_gelu, K\nfrom bert4keras.tokenizers import Tokenizer\nfrom bert4keras.models import build_transformer_model\nfrom bert4keras.optimizers import Adam\nfrom bert4keras.snippets import sequence_padding, DataGenerator\nfrom bert4keras.snippets import open\nfrom keras.layers import Dropout, Dense","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:51:23.712106Z","iopub.execute_input":"2024-12-25T00:51:23.712334Z","iopub.status.idle":"2024-12-25T00:51:26.527406Z","shell.execute_reply.started":"2024-12-25T00:51:23.712294Z","shell.execute_reply":"2024-12-25T00:51:26.526656Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(tf.__version__)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:51:26.529176Z","iopub.execute_input":"2024-12-25T00:51:26.529506Z","iopub.status.idle":"2024-12-25T00:51:26.533763Z","shell.execute_reply.started":"2024-12-25T00:51:26.529447Z","shell.execute_reply":"2024-12-25T00:51:26.533149Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"set_gelu('tanh')  # 切换gelu版本\n\nmaxlen = 128\nbatch_size = 32  # 64\nconfig_path = '/kaggle/input/chinese-bert/chinese_L-12_H-768_A-12/bert_config.json'\ncheckpoint_path = '/kaggle/input/chinese-bert/chinese_L-12_H-768_A-12/bert_model.ckpt'\ndict_path = '/kaggle/input/chinese-bert/chinese_L-12_H-768_A-12/vocab.txt'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:51:26.535025Z","iopub.execute_input":"2024-12-25T00:51:26.535218Z","iopub.status.idle":"2024-12-25T00:51:26.545827Z","shell.execute_reply.started":"2024-12-25T00:51:26.535183Z","shell.execute_reply":"2024-12-25T00:51:26.545085Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import pandas as pd\n\n# # 读取 CSV 文件\n# ATEC_TRAIN = '/kaggle/input/atec-cn/ATEC.train.data'\n# ATEC_VALID = '/kaggle/input/atec-cn/ATEC.valid.data'\n# ATEC_TEST = '/kaggle/input/atec-cn/ATEC.test.data'\n# atec_train = pd.read_csv(ATEC_TRAIN)\n# atec_valid = pd.read_csv(ATEC_VALID)\n# atec_test = pd.read_csv(ATEC_TEST)\n\n# BQ_TRAIN = '/kaggle/input/bq-cn/BQ.train.data'\n# BQ_VALID = '/kaggle/input/bq-cn/BQ.valid.data'\n# BQ_TEST  = '/kaggle/input/bq-cn/BQ.test.data'\n# bq_train = pd.read_csv(BQ_TRAIN)\n# bq_valid = pd.read_csv(BQ_VALID)\n# bq_test = pd.read_csv(BQ_TEST)\n\n# LCQMC_TRAIN = '/kaggle/input/lcqmc-cn/LCQMC.train.data'\n# LCQMC_VALID = '/kaggle/input/lcqmc-cn/LCQMC.valid.data'\n# LCQMC_TEST  = '/kaggle/input/lcqmc-cn/LCQMC.test.data'\n# lcqmc_train = pd.read_csv(LCQMC_TRAIN)\n# lcqmc_valid = pd.read_csv(LCQMC_VALID)\n# lcqmc_test = pd.read_csv(LCQMC_TEST)\n\n# STS_B_TRAIN = '/kaggle/input/sts-b-cn/STS-B.train.data'\n# STS_B_VALID = '/kaggle/input/sts-b-cn/STS-B.valid.data'\n# STS_B_TEST  = '/kaggle/input/sts-b-cn/STS-B.test.data'\n# sts_b_train = pd.read_csv(STS_B_TRAIN)\n# sts_b_valid = pd.read_csv(STS_B_VALID)\n# sts_b_test = pd.read_csv(STS_B_TEST)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:51:26.548266Z","iopub.execute_input":"2024-12-25T00:51:26.548489Z","iopub.status.idle":"2024-12-25T00:51:26.557782Z","shell.execute_reply.started":"2024-12-25T00:51:26.548445Z","shell.execute_reply":"2024-12-25T00:51:26.55714Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport gc\n\nLCQMC_TRAIN = '/kaggle/input/text-matching-baidu/lcqmc_train.csv'\nLCQMC_VALID = '/kaggle/input/text-matching-baidu/lcqmc_valid.csv'\nLCQMC_TEST  = '/kaggle/input/text-matching-baidu/lcqmc_test.csv'\nlcqmc_train = pd.read_csv(LCQMC_TRAIN)\nlcqmc_valid = pd.read_csv(LCQMC_VALID)\nlcqmc_test = pd.read_csv(LCQMC_TEST)\n\n# 适配 load_data 函数需要的格式\ndef convert_to_data_format(df):\n    \"\"\"将 DataFrame 转换为 (文本1, 文本2, 标签id) 的格式\"\"\"\n    data = []\n    for _, row in df.iterrows():\n        text1 = row['query1']\n        text2 = row['query2']\n        label = int(row['label'])\n        data.append((text1, text2, label))\n    return data\n\n# 将 DataFrame 转换为 (文本1, 文本2, 标签id) 格式\ntrain_data = convert_to_data_format(lcqmc_train)\nvalid_data = convert_to_data_format(lcqmc_valid)\ntest_data = convert_to_data_format(lcqmc_test)\n\n# 打印一下转化后的数据，检查格式\nprint(train_data[:3])  # 打印前3个数据\n\ndel lcqmc_train\ndel lcqmc_valid\ndel lcqmc_test\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:51:26.560103Z","iopub.execute_input":"2024-12-25T00:51:26.560389Z","iopub.status.idle":"2024-12-25T00:51:55.192284Z","shell.execute_reply.started":"2024-12-25T00:51:26.560337Z","shell.execute_reply":"2024-12-25T00:51:55.191571Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(len(train_data))\nprint(len(valid_data))\nprint(len(test_data))\ntrain_data = train_data[:]  # 这里控制读入数据量\nvalid_data = valid_data[:]\ntest_data = test_data[:]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:51:55.193537Z","iopub.execute_input":"2024-12-25T00:51:55.193756Z","iopub.status.idle":"2024-12-25T00:51:55.201768Z","shell.execute_reply.started":"2024-12-25T00:51:55.193719Z","shell.execute_reply":"2024-12-25T00:51:55.20102Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 建立分词器\ntokenizer = Tokenizer(dict_path, do_lower_case=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:51:55.202789Z","iopub.execute_input":"2024-12-25T00:51:55.203038Z","iopub.status.idle":"2024-12-25T00:51:55.246507Z","shell.execute_reply.started":"2024-12-25T00:51:55.202978Z","shell.execute_reply":"2024-12-25T00:51:55.246025Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class data_generator(DataGenerator):\n    \"\"\"数据生成器\n    \"\"\"\n    def __iter__(self, random=False):\n        batch_token_ids, batch_segment_ids, batch_labels = [], [], []\n        for is_end, (text1, text2, label) in self.sample(random):\n            token_ids, segment_ids = tokenizer.encode(\n                text1, text2, maxlen=maxlen\n            )\n            batch_token_ids.append(token_ids)\n            batch_segment_ids.append(segment_ids)\n            batch_labels.append([label])\n            if len(batch_token_ids) == self.batch_size or is_end:\n                batch_token_ids = sequence_padding(batch_token_ids)\n                batch_segment_ids = sequence_padding(batch_segment_ids)\n                batch_labels = sequence_padding(batch_labels)\n                yield [batch_token_ids, batch_segment_ids], batch_labels\n                batch_token_ids, batch_segment_ids, batch_labels = [], [], []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:51:55.247712Z","iopub.execute_input":"2024-12-25T00:51:55.248025Z","iopub.status.idle":"2024-12-25T00:51:55.25483Z","shell.execute_reply.started":"2024-12-25T00:51:55.247971Z","shell.execute_reply":"2024-12-25T00:51:55.254073Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 转换数据集\ntrain_generator = data_generator(train_data, batch_size)\nvalid_generator = data_generator(valid_data, batch_size)\ntest_generator = data_generator(test_data, batch_size)\n\ndel train_data\ndel valid_data\ndel test_data\ngc.collect()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:51:55.256033Z","iopub.execute_input":"2024-12-25T00:51:55.256313Z","iopub.status.idle":"2024-12-25T00:51:55.361139Z","shell.execute_reply.started":"2024-12-25T00:51:55.256261Z","shell.execute_reply":"2024-12-25T00:51:55.360411Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 加载预训练模型\nbert = build_transformer_model(\n    config_path=config_path,\n    checkpoint_path=checkpoint_path,\n    with_pool=True,\n    return_keras_model=False,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:51:55.3623Z","iopub.execute_input":"2024-12-25T00:51:55.362536Z","iopub.status.idle":"2024-12-25T00:52:07.23547Z","shell.execute_reply.started":"2024-12-25T00:51:55.362473Z","shell.execute_reply":"2024-12-25T00:52:07.234631Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"output = Dropout(rate=0.1)(bert.model.output)\noutput = Dense(\n    units=2, activation='softmax', kernel_initializer=bert.initializer\n)(output)\n\nmodel = keras.models.Model(bert.model.input, output)\nmodel.summary()\n\nmodel.compile(\n    loss='sparse_categorical_crossentropy',\n    optimizer=Adam(2e-5),  # 用足够小的学习率 2e-5\n    # optimizer=PiecewiseLinearLearningRate(Adam(5e-5), {10000: 1, 30000: 0.1}),\n    metrics=['accuracy'],\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:52:07.236874Z","iopub.execute_input":"2024-12-25T00:52:07.237198Z","iopub.status.idle":"2024-12-25T00:52:07.387773Z","shell.execute_reply.started":"2024-12-25T00:52:07.237144Z","shell.execute_reply":"2024-12-25T00:52:07.387134Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate(data):\n    total, right = 0., 0.\n    for x_true, y_true in data:\n        y_pred = model.predict(x_true).argmax(axis=1)\n        y_true = y_true[:, 0]\n        total += len(y_true)\n        right += (y_true == y_pred).sum()\n    return right / total","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:52:07.388696Z","iopub.execute_input":"2024-12-25T00:52:07.388945Z","iopub.status.idle":"2024-12-25T00:52:07.393364Z","shell.execute_reply.started":"2024-12-25T00:52:07.388885Z","shell.execute_reply":"2024-12-25T00:52:07.392519Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Evaluator(keras.callbacks.Callback):\n    \"\"\"评估与保存\n    \"\"\"\n    def __init__(self):\n        self.best_val_acc = 0.\n\n    def on_epoch_end(self, epoch, logs=None):\n        val_acc = evaluate(valid_generator)\n        if val_acc > self.best_val_acc:\n            self.best_val_acc = val_acc\n            model.save_weights('best_model.weights')\n        test_acc = evaluate(test_generator)\n        print(\n            u'val_acc: %.5f, best_val_acc: %.5f, test_acc: %.5f\\n' %\n            (val_acc, self.best_val_acc, test_acc)\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:52:07.394399Z","iopub.execute_input":"2024-12-25T00:52:07.394594Z","iopub.status.idle":"2024-12-25T00:52:07.409283Z","shell.execute_reply.started":"2024-12-25T00:52:07.394555Z","shell.execute_reply":"2024-12-25T00:52:07.408719Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!rm -rf /kaggle/working/best_model.weights","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:52:07.41039Z","iopub.execute_input":"2024-12-25T00:52:07.410688Z","iopub.status.idle":"2024-12-25T00:52:08.459201Z","shell.execute_reply.started":"2024-12-25T00:52:07.410631Z","shell.execute_reply":"2024-12-25T00:52:08.457681Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"evaluator = Evaluator()\n\nmodel.fit(\n    train_generator.forfit(),\n    steps_per_epoch=len(train_generator),\n    epochs=1,  # 20\n    callbacks=[evaluator]\n)\n\nmodel.load_weights('best_model.weights')\nprint(u'final test acc: %05f\\n' % (evaluate(test_generator)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T00:52:08.460709Z","iopub.execute_input":"2024-12-25T00:52:08.460964Z","iopub.status.idle":"2024-12-25T01:22:04.773801Z","shell.execute_reply.started":"2024-12-25T00:52:08.460899Z","shell.execute_reply":"2024-12-25T01:22:04.77297Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Predict","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport json\n\n# 定义文件路径\nfile_path = \"/kaggle/input/text-matching-baidu/similarity_ch/similarity_ch/similarity_ch\"  # 将其替换为实际文件路径\n\n# 创建一个空的列表来存储每一行的JSON数据\ndata = []\n\n# 打开文件逐行读取\nwith open(file_path, \"r\", encoding=\"utf-8\") as file:\n    for line in file:\n        # 使用 json.loads 解析每一行的JSON数据\n        try:\n            data.append(json.loads(line.strip()))\n        except json.JSONDecodeError:\n            print(f\"JSONDecodeError: {line}\")\n\n# 转换为 DataFrame\nsubmit = pd.DataFrame(data)\n\n# 显示前几行\nprint(submit.head())\n\n# 检查列类型\nprint(submit.info())\n\n# 将 DataFrame 转换为 (文本1, 文本2, 标签id) 格式\nsubmit_data = submit[['query','title']]\nsubmit_data.columns = ['query1','query2']\nsubmit_data['label'] = 0\n\n# 保存为 CSV 文件\n# submit_data.to_csv('submit_data.csv', index=False)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T01:22:04.775389Z","iopub.execute_input":"2024-12-25T01:22:04.775711Z","iopub.status.idle":"2024-12-25T01:22:04.864208Z","shell.execute_reply.started":"2024-12-25T01:22:04.775644Z","shell.execute_reply":"2024-12-25T01:22:04.863472Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 使用之前创建的 tokenizer 对数据进行编码\ndef encode_data(data):\n    \"\"\"对输入数据进行编码\"\"\"\n    all_token_ids, all_segment_ids = [], []\n    for text1, text2 in zip(data['query1'], data['query2']):\n        token_ids, segment_ids = tokenizer.encode(text1, text2, maxlen=maxlen)\n        all_token_ids.append(token_ids)\n        all_segment_ids.append(segment_ids)\n\n    # 使用 sequence_padding 函数来填充序列至相同长度\n    all_token_ids = sequence_padding(all_token_ids)\n    all_segment_ids = sequence_padding(all_segment_ids)\n    return [all_token_ids, all_segment_ids]\n\n# 编码 submit_data 数据\nencoded_data = encode_data(submit_data)\nencoded_data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T01:22:04.865722Z","iopub.execute_input":"2024-12-25T01:22:04.86605Z","iopub.status.idle":"2024-12-25T01:22:05.461675Z","shell.execute_reply.started":"2024-12-25T01:22:04.865992Z","shell.execute_reply":"2024-12-25T01:22:05.460982Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 加载最佳模型权重\n# model.load_weights('best_model.weights')\n\n# 进行预测\npredictions = model.predict(encoded_data)\n\n# 获取预测的类别标签\npredicted_labels = predictions.argmax(axis=1)\n\npredicted_labels","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T01:22:05.462945Z","iopub.execute_input":"2024-12-25T01:22:05.463199Z","iopub.status.idle":"2024-12-25T01:22:14.817833Z","shell.execute_reply.started":"2024-12-25T01:22:05.46315Z","shell.execute_reply":"2024-12-25T01:22:14.817023Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport zipfile\n\n# 添加预测结果到 submit\nsubmit['label'] = predicted_labels\n\n# 模拟生成预测结果\n# 假设 label 为 1，证据字段为分词后的索引列表（这里为简单示例）\n# submit[\"label\"] = 1  # 全部设为 1\nsubmit[\"text_q_ids\"] = submit[\"text_q_seg\"].apply(lambda x: \",\".join(map(str, range(len(x)))))\nsubmit[\"text_t_ids\"] = submit[\"text_t_seg\"].apply(lambda x: \",\".join(map(str, range(len(x)))))\n\n# 格式化提交文件内容\nsubmit_lines = submit.apply(\n    lambda row: f\"{row['id']}\\t{row['label']}\\t{row['text_q_ids']}\\t{row['text_t_ids']}\", axis=1\n).tolist()\n\n# 写入文件 sim_rationale.txt\nsubmit_file = \"sim_rationale.txt\"\nwith open(submit_file, \"w\", encoding=\"utf-8\") as f:\n    f.write(\"\\n\".join(submit_lines))\n\n# 压缩为 sim_rationale.txt.zip\nzip_file = \"sim_rationale.txt.zip\"\nwith zipfile.ZipFile(zip_file, \"w\", zipfile.ZIP_DEFLATED) as zipf:\n    zipf.write(submit_file)\n\nprint(f\"文件已生成并压缩为：{zip_file}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T01:22:14.819294Z","iopub.execute_input":"2024-12-25T01:22:14.819607Z","iopub.status.idle":"2024-12-25T01:22:14.998997Z","shell.execute_reply.started":"2024-12-25T01:22:14.819549Z","shell.execute_reply":"2024-12-25T01:22:14.998193Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\n\n# 假设这是你的预测标签数组\npredictions = predicted_labels\n\n# 统计正样本（1）的数量\npositive_count = np.sum(predictions == 1)\n\n# 统计负样本（0）的数量\nnegative_count = np.sum(predictions == 0)\n\n# 计算总样本量\ntotal_count = len(predictions)\n\n# 计算正样本比例\npositive_ratio = positive_count / total_count\n\n# 计算负样本比例\nnegative_ratio = negative_count / total_count\n\nprint(f\"正样本数量: {positive_count}, 负样本数量: {negative_count}\")\nprint(f\"正样本比例: {positive_ratio:.2f}, 负样本比例: {negative_ratio:.2f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-25T01:22:15.000103Z","iopub.execute_input":"2024-12-25T01:22:15.000305Z","iopub.status.idle":"2024-12-25T01:22:15.005686Z","shell.execute_reply.started":"2024-12-25T01:22:15.000269Z","shell.execute_reply":"2024-12-25T01:22:15.00497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}