// SPDX-FileCopyrightText: 2026 Filip Leonarski, Paul Scherrer Institute // SPDX-License-Identifier: GPL-3.0-only #include "ModelMaskGPU.h" #include #include #include #include "../../common/JFJochException.h" namespace { void cuda_err(cudaError_t val) { if (val != cudaSuccess) throw JFJochException(JFJochExceptionCategory::GPUCUDAError, cudaGetErrorString(val)); } constexpr int THREADS = 256; constexpr int IMAGE_THREADS = 64; // gemmi SolventMasker(AtomicRadiiSet::Refmac): rshrink and island_min_volume. constexpr double R_SHRINK = 0.8; constexpr double ISLAND_MIN_VOLUME = 50.0; __device__ int wrap(int a, int n) { const int r = a % n; return r < 0 ? r + n : r; } __global__ void fill_solvent(float *mask, size_t n) { const size_t i = blockIdx.x * static_cast(blockDim.x) + threadIdx.x; if (i < n) mask[i] = 1.0f; } // One block per image (atom, operator): gemmi's use_points_around(fpos, radius) around the image, // setting to 0 every point of the box with !(dist_sq > radius^2). The box is sized from the interplanar // spacing, clamped to n - 1 and centred on the rounded grid position, as gemmi does it. __global__ void mask_images(const ModelMaskAtom *atoms, const ModelMaskOp *ops, int n_ops, ModelMaskGPU::Geometry g, float *mask) { const ModelMaskAtom a = atoms[blockIdx.x / n_ops]; // The image's box, one axis per thread: double arithmetic is slow on the GPU, so it is done once // per image rather than once per point. __shared__ int lo[3], len[3]; __shared__ float c[3]; // the image's grid position less the box's first point if (threadIdx.x < 3) { const int k = threadIdx.x; const ModelMaskOp &op = ops[blockIdx.x % n_ops]; const int n = k == 0 ? g.nu : k == 1 ? g.nv : g.nw; const double x = op.rot[3 * k] * a.x + op.rot[3 * k + 1] * a.y + op.rot[3 * k + 2] * a.z + op.tran[k]; const double nx = x * n; const int centre = static_cast(round(nx)); const int d = min(static_cast(ceil(a.radius / g.spacing[k])), n - 1); lo[k] = centre - d; len[k] = 2 * d + 1; c[k] = static_cast(nx - centre) + static_cast(d); } __syncthreads(); const float r2 = static_cast(a.radius * a.radius); const int total = len[0] * len[1] * len[2]; for (int t = threadIdx.x; t < total; t += blockDim.x) { const int iu = t % len[0], iv = (t / len[0]) % len[1], iw = t / (len[0] * len[1]); const float fu = c[0] - iu, fv = c[1] - iv, fw = c[2] - iw; const float x = g.orth_n[0] * fu + g.orth_n[1] * fv + g.orth_n[2] * fw; const float y = g.orth_n[3] * fu + g.orth_n[4] * fv + g.orth_n[5] * fw; const float z = g.orth_n[6] * fu + g.orth_n[7] * fv + g.orth_n[8] * fw; if (x * x + y * y + z * z > r2) continue; const size_t idx = wrap(lo[0] + iu, g.nu) + static_cast(g.nu) * (wrap(lo[1] + iv, g.nv) + static_cast(g.nv) * wrap(lo[2] + iw, g.nw)); mask[idx] = 0.0f; } } // Union-find labelling of the solvent after Playne & Hawick (2018) IEEE TPDS 29, 1217-1230: every // solvent point starts as its own set, neighbouring points are united with atomicCAS, always linking the // larger root under the smaller, so each component ends up rooted at its lowest index whatever the // order of the unions. // Path halving: a non-root entry is pointed at its grandparent. Entries only ever point to a lower // index of the same set, so the racing writes of other threads are harmless and a root is changed // only by the atomicCAS in unite(). __device__ int find_root(int *label, int i) { int p = label[i]; while (p != i) { const int gp = label[p]; if (gp != p) label[i] = gp; i = p; p = gp; } return i; } __device__ void unite(int *label, int a, int b) { while (true) { a = find_root(label, a); b = find_root(label, b); if (a == b) return; if (a < b) { const int t = a; a = b; b = t; } const int old = atomicCAS(&label[a], a, b); if (old == a) return; a = old; } } __global__ void uf_init(const float *mask, int *label, size_t n) { const size_t i = blockIdx.x * static_cast(blockDim.x) + threadIdx.x; if (i < n) label[i] = mask[i] == 1.0f ? static_cast(i) : -1; } // gemmi's FloodFill joins a run of points to the runs of the rows v +- 1, w +- 1 over u - 1 .. u + len, // with wrap-around: 26-connectivity on the periodic grid. Each point unites with the 13 neighbours // ahead of it; the other 13 unite with it from their side. __global__ void uf_merge(const float *mask, int *label, int nu, int nv, int nw) { const size_t n = static_cast(nu) * nv * nw; const size_t i = blockIdx.x * static_cast(blockDim.x) + threadIdx.x; if (i >= n || mask[i] != 1.0f) return; const int u = static_cast(i % nu); const int v = static_cast((i / nu) % nv); const int w = static_cast(i / (static_cast(nu) * nv)); const int off[13][3] = {{1, 0, 0}, {-1, 1, 0}, {0, 1, 0}, {1, 1, 0}, {-1, -1, 1}, {0, -1, 1}, {1, -1, 1}, {-1, 0, 1}, {0, 0, 1}, {1, 0, 1}, {-1, 1, 1}, {0, 1, 1}, {1, 1, 1}}; for (int k = 0; k < 13; k++) { const size_t j = wrap(u + off[k][0], nu) + static_cast(nu) * (wrap(v + off[k][1], nv) + static_cast(nv) * wrap(w + off[k][2], nw)); if (mask[j] == 1.0f) unite(label, static_cast(i), static_cast(j)); } } // The roots, read-only and into a separate array: writing them over the labels while other threads // still walk the labels would let a path-halving write replace a written root with a mere ancestor. __global__ void uf_roots(const int *label, int *root, size_t n) { const size_t i = blockIdx.x * static_cast(blockDim.x) + threadIdx.x; if (i >= n) return; int r = label[i]; if (r >= 0) while (label[r] != r) r = label[r]; root[i] = r; } // Component sizes, counted only while they are at most limit: every island is counted in full, and the // one large component stops being counted almost at once instead of queueing millions of atomics on one // address. Only size <= limit is ever read, which the capping does not change. __global__ void uf_count(const int *root, int *size, int limit, size_t n) { const size_t i = blockIdx.x * static_cast(blockDim.x) + threadIdx.x; if (i >= n) return; const int r = root[i]; if (r >= 0 && reinterpret_cast(size)[r] <= limit) atomicAdd(&size[r], 1); } __global__ void uf_remove(float *mask, const int *root, const int *size, int limit, size_t n) { const size_t i = blockIdx.x * static_cast(blockDim.x) + threadIdx.x; if (i < n && root[i] >= 0 && size[root[i]] <= limit) mask[i] = 0.0f; } unsigned int Blocks(size_t n) { return static_cast((n + THREADS - 1) / THREADS); } } // namespace size_t ModelMaskGPU::DeviceBytes(size_t max_points) { return 2 * max_points * sizeof(int) + MAX_OPS * sizeof(ModelMaskOp); } ModelMaskGPU::ModelMaskGPU(cudaStream_t stream, size_t max_points) : stream(stream), max_points(max_points), ops_d(MAX_OPS), label(max_points), root(max_points) { if (max_points > INT32_MAX) throw JFJochException(JFJochExceptionCategory::InputParameterInvalid, "Solvent mask grid too large for int32 labels"); } void ModelMaskGPU::SetGrid(const ModelMaskGrid &grid, const std::vector &ops) { npoints = static_cast(grid.nu) * grid.nv * grid.nw; if (npoints > max_points) throw JFJochException(JFJochExceptionCategory::InputParameterInvalid, "Solvent mask grid larger than the buffers"); if (ops.empty() || ops.size() > MAX_OPS) throw JFJochException(JFJochExceptionCategory::InputParameterInvalid, "Solvent mask needs 1 to " + std::to_string(MAX_OPS) + " operators"); const int n[3] = {grid.nu, grid.nv, grid.nw}; geom.nu = grid.nu; geom.nv = grid.nv; geom.nw = grid.nw; for (int i = 0; i < 3; i++) for (int j = 0; j < 3; j++) geom.orth_n[3 * i + j] = static_cast(grid.orth[3 * i + j] / n[j]); // gemmi Grid::calculate_spacing(): 1 / (n |a*|), where a* is row k of orth's inverse, which is the // cross product of the other two columns over the determinant (= volume). const double *o = grid.orth; const double col[3][3] = {{o[0], o[3], o[6]}, {o[1], o[4], o[7]}, {o[2], o[5], o[8]}}; for (int k = 0; k < 3; k++) { const double *p = col[(k + 1) % 3], *q = col[(k + 2) % 3]; const double cx = p[1] * q[2] - p[2] * q[1], cy = p[2] * q[0] - p[0] * q[2], cz = p[0] * q[1] - p[1] * q[0]; geom.spacing[k] = grid.volume / (n[k] * std::sqrt(cx * cx + cy * cy + cz * cz)); } // gemmi's shrink, set_margin_around(r_shrink), only ever looks at the grid offsets other than 0 // within r_shrink, in the box of floor(r_shrink / spacing). No such offset, no change. const int d[3] = {static_cast(std::floor(R_SHRINK / geom.spacing[0])), static_cast(std::floor(R_SHRINK / geom.spacing[1])), static_cast(std::floor(R_SHRINK / geom.spacing[2]))}; for (int w = -d[2]; w <= d[2]; w++) for (int v = -d[1]; v <= d[1]; v++) for (int u = -d[0]; u <= d[0]; u++) { const double f[3] = {static_cast(u) / n[0], static_cast(v) / n[1], static_cast(w) / n[2]}; double r2 = 0; for (int i = 0; i < 3; i++) { const double x = o[3 * i] * f[0] + o[3 * i + 1] * f[1] + o[3 * i + 2] * f[2]; r2 += x * x; } if (r2 <= R_SHRINK * R_SHRINK && r2 != 0.0) throw JFJochException(JFJochExceptionCategory::InputParameterInvalid, "Solvent mask grid too fine: the GPU mask has no shrink step"); } // gemmi SolventMasker::remove_islands(), rounded the same way. island_limit = static_cast(static_cast(ISLAND_MIN_VOLUME * npoints / grid.volume)); n_ops = static_cast(ops.size()); cuda_err(cudaMemcpyAsync(ops_d, ops.data(), ops.size() * sizeof(ModelMaskOp), cudaMemcpyHostToDevice, stream)); } void ModelMaskGPU::Compute(const ModelMaskAtom *d_atoms, int n_atoms, float *d_mask) { fill_solvent<<>>(d_mask, npoints); cuda_err(cudaGetLastError()); if (n_atoms > 0) { mask_images<<(n_atoms) * n_ops, IMAGE_THREADS, 0, stream>>>( d_atoms, ops_d, n_ops, geom, d_mask); cuda_err(cudaGetLastError()); } RemoveIslands(d_mask); } // gemmi SolventMasker::remove_islands(): every component of solvent with at most island_limit points // becomes macromolecule. void ModelMaskGPU::RemoveIslands(float *d_mask) { uf_init<<>>(d_mask, label, npoints); cuda_err(cudaGetLastError()); uf_merge<<>>(d_mask, label, geom.nu, geom.nv, geom.nw); cuda_err(cudaGetLastError()); uf_roots<<>>(label, root, npoints); cuda_err(cudaGetLastError()); int *size = label; // the parents are no longer needed cuda_err(cudaMemsetAsync(size, 0, npoints * sizeof(int), stream)); uf_count<<>>(root, size, island_limit, npoints); cuda_err(cudaGetLastError()); uf_remove<<>>(d_mask, root, size, island_limit, npoints); cuda_err(cudaGetLastError()); }