Skip to content

Commit d3b5bb8

Browse files
author
abacus_fixer
committed
Merge remote-tracking branch 'upstream/develop' into 2026-08-10-b
2 parents 6821497 + 6e89138 commit d3b5bb8

7 files changed

Lines changed: 560 additions & 89 deletions

File tree

source/source_cell/module_neighlist/bin_manager.cpp

Lines changed: 112 additions & 54 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,15 @@
55
#include <stdexcept>
66
#include "bin_manager.h"
77

8+
#ifdef _OPENMP
9+
#include <omp.h>
10+
#endif
11+
12+
namespace
13+
{
14+
constexpr int neighbor_build_openmp_threshold = 256;
15+
}
16+
817
// ========== Bin class implementation ==========
918

1019
const std::vector<ModuleNeighList::LocalAtomIndex>& Bin::get_atom_indices() const {
@@ -176,6 +185,63 @@ int BinManager::bin_index(int ix, int iy, int iz) const {
176185
return ix * nbiny_ * nbinz_ + iy * nbinz_ + iz;
177186
}
178187

188+
template <typename Emit>
189+
void BinManager::visit_neighbors(const NeighborAtom& atom,
190+
const std::vector<NeighborAtom>& binned_atoms,
191+
double sradius2,
192+
const Emit& emit) const
193+
{
194+
const int ix = std::min(
195+
std::max(int((atom.position_x - x_min_) / bin_sizex_), 0),
196+
nbinx_ - 1
197+
);
198+
199+
const int iy = std::min(
200+
std::max(int((atom.position_y - y_min_) / bin_sizey_), 0),
201+
nbiny_ - 1
202+
);
203+
204+
const int iz = std::min(
205+
std::max(int((atom.position_z - z_min_) / bin_sizez_), 0),
206+
nbinz_ - 1
207+
);
208+
209+
for (int dx = -1; dx <= 1; dx++)
210+
{
211+
for (int dy = -1; dy <= 1; dy++)
212+
{
213+
for (int dz = -1; dz <= 1; dz++)
214+
{
215+
const int jx = ix + dx;
216+
const int jy = iy + dy;
217+
const int jz = iz + dz;
218+
219+
if (jx < 0 || jx >= nbinx_ ||
220+
jy < 0 || jy >= nbiny_ ||
221+
jz < 0 || jz >= nbinz_)
222+
{
223+
continue;
224+
}
225+
226+
const int nidx = bin_index(jx, jy, jz);
227+
for (const ModuleNeighList::LocalAtomIndex binned_atom_index : bins_[nidx].get_atom_indices())
228+
{
229+
const NeighborAtom& natom = binned_atoms[static_cast<std::size_t>(binned_atom_index)];
230+
const double delta_x = atom.position_x - natom.position_x;
231+
const double delta_y = atom.position_y - natom.position_y;
232+
const double delta_z = atom.position_z - natom.position_z;
233+
const double dist2 = delta_x * delta_x + delta_y * delta_y + delta_z * delta_z;
234+
235+
if (natom.atom_id != atom.atom_id && dist2 <= sradius2)
236+
{
237+
emit(natom.atom_id);
238+
}
239+
}
240+
}
241+
}
242+
}
243+
}
244+
179245
void BinManager::build_atom_neighbors(
180246
NeighborList& neighbor_list,
181247
const std::vector<NeighborAtom>& atoms,
@@ -184,71 +250,63 @@ void BinManager::build_atom_neighbors(
184250
{
185251
assert(atoms.size() == static_cast<size_t>(neighbor_list.get_nlocal()));
186252

187-
double sradius2 = sradius_ * sradius_;
253+
const double sradius2 = sradius_ * sradius_;
188254

189255
neighbor_list.reset();
190256

191-
std::vector<int> neigh_tmp;
192-
193257
const int nlocal = neighbor_list.get_nlocal();
194-
for (int i = 0; i < nlocal; i++)
195-
{
196-
neigh_tmp.clear();
197-
const NeighborAtom& atom = atoms[i];
198-
199-
int ix = std::min(
200-
std::max(int((atom.position_x - x_min_) / bin_sizex_), 0),
201-
nbinx_ - 1
202-
);
203258

204-
int iy = std::min(
205-
std::max(int((atom.position_y - y_min_) / bin_sizey_), 0),
206-
nbiny_ - 1
207-
);
208-
209-
int iz = std::min(
210-
std::max(int((atom.position_z - z_min_) / bin_sizez_), 0),
211-
nbinz_ - 1
212-
);
259+
#ifdef _OPENMP
260+
const bool use_parallel = nlocal >= neighbor_build_openmp_threshold && omp_get_max_threads() > 1;
261+
if (use_parallel)
262+
{
263+
std::vector<std::size_t> neighbor_counts(static_cast<std::size_t>(nlocal), 0);
213264

214-
for (int dx = -1; dx <= 1; dx++)
265+
#pragma omp parallel for schedule(static)
266+
for (int i = 0; i < nlocal; i++)
215267
{
216-
for (int dy = -1; dy <= 1; dy++)
217-
{
218-
for (int dz = -1; dz <= 1; dz++)
219-
{
220-
int jx = ix + dx;
221-
int jy = iy + dy;
222-
int jz = iz + dz;
223-
224-
if (jx < 0 || jx >= nbinx_ ||
225-
jy < 0 || jy >= nbiny_ ||
226-
jz < 0 || jz >= nbinz_)
227-
continue;
228-
229-
int nidx = bin_index(jx, jy, jz);
268+
std::size_t count = 0;
269+
visit_neighbors(atoms[i], binned_atoms, sradius2,
270+
[&count](ModuleNeighList::LocalAtomIndex) { ++count; });
271+
neighbor_counts[static_cast<std::size_t>(i)] = count;
272+
}
230273

231-
for (const ModuleNeighList::LocalAtomIndex binned_atom_index : bins_[nidx].get_atom_indices())
232-
{
233-
const NeighborAtom& natom = binned_atoms[static_cast<std::size_t>(binned_atom_index)];
234-
double dx = atom.position_x - natom.position_x;
235-
double dy = atom.position_y - natom.position_y;
236-
double dz = atom.position_z - natom.position_z;
274+
for (int i = 0; i < nlocal; i++)
275+
{
276+
const int n = ModuleNeighList::checked_int_size(
277+
neighbor_counts[static_cast<std::size_t>(i)],
278+
"BinManager neighbor count"
279+
);
280+
neighbor_list.firstneigh_[i] = neighbor_list.allocator_.allocate(n);
281+
neighbor_list.numneigh_[i] = n;
282+
}
237283

238-
double dist2 = dx * dx + dy * dy + dz * dz;
284+
#pragma omp parallel for schedule(static)
285+
for (int i = 0; i < nlocal; i++)
286+
{
287+
int* ptr = neighbor_list.firstneigh_[i];
288+
int k = 0;
289+
visit_neighbors(atoms[i], binned_atoms, sradius2,
290+
[&](ModuleNeighList::LocalAtomIndex atom_id)
291+
{
292+
assert(ptr != nullptr);
293+
ptr[k++] = atom_id;
294+
});
295+
assert(k == neighbor_list.numneigh_[i]);
296+
}
297+
return;
298+
}
299+
#endif
239300

240-
if (natom.atom_id == atom.atom_id)
241-
{
242-
continue;
243-
}
244-
if (dist2 <= sradius2)
301+
std::vector<int> neigh_tmp;
302+
for (int i = 0; i < nlocal; i++)
303+
{
304+
neigh_tmp.clear();
305+
visit_neighbors(atoms[i], binned_atoms, sradius2,
306+
[&neigh_tmp](ModuleNeighList::LocalAtomIndex atom_id)
245307
{
246-
neigh_tmp.push_back(natom.atom_id);
247-
}
248-
}
249-
}
250-
}
251-
}
308+
neigh_tmp.push_back(atom_id);
309+
});
252310

253311
const int n = ModuleNeighList::checked_int_size(neigh_tmp.size(), "BinManager neighbor count");
254312

source/source_cell/module_neighlist/bin_manager.h

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -198,6 +198,21 @@ class BinManager
198198
* @return Flat index in the bins_ array.
199199
*/
200200
int bin_index(int ix, int iy, int iz) const;
201+
202+
/**
203+
* @brief Visit neighbors of one atom in the existing deterministic bin order.
204+
*
205+
* @tparam Emit Callable accepting a rank-local neighbor atom ID.
206+
* @param atom Atom used as the neighbor-list center.
207+
* @param binned_atoms All atoms assigned to bins by do_binning().
208+
* @param sradius2 Squared search radius.
209+
* @param emit Callback invoked once for every accepted neighbor.
210+
*/
211+
template <typename Emit>
212+
void visit_neighbors(const NeighborAtom& atom,
213+
const std::vector<NeighborAtom>& binned_atoms,
214+
double sradius2,
215+
const Emit& emit) const;
201216
};
202217

203218
#endif // BIN_MANAGER_H

source/source_cell/module_neighlist/neighbor_search.cpp

Lines changed: 109 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
#include "source_cell/module_neighlist/neighbor_search.h"
22
#include <cmath>
33
#include <algorithm>
4+
#include <cstdint>
45
#include <limits>
56
#include <cassert>
67
#include <array>
@@ -112,13 +113,14 @@ void NeighborSearch::init(const AtomProvider& ucell, double sr)
112113
{
113114
for (int j = 0; j < ucell.get_na(i); j++)
114115
{
116+
const ModuleBase::Vector3<double> position = ucell.get_tau(i, j);
115117
const ModuleNeighList::LocalAtomIndex atom_count
116118
= ModuleNeighList::checked_local_atom_index(all_atoms_.size(),
117119
"NeighborSearch atom id");
118120
NeighborAtom atom(
119-
ucell.get_tau(i,j).x,
120-
ucell.get_tau(i,j).y,
121-
ucell.get_tau(i,j).z,
121+
position.x,
122+
position.y,
123+
position.z,
122124
i,
123125
j,
124126
atom_count
@@ -189,37 +191,111 @@ void NeighborSearch::check_expand_condition(const AtomProvider& ucell, int& glay
189191

190192
void NeighborSearch::set_member_variables(const AtomProvider& ucell, int glayerX_minus, int glayerX, int glayerY_minus, int glayerY, int glayerZ_minus, int glayerZ)
191193
{
192-
ModuleBase::Vector3<double> vec1(ucell.get_latvec().e11, ucell.get_latvec().e12, ucell.get_latvec().e13);
193-
ModuleBase::Vector3<double> vec2(ucell.get_latvec().e21, ucell.get_latvec().e22, ucell.get_latvec().e23);
194-
ModuleBase::Vector3<double> vec3(ucell.get_latvec().e31, ucell.get_latvec().e32, ucell.get_latvec().e33);
194+
if (glayerX_minus < 0 || glayerX < 0 || glayerY_minus < 0 || glayerY < 0
195+
|| glayerZ_minus < 0 || glayerZ < 0)
196+
{
197+
throw std::invalid_argument("NeighborSearch periodic image layers must be non-negative.");
198+
}
195199

196-
for (int ix = -glayerX_minus; ix < glayerX; ix++)
200+
const ModuleBase::Matrix3& lattice = ucell.get_latvec();
201+
const ModuleBase::Vector3<double> vec1(lattice.e11, lattice.e12, lattice.e13);
202+
const ModuleBase::Vector3<double> vec2(lattice.e21, lattice.e22, lattice.e23);
203+
const ModuleBase::Vector3<double> vec3(lattice.e31, lattice.e32, lattice.e33);
204+
205+
const std::size_t image_count_x
206+
= ModuleNeighList::checked_size_sum(static_cast<std::size_t>(glayerX_minus),
207+
static_cast<std::size_t>(glayerX),
208+
"NeighborSearch x image count");
209+
const std::size_t image_count_y
210+
= ModuleNeighList::checked_size_sum(static_cast<std::size_t>(glayerY_minus),
211+
static_cast<std::size_t>(glayerY),
212+
"NeighborSearch y image count");
213+
const std::size_t image_count_z
214+
= ModuleNeighList::checked_size_sum(static_cast<std::size_t>(glayerZ_minus),
215+
static_cast<std::size_t>(glayerZ),
216+
"NeighborSearch z image count");
217+
const std::size_t image_count_yz
218+
= ModuleNeighList::checked_size_product(image_count_y,
219+
image_count_z,
220+
"NeighborSearch yz image count");
221+
const std::size_t image_count
222+
= ModuleNeighList::checked_size_product(image_count_x,
223+
image_count_yz,
224+
"NeighborSearch periodic image count");
225+
if (image_count == 0)
197226
{
198-
for (int iy = -glayerY_minus; iy < glayerY; iy++)
199-
{
200-
for (int iz = -glayerZ_minus; iz < glayerZ; iz++)
201-
{
202-
if(ix==0 && iy==0 && iz==0)
203-
{
204-
continue;
205-
}
206-
for (int i = 0; i < ucell.get_ntype(); i++)
207-
{
208-
for (int j = 0; j < ucell.get_na(i); j++)
209-
{
210-
double atom_x = ucell.get_tau(i,j).x + vec1[0] * ix + vec2[0] * iy + vec3[0] * iz;
211-
double atom_y = ucell.get_tau(i,j).y + vec1[1] * ix + vec2[1] * iy + vec3[1] * iz;
212-
double atom_z = ucell.get_tau(i,j).z + vec1[2] * ix + vec2[2] * iy + vec3[2] * iz;
213-
214-
const ModuleNeighList::LocalAtomIndex atom_count
215-
= ModuleNeighList::checked_local_atom_index(all_atoms_.size(),
216-
"NeighborSearch atom id");
217-
NeighborAtom atom(atom_x, atom_y, atom_z, i, j, atom_count);
218-
ghost_atoms_.push_back(atom);
219-
all_atoms_.push_back(atom);
220-
}
221-
}
222-
}
223-
}
227+
throw std::invalid_argument("NeighborSearch periodic image range must contain the primary cell.");
228+
}
229+
230+
const std::size_t origin_xy
231+
= ModuleNeighList::checked_size_sum(
232+
ModuleNeighList::checked_size_product(static_cast<std::size_t>(glayerX_minus),
233+
image_count_y,
234+
"NeighborSearch primary image index"),
235+
static_cast<std::size_t>(glayerY_minus),
236+
"NeighborSearch primary image index");
237+
const std::size_t origin_image
238+
= ModuleNeighList::checked_size_sum(
239+
ModuleNeighList::checked_size_product(origin_xy,
240+
image_count_z,
241+
"NeighborSearch primary image index"),
242+
static_cast<std::size_t>(glayerZ_minus),
243+
"NeighborSearch primary image index");
244+
245+
const std::size_t base_atom_count = inside_atoms_.size();
246+
const std::size_t ghost_image_count = image_count - 1;
247+
const std::size_t generated_atom_count
248+
= ModuleNeighList::checked_size_product(ghost_image_count,
249+
base_atom_count,
250+
"NeighborSearch generated atom count");
251+
const std::size_t final_atom_count
252+
= ModuleNeighList::checked_size_sum(base_atom_count,
253+
generated_atom_count,
254+
"NeighborSearch total atom count");
255+
if (final_atom_count > static_cast<std::size_t>(std::numeric_limits<ModuleNeighList::LocalAtomIndex>::max()))
256+
{
257+
throw std::overflow_error("NeighborSearch total atom count exceeds local atom index range.");
258+
}
259+
if (generated_atom_count == 0)
260+
{
261+
return;
262+
}
263+
264+
const NeighborAtom placeholder(0.0, 0.0, 0.0, 0, 0, 0);
265+
all_atoms_.insert(all_atoms_.end(), generated_atom_count, placeholder);
266+
ghost_atoms_.assign(generated_atom_count, placeholder);
267+
268+
// Thread creation costs more than it saves for small unit cells.
269+
const std::size_t parallel_threshold = 100000;
270+
#ifdef _OPENMP
271+
#pragma omp parallel for schedule(static) if(generated_atom_count >= parallel_threshold)
272+
#endif
273+
for (std::int64_t generated_index = 0;
274+
generated_index < static_cast<std::int64_t>(generated_atom_count);
275+
++generated_index)
276+
{
277+
const std::size_t output_index = static_cast<std::size_t>(generated_index);
278+
const std::size_t ghost_image = output_index / base_atom_count;
279+
const std::size_t base_atom = output_index % base_atom_count;
280+
const std::size_t image = ghost_image < origin_image ? ghost_image : ghost_image + 1;
281+
const std::size_t image_x = image / image_count_yz;
282+
const std::size_t image_yz = image % image_count_yz;
283+
const std::size_t image_y = image_yz / image_count_z;
284+
const std::size_t image_z = image_yz % image_count_z;
285+
const double ix = static_cast<double>(image_x) - static_cast<double>(glayerX_minus);
286+
const double iy = static_cast<double>(image_y) - static_cast<double>(glayerY_minus);
287+
const double iz = static_cast<double>(image_z) - static_cast<double>(glayerZ_minus);
288+
289+
const NeighborAtom& source = inside_atoms_[base_atom];
290+
const ModuleNeighList::LocalAtomIndex atom_id
291+
= static_cast<ModuleNeighList::LocalAtomIndex>(base_atom_count + output_index);
292+
const NeighborAtom atom(source.position_x + vec1[0] * ix + vec2[0] * iy + vec3[0] * iz,
293+
source.position_y + vec1[1] * ix + vec2[1] * iy + vec3[1] * iz,
294+
source.position_z + vec1[2] * ix + vec2[2] * iy + vec3[2] * iz,
295+
source.atom_type,
296+
source.atom_index,
297+
atom_id);
298+
ghost_atoms_[output_index] = atom;
299+
all_atoms_[base_atom_count + output_index] = atom;
224300
}
225301
}

0 commit comments

Comments
 (0)