{"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":"# GASLFR - SVG visualisation\n\nConverting a matplolib animation to a HTML video being extremely slow, I tried to develop a faster solution using SVG.","metadata":{}},{"cell_type":"code","source":"! pip install mediapipe -qq","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-06-21T03:00:17.769977Z","iopub.execute_input":"2023-06-21T03:00:17.770421Z","iopub.status.idle":"2023-06-21T03:00:32.049883Z","shell.execute_reply.started":"2023-06-21T03:00:17.770385Z","shell.execute_reply":"2023-06-21T03:00:32.048487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import json\nimport os\n\nfrom IPython.display import HTML\nfrom mediapipe import solutions as mp\nimport numpy as np\nimport pandas as pd\n\n\nINPUT_DIR = \"/kaggle/input/asl-fingerspelling\"\nOUTPUT_DIR = \"/kaggle/working\"\n\nHAND_CONNECTIONS = mp.hands_connections.HAND_CONNECTIONS\nPOSE_CONNECTIONS = mp.pose_connections.POSE_CONNECTIONS\nFACE_CONNECTIONS = mp.face_mesh_connections.FACEMESH_CONTOURS","metadata":{"execution":{"iopub.status.busy":"2023-06-21T03:00:32.052277Z","iopub.execute_input":"2023-06-21T03:00:32.052647Z","iopub.status.idle":"2023-06-21T03:00:40.845433Z","shell.execute_reply.started":"2023-06-21T03:00:32.05261Z","shell.execute_reply":"2023-06-21T03:00:40.844655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(os.path.join(INPUT_DIR, \"train.csv\"))\nsup_df = pd.read_csv(os.path.join(INPUT_DIR, \"supplemental_metadata.csv\"))","metadata":{"execution":{"iopub.status.busy":"2023-06-21T03:00:40.846589Z","iopub.execute_input":"2023-06-21T03:00:40.847362Z","iopub.status.idle":"2023-06-21T03:00:41.09675Z","shell.execute_reply.started":"2023-06-21T03:00:40.84733Z","shell.execute_reply":"2023-06-21T03:00:41.095913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_full_landmark_path(path):\n    return os.path.join(INPUT_DIR, path)\n\nlandmarks = pd.read_parquet(get_full_landmark_path(df.loc[0][\"path\"]))\n\nselected_columns = list(landmarks.columns[~landmarks.columns.str.startswith('z_')])\ninference_args = {\"selected_columns\": selected_columns}\n\nwith open(os.path.join(OUTPUT_DIR, \"inference_args.json\"), \"w\") as f:\n    json.dump(inference_args, f)\n\nwith open(os.path.join(OUTPUT_DIR, \"inference_args.json\")) as f:\n    inference_args = json.load(f)\n    SELECTED_COLUMNS = inference_args[\"selected_columns\"]\n\ndef get_nan_ranges(sequence):\n    \"\"\"Get list of frame indices that start or end a nan sequence.\n\n    The first index corresponds to the frame of the first missing value. The following\n    indices correspond to every bound of missing or valid value range.\n\n    Args:\n        sequence (pd.DataFrame): The sequence landmarks data.\n\n    Returns:\n        dict: Dictionary with landmarks as index and lists of frame indices as values.\n    \"\"\"\n    # keep only x since missing x => missing y\n    s = sequence[\"x\"]\n\n    # replace keypoints with landmarks since a missing keypoint => missing landmark\n    s = s.reset_index(\"keypoint\")\n    s[\"keypoint\"] = s[\"keypoint\"].str.replace(r\"_[0-9]+\", \"\", regex=True)\n    s.rename({\"keypoint\": \"landmark\"}, axis=1, inplace=True)\n    s = s.reset_index().set_index([\"landmark\", \"frame\"])\n    s = s.groupby([\"landmark\", \"frame\"]).mean()[\"x\"]\n\n    # for each landmark and range of missing value returns a df with landmark as key\n    # and start and end frame of the range as values\n    m = s.isnull()\n    nan_ranges = [\n        g.reset_index(\"frame\").groupby(\"landmark\").agg(start=(\"frame\", \"min\"), end=(\"frame\", \"max\"))\n        for _, g in s[m].groupby((~m).groupby(level=\"landmark\").cumsum())\n    ]\n\n    # extract start and end and flatten\n    nan_ranges = {\n        r.index[0]: [\n            rrr for rr in nan_ranges for rrr in (rr.start.squeeze(), rr.end.squeeze()) if rr.index[0] == r.index[0]\n        ]\n        for r in nan_ranges\n    }\n\n    return nan_ranges\n\n\ndef animate_hidden_attribute(landmark_nan_ranges, n_frames, dur):\n    \"\"\"Generate an <animate /> tag that hides the landmark when its coordinates are nan.\n\n    Args:\n        landmark_nan_ranges (list): List of frame indices that start or end a nan sequence.\n        n_frames (int): Length of the sequence.\n        dur (float): Frame duration (in seconds).\n\n    Returns:\n        str: <animate /> tag for the hidden attribute.\n    \"\"\"\n    landmark_nan_ranges = list(landmark_nan_ranges / n_frames)\n\n    v = (\";\").join([\"visible\"] + [\"hidden\", \"visible\"] * (len(landmark_nan_ranges) // 2) + [\"visible\"])\n    kt = (\";\").join([str(k) for k in [0.0] + landmark_nan_ranges + [1.0]])\n\n    return f'<animate attributeName=\"visibility\" values=\"{v}\" dur=\"{dur}s\" keyTimes=\"{kt}\" repeatCount=\"indefinite\"/>'\n\n\ndef animate_attribute(attribute_name, values, dur):\n    \"\"\"Generate an <animate /> tag for the given attribute.\n\n    Args:\n        attribute_name (str): Name of the attribute to animate.\n        values (str): Values to be taken by the attribute, separated by semicolons.\n        dur (float): Frame duration (in seconds).\n\n    Returns:\n        str: <animate /> tag for the given attribute.\n    \"\"\"\n    return f'<animate attributeName=\"{attribute_name}\" values=\"{values}\" dur=\"{dur}\" repeatCount=\"indefinite\"/>'\n\n\ndef draw_connection(x0, x1, y0, y1, landmark_nan_ranges, n_frames, fps, color=\"black\", stroke_width=0.001):\n    \"\"\"Generate a <line></line> block for a connection its animation.\n\n    Args:\n        x0 (float): x0 coordinate.\n        x1 (float): x1 coordinate.\n        y0 (float): y0 coordinate.\n        y1 (float): y1 coordinate.\n        landmark_nan_ranges (list): List of frame indices that start or end a nan sequence.\n        n_frames (int): Number of frames in the sequence.\n        fps (int): FPS used for the animation.\n        color (str, optional): Color of the connection. Defaults to \"black\".\n        stroke_width (float, optional): Width of the stroke. Defaults to 0.001.\n\n    Returns:\n        str: <line></line> block for the connection.\n    \"\"\"\n    dur = n_frames / fps\n    animation = \"\".join(\n        [\n            animate_attribute(\"x1\", \";\".join(x0.values), dur),\n            animate_attribute(\"x2\", \";\".join(x1.values), dur),\n            animate_attribute(\"y1\", \";\".join(y0.values), dur),\n            animate_attribute(\"y2\", \";\".join(y1.values), dur),\n        ]\n    )\n    set_hidden = animate_hidden_attribute(landmark_nan_ranges, n_frames, dur)\n\n    return f'<line stroke=\"{color}\" stroke-width=\"{stroke_width}\">{animation}{set_hidden}</line>'","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2023-06-21T03:00:41.097884Z","iopub.execute_input":"2023-06-21T03:00:41.09818Z","iopub.status.idle":"2023-06-21T03:00:56.326091Z","shell.execute_reply.started":"2023-06-21T03:00:41.098154Z","shell.execute_reply":"2023-06-21T03:00:56.325073Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_landmarks(path, selected_colums=SELECTED_COLUMNS):\n    return pd.read_parquet(get_full_landmark_path(path), columns=selected_colums)\n\ndef load_sequence(sequence_id, data, selected_columns=SELECTED_COLUMNS):\n    path = data.loc[data[\"sequence_id\"] == sequence_id, \"path\"].squeeze()\n    landmarks = load_landmarks(path, selected_columns)\n    return landmarks.loc[sequence_id]\n\ndef visualise_sequence(sequence_id, data, selected_columns=None, fps=10):\n    \"\"\"Generate an animated SVG file for visualising the sequence.\n\n    Args:\n        sequence_id (int): Sequence ID.\n        data (pd.DataFrame): Metadata from where to find the path to the landmarks file.\n        data_dir (str): Path to the data directory.\n        selected_columns (list, optional): Subset of columns to be loaded. Defaults to None.\n        fps (int, optional): FPS used for the animation. Defaults to 10.\n\n    Returns:\n        HTML: SVG animation of the sequence.\n    \"\"\"\n    if selected_columns is not None and \"frame\" not in selected_columns:\n        selected_columns.append(\"frame\")\n\n    phrase = data.loc[data[\"sequence_id\"] == sequence_id, \"phrase\"].squeeze()\n    sequence = load_sequence(sequence_id, data, selected_columns)\n    sequence = sequence.loc[:, ~sequence.columns.str.startswith(\"z_\")]\n\n    sequence.sort_values(\"frame\")\n\n    n_frames = sequence[\"frame\"].max()\n\n    sequence_m = sequence.melt(\"frame\").rename({\"variable\": \"keypoint\"}, axis=1)\n    sequence_m.set_index(\"keypoint\", inplace=True)\n    sequence_m.index = sequence_m.index.str.split(pat=\"_\", n=1, expand=True)\n    sequence_m.reset_index(inplace=True)\n    sequence_m.rename({\"level_0\": \"axis\", \"level_1\": \"keypoint\"}, axis=1, inplace=True)\n    mux = pd.MultiIndex.from_product([sequence_m[\"keypoint\"].unique(), range(n_frames)], names=[\"keypoint\", \"frame\"])\n    sequence_m = sequence_m.pivot(index=[\"frame\", \"keypoint\"], columns=\"axis\").reorder_levels([\"keypoint\", \"frame\"])\n    sequence_m = sequence_m.reindex(mux)\n    sequence_m.sort_index(level=[\"keypoint\", \"frame\"], ascending=[1, 1], inplace=True)\n    sequence_m.columns = [\"x\", \"y\"]\n\n    min_x = np.nanmin(sequence_m[\"x\"])\n    max_x = np.nanmax(sequence_m[\"x\"])\n    min_y = np.nanmin(sequence_m[\"y\"])\n    max_y = np.nanmax(sequence_m[\"y\"])\n    pad = 0.1\n\n    nan_ranges = get_nan_ranges(sequence_m)\n\n    sequence_str = sequence_m.groupby(level=\"keypoint\").ffill().astype(str)\n\n    lines = \"\"\n    for _, (i, j) in enumerate(POSE_CONNECTIONS):\n        try:\n            lines = lines + draw_connection(\n                sequence_str.loc[(\"pose_{}\".format(i), slice(None)), \"x\"],\n                sequence_str.loc[(\"pose_{}\".format(j), slice(None)), \"x\"],\n                sequence_str.loc[(\"pose_{}\".format(i), slice(None)), \"y\"],\n                sequence_str.loc[(\"pose_{}\".format(j), slice(None)), \"y\"],\n                nan_ranges.get(\"pose\", []),\n                n_frames,\n                fps,\n                color=\"black\",\n            )\n        except KeyError:\n            pass\n    for _, (i, j) in enumerate(FACE_CONNECTIONS):\n        try:\n            lines = lines + draw_connection(\n                sequence_str.loc[(\"face_{}\".format(i), slice(None)), \"x\"],\n                sequence_str.loc[(\"face_{}\".format(j), slice(None)), \"x\"],\n                sequence_str.loc[(\"face_{}\".format(i), slice(None)), \"y\"],\n                sequence_str.loc[(\"face_{}\".format(j), slice(None)), \"y\"],\n                nan_ranges.get(\"face\", []),\n                n_frames,\n                fps,\n                color=\"black\",\n            )\n        except KeyError:\n            pass\n    for _, (i, j) in enumerate(HAND_CONNECTIONS):\n        try:\n            lines = (\n                lines\n                + draw_connection(\n                    sequence_str.loc[(\"right_hand_{}\".format(i), slice(None)), \"x\"],\n                    sequence_str.loc[(\"right_hand_{}\".format(j), slice(None)), \"x\"],\n                    sequence_str.loc[(\"right_hand_{}\".format(i), slice(None)), \"y\"],\n                    sequence_str.loc[(\"right_hand_{}\".format(j), slice(None)), \"y\"],\n                    nan_ranges.get(\"right_hand\", []),\n                    n_frames,\n                    fps,\n                    color=\"red\",\n                    stroke_width=0.01,\n                )\n                + draw_connection(\n                    sequence_str.loc[(\"left_hand_{}\".format(i), slice(None)), \"x\"],\n                    sequence_str.loc[(\"left_hand_{}\".format(j), slice(None)), \"x\"],\n                    sequence_str.loc[(\"left_hand_{}\".format(i), slice(None)), \"y\"],\n                    sequence_str.loc[(\"left_hand_{}\".format(j), slice(None)), \"y\"],\n                    nan_ranges.get(\"left_hand\", []),\n                    n_frames,\n                    fps,\n                    color=\"red\",\n                    stroke_width=0.01,\n                )\n            )\n        except KeyError:\n            pass\n\n    html = f'<h2>{sequence_id} - \"{phrase}\"</h2>'\n    html += '<svg height=\"600\" width=\"{}\" viewBox=\"{} {} {} {}\">{}</svg>'.format(\n        600 * (max_x - min_x + 2 * pad) / (max_y - min_y + 2 * pad),\n        min_x - pad,\n        min_y - pad,\n        max_x + pad,\n        max_y + pad,\n        lines,\n    )\n\n    return HTML(html.format(lines))","metadata":{"execution":{"iopub.status.busy":"2023-06-21T03:00:56.33028Z","iopub.execute_input":"2023-06-21T03:00:56.330651Z","iopub.status.idle":"2023-06-21T03:00:56.353757Z","shell.execute_reply.started":"2023-06-21T03:00:56.330622Z","shell.execute_reply":"2023-06-21T03:00:56.352592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i, row in df.head(3).iterrows():\n    display(visualise_sequence(row.sequence_id, df))","metadata":{"execution":{"iopub.status.busy":"2023-06-21T03:00:56.355099Z","iopub.execute_input":"2023-06-21T03:00:56.355483Z","iopub.status.idle":"2023-06-21T03:01:07.015408Z","shell.execute_reply.started":"2023-06-21T03:00:56.355451Z","shell.execute_reply":"2023-06-21T03:01:07.013934Z"},"trusted":true},"execution_count":null,"outputs":[]}]}