@@ -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
151151template <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