Skip to content

Commit 4ded701

Browse files
committed
refactor(io): complete phase-2 schemas and transforms migration
1 parent 9497859 commit 4ded701

7 files changed

Lines changed: 209 additions & 90 deletions

File tree

deepks/core/ml/eval/evaluator.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,7 @@
77
import deepks
88
except ImportError as e:
99
sys.path.append(os.path.dirname(os.path.realpath(__file__)) + "/../../")
10-
from deepks.io.readers.group_reader import generalized_eigh
10+
from deepks.io.transforms.linalg import generalized_eigh
1111
from deepks.model.utils import get_density_matrix, cal_phi_loss, cal_v_delta, get_occ_func, make_loss, get_gedm, cal_vdr, loss_hr
1212

1313
class Evaluator:

deepks/io/readers/group_reader.py

Lines changed: 69 additions & 87 deletions
Original file line numberDiff line numberDiff line change
@@ -1,89 +1,75 @@
1-
import os,time,sys
1+
import os
2+
import sys
23
import numpy as np
34
import torch
4-
from deepks.model.utils import make_integrator, cal_nb_overlap
5-
6-
def concat_batch(tdicts, dim=0):
7-
keys = tdicts[0].keys()
8-
assert all(d.keys() == keys for d in tdicts)
9-
return {
10-
k: torch.cat([d[k] for d in tdicts], dim)
11-
for k in keys
12-
}
13-
14-
def split_batch(tdict, size, dim=0, global_keys=None):
15-
if global_keys is None:
16-
global_keys = {"data_shape"}
17-
dsplit = {}
18-
for k,v in tdict.items():
19-
if k in global_keys:
20-
dsplit[k] = v
21-
elif isinstance(v, torch.Tensor):
22-
dsplit[k] = torch.split(v, size, dim)
23-
elif isinstance(v, np.ndarray):
24-
# support only for dim=0
25-
assert dim == 0, "numpy.ndarray supports only for dim=0 split"
26-
dsplit[k] = np.array_split(v, range(size, v.shape[0], size), axis=0)
27-
elif isinstance(v, list):
28-
# support only for dim=0
29-
assert dim == 0, "list supports only for dim=0 split"
30-
dsplit[k] = [v[i:i+size] for i in range(0, len(v), size)]
31-
else:
32-
raise TypeError(f"Unsupported type for split_batch: {type(v)}")
33-
# dsplit = {k: torch.split(v, size, dim) for k,v in tdict.items()}
34-
nsecs = [len(v) for k, v in dsplit.items() if k not in global_keys]
35-
assert all(ns == nsecs[0] for ns in nsecs)
36-
return [
37-
{k: (v[i] if k not in global_keys else v) for k, v in dsplit.items()}
38-
for i in range(nsecs[0])
39-
]
405

41-
def generalized_eigh(h,L_inv):
42-
symm_h=L_inv @ h @ L_inv.mT
43-
e,v=torch.linalg.eigh(symm_h)
44-
phi=L_inv.mT @ v
45-
return e,phi
6+
from deepks.io.schemas.reader_fields import (
7+
DEFAULT_READER_FIELD_NAMES,
8+
ReaderFieldNames,
9+
resolve_reader_paths,
10+
)
11+
from deepks.io.transforms.batch import concat_batch, split_batch
12+
from deepks.io.transforms.linalg import generalized_eigh
13+
from deepks.model.utils import make_integrator, cal_nb_overlap
4614

