File size: 1,545 Bytes
f15d29e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
import argparse
import json
import os
from pathlib import Path

from model.common.utils.data_classes import MatterGenCheckpointInfo
from model.generator import CrystalGenerator


def main():
    parser = argparse.ArgumentParser(description="Generate crystals with MatterGen")
    parser.add_argument("--checkpoint", required=True)
    parser.add_argument("--output", default="outputs/mattergen")
    parser.add_argument("--batch-size", type=int, default=1)
    parser.add_argument("--num-batches", type=int, default=1)
    parser.add_argument(
        "--properties",
        type=json.loads,
        default=None,
        help='Condition values as JSON, for example {"dft_mag_density": 0.15}.',
    )
    parser.add_argument("--record-trajectories", action="store_true")
    args = parser.parse_args()
    if not os.path.isdir(args.checkpoint):
        parser.error(f"checkpoint directory does not exist: {args.checkpoint}")
    checkpoint_info = MatterGenCheckpointInfo(
        model_path=Path(args.checkpoint).expanduser().resolve(),
        load_epoch="last",
    )
    generator = CrystalGenerator(
        checkpoint_info=checkpoint_info,
        batch_size=args.batch_size,
        num_batches=args.num_batches,
        properties_to_condition_on=args.properties,
        record_trajectories=args.record_trajectories,
    )
    structures = generator.generate(
        output_dir=Path(args.output).expanduser().resolve()
    )
    print(f"Generated {len(structures)} structures in {args.output}")


if __name__ == "__main__":
    main()