Files
Jungfraujoch/image_analysis/structure_refinement/ModelDensityGPU.cu
T
leonarski_fandClaude Opus 5.5 d5fcf2f05b rugnux: model validation's structure factors and maps on the GPU
ModelStructureFactorsGPU computes what compute_model_factors() and
map_from_coefficients() compute on the CPU - F_calc from the model's
density (IT92, Refmac-compatible blur, unblurred as prepare_asu_data()
does) and F_mask from the Refmac bulk-solvent mask, both on the
reflections prepare_asu_data(d_min) lists, in its order; and a map from
ASU coefficients on the grid get_size_for_hkl(coef, 0, 3.0) sizes - on a
device. Made once per cell, group, resolution and model, then evaluated
as often as the coordinates change, so refinement or MR can call it in a
loop. The device is an explicit parameter; every call leaves the calling
thread's current device as it found it.

Pieces:
- ModelDensityGPU: the rigid body's deterministic brick gather, moved
  out of RigidBodyGPU.cu into a component of its own (ModelMaskGPU's
  pattern); the rigid body uses it unchanged. MAX_BRICKS_PER_AXIS 8 ->
  16, so fine grids with high-B atoms (lysozyme at 1.2 A, a 0.9 A P1
  cell) are no longer refused; existing zones are gridded identically.
- One copy of the content is gridded and the symmetry composed in
  reciprocal space (SymmetryComposition), operators applied on the fly;
  the mask is ModelMaskGPU (every image of every atom, islands, shrink).
- Maps: gemmi's get_f_phi_on_grid() in ZYX order on the host (the
  coefficients written are the same), in-place cuFFT c2r, transposed back
  to XYZ on the device. One map at a time, in the engine's buffers.