4715
class Reader(object):
4816
def __init__(self, data_path, batch_size,
49-
e_name="l_e_delta", d_name="dm_eig",
50-
f_name="l_f_delta", gvx_name="grad_vx",
51-
s_name="l_s_delta", gvepsl_name="grad_vepsl",
52-
o_name="l_o_delta", op_name="orbital_precalc",
53-
h_name="l_h_delta", vdp_name="v_delta_precalc",
54-
vdrp_name="vdr_precalc", phialpha_name="phialpha",
55-
gevdm_name="grad_evdm", hr_name="l_hr_delta",
56-
h_base_name="h_base", h_ref_name="hamiltonian",
57-
read_overlap = False, overlap_name="overlap",
58-
eg_name="eg_base", gveg_name="grad_veg",
59-
gldv_name="grad_ldv", conv_name="conv",
60-
atom_name="atom", box_name="box",
17+
e_name=DEFAULT_READER_FIELD_NAMES.e_name,
18+
d_name=DEFAULT_READER_FIELD_NAMES.d_name,
19+
f_name=DEFAULT_READER_FIELD_NAMES.f_name,
20+
gvx_name=DEFAULT_READER_FIELD_NAMES.gvx_name,
21+
s_name=DEFAULT_READER_FIELD_NAMES.s_name,
22+
gvepsl_name=DEFAULT_READER_FIELD_NAMES.gvepsl_name,
23+
o_name=DEFAULT_READER_FIELD_NAMES.o_name,
24+
op_name=DEFAULT_READER_FIELD_NAMES.op_name,
25+
h_name=DEFAULT_READER_FIELD_NAMES.h_name,
26+
vdp_name=DEFAULT_READER_FIELD_NAMES.vdp_name,
27+
vdrp_name=DEFAULT_READER_FIELD_NAMES.vdrp_name,
28+
phialpha_name=DEFAULT_READER_FIELD_NAMES.phialpha_name,
29+
gevdm_name=DEFAULT_READER_FIELD_NAMES.gevdm_name,
30+
hr_name=DEFAULT_READER_FIELD_NAMES.hr_name,
31+
h_base_name=DEFAULT_READER_FIELD_NAMES.h_base_name,
32+
h_ref_name=DEFAULT_READER_FIELD_NAMES.h_ref_name,
33+
read_overlap=False,
34+
overlap_name=DEFAULT_READER_FIELD_NAMES.overlap_name,
35+
eg_name=DEFAULT_READER_FIELD_NAMES.eg_name,
36+
gveg_name=DEFAULT_READER_FIELD_NAMES.gveg_name,
37+
gldv_name=DEFAULT_READER_FIELD_NAMES.gldv_name,
38+
conv_name=DEFAULT_READER_FIELD_NAMES.conv_name,
39+
atom_name=DEFAULT_READER_FIELD_NAMES.atom_name,
40+
box_name=DEFAULT_READER_FIELD_NAMES.box_name,
6141
orb_list=None, alpha_list=None, **kwargs):
6242
self.data_path = data_path
6343
self.batch_size = batch_size
64-
self.e_path = self.check_exist(e_name+".npy")
65-
self.f_path = self.check_exist(f_name+".npy")
66-
self.s_path = self.check_exist(s_name+".npy")
67-
self.o_path = self.check_exist(o_name+".npy")
68-
self.h_path = self.check_exist(h_name+".npy")
69-
self.hr_path = self.check_exist(hr_name+".npy")
70-
self.h_base_path = self.check_exist(h_base_name+".npy")
71-
self.h_ref_path = self.check_exist(h_ref_name+".npy")
72-
self.overlap_path = self.check_exist(overlap_name+".npy")
73-
self.phialpha_path = self.check_exist(phialpha_name+".npy")
74-
self.gevdm_path = self.check_exist(gevdm_name+".npy")
75-
self.d_path = self.check_exist(d_name+".npy")
76-
self.gvx_path = self.check_exist(gvx_name+".npy")
77-
self.gvepsl_path = self.check_exist(gvepsl_name+".npy")
78-
self.op_path = self.check_exist(op_name+".npy")
79-
self.vdp_path = self.check_exist(vdp_name+".npy")
80-
self.vdrp_path = self.check_exist(vdrp_name+".npy")
81-
self.eg_path = self.check_exist(eg_name+".npy")
82-
self.gveg_path = self.check_exist(gveg_name+".npy")
83-
self.gldv_path = self.check_exist(gldv_name+".npy")
84-
self.c_path = self.check_exist(conv_name+".npy")
85-
self.a_path = self.check_exist(atom_name+".npy")
86-
self.b_path = self.check_exist(box_name+".npy")
44+
field_names = ReaderFieldNames(
45+
e_name=e_name,
46+
d_name=d_name,
47+
f_name=f_name,
48+
gvx_name=gvx_name,
49+
s_name=s_name,
50+
gvepsl_name=gvepsl_name,
51+
o_name=o_name,
52+
op_name=op_name,
53+
h_name=h_name,
54+
vdp_name=vdp_name,
55+
vdrp_name=vdrp_name,
56+
phialpha_name=phialpha_name,
57+
gevdm_name=gevdm_name,
58+
hr_name=hr_name,
59+
h_base_name=h_base_name,
60+
h_ref_name=h_ref_name,
61+
overlap_name=overlap_name,
62+
eg_name=eg_name,
63+
gveg_name=gveg_name,
64+
gldv_name=gldv_name,
65+
conv_name=conv_name,
66+
atom_name=atom_name,
67+
box_name=box_name,
68+
)
69+
for path_name, path in resolve_reader_paths(self.data_path, field_names).items():
70+
setattr(self, path_name, path)
71+
72+
self.system_raw_path = os.path.join(self.data_path, "system.raw")
8773
self.read_overlap = read_overlap
8874
self.orb_list = ["../../" + orb for orb in orb_list] if orb_list is not None else None
8975
self.alpha_list = ["../../" + alpha for alpha in alpha_list] if alpha_list is not None else None
@@ -93,16 +79,9 @@ def __init__(self, data_path, batch_size,
9379
# initialize sample index queue
9480
self.idx_queue = []
9581

96-
def check_exist(self, fname):
97-
if fname is None:
98-
return None
99-
fpath = os.path.join(self.data_path, fname)
100-
if os.path.exists(fpath):
101-
return fpath
102-
10382
def load_meta(self):
10483
try:
105-
sys_meta = np.loadtxt(self.check_exist('system.raw'), converters = float).astype(int).reshape([-1])
84+
sys_meta = np.loadtxt(self.system_raw_path, converters=float).astype(int).reshape([-1])
10685
self.natm = sys_meta[0]
10786
self.nproj = sys_meta[-1]
10887
except:
@@ -526,8 +505,11 @@ def revert_elem_const(self):
526505

527506
class SimpleReader(object):
528507
def __init__(self, data_path, batch_size,
529-
e_name="l_e_delta", d_name="dm_eig",
530-
conv_filter=True, conv_name="conv", **kwargs):
508+
e_name=DEFAULT_READER_FIELD_NAMES.e_name,
509+
d_name=DEFAULT_READER_FIELD_NAMES.d_name,
510+
conv_filter=True,
511+
conv_name=DEFAULT_READER_FIELD_NAMES.conv_name,
512+
**kwargs):
531513
# copy from config
532514
self.data_path = data_path
533515
self.batch_size = batch_size

