(args)
| 882 | return args |
| 883 | |
| 884 | def main(args): |
| 885 | pl.seed_everything(args.seed, workers=True) |
| 886 | |
| 887 | # initialize data module |
| 888 | data = DataModule(args) |
| 889 | data.prepare_dataset() |
| 890 | |
| 891 | default = ",".join(str(i) for i in range(torch.cuda.device_count())) |
| 892 | cuda_visible_devices = os.getenv("CUDA_VISIBLE_DEVICES", default=default).split(",") |
| 893 | dir_name = f"output_ngpus_{len(cuda_visible_devices)}_bs_{args.batch_size}_lr_{args.lr}_seed_{args.seed}" + \ |
| 894 | f"_reload_{args.reload}_lmax_{args.lmax}_vnorm_{args.vecnorm_type}" + \ |
| 895 | f"_vertex_{args.vertex_type}_L{args.num_layers}_D{args.embedding_dimension}_H{args.num_heads}" + \ |
| 896 | f"_cutoff_{args.cutoff}_E{args.energy_weight}_F{args.force_weight}_loss_{args.loss_type}" |
| 897 | |
| 898 | if args.load_model is None: |
| 899 | args.log_dir = os.path.join(args.out_dir, args.log_dir , dir_name) |
| 900 | if os.path.exists(args.log_dir): |
| 901 | if os.path.exists(os.path.join(args.log_dir, "last.ckpt")): |
| 902 | args.load_model = os.path.join(args.log_dir, "last.ckpt") |
| 903 | csv_path = os.path.join(args.log_dir, "metrics.csv") |
| 904 | while os.path.exists(csv_path): |
| 905 | csv_path = csv_path + '.bak' |
| 906 | if os.path.exists(os.path.join(args.log_dir, "metrics.csv")): |
| 907 | os.rename(os.path.join(args.log_dir, "metrics.csv"), csv_path) |
| 908 | |
| 909 | prior = None |
| 910 | if args.prior_model: |
| 911 | assert hasattr(priors, args.prior_model), ( |
| 912 | f"Unknown prior model {args['prior_model']}. " |
| 913 | f"Available models are {', '.join(priors.__all__)}" |
| 914 | ) |
| 915 | # initialize the prior model |
| 916 | prior = getattr(priors, args.prior_model)(dataset=data.dataset) |
| 917 | args.prior_args = prior.get_init_args() |
| 918 | |
| 919 | # initialize lightning module |
| 920 | model = LNNP(args, prior_model=prior, mean=data.mean, std=data.std) |
| 921 | |
| 922 | dir_path = os.path.join(args.out_dir, "ckpt") |
| 923 | |
| 924 | if args.task == "train": |
| 925 | |
| 926 | checkpoint_callback = ModelCheckpoint( |
| 927 | dirpath=dir_path, |
| 928 | monitor="val_loss", |
| 929 | save_top_k=1, |
| 930 | save_last=True, |
| 931 | every_n_epochs=args.save_interval, |
| 932 | filename="best", |
| 933 | ) |
| 934 | |
| 935 | early_stopping = EarlyStopping("val_loss", patience=args.early_stopping_patience) |
| 936 | |
| 937 | tb_logger = TensorBoardLogger(os.getenv("TENSORBOARD_LOG_PATH", "/tensorboard_logs/"), name="", version="", default_hp_metric=False) |
| 938 | csv_logger = CSVLogger(args.log_dir, name="", version="") |
| 939 | ddp_plugin = DDPStrategy(find_unused_parameters=False) |
| 940 | |
| 941 | trainer = pl.Trainer( |
no test coverage detected