Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
37 commits
Select commit Hold shift + click to select a range
2530a3a
add l2mae, ordered sampler, batch training
alphalm4 Jan 28, 2026
51c1b2b
add train_shift and train_scale
alphalm4 Feb 1, 2026
a50958f
l2reg; tmp
alphalm4 Feb 1, 2026
5321598
add aselmdb
Jaesun0912 Feb 6, 2026
9025cce
fix error_recorder bug
Jaesun0912 Feb 6, 2026
c987e34
Merge remote-tracking branch 'upstream/main' into omni-train
alphalm4 Feb 27, 2026
c6a5ce3
fix
alphalm4 Feb 27, 2026
0e2af5e
Merge branch 'main' into omni-train
alphalm4 Mar 15, 2026
3c665d9
rebase changelog
alphalm4 Mar 15, 2026
66071f7
Merge branch 'main' into omni-train
alphalm4 May 27, 2026
81cce2a
lmdb ver less than 2 required for simultaneous open
alphalm4 May 28, 2026
04052f2
port load_validset_sequence
alphalm4 May 28, 2026
bbd043b
remove deprecated lines in aselmdb_dataset.py
alphalm4 May 28, 2026
1c83a3c
fix validset subsampling
alphalm4 May 28, 2026
39a48c9
restore error_recorder (loss dict is build from loss.py)
alphalm4 May 30, 2026
8bd596d
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] May 30, 2026
cdca4ac
lint
alphalm4 May 30, 2026
eff657b
continue refactor
YutackPark Jun 9, 2026
01341be
refactor train.py
YutackPark Jun 9, 2026
cbe70d1
fix
alphalm4 Jun 9, 2026
9a41f0a
refactor loss.py
YutackPark Jun 9, 2026
66aca11
Merge branch 'omni-train' of github.com:MDIL-SNU/SevenNet into omni-t…
YutackPark Jun 9, 2026
15f4767
bugfix unittest
YutackPark Jun 9, 2026
295124f
add fallback in sequence loader
alphalm4 Jun 22, 2026
8e98a18
Merge remote-tracking branch 'origin/main' into omni-train
alphalm4 Jun 22, 2026
3761771
[pre-commit.ci] auto fixes from pre-commit.com hooks
pre-commit-ci[bot] Jun 22, 2026
788a0eb
lint
alphalm4 Jun 22, 2026
d8f2139
append modal L2 reg into get_loss_functions_from_config
alphalm4 Jun 22, 2026
a520184
pass only loss_functions to ErrorRecorder.from_config
alphalm4 Jun 22, 2026
e3e9bda
fix lc.csv
alphalm4 Jun 24, 2026
fc99720
refactor
alphalm4 Jun 24, 2026
4632096
refactor
YutackPark Jun 26, 2026
809517b
typo
YutackPark Jun 26, 2026
4c38c69
remove import in __init__ of reewc
YutackPark Jun 26, 2026
843433a
fix train_v2 working_dir passed positionally to processing_epoch_v2
YutackPark Jun 26, 2026
b2da516
typo
alphalm4 Jul 10, 2026
6c658bf
add changelog
alphalm4 Jul 10, 2026
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: 6 additions & 2 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,11 @@
All notable changes to this project will be documented in this file.

## [0.13.1.dev]
### Added
- Training features used for SevenNet-Omni: batch training, `OrderedSampler`, `grad_clip`, `onecyclelr`
- Loss: MAE, L2MAE
- Dataset type: `aselmdb`, `custom`

### Fixed
- `D3Calculator()` segfault bug when reusing the calculator within different sized `Atoms`.

Expand Down Expand Up @@ -39,13 +44,12 @@ All notable changes to this project will be documented in this file.

### Added
- TorchSim interface and docs
- SevenNet-Omni-i8, SevenNet-Omni-i12

## [0.12.0]
### Added
- Documentation moved to RTD
- LAMMPS-MLIAP integration with GhostExchangeOp

### Added
- SevenNet-Omni
- Example config for fine-tuning the SevenNet-MF-ompa model
- FlashTP support (https://github.com/SNU-ARC/flashTP)
Expand Down
4 changes: 3 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,9 @@ dependencies = [
"pandas",
"requests",
"ninja",
"setuptools>=61.0"
"setuptools>=61.0",
"lmdb<2.0.0",
"orjson"
]
[project.optional-dependencies]
test = ["pytest", "pytest-cov>=5", "ipython"]
Expand Down
2 changes: 1 addition & 1 deletion setup.cfg
Original file line number Diff line number Diff line change
Expand Up @@ -10,5 +10,5 @@ include_trailing_comma=True
force_grid_wrap=0
use_parentheses=True
line_length=80
known_third_party=ase,braceexpand,e3nn,numpy,packaging,pandas,pytest,requests,sklearn,torch,torch_geometric,tqdm,yaml
known_third_party=ase,braceexpand,e3nn,lmdb,numpy,orjson,packaging,pandas,pytest,requests,sklearn,torch,torch_geometric,tqdm,yaml
known_first_party=
28 changes: 26 additions & 2 deletions sevenn/_const.py
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,13 @@
IMPLEMENTED_SELF_CONNECTION_TYPE = ['nequip', 'linear']
IMPLEMENTED_INTERACTION_TYPE = ['nequip']

IMPLEMENTED_MODAL_MODULE_DICT = {
KEY.USE_MODAL_NODE_EMBEDDING: 'onehot_to_feature_x',
KEY.USE_MODAL_SELF_INTER_INTRO: 'self_interaction_1',
KEY.USE_MODAL_SELF_INTER_OUTRO: 'self_interaction_2',
KEY.USE_MODAL_OUTPUT_BLOCK: 'reduce_input_to_hidden',
}

IMPLEMENTED_SHIFT = ['per_atom_energy_mean', 'elemwise_reference_energies']
IMPLEMENTED_SCALE = ['force_rms', 'per_atom_energy_std', 'elemwise_force_rms']

Expand All @@ -26,6 +33,8 @@
'Stress',
'Stress_GPa',
'TotalLoss',
'L2_modal',
'Modal_cos',
]

