Skip to content
Open
Show file tree
Hide file tree
Changes from 2 commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 8 additions & 0 deletions .gitignore
Original file line number Diff line number Diff line change
@@ -0,0 +1,8 @@
__pycache__/
*.pyc
*.pyo
*.egg-info/
*.egg
.DS_Store
dist/
build/
108 changes: 108 additions & 0 deletions deeph/compat_scatter.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,108 @@
"""
Drop-in replacement for torch_scatter.scatter and torch_scatter.scatter_add
using native PyTorch operations. This avoids the need to install torch_scatter.
"""
import torch


def broadcast_index(index: torch.Tensor, src: torch.Tensor, dim: int) -> torch.Tensor:
"""Broadcast index to match the dimensions of src for scatter operations."""
if index.dim() == src.dim():
return index
# index is 1D, src has more dimensions
# Expand index to match src shape
for _ in range(src.dim() - 1):
index = index.unsqueeze(-1)
# Expand to match src shape
shape = list(src.shape)
shape[dim] = -1
return index.expand(shape)


def scatter_add(src: torch.Tensor, index: torch.Tensor, dim: int = 0,
out: torch.Tensor = None, dim_size: int = None) -> torch.Tensor:
"""
Drop-in replacement for torch_scatter.scatter_add.
"""
if dim_size is None:
dim_size = int(index.max().item()) + 1

if out is None:
size = list(src.shape)
size[dim] = dim_size
out = torch.zeros(size, dtype=src.dtype, device=src.device)

index = broadcast_index(index, src, dim)
out.scatter_add_(dim, index, src)
return out


def scatter(src: torch.Tensor, index: torch.Tensor, dim: int = 0,
out: torch.Tensor = None, dim_size: int = None,
reduce: str = 'sum') -> torch.Tensor:
"""
Drop-in replacement for torch_scatter.scatter.
Supports reduce modes: 'sum', 'add', 'mean', 'min', 'max', 'mul'.
"""
if reduce == 'add':
reduce = 'sum'

if dim_size is None:
dim_size = int(index.max().item()) + 1

if out is None:
size = list(src.shape)
size[dim] = dim_size
out = torch.zeros(size, dtype=src.dtype, device=src.device)

if reduce == 'sum':
index = broadcast_index(index, src, dim)
out.scatter_add_(dim, index, src)
elif reduce == 'mean':
index = broadcast_index(index, src, dim)
out.scatter_add_(dim, index, src)
# Count elements per index for averaging
count = torch.zeros(dim_size, dtype=src.dtype, device=src.device)
ones = torch.ones(index.shape[0], dtype=src.dtype, device=src.device)
count.scatter_add_(0, index[:, 0], ones)
count = count.clamp(min=1)
# Reshape count for broadcasting
for _ in range(out.dim() - 1):
count = count.unsqueeze(-1)
out = out / count
elif reduce == 'min':
index = broadcast_index(index, src, dim)
out.fill_(float('inf'))
# Use scatter_reduce if available (PyTorch >= 1.12), otherwise manual
if hasattr(out, 'scatter_reduce_'):
out.scatter_reduce_(dim, index, src, reduce='amin')
else:
# Fallback: use a loop (slow but correct)
for i in range(index.shape[0]):
idx = index[i, 0].item()
out[idx] = torch.min(out[idx], src[i])
elif reduce == 'max':
index = broadcast_index(index, src, dim)
out.fill_(float('-inf'))
if hasattr(out, 'scatter_reduce_'):
out.scatter_reduce_(dim, index, src, reduce='amax')
else:
for i in range(index.shape[0]):
idx = index[i, 0].item()
out[idx] = torch.max(out[idx], src[i])
elif reduce == 'mul':
index = broadcast_index(index, src, dim)
out.fill_(1.0)
for i in range(index.shape[0]):
idx = index[i, 0].item()
out[idx] = out[idx] * src[i]
else:
raise ValueError(f"Unknown reduce mode: {reduce}")

return out


