The GPU bulk-solvent mask had no shrink step (SolventMasker::shrink(), set_margin_around() with Refmac's r_shrink = 0.8 A), so every rigid-body zone whose grid has an offset within 0.8 A was sent to the CPU. That is not only fine grids: in an oblique setting (P2_1, beta ~141 deg) the 3.5 A zone's lattice-plane spacing is 0.69 A, and on such a set the real fit and all nine null replicates ran their rigid bodies on the CPU. The shrink is now two passes over the grid: mark the solvent points with a macromolecule point among gemmi's near offsets, then turn into solvent every macromolecule point with such an edge point at any stencil offset (near or far, gemmi's split at the coarsest grid step). That is gemmi's rule in both of its branches, and the stencil is built on the host in double exactly as gemmi builds it. ModelMaskGPU::ShrinkIsNoOp() and RigidBodyGPUEngine::MaskSupports() are gone; the rigid body no longer refuses such zones. Verified (ModelMaskGPU_ShrinkMatchesGemmi): the shrink alone on gemmi's post-island mask is bit-identical, and the whole GPU mask equals gemmi's put_mask_on_grid() bit for bit (0 differing points) on the five test groups at 1.5 A and an oblique P2_1 cell at 3.5 A; repeats are identical. Measured on the oblique-setting set of the open arm (P2_1, 14k atoms, 1.66 A): model validation 17.2 s -> 5.2 s with the GPU structure factors of the next commit in place (the first validation's ten rigid bodies, ~12 s on the CPU, now take well under a second). p.mtz md5-identical; the null moves from +73.1 to +86.1 sigma (the replicates are refined on the GPU, which agrees with the CPU to rounding), verdict unchanged. Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_01SVmAWnzCmRKAXVUCdc4iNi
372 lines
17 KiB
Plaintext
372 lines
17 KiB
Plaintext
// SPDX-FileCopyrightText: 2026 Filip Leonarski, Paul Scherrer Institute <filip.leonarski@psi.ch>
|
|
// SPDX-License-Identifier: GPL-3.0-only
|
|
|
|
// The masking, the island removal and the shrink follow SolventMasker and set_margin_around() of GEMMI's
|
|
// solmask.hpp (https://github.com/project-gemmi/gemmi/blob/master/include/gemmi/solmask.hpp)
|
|
// (c) Global Phasing Ltd., Mozilla Public License Version 2.0
|
|
|
|
#include "ModelMaskGPU.h"
|
|
|
|
#include <cmath>
|
|
#include <cstdint>
|
|
#include <string>
|
|
#include <vector>
|
|
#include <algorithm>
|
|
|
|
#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<size_t>(blockDim.x) + threadIdx.x;
|
|
if (i < n)
|
|
mask[i] = 1.0f;
|
|
}
|
|
|
|
// One block per image (atom, operator): gemmi's use_points_around<true>(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<int>(round(nx));
|
|
const int d = min(static_cast<int>(ceil(a.radius / g.spacing[k])), n - 1);
|
|
lo[k] = centre - d;
|
|
len[k] = 2 * d + 1;
|
|
c[k] = static_cast<float>(nx - centre) + static_cast<float>(d);
|
|
}
|
|
__syncthreads();
|
|
const float r2 = static_cast<float>(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<size_t>(g.nu) * (wrap(lo[1] + iv, g.nv)
|
|
+ static_cast<size_t>(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<size_t>(blockDim.x) + threadIdx.x;
|
|
if (i < n)
|
|
label[i] = mask[i] == 1.0f ? static_cast<int>(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<size_t>(nu) * nv * nw;
|
|
const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
|
|
if (i >= n || mask[i] != 1.0f)
|
|
return;
|
|
const int u = static_cast<int>(i % nu);
|
|
const int v = static_cast<int>((i / nu) % nv);
|
|
const int w = static_cast<int>(i / (static_cast<size_t>(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<size_t>(nu) * (wrap(v + off[k][1], nv)
|
|
+ static_cast<size_t>(nv) * wrap(w + off[k][2], nw));
|
|
if (mask[j] == 1.0f)
|
|
unite(label, static_cast<int>(i), static_cast<int>(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<size_t>(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<size_t>(blockDim.x) + threadIdx.x;
|
|
if (i >= n)
|
|
return;
|
|
const int r = root[i];
|
|
if (r >= 0 && reinterpret_cast<volatile int *>(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<size_t>(blockDim.x) + threadIdx.x;
|
|
if (i < n && root[i] >= 0 && size[root[i]] <= limit)
|
|
mask[i] = 0.0f;
|
|
}
|
|
|
|
// gemmi's SolventMasker::shrink(), set_margin_around(mask, r_shrink, 1, -1) followed by -1 -> 1, in two
|
|
// passes. gemmi walks the solvent points and turns into solvent every macromolecule point at a stencil
|
|
// offset from one; the offsets are split at the grid spacing into near ones (stencil1) and far ones
|
|
// (stencil2), and the far ones are only taken from a solvent point that has a macromolecule point among
|
|
// its near ones - a solvent point at the edge. (With no far offsets gemmi walks the macromolecule points
|
|
// instead, asking whether any near offset is solvent, which is the same rule.) So: first mark the edge
|
|
// points, then turn every macromolecule point with an edge point at any stencil offset into solvent.
|
|
// The new solvent points are not edge points, so the second pass can write the mask in place.
|
|
__global__ void shrink_edges(const float *mask, int *edge, const int3 *stencil, int n_near, int nu, int nv, int nw) {
|
|
const size_t n = static_cast<size_t>(nu) * nv * nw;
|
|
const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
|
|
if (i >= n)
|
|
return;
|
|
int is_edge = 0;
|
|
if (mask[i] >= 1.0f) {
|
|
const int u = static_cast<int>(i % nu);
|
|
const int v = static_cast<int>((i / nu) % nv);
|
|
const int w = static_cast<int>(i / (static_cast<size_t>(nu) * nv));
|
|
for (int k = 0; k < n_near && !is_edge; k++) {
|
|
const size_t j = wrap(u + stencil[k].x, nu)
|
|
+ static_cast<size_t>(nu) * (wrap(v + stencil[k].y, nv)
|
|
+ static_cast<size_t>(nv) * wrap(w + stencil[k].z, nw));
|
|
is_edge = mask[j] < 1.0f;
|
|
}
|
|
}
|
|
edge[i] = is_edge;
|
|
}
|
|
|
|
// The stencil is symmetric (an offset and its negative are the same distance), so a macromolecule point
|
|
// that an edge point reaches is one that reaches an edge point.
|
|
__global__ void shrink_margin(float *mask, const int *edge, const int3 *stencil, int n_stencil, int nu, int nv, int nw) {
|
|
const size_t n = static_cast<size_t>(nu) * nv * nw;
|
|
const size_t i = blockIdx.x * static_cast<size_t>(blockDim.x) + threadIdx.x;
|
|
if (i >= n || mask[i] >= 1.0f)
|
|
return;
|
|
const int u = static_cast<int>(i % nu);
|
|
const int v = static_cast<int>((i / nu) % nv);
|
|
const int w = static_cast<int>(i / (static_cast<size_t>(nu) * nv));
|
|
for (int k = 0; k < n_stencil; k++) {
|
|
const size_t j = wrap(u + stencil[k].x, nu)
|
|
+ static_cast<size_t>(nu) * (wrap(v + stencil[k].y, nv)
|
|
+ static_cast<size_t>(nv) * wrap(w + stencil[k].z, nw));
|
|
if (edge[j]) {
|
|
mask[i] = 1.0f;
|
|
return;
|
|
}
|
|
}
|
|
}
|
|
|
|
unsigned int Blocks(size_t n) {
|
|
return static_cast<unsigned int>((n + THREADS - 1) / THREADS);
|
|
}
|
|
|
|
} // namespace
|
|
|
|
size_t ModelMaskGPU::DeviceBytes(size_t max_points) {
|
|
return 2 * max_points * sizeof(int) + MAX_OPS * sizeof(ModelMaskOp) + MAX_STENCIL * sizeof(int3);
|
|
}
|
|
|
|
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),
|
|
stencil_d(MAX_STENCIL) {
|
|
if (max_points > INT32_MAX)
|
|
throw JFJochException(JFJochExceptionCategory::InputParameterInvalid,
|
|
"Solvent mask grid too large for int32 labels");
|
|
}
|
|
|
|
namespace {
|
|
|
|
// 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).
|
|
void GridSpacing(const ModelMaskGrid &grid, double spacing[3]) {
|
|
const int n[3] = {grid.nu, grid.nv, grid.nw};
|
|
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];
|
|
spacing[k] = grid.volume / (n[k] * std::sqrt(cx * cx + cy * cy + cz * cz));
|
|
}
|
|
}
|
|
|
|
} // namespace
|
|
|
|
void ModelMaskGPU::SetGrid(const ModelMaskGrid &grid, const std::vector<ModelMaskOp> &ops) {
|
|
npoints = static_cast<size_t>(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<float>(grid.orth[3 * i + j] / n[j]);
|
|
|
|
GridSpacing(grid, geom.spacing);
|
|
SetShrinkStencil(grid);
|
|
|
|
// gemmi SolventMasker::remove_islands(), rounded the same way.
|
|
island_limit = static_cast<int>(static_cast<size_t>(ISLAND_MIN_VOLUME * npoints / grid.volume));
|
|
|
|
n_ops = static_cast<int>(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<<<Blocks(npoints), THREADS, 0, stream>>>(d_mask, npoints);
|
|
cuda_err(cudaGetLastError());
|
|
if (n_atoms > 0) {
|
|
mask_images<<<static_cast<unsigned int>(n_atoms) * n_ops, IMAGE_THREADS, 0, stream>>>(
|
|
d_atoms, ops_d, n_ops, geom, d_mask);
|
|
cuda_err(cudaGetLastError());
|
|
}
|
|
RemoveIslands(d_mask);
|
|
Shrink(d_mask);
|
|
}
|
|
|
|
// gemmi's set_margin_around() stencil for r_shrink: every grid offset other than 0 within it, near
|
|
// (closer than the coarsest grid step) first and far after, in gemmi's order.
|
|
void ModelMaskGPU::SetShrinkStencil(const ModelMaskGrid &grid) {
|
|
const int n[3] = {grid.nu, grid.nv, grid.nw};
|
|
const int d[3] = {static_cast<int>(std::floor(R_SHRINK / geom.spacing[0])),
|
|
static_cast<int>(std::floor(R_SHRINK / geom.spacing[1])),
|
|
static_cast<int>(std::floor(R_SHRINK / geom.spacing[2]))};
|
|
if (2 * d[0] >= n[0] || 2 * d[1] >= n[1] || 2 * d[2] >= n[2])
|
|
throw JFJochException(JFJochExceptionCategory::InputParameterInvalid,
|
|
"Solvent mask grid smaller than the shrink radius");
|
|
const double *o = grid.orth;
|
|
// The cell edges a, b, c are the lengths of orth's columns.
|
|
double max_step = 0;
|
|
for (int k = 0; k < 3; k++)
|
|
max_step = std::max(max_step, std::sqrt(o[k] * o[k] + o[3 + k] * o[3 + k] + o[6 + k] * o[6 + k]) / n[k]);
|
|
const double max_spacing2 = max_step * max_step + 1e-6;
|
|
std::vector<int3> near, far;
|
|
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<double>(u) / n[0], static_cast<double>(v) / n[1],
|
|
static_cast<double>(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)
|
|
(r2 < max_spacing2 ? near : far).push_back(make_int3(u, v, w));
|
|
}
|
|
n_near = static_cast<int>(near.size());
|
|
n_stencil = static_cast<int>(near.size() + far.size());
|
|
if (n_stencil > MAX_STENCIL)
|
|
throw JFJochException(JFJochExceptionCategory::InputParameterInvalid, "Solvent mask grid too fine for the shrink");
|
|
near.insert(near.end(), far.begin(), far.end());
|
|
if (n_stencil > 0) {
|
|
cuda_err(cudaMemcpyAsync(stencil_d, near.data(), near.size() * sizeof(int3), cudaMemcpyHostToDevice, stream));
|
|
cuda_err(cudaStreamSynchronize(stream));
|
|
}
|
|
}
|
|
|
|
void ModelMaskGPU::Shrink(float *d_mask) {
|
|
if (n_stencil == 0)
|
|
return;
|
|
int *edge = label; // free once the islands are gone
|
|
shrink_edges<<<Blocks(npoints), THREADS, 0, stream>>>(d_mask, edge, stencil_d, n_near, geom.nu, geom.nv, geom.nw);
|
|
cuda_err(cudaGetLastError());
|
|
shrink_margin<<<Blocks(npoints), THREADS, 0, stream>>>(d_mask, edge, stencil_d, n_stencil, geom.nu, geom.nv,
|
|
geom.nw);
|
|
cuda_err(cudaGetLastError());
|
|
}
|
|
|
|
// gemmi SolventMasker::remove_islands(): every component of solvent with at most island_limit points
|
|
// becomes macromolecule.
|
|
void ModelMaskGPU::RemoveIslands(float *d_mask) {
|
|
uf_init<<<Blocks(npoints), THREADS, 0, stream>>>(d_mask, label, npoints);
|
|
cuda_err(cudaGetLastError());
|
|
uf_merge<<<Blocks(npoints), THREADS, 0, stream>>>(d_mask, label, geom.nu, geom.nv, geom.nw);
|
|
cuda_err(cudaGetLastError());
|
|
uf_roots<<<Blocks(npoints), THREADS, 0, stream>>>(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<<<Blocks(npoints), THREADS, 0, stream>>>(root, size, island_limit, npoints);
|
|
cuda_err(cudaGetLastError());
|
|
uf_remove<<<Blocks(npoints), THREADS, 0, stream>>>(d_mask, root, size, island_limit, npoints);
|
|
cuda_err(cudaGetLastError());
|
|
}
|