Download VESBD-ODE/train.yaml from OMatG/MPTS-52-CSP: direct link, hf CLI and curl.
- Browser
- Download file 4.96 kB
-
https://huggingface.co/OMatG/MPTS-52-CSP/resolve/main/VESBD-ODE/train.yaml
- Command line
-
hf download hf://OMatG/MPTS-52-CSP/VESBD-ODE/train.yaml
-
curl -L -o train.yaml https://huggingface.co/OMatG/MPTS-52-CSP/resolve/main/VESBD-ODE/train.yaml
4.96 kB
| model: | |
| si: | |
| class_path: omg.si.stochastic_interpolants.StochasticInterpolants | |
| init_args: | |
| stochastic_interpolants: | |
| # chemical species | |
| - class_path: omg.si.single_stochastic_interpolant_identity.SingleStochasticInterpolantIdentity | |
| # fractional coordinates | |
| - class_path: omg.si.single_stochastic_interpolant_os.SingleStochasticInterpolantOS | |
| init_args: | |
| interpolant: | |
| class_path: omg.si.interpolants.PeriodicScoreBasedDiffusionModelInterpolantVE | |
| init_args: | |
| sigma: | |
| class_path: omg.si.sigma.GeometricSigma | |
| init_args: | |
| sigma_min: 0.004705415831077799 | |
| sigma_max: 0.9967130801483843 | |
| epsilon: null | |
| differential_equation_type: "ODE" | |
| integrator_kwargs: | |
| method: "euler" | |
| velocity_annealing_factor: 8.284579088906593 | |
| correct_center_of_mass_motion: true | |
| predict_velocity: true | |
| # lattice vectors | |
| - class_path: omg.si.single_stochastic_interpolant.SingleStochasticInterpolant | |
| init_args: | |
| interpolant: omg.si.interpolants.LinearInterpolant | |
| gamma: | |
| class_path: omg.si.gamma.LatentGammaSqrt | |
| init_args: | |
| a: 0.016616684357970132 | |
| epsilon: | |
| class_path: omg.si.epsilon.VanishingEpsilon | |
| init_args: | |
| c: 3.9372558236242052 | |
| mu: 0.2649556265396099 | |
| sigma: 0.03578203230805775 | |
| differential_equation_type: "SDE" | |
| integrator_kwargs: | |
| method: "euler" | |
| dt: 0.0015144158387556672 | |
| velocity_annealing_factor: 0.42775377056075214 | |
| correct_center_of_mass_motion: false | |
| data_fields: | |
| # if the order of the data_fields changes, | |
| # the order of the above StochasticInterpolant inputs must also change | |
| - "species" | |
| - "pos" | |
| - "cell" | |
| integration_time_steps: 660 | |
| relative_si_costs: | |
| species_loss: 0.0 | |
| pos_loss_b: 0.9813067351598369 | |
| cell_loss_b: 0.0005256953168558359 | |
| cell_loss_z: 0.018167569523307267 | |
| sampler: | |
| class_path: omg.sampler.IndependentSampler | |
| init_args: | |
| pos_distribution: | |
| class_path: omg.sampler.position_distributions.NormalPositionDistribution | |
| init_args: | |
| scale: 9.77149759679434 | |
| cell_distribution: | |
| class_path: omg.sampler.cell_distributions.InformedLatticeDistribution | |
| init_args: | |
| dataset_name: mpts_52 | |
| species_distribution: | |
| class_path: omg.sampler.species_distributions.MirrorSpecies | |
| model: | |
| class_path: omg.model.model.Model | |
| init_args: | |
| encoder: | |
| class_path: omg.model.encoders.cspnet_full.CSPNetFull | |
| head: | |
| class_path: omg.model.heads.pass_through.PassThrough | |
| time_embedder: | |
| class_path: omg.model.model_utils.SinusoidalTimeEmbeddings | |
| init_args: | |
| dim: 256 | |
| use_min_perm_dist: False | |
| float_32_matmul_precision: "high" | |
| validation_mode: "match_rate" | |
| number_cpus: 7 | |
| dataset_name: "mpts_52" | |
| data: | |
| train_dataset: | |
| class_path: omg.datamodule.StructureDataset | |
| init_args: | |
| file_path: "data/mpts_52/train.lmdb" | |
| lazy_storage: True | |
| niggli_reduce: False | |
| val_dataset: | |
| class_path: omg.datamodule.StructureDataset | |
| init_args: | |
| file_path: "data/mpts_52/val.lmdb" | |
| lazy_storage: True | |
| niggli_reduce: False | |
| pred_dataset: | |
| class_path: omg.datamodule.StructureDataset | |
| init_args: | |
| file_path: "data/mpts_52/test.lmdb" | |
| lazy_storage: True | |
| niggli_reduce: False | |
| batch_size: 256 | |
| num_workers: 4 | |
| pin_memory: True | |
| persistent_workers: True | |
| trainer: | |
| callbacks: | |
| - class_path: lightning.pytorch.callbacks.ModelCheckpoint | |
| init_args: | |
| filename: "best_val_loss_total" | |
| save_top_k: 1 | |
| monitor: "val_loss_total" | |
| save_weights_only: true | |
| - class_path: lightning.pytorch.callbacks.ModelCheckpoint | |
| init_args: | |
| filename: "best_val_match_rate" | |
| save_top_k: 1 | |
| monitor: "match_rate" | |
| save_weights_only: true | |
| mode: 'max' | |
| - class_path: lightning.pytorch.callbacks.ModelCheckpoint | |
| init_args: | |
| filename: "best_val_rmsd" | |
| save_top_k: 1 | |
| monitor: "mean_rmsd" | |
| save_weights_only: true | |
| - class_path: lightning.pytorch.callbacks.ModelCheckpoint | |
| init_args: | |
| save_top_k: -1 # Store every checkpoint after 100 epochs. | |
| monitor: "val_loss_total" | |
| every_n_epochs: 100 | |
| save_weights_only: false | |
| gradient_clip_val: 0.5 | |
| num_sanity_val_steps: 0 | |
| precision: "32-true" | |
| max_epochs: 2000 | |
| enable_progress_bar: true | |
| limit_val_batches: 0.5 | |
| check_val_every_n_epoch: 100 | |
| optimizer: | |
| class_path: torch.optim.Adam | |
| init_args: | |
| lr: 0.000296636127734534 | |