#include "psi.h"
#include "source_base/global_variable.h"
#include "source_base/module_device/device.h"
#include "source_base/tool_quit.h"
#include "source_io/module_parameter/parameter.h"
#include
#include
#include
namespace psi
{
Range::Range(const size_t range_in)
{
k_first = true;
index_1 = 0;
range_1 = range_in;
range_2 = range_in;
}
Range::Range(const bool k_first_in, const size_t index_1_in, const size_t range_1_in, const size_t range_2_in)
{
k_first = k_first_in;
index_1 = index_1_in;
range_1 = range_1_in;
range_2 = range_2_in;
}
// Constructor 0: basic
template
Psi::Psi()
{
}
template
Psi::~Psi()
{
if (this->allocate_inside)
{
delete_memory_op()(this->psi);
}
}
// Constructor 1:
template
Psi::Psi(const int nk_in,
const int nbd_in,
const int nbs_in,
const std::vector& ngk_in,
const bool k_first_in)
{
assert(nk_in > 0);
assert(nbd_in >= 0);
assert(nbs_in > 0);
this->k_first = k_first_in;
this->allocate_inside = true;
this->ngk = ngk_in.data(); // modify later
// This function will delete the psi array first(if psi exist), then malloc a new memory for it.
resize_memory_op()(this->psi, nk_in * static_cast(nbd_in) * nbs_in, "no_record");
this->nk = nk_in;
this->nbands = nbd_in;
this->nbasis = nbs_in;
this->current_b = 0;
this->current_k = 0;
this->current_nbasis = nbs_in;
this->psi_current = this->psi;
this->psi_bias = 0;
// Currently only GPU's implementation is supported for device recording!
base_device::information::print_device_info(this->ctx, GlobalV::ofs_device);
base_device::information::record_device_memory(this->ctx,
GlobalV::ofs_device,
"Psi->resize()",
sizeof(T) * nk_in * nbd_in * nbs_in);
}
// Constructor 3-1: 2D Psi version
template
Psi::Psi(T* psi_pointer,
const int nk_in,
const int nbd_in,
const int nbs_in,
const int current_nbasis_in,
const bool k_first_in)
{
// Currently this function only supports nk_in == 1 when called within diagH_subspace_init.
// assert(nk_in == 1); // NOTE because lr/utils/lr_uril.hpp func & get_psi_spin func
this->k_first = k_first_in;
this->allocate_inside = false;
this->ngk = nullptr;
this->psi = psi_pointer;
this->nk = nk_in;
this->nbands = nbd_in;
this->nbasis = nbs_in;
this->current_k = 0;
this->current_b = 0;
this->current_nbasis = current_nbasis_in;
this->psi_current = psi_pointer;
this->psi_bias = 0;
// Currently only GPU's implementation is supported for device recording!
base_device::information::print_device_info(this->ctx, GlobalV::ofs_device);
}
// Constructor 3-2: 2D Psi version
template
Psi::Psi(const int nk_in,
const int nbd_in,
const int nbs_in,
const int current_nbasis_in,
const bool k_first_in)
{
// Currently this function only supports nk_in == 1 when called within diagH_subspace_init.
// assert(nk_in == 1);
this->k_first = k_first_in;
this->allocate_inside = true;
this->ngk = nullptr;
assert(nk_in > 0 && nbd_in >= 0 && nbs_in > 0);
resize_memory_op()(this->psi, nk_in * static_cast(nbd_in) * nbs_in, "no_record");
this->nk = nk_in;
this->nbands = nbd_in;
this->nbasis = nbs_in;
this->current_k = 0;
this->current_b = 0;
this->current_nbasis = current_nbasis_in;
this->psi_current = this->psi;
this->psi_bias = 0;
// Currently only GPU's implementation is supported for device recording!
base_device::information::print_device_info(this->ctx, GlobalV::ofs_device);
base_device::information::record_device_memory(this->ctx,
GlobalV::ofs_device,
"Psi->resize()",
sizeof(T) * nk_in * nbd_in * nbs_in);
}
// Constructor 2-1:
template
Psi::Psi(const Psi& psi_in)
{
this->ngk = psi_in.ngk;
this->nk = psi_in.get_nk();
this->nbands = psi_in.get_nbands();
this->nbasis = psi_in.get_nbasis();
this->current_k = psi_in.get_current_k();
this->current_b = psi_in.get_current_b();
this->k_first = psi_in.get_k_first();
// this function will copy psi_in.psi to this->psi no matter the device types of each other.
this->resize(psi_in.get_nk(), psi_in.get_nbands(), psi_in.get_nbasis());
base_device::memory::synchronize_memory_op()(this->psi,
psi_in.get_pointer() - psi_in.get_psi_bias(),
psi_in.size());
this->psi_bias = psi_in.get_psi_bias();
this->current_nbasis = psi_in.get_current_nbas();
this->psi_current = this->psi + psi_in.get_psi_bias();
}
// Constructor 2-2:
template
template
Psi::Psi(const Psi& psi_in)
{
this->ngk = psi_in.get_ngk_pointer();
this->nk = psi_in.get_nk();
this->nbands = psi_in.get_nbands();
this->nbasis = psi_in.get_nbasis();
this->current_k = psi_in.get_current_k();
this->current_b = psi_in.get_current_b();
this->k_first = psi_in.get_k_first();
// this function will copy psi_in.psi to this->psi no matter the device types of each other.
this->resize(psi_in.get_nk(), psi_in.get_nbands(), psi_in.get_nbasis());
// Specifically, if the Device_in type is CPU and the Device type is GPU.
// Which means we need to initialize a GPU psi from a given CPU psi.
// We first malloc a memory in CPU, then cast the memory from T_in to T in CPU.
// Finally, synchronize the memory from CPU to GPU.
// This could help to reduce the peak memory usage of device.
if (std::is_same::value && std::is_same::value)
{
auto* arr = (T*)malloc(sizeof(T) * psi_in.size());
// cast the memory from T_in to T in CPU
base_device::memory::cast_memory_op()(arr,
psi_in.get_pointer()
- psi_in.get_psi_bias(),
psi_in.size());
// synchronize the memory from CPU to GPU
base_device::memory::synchronize_memory_op()(this->psi,
arr,
psi_in.size());
free(arr);
}
else
{
base_device::memory::cast_memory_op()(this->psi,
psi_in.get_pointer() - psi_in.get_psi_bias(),
psi_in.size());
}
this->psi_bias = psi_in.get_psi_bias();
this->current_nbasis = psi_in.get_current_nbas();
this->psi_current = this->psi + psi_in.get_psi_bias();
}
template
void Psi::set_all_psi(const T* another_pointer, const std::size_t size_in)
{
assert(size_in == this->size());
synchronize_memory_op()(this->psi, another_pointer, this->size());
}
template
Psi& Psi::operator=(const Psi& psi_in)
{
// printf("%d\n", &psi_in);
this->ngk = psi_in.ngk;
this->nk = psi_in.get_nk();
this->nbands = psi_in.get_nbands();
this->nbasis = psi_in.get_nbasis();
this->current_k = psi_in.get_current_k();
this->current_b = psi_in.get_current_b();
this->k_first = psi_in.get_k_first();
// this function will copy psi_in.psi to this->psi no matter the device types of each other.
this->resize(psi_in.get_nk(), psi_in.get_nbands(), psi_in.get_nbasis());
base_device::memory::synchronize_memory_op()(this->psi,
psi_in.psi,
psi_in.size());
this->psi_bias = psi_in.get_psi_bias();
this->current_nbasis = psi_in.get_current_nbas();
this->psi_current = this->psi + psi_in.get_psi_bias();
return *this;
}
template
void Psi::resize(const int nks_in, const int nbands_in, const int nbasis_in)
{
assert(nks_in > 0 && nbands_in >= 0 && nbasis_in > 0);
// This function will delete the psi array first(if psi exist), then malloc a new memory for it.
resize_memory_op()(this->psi, nks_in * static_cast(nbands_in) * nbasis_in, "no_record");
// this->zero_out();
this->nk = nks_in;
this->nbands = nbands_in;
this->nbasis = nbasis_in;
this->current_nbasis = nbasis_in;
this->psi_current = this->psi;
// GlobalV::ofs_device = 0);
assert(this->k_first ? ikb < this->nbands : ikb < this->nk);
return this->psi_current + ikb * this->nbasis;
}
template
const bool& Psi::get_k_first() const
{
return this->k_first;
}
template
const Device* Psi::get_device() const
{
return this->ctx;
}
template
const int* Psi::get_ngk_pointer() const
{
return this->ngk;
}
template
const size_t& Psi::get_psi_bias() const
{
return this->psi_bias;
}
template
const int& Psi::get_current_ngk() const
{
if (this->get_npol() == 1)
{
return this->current_nbasis;
}
else
{
return this->nbasis;
}
}
template
int Psi::get_npol() const
{
if (PARAM.inp.nspin == 4)
{
return 2;
}
else
{
return 1;
}
}
template
const int& Psi::get_nk() const
{
return this->nk;
}
template
const int& Psi::get_nbands() const
{
return this->nbands;
}
template
const int& Psi::get_nbasis() const
{
return this->nbasis;
}
template
std::size_t Psi::size() const
{
if (this->psi == nullptr)
{
return 0;
}
return this->nk * static_cast(this->nbands) * this->nbasis;
}
template
void Psi::fix_k(const int ik) const
{
assert(ik >= 0);
this->current_k = ik;
if (this->ngk != nullptr)
{
this->current_nbasis = this->ngk[ik];
}
else
{
this->current_nbasis = this->nbasis;
}
if (this->k_first)
{
this->current_b = 0;
}
int base = this->current_b * this->nk * this->nbasis;
if (ik >= this->nk)
{
// mem_saver: fix to base
this->psi_bias = base;
this->psi_current = const_cast(&(this->psi[base]));
}
else
{
this->psi_bias = k_first ? ik * this->nbands * this->nbasis : base + ik * this->nbasis;
this->psi_current = const_cast(&(this->psi[psi_bias]));
}
}
template
void Psi::fix_b(const int ib) const
{
assert(ib >= 0);
this->current_b = ib;
if (!this->k_first)
{
this->current_k = 0;
}
int base = this->current_k * this->nbands * this->nbasis;
if (ib >= this->nbands)
{
// mem_saver: fix to base
this->psi_bias = base;
this->psi_current = const_cast(&(this->psi[base]));
}
else
{
this->psi_bias = k_first ? base + ib * this->nbasis : ib * this->nk * this->nbasis;
this->psi_current = const_cast(&(this->psi[psi_bias]));
}
}
template
void Psi::fix_kb(const int ik, const int ib) const
{
assert(ik >= 0 && ib >= 0);
this->current_k = ik;
this->current_b = ib;
if (ik >= this->nk || ib >= this->nbands)
{ // fix to 0
this->psi_bias = 0;
this->psi_current = const_cast(&(this->psi[0]));
}
else
{
this->psi_bias = k_first ? (ik * this->nbands + ib) * this->nbasis : (ib * this->nk + ik) * this->nbasis;
this->psi_current = const_cast(&(this->psi[psi_bias]));
}
}
template
T& Psi::operator()(const int ikb1, const int ikb2, const int ibasis) const
{
assert(ikb1 >= 0 && ikb2 >= 0 && ibasis >= 0);
assert(this->k_first ? ikb1 < this->nk && ikb2 < this->nbands : ikb1 < this->nbands && ikb2 < this->nk);
return this->k_first ? this->psi[(ikb1 * this->nbands + ikb2) * this->nbasis + ibasis]
: this->psi[(ikb1 * this->nk + ikb2) * this->nbasis + ibasis];
}
template
T& Psi::operator()(const int ikb2, const int ibasis) const
{
assert(this->k_first ? this->current_b == 0 : this->current_k == 0);
assert(this->k_first ? ikb2 >= 0 && ikb2 < this->nbands : ikb2 >= 0 && ikb2 < this->nk);
assert(ibasis >= 0 && ibasis < this->nbasis);
return this->psi_current[ikb2 * this->nbasis + ibasis];
}
template
T& Psi::operator()(const int ibasis) const
{
assert(ibasis >= 0 && ibasis < this->nbasis);
return this->psi_current[ibasis];
}
template
int Psi::get_current_k() const
{
return this->current_k;
}
template
int Psi::get_current_b() const
{
return this->current_b;
}
template
int Psi::get_current_nbas() const
{
return this->current_nbasis;
}
template
const int& Psi::get_ngk(const int ik_in) const
{
assert(this->ngk != nullptr);
return this->ngk[ik_in];
}
template
void Psi::zero_out()
{
// this->psi.assign(this->psi.size(), T(0));
set_memory_op()(this->psi, 0, this->size());
}
template
std::tuple Psi::to_range(const Range& range) const
{
const int& i1 = range.index_1;
const int& r1 = range.range_1;
const int& r2 = range.range_2;
if (range.k_first != this->k_first || r1 < 0
|| r2 < r1
// || (range.k_first && (r2 >= this->nbands || i1 >= this->nk))
// || (!range.k_first && (r2 >= this->nk || i1 >= this->nbands)))
|| (range.k_first ? (i1 >= this->nk) : (i1 >= this->nbands)) // illegal index 1
|| (range.k_first ? (i1 > 0 && r2 >= this->nbands) : (i1 > 0 && r2 >= this->nk)) // illegal range of index 2
|| (range.k_first ? (i1 < 0 && r2 >= this->nk) : (i1 < 0 && r2 >= this->nbands))) // illegal range of index 1
{
return std::tuple(nullptr, 0);
}
else if (i1 < 0) // [r1, r2] is the range of index1 with length m
{
const T* p = &this->psi[r1 * (k_first ? this->nbands : this->nk) * this->nbasis];
int m = (r2 - r1 + 1) * this->get_npol();
return std::tuple(p, m);
}
else // [r1, r2] is the range of index2 with length m
{
const T* p = &this->psi[(i1 * (k_first ? this->nbands : this->nk) + r1) * this->nbasis];
int m = (r2 - r1 + 1) * this->get_npol();
return std::tuple(p, m);
}
}
template class Psi;
template class Psi;
template class Psi;
template class Psi;
template Psi::Psi(
const Psi&);
template Psi::Psi(
const Psi&);
#if ((defined __CUDA) || (defined __ROCM))
template class Psi;
template class Psi;
template Psi::Psi(const Psi&);
template Psi::Psi(const Psi&);
template Psi::Psi(
const Psi&);
template Psi::Psi(
const Psi&);
template class Psi;
template class Psi;
template Psi::Psi(const Psi&);
template Psi::Psi(const Psi&);
template Psi::Psi(
const Psi&);
template Psi::Psi(
const Psi&);
template Psi::Psi(
const Psi&);
template Psi::Psi(
const Psi&);
template Psi::Psi(
const Psi&);
#endif
} // namespace psi