{"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":"code","source":"import numpy as np\nimport pandas as pd\nfrom pandas.testing import assert_index_equal\nfrom typing import Dict, Tuple\n\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}\n            \n\ndef filter_detections(\n        detections: pd.DataFrame, intervals: pd.DataFrame\n) -> pd.DataFrame:\n    \"\"\"Drop detections not inside a scoring interval.\"\"\"\n    detection_time = detections.loc[:, 'time'].sort_values().to_numpy()\n    intervals = intervals.to_numpy()\n    is_scored = np.full_like(detection_time, False, dtype=bool)\n\n    i, j = 0, 0\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        if time < int_.left:\n            i += 1\n        # If the detection is inside the interval, keep it and go to the next detection.        \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        else:\n            j += 1\n\n    return detections.loc[is_scored].reset_index(drop=True)\n\n\ndef match_detections(\n        tolerance: float, ground_truths: pd.DataFrame, detections: pd.DataFrame\n) -> pd.DataFrame:\n    \"\"\"Match detections to ground truth events. Arguments are taken from a common event x tolerance x video evaluation group.\"\"\"\n    detections_sorted = detections.sort_values('score', ascending=False).dropna()\n\n    is_matched = np.full_like(detections_sorted['event'], False, dtype=bool)\n    gts_matched = set()\n    for i, det in enumerate(detections_sorted.itertuples(index=False)):\n        best_error = tolerance\n        best_gt = None\n\n        for gt in ground_truths.itertuples(index=False):\n            error = abs(det.time - gt.time)\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    detections_sorted['matched'] = is_matched\n\n    return detections_sorted\n\n\ndef precision_recall_curve(\n        matches: np.ndarray, scores: np.ndarray, p: int\n) -> Tuple[np.ndarray, np.ndarray, np.ndarray]:\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]\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]\n    fps = np.cumsum(~matches)[threshold_idxs]\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]\n\n\ndef 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\n\ndef event_detection_ap(\n        solution: pd.DataFrame,\n        submission: pd.DataFrame,\n        tolerances: Dict[str, float],\n) -> float:\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    # Ensure solution and submission are sorted properly\n    solution = solution.sort_values(['video_id', 'time'])\n    submission = submission.sort_values(['video_id', 'time'])\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)\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(\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-08-08T16:05:35.622197Z","iopub.execute_input":"2022-08-08T16:05:35.622681Z","iopub.status.idle":"2022-08-08T16:05:35.685245Z","shell.execute_reply.started":"2022-08-08T16:05:35.622582Z","shell.execute_reply":"2022-08-08T16:05:35.684122Z"},"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-08-08T16:05:35.686643Z","iopub.execute_input":"2022-08-08T16:05:35.686985Z","iopub.status.idle":"2022-08-08T16:05:35.692954Z","shell.execute_reply.started":"2022-08-08T16:05:35.686955Z","shell.execute_reply":"2022-08-08T16:05:35.691625Z"},"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-08-08T16:05:35.694632Z","iopub.execute_input":"2022-08-08T16:05:35.695041Z","iopub.status.idle":"2022-08-08T16:05:35.756043Z","shell.execute_reply.started":"2022-08-08T16:05:35.694989Z","shell.execute_reply":"2022-08-08T16:05:35.755033Z"},"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-08-08T16:05:35.758058Z","iopub.execute_input":"2022-08-08T16:05:35.758392Z","iopub.status.idle":"2022-08-08T16:05:35.787163Z","shell.execute_reply.started":"2022-08-08T16:05:35.758363Z","shell.execute_reply":"2022-08-08T16:05:35.785942Z"},"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-08-08T16:05:35.788966Z","iopub.execute_input":"2022-08-08T16:05:35.789756Z","iopub.status.idle":"2022-08-08T16:05:53.504391Z","shell.execute_reply.started":"2022-08-08T16:05:35.789720Z","shell.execute_reply":"2022-08-08T16:05:53.502775Z"},"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-08-08T16:14:49.483561Z","iopub.execute_input":"2022-08-08T16:14:49.484021Z","iopub.status.idle":"2022-08-08T16:15:07.528367Z","shell.execute_reply.started":"2022-08-08T16:14:49.483984Z","shell.execute_reply":"2022-08-08T16:15:07.527020Z"},"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-08-08T16:06:10.447897Z","iopub.execute_input":"2022-08-08T16:06:10.448259Z","iopub.status.idle":"2022-08-08T16:06:27.559226Z","shell.execute_reply.started":"2022-08-08T16:06:10.448226Z","shell.execute_reply":"2022-08-08T16:06:27.558076Z"},"trusted":true},"execution_count":null,"outputs":[]}]}