sevenn_get_model.py 1.74 KB
Newer Older
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
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
import argparse
import os

from sevenn import __version__

description_get_model = (
    'deploy LAMMPS model from the checkpoint'
)
checkpoint_help = (
    'path to the checkpoint | SevenNet-0 | 7net-0 |'
    ' {SevenNet-0|7net-0}_{11July2024|22May2024}'
)
output_name_help = 'filename prefix'
get_parallel_help = 'deploy parallel model'


def add_parser(subparsers):
    ag = subparsers.add_parser(
        'get_model', help=description_get_model, aliases=['deploy']
    )
    add_args(ag)


def add_args(parser):
    ag = parser
    ag.add_argument('checkpoint', help=checkpoint_help, type=str)
    ag.add_argument(
        '-o', '--output_prefix', nargs='?', help=output_name_help, type=str
    )
    ag.add_argument(
        '-p', '--get_parallel', help=get_parallel_help, action='store_true'
    )
    ag.add_argument(
        '-m',
        '--modal',
        help='Modality of multi-modal model',
        type=str,
    )


def run(args):
    import sevenn.util
    from sevenn.scripts.deploy import deploy, deploy_parallel

    checkpoint = args.checkpoint
    output_prefix = args.output_prefix
    get_parallel = args.get_parallel
    get_serial = not get_parallel
    modal = args.modal

    if output_prefix is None:
        output_prefix = 'deployed_parallel' if not get_serial else 'deployed_serial'

    checkpoint_path = None
    if os.path.isfile(checkpoint):
        checkpoint_path = checkpoint
    else:
        checkpoint_path = sevenn.util.pretrained_name_to_path(checkpoint)

    if get_serial:
        deploy(checkpoint_path, output_prefix, modal)
    else:
        deploy_parallel(checkpoint_path, output_prefix, modal)


# legacy way
def main():
    ag = argparse.ArgumentParser(description=description_get_model)
    add_args(ag)
    run(ag.parse_args())