{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":117682,"databundleVersionId":15062069,"sourceType":"competition"}],"dockerImageVersionId":31239,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"For discussion, refer to:   \nhttps://www.kaggle.com/competitions/vesuvius-challenge-surface-detection/discussion/667135","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\n","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Helpers: coords + grid_sample\n\n\ndef zyx_to_grid_norm(x_zyx: torch.Tensor, vol_shape_zyx):\n    \"\"\"\n    x_zyx: (N,3) in voxel coords, z in [0,Z-1], y in [0,Y-1], x in [0,X-1]\n    returns: (1,N,1,1,3) normalized coords for grid_sample in order (x,y,z)\n    \"\"\"\n    Z, Y, X = vol_shape_zyx\n    z = x_zyx[:, 0]\n    y = x_zyx[:, 1]\n    x = x_zyx[:, 2]\n\n    # normalize to [-1,1]\n    # align_corners=True -> -1 maps to 0, +1 maps to size-1\n    x_n = 2.0 * x / (X - 1) - 1.0\n    y_n = 2.0 * y / (Y - 1) - 1.0\n    z_n = 2.0 * z / (Z - 1) - 1.0\n\n    grid = torch.stack([x_n, y_n, z_n], dim=-1)  # (N,3) in (x,y,z)\n    grid = grid.view(1, -1, 1, 1, 3)\n    return grid\n\ndef sample_velocity(u_grid_zyx: torch.Tensor, x_zyx: torch.Tensor, vol_shape_zyx):\n    \"\"\"\n    u_grid_zyx: (1,3,Zg,Yg,Xg) velocity in ZYX component order (dz,dy,dx) at grid resolution.\n    x_zyx: (N,3) points in voxel coords (ZYX) at volume resolution.\n    returns: (N,3) velocity at points in voxel units per unit time.\n    \"\"\"\n    # grid_sample expects input shape (N,C,D,H,W), grid (N, Dout, Hout, Wout, 3)\n    # We'll treat points as (Dout=N, Hout=1, Wout=1)\n    # Need normalized coords w.r.t. u_grid resolution, not volume resolution.\n    Zg, Yg, Xg = u_grid_zyx.shape[-3], u_grid_zyx.shape[-2], u_grid_zyx.shape[-1]\n\n    # Convert x_zyx (volume voxel coords) to u-grid voxel coords\n    Z, Y, X = vol_shape_zyx\n    z = x_zyx[:, 0] * (Zg - 1) / max(Z - 1, 1)\n    y = x_zyx[:, 1] * (Yg - 1) / max(Y - 1, 1)\n    x = x_zyx[:, 2] * (Xg - 1) / max(X - 1, 1)\n    xg_zyx = torch.stack([z, y, x], dim=-1)\n\n    # normalized coords for u-grid\n    grid = zyx_to_grid_norm(xg_zyx, (Zg, Yg, Xg))  # (1,N,1,1,3) in (x,y,z) normalized\n\n    # sample: output (1,3,N,1,1)\n    v = F.grid_sample(\n        u_grid_zyx, grid,\n        mode=\"bilinear\",\n        padding_mode=\"border\",\n        align_corners=True\n    )\n    v = v.view(3, -1).transpose(0, 1).contiguous()  # (N,3), still in (dz,dy,dx)\n    return v\n\ndef euler_integrate(u_grid_zyx: torch.Tensor, x0_zyx: torch.Tensor, vol_shape_zyx,\n                    n_steps=16, dt=1.0/16.0):\n    \"\"\"\n    Integrate dx/dt = u(x) with Euler steps.\n    x0_zyx: (N,3) in voxel coords\n    \"\"\"\n    x = x0_zyx\n    for _ in range(n_steps):\n        v = sample_velocity(u_grid_zyx, x, vol_shape_zyx)  # (N,3) in voxel units\n        x = x + dt * v\n    return x\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#self = canoical\n\nclass CanoicalFitter (nn.Module):\n    def __init__(self, vol_shape_zyx, grid_shape_zyx=(40,40,40),\n                 n_steps=16):\n        super().__init__()\n        self.vol_shape_zyx = tuple(vol_shape_zyx)\n        self.n_steps = n_steps\n\n        Zg, Yg, Xg = grid_shape_zyx\n        # Learnable velocity field on coarse grid: (1,3,Zg,Yg,Xg) in (dz,dy,dx)\n        self.u = nn.Parameter(torch.zeros(1, 3, Zg, Yg, Xg))\n\n        #---\n        # todo ... initialisae surface to median of data ....\n\n        \n        # Learnable sheet gap (positive)\n        self.log_g = nn.Parameter(torch.tensor(0.0))\n\n    @torch.no_grad()\n    def set_canonical_frame(self, mu_zyx, n_zyx):\n        mu_zyx = torch.as_tensor(mu_zyx, dtype=torch.float32, device=self.mu_zyx.device)\n        n_zyx  = torch.as_tensor(n_zyx, dtype=torch.float32, device=self.n_zyx.device)\n        n_zyx  = n_zyx / (n_zyx.norm() + 1e-8)\n        self.mu_zyx.copy_(mu_zyx)\n        self.n_zyx.copy_(n_zyx)\n\n    def V_to_S(self, xV_zyx):\n        \"\"\"\n        Approx inverse map by integrating with -u (paper trick).\n        \"\"\"\n        xS_zyx = euler_integrate(-self.u, xV_zyx, self.vol_shape_zyx, n_steps=self.n_steps, dt=1.0/self.n_steps)\n        return xS_zyx\n\n    def sheet_distance(self, xS_zyx):\n        \"\"\"\n        Distance of canonical-mapped points to the nearest of two parallel sheets\n        along canonical normal axis n (stored).\n        \"\"\"\n        g = torch.exp(self.log_g) + 1e-6  # gap in voxel units along n\n        # canonical coordinate w = n·(x - mu)\n        w = ((xS_zyx - self.mu_zyx) * self.n_zyx).sum(dim=-1)  # (N,)\n        d1 = (w + 0.5 * g).abs()\n        d2 = (w - 0.5 * g).abs()\n\n        # softmin is smoother than min()\n        tau = 0.25\n        d = -tau * torch.log(torch.exp(-d1/tau) + torch.exp(-d2/tau))\n        return d, w, g\n\n    def smoothness_loss(self):\n        \"\"\"\n        Simple ||∇u||^2 on the coarse grid.\n        \"\"\"\n        u = self.u\n        dz = u[:,:,1:,:,:] - u[:,:,:-1,:,:]\n        dy = u[:,:,:,1:,:] - u[:,:,:,:-1,:]\n        dx = u[:,:,:,:,1:] - u[:,:,:,:,:-1]\n        return (dz.pow(2).mean() + dy.pow(2).mean() + dx.pow(2).mean())\n\n\ndef fit_flow(\n    point_zyx,\n    canoical\n):\n\n    #canoical.parameters is the surface \n    model = canoical\n    optimizer = torch.optim.Adam(model.parameters(), lr=lr)\n\n    P = point_zyx\n\n    for it in range(iters):\n        idx = torch.randint(0, P.shape[0], (batch_size,), device=device)\n        xV = P[idx]  # (B,3)\n\n        xS = model.V_to_S(xV)\n        d, w, g = model.distance_loss(xS)\n\n        loss_data   = d.mean()\n        loss_smooth = model.smoothness_loss()\n        loss_small  = model.u.pow(2).mean()\n        loss_gap    = 1.0 / (g + 1e-3)  # prevents collapse to g->0\n\n        loss = 1.0*loss_data + 0.05*loss_smooth + 0.001*loss_small + 0.01*loss_gap\n\n        optimizer.zero_grad(set_to_none=True)\n        loss.backward()\n        optimizer.step()\n\n\n        ","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}