Spaces:
Paused
Paused
File size: 1,760 Bytes
5716801 | 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 | from argparse import ArgumentError, ArgumentParser, Namespace
import sys
import logging
from pathlib import Path
from chemprop.cli.utils import Subcommand
from chemprop.utils.v1_to_v2 import convert_model_file_v1_to_v2
logger = logging.getLogger(__name__)
class ConvertSubcommand(Subcommand):
COMMAND = "convert"
HELP = "convert a v1 model checkpoint (.pt) to a v2 model checkpoint (.ckpt)"
@classmethod
def add_args(cls, parser: ArgumentParser) -> ArgumentParser:
parser.add_argument(
"-i",
"--input-path",
required=True,
type=Path,
help="The path to a v1 model .pt checkpoint file.",
)
parser.add_argument(
"-o",
"--output-path",
type=Path,
help="The path to which the converted model will be saved. Defaults to 'CURRENT_DIRECTORY/STEM_OF_INPUT_v2.ckpt'",
)
return parser
@classmethod
def func(cls, args: Namespace):
if args.output_path is None:
args.output_path = Path(args.input_path.stem + "_v2.ckpt")
if args.output_path.suffix != ".ckpt":
raise ArgumentError(
argument=None, message=f"Output must be a `.ckpt` file. Got {args.output_path}"
)
logger.info(
f"Converting v1 model checkpoint '{args.input_path}' to v2 model checkpoint '{args.output_path}'..."
)
convert_model_file_v1_to_v2(args.input_path, args.output_path)
if __name__ == "__main__":
parser = ArgumentParser()
parser = ConvertSubcommand.add_args(parser)
logging.basicConfig(stream=sys.stdout, level=logging.DEBUG, force=True)
args = parser.parse_args()
ConvertSubcommand.func(args)
|