Decided once, up front, per card, from its TOTAL memory: the engine's
bytes (16 N + cuFFT work + reflections, N the larger of the structure-
factor and map grids) must be at most half the card - the rigid body's
engines take at most a quarter beside it. Otherwise, or where the gather
cannot grid the cell, the CPU path runs, logged with needed vs total.
Anything to a resolution other than d_min (the null's 3.5 A fits) stays
on the CPU, so all replicates and the real model's side of the null are
computed the same way. A CUDA failure takes the existing path: the
validation restarts on the CPU.

Measured, model validation total per run (CPU path -> GPU), 16 GB card:
  F432 215 A cubic, 1.30 A, 500^3 grid: 47.7 -> 15.6 s (two validations;
     14.3 -> 2.4 and 33.4 -> 13.2, the rest of the second is writing the
     three 0.5 GB maps); F_calc + F_mask 5.7 s -> 0.05 s
  C2 1.11 A: 23.9 -> 13.6 s; P3_2 1.55 A: 18.2 -> 10.2 s;
  P2_1 1.25 A: 13.4 -> 6.5 s; F4_132 328 A: 13.0 -> 5.5 s;
  P6_5: 8.8 -> 4.2 s; P4_3 0.97 A: 4.6 -> 2.5 s; P1 0.92 A: 4.2 -> 2.2 s;
  small P1: 3.2 -> 1.3 s; P6_1: 8.8 -> 6.2 s; lysozyme: 1.8 -> 1.4 s.
p.mtz md5-identical on all 13 sets. Against the CPU path: FC within
1e-4 of mean |F|, phases of the strong half within 0.003 deg, maps within
1e-4 (2mFo-DFc) and 7e-4 (mFo-DFc) of their rms; every logged R, CC,
FOM, k_sol and anomalous site list identical at the printed precision,
except where a rigid-body commit sat on an exact R-free tie (0.2155 ->
0.2155) and fell the other way (R-work 0.2127 vs 0.2129). The GPU result
is bit-identical run to run and with -N 8 (maps, map MTZ, placed model).
Peak device memory of the engine: 2.5 GB at 500^3 (process total peaked
at 14.4 GB with what the merge still holds).

Tests: ModelStructureFactorsGPU_MatchesCPU (five groups, 3.5 and 1.5 A:
same reflections, F_calc <= 1e-5 of mean |F|, F_mask 2e-7 rms, repeat
bit-identical), ModelStructureFactorsGPU_MapMatchesCPU (<= 5e-6 of rms).

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01SVmAWnzCmRKAXVUCdc4iNi
2026-10-09 12:38:37 +02:00

351 lines
17 KiB
Plaintext

// SPDX-FileCopyrightText: 2026 Filip Leonarski, Paul Scherrer Institute <filip.leonarski@psi.ch>
// SPDX-License-Identifier: GPL-3.0-only
#include "ModelDensityGPU.h"
#include <algorithm>
#include <cmath>
#include <cub/cub.cuh>
#include "../../common/JFJochException.h"
namespace {
void cuda_err(cudaError_t val) {
if (val != cudaSuccess)
throw JFJochException(JFJochExceptionCategory::GPUCUDAError, cudaGetErrorString(val));
}
constexpr int BRICK = 8; // the gather's bricks are BRICK^3 grid points, one block each
constexpr int GATHER_TILE = 128; // atoms staged in shared memory at a time
constexpr int MAX_BRICKS_PER_AXIS = 16; // an atom's box touches at most this many bricks along an axis
constexpr int THREADS = 256;
__device__ __forceinline__ int imod(int a, int n) {
const int r = a % n;
return r < 0 ? r + n : r;
}
// gemmi's unsafe_expapprox() (formfact.hpp), the exponential the density is computed with.
__device__ __forceinline__ float expapprox(float x) {
const float val = 12102203.1615614f * x + 1065353216.f;
const int vali = static_cast<int>(val);
const float a = __int_as_float(vali & 0x7F800000);
const float b = __int_as_float((vali & 0x7FFFFF) | 0x3F800000);
return a * (0.509871020f + b * (0.312146713f + b * (0.166617139f + b * (-2.190619930e-3f + b * 1.3555747234e-2f))));
}
// The distinct bricks the points c - d ... c + d of one axis fall in, wrapped into the cell. A grid size
// that is not a multiple of BRICK leaves the last brick partial, which is why this is done on wrapped
// points and not in brick coordinates. -1 where there are more than MAX_BRICKS_PER_AXIS of them.
__device__ int axis_bricks(int c, int d, int n, int *out) {
int k = 0;
for (int p = c - d; p <= c + d; p++) {
const int b = imod(p, n) / BRICK;
bool seen = false;
for (int j = 0; j < k; j++)
seen = seen || out[j] == b;
if (seen)
continue;
if (k == MAX_BRICKS_PER_AXIS)
return -1;
out[k++] = b;
}
return k;
}
// The bricks atom i's box touches, per axis. The box is gemmi's: `box` points either side of the
// point nearest the atom.
__device__ void atom_bricks(const float4 &pos, const int *box, int nu, int nv, int nw, int *bu, int &nbu, int *bv,
int &nbv, int *bw, int &nbw) {
nbu = axis_bricks(__float2int_rn(pos.x), box[0], nu, bu);
nbv = axis_bricks(__float2int_rn(pos.y), box[1], nv, bv);
nbw = axis_bricks(__float2int_rn(pos.z), box[2], nw, bw);
}
__global__ void count_pairs_kernel(const float4 *pos, const int *box, int n, int nu, int nv, int nw, int *count,
int *overflow) {
const int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i > n)
return;
if (i == n) {
count[n] = 0; // so the exclusive scan's last entry is the total
return;
}
int bu[MAX_BRICKS_PER_AXIS], bv[MAX_BRICKS_PER_AXIS], bw[MAX_BRICKS_PER_AXIS], a, b, c;
atom_bricks(pos[i], box + 3 * i, nu, nv, nw, bu, a, bv, b, bw, c);
if (a < 0 || b < 0 || c < 0) {
*overflow = 1;
a = b = c = 0;
}
count[i] = a * b * c;
}
// The (brick, atom) pairs, atom-major: sorted stably by brick, each brick's atoms stay in model order.
__global__ void fill_pairs_kernel(const float4 *pos, const int *box, int n, int nu, int nv, int nw,
const int *offset, int *key, int *value) {
const int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= n)
return;
int bu[MAX_BRICKS_PER_AXIS], bv[MAX_BRICKS_PER_AXIS], bw[MAX_BRICKS_PER_AXIS], a, b, c;
atom_bricks(pos[i], box + 3 * i, nu, nv, nw, bu, a, bv, b, bw, c);
const int nbu = (nu + BRICK - 1) / BRICK, nbv = (nv + BRICK - 1) / BRICK;
int o = offset[i];
for (int k = 0; k < c; k++)
for (int j = 0; j < b; j++)
for (int l = 0; l < a; l++) {
key[o] = bu[l] + nbu * (bv[j] + nbv * bw[k]);
value[o] = i;
o++;
}
}
__global__ void brick_ranges_kernel(const int *key, int npairs, int *start, int *end) {
const int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i >= npairs)
return;
const int k = key[i];
if (i == 0 || key[i - 1] != k)
start[k] = i;
if (i == npairs - 1 || key[i + 1] != k)
end[k] = i + 1;
}
// The density of the model on the grid, as gemmi's do_add_atom_density_to_grid() puts it there: every
// point adds, in model order, each atom whose sphere it is inside. One block per brick; the brick's atoms
// are staged in shared memory with the position of their image nearest to the brick, relative to the
// brick's first point. Where the cell is wider than an atom's box plus a brick, that is the only image
// that reaches any point of the brick; on a narrower cell each point finds its own nearest image.
__global__ void gather_kernel(const ModelDensityAtom *atoms, const float4 *pos, const int *pair_atom,
const int *brick_start, const int *brick_end, ModelDensityGPU::Geometry g,
float *grid) {
const int nbu = (g.nu + BRICK - 1) / BRICK, nbv = (g.nv + BRICK - 1) / BRICK;
const int brick = blockIdx.x;
const int s = brick_start[brick], e = brick_end[brick]; // s == e: an empty brick, written as zeros
const int u0 = (brick % nbu) * BRICK, v0 = (brick / nbu % nbv) * BRICK, w0 = (brick / (nbu * nbv)) * BRICK;
const int u = u0 + threadIdx.x, v = v0 + threadIdx.y, w = w0 + threadIdx.z;
const bool inside = u < g.nu && v < g.nv && w < g.nw;
const int tid = threadIdx.x + BRICK * (threadIdx.y + BRICK * threadIdx.z);
// This point relative to the brick's first one, in Cartesian coordinates.
const float tx = g.orth_n[0] * threadIdx.x + g.orth_n[1] * threadIdx.y + g.orth_n[2] * threadIdx.z;
const float ty = g.orth_n[3] * threadIdx.x + g.orth_n[4] * threadIdx.y + g.orth_n[5] * threadIdx.z;
const float tz = g.orth_n[6] * threadIdx.x + g.orth_n[7] * threadIdx.y + g.orth_n[8] * threadIdx.z;
constexpr int WORDS = sizeof(ModelDensityAtom) / sizeof(int);
__shared__ int sh_index[GATHER_TILE];
__shared__ ModelDensityAtom sh_atom[GATHER_TILE];
__shared__ float3 sh_centre[GATHER_TILE]; // Cartesian
__shared__ float3 sh_grid[GATHER_TILE]; // the same, in grid units
float acc = 0.0f;
for (int c = s; c < e; c += GATHER_TILE) {
const int m = min(GATHER_TILE, e - c);
__syncthreads();
if (tid < m) {
const int a = pair_atom[c + tid];
sh_index[tid] = a;
// The atom's image nearest to the brick's centre, in grid units from its first point.
const float4 p = pos[a];
float fx = p.x - u0, fy = p.y - v0, fz = p.z - w0;
fx -= g.nu * rintf((fx - 0.5f * (BRICK - 1)) / g.nu);
fy -= g.nv * rintf((fy - 0.5f * (BRICK - 1)) / g.nv);
fz -= g.nw * rintf((fz - 0.5f * (BRICK - 1)) / g.nw);
sh_grid[tid] = make_float3(fx, fy, fz);
sh_centre[tid] = make_float3(g.orth_n[0] * fx + g.orth_n[1] * fy + g.orth_n[2] * fz,
g.orth_n[3] * fx + g.orth_n[4] * fy + g.orth_n[5] * fz,
g.orth_n[6] * fx + g.orth_n[7] * fy + g.orth_n[8] * fz);
}
__syncthreads();
for (int i = tid; i < m * WORDS; i += BRICK * BRICK * BRICK)
reinterpret_cast<int *>(sh_atom)[i] = reinterpret_cast<const int *>(atoms + sh_index[i / WORDS])[i % WORDS];
__syncthreads();
if (!inside)
continue;
for (int k = 0; k < m; k++) {
const ModelDensityAtom &at = sh_atom[k];
float x = sh_centre[k].x - tx, y = sh_centre[k].y - ty, z = sh_centre[k].z - tz;
if (g.narrow) {
float fx = sh_grid[k].x - threadIdx.x, fy = sh_grid[k].y - threadIdx.y, fz = sh_grid[k].z - threadIdx.z;
fx -= g.nu * rintf(fx / g.nu);
fy -= g.nv * rintf(fy / g.nv);
fz -= g.nw * rintf(fz / g.nw);
x = g.orth_n[0] * fx + g.orth_n[1] * fy + g.orth_n[2] * fz;
y = g.orth_n[3] * fx + g.orth_n[4] * fy + g.orth_n[5] * fz;
z = g.orth_n[6] * fx + g.orth_n[7] * fy + g.orth_n[8] * fz;
}
const float r2 = x * x + y * y + z * z;
if (r2 > at.radius * at.radius)
continue;
float density = 0.0f;
if (!at.aniso) {
for (int q = 0; q < 5; q++)
density += at.a[q] * expapprox(fmaxf(at.b[q][0] * r2, -88.f));
} else {
for (int q = 0; q < 5; q++) {
const float *b = at.b[q];
const float rur = x * x * b[0] + y * y * b[1] + z * z * b[2] + 2 * (x * y * b[3] + x * z * b[4] + y * z * b[5]);
density += at.a[q] * expapprox(fmaxf(rur, -88.f));
}
}
acc += at.occ * density;
}
}
if (inside)
grid[u + static_cast<size_t>(g.nu) * (v + static_cast<size_t>(g.nv) * w)] = acc;
}
int blocks(size_t n) {
return static_cast<int>((n + THREADS - 1) / THREADS);
}
// cub's temporary storage for the pair scan and sort at the capacity.
size_t CubBytes(size_t atoms, size_t pairs) {
size_t scan = 0, sort = 0;
cuda_err(cub::DeviceScan::ExclusiveSum(nullptr, scan, static_cast<int *>(nullptr), static_cast<int *>(nullptr),
static_cast<int>(atoms + 1)));
cuda_err(cub::DeviceRadixSort::SortPairs(nullptr, sort, static_cast<int *>(nullptr), static_cast<int *>(nullptr),
static_cast<int *>(nullptr), static_cast<int *>(nullptr),
static_cast<int>(pairs)));
return std::max(scan, sort);
}
} // namespace
// The most bricks the 2 d + 1 wrapped points of a box can fall in along an axis of n points: one more
// than the points span for where they start in a brick, and one more again where they wrap past a last
// brick that is partial (n = 17: the points 15, 16, 0 fall in bricks 1, 2 and 0).
size_t ModelDensityGPU::AxisBrickBound(int d, int n) {
const int nb = (n + BRICK - 1) / BRICK;
return std::min(nb, (2 * d + 1 + BRICK - 1) / BRICK + 2);
}
size_t ModelDensityGPU::Bricks(int nu, int nv, int nw) {
return static_cast<size_t>((nu + BRICK - 1) / BRICK) * ((nv + BRICK - 1) / BRICK) * ((nw + BRICK - 1) / BRICK);
}
std::vector<int> ModelDensityGPU::Boxes(const ModelDensityGrid &grid, const std::vector<ModelDensityAtom> &atoms) {
const int n3[3] = {grid.nu, grid.nv, grid.nw};
double spacing[3];
for (int k = 0; k < 3; k++)
spacing[k] = 1.0 / (n3[k] * std::sqrt(grid.frac[3 * k] * grid.frac[3 * k] + grid.frac[3 * k + 1] * grid.frac[3 * k + 1] +
grid.frac[3 * k + 2] * grid.frac[3 * k + 2]));
std::vector<int> box(3 * atoms.size());
for (size_t i = 0; i < atoms.size(); i++)
for (int k = 0; k < 3; k++)
box[3 * i + k] = static_cast<int>(std::ceil(atoms[i].radius / spacing[k]));
return box;
}
size_t ModelDensityGPU::PairBound(const ModelDensityGrid &grid, const std::vector<ModelDensityAtom> &atoms) {
const std::vector<int> box = Boxes(grid, atoms);
const int n3[3] = {grid.nu, grid.nv, grid.nw};
size_t pairs = 0;
for (size_t i = 0; i < atoms.size(); i++) {
size_t product = 1;
for (int k = 0; k < 3; k++)
product *= AxisBrickBound(box[3 * i + k], n3[k]);
pairs += product;
}
return pairs;
}
bool ModelDensityGPU::Supports(const ModelDensityGrid &grid, const std::vector<ModelDensityAtom> &atoms) {
const std::vector<int> box = Boxes(grid, atoms);
const int n3[3] = {grid.nu, grid.nv, grid.nw};
for (size_t i = 0; i < box.size(); i++)
if (2 * box[i] + 1 > n3[i % 3] || AxisBrickBound(box[i], n3[i % 3]) > MAX_BRICKS_PER_AXIS)
return false;
return true;
}
size_t ModelDensityGPU::DeviceBytes(size_t max_atoms, size_t max_pairs, size_t max_bricks) {
const size_t na = max_atoms + 1;
return na * (sizeof(ModelDensityAtom) + 3 * sizeof(int) + 2 * sizeof(int)) + 4 * max_pairs * sizeof(int) +
2 * max_bricks * sizeof(int) + CubBytes(na, max_pairs + 1);
}
ModelDensityGPU::ModelDensityGPU(cudaStream_t stream, size_t max_atoms_, size_t max_pairs_, size_t max_bricks_)
: stream(stream), max_atoms(max_atoms_), max_pairs(max_pairs_), max_bricks(max_bricks_) {
const auto sync = CudaAlloc::Synchronous;
const size_t na = std::max<size_t>(max_atoms, 1), np = std::max<size_t>(max_pairs, 1),
nb = std::max<size_t>(max_bricks, 1);
atoms = CudaDevicePtr<ModelDensityAtom>(na, sync);
box = CudaDevicePtr<int>(3 * na, sync);
overflow = CudaDevicePtr<int>(1, sync);
count = CudaDevicePtr<int>(na + 1, sync);
offset = CudaDevicePtr<int>(na + 1, sync);
key = CudaDevicePtr<int>(np, sync);
value = CudaDevicePtr<int>(np, sync);
key_sorted = CudaDevicePtr<int>(np, sync);
value_sorted = CudaDevicePtr<int>(np, sync);
brick_start = CudaDevicePtr<int>(nb, sync);
brick_end = CudaDevicePtr<int>(nb, sync);
cub_bytes = CubBytes(na, np);
cub_temp = CudaDevicePtr<char>(std::max<size_t>(cub_bytes, 1), sync);
}
void ModelDensityGPU::SetAtoms(const ModelDensityGrid &grid, const std::vector<ModelDensityAtom> &atom_list) {
if (atom_list.size() > max_atoms || Bricks(grid.nu, grid.nv, grid.nw) > max_bricks)
throw JFJochException(JFJochExceptionCategory::GPUCUDAError, "model density: grid or atoms over the reserve");
// An atom whose box is wider than the cell would reach a point through more than one image, which
// gemmi's box walk adds and the gather's nearest image does not; such a cell is left to the CPU.
if (!Supports(grid, atom_list))
throw JFJochException(JFJochExceptionCategory::GPUCUDAError, "model density: an atom wider than the cell");
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] / (j == 0 ? grid.nu : j == 1 ? grid.nv : grid.nw));
const std::vector<int> b = Boxes(grid, atom_list);
geom.narrow = 0;
for (size_t i = 0; i < b.size(); i++)
if (2 * b[i] + BRICK > (i % 3 == 0 ? grid.nu : i % 3 == 1 ? grid.nv : grid.nw))
geom.narrow = 1;
n_atoms = static_cast<int>(atom_list.size());
if (n_atoms > 0) {
cuda_err(cudaMemcpyAsync(atoms, atom_list.data(), atom_list.size() * sizeof(ModelDensityAtom),
cudaMemcpyHostToDevice, stream));
cuda_err(cudaMemcpyAsync(box, b.data(), b.size() * sizeof(int), cudaMemcpyHostToDevice, stream));
}
// The uploads read host vectors that end with this call.
cuda_err(cudaStreamSynchronize(stream));
}
void ModelDensityGPU::Compute(const float4 *d_pos, float *d_grid) {
cuda_err(cudaMemsetAsync(overflow, 0, sizeof(int), stream));
count_pairs_kernel<<<blocks(n_atoms + 1), THREADS, 0, stream>>>(d_pos, box, n_atoms, geom.nu, geom.nv, geom.nw,
count, overflow);
cuda_err(cudaGetLastError());
size_t temp = cub_bytes;
cuda_err(cub::DeviceScan::ExclusiveSum(cub_temp.get(), temp, count.get(), offset.get(), n_atoms + 1, stream));
int npairs = 0;
int over = 0;
cuda_err(cudaMemcpyAsync(&npairs, offset.get() + n_atoms, sizeof(int), cudaMemcpyDeviceToHost, stream));
cuda_err(cudaMemcpyAsync(&over, overflow.get(), sizeof(int), cudaMemcpyDeviceToHost, stream));
cuda_err(cudaStreamSynchronize(stream));
if (over)
throw JFJochException(JFJochExceptionCategory::GPUCUDAError, "model density: an atom's box over the bricks per axis");
if (static_cast<size_t>(npairs) > max_pairs)
throw JFJochException(JFJochExceptionCategory::GPUCUDAError, "model density: gather pairs over the reserve");
fill_pairs_kernel<<<blocks(n_atoms), THREADS, 0, stream>>>(d_pos, box, n_atoms, geom.nu, geom.nv, geom.nw,
offset, key, value);
cuda_err(cudaGetLastError());
const int nb = static_cast<int>(Bricks(geom.nu, geom.nv, geom.nw));
int bits = 1;
while ((1 << bits) < nb)
bits++;
temp = cub_bytes;
cuda_err(cub::DeviceRadixSort::SortPairs(cub_temp.get(), temp, key.get(), key_sorted.get(), value.get(),
value_sorted.get(), npairs, 0, bits, stream));
cuda_err(cudaMemsetAsync(brick_start, 0, nb * sizeof(int), stream));
cuda_err(cudaMemsetAsync(brick_end, 0, nb * sizeof(int), stream));
if (npairs > 0) {
brick_ranges_kernel<<<blocks(npairs), THREADS, 0, stream>>>(key_sorted, npairs, brick_start, brick_end);
cuda_err(cudaGetLastError());
}
gather_kernel<<<nb, dim3(BRICK, BRICK, BRICK), 0, stream>>>(atoms, d_pos, value_sorted, brick_start, brick_end,
geom, d_grid);
cuda_err(cudaGetLastError());
}