{"cells":[{"metadata":{"trusted":true},"cell_type":"code","source":"import tensorflow as tf\nimport pandas as pd\nimport cv2\nfrom tqdm import tqdm\n\nBASE_PATH = \"../data\"\nID_COL = \"StudyInstanceUID\"\nSPLITS = 5\n\ntrain = pd.read_csv(f'{BASE_PATH}/train.csv')\ntest = pd.read_csv(f'{BASE_PATH}/sample_submission.csv')\ntrain[\"path\"] = f\"{BASE_PATH}/train/\" + train[ID_COL].astype(str) + \".jpg\"\ntest[\"path\"] = f\"{BASE_PATH}/test/\" + test[ID_COL].astype(str) + \".jpg\"\ntarget_cols = train.drop(columns=[ID_COL] + [\"PatientID\", \"path\"]).columns\n\ntrain[\"fold\"] = pd.cut(train.index, SPLITS, labels=False)\ntest[\"fold\"] = pd.cut(test.index, SPLITS, labels=False)\n\nfilename_train = \"train.tfrecord\"\nfilename_test = \"test.tfrecord\"\n\n\n# 下記の関数を使うと値を tf.Example と互換性の有る型に変換できる\n\ndef _bytes_feature(value):\n    \"\"\"string / byte 型から byte_list を返す\"\"\"\n    if isinstance(value, type(tf.constant(0))):\n        value = value.numpy()  # BytesList won't unpack a string from an EagerTensor.\n    return tf.train.Feature(bytes_list=tf.train.BytesList(value=[value]))\n\n\ndef _float_feature(value):\n    \"\"\"float / double 型から float_list を返す\"\"\"\n    return tf.train.Feature(float_list=tf.train.FloatList(value=[value]))\n\n\ndef _int64_feature(value):\n    \"\"\"bool / enum / int / uint 型から Int64_list を返す\"\"\"\n    return tf.train.Feature(int64_list=tf.train.Int64List(value=[value]))\n\n\ndef _int64_list_feature(value):\n    \"\"\"List[bool / enum / int / uint] 型から Int64_list を返す\"\"\"\n    return tf.train.Feature(int64_list=tf.train.Int64List(value=value))\n\n\ndef serialize_train(image, target, image_name):\n    feature = {\n        'image': _bytes_feature(image),\n        \"target\": _int64_list_feature(target.tolist()),\n        'image_name': _bytes_feature(image_name)\n    }\n    example_proto = tf.train.Example(features=tf.train.Features(feature=feature))\n    return example_proto.SerializeToString()\n\n\ndef serialize_test(image, image_name):\n    feature = {\n        'image': _bytes_feature(image),\n        'image_name': _bytes_feature(image_name)\n    }\n    example_proto = tf.train.Example(features=tf.train.Features(feature=feature))\n    return example_proto.SerializeToString()\n\n\ndef save_tfrecord(filename, data, train=True):\n    with tf.io.TFRecordWriter(filename) as writer:\n        for i, row in tqdm(data.iterrows()):\n            if train:\n                label = row[target_cols]\n            path = row[\"path\"]\n\n            img = cv2.imread(path)\n            img = cv2.imencode('.jpg', img, (cv2.IMWRITE_JPEG_QUALITY, 100))[1].tobytes()\n            if train:\n                example = serialize_train(img, label, str.encode(row[ID_COL]))\n            else:\n                example = serialize_test(img, str.encode(row[ID_COL]))\n            writer.write(example)\n\n\nfor i in range(SPLITS):\n    filename = f\"train_{i}.tfrecord\"\n    use_data = train[train[\"fold\"] == i]\n    save_tfrecord(filename, use_data, train=True)\n\nfor i in range(SPLITS):\n    filename = f\"test_{i}.tfrecord\"\n    use_data = test[test[\"fold\"] == i]\n    save_tfrecord(filename, use_data, train=False)\n","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}