deepks/io/schemas/__init__.py

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1,3 @@
1-
"""Scaffold package for refactor architecture."""
1+
"""Schema definitions for DeepKS I/O layer."""
2+
3+
from .reader_fields import * # noqa: F401,F403

deepks/io/schemas/reader_fields.py

Lines changed: 78 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,78 @@
1+
"""Reader file-name schema utilities for DeePKS data loading."""
2+
3+
from dataclasses import asdict, dataclass
4+
import os
5+
from typing import Dict, Optional
6+
7+
8+
@dataclass(frozen=True)
9+
class ReaderFieldNames:
10+
e_name: str = "l_e_delta"
11+
d_name: str = "dm_eig"
12+
f_name: str = "l_f_delta"
13+
gvx_name: str = "grad_vx"
14+
s_name: str = "l_s_delta"
15+
gvepsl_name: str = "grad_vepsl"
16+
o_name: str = "l_o_delta"
17+
op_name: str = "orbital_precalc"
18+
h_name: str = "l_h_delta"
19+
vdp_name: str = "v_delta_precalc"
20+
vdrp_name: str = "vdr_precalc"
21+
phialpha_name: str = "phialpha"
22+
gevdm_name: str = "grad_evdm"
23+
hr_name: str = "l_hr_delta"
24+
h_base_name: str = "h_base"
25+
h_ref_name: str = "hamiltonian"
26+
overlap_name: str = "overlap"
27+
eg_name: str = "eg_base"
28+
gveg_name: str = "grad_veg"
29+
gldv_name: str = "grad_ldv"
30+
conv_name: str = "conv"
31+
atom_name: str = "atom"
32+
box_name: str = "box"
33+
34+
35+
DEFAULT_READER_FIELD_NAMES = ReaderFieldNames()
36+
37+
38+
READER_PATH_ATTR_MAP = {
39+
"e_name": "e_path",
40+
"d_name": "d_path",
41+
"f_name": "f_path",
42+
"gvx_name": "gvx_path",
43+
"s_name": "s_path",
44+
"gvepsl_name": "gvepsl_path",
45+
"o_name": "o_path",
46+
"op_name": "op_path",
47+
"h_name": "h_path",
48+
"vdp_name": "vdp_path",
49+
"vdrp_name": "vdrp_path",
50+
"phialpha_name": "phialpha_path",
51+
"gevdm_name": "gevdm_path",
52+
"hr_name": "hr_path",
53+
"h_base_name": "h_base_path",
54+
"h_ref_name": "h_ref_path",
55+
"overlap_name": "overlap_path",
56+
"eg_name": "eg_path",
57+
"gveg_name": "gveg_path",
58+
"gldv_name": "gldv_path",
59+
"conv_name": "c_path",
60+
"atom_name": "a_path",
61+
"box_name": "b_path",
62+
}
63+
64+
65+
def resolve_numpy_path(data_path: str, stem: Optional[str]) -> Optional[str]:
66+
if stem is None:
67+
return None
68+
fpath = os.path.join(data_path, f"{stem}.npy")
69+
if os.path.exists(fpath):
70+
return fpath
71+
return None
72+
73+
74+
def resolve_reader_paths(data_path: str, names: ReaderFieldNames) -> Dict[str, Optional[str]]:
75+
return {
76+
READER_PATH_ATTR_MAP[name]: resolve_numpy_path(data_path, stem)
77+
for name, stem in asdict(names).items()
78+
}

