From 1393311b80f58562f618fa958b03ffc8a4920b3a Mon Sep 17 00:00:00 2001 From: alphalm4 Date: Sun, 12 Jul 2026 18:42:25 +0900 Subject: [PATCH 1/5] Add mixed dtype shift-scale support --- sevenn/_const.py | 2 ++ sevenn/_keys.py | 1 + sevenn/calculator.py | 17 +++++++++- sevenn/checkpoint.py | 8 +++++ sevenn/main/sevenn_get_model.py | 21 ++++++++++-- sevenn/mliap.py | 5 ++- sevenn/model_build.py | 6 +++- sevenn/nn/scale.py | 41 ++++++++++++++++------- sevenn/pair_e3gnn/pair_e3gnn.cpp | 12 +++---- sevenn/pair_e3gnn/pair_e3gnn_parallel.cpp | 13 +++---- sevenn/scripts/deploy.py | 7 ++++ sevenn/torchsim.py | 22 +++++++----- sevenn/train/loss.py | 4 +++ sevenn/util.py | 6 +++- 14 files changed, 126 insertions(+), 39 deletions(-) diff --git a/sevenn/_const.py b/sevenn/_const.py index 6c78f9c8..e16e19a4 100644 --- a/sevenn/_const.py +++ b/sevenn/_const.py @@ -127,6 +127,7 @@ def error_record_condition(x): KEY.TRAIN_SHIFT_SCALE: False, KEY.TRAIN_SHIFT: False, KEY.TRAIN_SCALE: False, + KEY.SHIFT_SCALE_DTYPE: 'double', # KEY.OPTIMIZE_BY_REDUCE: True, # deprecated, always True KEY.USE_BIAS_IN_LINEAR: False, KEY.USE_MODAL_NODE_EMBEDDING: False, @@ -171,6 +172,7 @@ def error_record_condition(x): KEY.TRAIN_SHIFT_SCALE: bool, KEY.TRAIN_SHIFT: bool, KEY.TRAIN_SCALE: bool, + KEY.SHIFT_SCALE_DTYPE: lambda x: x in ['single', 'double'], KEY.TRAIN_DENOMINTAOR: bool, KEY.USE_BIAS_IN_LINEAR: bool, KEY.USE_MODAL_NODE_EMBEDDING: bool, diff --git a/sevenn/_keys.py b/sevenn/_keys.py index 8773a099..c2b64b0a 100644 --- a/sevenn/_keys.py +++ b/sevenn/_keys.py @@ -229,6 +229,7 @@ CONV_DENOMINATOR = 'conv_denominator' SHIFT = 'shift' SCALE = 'scale' +SHIFT_SCALE_DTYPE = 'shift_scale_dtype' LOADER_KWARGS = 'loader_kwargs' USE_SPECIES_WISE_SHIFT_SCALE = 'use_species_wise_shift_scale' diff --git a/sevenn/calculator.py b/sevenn/calculator.py index 803ad779..e6642fd5 100644 --- a/sevenn/calculator.py +++ b/sevenn/calculator.py @@ -37,6 +37,7 @@ def __init__( enable_oeq: bool = False, compute_atomic_virial: bool = False, sevennet_config: Optional[Dict] = None, # Not used in logic, just meta info + shift_scale_dtype: str = 'double', **kwargs, ) -> None: """Initialize SevenNetCalculator. @@ -67,10 +68,14 @@ def __init__( Not used, but can be used to carry meta information of this calculator compute_atomic_virial: bool, default=False If True, request per-atom virial output (`stresses`) at runtime. + shift_scale_dtype: str, default='double' + dtype of the final shift/scale (rescale) parameters used at inference. + 'single' is only for backward reproducibility. """ super().__init__(**kwargs) self.sevennet_config = None self.compute_atomic_virial = compute_atomic_virial + self.shift_scale_dtype = shift_scale_dtype if isinstance(model, pathlib.PurePath): model = str(model) @@ -118,7 +123,8 @@ def __init__( cp = util.load_checkpoint(model) model_loaded = cp.build_model( - enable_cueq=enable_cueq, enable_flash=enable_flash, enable_oeq=enable_oeq # noqa: E501 + enable_cueq=enable_cueq, enable_flash=enable_flash, enable_oeq=enable_oeq, # noqa: E501 + shift_scale_dtype=shift_scale_dtype, ) model_loaded.set_is_batch_data(False) @@ -248,6 +254,7 @@ def __init__( functional_name: str = 'pbe', vdw_cutoff: float = 9000, # au^2, 0.52917726 angstrom = 1 au cn_cutoff: float = 1600, # au^2, 0.52917726 angstrom = 1 au + shift_scale_dtype: Optional[str] = 'double', **kwargs, # pass extra kwargs to both calculators ) -> None: """Initialize SevenNetD3Calculator. CUDA required. @@ -308,11 +315,19 @@ def __init__( enable_flash=enable_flash, enable_oeq=enable_oeq, sevennet_config=sevennet_config, + shift_scale_dtype=shift_scale_dtype, **kwargs, ) super().__init__([sevennet_calc, d3_calc]) + def calculate(self, atoms=None, properties=None, system_changes=all_changes): + # To match the precision of SevenNetCalculator output + super().calculate(atoms, properties, system_changes) + for key in ('forces', 'stress', 'stresses'): + if key in self.results: + self.results[key] = np.asarray(self.results[key], dtype=np.float32) + def _load(name: str) -> ctypes.CDLL: from torch.utils.cpp_extension import LIB_EXT, _get_build_directory, load diff --git a/sevenn/checkpoint.py b/sevenn/checkpoint.py index 8e4a704e..18f9fa3b 100644 --- a/sevenn/checkpoint.py +++ b/sevenn/checkpoint.py @@ -357,6 +357,7 @@ def build_model( enable_flash: bool = False, enable_oeq: bool = False, _flash_lammps: bool = False, + shift_scale_dtype: str = 'double', ) -> AtomGraphSequential: from .model_build import build_E3_equivariant_model @@ -376,6 +377,13 @@ def build_model( cfg_new['_flash_lammps'] = _flash_lammps cfg_new[KEY.USE_OEQ] = enable_oeq + if shift_scale_dtype not in ('single', 'double'): + raise ValueError( + "shift_scale_dtype must be 'single' or 'double', " + f'got {shift_scale_dtype!r}' + ) + cfg_new[KEY.SHIFT_SCALE_DTYPE] = shift_scale_dtype + if (cp_using_cueq, cp_using_flash, cp_using_oeq) == ( enable_cueq, enable_flash, diff --git a/sevenn/main/sevenn_get_model.py b/sevenn/main/sevenn_get_model.py index 8f7bcdb6..d8fe9f0f 100644 --- a/sevenn/main/sevenn_get_model.py +++ b/sevenn/main/sevenn_get_model.py @@ -62,6 +62,15 @@ def add_args(parser): help='Use LAMMPS ML-IAP interface.', action='store_true', ) + ag.add_argument( + '--shift_scale_dtype', + choices=('double', 'single'), + default='double', + help=( + 'dtype of final shift/scale energy rescaling at inference. ' + 'Default is double; single is for backward compatibility.' + ), + ) def run(args): @@ -76,6 +85,7 @@ def run(args): use_cueq = args.enable_cueq use_oeq = args.enable_oeq use_mliap = args.use_mliap + shift_scale_dtype = args.shift_scale_dtype # Check dependencies if use_flash: @@ -119,9 +129,15 @@ def run(args): from sevenn.scripts.deploy import deploy, deploy_parallel if get_serial: - deploy(checkpoint_path, output_prefix, modal, use_flash=use_flash, use_oeq=use_oeq) # noqa: E501 + deploy( + checkpoint_path, output_prefix, modal, use_flash=use_flash, use_oeq=use_oeq, # noqa: E501 + shift_scale_dtype=shift_scale_dtype + ) else: - deploy_parallel(checkpoint_path, output_prefix, modal, use_flash=use_flash, use_oeq=use_oeq) # noqa: E501 + deploy_parallel( + checkpoint_path, output_prefix, modal, use_flash=use_flash, use_oeq=use_oeq, # noqa: E501 + shift_scale_dtype=shift_scale_dtype + ) else: from sevenn import mliap @@ -139,6 +155,7 @@ def run(args): use_cueq=use_cueq, use_flash=use_flash, use_oeq=use_oeq, + shift_scale_dtype=shift_scale_dtype, ) torch.save(mliap_module, output_prefix) diff --git a/sevenn/mliap.py b/sevenn/mliap.py index 513d55c7..70fb9d0b 100644 --- a/sevenn/mliap.py +++ b/sevenn/mliap.py @@ -87,6 +87,7 @@ def __init__( modal: Optional[str] = None use_cueq: bool = False use_flash: bool = False + shift_scale_dtype: str = 'double' """ super().__init__() @@ -112,6 +113,7 @@ def __init__( self.use_cueq = kwargs.get('use_cueq', False) self.use_flash = kwargs.get('use_flash', False) self.use_oeq = kwargs.get('use_oeq', False) + self.shift_scale_dtype = kwargs.get('shift_scale_dtype', 'double') self.modal = kwargs.get('modal', None) # extract configs @@ -153,7 +155,8 @@ def _ensure_model_initialized(self): print('[INFO] Lazy initializing SevenNet model...', flush=True) print(f'[INFO] cueq={self.use_cueq}, flashTP={self.use_flash}, oeq={self.use_oeq}', flush=True) # noqa: E501 model = self.cp.build_model( - enable_cueq=self.use_cueq, enable_flash=self.use_flash, enable_oeq=self.use_oeq # noqa: E501 + enable_cueq=self.use_cueq, enable_flash=self.use_flash, enable_oeq=self.use_oeq, # noqa: E501 + shift_scale_dtype=self.shift_scale_dtype ) for k, module in model._modules.items(): diff --git a/sevenn/model_build.py b/sevenn/model_build.py index f0482f11..92582e64 100644 --- a/sevenn/model_build.py +++ b/sevenn/model_build.py @@ -167,7 +167,11 @@ def init_shift_scale( shift_scale.append(s) shift, scale = shift_scale - ss_kwargs = {'train_shift': train_shift, 'train_scale': train_scale} + ss_kwargs = { + 'train_shift': train_shift, + 'train_scale': train_scale, + 'shift_scale_dtype': config.get(KEY.SHIFT_SCALE_DTYPE, 'double'), + } rescale_module = None if config.get(KEY.USE_MODALITY, False): rescale_module = ModalWiseRescale.from_mappers( # type: ignore diff --git a/sevenn/nn/scale.py b/sevenn/nn/scale.py index 551355d4..0b975834 100644 --- a/sevenn/nn/scale.py +++ b/sevenn/nn/scale.py @@ -8,6 +8,14 @@ from sevenn._const import NUM_UNIV_ELEMENT, AtomGraphDataType +def _resolve_shift_scale_dtype(dtype: str) -> torch.dtype: + if dtype == 'single': + return torch.float32 + if dtype == 'double': + return torch.float64 + raise ValueError(f'Unsupported shift/scale dtype: {dtype}') + + def _as_univ( ss: List[float], type_map: Dict[int, int], default: float ) -> List[float]: @@ -33,6 +41,7 @@ def __init__( train_shift: bool = False, train_scale: bool = False, train_shift_scale: bool = False, + shift_scale_dtype: str = 'double', **kwargs, ) -> None: assert isinstance(shift, float) and isinstance(scale, float) @@ -40,11 +49,12 @@ def __init__( if train_shift_scale: train_shift = True train_scale = True + dtype = _resolve_shift_scale_dtype(shift_scale_dtype) self.shift = nn.Parameter( - torch.FloatTensor([shift]), requires_grad=train_shift + torch.tensor([shift], dtype=dtype), requires_grad=train_shift ) self.scale = nn.Parameter( - torch.FloatTensor([scale]), requires_grad=train_scale + torch.tensor([scale], dtype=dtype), requires_grad=train_scale ) self.key_input = data_key_in self.key_output = data_key_out @@ -56,7 +66,8 @@ def get_scale(self) -> float: return self.scale.detach().cpu().tolist()[0] def forward(self, data: AtomGraphDataType) -> AtomGraphDataType: - data[self.key_output] = data[self.key_input] * self.scale + self.shift + inp = data[self.key_input].to(self.shift.dtype) + data[self.key_output] = inp * self.scale + self.shift return data @@ -79,11 +90,13 @@ def __init__( train_shift: bool = False, train_scale: bool = False, train_shift_scale: bool = False, + shift_scale_dtype: str = 'double', ) -> None: super().__init__() if train_shift_scale: train_shift = True train_scale = True + dtype = _resolve_shift_scale_dtype(shift_scale_dtype) assert isinstance(shift, float) or isinstance(shift, list) assert isinstance(scale, float) or isinstance(scale, list) @@ -105,10 +118,10 @@ def __init__( scale = [scale] * num_species if isinstance(scale, float) else scale self.shift = nn.Parameter( - torch.FloatTensor(shift), requires_grad=train_shift + torch.tensor(shift, dtype=dtype), requires_grad=train_shift ) self.scale = nn.Parameter( - torch.FloatTensor(scale), requires_grad=train_scale + torch.tensor(scale, dtype=dtype), requires_grad=train_scale ) self.key_input = data_key_in self.key_output = data_key_out @@ -165,9 +178,10 @@ def from_mappers( def forward(self, data: AtomGraphDataType) -> AtomGraphDataType: indices = data[self.key_indices] - data[self.key_output] = data[self.key_input] * self.scale[indices].view( - -1, 1 - ) + self.shift[indices].view(-1, 1) + scale = self.scale[indices].view(-1, 1) + shift = self.shift[indices].view(-1, 1) + inp = data[self.key_input].to(shift.dtype) + data[self.key_output] = inp * scale + shift return data @@ -193,16 +207,18 @@ def __init__( train_shift: bool = False, train_scale: bool = False, train_shift_scale: bool = False, + shift_scale_dtype: str = 'double', ) -> None: super().__init__() if train_shift_scale: train_shift = True train_scale = True + dtype = _resolve_shift_scale_dtype(shift_scale_dtype) self.shift = nn.Parameter( - torch.FloatTensor(shift), requires_grad=train_shift + torch.tensor(shift, dtype=dtype), requires_grad=train_shift ) self.scale = nn.Parameter( - torch.FloatTensor(scale), requires_grad=train_scale + torch.tensor(scale, dtype=dtype), requires_grad=train_scale ) self.key_input = data_key_in self.key_output = data_key_out @@ -371,9 +387,8 @@ def forward(self, data: AtomGraphDataType) -> AtomGraphDataType: if self.use_modal_wise_scale else self.scale[atom_indices] ) - data[self.key_output] = data[self.key_input] * scale.view( - -1, 1 - ) + shift.view(-1, 1) + inp = data[self.key_input].to(shift.dtype) + data[self.key_output] = inp * scale.view(-1, 1) + shift.view(-1, 1) return data diff --git a/sevenn/pair_e3gnn/pair_e3gnn.cpp b/sevenn/pair_e3gnn/pair_e3gnn.cpp index 587f4273..3912b41a 100644 --- a/sevenn/pair_e3gnn/pair_e3gnn.cpp +++ b/sevenn/pair_e3gnn/pair_e3gnn.cpp @@ -204,10 +204,10 @@ void PairE3GNN::compute(int eflag, int vflag) { // dE_dr auto grads = torch::autograd::grad({energy_tensor}, {edge_vec_device}); - torch::Tensor dE_dr = grads[0].to(torch::kCPU); + torch::Tensor dE_dr = grads[0].to(torch::kCPU).to(torch::kFloat); - eng_vdwl += energy_tensor.detach().to(torch::kCPU).item(); - torch::Tensor force_tensor = torch::zeros({nlocal, 3}); + eng_vdwl += energy_tensor.detach().to(torch::kCPU).item(); + torch::Tensor force_tensor = torch::zeros({nlocal, 3}, FLOAT_TYPE); auto _edge_idx_src_tensor = edge_idx_src_tensor.repeat_interleave(3).view({nedges, 3}); @@ -237,7 +237,7 @@ void PairE3GNN::compute(int eflag, int vflag) { diag, s12.unsqueeze(-1), s23.unsqueeze(-1), s31.unsqueeze(-1)}; auto voigt = torch::cat(voigt_list, 1); - torch::Tensor per_atom_stress_tensor = torch::zeros({nlocal, 6}); + torch::Tensor per_atom_stress_tensor = torch::zeros({nlocal, 6}, FLOAT_TYPE); auto _edge_idx_dst6_tensor = edge_idx_dst_tensor.repeat_interleave(6).view({nedges, 6}); per_atom_stress_tensor.scatter_reduce_(0, _edge_idx_dst6_tensor, voigt, @@ -271,8 +271,8 @@ void PairE3GNN::compute(int eflag, int vflag) { if (eflag_atom) { torch::Tensor atomic_energy_tensor = - output.at("atomic_energy").toTensor().to(torch::kCPU).view({nlocal}); - auto atomic_energy = atomic_energy_tensor.accessor(); + output.at("atomic_energy").toTensor().to(torch::kCPU).to(torch::kDouble).view({nlocal}); + auto atomic_energy = atomic_energy_tensor.accessor(); for (int gi = 0; gi < nlocal; gi++) { const int i = graph_index_to_i[gi]; eatom[i] += atomic_energy[gi]; diff --git a/sevenn/pair_e3gnn/pair_e3gnn_parallel.cpp b/sevenn/pair_e3gnn/pair_e3gnn_parallel.cpp index de8c6a11..b8f49e46 100644 --- a/sevenn/pair_e3gnn/pair_e3gnn_parallel.cpp +++ b/sevenn/pair_e3gnn/pair_e3gnn_parallel.cpp @@ -456,10 +456,10 @@ void PairE3GNNParallel::compute(int eflag, int vflag) { std::cout << world_rank << " Used/GraphSize: " << Mused / graph_size << "\n" << std::endl; } - eng_vdwl += energy_tensor.item(); // accumulate energy + eng_vdwl += energy_tensor.detach().to(torch::kCPU).item(); - dE_dr = dE_dr.to(torch::kCPU); - torch::Tensor force_tensor = torch::zeros({graph_indexer, 3}); + dE_dr = dE_dr.to(torch::kCPU).to(torch::kFloat); + torch::Tensor force_tensor = torch::zeros({graph_indexer, 3}, FLOAT_TYPE); auto _edge_idx_src_tensor = edge_idx_src_tensor.repeat_interleave(3).view({nedges, 3}); @@ -488,7 +488,8 @@ void PairE3GNNParallel::compute(int eflag, int vflag) { diag, s12.unsqueeze(-1), s23.unsqueeze(-1), s31.unsqueeze(-1)}; auto voigt = torch::cat(voigt_list, 1); - torch::Tensor per_atom_stress_tensor = torch::zeros({graph_indexer, 6}); + torch::Tensor per_atom_stress_tensor = + torch::zeros({graph_indexer, 6}, FLOAT_TYPE); auto _edge_idx_dst6_tensor = edge_idx_dst_tensor.repeat_interleave(6).view({nedges, 6}); per_atom_stress_tensor.scatter_reduce_(0, _edge_idx_dst6_tensor, voigt, @@ -507,8 +508,8 @@ void PairE3GNNParallel::compute(int eflag, int vflag) { if (eflag_atom) { torch::Tensor atomic_energy_tensor = - output.at("atomic_energy").toTensor().cpu().view({nlocal}); - auto atomic_energy = atomic_energy_tensor.accessor(); + output.at("atomic_energy").toTensor().cpu().to(torch::kDouble).view({nlocal}); + auto atomic_energy = atomic_energy_tensor.accessor(); for (int graph_idx = 0; graph_idx < nlocal; graph_idx++) { int i = graph_index_to_i[graph_idx]; eatom[i] += atomic_energy[graph_idx]; diff --git a/sevenn/scripts/deploy.py b/sevenn/scripts/deploy.py index ca718c31..594cdb0f 100644 --- a/sevenn/scripts/deploy.py +++ b/sevenn/scripts/deploy.py @@ -19,6 +19,7 @@ def deploy( modal: Optional[str] = None, use_flash: bool = False, use_oeq: bool = False, + shift_scale_dtype: str = 'double', ) -> None: if not (use_flash or use_oeq): warn_no_tp_accelerator('LAMMPS TorchScript deployment') @@ -30,6 +31,7 @@ def deploy( enable_flash=use_flash, enable_oeq=use_oeq, _flash_lammps=use_flash, + shift_scale_dtype=shift_scale_dtype, ), cp.config, ) @@ -69,6 +71,7 @@ def deploy( ) md_configs.update({'version': __version__}) md_configs.update({'dtype': config.pop(KEY.DTYPE, 'single')}) + md_configs.update({'shift_scale_dtype': shift_scale_dtype}) md_configs.update({'time': datetime.now().strftime('%Y-%m-%d')}) if fname.endswith('.pt') is False: @@ -83,6 +86,7 @@ def deploy_parallel( modal: Optional[str] = None, use_flash: bool = False, use_oeq: bool = False, + shift_scale_dtype: str = 'double', ) -> None: if not (use_flash or use_oeq): warn_no_tp_accelerator( @@ -100,12 +104,14 @@ def deploy_parallel( enable_flash=use_flash, enable_oeq=use_oeq, _flash_lammps=use_flash, + shift_scale_dtype=shift_scale_dtype, ), cp.config, ) config[KEY.CUEQUIVARIANCE_CONFIG] = {'use': False} config[KEY.USE_FLASH_TP] = use_flash config[KEY.USE_OEQ] = use_oeq + config[KEY.SHIFT_SCALE_DTYPE] = shift_scale_dtype config['_flash_lammps'] = use_flash model_state_dct = model.state_dict() @@ -164,6 +170,7 @@ def deploy_parallel( ) md_configs.update({'version': __version__}) md_configs.update({'dtype': config.pop(KEY.DTYPE, 'single')}) + md_configs.update({'shift_scale_dtype': shift_scale_dtype}) md_configs.update({'time': datetime.now().strftime('%Y-%m-%d')}) os.makedirs(fname, exist_ok=True) diff --git a/sevenn/torchsim.py b/sevenn/torchsim.py index 2f375542..cdc18bdd 100644 --- a/sevenn/torchsim.py +++ b/sevenn/torchsim.py @@ -81,6 +81,7 @@ def __init__( compute_atomic_virial: bool = False, device: torch.device | str = 'auto', dtype: torch.dtype = torch.float32, + shift_scale_dtype: str = 'double', ) -> None: """Initialize the SevenNetModel with specified configuration. @@ -101,7 +102,11 @@ def __init__( neighbor_list_fn (Callable): Neighbor list function to use. Default is torch_nl_linked_cell. device (torch.device | str): Device to run the model on - dtype (torch.dtype): Data type for computation + dtype (torch.dtype): TorchSim interface dtype. Only float32 is + supported; wrap with Float64Wrapper if outer TorchSim cell/state + arithmetic must run in double precision. + shift_scale_dtype (str): dtype for final energy shift/scale rescaling. + Defaults to 'double'; 'single' keeps legacy energy dtype. Raises: ImportError: if torch_sim is not installed @@ -144,6 +149,7 @@ def __init__( enable_flash=enable_flash, enable_cueq=enable_cueq, enable_oeq=enable_oeq, + shift_scale_dtype=shift_scale_dtype, ) _validate(model, modal) @@ -167,9 +173,6 @@ def __init__( self.model = model.to(self._device) self.model = self.model.eval() - if self._dtype is not None: - self.model = self.model.to(dtype=self._dtype) - self.implemented_properties = ['energy', 'forces', 'stress'] @property @@ -282,13 +285,13 @@ def forward(self, state: ts.SimState, **kwargs) -> dict[str, torch.Tensor]: forces = output[key.PRED_FORCE] if forces is not None: - results['forces'] = forces + results['forces'] = forces.to(dtype=self._dtype) stress = output[key.PRED_STRESS] if stress is not None: results['stress'] = -voigt_6_to_full_3x3_stress( stress[..., [0, 1, 2, 4, 5, 3]], - ) + ).to(dtype=self._dtype) results = {k: v.detach() for k, v in results.items()} @@ -323,6 +326,7 @@ def __init__( neighbor_list_fn: Callable | None = None, device: torch.device | str = 'auto', dtype: torch.dtype = torch.float32, + shift_scale_dtype: str = 'double', d3_mode: str = 'auto', d3_batch_threshold: int = 4, damping_type: str = 'damp_bj', @@ -346,6 +350,7 @@ def __init__( neighbor_list_fn=neighbor_list_fn, device=device, dtype=dtype, + shift_scale_dtype=shift_scale_dtype, ) self.d3_mode = d3_mode @@ -477,7 +482,7 @@ def _apply_batch_d3( ) results['energy'] += torch.from_numpy(d3_energy).to( - device=self._device, dtype=self._dtype, + device=self._device, dtype=results['energy'].dtype, ) results['forces'] += torch.from_numpy( np.ascontiguousarray(d3_forces), @@ -505,7 +510,8 @@ class Float64Wrapper(ModelInterface): Casts state tensors to float32 before calling the wrapped model, then casts outputs back to float64. Reports ``dtype=float64`` to torch-sim - so all optimizer / integrator arithmetic is done in double precision. + so all cell algebra / optimizer / integrator arithmetic is done in double. + The wrapped model still receives a float32 ``SimState``. This is needed because ``SumModel`` requires all children to share the same dtype, and ``D3DispersionModel`` defaults to float64. diff --git a/sevenn/train/loss.py b/sevenn/train/loss.py index 941eabcd..040fa064 100644 --- a/sevenn/train/loss.py +++ b/sevenn/train/loss.py @@ -75,6 +75,10 @@ def get_loss(self, batch_data: Dict[str, Any], model: Optional[Callable] = None) assert self.ref_key is not None return torch.zeros(1, device=batch_data[self.ref_key].device) + ref = ref.to(dtype=pred.dtype) + if w_tensor is not None: + w_tensor = w_tensor.to(dtype=pred.dtype) + loss = self.criterion(pred, ref) if self.use_weight: loss = torch.mean(loss * w_tensor) diff --git a/sevenn/util.py b/sevenn/util.py index 89d54065..1c323cf6 100644 --- a/sevenn/util.py +++ b/sevenn/util.py @@ -122,10 +122,14 @@ def model_from_checkpoint( enable_cueq: bool = False, enable_flash: bool = False, enable_oeq: bool = False, + shift_scale_dtype: str = 'double', ) -> Tuple[torch.nn.Module, Dict[str, Any]]: cp = load_checkpoint(checkpoint) model = cp.build_model( - enable_cueq=enable_cueq, enable_flash=enable_flash, enable_oeq=enable_oeq + enable_cueq=enable_cueq, + enable_flash=enable_flash, + enable_oeq=enable_oeq, + shift_scale_dtype=shift_scale_dtype, ) return model, cp.config From 2dc2434e8ad1ba2ee1d0abf5c74b0ea9f4adb851 Mon Sep 17 00:00:00 2001 From: alphalm4 Date: Thu, 16 Jul 2026 21:06:58 +0900 Subject: [PATCH 2/5] set stable lammps version for lmp-mliap --- docs/source/user_guide/lammps_mliap.md | 8 +++++--- 1 file changed, 5 insertions(+), 3 deletions(-) diff --git a/docs/source/user_guide/lammps_mliap.md b/docs/source/user_guide/lammps_mliap.md index 4b564295..f0f9a69a 100644 --- a/docs/source/user_guide/lammps_mliap.md +++ b/docs/source/user_guide/lammps_mliap.md @@ -1,7 +1,7 @@ # LAMMPS: ML-IAP :::{caution} -Currently the parallel implementation of LAMMPS/ML-IAP is not tested. +Parallel execution of LAMMPS/ML-IAP using CUDA-aware MPI is supported, but not yet sufficiently tested. For multi-rank LAMMPS runs, use the TorchScript `e3gnn/parallel` workflow described in {doc}`lammps_torch`. ::: ## Requirements @@ -22,10 +22,12 @@ Get LAMMPS source code: ```bash git clone https://github.com/lammps/lammps lammps-mliap cd lammps-mliap -git checkout ccca772 +git checkout stable_22Jul2025_update4 ``` :::{note} -We found that some of the latest versions of LAMMPS produce inconsistent energies. Therefore, we highly recommend using this specific commit. This restriction will be relaxed once consistency checks are completed. +SevenNet's ML-IAP tests are validated against the LAMMPS +`stable_22Jul2025_update4` release. Other LAMMPS versions may work, but should +be validated with the ML-IAP test suite before production use. ::: From e176ecf80baeecf912f62469140f5aa3e1b4e9fc Mon Sep 17 00:00:00 2001 From: alphalm4 Date: Thu, 16 Jul 2026 21:28:09 +0900 Subject: [PATCH 3/5] type matching --- tests/unit_tests/test_pretrained.py | 2 +- tests/unit_tests/test_shift_scale.py | 14 +++++++++----- 2 files changed, 10 insertions(+), 6 deletions(-) diff --git a/tests/unit_tests/test_pretrained.py b/tests/unit_tests/test_pretrained.py index 2bc2c645..7398fe2b 100644 --- a/tests/unit_tests/test_pretrained.py +++ b/tests/unit_tests/test_pretrained.py @@ -11,7 +11,7 @@ def acl(a, b, atol=1e-6): - return torch.allclose(a, b, atol=atol) + return torch.allclose(a.float(), b.float(), atol=atol) @pytest.fixture diff --git a/tests/unit_tests/test_shift_scale.py b/tests/unit_tests/test_shift_scale.py index 3d5a4eba..5806fedb 100644 --- a/tests/unit_tests/test_shift_scale.py +++ b/tests/unit_tests/test_shift_scale.py @@ -10,6 +10,10 @@ get_resolved_shift_scale, ) + +def acl(a, b): + return torch.allclose(a.float(), b.float()) + ################################################################################ # Tests for Rescale # ################################################################################ @@ -45,7 +49,7 @@ def test_rescale_forward(): # Check correctness expected_output = input_data * scale + shift - assert torch.allclose(out_data[KEY.ATOMIC_ENERGY], expected_output) + assert acl(out_data[KEY.ATOMIC_ENERGY], expected_output) def test_rescale_get_shift_and_scale(): @@ -86,8 +90,8 @@ def test_specieswise_rescale_init_list(): module = SpeciesWiseRescale(shift=shift, scale=scale) assert len(module.shift) == 3 assert len(module.scale) == 3 - assert torch.allclose(module.shift, torch.tensor([1.0, 2.0, 3.0])) - assert torch.allclose(module.scale, torch.tensor([2.0, 3.0, 4.0])) + assert acl(module.shift, torch.tensor([1.0, 2.0, 3.0])) + assert acl(module.scale, torch.tensor([2.0, 3.0, 4.0])) def test_specieswise_rescale_forward(): @@ -123,7 +127,7 @@ def test_specieswise_rescale_forward(): # For atom 2: scale=2, shift=1, input=3 => 3*2+1=7 expected = torch.tensor([[3.0], [15.0], [7.0]]) - assert torch.allclose(out['out'], expected) + assert acl(out['out'], expected) def test_specieswise_rescale_get_shift_scale(): @@ -212,7 +216,7 @@ def test_modalwise_rescale_forward(): # i=2 => modal_idx=1, atom_idx=0 => shift=5.0, scale=10.0 => out=2*10+5=25 # i=3 => modal_idx=1, atom_idx=1 => shift=15.0, scale=20.0 => out=2*20+15=55 expected = torch.tensor([[1.0], [12.0], [25.0], [55.0]]) - assert torch.allclose(out['out'], expected) + assert acl(out['out'], expected) def test_modalwise_rescale_get_shift_scale(): From 99bef7c10b3d96a714484fa11854feacdc6608a8 Mon Sep 17 00:00:00 2001 From: YutackPark Date: Tue, 21 Jul 2026 15:48:29 +0900 Subject: [PATCH 4/5] Replace shift_scale_dtype plumbing with SEVENN_SHIFT_SCALE_DTYPE env var The single-precision rescale path is a rarely-used backward-compat escape hatch, so threading a shift_scale_dtype argument through the calculator, checkpoint, deploy, mliap, torchsim, and CLI layers was excessive. Read the dtype from the SEVENN_SHIFT_SCALE_DTYPE env var at the single point where the shift/scale parameters are created (sevenn/nn/scale.py), matching the existing SEVENN_DEBUG-style toggles. Reverts the argument plumbing in _const, _keys, checkpoint, mliap, model_build, sevenn_get_model, and util (now identical to base). Keeps the behavioral changes: double-precision rescale forward casts, C++ energy in double, torchsim output casts, D3 output float32 cast, and loss dtype alignment. deploy still records the resolved dtype in metadata via resolve_shift_scale_dtype(). Co-Authored-By: Claude Opus 4.8 (1M context) --- sevenn/_const.py | 2 -- sevenn/_keys.py | 1 - sevenn/calculator.py | 10 +--------- sevenn/checkpoint.py | 8 -------- sevenn/main/sevenn_get_model.py | 21 ++------------------- sevenn/mliap.py | 5 +---- sevenn/model_build.py | 6 +----- sevenn/nn/scale.py | 22 ++++++++++++++-------- sevenn/scripts/deploy.py | 16 +++++++++------- sevenn/torchsim.py | 6 ------ sevenn/util.py | 6 +----- 11 files changed, 29 insertions(+), 74 deletions(-) diff --git a/sevenn/_const.py b/sevenn/_const.py index e16e19a4..6c78f9c8 100644 --- a/sevenn/_const.py +++ b/sevenn/_const.py @@ -127,7 +127,6 @@ def error_record_condition(x): KEY.TRAIN_SHIFT_SCALE: False, KEY.TRAIN_SHIFT: False, KEY.TRAIN_SCALE: False, - KEY.SHIFT_SCALE_DTYPE: 'double', # KEY.OPTIMIZE_BY_REDUCE: True, # deprecated, always True KEY.USE_BIAS_IN_LINEAR: False, KEY.USE_MODAL_NODE_EMBEDDING: False, @@ -172,7 +171,6 @@ def error_record_condition(x): KEY.TRAIN_SHIFT_SCALE: bool, KEY.TRAIN_SHIFT: bool, KEY.TRAIN_SCALE: bool, - KEY.SHIFT_SCALE_DTYPE: lambda x: x in ['single', 'double'], KEY.TRAIN_DENOMINTAOR: bool, KEY.USE_BIAS_IN_LINEAR: bool, KEY.USE_MODAL_NODE_EMBEDDING: bool, diff --git a/sevenn/_keys.py b/sevenn/_keys.py index c2b64b0a..8773a099 100644 --- a/sevenn/_keys.py +++ b/sevenn/_keys.py @@ -229,7 +229,6 @@ CONV_DENOMINATOR = 'conv_denominator' SHIFT = 'shift' SCALE = 'scale' -SHIFT_SCALE_DTYPE = 'shift_scale_dtype' LOADER_KWARGS = 'loader_kwargs' USE_SPECIES_WISE_SHIFT_SCALE = 'use_species_wise_shift_scale' diff --git a/sevenn/calculator.py b/sevenn/calculator.py index e6642fd5..a1256b5f 100644 --- a/sevenn/calculator.py +++ b/sevenn/calculator.py @@ -37,7 +37,6 @@ def __init__( enable_oeq: bool = False, compute_atomic_virial: bool = False, sevennet_config: Optional[Dict] = None, # Not used in logic, just meta info - shift_scale_dtype: str = 'double', **kwargs, ) -> None: """Initialize SevenNetCalculator. @@ -68,14 +67,10 @@ def __init__( Not used, but can be used to carry meta information of this calculator compute_atomic_virial: bool, default=False If True, request per-atom virial output (`stresses`) at runtime. - shift_scale_dtype: str, default='double' - dtype of the final shift/scale (rescale) parameters used at inference. - 'single' is only for backward reproducibility. """ super().__init__(**kwargs) self.sevennet_config = None self.compute_atomic_virial = compute_atomic_virial - self.shift_scale_dtype = shift_scale_dtype if isinstance(model, pathlib.PurePath): model = str(model) @@ -123,8 +118,7 @@ def __init__( cp = util.load_checkpoint(model) model_loaded = cp.build_model( - enable_cueq=enable_cueq, enable_flash=enable_flash, enable_oeq=enable_oeq, # noqa: E501 - shift_scale_dtype=shift_scale_dtype, + enable_cueq=enable_cueq, enable_flash=enable_flash, enable_oeq=enable_oeq # noqa: E501 ) model_loaded.set_is_batch_data(False) @@ -254,7 +248,6 @@ def __init__( functional_name: str = 'pbe', vdw_cutoff: float = 9000, # au^2, 0.52917726 angstrom = 1 au cn_cutoff: float = 1600, # au^2, 0.52917726 angstrom = 1 au - shift_scale_dtype: Optional[str] = 'double', **kwargs, # pass extra kwargs to both calculators ) -> None: """Initialize SevenNetD3Calculator. CUDA required. @@ -315,7 +308,6 @@ def __init__( enable_flash=enable_flash, enable_oeq=enable_oeq, sevennet_config=sevennet_config, - shift_scale_dtype=shift_scale_dtype, **kwargs, ) diff --git a/sevenn/checkpoint.py b/sevenn/checkpoint.py index 18f9fa3b..8e4a704e 100644 --- a/sevenn/checkpoint.py +++ b/sevenn/checkpoint.py @@ -357,7 +357,6 @@ def build_model( enable_flash: bool = False, enable_oeq: bool = False, _flash_lammps: bool = False, - shift_scale_dtype: str = 'double', ) -> AtomGraphSequential: from .model_build import build_E3_equivariant_model @@ -377,13 +376,6 @@ def build_model( cfg_new['_flash_lammps'] = _flash_lammps cfg_new[KEY.USE_OEQ] = enable_oeq - if shift_scale_dtype not in ('single', 'double'): - raise ValueError( - "shift_scale_dtype must be 'single' or 'double', " - f'got {shift_scale_dtype!r}' - ) - cfg_new[KEY.SHIFT_SCALE_DTYPE] = shift_scale_dtype - if (cp_using_cueq, cp_using_flash, cp_using_oeq) == ( enable_cueq, enable_flash, diff --git a/sevenn/main/sevenn_get_model.py b/sevenn/main/sevenn_get_model.py index d8fe9f0f..8f7bcdb6 100644 --- a/sevenn/main/sevenn_get_model.py +++ b/sevenn/main/sevenn_get_model.py @@ -62,15 +62,6 @@ def add_args(parser): help='Use LAMMPS ML-IAP interface.', action='store_true', ) - ag.add_argument( - '--shift_scale_dtype', - choices=('double', 'single'), - default='double', - help=( - 'dtype of final shift/scale energy rescaling at inference. ' - 'Default is double; single is for backward compatibility.' - ), - ) def run(args): @@ -85,7 +76,6 @@ def run(args): use_cueq = args.enable_cueq use_oeq = args.enable_oeq use_mliap = args.use_mliap - shift_scale_dtype = args.shift_scale_dtype # Check dependencies if use_flash: @@ -129,15 +119,9 @@ def run(args): from sevenn.scripts.deploy import deploy, deploy_parallel if get_serial: - deploy( - checkpoint_path, output_prefix, modal, use_flash=use_flash, use_oeq=use_oeq, # noqa: E501 - shift_scale_dtype=shift_scale_dtype - ) + deploy(checkpoint_path, output_prefix, modal, use_flash=use_flash, use_oeq=use_oeq) # noqa: E501 else: - deploy_parallel( - checkpoint_path, output_prefix, modal, use_flash=use_flash, use_oeq=use_oeq, # noqa: E501 - shift_scale_dtype=shift_scale_dtype - ) + deploy_parallel(checkpoint_path, output_prefix, modal, use_flash=use_flash, use_oeq=use_oeq) # noqa: E501 else: from sevenn import mliap @@ -155,7 +139,6 @@ def run(args): use_cueq=use_cueq, use_flash=use_flash, use_oeq=use_oeq, - shift_scale_dtype=shift_scale_dtype, ) torch.save(mliap_module, output_prefix) diff --git a/sevenn/mliap.py b/sevenn/mliap.py index 70fb9d0b..513d55c7 100644 --- a/sevenn/mliap.py +++ b/sevenn/mliap.py @@ -87,7 +87,6 @@ def __init__( modal: Optional[str] = None use_cueq: bool = False use_flash: bool = False - shift_scale_dtype: str = 'double' """ super().__init__() @@ -113,7 +112,6 @@ def __init__( self.use_cueq = kwargs.get('use_cueq', False) self.use_flash = kwargs.get('use_flash', False) self.use_oeq = kwargs.get('use_oeq', False) - self.shift_scale_dtype = kwargs.get('shift_scale_dtype', 'double') self.modal = kwargs.get('modal', None) # extract configs @@ -155,8 +153,7 @@ def _ensure_model_initialized(self): print('[INFO] Lazy initializing SevenNet model...', flush=True) print(f'[INFO] cueq={self.use_cueq}, flashTP={self.use_flash}, oeq={self.use_oeq}', flush=True) # noqa: E501 model = self.cp.build_model( - enable_cueq=self.use_cueq, enable_flash=self.use_flash, enable_oeq=self.use_oeq, # noqa: E501 - shift_scale_dtype=self.shift_scale_dtype + enable_cueq=self.use_cueq, enable_flash=self.use_flash, enable_oeq=self.use_oeq # noqa: E501 ) for k, module in model._modules.items(): diff --git a/sevenn/model_build.py b/sevenn/model_build.py index 92582e64..f0482f11 100644 --- a/sevenn/model_build.py +++ b/sevenn/model_build.py @@ -167,11 +167,7 @@ def init_shift_scale( shift_scale.append(s) shift, scale = shift_scale - ss_kwargs = { - 'train_shift': train_shift, - 'train_scale': train_scale, - 'shift_scale_dtype': config.get(KEY.SHIFT_SCALE_DTYPE, 'double'), - } + ss_kwargs = {'train_shift': train_shift, 'train_scale': train_scale} rescale_module = None if config.get(KEY.USE_MODALITY, False): rescale_module = ModalWiseRescale.from_mappers( # type: ignore diff --git a/sevenn/nn/scale.py b/sevenn/nn/scale.py index 0b975834..afe4a532 100644 --- a/sevenn/nn/scale.py +++ b/sevenn/nn/scale.py @@ -1,3 +1,4 @@ +import os from typing import Any, Dict, List, Optional, Union import torch @@ -7,13 +8,21 @@ import sevenn._keys as KEY from sevenn._const import NUM_UNIV_ELEMENT, AtomGraphDataType +# Precision of the final energy shift/scale (rescale) parameters. Defaults to +# double for numerical consistency; set SEVENN_SHIFT_SCALE_DTYPE='single' only +# to reproduce the legacy float32 behavior. The dtype is frozen into the +# parameters when the model is built, so it must be set at build/deploy time +# (not at inference time). +SHIFT_SCALE_DTYPE_ENV = 'SEVENN_SHIFT_SCALE_DTYPE' -def _resolve_shift_scale_dtype(dtype: str) -> torch.dtype: + +def resolve_shift_scale_dtype() -> torch.dtype: + dtype = os.environ.get(SHIFT_SCALE_DTYPE_ENV, 'double') if dtype == 'single': return torch.float32 if dtype == 'double': return torch.float64 - raise ValueError(f'Unsupported shift/scale dtype: {dtype}') + raise ValueError(f'Unsupported shift/scale dtype: {dtype!r}') def _as_univ( @@ -41,7 +50,6 @@ def __init__( train_shift: bool = False, train_scale: bool = False, train_shift_scale: bool = False, - shift_scale_dtype: str = 'double', **kwargs, ) -> None: assert isinstance(shift, float) and isinstance(scale, float) @@ -49,7 +57,7 @@ def __init__( if train_shift_scale: train_shift = True train_scale = True - dtype = _resolve_shift_scale_dtype(shift_scale_dtype) + dtype = resolve_shift_scale_dtype() self.shift = nn.Parameter( torch.tensor([shift], dtype=dtype), requires_grad=train_shift ) @@ -90,13 +98,12 @@ def __init__( train_shift: bool = False, train_scale: bool = False, train_shift_scale: bool = False, - shift_scale_dtype: str = 'double', ) -> None: super().__init__() if train_shift_scale: train_shift = True train_scale = True - dtype = _resolve_shift_scale_dtype(shift_scale_dtype) + dtype = resolve_shift_scale_dtype() assert isinstance(shift, float) or isinstance(shift, list) assert isinstance(scale, float) or isinstance(scale, list) @@ -207,13 +214,12 @@ def __init__( train_shift: bool = False, train_scale: bool = False, train_shift_scale: bool = False, - shift_scale_dtype: str = 'double', ) -> None: super().__init__() if train_shift_scale: train_shift = True train_scale = True - dtype = _resolve_shift_scale_dtype(shift_scale_dtype) + dtype = resolve_shift_scale_dtype() self.shift = nn.Parameter( torch.tensor(shift, dtype=dtype), requires_grad=train_shift ) diff --git a/sevenn/scripts/deploy.py b/sevenn/scripts/deploy.py index 594cdb0f..af02e300 100644 --- a/sevenn/scripts/deploy.py +++ b/sevenn/scripts/deploy.py @@ -10,6 +10,7 @@ import sevenn._keys as KEY from sevenn import __version__ from sevenn.model_build import build_E3_equivariant_model +from sevenn.nn.scale import resolve_shift_scale_dtype from sevenn.util import load_checkpoint, warn_no_tp_accelerator @@ -19,7 +20,6 @@ def deploy( modal: Optional[str] = None, use_flash: bool = False, use_oeq: bool = False, - shift_scale_dtype: str = 'double', ) -> None: if not (use_flash or use_oeq): warn_no_tp_accelerator('LAMMPS TorchScript deployment') @@ -31,7 +31,6 @@ def deploy( enable_flash=use_flash, enable_oeq=use_oeq, _flash_lammps=use_flash, - shift_scale_dtype=shift_scale_dtype, ), cp.config, ) @@ -69,9 +68,12 @@ def deploy( md_configs.update( {'model_type': config.pop(KEY.MODEL_TYPE, 'E3_equivariant_model')} ) + ss_dtype = resolve_shift_scale_dtype() md_configs.update({'version': __version__}) md_configs.update({'dtype': config.pop(KEY.DTYPE, 'single')}) - md_configs.update({'shift_scale_dtype': shift_scale_dtype}) + md_configs.update( + {'shift_scale_dtype': 'single' if ss_dtype == torch.float32 else 'double'} + ) md_configs.update({'time': datetime.now().strftime('%Y-%m-%d')}) if fname.endswith('.pt') is False: @@ -86,7 +88,6 @@ def deploy_parallel( modal: Optional[str] = None, use_flash: bool = False, use_oeq: bool = False, - shift_scale_dtype: str = 'double', ) -> None: if not (use_flash or use_oeq): warn_no_tp_accelerator( @@ -104,14 +105,12 @@ def deploy_parallel( enable_flash=use_flash, enable_oeq=use_oeq, _flash_lammps=use_flash, - shift_scale_dtype=shift_scale_dtype, ), cp.config, ) config[KEY.CUEQUIVARIANCE_CONFIG] = {'use': False} config[KEY.USE_FLASH_TP] = use_flash config[KEY.USE_OEQ] = use_oeq - config[KEY.SHIFT_SCALE_DTYPE] = shift_scale_dtype config['_flash_lammps'] = use_flash model_state_dct = model.state_dict() @@ -168,9 +167,12 @@ def deploy_parallel( md_configs.update( {'model_type': config.pop(KEY.MODEL_TYPE, 'E3_equivariant_model')} ) + ss_dtype = resolve_shift_scale_dtype() md_configs.update({'version': __version__}) md_configs.update({'dtype': config.pop(KEY.DTYPE, 'single')}) - md_configs.update({'shift_scale_dtype': shift_scale_dtype}) + md_configs.update( + {'shift_scale_dtype': 'single' if ss_dtype == torch.float32 else 'double'} + ) md_configs.update({'time': datetime.now().strftime('%Y-%m-%d')}) os.makedirs(fname, exist_ok=True) diff --git a/sevenn/torchsim.py b/sevenn/torchsim.py index cdc18bdd..f4706738 100644 --- a/sevenn/torchsim.py +++ b/sevenn/torchsim.py @@ -81,7 +81,6 @@ def __init__( compute_atomic_virial: bool = False, device: torch.device | str = 'auto', dtype: torch.dtype = torch.float32, - shift_scale_dtype: str = 'double', ) -> None: """Initialize the SevenNetModel with specified configuration. @@ -105,8 +104,6 @@ def __init__( dtype (torch.dtype): TorchSim interface dtype. Only float32 is supported; wrap with Float64Wrapper if outer TorchSim cell/state arithmetic must run in double precision. - shift_scale_dtype (str): dtype for final energy shift/scale rescaling. - Defaults to 'double'; 'single' keeps legacy energy dtype. Raises: ImportError: if torch_sim is not installed @@ -149,7 +146,6 @@ def __init__( enable_flash=enable_flash, enable_cueq=enable_cueq, enable_oeq=enable_oeq, - shift_scale_dtype=shift_scale_dtype, ) _validate(model, modal) @@ -326,7 +322,6 @@ def __init__( neighbor_list_fn: Callable | None = None, device: torch.device | str = 'auto', dtype: torch.dtype = torch.float32, - shift_scale_dtype: str = 'double', d3_mode: str = 'auto', d3_batch_threshold: int = 4, damping_type: str = 'damp_bj', @@ -350,7 +345,6 @@ def __init__( neighbor_list_fn=neighbor_list_fn, device=device, dtype=dtype, - shift_scale_dtype=shift_scale_dtype, ) self.d3_mode = d3_mode diff --git a/sevenn/util.py b/sevenn/util.py index 1c323cf6..89d54065 100644 --- a/sevenn/util.py +++ b/sevenn/util.py @@ -122,14 +122,10 @@ def model_from_checkpoint( enable_cueq: bool = False, enable_flash: bool = False, enable_oeq: bool = False, - shift_scale_dtype: str = 'double', ) -> Tuple[torch.nn.Module, Dict[str, Any]]: cp = load_checkpoint(checkpoint) model = cp.build_model( - enable_cueq=enable_cueq, - enable_flash=enable_flash, - enable_oeq=enable_oeq, - shift_scale_dtype=shift_scale_dtype, + enable_cueq=enable_cueq, enable_flash=enable_flash, enable_oeq=enable_oeq ) return model, cp.config From a52fe390e0d90b8d71c5c1f34f1d0f6e5dc715fe Mon Sep 17 00:00:00 2001 From: YutackPark Date: Tue, 21 Jul 2026 16:17:57 +0900 Subject: [PATCH 5/5] refactor --- sevenn/nn/scale.py | 7 ++----- sevenn/scripts/deploy.py | 9 --------- 2 files changed, 2 insertions(+), 14 deletions(-) diff --git a/sevenn/nn/scale.py b/sevenn/nn/scale.py index afe4a532..90073c30 100644 --- a/sevenn/nn/scale.py +++ b/sevenn/nn/scale.py @@ -10,14 +10,11 @@ # Precision of the final energy shift/scale (rescale) parameters. Defaults to # double for numerical consistency; set SEVENN_SHIFT_SCALE_DTYPE='single' only -# to reproduce the legacy float32 behavior. The dtype is frozen into the -# parameters when the model is built, so it must be set at build/deploy time -# (not at inference time). -SHIFT_SCALE_DTYPE_ENV = 'SEVENN_SHIFT_SCALE_DTYPE' +# to reproduce the legacy float32 behavior. def resolve_shift_scale_dtype() -> torch.dtype: - dtype = os.environ.get(SHIFT_SCALE_DTYPE_ENV, 'double') + dtype = os.environ.get('SEVENN_SHIFT_SCALE_DTYPE', 'double') if dtype == 'single': return torch.float32 if dtype == 'double': diff --git a/sevenn/scripts/deploy.py b/sevenn/scripts/deploy.py index af02e300..ca718c31 100644 --- a/sevenn/scripts/deploy.py +++ b/sevenn/scripts/deploy.py @@ -10,7 +10,6 @@ import sevenn._keys as KEY from sevenn import __version__ from sevenn.model_build import build_E3_equivariant_model -from sevenn.nn.scale import resolve_shift_scale_dtype from sevenn.util import load_checkpoint, warn_no_tp_accelerator @@ -68,12 +67,8 @@ def deploy( md_configs.update( {'model_type': config.pop(KEY.MODEL_TYPE, 'E3_equivariant_model')} ) - ss_dtype = resolve_shift_scale_dtype() md_configs.update({'version': __version__}) md_configs.update({'dtype': config.pop(KEY.DTYPE, 'single')}) - md_configs.update( - {'shift_scale_dtype': 'single' if ss_dtype == torch.float32 else 'double'} - ) md_configs.update({'time': datetime.now().strftime('%Y-%m-%d')}) if fname.endswith('.pt') is False: @@ -167,12 +162,8 @@ def deploy_parallel( md_configs.update( {'model_type': config.pop(KEY.MODEL_TYPE, 'E3_equivariant_model')} ) - ss_dtype = resolve_shift_scale_dtype() md_configs.update({'version': __version__}) md_configs.update({'dtype': config.pop(KEY.DTYPE, 'single')}) - md_configs.update( - {'shift_scale_dtype': 'single' if ss_dtype == torch.float32 else 'double'} - ) md_configs.update({'time': datetime.now().strftime('%Y-%m-%d')}) os.makedirs(fname, exist_ok=True)