tidalove/yolox
0
1#!/usr/bin/env python32# -*- coding:utf-8 -*-3# Copyright (c) Megvii, Inc. and its affiliates.4 5import argparse6import os7from loguru import logger8 9import torch10 11from yolox.exp import get_exp12 13 14def make_parser():15 parser = argparse.ArgumentParser("YOLOX torchscript deploy")16 parser.add_argument(17 "--output-name", type=str, default="yolox.torchscript.pt", help="output name of models"18 )19 parser.add_argument("--batch-size", type=int, default=1, help="batch size")20 parser.add_argument(21 "-f",22 "--exp_file",23 default=None,24 type=str,25 help="experiment description file",26 )27 parser.add_argument("-expn", "--experiment-name", type=str, default=None)28 parser.add_argument("-n", "--name", type=str, default=None, help="model name")29 parser.add_argument("-c", "--ckpt", default=None, type=str, help="ckpt path")30 parser.add_argument(31 "--decode_in_inference",32 action="store_true",33 help="decode in inference or not"34 )35 parser.add_argument(36 "opts",37 help="Modify config options using the command-line",38 default=None,39 nargs=argparse.REMAINDER,40 )41 42 return parser43 44 45@logger.catch46def main():47 args = make_parser().parse_args()48 logger.info("args value: {}".format(args))49 exp = get_exp(args.exp_file, args.name)50 exp.merge(args.opts)51 52 if not args.experiment_name:53 args.experiment_name = exp.exp_name54 55 model = exp.get_model()56 if args.ckpt is None:57 file_name = os.path.join(exp.output_dir, args.experiment_name)58 ckpt_file = os.path.join(file_name, "best_ckpt.pth")59 else:60 ckpt_file = args.ckpt61 62 # load the model state dict63 ckpt = torch.load(ckpt_file, map_location="cpu")64 65 model.eval()66 if "model" in ckpt:67 ckpt = ckpt["model"]68 model.load_state_dict(ckpt)69 model.head.decode_in_inference = args.decode_in_inference70 71 logger.info("loading checkpoint done.")72 dummy_input = torch.randn(args.batch_size, 3, exp.test_size[0], exp.test_size[1])73 74 mod = torch.jit.trace(model, dummy_input)75 mod.save(args.output_name)76 logger.info("generated torchscript model named {}".format(args.output_name))77 78 79if __name__ == "__main__":80 main()81 