Skip to content
Closed
Show file tree
Hide file tree
Changes from all 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
5 changes: 5 additions & 0 deletions sevenn/_const.py
Original file line number Diff line number Diff line change
Expand Up @@ -132,6 +132,9 @@ def error_record_condition(x):
KEY.USE_FLASH_TP: False,
KEY.CUEQUIVARIANCE_CONFIG: {},
KEY.USE_OEQ: False,
KEY.USE_LES: False,
# les_config keys: les_args (dict), compute_bec (bool), bec_output_index (int|None)
KEY.LES_CONFIG: {},
}


Expand Down Expand Up @@ -180,6 +183,8 @@ def error_record_condition(x):
KEY.USE_FLASH_TP: bool,
KEY.CUEQUIVARIANCE_CONFIG: dict,
KEY.USE_OEQ: bool,
KEY.USE_LES: bool,
KEY.LES_CONFIG: dict,
}


Expand Down
10 changes: 10 additions & 0 deletions sevenn/_keys.py
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,12 @@
ATOMIC_ENERGY: Final[str] = 'atomic_energy'
PRED_TOTAL_ENERGY: Final[str] = 'inferred_total_energy'

# LES (Latent Ewald Summation) outputs
LES_Q: Final[str] = 'les_latent_charge' # (N_atoms, n_charges) per-atom latent charges
SR_ENERGY: Final[str] = 'les_sr_energy' # (n_graphs,) short-range energy sum
LR_ENERGY: Final[str] = 'les_lr_energy' # (n_graphs,) long-range Ewald energy
LES_BEC: Final[str] = 'les_born_eff_charge' # (N_atoms, 3, 3) Born effective charges

PRED_PER_ATOM_ENERGY: Final[str] = 'inferred_per_atom_energy'
PER_ATOM_ENERGY: Final[str] = 'per_atom_energy'

Expand Down Expand Up @@ -225,6 +231,10 @@
CUEQUIVARIANCE_CONFIG = 'cuequivariance_config'
USE_OEQ = 'use_oeq'

# LES model configuration keys
USE_LES = 'use_les'
LES_CONFIG = 'les_config'

_NORMALIZE_SPH = '_normalize_sph'
OPTIMIZE_BY_REDUCE = 'optimize_by_reduce'