def scatter_mean(src: torch.Tensor, index: torch.Tensor, dim: int = 0,
out: torch.Tensor = None, dim_size: int = None) -> torch.Tensor:
"""Drop-in replacement for torch_scatter.scatter_mean."""
return scatter(src, index, dim=dim, out=out, dim_size=dim_size, reduce='mean')
51 changes: 34 additions & 17 deletions deeph/data.py
Original file line number Diff line number Diff line change
Expand Up @@ -105,7 +105,19 @@ def __init__(self, raw_data_dir: str, graph_dir: str, interface: str, target: st
loaded_data = torch.load(self.data_file)
except AttributeError:
raise RuntimeError('Error in loading graph data file, try to delete it and generate the graph file with the current version of PyG')
if len(loaded_data) == 2:
if isinstance(loaded_data, tuple) and len(loaded_data) == 2 and isinstance(loaded_data[0], list):
# New format: (data_list, info) β€” collate now at load time
data_list, tmp = loaded_data
print(f'Collating {len(data_list)} graph structures on load...')
import gc as _gc
_gc.collect()
self.data, self.slices = self.collate(data_list)
self.info = tmp
# Free data_list AFTER collation
del data_list
_gc.collect()
print(f'Collation done. Atomic types: {self.info["index_to_Z"].tolist()}')
elif len(loaded_data) == 2:
warnings.warn('You are using the graph data file with an old version')
self.data, self.slices = loaded_data
self.info = {
Expand Down Expand Up @@ -168,43 +180,48 @@ def process(self):
folder_list = folder_list[500:5000:3]
if self.dataset_name == 'bp_bilayer':
folder_list = folder_list[:600]
if 'C2N' in self.dataset_name:
# P0 optimization: use 600 structures (proven stable)
folder_list = folder_list[:600]
print(f'Dataset {self.dataset_name}: using {len(folder_list)} structures')
assert len(folder_list) != 0, "Can not find any structure"
print('Found %d structures, have cost %d seconds' % (len(folder_list), time.time() - begin))

import gc as _gc

# Phase 1: Process all structures using ONE pool (proven to reach 100%)
if self.multiprocessing == 0:
print(f'Use multiprocessing (nodes = num_processors x num_threads = 1 x {torch.get_num_threads()})')
data_list = [self.process_worker(folder) for folder in tqdm.tqdm(folder_list)]
print(f'Use multiprocessing (1 x {torch.get_num_threads()})')
data_list = [self.process_worker(f) for f in tqdm.tqdm(folder_list)]
else:
pool_dict = {} if self.multiprocessing < 0 else {'nodes': self.multiprocessing}
# BS (2023.06.06):
# The keyword "num_threads" in kernel.py can be used to set the torch threads.
# The multiprocessing in the "process_worker" is in contradiction with the num_threads utilized in torch.
# To avoid this conflict, I limit the number of torch threads to one,
# and recover it when finishing the process_worker.
torch_num_threads = torch.get_num_threads()
torch_num_threads_saved = torch.get_num_threads()
torch.set_num_threads(1)

with Pool(**pool_dict) as pool:
nodes = pool.nodes
print(f'Use multiprocessing (nodes = num_processors x num_threads = {nodes} x {torch.get_num_threads()})')
print(f'Use multiprocessing (nodes = {pool.nodes} x {torch.get_num_threads()})')
data_list = list(tqdm.tqdm(pool.imap(self.process_worker, folder_list), total=len(folder_list)))
torch.set_num_threads(torch_num_threads)
torch.set_num_threads(torch_num_threads_saved)

print('Finish processing %d structures, have cost %d seconds' % (len(data_list), time.time() - begin))

if self.pre_filter is not None:
data_list = [d for d in data_list if self.pre_filter(d)]
if self.pre_transform is not None:
data_list = [self.pre_transform(d) for d in data_list]

# Force GC to recover pool worker memory
_gc.collect()
print('Memory after GC: data_list has %d structures' % len(data_list))

# Compute global metadata β€” must process ALL data to remap element indices
index_to_Z, Z_to_index = self.element_statistics(data_list)
spinful = data_list[0].spinful
for d in data_list:
assert spinful == d.spinful

data, slices = self.collate(data_list)
torch.save((data, slices, dict(spinful=spinful, index_to_Z=index_to_Z, Z_to_index=Z_to_index)), self.data_file)
print('Finish saving %d structures to %s, have cost %d seconds' % (
len(data_list), self.data_file, time.time() - begin))
# Save data_list directly (collation happens at load time)
torch.save((data_list, dict(spinful=spinful, index_to_Z=index_to_Z, Z_to_index=Z_to_index)), self.data_file)
print(f'Finish saving {len(data_list)} structures to {self.data_file}, cost {time.time()-begin:.0f}s')

def element_statistics(self, data_list):
index_to_Z, inverse_indices = torch.unique(data_list[0].x, sorted=True, return_inverse=True)
Expand Down
5 changes: 4 additions & 1 deletion deeph/from_PyG_future/graph_norm.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,10 @@

import torch
from torch import Tensor
from torch_scatter import scatter_mean
try:
from torch_scatter import scatter_mean
except ImportError:
from ..compat_scatter import scatter_mean

from torch_geometric.nn.inits import zeros, ones

Expand Down
5 changes: 3 additions & 2 deletions deeph/graph.py
Original file line number Diff line number Diff line change
Expand Up @@ -332,8 +332,9 @@ def __init__(self, if_lcmp):
def __call__(self, graph_list):
if self.if_lcmp:
flag_dict = hasattr(graph_list[0], 'subgraph_dict')
if self.flag_pyg2:
assert flag_dict, 'Please generate the graph file with the current version of PyG'
if self.flag_pyg2 and not flag_dict:
# For single-sample inference without DFT data, skip LCMP subgraph
return Batch.from_data_list(graph_list), None
Comment thread
liu-687 marked this conversation as resolved.
Outdated
batch = Batch.from_data_list(graph_list)

subgraph_atom_idx_batch = []
Expand Down
12 changes: 10 additions & 2 deletions deeph/inference/pred_ham.py
Original file line number Diff line number Diff line change
Expand Up @@ -163,8 +163,16 @@ def predict(input_dir: str, output_dir: str, disable_cuda: bool, device: str,
sys.stderr = sys.stderr.terminal

if restore_blocks_py:
for hamiltonian in hoppings_pred.values():
assert np.all(np.isnan(hamiltonian) == False)
# NaN entries correspond to orbital pairs not covered by the model;
# replace them with zeros (they correspond to negligible matrix elements).
nan_count = 0
for key, hamiltonian in hoppings_pred.items():
nan_mask = np.isnan(hamiltonian)
if nan_mask.any():
nan_count += nan_mask.sum()
hamiltonian[nan_mask] = 0.0
if nan_count > 0:
print(f'Filled {nan_count} NaN entries with zeros (uncovered orbital pairs)')
write_ham_h5(hoppings_pred, path=os.path.join(output_dir, 'rh_pred.h5'))
else:
block_without_restoration['num_model'] = index_model
Expand Down
37 changes: 27 additions & 10 deletions deeph/kernel.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,10 @@
from torch.utils.data import SubsetRandomSampler, DataLoader
from torch.nn.utils import clip_grad_norm_
from torch.utils.tensorboard import SummaryWriter
from torch_scatter import scatter_add
try:
from torch_scatter import scatter_add
except ImportError:
from .compat_scatter import scatter_add
import numpy as np
from psutil import cpu_count

Expand Down Expand Up @@ -372,6 +375,8 @@ def get_dataset(self, only_get_graph=False):
def make_mask(self, dataset):
dataset_mask = []
for data in dataset:
# FIX: Don't move entire data to GPU to avoid OOM
# Only individual tensors will be moved as needed
if self.target == 'hamiltonian' or self.target == 'phiVdphi' or self.target == 'density_matrix':
Oij_value = data.term_real
if data.term_real is not None:
Expand Down Expand Up @@ -406,11 +411,18 @@ def make_mask(self, dataset):
out_fea_len = self.num_orbital * 3
else:
out_fea_len = self.num_orbital
mask = torch.zeros(data.edge_attr.shape[0], out_fea_len, dtype=torch.int8)
label = torch.zeros(data.edge_attr.shape[0], out_fea_len, dtype=torch.get_default_dtype())
mask = torch.zeros(data.edge_attr.shape[0], out_fea_len, dtype=torch.int8, device=self.device)
label = torch.zeros(data.edge_attr.shape[0], out_fea_len, dtype=torch.get_default_dtype(), device=self.device)

atomic_number_edge_i = self.index_to_Z[data.x[data.edge_index[0]]]
atomic_number_edge_j = self.index_to_Z[data.x[data.edge_index[1]]]
idx_to_Z = self.index_to_Z.to(self.device)
x_gpu = data.x.to(self.device)
edge_index_gpu = data.edge_index.to(self.device)
atomic_number_edge_i = idx_to_Z[x_gpu[edge_index_gpu[0]]]
atomic_number_edge_j = idx_to_Z[x_gpu[edge_index_gpu[1]]]

# Move Oij to GPU for computation
if if_only_rc == False:
Oij_value = Oij_value.to(self.device)

for index_out, orbital_dict in enumerate(self.orbital):
for N_M_str, a_b in orbital_dict.items():
Expand Down Expand Up @@ -453,29 +465,29 @@ def make_mask(self, dataset):
(atomic_number_edge_i == condition_atomic_number_i)
& (atomic_number_edge_j == condition_atomic_number_j),
Oij_value[:, condition_orbital_i, condition_orbital_j].t(),
torch.zeros(8, data.edge_attr.shape[0], dtype=torch.get_default_dtype())
torch.zeros(8, data.edge_attr.shape[0], dtype=torch.get_default_dtype(), device=self.device)
).t()
else:
if self.target == 'phiVdphi':
label[:, 3 * index_out:3 * (index_out + 1)] = torch.where(
(atomic_number_edge_i == condition_atomic_number_i)
& (atomic_number_edge_j == condition_atomic_number_j),
Oij_value[:, condition_orbital_i, condition_orbital_j].t(),
torch.zeros(3, data.edge_attr.shape[0], dtype=torch.get_default_dtype())
torch.zeros(3, data.edge_attr.shape[0], dtype=torch.get_default_dtype(), device=self.device)
).t()
else:
label[:, index_out] += torch.where(
(atomic_number_edge_i == condition_atomic_number_i)
& (atomic_number_edge_j == condition_atomic_number_j),
Oij_value[:, condition_orbital_i, condition_orbital_j],
torch.zeros(data.edge_attr.shape[0], dtype=torch.get_default_dtype())
torch.zeros(data.edge_attr.shape[0], dtype=torch.get_default_dtype(), device=self.device)
)
assert len(torch.where((mask != 1) & (mask != 0))[0]) == 0
mask = mask.bool()
data.mask = mask
data.mask = mask.cpu() # FIX: Keep mask on CPU
del data.term_mask
if if_only_rc == False:
data.label = label
data.label = label.cpu() # FIX: Keep label on CPU
if self.target == 'hamiltonian' or self.target == 'density_matrix':
del data.term_real
elif self.target == 'O_ij':
Expand All @@ -485,6 +497,11 @@ def make_mask(self, dataset):
del data.rvxc
del data.rvna
dataset_mask.append(data)
# FIX: Free GPU memory after each sample
del x_gpu, edge_index_gpu, mask, label
if if_only_rc == False:
del Oij_value
torch.cuda.empty_cache()
Comment thread
liu-687 marked this conversation as resolved.
Outdated
return dataset_mask

def train(self, train_loader, val_loader, test_loader):
Expand Down
7 changes: 5 additions & 2 deletions deeph/model.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,10 @@
from torch_geometric.nn.inits import glorot, zeros
from torch_geometric.utils import softmax
from torch_geometric.nn.models.dimenet import BesselBasisLayer
from torch_scatter import scatter_add, scatter
try:
from torch_scatter import scatter_add, scatter
except ImportError:
from .compat_scatter import scatter_add, scatter
import numpy as np
from scipy.special import comb

Expand Down Expand Up @@ -141,7 +144,7 @@ def forward(self, x: Union[torch.Tensor, PairTensor], edge_index: Adj,
if isinstance(x, torch.Tensor):
x: PairTensor = (x, x)

# propagate_type: (x: PairTensor, edge_attr: OptTensor)
# propagate_type: (x: PairTensor, edge_attr: OptTensor, distance: Tensor)
out = self.propagate(edge_index, x=x, edge_attr=edge_attr, distance=distance, size=size)
if self.normalization == 'BatchNorm':
out = self.bn(out)
Expand Down
2 changes: 1 addition & 1 deletion deeph/scripts/evaluate.py
Original file line number Diff line number Diff line change
Expand Up @@ -118,7 +118,7 @@ def main():

label = batch.label
mask = batch.mask
output = output.cpu().reshape(label.shape)
output = output.reshape(label.shape).to(label.device)

assert label.shape == output.shape == mask.shape
mse = torch.pow(label - output, 2)
Expand Down