{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceType":"competition","sourceId":127283,"databundleVersionId":15634477}],"dockerImageVersionId":31287,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install torch_geometric","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T06:29:37.522999Z","iopub.execute_input":"2026-03-12T06:29:37.523275Z","iopub.status.idle":"2026-03-12T06:29:42.962623Z","shell.execute_reply.started":"2026-03-12T06:29:37.523250Z","shell.execute_reply":"2026-03-12T06:29:42.961914Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Traffic Accident Prediction with ST-GCN++\n\nThis project implements a **multitask spatio-temporal graph convolutional network (ST-GCN++)** for traffic accident prediction and classification using only **vehicle bounding box annotations** from video data.\n\n- **Graph Representation:** Each vehicle is a node with normalized position, motion, and type (car, truck, bike), forming a temporal graph across frames.\n- **Multitask Learning:** The model jointly predicts:\n  - **Accident time and location** (regression)\n  - **Accident type** (classification)\n- **Efficiency:** Operates solely on structured bounding-box data without relying on RGB frames, making it lightweight and scalable for large datasets.\n\nThis approach provides a structured, temporal understanding of traffic interactions to accurately anticipate accidents and their characteristics.","metadata":{}},{"cell_type":"code","source":"import os\nimport json\nimport pandas as pd\nimport torch\nimport numpy as np\nfrom torch_geometric.data import Data, Dataset\n\n\nclass AccidentGraphDataset(Dataset):\n    def __init__(\n        self,\n        labels_csv=\"/kaggle/input/competitions/accident/sim_dataset/labels.csv\",\n        annotations_dir=\"/kaggle/input/competitions/accident/sim_dataset/video_annotations\",\n        num_frames=50,\n        max_vehicles=20,\n        limit_per_class=None\n    ):\n\n        self.labels_path = labels_csv\n        self.anno_dir = annotations_dir\n        self.labels_df = pd.read_csv(self.labels_path)\n\n        # Fixed Spatio-Temporal Dimensions\n        self.num_frames = num_frames\n        self.max_vehicles = max_vehicles\n\n        # Vehicle one-hot mapping\n        self.veh_tag_map = {\n            14: [1, 0, 0],  # car\n            15: [0, 1, 0],  # truck\n            18: [0, 0, 1]   # bike\n        }\n\n        # Accident type mapping\n        self.type_map = {\n            \"sideswipe\": 0,\n            \"rear-end\": 1,\n            \"single\": 2,\n            \"head-on\": 3,\n            \"t-bone\": 4\n        }\n\n        # Normalization constants\n        self.ref_width = 1920\n        self.ref_height = 1080\n        self.ref_duration = 30.0\n        self.ref_max_frames = 500\n\n        if limit_per_class:\n            self.labels_df = (\n                self.labels_df.groupby(\"type\")\n                .head(limit_per_class)\n                .reset_index(drop=True)\n            )\n\n        super().__init__(None)\n\n    def len(self):\n        return len(self.labels_df)\n\n    def get(self, idx):\n\n        row = self.labels_df.iloc[idx]\n\n        u = torch.tensor([\n            row[\"width\"] / self.ref_width,\n            row[\"height\"] / self.ref_height,\n            row[\"duration\"] / self.ref_duration,\n            row[\"no_frames\"] / self.ref_max_frames\n        ], dtype=torch.float).view(1, -1)\n\n        video_name = os.path.basename(row[\"rgb_path\"]).replace(\".mp4\", \"\")\n\n        anno_path = os.path.join(\n            self.anno_dir,\n            f\"{video_name}.json\",\n            f\"{video_name}.json\"\n        )\n\n        node_tensor = np.zeros(\n            (self.num_frames, self.max_vehicles, 11),\n            dtype=np.float32\n        )\n\n        y_targets = torch.tensor([\n            row[\"accident_time\"],\n            row[\"center_x\"],\n            row[\"center_y\"]\n        ], dtype=torch.float)\n\n        y_type = torch.tensor(\n            [self.type_map[row[\"type\"]]],\n            dtype=torch.long\n        )\n\n        if not os.path.exists(anno_path):\n            return Data(\n                x=torch.from_numpy(node_tensor).view(-1, 11),\n                u=u,\n                y=y_targets,\n                y_type=y_type\n            )\n\n        with open(anno_path, \"r\") as f:\n            gt_data = json.load(f)\n\n        base_frames = gt_data[\"base\"]\n\n        if len(base_frames) == 0:\n            return Data(\n                x=torch.from_numpy(node_tensor).view(-1, 11),\n                u=u,\n                y=y_targets,\n                y_type=y_type\n            )\n\n        # Temporal uniform sampling\n        inds = np.linspace(\n            0,\n            len(base_frames) - 1,\n            self.num_frames\n        ).astype(int)\n\n        sampled_frames = [base_frames[i] for i in inds]\n\n        min_iter = min([item[\"iteration\"] for item in base_frames])\n\n        # Find top vehicles\n        unique_ids = {}\n\n        for item in sampled_frames:\n            for obj in item[\"objects\"]:\n                unique_ids[obj[\"id\"]] = unique_ids.get(obj[\"id\"], 0) + 1\n\n        top_ids = sorted(\n            unique_ids.items(),\n            key=lambda x: x[1],\n            reverse=True\n        )[:self.max_vehicles]\n\n        id_to_idx = {\n            obj_id: i\n            for i, (obj_id, _) in enumerate(top_ids)\n        }\n\n        id_to_prev_center = {}\n\n        for t, item in enumerate(sampled_frames):\n\n            frame_idx = item[\"iteration\"] - min_iter\n\n            rel_time = (\n                frame_idx / row[\"no_frames\"]\n                if row[\"no_frames\"] > 0\n                else 0.0\n            )\n\n            for obj in item[\"objects\"]:\n\n                if obj[\"id\"] not in id_to_idx:\n                    continue\n\n                v_idx = id_to_idx[obj[\"id\"]]\n\n                bbox = obj[\"2d_bbox\"]\n\n                x1 = bbox[0][0] / row[\"width\"]\n                y1 = bbox[0][1] / row[\"height\"]\n                x2 = bbox[1][0] / row[\"width\"]\n                y2 = bbox[1][1] / row[\"height\"]\n\n                cx = (x1 + x2) / 2\n                cy = (y1 + y2) / 2\n\n                veh_type = self.veh_tag_map.get(\n                    obj[\"tag\"],\n                    [0, 0, 0]\n                )\n\n                vx = vy = speed = 0.0\n\n                if obj[\"id\"] in id_to_prev_center:\n                    px, py = id_to_prev_center[obj[\"id\"]]\n                    vx = cx - px\n                    vy = cy - py\n                    speed = np.sqrt(vx**2 + vy**2)\n\n                id_to_prev_center[obj[\"id\"]] = (cx, cy)\n\n                node_tensor[t, v_idx, :] = [\n                    x1, y1, x2, y2,\n                    rel_time,\n                    vx, vy, speed,\n                    *veh_type\n                ]\n\n        x = torch.from_numpy(node_tensor).view(-1, 11)\n\n        return Data(\n            x=x,\n            edge_index=torch.empty((2, 0), dtype=torch.long),\n            u=u,\n            y=y_targets,\n            y_type=y_type\n        )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T06:29:42.964120Z","iopub.execute_input":"2026-03-12T06:29:42.964341Z","iopub.status.idle":"2026-03-12T06:30:01.636670Z","shell.execute_reply.started":"2026-03-12T06:29:42.964316Z","shell.execute_reply":"2026-03-12T06:30:01.636131Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nfrom torch_geometric.loader import DataLoader\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import Subset\nfrom tqdm import tqdm\n\n# ==========================================\n# 1. TRAFFIC GRAPH DEFINITION\n# ==========================================\nclass TrafficGraph:\n    def __init__(self, num_node=20):\n        self.num_node = num_node\n        self.A = self.get_adjacency_matrix()\n\n    def get_adjacency_matrix(self):\n        adj = np.ones((self.num_node, self.num_node))\n        Dl = np.sum(adj, axis=0)\n\n        Dn = np.zeros((self.num_node, self.num_node))\n\n        for i in range(self.num_node):\n            if Dl[i] > 0:\n                Dn[i, i] = Dl[i]**(-1)\n\n        AD = np.dot(adj, Dn)\n\n        return np.expand_dims(AD, 0)\n\n# ==========================================\n# 2. ST-GCN++ BLOCKS\n# ==========================================\nclass STGCNPlusPlusSpatial(nn.Module):\n    def __init__(self, in_channels, out_channels, A):\n        super().__init__()\n\n        self.PA = nn.Parameter(torch.from_numpy(A.astype(np.float32)))\n        self.num_subset = A.shape[0]\n\n        self.conv = nn.Conv2d(\n            in_channels,\n            out_channels * self.num_subset,\n            kernel_size=1\n        )\n\n        if in_channels != out_channels:\n            self.down = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, 1),\n                nn.BatchNorm2d(out_channels)\n            )\n        else:\n            self.down = nn.Identity()\n\n        self.bn = nn.BatchNorm2d(out_channels)\n        self.relu = nn.ReLU()\n\n    def forward(self, x):\n\n        N, C, T, V = x.size()\n\n        x_conv = self.conv(x).view(\n            N,\n            self.num_subset,\n            -1,\n            T,\n            V\n        )\n\n        y = torch.einsum(\n            'nkctv,kvw->nctw',\n            x_conv,\n            self.PA\n        )\n\n        return self.relu(self.bn(y) + self.down(x))\n\n\nclass MSTCN(nn.Module):\n    def __init__(self, in_channels, out_channels, stride=1):\n\n        super().__init__()\n\n        self.branch1 = nn.Conv2d(in_channels, out_channels, 1)\n\n        self.branch2 = nn.Sequential(\n            nn.MaxPool2d((3, 1), stride=(stride, 1), padding=(1, 0)),\n            nn.BatchNorm2d(in_channels),\n            nn.ReLU(),\n            nn.Conv2d(in_channels, out_channels, 1)\n        )\n\n        mid_c = out_channels // 4\n\n        self.transform = nn.Conv2d(in_channels, mid_c * 4, 1)\n\n        self.branch_convs = nn.ModuleList()\n\n        for d in [1, 2, 3, 4]:\n            self.branch_convs.append(\n                nn.Sequential(\n                    nn.Conv2d(\n                        mid_c,\n                        mid_c,\n                        kernel_size=(3,1),\n                        stride=(stride,1),\n                        padding=((3+(3-1)*(d-1))//2,0),\n                        dilation=(d,1)\n                    ),\n                    nn.BatchNorm2d(mid_c),\n                    nn.ReLU()\n                )\n            )\n\n        self.agg = nn.Conv2d(\n            out_channels * 2 + mid_c * 4,\n            out_channels,\n            1\n        )\n\n        self.bn = nn.BatchNorm2d(out_channels)\n\n        if in_channels != out_channels or stride != 1:\n            self.residual = nn.Sequential(\n                nn.Conv2d(\n                    in_channels,\n                    out_channels,\n                    1,\n                    stride=(stride,1)\n                ),\n                nn.BatchNorm2d(out_channels)\n            )\n        else:\n            self.residual = nn.Identity()\n\n        self.relu = nn.ReLU()\n\n    def forward(self, x):\n\n        if self.branch2[0].stride[0] > 1:\n            b1 = F.avg_pool2d(\n                self.branch1(x),\n                kernel_size=(3,1),\n                stride=(2,1),\n                padding=(1,0)\n            )\n        else:\n            b1 = self.branch1(x)\n\n        b2 = self.branch2(x)\n\n        x_splits = torch.chunk(self.transform(x), 4, dim=1)\n\n        branches = [\n            conv(split)\n            for conv, split in zip(self.branch_convs, x_splits)\n        ]\n\n        out = torch.cat([b1, b2] + branches, dim=1)\n\n        out = self.bn(self.agg(out))\n\n        return self.relu(out + self.residual(x))\n\n\nclass STGCNPlusPlusBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, A, stride=1):\n\n        super().__init__()\n\n        self.gcn = STGCNPlusPlusSpatial(in_channels, out_channels, A)\n\n        self.tcn = MSTCN(out_channels, out_channels, stride=stride)\n\n    def forward(self, x):\n\n        return self.tcn(self.gcn(x))\n\n\n# ==========================================\n# 3. MULTITASK MODEL\n# ==========================================\nclass AccidentSTGCNMultitask(nn.Module):\n\n    def __init__(\n        self,\n        in_channels=11,\n        num_point=20,\n        num_frame=50,\n        base_channels=64,\n        u_dim=4,\n        num_classes=5\n    ):\n\n        super().__init__()\n\n        self.num_point = num_point\n        self.num_frame = num_frame\n\n        self.A = TrafficGraph(num_node=num_point).A\n\n        self.data_bn = nn.BatchNorm1d(in_channels * num_point)\n\n        self.layers = nn.ModuleList([\n            STGCNPlusPlusBlock(in_channels, base_channels, self.A, 1),\n            STGCNPlusPlusBlock(base_channels, base_channels, self.A, 1),\n            STGCNPlusPlusBlock(base_channels, base_channels*2, self.A, 2),\n            STGCNPlusPlusBlock(base_channels*2, base_channels*2, self.A, 1),\n            STGCNPlusPlusBlock(base_channels*2, base_channels*4, self.A, 2),\n            STGCNPlusPlusBlock(base_channels*4, base_channels*4, self.A, 1)\n        ])\n\n        self.regressor = nn.Sequential(\n            nn.Linear(base_channels*4 + u_dim, 128),\n            nn.ReLU(),\n            nn.Linear(128, 3)\n        )\n\n        self.classifier = nn.Sequential(\n            nn.Linear(base_channels*4 + u_dim, 128),\n            nn.ReLU(),\n            nn.Linear(128, num_classes)\n        )\n\n    def forward(self, data):\n\n        x = data.x\n        u = data.u\n        batch_size = data.num_graphs\n\n        V = self.num_point\n        T = self.num_frame\n        C = x.shape[1]\n\n        x = x.view(batch_size, T, V, C)\n\n        x = x.permute(0, 3, 1, 2).contiguous()\n\n        N, C_dim, T_dim, V_dim = x.size()\n\n        x = x.permute(0,1,3,2).contiguous().view(N, C_dim*V_dim, T_dim)\n\n        x = self.data_bn(x)\n\n        x = x.view(N, C_dim, V_dim, T_dim)\n\n        x = x.permute(0,1,3,2).contiguous()\n\n        for layer in self.layers:\n            x = layer(x)\n\n        x = F.avg_pool2d(x, kernel_size=x.size()[2:]).view(N, -1)\n\n        feat = torch.cat([x, u], dim=1)\n\n        return self.regressor(feat), self.classifier(feat)\n\n\n# ==========================================\n# 4. METRICS\n# ==========================================\ndef evaluate_metrics(reg_preds, reg_targets, cls_preds, cls_targets, sigma_t=1.0, sigma_s=0.1):\n\n    T = torch.exp(-((reg_preds[:,0]-reg_targets[:,0])**2)/(2*sigma_t**2))\n\n    S = torch.exp(-((reg_preds[:,1]-reg_targets[:,1])**2 +\n                    (reg_preds[:,2]-reg_targets[:,2])**2)/(2*sigma_s**2))\n\n    correct = (cls_preds.argmax(dim=1) == cls_targets).float().mean().item()\n\n    T_avg = max(T.mean().item(), 1e-6)\n    S_avg = max(S.mean().item(), 1e-6)\n    C_avg = max(correct, 1e-6)\n\n    harmonic_mean = 3/(1/T_avg + 1/S_avg + 1/C_avg)\n\n    return harmonic_mean, T_avg, S_avg, correct\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T06:30:01.637468Z","iopub.execute_input":"2026-03-12T06:30:01.637916Z","iopub.status.idle":"2026-03-12T06:30:02.520298Z","shell.execute_reply.started":"2026-03-12T06:30:01.637887Z","shell.execute_reply":"2026-03-12T06:30:02.519674Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ==========================================\n# 5. TRAINING\n# ==========================================\n# ==========================================\n# 5. TRAINING WITH VALIDATION\n# ==========================================\ndef train():\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n    print(\"Training Multitask ST-GCN++ on\", device)\n\n    dataset = AccidentGraphDataset(num_frames=50, max_vehicles=20)\n\n    indices = list(range(len(dataset)))\n    labels = [dataset.get(i).y_type.item() for i in indices]\n\n    train_idx, val_idx = train_test_split(\n        indices, test_size=0.20, stratify=labels, random_state=42\n    )\n\n    train_dataset = Subset(dataset, train_idx)\n    val_dataset = Subset(dataset, val_idx)\n\n    train_loader = DataLoader(\n        train_dataset, batch_size=8, shuffle=True, num_workers=4\n    )\n    val_loader = DataLoader(\n        val_dataset, batch_size=8, shuffle=False, num_workers=4\n    )\n\n    model = AccidentSTGCNMultitask().to(device)\n    optimizer = torch.optim.Adam(model.parameters(), lr=0.001)\n    best_score = 0.0\n\n    for epoch in range(1, 150):\n        # ------------------ TRAIN ------------------\n        model.train()\n        train_loss = 0\n        train_reg_preds, train_reg_targets = [], []\n        train_cls_preds, train_cls_targets = [], []\n\n        pbar = tqdm(train_loader, desc=f\"Epoch {epoch} [Train]\")\n        for data in pbar:\n            data = data.to(device)\n            optimizer.zero_grad()\n\n            reg_out, cls_out = model(data)\n            y = data.y.view(-1, 3)\n\n            loss_t = F.mse_loss(reg_out[:, 0], y[:, 0])\n            loss_s = F.mse_loss(reg_out[:, 1:], y[:, 1:])\n            loss_cls = F.cross_entropy(cls_out, data.y_type.view(-1))\n\n            loss = loss_t + 20.0 * loss_s + loss_cls\n            loss.backward()\n            optimizer.step()\n\n            train_loss += loss.item()\n            train_reg_preds.append(reg_out.detach().cpu())\n            train_reg_targets.append(y.cpu())\n            train_cls_preds.append(cls_out.detach().cpu())\n            train_cls_targets.append(data.y_type.view(-1).cpu())\n\n            pbar.set_postfix(loss=f\"{loss.item():.4f}\")\n\n        tr_score, tr_T, tr_S, tr_acc = evaluate_metrics(\n            torch.cat(train_reg_preds),\n            torch.cat(train_reg_targets),\n            torch.cat(train_cls_preds),\n            torch.cat(train_cls_targets)\n        )\n\n        # ------------------ VALIDATION ------------------\n        model.eval()\n        val_loss = 0\n        val_reg_preds, val_reg_targets = [], []\n        val_cls_preds, val_cls_targets = [], []\n\n        with torch.no_grad():\n            for data in val_loader:\n                data = data.to(device)\n                reg_out, cls_out = model(data)\n                y = data.y.view(-1, 3)\n\n                loss_t = F.mse_loss(reg_out[:, 0], y[:, 0])\n                loss_s = F.mse_loss(reg_out[:, 1:], y[:, 1:])\n                loss_cls = F.cross_entropy(cls_out, data.y_type.view(-1))\n                val_loss += (loss_t + 20.0 * loss_s + loss_cls).item()\n\n                val_reg_preds.append(reg_out.cpu())\n                val_reg_targets.append(y.cpu())\n                val_cls_preds.append(cls_out.cpu())\n                val_cls_targets.append(data.y_type.view(-1).cpu())\n\n        val_score, val_T, val_S, val_acc = evaluate_metrics(\n            torch.cat(val_reg_preds),\n            torch.cat(val_reg_targets),\n            torch.cat(val_cls_preds),\n            torch.cat(val_cls_targets)\n        )\n\n        # ------------------ SUMMARY ------------------\n        print(f\"\\nEpoch {epoch:02d} Summary\")\n        print(f\"TRAIN | Loss: {train_loss/len(train_loader):.4f} | \"\n              f\"T: {tr_T:.4f} | S: {tr_S:.4f} | CLS Acc: {tr_acc*100:.2f}% | Harmonic Score: {tr_score:.4f}\")\n        print(f\"VAL   | Loss: {val_loss/len(val_loader):.4f} | \"\n              f\"T: {val_T:.4f} | S: {val_S:.4f} | CLS Acc: {val_acc*100:.2f}% | Harmonic Score: {val_score:.4f}\")\n\n        if val_score > best_score:\n            best_score = val_score\n            torch.save(model.state_dict(), \"stgcn_multitask_best.pth\")\n            print(f\"⭐ New Best Model Saved: {best_score:.4f}\")\n        print(\"-\" * 50 + \"\\n\")\n\n\nif __name__ == \"__main__\":\n    train()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-03-12T06:30:02.521655Z","iopub.execute_input":"2026-03-12T06:30:02.522093Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}