Expand Down
134 changes: 134 additions & 0 deletions sevenn/checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -369,6 +369,140 @@ def build_model(

return model

def build_model_with_les(
self,
les_config: Optional[Dict[str, Any]] = None,
freeze_sr: bool = True,
*,
enable_cueq: Optional[bool] = None,
enable_flash: Optional[bool] = None,
enable_oeq: Optional[bool] = None,
) -> AtomGraphSequential:
"""
Build a LES-equipped model initialised from this (non-LES) checkpoint.

New LES parameters not present in the pretrained state dict use their
default construction-time initialisation:
- ``les_charge_readout.linear.weight``: e3nn default (non-zero)
so that the Ewald gradient is non-zero and fine-tuning starts
immediately. Pass ``les_config={'zero_init': True}`` to force
zero weights, which preserves the SR baseline exactly at epoch 0
but sets all charge gradients to zero — useful for verification
tests, not for real training.
- ``les_lr_energy.les.*``: Les() default init

All other parameters are loaded from the checkpoint unchanged.

Args:
les_config: LES configuration dict forwarded to the model builder as
``LES_CONFIG``. Supported keys:
les_args (dict) — kwargs for ``Les()``,
default ``{'use_atomwise': False}``
n_charges (int) — number of latent charge channels
per atom; the Ewald energy is the
sum of n_charges independent
Coulomb interactions, default 1
hidden_channels (list) — hidden layer widths for the charge
readout MLP, e.g. ``[128]`` for a
two-layer 128→128→1 network,
default ``[]`` (single layer)
zero_init (bool) — zero-initialise the charge readout
weights, default ``False``
compute_bec (bool) — compute Born effective charges,
default ``False``
bec_output_index (int) — 0/1/2 for BEC component,
default ``None``
freeze_sr: if ``True`` (default), set ``requires_grad=False`` on
all parameters whose top-level module is not
``les_charge_readout`` or ``les_lr_energy``. The Trainer
already filters parameters by ``requires_grad``, so no
Trainer changes are needed.
enable_cueq/enable_flash/enable_oeq: backend overrides (same
semantics as :meth:`build_model`).

Returns:
:class:`~sevenn.nn.sequential.AtomGraphSequential` with LES layers
added and SR weights loaded from the checkpoint.

Raises:
RuntimeError: if any *non*-LES keys are missing after loading,
indicating an unexpected model mismatch.
"""
from .model_build import build_E3_equivariant_model

if les_config is None:
les_config = {}

# Resolve backend flags using the same logic as build_model()
try:
cp_using_cueq = self.config[KEY.CUEQUIVARIANCE_CONFIG]['use']
except KeyError:
cp_using_cueq = False
final_cueq = cp_using_cueq if enable_cueq is None else enable_cueq

cp_using_flash = self.config.get(KEY.USE_FLASH_TP, False)
final_flash = cp_using_flash if enable_flash is None else enable_flash

cp_using_oeq = self.config.get(KEY.USE_OEQ, False)
final_oeq = cp_using_oeq if enable_oeq is None else enable_oeq

if sum([final_cueq, final_flash, final_oeq]) > 1:
raise ValueError('Only one TP accelerator can be enabled.')

# Step 1: use build_model() to get a correctly backend-converted non-LES
# state dict. This handles FlashTP↔e3nn and cueq↔e3nn in one place
# without duplicating the conversion logic here.
sr_state_dict = self.build_model(
enable_cueq=enable_cueq,
enable_flash=enable_flash,
enable_oeq=enable_oeq,
).state_dict()

# Step 2: build the LES model with the identical resolved backend.
cfg_new = self.config
cfg_new[KEY.USE_LES] = True
cfg_new[KEY.LES_CONFIG] = les_config
cfg_new[KEY.USE_OEQ] = final_oeq
cfg_new[KEY.CUEQUIVARIANCE_CONFIG] = {'use': final_cueq}
cfg_new[KEY.USE_FLASH_TP] = final_flash

model = build_E3_equivariant_model(cfg_new)

# Step 3: load the already-converted SR state dict into the LES model.
# strict=False because LES-specific keys are absent from sr_state_dict.
missing, not_used = model.load_state_dict(sr_state_dict, strict=False)

# Only LES-specific keys should be absent from the pretrained checkpoint.
# Any other missing key means an unexpected model mismatch.
les_prefixes = ('les_charge_readout.', 'les_lr_energy.')
unexpected_missing = [
k for k in missing if not k.startswith(les_prefixes)
]
if unexpected_missing:
raise RuntimeError(
'Unexpected missing keys when loading SR checkpoint into LES '
f'model: {unexpected_missing}'
)
if missing:
warnings.warn(
'LES parameters initialised from scratch (not in checkpoint): '
f'{missing}',
UserWarning,
)
if not_used:
warnings.warn(
f'Checkpoint keys not used in LES model: {not_used}',
UserWarning,
)

if freeze_sr:
les_module_names = {'les_charge_readout', 'les_lr_energy'}
for name, param in model.named_parameters():
if name.split('.')[0] not in les_module_names:
param.requires_grad_(False)

return model

def yaml_dict(self, mode: str) -> Dict[str, Any]:
"""
Return dict for input.yaml from checkpoint config
Expand Down
88 changes: 77 additions & 11 deletions sevenn/model_build.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,11 +14,12 @@
from .nn.edge_embedding import (
BesselBasis,
EdgeEmbedding,
EdgePreprocess,
PolynomialCutoff,
SphericalEncoding,
XPLORCutoff,
)
from .nn.force_output import ForceStressOutputFromEdge
from .nn.force_output import ForceStressOutput, ForceStressOutputFromEdge
from .nn.interaction_blocks import NequIP_interaction_block
from .nn.linear import AtomReduce, FCN_e3nn, IrrepsLinear
from .nn.node_embedding import OnehotEmbedding
Expand All @@ -28,6 +29,11 @@
SelfConnectionLinearIntro,
SelfConnectionOutro,
)
from .nn.les import (
AddLREnergy,
LatentChargeReadout,
LatentEwaldSum,
)
from .nn.sequential import AtomGraphSequential

# warning from PyTorch, about e3nn type annotations
Expand Down Expand Up @@ -458,8 +464,22 @@ def build_E3_equivariant_model(

for data w/o cell volume, pred_stress has garbage values
"""
if parallel and config.get(KEY.USE_LES, False):
raise NotImplementedError(
'LES parallel mode is not supported on LES-legacy branch. '
'LES-legacy uses EdgePreprocess (serial) for unified pos/strain '
'gradients. Use the LES branch for training / parallel deployment.'
)

layers = OrderedDict()

# Legacy serial LES: insert EdgePreprocess before edge_embedding so that
# pos → EDGE_VEC is live in the autograd graph. The strain leaf created
# here is shared by LatentEwaldSum (cell) and ForceStressOutput (forces +
# complete stress), replacing the separate Path-2 / LES_STRAIN approach.
if config.get(KEY.USE_LES, False):
layers['edge_preprocess'] = EdgePreprocess(is_stress=True)

cutoff = config[KEY.CUTOFF]
num_species = config[KEY.NUM_SPECIES]
feature_multiplicity = config[KEY.NODE_FEATURE_MULTIPLICITY]
Expand Down Expand Up @@ -599,19 +619,65 @@ def build_E3_equivariant_model(
layers.update(interaction_builder(**param_interaction_block))
irreps_x = irreps_out

if config.get(KEY.USE_LES, False):
les_cfg = config.get(KEY.LES_CONFIG, {})
# Latent charge readout must sit BEFORE init_feature_reduce because
# reduce_input_to_hidden overwrites KEY.NODE_FEATURE in-place.
# Reading here guarantees we use the full (pre-projection) node features.
layers['les_charge_readout'] = LatentChargeReadout(
irreps_in=irreps_x, # type: ignore
data_key_in=KEY.NODE_FEATURE,
data_key_out=KEY.LES_Q,
n_charges=les_cfg.get('n_charges', 1),
hidden_channels=les_cfg.get('hidden_channels', None),
zero_init=les_cfg.get('zero_init', False),
)

layers.update(init_feature_reduce(config, irreps_x)) # type: ignore

layers.update(
{
'rescale_atomic_energy': init_shift_scale(config),
'reduce_total_enegy': AtomReduce(
data_key_in=KEY.ATOMIC_ENERGY,
data_key_out=KEY.PRED_TOTAL_ENERGY,
),
}
)
if config.get(KEY.USE_LES, False):
layers.update(
{
'rescale_atomic_energy': init_shift_scale(config),
# SR energy: sum of per-atom (local) energies
'reduce_sr_energy': AtomReduce(
data_key_in=KEY.ATOMIC_ENERGY,
data_key_out=KEY.SR_ENERGY,
),
# LR energy: Ewald summation on latent charges
'les_lr_energy': LatentEwaldSum(
les_args=les_cfg.get('les_args', {'use_atomwise': False}),
data_key_in=KEY.LES_Q,
data_key_out=KEY.LR_ENERGY,
compute_bec=les_cfg.get('compute_bec', False),
bec_output_index=les_cfg.get('bec_output_index', None),
),
# Total = SR + LR
'add_lr_to_total': AddLREnergy(
key_sr=KEY.SR_ENERGY,
key_lr=KEY.LR_ENERGY,
data_key_out=KEY.PRED_TOTAL_ENERGY,
),
}
)
else:
layers.update(
{
'rescale_atomic_energy': init_shift_scale(config),
'reduce_total_enegy': AtomReduce(
data_key_in=KEY.ATOMIC_ENERGY,
data_key_out=KEY.PRED_TOTAL_ENERGY,
),
}
)

gradient_module = ForceStressOutputFromEdge()
if config.get(KEY.USE_LES, False):
# Legacy serial mode: unified pos + _strain gradients via EdgePreprocess.
# ForceStressOutput differentiates w.r.t. strained pos (force) and
# _strain (complete stress — SR virial + LR positional + LR cell).
gradient_module = ForceStressOutput()
else:
gradient_module = ForceStressOutputFromEdge()
grad_key = gradient_module.get_grad_key()
layers.update({'force_output': gradient_module})

Expand Down
24 changes: 21 additions & 3 deletions sevenn/nn/edge_embedding.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,8 +12,13 @@
@compile_mode('script')
class EdgePreprocess(nn.Module):
"""
preprocessing pos to edge vectors and edge lengths
currently used in sevenn/scripts/deploy for lammps serial model
Computes edge vectors and lengths from atomic positions, with optional
affine strain for stress computation.

Used exclusively by LES models (LES branch) as the first model
layer, restoring the pos → EDGE_VEC autograd connection so that
ForceStressOutput can recover forces and stress via positional and strain
gradients rather than edge-virial decomposition.
"""

def __init__(self, is_stress: bool) -> None:
Expand All @@ -30,9 +35,9 @@ def forward(self, data: AtomGraphDataType) -> AtomGraphDataType:
cell_shift = data[KEY.CELL_SHIFT]
pos = data[KEY.POS]

batch = data[KEY.BATCH] # for deploy, must be defined first
if self.is_stress:
if self._is_batch_data:
batch = data[KEY.BATCH]
num_batch = int(batch.max().cpu().item()) + 1
strain = torch.zeros(
(num_batch, 3, 3),
Expand Down Expand Up @@ -60,12 +65,25 @@ def forward(self, data: AtomGraphDataType) -> AtomGraphDataType:
pos = pos + torch.mm(pos, sym_strain)
cell = cell + torch.mm(cell, sym_strain)

# Write strained pos and cell back so that downstream modules
# (LatentEwaldSum, ForceStressOutput) receive tensors that are
# connected to _strain in the autograd graph. This enables a
# single d(E)/d(_strain) call to capture the complete stress
# (SR virial + LR positional + LR cell/k-space) without a
# separate LES_STRAIN leaf inside LatentEwaldSum.
data[KEY.POS] = pos
if self._is_batch_data:
data[KEY.CELL] = cell.reshape(-1, 3)
else:
data[KEY.CELL] = cell

idx_src = data[KEY.EDGE_IDX][0]
idx_dst = data[KEY.EDGE_IDX][1]

edge_vec = pos[idx_dst] - pos[idx_src]

if self._is_batch_data:
batch = data[KEY.BATCH]
edge_vec = edge_vec + torch.einsum(
'ni,nij->nj', cell_shift, cell[batch[idx_src]]
)
Expand Down
Loading
Loading