IMPLEMENTED_MODEL = ['E3_equivariant_model']
Expand Down Expand Up @@ -116,6 +125,8 @@ def error_record_condition(x):
KEY.CONV_DENOMINATOR: 'avg_num_neigh',
KEY.TRAIN_DENOMINTAOR: False,
KEY.TRAIN_SHIFT_SCALE: False,
KEY.TRAIN_SHIFT: False,
KEY.TRAIN_SCALE: False,
# KEY.OPTIMIZE_BY_REDUCE: True, # deprecated, always True
KEY.USE_BIAS_IN_LINEAR: False,
KEY.USE_MODAL_NODE_EMBEDDING: False,
Expand Down Expand Up @@ -158,6 +169,8 @@ def error_record_condition(x):
],
KEY.CONVOLUTION_WEIGHT_NN_HIDDEN_NEURONS: list,
KEY.TRAIN_SHIFT_SCALE: bool,
KEY.TRAIN_SHIFT: bool,
KEY.TRAIN_SCALE: bool,
KEY.TRAIN_DENOMINTAOR: bool,
KEY.USE_BIAS_IN_LINEAR: bool,
KEY.USE_MODAL_NODE_EMBEDDING: bool,
Expand Down Expand Up @@ -215,6 +228,7 @@ def model_defaults(config):
KEY.REHEARSAL: False,
KEY.MEM_BATCH_SIZE: 0,
KEY.MEM_RATIO: 1,
KEY.LOADER_KWARGS: {},
# KEY.DATA_SHUFFLE: True,
# KEY.DATA_WEIGHT: False,
# KEY.DATA_MODALITY: False,
Expand All @@ -230,7 +244,7 @@ def model_defaults(config):
KEY.RATIO: float,
KEY.BATCH_SIZE: int,
KEY.PREPROCESS_NUM_CORES: int,
KEY.DATASET_TYPE: lambda x: x in ['graph', 'atoms'],
KEY.DATASET_TYPE: lambda x: x in ['graph', 'atoms', 'custom', 'aselmdb'],
# KEY.USE_SPECIES_WISE_SHIFT_SCALE: bool,
KEY.SHIFT: lambda x: type(x) in [float, list] or x in IMPLEMENTED_SHIFT,
KEY.SCALE: lambda x: type(x) in [float, list] or x in IMPLEMENTED_SCALE,
Expand Down Expand Up @@ -264,15 +278,20 @@ def data_defaults(config):
KEY.OPTIM_PARAM: {},
KEY.SCHEDULER: 'exponentiallr',
KEY.SCHEDULER_PARAM: {},
KEY.ENERGY_WEIGHT: 1.0,
KEY.FORCE_WEIGHT: 0.1,
KEY.STRESS_WEIGHT: 1e-6, # SIMPLE-NN default
KEY.GRAD_CLIP: None,
KEY.REG_PARAM: {},
KEY.PER_EPOCH: 5,
KEY.TRAIN_BY_BATCH: False,
# KEY.USE_TESTSET: False,
KEY.CONTINUE: {
KEY.CHECKPOINT: False,
KEY.RESET_OPTIMIZER: False,
KEY.RESET_SCHEDULER: False,
KEY.RESET_EPOCH: False,
KEY.RESET_DATA_PROGRESS: True,
KEY.USE_STATISTIC_VALUES_OF_CHECKPOINT: True,
KEY.USE_STATISTIC_VALUES_FOR_CP_MODAL_ONLY: True,
},
Expand All @@ -296,16 +315,21 @@ def data_defaults(config):
TRAINING_CONFIG_CONDITION = {
KEY.RANDOM_SEED: int,
KEY.EPOCH: int,
KEY.ENERGY_WEIGHT: float,
KEY.FORCE_WEIGHT: float,
KEY.STRESS_WEIGHT: float,
KEY.GRAD_CLIP: lambda x: x is None or (type(x) in [float, int] and x > 0),
KEY.REG_PARAM: dict,
KEY.USE_TESTSET: None, # Not used
KEY.NUM_WORKERS: int,
KEY.PER_EPOCH: int,
KEY.PER_EPOCH: lambda x: type(x) in [float, int],
KEY.TRAIN_BY_BATCH: bool,
KEY.CONTINUE: {
KEY.CHECKPOINT: str,
KEY.RESET_OPTIMIZER: bool,
KEY.RESET_SCHEDULER: bool,
KEY.RESET_EPOCH: bool,
KEY.RESET_DATA_PROGRESS: bool,
KEY.USE_STATISTIC_VALUES_OF_CHECKPOINT: bool,
KEY.USE_STATISTIC_VALUES_FOR_CP_MODAL_ONLY: bool,
},
Expand Down
17 changes: 17 additions & 0 deletions sevenn/_keys.py
Original file line number Diff line number Diff line change
Expand Up @@ -117,12 +117,17 @@
EPOCH = 'epoch'
LOSS = 'loss'
LOSS_PARAM = 'loss_param'
LOSS_TYPE = 'loss_type'
LOSS_WEIGHT = 'loss_weight'
OPTIMIZER = 'optimizer'
OPTIM_PARAM = 'optim_param'
SCHEDULER = 'scheduler'
SCHEDULER_PARAM = 'scheduler_param'
SCHEDULER_BATCH_MODE = 'scheduler_batch_mode'
ENERGY_WEIGHT = 'energy_loss_weight'
FORCE_WEIGHT = 'force_loss_weight'
STRESS_WEIGHT = 'stress_loss_weight'
GRAD_CLIP = 'grad_clip'
DEVICE = 'device'
DTYPE = 'dtype'

Expand All @@ -135,6 +140,7 @@
RESET_OPTIMIZER = 'reset_optimizer'
RESET_SCHEDULER = 'reset_scheduler'
RESET_EPOCH = 'reset_epoch'
RESET_DATA_PROGRESS = 'reset_data_progress'
USE_STATISTIC_VALUES_OF_CHECKPOINT = 'use_statistic_values_of_checkpoint'
USE_STATISTIC_VALUES_FOR_CP_MODAL_ONLY = (
'use_statistic_values_for_cp_modal_only'
Expand All @@ -157,6 +163,11 @@
DDP_BACKEND = 'ddp_backend'
PER_EPOCH = 'per_epoch'

TRAIN_BY_BATCH = 'train_by_batch'
TOTAL_DATA_NUM = 'total_data_num'
CURRENT_DATA_IDX = 'current_data_index'
NUMPY_RNG_STATE = 'numpy_rng_state'

USE_WEIGHT = 'use_weight'
USE_MODALITY = 'use_modality'
DEFAULT_MODAL = 'default_modal'
Expand Down Expand Up @@ -218,12 +229,15 @@
CONV_DENOMINATOR = 'conv_denominator'
SHIFT = 'shift'
SCALE = 'scale'
LOADER_KWARGS = 'loader_kwargs'

USE_SPECIES_WISE_SHIFT_SCALE = 'use_species_wise_shift_scale'
USE_MODAL_WISE_SHIFT = 'use_modal_wise_shift'
USE_MODAL_WISE_SCALE = 'use_modal_wise_scale'

TRAIN_SHIFT_SCALE = 'train_shift_scale'
TRAIN_SHIFT = 'train_shift'
TRAIN_SCALE = 'train_scale'
TRAIN_DENOMINTAOR = 'train_denominator'
INTERACTION_TYPE = 'interaction_type'
TRAIN_AVG_NUM_NEIGH = 'train_avg_num_neigh' # deprecated
Expand All @@ -232,6 +246,9 @@
CUEQUIVARIANCE_CONFIG = 'cuequivariance_config'
USE_OEQ = 'use_oeq'

REG_PARAM = 'regularization_param'
REG_WEIGHT = 'regularization_weight'

_NORMALIZE_SPH = '_normalize_sph'
OPTIMIZE_BY_REDUCE = 'optimize_by_reduce'

Expand Down
8 changes: 8 additions & 0 deletions sevenn/checkpoint.py
Original file line number Diff line number Diff line change
Expand Up @@ -226,6 +226,7 @@ def __init__(self, checkpoint_path: Union[pathlib.Path, str]) -> None:
self._checkpoint_path = os.path.abspath(checkpoint_path)
self._config = None
self._epoch = None
self._data_progress = None
self._model_state_dict = None
self._optimizer_state_dict = None
self._scheduler_state_dict = None
Expand Down Expand Up @@ -304,6 +305,12 @@ def epoch(self) -> Optional[int]:
self._load()
return self._epoch

@property
def data_progress(self) -> Optional[Dict[str, int]]:
if not self._loaded:
self._load()
return self._data_progress

@property
def time(self) -> str:
if not self._loaded:
Expand All @@ -328,6 +335,7 @@ def _load(self) -> None:
self._optimizer_state_dict = cp.get('optimizer_state_dict', {})
self._scheduler_state_dict = cp.get('scheduler_state_dict', {})
self._epoch = cp.get('epoch', None)
self._data_progress = cp.get('data_progress', None)
self._time = cp.get('time', 'Not found')
self._hash = cp.get('hash', 'Not found')

Expand Down
Loading
Loading