duan8/ufld/gen_wts.py
tsieyy 30de4562bc
Adapt ufld model to tensorrt8 (#1288)
* make ufld adapt tensorrt8

* make ufld adapt tensorrt8 and tensorrt7
2023-04-19 12:45:43 +08:00

22 lines
613 B
Python

import torch
import struct
#import models.crnn as crnn
from model.model import parsingNet
# Initialize
model = parsingNet(pretrained = False, backbone='18', cls_dim = (101, 56, 4), use_aux=False)
device = 'cpu'
# Load model
state_dict = torch.load('tusimple_18.pth', map_location='cpu')['model']
model.to(device).eval()
f = open('lane.wts', 'w')
f.write('{}\n'.format(len(state_dict.keys())))
for k, v in state_dict.items():
vr = v.reshape(-1).cpu().numpy()
f.write('{} {} '.format(k, len(vr)))
for vv in vr:
f.write(' ')
f.write(struct.pack('>f',float(vv)).hex())
f.write('\n')