{"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":"markdown","source":"# DFL Event Detection Average Precision #","metadata":{}},{"cell_type":"markdown","source":"[Objective/目的]  \n- Review the contents of the function in your notebook to understand the competition criteria.  \n- コンペティションの判断基準の理解のため、ノートブックで関数の中身を確認します。  \n\n[Competition Metric - DFL Event Detection AP]  \nhttps://www.kaggle.com/code/ryanholbrook/competition-metric-dfl-event-detection-ap","metadata":{}},{"cell_type":"markdown","source":"<h1 align='left'>Table of Contents 📜</h1>\n<ul style=\"list-style-type:square\">\n    <li><a href=\"#1\">Competition Metric - DFL Event Detection AP </a></li>\n    <li><a href=\"#2\">関数の中身理解 / Understanding the content of functions  </a></li>","metadata":{}},{"cell_type":"markdown","source":"<a id='#1'></a>\n<h1 align='left'>Competition Metric - DFL Event Detection AP</h1>\n<ul style=\"list-style-type:square\">","metadata":{}},{"cell_type":"code","source":"#######################\n# Import libraries\n#######################\nimport numpy as np\nimport pandas as pd\nfrom pandas.testing import assert_index_equal\nfrom typing import Dict, Tuple","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:08:54.765979Z","iopub.execute_input":"2022-07-31T03:08:54.767019Z","iopub.status.idle":"2022-07-31T03:08:54.774356Z","shell.execute_reply.started":"2022-07-31T03:08:54.766976Z","shell.execute_reply":"2022-07-31T03:08:54.773178Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# match_detections関数の中で使うbest_errorの初期値設定に用いる辞書\ntolerances = {\n    \"challenge\": [0.3, 0.4, 0.5, 0.6, 0.7],\n    \"play\": [0.15, 0.20, 0.25, 0.30, 0.35],\n    \"throwin\": [0.15, 0.20, 0.25, 0.30, 0.35],\n}","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:08:54.775616Z","iopub.execute_input":"2022-07-31T03:08:54.776444Z","iopub.status.idle":"2022-07-31T03:08:54.786520Z","shell.execute_reply.started":"2022-07-31T03:08:54.776407Z","shell.execute_reply":"2022-07-31T03:08:54.785397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def filter_detections(\n    detections: pd.DataFrame, \n    intervals: pd.DataFrame\n) -> pd.DataFrame:\n    \"\"\"Drop detections not inside a scoring interval.\n       start-endの間にdetection_timeが無いものをドロップする関数    \n    \"\"\"\n    detection_time = detections.loc[:, 'time'].sort_values().to_numpy()\n    intervals = intervals.to_numpy()\n    # np.full_like: detection_timeと同じshapeでFalseを埋める\n    is_scored = np.full_like(detection_time, False, dtype=bool)\n\n    i, j = 0, 0\n    # while文の中では、detection_timeが startとendの間にあると、False->Trueに変換\n    while i < len(detection_time) and j < len(intervals):\n        time = detection_time[i]\n        int_ = intervals[j]\n\n        # If the detection is prior in time to the interval, go to the next detection.\n        # int_.leftはstartの時刻を示す。それよりもtimeが小さいならば、iをインクリメント\n        if time < int_.left:\n            i += 1\n        # If the detection is inside the interval, keep it and go to the next detection.    \n        # int_にtimeが含まれていればis_scored[i]をFalse->Trueに変える\n        elif time in int_:\n            is_scored[i] = True\n            i += 1\n        # If the detection is later in time, go to the next interval.\n        # time > int_.rightならばjをインクリメント\n        else:\n            j += 1\n\n    return detections.loc[is_scored].reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:08:54.788068Z","iopub.execute_input":"2022-07-31T03:08:54.788953Z","iopub.status.idle":"2022-07-31T03:08:54.800464Z","shell.execute_reply.started":"2022-07-31T03:08:54.788914Z","shell.execute_reply":"2022-07-31T03:08:54.799169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def match_detections(\n    tolerance: float, \n    ground_truths: pd.DataFrame, \n    detections: pd.DataFrame\n) -> pd.DataFrame:\n    \"\"\"Match detections to ground truth events. \n    Arguments are taken from a common event x tolerance x video evaluation group.\n    detectionsとground truth(正解)を照合する。\n    eventとtolerance毎に照合を行う。\n    \"\"\"\n    # Scoreでソート\n    detections_sorted = detections.sort_values('score', ascending=False).dropna()\n    # is_matchedというdetections_sorted['event']と同じ行数のFalseだけのarrayを用意\n    is_matched = np.full_like(detections_sorted['event'], False, dtype=bool)\n    \n    # 繰り返し処理\n    gts_matched = set()\n    for i, det in enumerate(detections_sorted.itertuples(index=False)):\n        best_error = tolerance # best_errorの初期値設定\n        best_gt = None\n\n        for gt in ground_truths.itertuples(index=False):\n            error = abs(det.time - gt.time) # timeのズレ分をerrorとする\n            # errorがbest かつ　gtが新規ならば更新\n            if error < best_error and not gt in gts_matched:\n                best_gt = gt\n                best_error = error\n            \n        if best_gt is not None:\n            is_matched[i] = True\n            gts_matched.add(best_gt)\n\n    # is_matchedを加える\n    detections_sorted['matched'] = is_matched\n\n    return detections_sorted","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:08:54.802295Z","iopub.execute_input":"2022-07-31T03:08:54.802789Z","iopub.status.idle":"2022-07-31T03:08:54.818315Z","shell.execute_reply.started":"2022-07-31T03:08:54.802739Z","shell.execute_reply":"2022-07-31T03:08:54.817160Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def precision_recall_curve(\n    matches: np.ndarray, \n    scores: np.ndarray, \n    p: int\n) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:\n    \"\"\"matches. scores, pからprecision、recallを求める\n    \"\"\"\n    if len(matches) == 0:\n        return [1], [0], []\n\n    # Sort matches by decreasing confidence\n    idxs = np.argsort(scores, kind='stable')[::-1]\n    scores = scores[idxs]\n    matches = matches[idxs]\n    \n    distinct_value_indices = np.where(np.diff(scores))[0]\n    threshold_idxs = np.r_[distinct_value_indices, matches.size - 1] # np.r_でdistinct_value_indicesとmatches.size - 1を結合\n    thresholds = scores[threshold_idxs]\n    \n    # Matches become TPs and non-matches FPs as confidence threshold decreases\n    tps = np.cumsum(matches)[threshold_idxs]  # Trueの総和\n    fps = np.cumsum(~matches)[threshold_idxs] # Falseの総和\n    \n    precision = tps / (tps + fps)\n    precision[np.isnan(precision)] = 0\n    recall = tps / p  # total number of ground truths might be different than total number of matches\n    \n    # Stop when full recall attained and reverse the outputs so recall is non-increasing.\n    last_ind = tps.searchsorted(tps[-1])\n    sl = slice(last_ind, None, -1)\n\n    # Final precision is 1 and final recall is 0\n    return np.r_[precision[sl], 1], np.r_[recall[sl], 0], thresholds[sl]","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:08:54.819908Z","iopub.execute_input":"2022-07-31T03:08:54.820397Z","iopub.status.idle":"2022-07-31T03:08:54.835988Z","shell.execute_reply.started":"2022-07-31T03:08:54.820352Z","shell.execute_reply":"2022-07-31T03:08:54.834987Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def average_precision_score(matches: np.ndarray, scores: np.ndarray, p: int) -> float:\n    precision, recall, _ = precision_recall_curve(matches, scores, p)\n    # Compute step integral\n    return -np.sum(np.diff(recall) * np.array(precision)[:-1])","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:08:54.837499Z","iopub.execute_input":"2022-07-31T03:08:54.838117Z","iopub.status.idle":"2022-07-31T03:08:54.852934Z","shell.execute_reply.started":"2022-07-31T03:08:54.838059Z","shell.execute_reply":"2022-07-31T03:08:54.851621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def event_detection_ap(\n        solution: pd.DataFrame,\n        submission: pd.DataFrame,\n        tolerances: Dict[str, float],\n) -> float:\n    \"\"\"メインのKPI\n    \"\"\"\n\n    assert_index_equal(solution.columns, pd.Index(['video_id', 'time', 'event']))\n    assert_index_equal(submission.columns, pd.Index(['video_id', 'time', 'event', 'score']))\n\n    # Extract scoring intervals.\n    intervals = (\n        solution\n        .query(\"event in ['start', 'end']\")\n        .assign(interval=lambda x: x.groupby(['video_id', 'event']).cumcount())\n        .pivot(index='interval', columns=['video_id', 'event'], values='time')\n        .stack('video_id')\n        .swaplevel()\n        .sort_index()\n        .loc[:, ['start', 'end']]\n        .apply(lambda x: pd.Interval(*x, closed='both'), axis=1)\n    )\n\n    # Extract ground-truth events.\n    ground_truths = (\n        solution\n        .query(\"event not in ['start', 'end']\")\n        .reset_index(drop=True)\n    )\n\n    # Map each event class to its prevalence (needed for recall calculation)\n    class_counts = ground_truths.value_counts('event').to_dict()\n\n    # Create table for detections with a column indicating a match to a ground-truth event\n    detections = submission.assign(matched = False)\n\n    # Remove detections outside of scoring intervals\n    detections_filtered = []\n    for (det_group, dets), (int_group, ints) in zip(\n        detections.groupby('video_id'), intervals.groupby('video_id')\n    ):\n        assert det_group == int_group\n        detections_filtered.append(filter_detections(dets, ints))\n    detections_filtered = pd.concat(detections_filtered, ignore_index=True)\n\n    # Create table of event-class x tolerance x video_id values\n    aggregation_keys = pd.DataFrame(\n        [(ev, tol, vid)\n         for ev in tolerances.keys()\n         for tol in tolerances[ev]\n         for vid in ground_truths['video_id'].unique()],\n        columns=['event', 'tolerance', 'video_id'],\n    )\n\n    # Create match evaluation groups: event-class x tolerance x video_id\n    detections_grouped = (\n        aggregation_keys\n        .merge(detections_filtered, on=['event', 'video_id'], how='left')\n        .groupby(['event', 'tolerance', 'video_id'])\n    )\n    ground_truths_grouped = (\n        aggregation_keys\n        .merge(ground_truths, on=['event', 'video_id'], how='left')\n        .groupby(['event', 'tolerance', 'video_id'])\n    )\n    \n    # Match detections to ground truth events by evaluation group\n    detections_matched = []\n    for key in aggregation_keys.itertuples(index=False):\n        dets = detections_grouped.get_group(key)\n        gts = ground_truths_grouped.get_group(key)\n        detections_matched.append(\n            match_detections(dets['tolerance'].iloc[0], gts, dets) # ここでmatch_detections関数を実行\n        )\n    detections_matched = pd.concat(detections_matched)\n    \n    # Compute AP per event x tolerance group\n    event_classes = ground_truths['event'].unique()\n    ap_table = (\n        detections_matched\n        .query(\"event in @event_classes\")\n        .groupby(['event', 'tolerance']).apply(\n        lambda group: average_precision_score(   # ここでaverage_precision_score関数を実行\n        group['matched'].to_numpy(),\n                group['score'].to_numpy(),\n                class_counts[group['event'].iat[0]],\n            )\n        )\n    )\n\n    # Average over tolerances, then over event classes\n    mean_ap = ap_table.groupby('event').mean().mean()\n\n    return mean_ap","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:08:54.854612Z","iopub.execute_input":"2022-07-31T03:08:54.855222Z","iopub.status.idle":"2022-07-31T03:08:54.874553Z","shell.execute_reply.started":"2022-07-31T03:08:54.855176Z","shell.execute_reply":"2022-07-31T03:08:54.873304Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Example #","metadata":{}},{"cell_type":"markdown","source":"Let's walk through a few examples, using the training set labels as a stand-in for the ground-truth labels.","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\n\ndata_dir = Path('../input/dfl-bundesliga-data-shootout')","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:08:54.875848Z","iopub.execute_input":"2022-07-31T03:08:54.876821Z","iopub.status.idle":"2022-07-31T03:08:54.890816Z","shell.execute_reply.started":"2022-07-31T03:08:54.876781Z","shell.execute_reply":"2022-07-31T03:08:54.889508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"solution = pd.read_csv(data_dir / 'train.csv', usecols=['video_id', 'time', 'event'])\nsolution.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:08:54.892499Z","iopub.execute_input":"2022-07-31T03:08:54.892860Z","iopub.status.idle":"2022-07-31T03:08:54.926700Z","shell.execute_reply.started":"2022-07-31T03:08:54.892829Z","shell.execute_reply":"2022-07-31T03:08:54.925433Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The submission should be the `video_id`, `time`, and `event` columns, dropping the `start` and `end` events which delimit the scoring intervals.","metadata":{}},{"cell_type":"code","source":"perfect_submission = (\n    solution\n    .query(\"event not in ['start', 'end']\")\n    .reset_index(drop=True)\n    .assign(score = 1.0)\n)\nperfect_submission.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:08:54.928246Z","iopub.execute_input":"2022-07-31T03:08:54.928681Z","iopub.status.idle":"2022-07-31T03:08:54.950026Z","shell.execute_reply.started":"2022-07-31T03:08:54.928636Z","shell.execute_reply":"2022-07-31T03:08:54.949226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Running this through our metric, we see indeed that this results in a perfect score.","metadata":{}},{"cell_type":"code","source":"event_detection_ap(solution, perfect_submission, tolerances)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:08:54.951340Z","iopub.execute_input":"2022-07-31T03:08:54.951823Z","iopub.status.idle":"2022-07-31T03:09:11.918670Z","shell.execute_reply.started":"2022-07-31T03:08:54.951791Z","shell.execute_reply":"2022-07-31T03:09:11.917418Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now let's see what happens if we shuffle about 10% of the event labels.","metadata":{}},{"cell_type":"code","source":"noisy_submission = perfect_submission.copy()\nidx = noisy_submission.sample(frac=0.1).index\nnoisy_submission.loc[idx, 'event'] = noisy_submission.loc[idx, 'event'].sort_index().to_numpy()\n\nevent_detection_ap(solution, noisy_submission, tolerances)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:11.921382Z","iopub.execute_input":"2022-07-31T03:09:11.922112Z","iopub.status.idle":"2022-07-31T03:09:28.992247Z","shell.execute_reply.started":"2022-07-31T03:09:11.922045Z","shell.execute_reply":"2022-07-31T03:09:28.991038Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"And finally we'll try shuffling the timestamps.","metadata":{}},{"cell_type":"code","source":"noisy_submission_2 = perfect_submission.copy()\ntime_noise = np.random.normal(loc=0.0, scale=0.15, size=noisy_submission.shape[0])\nnoisy_submission_2['time'] = noisy_submission_2['time'] + time_noise\n\nevent_detection_ap(solution, noisy_submission_2, tolerances)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:28.993625Z","iopub.execute_input":"2022-07-31T03:09:28.994055Z","iopub.status.idle":"2022-07-31T03:09:45.676471Z","shell.execute_reply.started":"2022-07-31T03:09:28.994021Z","shell.execute_reply":"2022-07-31T03:09:45.675287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---","metadata":{}},{"cell_type":"markdown","source":"<a id='#2'></a>\n<h1 align='left'>関数の中身理解 / Understanding the content of functions </h1>","metadata":{}},{"cell_type":"markdown","source":"## MAIN： event_detection_ap(solution, noisy_submission, tolerances)  \nこの関数がメインなので、これを上から順番に確認していきます。  \nSince this function is the main, we will check it in order from the top.","metadata":{}},{"cell_type":"code","source":"# solution\nprint(perfect_submission.shape)\nperfect_submission.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:45.678252Z","iopub.execute_input":"2022-07-31T03:09:45.678976Z","iopub.status.idle":"2022-07-31T03:09:45.693856Z","shell.execute_reply.started":"2022-07-31T03:09:45.678928Z","shell.execute_reply":"2022-07-31T03:09:45.692664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# noisy_submission\nnoisy_submission = perfect_submission.copy()\nidx = noisy_submission.sample(frac=0.1).index # 10%だけidxを選ぶ\nidx","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:45.695569Z","iopub.execute_input":"2022-07-31T03:09:45.696471Z","iopub.status.idle":"2022-07-31T03:09:45.712600Z","shell.execute_reply.started":"2022-07-31T03:09:45.696423Z","shell.execute_reply":"2022-07-31T03:09:45.711344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# noisy_submissionの'event'列をシャッフルする。\nnoisy_submission.loc[idx, 'event'] = noisy_submission.loc[idx, 'event'].sort_index().to_numpy()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:45.730809Z","iopub.execute_input":"2022-07-31T03:09:45.731969Z","iopub.status.idle":"2022-07-31T03:09:45.741931Z","shell.execute_reply.started":"2022-07-31T03:09:45.731924Z","shell.execute_reply":"2022-07-31T03:09:45.740565Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"solution","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:45.743718Z","iopub.execute_input":"2022-07-31T03:09:45.744638Z","iopub.status.idle":"2022-07-31T03:09:45.763829Z","shell.execute_reply.started":"2022-07-31T03:09:45.744563Z","shell.execute_reply":"2022-07-31T03:09:45.762425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"noisy_submission","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:45.766034Z","iopub.execute_input":"2022-07-31T03:09:45.766695Z","iopub.status.idle":"2022-07-31T03:09:45.790967Z","shell.execute_reply.started":"2022-07-31T03:09:45.766658Z","shell.execute_reply":"2022-07-31T03:09:45.789050Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# indexが'video_id', 'time', 'event'を含むかを確認\nassert_index_equal(solution.columns, pd.Index(['video_id', 'time', 'event']))\n# AssertionError: Index are different\n\n# エラー発生時は下記のように出力\n    # Index values are different (33.33333 %)\n    # [left]:  Index(['video_id', 'time', 'event'], dtype='object')\n    # [right]: Index(['video_id', 'time', 'event__'], dtype='object')","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:45.792927Z","iopub.execute_input":"2022-07-31T03:09:45.793844Z","iopub.status.idle":"2022-07-31T03:09:45.801851Z","shell.execute_reply.started":"2022-07-31T03:09:45.793797Z","shell.execute_reply":"2022-07-31T03:09:45.800716Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 問題ない場合は出力は何もない\nsubmission = noisy_submission\nassert_index_equal(submission.columns, pd.Index(['video_id', 'time', 'event', 'score']))","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:45.805068Z","iopub.execute_input":"2022-07-31T03:09:45.806147Z","iopub.status.idle":"2022-07-31T03:09:45.813716Z","shell.execute_reply.started":"2022-07-31T03:09:45.806105Z","shell.execute_reply":"2022-07-31T03:09:45.812873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Extract scoring intervals.\nintervals = (\n    solution\n    .query(\"event in ['start', 'end']\")\n    .assign(interval=lambda x: x.groupby(['video_id', 'event']).cumcount())\n    .pivot(index='interval', columns=['video_id', 'event'], values='time')\n    .stack('video_id')\n    .swaplevel()\n    .sort_index()\n    .loc[:, ['start', 'end']]\n    .apply(lambda x: pd.Interval(*x, closed='both'), axis=1)\n)\n\nintervals","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:45.815756Z","iopub.execute_input":"2022-07-31T03:09:45.816488Z","iopub.status.idle":"2022-07-31T03:09:46.041250Z","shell.execute_reply.started":"2022-07-31T03:09:45.816445Z","shell.execute_reply":"2022-07-31T03:09:46.040120Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# startとendのみを切り出し\nintervals_ = (\n    solution\n    .query(\"event in ['start', 'end']\")\n#     .assign(interval=lambda x: x.groupby(['video_id', 'event']).cumcount())\n#     .pivot(index='interval', columns=['video_id', 'event'], values='time')\n#     .stack('video_id')\n#     .swaplevel()\n#     .sort_index()\n#     .loc[:, ['start', 'end']]\n#     .apply(lambda x: pd.Interval(*x, closed='both'), axis=1)\n)\n\nintervals_","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.042464Z","iopub.execute_input":"2022-07-31T03:09:46.042789Z","iopub.status.idle":"2022-07-31T03:09:46.063849Z","shell.execute_reply.started":"2022-07-31T03:09:46.042760Z","shell.execute_reply":"2022-07-31T03:09:46.062541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# cumcountで累積回数をナンバリング\nintervals_ = (\n    solution\n    .query(\"event in ['start', 'end']\")\n    .assign(interval=lambda x: x.groupby(['video_id', 'event']).cumcount())\n#     .pivot(index='interval', columns=['video_id', 'event'], values='time')\n#     .stack('video_id')\n#     .swaplevel()\n#     .sort_index()\n#     .loc[:, ['start', 'end']]\n#     .apply(lambda x: pd.Interval(*x, closed='both'), axis=1)\n)\n\nintervals_","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.065474Z","iopub.execute_input":"2022-07-31T03:09:46.067374Z","iopub.status.idle":"2022-07-31T03:09:46.093689Z","shell.execute_reply.started":"2022-07-31T03:09:46.067334Z","shell.execute_reply":"2022-07-31T03:09:46.092815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# index='interval', columns=['video_id', 'event'], values='time'でinterval毎にピボット\nintervals_ = (\n    solution\n    .query(\"event in ['start', 'end']\")\n    .assign(interval=lambda x: x.groupby(['video_id', 'event']).cumcount())\n    .pivot(index='interval', columns=['video_id', 'event'], values='time')\n#     .stack('video_id')\n#     .swaplevel()\n#     .sort_index()\n#     .loc[:, ['start', 'end']]\n#     .apply(lambda x: pd.Interval(*x, closed='both'), axis=1)\n)\n\nintervals_","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.095185Z","iopub.execute_input":"2022-07-31T03:09:46.096263Z","iopub.status.idle":"2022-07-31T03:09:46.150497Z","shell.execute_reply.started":"2022-07-31T03:09:46.096184Z","shell.execute_reply":"2022-07-31T03:09:46.149504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# video_id毎にスタックする。\nintervals_ = (\n    solution\n    .query(\"event in ['start', 'end']\")\n    .assign(interval=lambda x: x.groupby(['video_id', 'event']).cumcount())\n    .pivot(index='interval', columns=['video_id', 'event'], values='time')\n    .stack('video_id')\n#     .swaplevel()\n#     .sort_index()\n#     .loc[:, ['start', 'end']]\n#     .apply(lambda x: pd.Interval(*x, closed='both'), axis=1)\n)\n\nintervals_","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.151949Z","iopub.execute_input":"2022-07-31T03:09:46.152501Z","iopub.status.idle":"2022-07-31T03:09:46.187548Z","shell.execute_reply.started":"2022-07-31T03:09:46.152465Z","shell.execute_reply":"2022-07-31T03:09:46.186425Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# video_idとintervalをswap\nintervals_ = (\n    solution\n    .query(\"event in ['start', 'end']\")\n    .assign(interval=lambda x: x.groupby(['video_id', 'event']).cumcount())\n    .pivot(index='interval', columns=['video_id', 'event'], values='time')\n    .stack('video_id')\n    .swaplevel()\n#     .sort_index()\n#     .loc[:, ['start', 'end']]\n#     .apply(lambda x: pd.Interval(*x, closed='both'), axis=1)\n)\n\nintervals_","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.189022Z","iopub.execute_input":"2022-07-31T03:09:46.189806Z","iopub.status.idle":"2022-07-31T03:09:46.227444Z","shell.execute_reply.started":"2022-07-31T03:09:46.189758Z","shell.execute_reply":"2022-07-31T03:09:46.226481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# intervalでソート\nintervals_ = (\n    solution\n    .query(\"event in ['start', 'end']\")\n    .assign(interval=lambda x: x.groupby(['video_id', 'event']).cumcount())\n    .pivot(index='interval', columns=['video_id', 'event'], values='time')\n    .stack('video_id')\n    .swaplevel()\n    .sort_index()\n#     .loc[:, ['start', 'end']]\n#     .apply(lambda x: pd.Interval(*x, closed='both'), axis=1)\n)\n\nintervals_","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.228975Z","iopub.execute_input":"2022-07-31T03:09:46.229330Z","iopub.status.idle":"2022-07-31T03:09:46.268787Z","shell.execute_reply.started":"2022-07-31T03:09:46.229298Z","shell.execute_reply":"2022-07-31T03:09:46.267439Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# pd.Intervalでスライス区切りのような形に変形\nintervals_ = (\n    solution\n    .query(\"event in ['start', 'end']\")\n    .assign(interval=lambda x: x.groupby(['video_id', 'event']).cumcount())\n    .pivot(index='interval', columns=['video_id', 'event'], values='time')\n    .stack('video_id')\n    .swaplevel()\n    .sort_index()\n    .loc[:, ['start', 'end']]\n    .apply(lambda x: pd.Interval(*x, closed='both'), axis=1)\n)\n\nintervals_","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.270316Z","iopub.execute_input":"2022-07-31T03:09:46.270636Z","iopub.status.idle":"2022-07-31T03:09:46.350972Z","shell.execute_reply.started":"2022-07-31T03:09:46.270606Z","shell.execute_reply":"2022-07-31T03:09:46.349870Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Extract ground-truth events. (event列でstart, end以外のものを抜き出す。)\nground_truths = (\n    solution\n    .query(\"event not in ['start', 'end']\")\n    .reset_index(drop=True)\n)\nground_truths","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.352426Z","iopub.execute_input":"2022-07-31T03:09:46.352749Z","iopub.status.idle":"2022-07-31T03:09:46.371916Z","shell.execute_reply.started":"2022-07-31T03:09:46.352720Z","shell.execute_reply":"2022-07-31T03:09:46.371019Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Map each event class to its prevalence (needed for recall calculation)\n# eventとその回数の辞書をつくる\nclass_counts = ground_truths.value_counts('event').to_dict()\nclass_counts","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.372988Z","iopub.execute_input":"2022-07-31T03:09:46.373521Z","iopub.status.idle":"2022-07-31T03:09:46.381941Z","shell.execute_reply.started":"2022-07-31T03:09:46.373487Z","shell.execute_reply":"2022-07-31T03:09:46.380861Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Create table for detections with a column indicating a match to a ground-truth event\n# matchedという列をsubmissionに追加し、全てFalseを入れて、detectionsというDataFrameにすうｒ．\ndetections = submission.assign(matched = False)\ndetections","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.383546Z","iopub.execute_input":"2022-07-31T03:09:46.383902Z","iopub.status.idle":"2022-07-31T03:09:46.409191Z","shell.execute_reply.started":"2022-07-31T03:09:46.383873Z","shell.execute_reply":"2022-07-31T03:09:46.407895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Remove detections outside of scoring intervals\ndetections_filtered = []\n\nfor (det_group, dets), (int_group, ints) in zip(\n    detections.groupby('video_id'), intervals.groupby('video_id')\n):\n    assert det_group == int_group\n    detections_filtered.append(filter_detections(dets, ints))\ndetections_filtered = pd.concat(detections_filtered, ignore_index=True)\n\ndetections_filtered","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.411162Z","iopub.execute_input":"2022-07-31T03:09:46.411624Z","iopub.status.idle":"2022-07-31T03:09:46.466391Z","shell.execute_reply.started":"2022-07-31T03:09:46.411579Z","shell.execute_reply":"2022-07-31T03:09:46.465155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# はじめの１データだけでforループを止めて中身を確認 \nfor (det_group, dets), (int_group, ints) in zip(\n    detections.groupby('video_id'), intervals.groupby('video_id')\n):\n    assert det_group == int_group \n    detections_filtered.append(filter_detections(dets, ints))\n    break ","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.468328Z","iopub.execute_input":"2022-07-31T03:09:46.468764Z","iopub.status.idle":"2022-07-31T03:09:46.482053Z","shell.execute_reply.started":"2022-07-31T03:09:46.468722Z","shell.execute_reply":"2022-07-31T03:09:46.481254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# detectionsから、det_group==video_idとdets->video_idのデータを呼び出す。\n(det_group, dets)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.483755Z","iopub.execute_input":"2022-07-31T03:09:46.484664Z","iopub.status.idle":"2022-07-31T03:09:46.498187Z","shell.execute_reply.started":"2022-07-31T03:09:46.484619Z","shell.execute_reply":"2022-07-31T03:09:46.497125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# intervals、int_group==video_idとdets->video_idのデータを呼び出す。\n(int_group, ints)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.499777Z","iopub.execute_input":"2022-07-31T03:09:46.500129Z","iopub.status.idle":"2022-07-31T03:09:46.511615Z","shell.execute_reply.started":"2022-07-31T03:09:46.500090Z","shell.execute_reply":"2022-07-31T03:09:46.510723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 互いのvideo_idが一致しているか確認 \nassert det_group == int_group ","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.513189Z","iopub.execute_input":"2022-07-31T03:09:46.513804Z","iopub.status.idle":"2022-07-31T03:09:46.519644Z","shell.execute_reply.started":"2022-07-31T03:09:46.513767Z","shell.execute_reply":"2022-07-31T03:09:46.518539Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# filter_detections関数を実行 \n# detections_filtered.append(filter_detections(dets, ints))\nfilter_detections(dets, ints)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.521856Z","iopub.execute_input":"2022-07-31T03:09:46.522732Z","iopub.status.idle":"2022-07-31T03:09:46.547858Z","shell.execute_reply.started":"2022-07-31T03:09:46.522686Z","shell.execute_reply":"2022-07-31T03:09:46.546739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---","metadata":{}},{"cell_type":"markdown","source":"## filter_detections関数の中身確認  \nここで、filter_detections関数を実行しています。その中身を確認します。  \nHere I am running the filter_detections function. Check its contents.","metadata":{}},{"cell_type":"code","source":"# detectionsのtimeをnp.arrayとしてdetection_timeに渡す \ndetections_ = dets\nintervals_  = ints\ndetection_time = detections_.loc[:, 'time'].sort_values().to_numpy()\nprint(len(detection_time))\ndetection_time[:4]","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.549530Z","iopub.execute_input":"2022-07-31T03:09:46.550022Z","iopub.status.idle":"2022-07-31T03:09:46.559808Z","shell.execute_reply.started":"2022-07-31T03:09:46.549987Z","shell.execute_reply":"2022-07-31T03:09:46.558646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# numpy化\nintervals_ = intervals_.to_numpy()\nintervals_[:4]","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.561227Z","iopub.execute_input":"2022-07-31T03:09:46.562135Z","iopub.status.idle":"2022-07-31T03:09:46.572706Z","shell.execute_reply.started":"2022-07-31T03:09:46.562063Z","shell.execute_reply":"2022-07-31T03:09:46.571499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# np.full_like: detection_timeと同じshapeでFalseを埋める\nis_scored = np.full_like(detection_time, False, dtype=bool)\nprint(len(is_scored))\nis_scored[:4]","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.575226Z","iopub.execute_input":"2022-07-31T03:09:46.576131Z","iopub.status.idle":"2022-07-31T03:09:46.584972Z","shell.execute_reply.started":"2022-07-31T03:09:46.576065Z","shell.execute_reply":"2022-07-31T03:09:46.583880Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# while文の中がどのように処理されるのか確認するための関数 \ndef print_ij(i, j):\n    if i < 10 and j < 10:\n        print(\n            f\"i: {i}, j: {j},\\n\\\n            detection_time[i]:{detection_time[i]},\\n\\\n            intervals_[j]: {intervals_[j]}\\n\\\n            is_scored[i]:{is_scored[i]} \"\n        )","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.586525Z","iopub.execute_input":"2022-07-31T03:09:46.586869Z","iopub.status.idle":"2022-07-31T03:09:46.594332Z","shell.execute_reply.started":"2022-07-31T03:09:46.586838Z","shell.execute_reply":"2022-07-31T03:09:46.592966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# while文の中では、detection_timeが startとendの間にあると、False->Trueに変換\ni, j = 0, 0\nwhile i < len(detection_time) and j < len(intervals_):\n    time = detection_time[i]\n    int_ = intervals_[j]\n\n    # If the detection is prior in time to the interval, go to the next detection.\n    # int_.leftはstartの時刻を示す。それよりもtimeが小さいならば、iをインクリメント\n    if time < int_.left:\n        i += 1\n        print_ij(i-1, j)\n    # If the detection is inside the interval, keep it and go to the next detection.  \n    # int_にtimeが含まれていればis_scored[i]をFalse->Trueに変える\n    elif time in int_:\n        is_scored[i] = True\n        i += 1\n        print_ij(i-1, j)\n    # If the detection is later in time, go to the next interval.\n    # time > int_.rightならばjをインクリメント\n    else:\n        j += 1\n        print_ij(i, j-1)\n\nis_scored[:4]","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.598271Z","iopub.execute_input":"2022-07-31T03:09:46.598633Z","iopub.status.idle":"2022-07-31T03:09:46.611420Z","shell.execute_reply.started":"2022-07-31T03:09:46.598601Z","shell.execute_reply":"2022-07-31T03:09:46.610341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# return detections.loc[is_scored].reset_index(drop=True)\ndetections_.loc[is_scored].reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.612398Z","iopub.execute_input":"2022-07-31T03:09:46.612749Z","iopub.status.idle":"2022-07-31T03:09:46.640915Z","shell.execute_reply.started":"2022-07-31T03:09:46.612714Z","shell.execute_reply":"2022-07-31T03:09:46.639953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### ↑  filter_detections関数の中身確認完了","metadata":{}},{"cell_type":"markdown","source":"---","metadata":{}},{"cell_type":"code","source":"# Create table of event-class x tolerance x video_id values\naggregation_keys = pd.DataFrame(\n    [(ev, tol, vid)\n     for ev in tolerances.keys()\n     for tol in tolerances[ev]\n     for vid in ground_truths['video_id'].unique()],\n    columns=['event', 'tolerance', 'video_id'],\n)\n\naggregation_keys","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.642495Z","iopub.execute_input":"2022-07-31T03:09:46.642825Z","iopub.status.idle":"2022-07-31T03:09:46.671339Z","shell.execute_reply.started":"2022-07-31T03:09:46.642785Z","shell.execute_reply":"2022-07-31T03:09:46.670161Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tolerances","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.673246Z","iopub.execute_input":"2022-07-31T03:09:46.673671Z","iopub.status.idle":"2022-07-31T03:09:46.680309Z","shell.execute_reply.started":"2022-07-31T03:09:46.673631Z","shell.execute_reply":"2022-07-31T03:09:46.679359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ground_truths['video_id'].unique()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.682007Z","iopub.execute_input":"2022-07-31T03:09:46.682440Z","iopub.status.idle":"2022-07-31T03:09:46.695540Z","shell.execute_reply.started":"2022-07-31T03:09:46.682408Z","shell.execute_reply":"2022-07-31T03:09:46.694709Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# detections_filtered やground_truthsをaggregation_keysにマージする。\n# Create match evaluation groups: event-class x tolerance x video_id\ndetections_grouped = (\n    aggregation_keys\n    .merge(detections_filtered, on=['event', 'video_id'], how='left')\n    .groupby(['event', 'tolerance', 'video_id'])\n)\nground_truths_grouped = (\n    aggregation_keys\n    .merge(ground_truths, on=['event', 'video_id'], how='left')\n    .groupby(['event', 'tolerance', 'video_id'])\n)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.696849Z","iopub.execute_input":"2022-07-31T03:09:46.697214Z","iopub.status.idle":"2022-07-31T03:09:46.719597Z","shell.execute_reply.started":"2022-07-31T03:09:46.697174Z","shell.execute_reply":"2022-07-31T03:09:46.718656Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"detections_grouped_ = (\n    aggregation_keys\n    .merge(detections_filtered, on=['event', 'video_id'], how='left')\n#     .groupby(['event', 'tolerance', 'video_id'])\n)\n\ndetections_grouped_","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.721211Z","iopub.execute_input":"2022-07-31T03:09:46.721981Z","iopub.status.idle":"2022-07-31T03:09:46.747974Z","shell.execute_reply.started":"2022-07-31T03:09:46.721935Z","shell.execute_reply":"2022-07-31T03:09:46.746792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ground_truths_grouped_ = (\n    aggregation_keys\n    .merge(ground_truths, on=['event', 'video_id'], how='left')\n#     .groupby(['event', 'tolerance', 'video_id'])\n)\nground_truths_grouped_","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.749385Z","iopub.execute_input":"2022-07-31T03:09:46.749803Z","iopub.status.idle":"2022-07-31T03:09:46.775056Z","shell.execute_reply.started":"2022-07-31T03:09:46.749768Z","shell.execute_reply":"2022-07-31T03:09:46.773697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Match detections to ground truth events by evaluation group\ndetections_matched = []\nfor key in aggregation_keys.itertuples(index=False):\n    dets = detections_grouped.get_group(key)\n    gts = ground_truths_grouped.get_group(key)\n    detections_matched.append(\n        match_detections(dets['tolerance'].iloc[0], gts, dets)\n    )\ndetections_matched = pd.concat(detections_matched)\n\ndetections_matched","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:09:46.776395Z","iopub.execute_input":"2022-07-31T03:09:46.776799Z","iopub.status.idle":"2022-07-31T03:10:03.848590Z","shell.execute_reply.started":"2022-07-31T03:09:46.776766Z","shell.execute_reply":"2022-07-31T03:10:03.847549Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Match detections to ground truth events by evaluation group\ndetections_matched_ = []\nfor key in aggregation_keys.itertuples(index=False):\n    dets = detections_grouped.get_group(key)\n    gts = ground_truths_grouped.get_group(key)\n    # ここでmatch_detections関数を使います。\n    detections_matched_.append(\n        match_detections(dets['tolerance'].iloc[0], gts, dets) \n    )\n    break","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:03.850034Z","iopub.execute_input":"2022-07-31T03:10:03.850395Z","iopub.status.idle":"2022-07-31T03:10:03.889915Z","shell.execute_reply.started":"2022-07-31T03:10:03.850364Z","shell.execute_reply":"2022-07-31T03:10:03.888763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dets.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:03.892259Z","iopub.execute_input":"2022-07-31T03:10:03.892984Z","iopub.status.idle":"2022-07-31T03:10:03.908541Z","shell.execute_reply.started":"2022-07-31T03:10:03.892932Z","shell.execute_reply":"2022-07-31T03:10:03.907419Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gts.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:03.909589Z","iopub.execute_input":"2022-07-31T03:10:03.910635Z","iopub.status.idle":"2022-07-31T03:10:03.923522Z","shell.execute_reply.started":"2022-07-31T03:10:03.910597Z","shell.execute_reply":"2022-07-31T03:10:03.922580Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"match_detections(dets['tolerance'].iloc[0], gts, dets)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:03.924780Z","iopub.execute_input":"2022-07-31T03:10:03.925328Z","iopub.status.idle":"2022-07-31T03:10:03.993971Z","shell.execute_reply.started":"2022-07-31T03:10:03.925294Z","shell.execute_reply":"2022-07-31T03:10:03.992743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---\n## match_detections関数の確認  \nここで、match_detections関数を実行しています。その中身を確認します。  \nHere I am running the match_detections function. Check its contents.","metadata":{}},{"cell_type":"code","source":"# def match_detections(\n#     tolerance: float, \n#     ground_truths: pd.DataFrame, \n#     detections: pd.DataFrame\n# ) -> pd.DataFrame:\ntolerance_, ground_truths_, detections_=dets['tolerance'].iloc[0], gts, dets","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:03.995578Z","iopub.execute_input":"2022-07-31T03:10:03.995938Z","iopub.status.idle":"2022-07-31T03:10:04.001027Z","shell.execute_reply.started":"2022-07-31T03:10:03.995905Z","shell.execute_reply":"2022-07-31T03:10:04.000255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Scoreでソート\ndetections_sorted = detections_.sort_values('score', ascending=False).dropna()\nprint(detections_sorted.shape)\ndetections_sorted.head()","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.002069Z","iopub.execute_input":"2022-07-31T03:10:04.002942Z","iopub.status.idle":"2022-07-31T03:10:04.026879Z","shell.execute_reply.started":"2022-07-31T03:10:04.002905Z","shell.execute_reply":"2022-07-31T03:10:04.025688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# is_matchedというFalseのarrayを作る\nis_matched = np.full_like(detections_sorted['event'], False, dtype=bool)\nprint(len(is_matched))\nis_matched","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.028717Z","iopub.execute_input":"2022-07-31T03:10:04.029849Z","iopub.status.idle":"2022-07-31T03:10:04.040854Z","shell.execute_reply.started":"2022-07-31T03:10:04.029801Z","shell.execute_reply":"2022-07-31T03:10:04.039657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"tolerance_","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.042497Z","iopub.execute_input":"2022-07-31T03:10:04.043148Z","iopub.status.idle":"2022-07-31T03:10:04.053296Z","shell.execute_reply.started":"2022-07-31T03:10:04.043105Z","shell.execute_reply":"2022-07-31T03:10:04.052000Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gts_matched = set()\n\nfor i, det in enumerate(detections_sorted.itertuples(index=False)):\n    if i < 5:\n        print(i)\n        print(f\"det: {det}\")\n    best_error = tolerance_ # best_errorの初期値設定\n    best_gt = None\n\n    j = 0\n    for gt in ground_truths.itertuples(index=False):      \n        error = abs(det.time - gt.time) # timeのズレ分をerrorとする\n   \n        if i<5 and j < 5:\n            print(f\"gt: {gt}\")\n            print(f\"error: {error}, best_error: {best_error}\")\n            j += 1\n\n        # errorがbest かつ　gtが新規ならば更新\n        if error < best_error and not gt in gts_matched:\n            best_gt = gt\n            best_error = error\n\n    if best_gt is not None:\n        is_matched[i] = True\n        gts_matched.add(best_gt)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.054773Z","iopub.execute_input":"2022-07-31T03:10:04.055418Z","iopub.status.idle":"2022-07-31T03:10:04.395257Z","shell.execute_reply.started":"2022-07-31T03:10:04.055383Z","shell.execute_reply":"2022-07-31T03:10:04.394106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gts_matched","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.396791Z","iopub.execute_input":"2022-07-31T03:10:04.397256Z","iopub.status.idle":"2022-07-31T03:10:04.406818Z","shell.execute_reply.started":"2022-07-31T03:10:04.397214Z","shell.execute_reply":"2022-07-31T03:10:04.405474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# detections_sorted['matched'] = is_matched\n# return detections_sorted\ndetections_sorted['matched'] = is_matched\ndetections_sorted.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.408763Z","iopub.execute_input":"2022-07-31T03:10:04.409252Z","iopub.status.idle":"2022-07-31T03:10:04.429626Z","shell.execute_reply.started":"2022-07-31T03:10:04.409210Z","shell.execute_reply":"2022-07-31T03:10:04.428530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### ↑ ここまででmatch_detections関数の中身確認完了","metadata":{}},{"cell_type":"markdown","source":"---","metadata":{}},{"cell_type":"code","source":"# Compute AP per event x tolerance group\nevent_classes = ground_truths['event'].unique()\nap_table = (\n    detections_matched\n    .query(\"event in @event_classes\")\n    .groupby(['event', 'tolerance']).apply(\n    lambda group: average_precision_score(   # average_precision_score関数をここで使っている。\n    group['matched'].to_numpy(),\n            group['score'].to_numpy(),\n            class_counts[group['event'].iat[0]],\n        )\n    )\n)\n\nap_table","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.431257Z","iopub.execute_input":"2022-07-31T03:10:04.431581Z","iopub.status.idle":"2022-07-31T03:10:04.465059Z","shell.execute_reply.started":"2022-07-31T03:10:04.431550Z","shell.execute_reply":"2022-07-31T03:10:04.464233Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---\n## average_precision_score関数の確認\nここで、average_precision_score関数を実行しています。その中身を確認します。  \nHere I am running the average_precision_score function. Check its contents.","metadata":{}},{"cell_type":"code","source":"ap_table_ = (\n    detections_matched\n    .query(\"event in @event_classes\")\n#     .groupby(['event', 'tolerance']) #.apply(\n#     lambda group: average_precision_score(   # average_precision_score関数をここで使っている。\n#     group['matched'].to_numpy(),\n#             group['score'].to_numpy(),\n#             class_counts[group['event'].iat[0]],\n#         )\n#     )\n)\nap_table_","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.466264Z","iopub.execute_input":"2022-07-31T03:10:04.467454Z","iopub.status.idle":"2022-07-31T03:10:04.495845Z","shell.execute_reply.started":"2022-07-31T03:10:04.467404Z","shell.execute_reply":"2022-07-31T03:10:04.494947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# group['matched'].to_numpy()\nmatches = ap_table_.query('event==\"challenge\" & tolerance==0.30')['matched'].to_numpy()\nprint(matches.shape)\nmatches[:3]","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.497019Z","iopub.execute_input":"2022-07-31T03:10:04.497897Z","iopub.status.idle":"2022-07-31T03:10:04.514314Z","shell.execute_reply.started":"2022-07-31T03:10:04.497858Z","shell.execute_reply":"2022-07-31T03:10:04.512707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# group['score'].to_numpy()\nscores = ap_table_.query('event==\"challenge\" & tolerance==0.30')['score'].to_numpy()\nprint(scores.shape)\nmatches[:3]","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.515600Z","iopub.execute_input":"2022-07-31T03:10:04.516892Z","iopub.status.idle":"2022-07-31T03:10:04.532205Z","shell.execute_reply.started":"2022-07-31T03:10:04.516836Z","shell.execute_reply":"2022-07-31T03:10:04.530785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class_counts[group['event'].iat[0]]\nprint(class_counts)\ntmp = ap_table_.query('event==\"challenge\"')[\"event\"]\np = class_counts[tmp.iat[0]] # 0行目をiat[0]で取得 -> この例ではclass_counts[\"challenge\"] となる\np","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.533814Z","iopub.execute_input":"2022-07-31T03:10:04.534484Z","iopub.status.idle":"2022-07-31T03:10:04.547252Z","shell.execute_reply.started":"2022-07-31T03:10:04.534443Z","shell.execute_reply":"2022-07-31T03:10:04.545757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def average_precision_score(matches: np.ndarray, scores: np.ndarray, p: int) -> float:\n#     precision, recall, _ = precision_recall_curve(matches, scores, p)\n#     # Compute step integral\n#     return -np.sum(np.diff(recall) * np.array(precision)[:-1])\n\nprecision, recall, _ = precision_recall_curve(matches, scores, p) # precision_recall_curve関数を実行\nprint(precision, recall, _)\nprint(-np.sum(np.diff(recall) * np.array(precision)[:-1]))","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.549357Z","iopub.execute_input":"2022-07-31T03:10:04.549818Z","iopub.status.idle":"2022-07-31T03:10:04.558745Z","shell.execute_reply.started":"2022-07-31T03:10:04.549770Z","shell.execute_reply":"2022-07-31T03:10:04.557434Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---\n## precision_recall_curve関数の確認  \nここで、precision_recall_curve関数を実行しています。その中身を確認します。  \nHere I am running the precision_recall_curve function. Check its contents.","metadata":{}},{"cell_type":"code","source":"# def precision_recall_curve(\n#     matches: np.ndarray, \n#     scores: np.ndarray, \n#     p: int\n# ) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:\n\n# matchesが何もない場合\n# if len(matches) == 0:\n#     return [1], [0], []","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.560367Z","iopub.execute_input":"2022-07-31T03:10:04.561594Z","iopub.status.idle":"2022-07-31T03:10:04.571225Z","shell.execute_reply.started":"2022-07-31T03:10:04.561550Z","shell.execute_reply":"2022-07-31T03:10:04.570032Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Sort matches by decreasing confidence\nidxs = np.argsort(scores, kind='stable')[::-1] # scoresをソート\nscores_ = scores[idxs]\nmatches_ = matches[idxs]\n\nprint(idxs[:3])\nprint(scores_[:3])\nprint(matches_[:3])","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.572926Z","iopub.execute_input":"2022-07-31T03:10:04.573333Z","iopub.status.idle":"2022-07-31T03:10:04.584628Z","shell.execute_reply.started":"2022-07-31T03:10:04.573296Z","shell.execute_reply":"2022-07-31T03:10:04.583415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"distinct_value_indices = np.where(np.diff(scores))[0]\nthreshold_idxs = np.r_[distinct_value_indices, matches.size - 1] # np.r_でdistinct_value_indicesとmatches.size - 1を結合\nthresholds = scores[threshold_idxs]\n\nprint(distinct_value_indices)\nprint(threshold_idxs)\nprint(thresholds)","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.585567Z","iopub.execute_input":"2022-07-31T03:10:04.585883Z","iopub.status.idle":"2022-07-31T03:10:04.597185Z","shell.execute_reply.started":"2022-07-31T03:10:04.585854Z","shell.execute_reply":"2022-07-31T03:10:04.596380Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Matches become TPs and non-matches FPs as confidence threshold decreases\ntps = np.cumsum(matches)[threshold_idxs]  # Trueの総和\nfps = np.cumsum(~matches)[threshold_idxs] # Falseの総和\n\nprecision = tps / (tps + fps)\nprecision[np.isnan(precision)] = 0\nrecall = tps / p  # total number of ground truths might be different than total number of matches\n\nprint(f\"p: {p}\")\nprint(f\"tps: {tps}, fps: {fps}\")\nprint(f\"precision: {precision}, recall: {recall}\")","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.598528Z","iopub.execute_input":"2022-07-31T03:10:04.598994Z","iopub.status.idle":"2022-07-31T03:10:04.610670Z","shell.execute_reply.started":"2022-07-31T03:10:04.598963Z","shell.execute_reply":"2022-07-31T03:10:04.609615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Stop when full recall attained and reverse the outputs so recall is non-increasing.\nlast_ind = tps.searchsorted(tps[-1])\nsl = slice(last_ind, None, -1)\n\nprint(f\"last_ind: {last_ind}\")\nprint(f\"sl: {sl}\")","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.612050Z","iopub.execute_input":"2022-07-31T03:10:04.612576Z","iopub.status.idle":"2022-07-31T03:10:04.624916Z","shell.execute_reply.started":"2022-07-31T03:10:04.612544Z","shell.execute_reply":"2022-07-31T03:10:04.624016Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Final precision is 1 and final recall is 0\n# return np.r_[precision[sl], 1], np.r_[recall[sl], 0], thresholds[sl]\nnp.r_[precision[sl], 1], np.r_[recall[sl], 0], thresholds[sl]","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.626328Z","iopub.execute_input":"2022-07-31T03:10:04.626677Z","iopub.status.idle":"2022-07-31T03:10:04.638979Z","shell.execute_reply.started":"2022-07-31T03:10:04.626647Z","shell.execute_reply":"2022-07-31T03:10:04.637999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### ↑ precision_recall_curve関数の確認が完了 \n### ↑ average_precision_score関数の確認が完了 ","metadata":{}},{"cell_type":"markdown","source":"---","metadata":{}},{"cell_type":"code","source":"event_classes","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.640980Z","iopub.execute_input":"2022-07-31T03:10:04.641458Z","iopub.status.idle":"2022-07-31T03:10:04.650046Z","shell.execute_reply.started":"2022-07-31T03:10:04.641412Z","shell.execute_reply":"2022-07-31T03:10:04.649171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ap_table","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.651311Z","iopub.execute_input":"2022-07-31T03:10:04.652585Z","iopub.status.idle":"2022-07-31T03:10:04.665759Z","shell.execute_reply.started":"2022-07-31T03:10:04.652397Z","shell.execute_reply":"2022-07-31T03:10:04.664206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Average over tolerances, then over event classes\nmean_ap = ap_table.groupby('event').mean().mean()\n# return mean_ap\nmean_ap","metadata":{"execution":{"iopub.status.busy":"2022-07-31T03:10:04.667024Z","iopub.execute_input":"2022-07-31T03:10:04.668040Z","iopub.status.idle":"2022-07-31T03:10:04.676310Z","shell.execute_reply.started":"2022-07-31T03:10:04.668002Z","shell.execute_reply":"2022-07-31T03:10:04.675152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Fin","metadata":{}}]}