@@ -40,7 +40,33 @@ void Setup_Psi_pw::before_runner_impl(
4040 }
4141
4242 if (inp.device == " gpu" || inp.precision == " single" ) {
43+ const int nks = kv.get_nks ();
44+ psi::PsiStorageMode target_mode = psi::PsiStorageMode::ALL_GPU ;
45+ if (inp.device_memory_mode == " paged" )
46+ {
47+ target_mode = psi::PsiStorageMode::PAGED_GPU ;
48+ }
49+ else if (inp.device_memory_mode == " full_gpu" )
50+ {
51+ target_mode = psi::PsiStorageMode::ALL_GPU ;
52+ }
53+ else if (nks > 10 )
54+ {
55+ target_mode = psi::PsiStorageMode::PAGED_GPU ;
56+ }
57+ if (target_mode != psi::PsiStorageMode::ALL_GPU )
58+ {
59+ this ->psi_cpu ->set_storage_mode (target_mode);
60+ }
61+ const char * mode_str = (target_mode == psi::PsiStorageMode::PAGED_GPU ) ? " PAGED_GPU" : " ALL_GPU" ;
62+ GlobalV::ofs_running << " GPU memory mode for Psi: " << mode_str
63+ << " (nks=" << nks << " , device_memory_mode=\" "
64+ << inp.device_memory_mode << " \" )" << std::endl;
4365 this ->psi_t = static_cast <void *>(new psi::Psi<T, Device>(this ->psi_cpu [0 ]));
66+ if (target_mode != psi::PsiStorageMode::ALL_GPU )
67+ {
68+ this ->psi_cpu ->set_storage_mode (psi::PsiStorageMode::ALL_GPU );
69+ }
4470 } else {
4571 this ->psi_t = static_cast <void *>(reinterpret_cast <psi::Psi<T, Device>*>(this ->psi_cpu ));
4672 }
@@ -178,9 +204,19 @@ template <typename T, typename Device>
178204void Setup_Psi_pw::copy_d2h_impl ()
179205{
180206 auto * psi_t = this ->get_psi_t <T, Device>();
181- this ->castmem_d2h_impl <T, Device>(this ->psi_cpu [0 ].get_pointer () - this ->psi_cpu [0 ].get_psi_bias (),
182- psi_t ->get_pointer () - psi_t ->get_psi_bias (),
183- this ->psi_cpu [0 ].size ());
207+ if (psi_t ->get_storage_mode () == psi::PsiStorageMode::PAGED_GPU )
208+ {
209+ // PAGED_GPU: full data is on CPU in psi_cpu_, just memcpy
210+ const size_t total_size = sizeof (T) * psi_t ->get_nk () * psi_t ->get_nbands () * psi_t ->get_nbasis ();
211+ std::memcpy (this ->psi_cpu [0 ].get_pointer () - this ->psi_cpu [0 ].get_psi_bias (),
212+ psi_t ->get_cpu_pointer (), total_size);
213+ }
214+ else
215+ {
216+ this ->castmem_d2h_impl <T, Device>(this ->psi_cpu [0 ].get_pointer () - this ->psi_cpu [0 ].get_psi_bias (),
217+ psi_t ->get_pointer () - psi_t ->get_psi_bias (),
218+ this ->psi_cpu [0 ].size ());
219+ }
184220}
185221
186222void Setup_Psi_pw::copy_d2h ()
0 commit comments