{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":35779,"databundleVersionId":3604062,"sourceType":"competition"}],"dockerImageVersionId":30191,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"GSD Challenge 2022(prediction)","metadata":{}},{"cell_type":"code","source":"import os\nimport pandas as pd\nimport numpy as np\nimport glob\nfrom dataclasses import dataclass\nfrom scipy.interpolate import InterpolatedUnivariateSpline\nimport matplotlib.pyplot as plt\nfrom matplotlib_venn import venn2, venn2_circles\nimport seaborn as sns\nfrom tqdm.notebook import tqdm\nimport pathlib\nimport plotly\nimport plotly.express as px\npd.set_option(\"max_columns\", 500)","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:31:30.708149Z","iopub.execute_input":"2025-03-23T17:31:30.708421Z","iopub.status.idle":"2025-03-23T17:31:33.825747Z","shell.execute_reply.started":"2025-03-23T17:31:30.708340Z","shell.execute_reply":"2025-03-23T17:31:33.824707Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trip_id = \"2020-05-15-US-MTV-1/GooglePixel4XL\"","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:31:33.827386Z","iopub.execute_input":"2025-03-23T17:31:33.827591Z","iopub.status.idle":"2025-03-23T17:31:33.831822Z","shell.execute_reply.started":"2025-03-23T17:31:33.827564Z","shell.execute_reply":"2025-03-23T17:31:33.830947Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"gt = pd.read_csv(\n    \"../input/smartphone-decimeter-2022/train/2020-05-15-US-MTV-1/GooglePixel4XL/ground_truth.csv\"\n)\ngnss = pd.read_csv(\n    \"../input/smartphone-decimeter-2022/train/2020-05-15-US-MTV-1/GooglePixel4XL/device_gnss.csv\"\n)\nimu = pd.read_csv(\n    \"../input/smartphone-decimeter-2022/train/2020-05-15-US-MTV-1/GooglePixel4XL/device_imu.csv\"\n)","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:31:33.832951Z","iopub.execute_input":"2025-03-23T17:31:33.833187Z","iopub.status.idle":"2025-03-23T17:31:36.077030Z","shell.execute_reply.started":"2025-03-23T17:31:33.833154Z","shell.execute_reply":"2025-03-23T17:31:36.076248Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"INPUT_PATH = \"../input/smartphone-decimeter-2022\"\n\nWGS84_SEMI_MAJOR_AXIS = 6378137.0\nWGS84_SEMI_MINOR_AXIS = 6356752.314245\nWGS84_SQUARED_FIRST_ECCENTRICITY = 6.69437999013e-3\nWGS84_SQUARED_SECOND_ECCENTRICITY = 6.73949674226e-3\n\nHAVERSINE_RADIUS = 6_371_000\n\n\n@dataclass\nclass ECEF:\n    x: np.array\n    y: np.array\n    z: np.array\n\n    def to_numpy(self):\n        return np.stack([self.x, self.y, self.z], axis=0)\n\n    @staticmethod\n    def from_numpy(pos):\n        x, y, z = [np.squeeze(w) for w in np.split(pos, 3, axis=-1)]\n        return ECEF(x=x, y=y, z=z)\n\n\n@dataclass\nclass BLH:\n    lat: np.array\n    lng: np.array\n    hgt: np.array\n\n\ndef ECEF_to_BLH(ecef):\n    a = WGS84_SEMI_MAJOR_AXIS\n    b = WGS84_SEMI_MINOR_AXIS\n    e2 = WGS84_SQUARED_FIRST_ECCENTRICITY\n    e2_ = WGS84_SQUARED_SECOND_ECCENTRICITY\n    x = ecef.x\n    y = ecef.y\n    z = ecef.z\n    r = np.sqrt(x**2 + y**2)\n    t = np.arctan2(z * (a / b), r)\n    B = np.arctan2(z + (e2_ * b) * np.sin(t) ** 3, r - (e2 * a) * np.cos(t) ** 3)\n    L = np.arctan2(y, x)\n    n = a / np.sqrt(1 - e2 * np.sin(B) ** 2)\n    H = (r / np.cos(B)) - n\n    return BLH(lat=B, lng=L, hgt=H)\n\n\ndef haversine_distance(blh_1, blh_2):\n    dlat = blh_2.lat - blh_1.lat\n    dlng = blh_2.lng - blh_1.lng\n    a = (\n        np.sin(dlat / 2) ** 2\n        + np.cos(blh_1.lat) * np.cos(blh_2.lat) * np.sin(dlng / 2) ** 2\n    )\n    dist = 2 * HAVERSINE_RADIUS * np.arcsin(np.sqrt(a))\n    return dist\n\n\ndef pandas_haversine_distance(df1, df2):\n    blh1 = BLH(\n        lat=np.deg2rad(df1[\"LatitudeDegrees\"].to_numpy()),\n        lng=np.deg2rad(df1[\"LongitudeDegrees\"].to_numpy()),\n        hgt=0,\n    )\n    blh2 = BLH(\n        lat=np.deg2rad(df2[\"LatitudeDegrees\"].to_numpy()),\n        lng=np.deg2rad(df2[\"LongitudeDegrees\"].to_numpy()),\n        hgt=0,\n    )\n    return haversine_distance(blh1, blh2)\n\n\ndef ecef_to_lat_lng(tripID, gnss_df, UnixTimeMillis):\n    ecef_columns = [\n        \"WlsPositionXEcefMeters\",\n        \"WlsPositionYEcefMeters\",\n        \"WlsPositionZEcefMeters\",\n    ]\n    columns = [\"utcTimeMillis\"] + ecef_columns\n    ecef_df = (\n        gnss_df.drop_duplicates(subset=\"utcTimeMillis\")[columns]\n        .dropna()\n        .reset_index(drop=True)\n    )\n    ecef = ECEF.from_numpy(ecef_df[ecef_columns].to_numpy())\n    blh = ECEF_to_BLH(ecef)\n\n    TIME = ecef_df[\"utcTimeMillis\"].to_numpy()\n    lat = InterpolatedUnivariateSpline(TIME, blh.lat, ext=3)(UnixTimeMillis)\n    lng = InterpolatedUnivariateSpline(TIME, blh.lng, ext=3)(UnixTimeMillis)\n    return pd.DataFrame(\n        {\n            \"tripId\": tripID,\n            \"UnixTimeMillis\": UnixTimeMillis,\n            \"LatitudeDegrees\": np.degrees(lat),\n            \"LongitudeDegrees\": np.degrees(lng),\n        }\n    )\n\n\ndef calc_score(tripID, pred_df, gt_df):\n    d = pandas_haversine_distance(pred_df, gt_df)\n    score = np.mean([np.quantile(d, 0.50), np.quantile(d, 0.95)])\n    return score","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:31:36.078381Z","iopub.execute_input":"2025-03-23T17:31:36.078564Z","iopub.status.idle":"2025-03-23T17:31:36.094318Z","shell.execute_reply.started":"2025-03-23T17:31:36.078542Z","shell.execute_reply":"2025-03-23T17:31:36.093454Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"ss = pd.read_csv(\"../input/smartphone-decimeter-2022/sample_submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:31:36.095432Z","iopub.execute_input":"2025-03-23T17:31:36.095963Z","iopub.status.idle":"2025-03-23T17:31:36.199327Z","shell.execute_reply.started":"2025-03-23T17:31:36.095928Z","shell.execute_reply":"2025-03-23T17:31:36.198561Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"trip_id = \"2020-05-15-US-MTV-1/GooglePixel4XL\"\nbaseline = ecef_to_lat_lng(trip_id, gnss, gt[\"UnixTimeMillis\"].values)","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:31:36.200433Z","iopub.execute_input":"2025-03-23T17:31:36.201085Z","iopub.status.idle":"2025-03-23T17:31:36.228609Z","shell.execute_reply.started":"2025-03-23T17:31:36.201048Z","shell.execute_reply":"2025-03-23T17:31:36.228124Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_traffic(\n    df,\n    lat_col=\"LatitudeDegrees\",\n    lon_col=\"LongitudeDegrees\",\n    center=None,\n    color_col=\"phone\",\n    label_col=\"tripId\",\n    zoom=9,\n    opacity=1,\n):\n    if center is None:\n        center = {\n            \"lat\": df[lat_col].mean(),\n            \"lon\": df[lon_col].mean(),\n        }\n    fig = px.scatter_mapbox(\n        df,\n        # Here, plotly gets, (x,y) coordinates\n        lat=lat_col,\n        lon=lon_col,\n        # Here, plotly detects color of series\n        color=color_col,\n        labels=label_col,\n        zoom=zoom,\n        center=center,\n        height=600,\n        width=800,\n        opacity=0.5,\n    )\n    fig.update_layout(mapbox_style=\"stamen-terrain\")\n    fig.update_layout(margin={\"r\": 0, \"t\": 0, \"l\": 0, \"b\": 0})\n    fig.update_layout(title_text=\"GPS trafic\")\n    fig.show()\n\n\ndef plot_gt_vs_baseline(tripId):\n    \"\"\"\n    Create a plot of the baseline predictions vs. the ground truth\n    for a given tripId\n    \"\"\"\n    # Pull Data for an example phone\n    gt = pd.read_csv(\n        f\"../input/smartphone-decimeter-2022/train/{tripId}/ground_truth.csv\"\n    )\n    gnss = pd.read_csv(\n        f\"../input/smartphone-decimeter-2022/train/{tripId}/device_gnss.csv\"\n    )\n    imu = pd.read_csv(\n        f\"../input/smartphone-decimeter-2022/train/{tripId}/device_imu.csv\"\n    )\n    baseline = ecef_to_lat_lng(trip_id, gnss, gt[\"UnixTimeMillis\"].values)\n    # Combine ground truth with baseline predictions\n    baseline[\"isGT\"] = False\n    gt[\"isGT\"] = True\n    gt[\"tripId\"] = tripId\n\n    combined = (\n        pd.concat([baseline, gt[baseline.columns]], axis=0)\n        .reset_index(drop=True)\n        .copy()\n    )\n\n    # Plotting the route\n    visualize_traffic(\n        combined,\n        lat_col=\"LatitudeDegrees\",\n        lon_col=\"LongitudeDegrees\",\n        color_col=\"isGT\",\n        zoom=10,\n    )","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:31:36.230371Z","iopub.execute_input":"2025-03-23T17:31:36.230570Z","iopub.status.idle":"2025-03-23T17:31:36.239677Z","shell.execute_reply.started":"2025-03-23T17:31:36.230547Z","shell.execute_reply":"2025-03-23T17:31:36.238934Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from glob import glob\n\ntrain_gts = glob(\"../input/smartphone-decimeter-2022/train/*/*/ground_truth.csv\")\ntrip_ids = [\"/\".join(p.split(\"/\")[-3:-1]) for p in train_gts]","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:31:36.240615Z","iopub.execute_input":"2025-03-23T17:31:36.240848Z","iopub.status.idle":"2025-03-23T17:31:37.063920Z","shell.execute_reply.started":"2025-03-23T17:31:36.240816Z","shell.execute_reply":"2025-03-23T17:31:37.063348Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tripId = trip_ids[10]\n# Pull Data for an example phone\ngt = pd.read_csv(f\"../input/smartphone-decimeter-2022/train/{tripId}/ground_truth.csv\")\ngnss = pd.read_csv(f\"../input/smartphone-decimeter-2022/train/{tripId}/device_gnss.csv\")\nimu = pd.read_csv(f\"../input/smartphone-decimeter-2022/train/{tripId}/device_imu.csv\")\nbaseline = ecef_to_lat_lng(trip_id, gnss, gt[\"UnixTimeMillis\"].values)\n# Combine ground truth with baseline predictions\nbaseline[\"isGT\"] = False\ngt[\"isGT\"] = True\ngt[\"tripId\"] = tripId\n\ncombined = (\n    pd.concat([baseline, gt[baseline.columns]], axis=0).reset_index(drop=True).copy()\n)","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:31:37.064812Z","iopub.execute_input":"2025-03-23T17:31:37.065031Z","iopub.status.idle":"2025-03-23T17:31:38.048654Z","shell.execute_reply.started":"2025-03-23T17:31:37.065006Z","shell.execute_reply":"2025-03-23T17:31:38.047904Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"offset = 150_000\nstart = 1607640760432 + offset\ncombined.query(\"UnixTimeMillis < @start\")\n\nvisualize_traffic(combined.query(\"UnixTimeMillis < @start\"), color_col=\"isGT\", zoom=16)","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:31:38.049627Z","iopub.execute_input":"2025-03-23T17:31:38.049819Z","iopub.status.idle":"2025-03-23T17:31:38.796810Z","shell.execute_reply.started":"2025-03-23T17:31:38.049795Z","shell.execute_reply":"2025-03-23T17:31:38.796118Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob\n\nINPUT_PATH = \"../input/smartphone-decimeter-2022\"\n\nsample_df = pd.read_csv(f\"{INPUT_PATH}/sample_submission.csv\")\npred_dfs = []\nfor dirname in tqdm(sorted(glob.glob(f\"{INPUT_PATH}/test/*/*\"))):\n    drive, phone = dirname.split(\"/\")[-2:]\n    tripID = f\"{drive}/{phone}\"\n    gnss_df = pd.read_csv(f\"{dirname}/device_gnss.csv\")\n    UnixTimeMillis = sample_df[sample_df[\"tripId\"] == tripID][\n        \"UnixTimeMillis\"\n    ].to_numpy()\n    pred_dfs.append(ecef_to_lat_lng(tripID, gnss_df, UnixTimeMillis))\nsub_df = pd.concat(pred_dfs)\n\nbaselines = []\ngts = []\nfor dirname in tqdm(sorted(glob.glob(f\"{INPUT_PATH}/train/*/*\"))):\n    drive, phone = dirname.split(\"/\")[-2:]\n    tripID = f\"{drive}/{phone}\"\n    gnss_df = pd.read_csv(f\"{dirname}/device_gnss.csv\", low_memory=False)\n    gt_df = pd.read_csv(f\"{dirname}/ground_truth.csv\", low_memory=False)\n    baseline_df = ecef_to_lat_lng(tripID, gnss_df, gt_df[\"UnixTimeMillis\"].to_numpy())\n    baselines.append(baseline_df)\n    gts.append(gt_df)\nbaselines = pd.concat(baselines)\ngts = pd.concat(gts)","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:31:38.797974Z","iopub.execute_input":"2025-03-23T17:31:38.798230Z","iopub.status.idle":"2025-03-23T17:34:39.052786Z","shell.execute_reply.started":"2025-03-23T17:31:38.798195Z","shell.execute_reply":"2025-03-23T17:34:39.052026Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"baselines[\"group\"] = \"train_baseline\"\nsub_df[\"group\"] = \"submission_baseline\"\ngts[\"group\"] = \"train_ground_truth\"\ncombined = pd.concat([baselines, sub_df, gts]).reset_index(drop=True).copy()","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:34:39.054076Z","iopub.execute_input":"2025-03-23T17:34:39.054699Z","iopub.status.idle":"2025-03-23T17:34:39.209368Z","shell.execute_reply.started":"2025-03-23T17:34:39.054660Z","shell.execute_reply":"2025-03-23T17:34:39.208578Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sf_paths = combined.query(\"LatitudeDegrees > 36\").copy()\nla_paths = combined.query(\"LatitudeDegrees < 36\").copy()","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:34:39.210279Z","iopub.execute_input":"2025-03-23T17:34:39.210697Z","iopub.status.idle":"2025-03-23T17:34:39.316850Z","shell.execute_reply.started":"2025-03-23T17:34:39.210666Z","shell.execute_reply":"2025-03-23T17:34:39.316116Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calc_haversine(lat1, lon1, lat2, lon2):\n    \"\"\"Calculates the great circle distance between two points\n    on the earth. Inputs are array-like and specified in decimal degrees.\n    \"\"\"\n    RADIUS = 6_367_000\n    lat1, lon1, lat2, lon2 = map(np.radians, [lat1, lon1, lat2, lon2])\n    dlat = lat2 - lat1\n    dlon = lon2 - lon1\n    a = np.sin(dlat / 2) ** 2 + np.cos(lat1) * np.cos(lat2) * np.sin(dlon / 2) ** 2\n    dist = 2 * RADIUS * np.arcsin(a**0.5)\n    return dist\n\n\ndef add_prev_post_shift(\n    df,\n    lat_col=\"LatitudeDegrees\",\n    lng_col=\"LongitudeDegrees\",\n    dist_suffix=\"\",\n    sortby=[\"tripId\", \"UnixTimeMillis\"],\n):\n    df = df.sort_values(sortby).reset_index(drop=True)\n    df[f\"{lat_col}_shift1\"] = df.groupby([\"tripId\"])[lat_col].shift(1)\n    df[f\"{lng_col}_shift1\"] = df.groupby([\"tripId\"])[lng_col].shift(1)\n    df[f\"{lat_col}_shift-1\"] = df.groupby([\"tripId\"])[lat_col].shift(-1)\n    df[f\"{lng_col}_shift-1\"] = df.groupby([\"tripId\"])[lng_col].shift(-1)\n\n    df[f\"UnixTimeMillis_shift1\"] = df.groupby([\"tripId\"])[\"UnixTimeMillis\"].shift(1)\n    df[f\"UnixTimeMillis_shift-1\"] = df.groupby([\"tripId\"])[\"UnixTimeMillis\"].shift(-1)\n\n    df[f\"dist_prev{dist_suffix}\"] = calc_haversine(\n        df[lat_col], df[lng_col], df[f\"{lat_col}_shift1\"], df[f\"{lng_col}_shift1\"]\n    )\n    df[f\"dist_post{dist_suffix}\"] = calc_haversine(\n        df[lat_col], df[lng_col], df[f\"{lat_col}_shift-1\"], df[f\"{lng_col}_shift-1\"]\n    )\n    return df","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:34:39.317845Z","iopub.execute_input":"2025-03-23T17:34:39.318050Z","iopub.status.idle":"2025-03-23T17:34:39.327121Z","shell.execute_reply.started":"2025-03-23T17:34:39.318025Z","shell.execute_reply":"2025-03-23T17:34:39.326415Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"baselines2 = add_prev_post_shift(baselines)\nbaselines2[\"UnixTimeMillis_prev_diff\"] = (\n    baselines2[\"UnixTimeMillis\"] - baselines2[\"UnixTimeMillis_shift1\"]\n)\nbaselines2[\"speed_calc\"] = (\n    baselines2[\"dist_prev\"] / baselines2[\"UnixTimeMillis_prev_diff\"]\n)\n\n\nsub_df = add_prev_post_shift(sub_df)\nsub_df[\"UnixTimeMillis_prev_diff\"] = (\n    sub_df[\"UnixTimeMillis\"] - sub_df[\"UnixTimeMillis_shift1\"]\n)\nsub_df[\"speed_calc\"] = sub_df[\"dist_prev\"] / sub_df[\"UnixTimeMillis_prev_diff\"]","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:34:39.328216Z","iopub.execute_input":"2025-03-23T17:34:39.328417Z","iopub.status.idle":"2025-03-23T17:34:39.738324Z","shell.execute_reply.started":"2025-03-23T17:34:39.328393Z","shell.execute_reply":"2025-03-23T17:34:39.737806Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def do_postprocess(sub_df, thres=0.5):  # Lowered threshold to 0.5\n    sub = sub_df.copy()\n    for c, sub_stopped in sub.groupby(\"tripId\"):\n        sub_stopped = sub_stopped.loc[sub_stopped[\"dist_prev\"] < thres].copy()\n        sub_stopped[\"UnixTimeMillis_diff\"] = sub_stopped[\"UnixTimeMillis\"].diff()\n        sub_stopped[\"big_timeshift\"] = sub_stopped[\"UnixTimeMillis_diff\"] > 2_000\n        sub_stopped[\"time_group\"] = sub_stopped[\"big_timeshift\"].astype(\"int\").cumsum()\n\n        for stop_group, d in sub_stopped.groupby(\"time_group\"):\n            tstart, tstop = d[\"UnixTimeMillis\"].min(), d[\"UnixTimeMillis\"].max()\n            stopped_len = len(\n                sub.loc[\n                    (sub[\"UnixTimeMillis\"] >= tstart) & (sub[\"UnixTimeMillis\"] <= tstop)\n                ]\n            )\n            if stopped_len >= 20:\n                buffer = 800\n                latDegmean = sub.loc[\n                    (sub[\"UnixTimeMillis\"] >= (tstart - buffer))\n                    & (sub[\"UnixTimeMillis\"] <= (tstop + buffer))\n                ][\"LatitudeDegrees\"].mean()\n                lngDegmean = sub.loc[\n                    (sub[\"UnixTimeMillis\"] >= (tstart - buffer))\n                    & (sub[\"UnixTimeMillis\"] <= (tstop + buffer))\n                ][\"LongitudeDegrees\"].mean()\n                sub.loc[\n                    (sub[\"UnixTimeMillis\"] >= tstart)\n                    & (sub[\"UnixTimeMillis\"] <= tstop),\n                    \"LatitudeDegrees\",\n                ] = latDegmean\n                sub.loc[\n                    (sub[\"UnixTimeMillis\"] >= tstart)\n                    & (sub[\"UnixTimeMillis\"] <= tstop),\n                    \"LongitudeDegrees\",\n                ] = lngDegmean\n                sub.loc[\n                    (sub[\"UnixTimeMillis\"] >= tstart)\n                    & (sub[\"UnixTimeMillis\"] <= tstop),\n                    \"stopped\",\n                ] = True\n    sub[\"stopped\"] = sub[\"stopped\"].fillna(False)\n    return sub","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:34:39.739349Z","iopub.execute_input":"2025-03-23T17:34:39.739558Z","iopub.status.idle":"2025-03-23T17:34:39.748567Z","shell.execute_reply.started":"2025-03-23T17:34:39.739522Z","shell.execute_reply":"2025-03-23T17:34:39.747816Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub = do_postprocess(sub_df, thres=1.5)","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:34:39.749434Z","iopub.execute_input":"2025-03-23T17:34:39.749620Z","iopub.status.idle":"2025-03-23T17:34:42.075343Z","shell.execute_reply.started":"2025-03-23T17:34:39.749597Z","shell.execute_reply":"2025-03-23T17:34:42.074767Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"baselines2 = add_prev_post_shift(baselines)\nbaselines2[\"UnixTimeMillis_prev_diff\"] = (\n    baselines2[\"UnixTimeMillis\"] - baselines2[\"UnixTimeMillis_shift1\"]\n)\nbaselines2[\"speed_calc\"] = (\n    baselines2[\"dist_prev\"] / baselines2[\"UnixTimeMillis_prev_diff\"]\n)\n\nbaselines2_pp = do_postprocess(baselines2, thres=1.5)\nscores = []\nfor tripID in baselines2_pp[\"tripId\"].unique():\n    score = calc_score(tripID, baselines2_pp, gts)\n    scores.append(score)\n\nmean_score = np.mean(scores)\nprint(f\"mean_score = {mean_score:.3f}\")","metadata":{"execution":{"iopub.status.busy":"2025-03-23T17:38:44.779702Z","iopub.execute_input":"2025-03-23T17:38:44.780461Z","iopub.status.idle":"2025-03-23T17:39:10.458099Z","shell.execute_reply.started":"2025-03-23T17:38:44.780420Z","shell.execute_reply":"2025-03-23T17:39:10.457319Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\n\n# Assuming the 'baselines2' dataframe is already prepared\nbaselines2 = add_prev_post_shift(baselines)\n\n# Calculate differences and speed (these might be used for your model's predictions)\nbaselines2[\"UnixTimeMillis_prev_diff\"] = (\n    baselines2[\"UnixTimeMillis\"] - baselines2[\"UnixTimeMillis_shift1\"]\n)\nbaselines2[\"speed_calc\"] = (\n    baselines2[\"dist_prev\"] / baselines2[\"UnixTimeMillis_prev_diff\"]\n)\n\n# Post-processing (filtering, thresholding, etc.)\nbaselines2_pp = do_postprocess(baselines2, thres=1.5)\n\n# Assuming your model has been trained already, or you are predicting something\n# Here, let's assume you're predicting 'speed_calc' or another feature.\n\n# Generate predictions for each trip (you can use your model or any other prediction logic here)\npredictions = []\ntrip_ids = baselines2_pp[\"tripId\"].unique()\n\n# For each trip, calculate or generate a prediction\nfor tripID in trip_ids:\n    # Example prediction logic (you can replace this with a model prediction)\n    trip_data = baselines2_pp[baselines2_pp[\"tripId\"] == tripID]\n    predicted_speed = np.mean(trip_data[\"speed_calc\"])  # Just an example of using mean speed for prediction\n    predictions.append(predicted_speed)\n\n# Create the submission dataframe\nsubmission = pd.DataFrame({\n    'tripId': trip_ids,   # Unique trip IDs\n    'prediction': predictions  # Your predictions for each trip\n})\n\n# Save the predictions to a CSV file\nsubmission.to_csv('submission.csv', index=False)\n\n# Print the mean prediction\nmean_prediction = np.mean(predictions)\nprint(f\"mean_prediction = {mean_prediction:.3f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T17:52:29.633382Z","iopub.execute_input":"2025-03-23T17:52:29.633925Z","iopub.status.idle":"2025-03-23T17:52:55.599912Z","shell.execute_reply.started":"2025-03-23T17:52:29.633873Z","shell.execute_reply":"2025-03-23T17:52:55.599119Z"}},"outputs":[],"execution_count":null}]}