CoolFace
Apppublic

tidalove/yolox

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
export_torchscript.py81 linesDownload Raw Back to tools
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