Skip to content

Commit 95d8a61

Browse files
committed
Feature: update OpenMP in RI_2D_Comm::split_m2D_ktoR_k()
1 parent 85ebe50 commit 95d8a61

2 files changed

Lines changed: 82 additions & 70 deletions

File tree

source/source_lcao/module_ri/RI_2D_Comm.h

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -46,7 +46,7 @@ extern std::vector<std::map<TA, std::map<TAC, RI::Tensor<Tdata>>>> split_m2D_kto
4646
const int nspin);
4747

4848
template <typename Tdata, typename Tmatrix>
49-
extern std::vector<std::map<TA, std::map<TAC, RI::Tensor<Tdata>>>> split_m2D_ktoR_general(
49+
extern std::vector<std::map<TA, std::map<TAC, RI::Tensor<Tdata>>>> split_m2D_ktoR_k(
5050
const UnitCell& ucell,
5151
const K_Vectors& kv,
5252
const std::vector<const Tmatrix*>& mks_2D,

source/source_lcao/module_ri/RI_2D_Comm.hpp

Lines changed: 81 additions & 69 deletions
Original file line numberDiff line numberDiff line change
@@ -50,7 +50,7 @@ auto RI_2D_Comm::split_m2D_ktoR(const UnitCell& ucell,
5050
std::vector<std::map<TA, std::map<TAC, RI::Tensor<Tdata>>>> mRs_a2D
5151
= (period == TC{1, 1, 1})
5252
? RI_2D_Comm::split_m2D_ktoR_gamma<Tdata, Tmatrix>(ucell, mks_2D, pv, nspin)
53-
: RI_2D_Comm::split_m2D_ktoR_general<Tdata, Tmatrix>(ucell, kv, mks_2D, pv, nspin, spgsym);
53+
: RI_2D_Comm::split_m2D_ktoR_k<Tdata, Tmatrix>(ucell, kv, mks_2D, pv, nspin, spgsym);
5454
ModuleBase::timer::end("RI_2D_Comm", "split_m2D_ktoR");
5555
return mRs_a2D;
5656
}
@@ -149,104 +149,116 @@ auto RI_2D_Comm::split_m2D_ktoR_gamma(const UnitCell& ucell,
149149
}
150150

151151
template<typename Tdata, typename Tmatrix>
152-
auto RI_2D_Comm::split_m2D_ktoR_general(const UnitCell& ucell,
152+
auto RI_2D_Comm::split_m2D_ktoR_k(const UnitCell& ucell,
153153
const K_Vectors& kv,
154154
const std::vector<const Tmatrix*>& mks_2D,
155155
const Parallel_2D& pv,
156156
const int nspin,
157157
const bool spgsym)
158158
-> std::vector<std::map<TA, std::map<TAC, RI::Tensor<Tdata>>>>
159159
{
160-
ModuleBase::TITLE("RI_2D_Comm","split_m2D_ktoR_general");
161-
ModuleBase::timer::start("RI_2D_Comm", "split_m2D_ktoR_general");
160+
ModuleBase::TITLE("RI_2D_Comm","split_m2D_ktoR_k");
161+
ModuleBase::timer::start("RI_2D_Comm", "split_m2D_ktoR_k");
162162

163163
const TC period = RI_Util::get_Born_vonKarmen_period(kv);
164164
const std::map<int,int> nspin_k = {{1,1}, {2,2}, {4,1}};
165165
const double SPIN_multiple = std::map<int, double>{ {1,0.5}, {2,1}, {4,1} }.at(nspin); // why?
166166

167167
std::vector<std::map<TA, std::map<TAC, RI::Tensor<Tdata>>>> mRs_a2D(nspin);
168-
for (int is_k = 0; is_k < nspin_k.at(nspin); ++is_k)
169-
{
170-
const std::vector<int> ik_list = RI_2D_Comm::get_ik_list(kv, is_k);
171-
const auto cells = RI_Util::get_Born_von_Karmen_cells(period);
172-
#ifdef _OPENMP
173-
#pragma omp parallel for schedule(dynamic)
174-
#endif
175-
for (size_t icell = 0; icell < cells.size(); ++icell)
168+
#ifdef _OPENMP
169+
#pragma omp parallel
170+
#endif
171+
{
172+
std::vector<std::map<TA, std::map<TAC, RI::Tensor<Tdata>>>> mRs_a2D_thread(nspin);
173+
for (int is_k = 0; is_k < nspin_k.at(nspin); ++is_k)
176174
{
177-
const TC& cell = cells[icell];
178-
RI::Tensor<Tdata> mR_2D;
179-
int ik_full = 0;
180-
for (const int ik : ik_list)
175+
const std::vector<int> ik_list = RI_2D_Comm::get_ik_list(kv, is_k);
176+
const auto cells = RI_Util::get_Born_von_Karmen_cells(period);
177+
#pragma omp for schedule(dynamic)
178+
for (size_t icell = 0; icell < cells.size(); ++icell)
181179
{
182-
auto set_mR_2D = [&mR_2D](auto&& mk_frac)
180+
const TC& cell = cells[icell];
181+
RI::Tensor<Tdata> mR_2D;
182+
int ik_full = 0;
183+
for (const int ik : ik_list)
183184
{
184-
if (mR_2D.empty())
185-
{ mR_2D = RI::Global_Func::convert<Tdata>(mk_frac); }
186-
else
187-
{ mR_2D = mR_2D + RI::Global_Func::convert<Tdata>(mk_frac); }
188-
};
189-
using Tdata_m = typename Tmatrix::value_type;
190-
if (!spgsym)
191-
{
192-
RI::Tensor<Tdata_m> mk_2D = RI_Util::Vector_to_Tensor<Tdata_m>(*mks_2D[ik], pv.get_col_size(), pv.get_row_size());
193-
const Tdata_m frac = SPIN_multiple
194-
* RI::Global_Func::convert<Tdata_m>(std::exp(
195-
-ModuleBase::TWO_PI * ModuleBase::IMAG_UNIT * (kv.kvec_c[ik] * (RI_Util::array3_to_Vector3(cell) * ucell.latvec))));
196-
if (static_cast<int>(std::round(SPIN_multiple * kv.wk[ik] * kv.get_nkstot_full())) == 2)
197-
{ set_mR_2D(mk_2D * (frac * 0.5) + tensor_conj(mk_2D * (frac * 0.5))); }
198-
else
199-
{ set_mR_2D(mk_2D * frac); }
200-
}
201-
else
202-
{ // traverse kstar, ik means ik_ibz
203-
for (auto& isym_kvd : kv.kstars[ik % ik_list.size()])
185+
auto set_mR_2D = [&mR_2D](auto&& mk_frac)
186+
{
187+
if (mR_2D.empty())
188+
{ mR_2D = RI::Global_Func::convert<Tdata>(mk_frac); }
189+
else
190+
{ mR_2D = mR_2D + RI::Global_Func::convert<Tdata>(mk_frac); }
191+
};
192+
using Tdata_m = typename Tmatrix::value_type;
193+
if (!spgsym)
204194
{
205-
RI::Tensor<Tdata_m> mk_2D = RI_Util::Vector_to_Tensor<Tdata_m>(*mks_2D[ik_full + is_k * kv.get_nkstot_full()], pv.get_col_size(), pv.get_row_size());
195+
RI::Tensor<Tdata_m> mk_2D = RI_Util::Vector_to_Tensor<Tdata_m>(*mks_2D[ik], pv.get_col_size(), pv.get_row_size());
206196
const Tdata_m frac = SPIN_multiple
207197
* RI::Global_Func::convert<Tdata_m>(std::exp(
208-
-ModuleBase::TWO_PI * ModuleBase::IMAG_UNIT * ((isym_kvd.second * ucell.G) * (RI_Util::array3_to_Vector3(cell) * ucell.latvec))));
209-
set_mR_2D(mk_2D * frac);
210-
++ik_full;
198+
-ModuleBase::TWO_PI * ModuleBase::IMAG_UNIT * (kv.kvec_c[ik] * (RI_Util::array3_to_Vector3(cell) * ucell.latvec))));
199+
if (static_cast<int>(std::round(SPIN_multiple * kv.wk[ik] * kv.get_nkstot_full())) == 2)
200+
{ set_mR_2D(mk_2D * (frac * 0.5) + tensor_conj(mk_2D * (frac * 0.5))); }
201+
else
202+
{ set_mR_2D(mk_2D * frac); }
211203
}
212-
}
213-
}
214-
for(int iwt0_2D=0; iwt0_2D!=mR_2D.shape[0]; ++iwt0_2D)
215-
{
216-
const int iwt0 =ModuleBase::GlobalFunc::IS_COLUMN_MAJOR_KS_SOLVER(PARAM.inp.ks_solver)
217-
? pv.local2global_col(iwt0_2D)
218-
: pv.local2global_row(iwt0_2D);
219-
int iat0, iw0_b, is0_b;
220-
std::tie(iat0,iw0_b,is0_b) = RI_2D_Comm::get_iat_iw_is_block(ucell,iwt0);
221-
const int it0 = ucell.iat2it[iat0];
222-
for(int iwt1_2D=0; iwt1_2D!=mR_2D.shape[1]; ++iwt1_2D)
223-
{
224-
const int iwt1 =ModuleBase::GlobalFunc::IS_COLUMN_MAJOR_KS_SOLVER(PARAM.inp.ks_solver)
225-
? pv.local2global_row(iwt1_2D)
226-
: pv.local2global_col(iwt1_2D);
227-
int iat1, iw1_b, is1_b;
228-
std::tie(iat1,iw1_b,is1_b) = RI_2D_Comm::get_iat_iw_is_block(ucell,iwt1);
229-
const int it1 = ucell.iat2it[iat1];
230-
231-
const int is_b = RI_2D_Comm::get_is_block(is_k, is0_b, is1_b);
232-
#ifdef _OPENMP
233-
#pragma omp critical(RI_split_m2D_ktoR_general)
234-
#endif
204+
else
205+
{ // traverse kstar, ik means ik_ibz
206+
for (auto& isym_kvd : kv.kstars[ik % ik_list.size()])
207+
{
208+
RI::Tensor<Tdata_m> mk_2D = RI_Util::Vector_to_Tensor<Tdata_m>(*mks_2D[ik_full + is_k * kv.get_nkstot_full()], pv.get_col_size(), pv.get_row_size());
209+
const Tdata_m frac = SPIN_multiple
210+
* RI::Global_Func::convert<Tdata_m>(std::exp(
211+
-ModuleBase::TWO_PI * ModuleBase::IMAG_UNIT * ((isym_kvd.second * ucell.G) * (RI_Util::array3_to_Vector3(cell) * ucell.latvec))));
212+
set_mR_2D(mk_2D * frac);
213+
++ik_full;
214+
}
215+
}
216+
} // end for ik
217+
for(int iwt0_2D=0; iwt0_2D!=mR_2D.shape[0]; ++iwt0_2D)
218+
{
219+
const int iwt0 =ModuleBase::GlobalFunc::IS_COLUMN_MAJOR_KS_SOLVER(PARAM.inp.ks_solver)
220+
? pv.local2global_col(iwt0_2D)
221+
: pv.local2global_row(iwt0_2D);
222+
int iat0, iw0_b, is0_b;
223+
std::tie(iat0,iw0_b,is0_b) = RI_2D_Comm::get_iat_iw_is_block(ucell,iwt0);
224+
const int it0 = ucell.iat2it[iat0];
225+
for(int iwt1_2D=0; iwt1_2D!=mR_2D.shape[1]; ++iwt1_2D)
235226
{
236-
RI::Tensor<Tdata>& mR_a2D = mRs_a2D[is_b][iat0][{iat1, cell}];
227+
const int iwt1 =ModuleBase::GlobalFunc::IS_COLUMN_MAJOR_KS_SOLVER(PARAM.inp.ks_solver)
228+
? pv.local2global_row(iwt1_2D)
229+
: pv.local2global_col(iwt1_2D);
230+
int iat1, iw1_b, is1_b;
231+
std::tie(iat1,iw1_b,is1_b) = RI_2D_Comm::get_iat_iw_is_block(ucell,iwt1);
232+
const int it1 = ucell.iat2it[iat1];
233+
234+
const int is_b = RI_2D_Comm::get_is_block(is_k, is0_b, is1_b);
235+
RI::Tensor<Tdata>& mR_a2D = mRs_a2D_thread[is_b][iat0][{iat1, cell}];
237236
if (mR_a2D.empty())
238237
{
239238
mR_a2D = RI::Tensor<Tdata>(
240239
{static_cast<size_t>(ucell.atoms[it0].nw),
241240
static_cast<size_t>(ucell.atoms[it1].nw)});
242241
}
243242
mR_a2D(iw0_b, iw1_b) = mR_2D(iwt0_2D, iwt1_2D);
243+
} // for iwt1_2D
244+
} // end for iwt0_2D
245+
} // end for icell
246+
} // end for is_k
247+
248+
#ifdef _OPENMP
249+
#pragma omp critical
250+
#endif
251+
{
252+
for(int is=0; is<nspin; ++is)
253+
for(auto &mRs_A : mRs_a2D_thread[is])
254+
for(auto &mRs_B : mRs_A.second)
255+
{
256+
assert(mRs_a2D[is][mRs_A.first][mRs_B.first].empty());
257+
mRs_a2D[is][mRs_A.first][mRs_B.first] = std::move(mRs_B.second);
244258
}
245-
}
246-
}
247259
}
248-
}
249-
ModuleBase::timer::end("RI_2D_Comm", "split_m2D_ktoR_general");
260+
} // end #pragma omp parallel
261+
ModuleBase::timer::end("RI_2D_Comm", "split_m2D_ktoR_k");
250262
return mRs_a2D;
251263
}
252264

0 commit comments

Comments
 (0)