deepks/io/transforms/__init__.py

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1 +1,4 @@
1-
"""Scaffold package for refactor architecture."""
1+
"""Data transformation helpers for DeepKS I/O layer."""
2+
3+
from .batch import * # noqa: F401,F403
4+
from .linalg import * # noqa: F401,F403

deepks/io/transforms/batch.py

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
"""Batch-level tensor/list splitting and concatenation helpers."""
2+
3+
from typing import Any, Dict, Iterable, Optional, Set
4+
5+
import numpy as np
6+
import torch
7+
8+
9+
def concat_batch(tdicts: Iterable[Dict[str, Any]], dim: int = 0) -> Dict[str, Any]:
10+
tdicts = list(tdicts)
11+
keys = tdicts[0].keys()
12+
assert all(d.keys() == keys for d in tdicts)
13+
return {k: torch.cat([d[k] for d in tdicts], dim) for k in keys}
14+
15+
16+
def split_batch(
17+
tdict: Dict[str, Any],
18+
size: int,
19+
dim: int = 0,
20+
global_keys: Optional[Set[str]] = None,
21+
):
22+
if global_keys is None:
23+
global_keys = {"data_shape"}
24+
dsplit = {}
25+
for k, v in tdict.items():
26+
if k in global_keys:
27+
dsplit[k] = v
28+
elif isinstance(v, torch.Tensor):
29+
dsplit[k] = torch.split(v, size, dim)
30+
elif isinstance(v, np.ndarray):
31+
assert dim == 0, "numpy.ndarray supports only for dim=0 split"
32+
dsplit[k] = np.array_split(v, range(size, v.shape[0], size), axis=0)
33+
elif isinstance(v, list):
34+
assert dim == 0, "list supports only for dim=0 split"
35+
dsplit[k] = [v[i : i + size] for i in range(0, len(v), size)]
36+
else:
37+
raise TypeError(f"Unsupported type for split_batch: {type(v)}")
38+
39+
nsecs = [len(v) for k, v in dsplit.items() if k not in global_keys]
40+
assert all(ns == nsecs[0] for ns in nsecs)
41+
return [
42+
{k: (v[i] if k not in global_keys else v) for k, v in dsplit.items()}
43+
for i in range(nsecs[0])
44+
]

deepks/io/transforms/linalg.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,10 @@
1+
"""Linear algebra helpers shared by reader/evaluator modules."""
2+
3+
import torch
4+
5+
6+
def generalized_eigh(h, l_inv):
7+
symm_h = l_inv @ h @ l_inv.mT
8+
e, v = torch.linalg.eigh(symm_h)
9+
phi = l_inv.mT @ v
10+
return e, phi

0 commit comments

Comments
 (0)