{"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":"!pip install ensemble-boxes\n# install WBF from https://github.com/ZFTurbo/Weighted-Boxes-Fusion","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-04-15T04:03:15.630397Z","iopub.execute_input":"2022-04-15T04:03:15.630723Z","iopub.status.idle":"2022-04-15T04:03:29.609292Z","shell.execute_reply.started":"2022-04-15T04:03:15.630633Z","shell.execute_reply":"2022-04-15T04:03:29.608342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import csv\nfrom pathlib import Path\nfrom tqdm import tqdm\nfrom multiprocessing import Pool, Manager\nimport numpy as np\nfrom ensemble_boxes import weighted_boxes_fusion\n","metadata":{"execution":{"iopub.status.busy":"2022-04-15T04:12:18.657208Z","iopub.execute_input":"2022-04-15T04:12:18.658308Z","iopub.status.idle":"2022-04-15T04:12:19.600796Z","shell.execute_reply.started":"2022-04-15T04:12:18.658233Z","shell.execute_reply":"2022-04-15T04:12:19.599767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def filter_size_box_infos(box_infos, min_size=4, image_size=4000):\n    # sort by area\n    box_infos.sort(key=lambda x: (x[3] - x[1]) * (x[4] - x[2]))\n    invalid_idxs = []\n    for i in range(len(box_infos)):\n        cls, x_min, y_min, x_max, y_max, score, _ = box_infos[i]\n        w = int((x_max - x_min) * image_size)\n        h = int((y_max - y_min) * image_size)\n        if max(w, h) < min_size:\n            invalid_idxs.append(i)\n    invalid_idxs = list(set(invalid_idxs))\n    invalid_idxs.sort()\n    for idx in invalid_idxs[::-1]:\n        del box_infos[idx]\n    return box_infos\n\n\ndef preprocess_box_infos(box_infos, expand_ratio=0.1):\n    # expand_box_info:\n    for idx in range(len(box_infos)):\n        cls, x_min, y_min, x_max, y_max, score, model_idx = box_infos[idx]\n        w_ = x_max - x_min\n        h_ = y_max - y_min\n        new_x_min = x_min - expand_ratio * w_\n        new_x_max = x_max + expand_ratio * w_\n        new_y_min = y_min - expand_ratio * h_\n        new_y_max = y_max + expand_ratio * h_\n        box_infos[idx] = [cls, new_x_min, new_y_min, new_x_max, new_y_max, score, model_idx]\n    return box_infos\n\n\ndef filter_overlap(box_infos):\n    def get_overlap_info(box1, box2):\n        cls1, x_min1, y_min1, x_max1, y_max1, score1, _ = box1\n        cls2, x_min2, y_min2, x_max2, y_max2, score2, _ = box2\n        if max(x_min1, x_min2) < min(x_max1, x_max2) and \\\n                max(y_min1, y_min2) < min(y_max1, y_max2):\n            return True, cls1 if score1 > score2 else cls2\n        return False, None\n\n    box_infos = filter_size_box_infos(box_infos, min_size=6)\n    box_infos = preprocess_box_infos(box_infos, expand_ratio=0.075)\n    box_infos.sort(key=lambda x: x[5], reverse=True)  # sort by confidence\n    valid_box_infos = []\n    travel_idxs = []\n    list_valid_number = []\n    for i in range(len(box_infos)):\n        if i in travel_idxs:\n            continue  # skip\n        travel_idxs.append(i)\n        for j in range(i + 1, len(box_infos)):\n            if j in travel_idxs:\n                continue\n            is_overlap__, cls_ = get_overlap_info(box_infos[i], box_infos[j])\n            if is_overlap__:\n                travel_idxs.append(j)\n        list_valid_number.append(int(box_infos[i][0]))\n        valid_box_infos.append(box_infos[i])\n    list_valid_number = list(map(int, list_valid_number))\n\n    invalid_idxs = []\n    for idx, box_info in enumerate(valid_box_infos):\n        cls, x_min, y_min, x_max, y_max, score, model_idx = box_info\n        w_ = x_max - x_min\n        h_ = y_max - y_min\n        if cls == 1 and h_ / w_ < 2.0 and score < 0.8:\n            invalid_idxs.append(idx)\n\n    invalid_idxs = list(set(invalid_idxs))\n    invalid_idxs.sort()\n    for idx in invalid_idxs[::-1]:\n        del valid_box_infos[idx]\n\n    ret = sum([int(info[0]) for info in valid_box_infos])\n\n    if ret > 27:\n        ret = 10\n    return ret, valid_box_infos\n\ndef get_acc_by_file(\n    file_name,\n    list_input_dir,\n    list_weights,\n    list_conf_thresh,\n    wfb=True,\n):\n    data_lines_dict = {'_'.join([str(e) for e in weights]): [] for weights in list_weights}\n    valid_box_infos = {'_'.join([str(e) for e in weights]): {} for weights in list_weights}\n    boxes_list = []\n    scores_list = []\n    labels_list = []\n    dir_list = []\n    iou_thr = 0.5\n    skip_box_thr = 0.0001\n    for idx_dir, input_dir in enumerate(list_input_dir):\n        txt_path = input_dir.joinpath(file_name)\n        if not txt_path.is_file():\n            boxes_list.append([])\n            scores_list.append([])\n            labels_list.append([])\n            continue\n        boxes = []\n        scores = []\n        labels = []\n        dirs = []\n        with open(str(txt_path), 'r') as f:\n            lines = [line.strip() for line in f.readlines()]\n\n        for line in lines:\n            class_idx, x_center, y_center, w, h, score = line.split()\n            x_center, y_center, w, h, score = list(map(float, [x_center, y_center, w, h, score]))\n            class_idx = int(class_idx)\n            if score < list_conf_thresh[idx_dir]:\n                continue\n            boxes.append([x_center - 0.5 * w, y_center - 0.5 * h, x_center + 0.5 * w, y_center + 0.5 * h])\n            scores.append(score)\n            labels.append(class_idx)\n            dirs.append(idx_dir)\n\n        boxes_list.append(boxes)\n        scores_list.append(scores)\n        labels_list.append(labels)\n        dir_list.append(dirs)\n\n    for weights in list_weights:\n        if wfb:\n            boxes, scores, labels = weighted_boxes_fusion(\n                boxes_list, scores_list, labels_list,\n                weights=weights, iou_thr=iou_thr, skip_box_thr=skip_box_thr)\n        else:\n            # flatten\n            boxes = [box for boxes in boxes_list for box in boxes]\n            scores = [score * weight for weight, scores in zip(weights, scores_list) for score in scores]\n            labels = [label for labels in labels_list for label in labels]\n            dirs = [dir_ for dirs in dir_list for dir_ in dirs]\n\n        box_infos = []\n        for box, score, label, dir_ in zip(boxes, scores, labels, dirs):\n            x1, y1, x2, y2 = box\n            box_infos.append([label, x1, y1, x2, y2, score, dir_])\n        digit_sum, final_box_infos = filter_overlap(box_infos)\n        weight_name = '_'.join([str(e) for e in weights])\n        data_lines_dict[weight_name].append([file_name[:-4], digit_sum])\n        valid_box_infos[weight_name][file_name[:-4]] = final_box_infos\n\n    return data_lines_dict, valid_box_infos\n\ndef merge_results(list_input_dir,\n                  save_dir='./submission_wfb',\n                  list_weights=None,\n                  list_conf_thresh=None,\n                  wfb=True,\n                  save_txt_dir=None):\n    if list_conf_thresh is None:\n        list_conf_thresh = [0.0] * len(list_input_dir)\n    if list_weights is None:\n        list_weights = [[1] * len(list_input_dir)]\n\n    # Get all file_name\n    list_file_name = []\n    for input_dir in list_input_dir:\n        list_file_name += [p.name for p in input_dir.glob('*.txt')]\n    list_file_name = list(set(list_file_name))\n\n    if isinstance(save_txt_dir, (str, Path)):\n        save_txt_dir = Path(save_txt_dir)\n\n    header = ['id', 'digit_sum']\n    data_lines_dict = {'_'.join([str(e) for e in weights]): [] for weights in list_weights}\n    valid_box_infos = {'_'.join([str(e) for e in weights]): {} for weights in list_weights}\n\n    with Pool(8) as pool:\n        results = [pool.apply_async(get_acc_by_file, args=(file_name,\n                                                           list_input_dir,\n                                                           list_weights,\n                                                           list_conf_thresh,\n                                                           wfb,))\n                   for file_name in list_file_name]\n        results = [tuple(ret.get()) for ret in tqdm(results)]\n    for sub_data_lines, sub_valid_box_info in results:\n        for k, v in sub_data_lines.items():\n            data_lines_dict[k] += v\n        for weight_name, dict_val in sub_valid_box_info.items():\n            for file_name, val in dict_val.items():\n                valid_box_infos[weight_name][file_name] = val\n\n    save_dir = Path(save_dir)\n    if not save_dir.is_dir():\n        save_dir.mkdir(parents=True)\n\n    # save all csv file\n    for weights in list_weights:\n        data_lines_name = '_'.join([str(e) for e in weights])\n        save_csv_path = save_dir.joinpath('submission_wfb_' + data_lines_name + '.csv')\n        with open(str(save_csv_path), 'w', encoding='UTF8', newline='') as f:\n            writer = csv.writer(f)\n            # write the header\n            writer.writerow(header)\n            # write multiple rows\n            writer.writerows(data_lines_dict[data_lines_name])\n    # save all txt_file\n    if save_txt_dir is not None:\n        for weights in list_weights:\n            data_lines_name = '_'.join([str(e) for e in weights])\n            save_txt_dir = save_txt_dir.joinpath(data_lines_name)\n            if not save_txt_dir.is_dir():\n                save_txt_dir.mkdir(parents=True)\n            for file_name, box_infos in valid_box_infos[data_lines_name].items():\n                save_txt_path = save_txt_dir.joinpath(file_name + '.txt')\n                lines = []\n                for info in box_infos:\n                    cls, x_min, y_min, x_max, y_max, score, _ = info\n                    x_center = (x_min + x_max) / 2\n                    y_center = (y_min + y_max) / 2\n                    w = x_max - x_min\n                    h = y_max - y_min\n                    lines.append(' '.join(list(map(str, [cls, x_center, y_center, w, h, score]))))\n                with open(str(save_txt_path), 'w') as f:\n                    f.write('\\n'.join(lines))\n","metadata":{"execution":{"iopub.status.busy":"2022-04-15T04:12:21.402180Z","iopub.execute_input":"2022-04-15T04:12:21.402499Z","iopub.status.idle":"2022-04-15T04:12:21.456366Z","shell.execute_reply.started":"2022-04-15T04:12:21.402470Z","shell.execute_reply":"2022-04-15T04:12:21.455698Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls ../input/ultra-mnist-asset/txt_predicts/txt_predicts/220414_fold0_e50_1280/220414_fold0_e50_1280 -1 | wc -l\n","metadata":{"execution":{"iopub.status.busy":"2022-04-15T04:13:23.172211Z","iopub.execute_input":"2022-04-15T04:13:23.173106Z","iopub.status.idle":"2022-04-15T04:13:24.275748Z","shell.execute_reply.started":"2022-04-15T04:13:23.173062Z","shell.execute_reply":"2022-04-15T04:13:24.274952Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"merge_results(\n    list_input_dir=[\n        Path('../input/ultra-mnist-asset/txt_predicts/txt_predicts/220327_normal_overlap_split_9split_1024/220327_normal_overlap_split_9split_1024'),\n        Path('../input/ultra-mnist-asset/txt_predicts/txt_predicts/220327_normal_overlap_split_16split_1024/220327_normal_overlap_split_16split_1024'),\n        Path('../input/ultra-mnist-asset/txt_predicts/txt_predicts/220326_60K_tiny_dataset_v2_overlap_split_25split/220326_60K_tiny_dataset_v2_overlap_split_25split'),\n        Path('../input/ultra-mnist-asset/txt_predicts/txt_predicts/220414_fold0_e50_1024/220414_fold0_e50_1024'),\n        Path('../input/ultra-mnist-asset/txt_predicts/txt_predicts/220414_fold0_e50_768/220414_fold0_e50_768'),\n        Path('../input/ultra-mnist-asset/txt_predicts/txt_predicts/220414_fold0_e50_1280/220414_fold0_e50_1280'),\n        Path('../input/ultra-mnist-asset/txt_predicts/txt_predicts/220414_fold0_e50_1536/220414_fold0_e50_1536'),\n    ],\n    save_dir='./220414_7merge',\n    list_weights=[[1, 1, 1, 1, 1, 1, 1]],\n    list_conf_thresh=[0.8, 0.8, 0.7, 0.9, 0.95, 0.9, 0.95],\n    wfb=False\n)\n","metadata":{"execution":{"iopub.status.busy":"2022-04-15T04:14:49.838699Z","iopub.execute_input":"2022-04-15T04:14:49.839000Z","iopub.status.idle":"2022-04-15T04:16:26.741208Z","shell.execute_reply.started":"2022-04-15T04:14:49.838970Z","shell.execute_reply":"2022-04-15T04:16:26.739861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls ./220414_7merge","metadata":{"execution":{"iopub.status.busy":"2022-04-15T04:16:37.393963Z","iopub.execute_input":"2022-04-15T04:16:37.394307Z","iopub.status.idle":"2022-04-15T04:16:38.156206Z","shell.execute_reply.started":"2022-04-15T04:16:37.394271Z","shell.execute_reply":"2022-04-15T04:16:38.155076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}