{"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    print(ap_table.groupby('event').mean())\n    mean_ap = ap_table.groupby('event').mean().mean()\n\n    return mean_ap","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:19:01.922729Z","iopub.execute_input":"2022-08-09T03:19:01.923253Z","iopub.status.idle":"2022-08-09T03:19:01.983252Z","shell.execute_reply.started":"2022-08-09T03:19:01.923141Z","shell.execute_reply":"2022-08-09T03:19:01.982082Z"},"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-09T03:19:01.985712Z","iopub.execute_input":"2022-08-09T03:19:01.986560Z","iopub.status.idle":"2022-08-09T03:19:01.992147Z","shell.execute_reply.started":"2022-08-09T03:19:01.986513Z","shell.execute_reply":"2022-08-09T03:19:01.991050Z"},"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-09T03:19:01.993673Z","iopub.execute_input":"2022-08-09T03:19:01.993997Z","iopub.status.idle":"2022-08-09T03:19:02.059186Z","shell.execute_reply.started":"2022-08-09T03:19:01.993965Z","shell.execute_reply":"2022-08-09T03:19:02.057888Z"},"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":"import glob\nfiles = glob.glob(\"/kaggle/input/dfl-bundesliga-data-shootout/train/*\")","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:19:02.062884Z","iopub.execute_input":"2022-08-09T03:19:02.063348Z","iopub.status.idle":"2022-08-09T03:19:02.072119Z","shell.execute_reply.started":"2022-08-09T03:19:02.063311Z","shell.execute_reply":"2022-08-09T03:19:02.070966Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\ntrain_submission = []\nfor f in files:\n    times = np.arange(0, 60*90, 0.5)\n    for event in [\"play\", \"throwin\"]:\n        df = pd.DataFrame({\n            \"video_id\": [os.path.basename(f).replace(\".mp4\", \"\")] * len(times),\n            \"time\": times,\n            \"event\": [event] * len(times),\n            \"score\": [1.0] * len(times)\n        })\n        train_submission.append(df)    \n    \n    times = np.arange(0, 60*90, 1)\n    for event in [\"challenge\"]:\n        df = pd.DataFrame({\n            \"video_id\": [os.path.basename(f).replace(\".mp4\", \"\")] * len(times),\n            \"time\": times,\n            \"event\": [event] * len(times),\n            \"score\": [1.0] * len(times)\n        })\n        train_submission.append(df)     \ntrain_submission = pd.concat(train_submission).reset_index(drop=True)\ntrain_submission.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:24:19.477485Z","iopub.execute_input":"2022-08-09T03:24:19.477975Z","iopub.status.idle":"2022-08-09T03:24:19.611005Z","shell.execute_reply.started":"2022-08-09T03:24:19.477938Z","shell.execute_reply":"2022-08-09T03:24:19.609411Z"},"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, train_submission, tolerances)","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:24:20.048915Z","iopub.execute_input":"2022-08-09T03:24:20.049414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"files_test = glob.glob(\"/kaggle/input/dfl-bundesliga-data-shootout/test/*\")","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:19:02.161713Z","iopub.status.idle":"2022-08-09T03:19:02.162185Z","shell.execute_reply.started":"2022-08-09T03:19:02.161933Z","shell.execute_reply":"2022-08-09T03:19:02.161951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_submission = []\nfor f in files_test:\n    times = np.arange(0, 60*90, 0.5)\n    for event in [\"play\", \"throwin\"]:\n        df = pd.DataFrame({\n            \"video_id\": [os.path.basename(f).replace(\".mp4\", \"\")] * len(times),\n            \"time\": times,\n            \"event\": [event] * len(times),\n            \"score\": [1.0] * len(times)\n        })\n        test_submission.append(df)\n    times = np.arange(0, 60*90, 1)\n    for event in [\"throwin\", \"challenge\"]:\n        df = pd.DataFrame({\n            \"video_id\": [os.path.basename(f).replace(\".mp4\", \"\")] * len(times),\n            \"time\": times,\n            \"event\": [event] * len(times),\n            \"score\": [1.0] * len(times)\n        })\n        test_submission.append(df)\n    \ntest_submission = pd.concat(test_submission)\ntest_submission.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:19:02.164543Z","iopub.status.idle":"2022-08-09T03:19:02.166211Z","shell.execute_reply.started":"2022-08-09T03:19:02.165746Z","shell.execute_reply":"2022-08-09T03:19:02.165795Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_submission.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-08-09T03:19:02.168400Z","iopub.status.idle":"2022-08-09T03:19:02.169046Z","shell.execute_reply.started":"2022-08-09T03:19:02.168723Z","shell.execute_reply":"2022-08-09T03:19:02.168752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}