8.2 KiB
8.2 KiB
In [ ]:
import torch
weights_file = 'BiRefNet-general-bb_swin_v1_tiny-epoch_232.pth' # https://github.com/ZhengPeng7/BiRefNet/releases/download/v1/BiRefNet-general-bb_swin_v1_tiny-epoch_232.pth
device = 'cuda' if torch.cuda.is_available() else 'cpu'In [ ]:
with open('config.py') as fp:
file_lines = fp.read()
if 'swin_v1_tiny' in weights_file:
print('Set `swin_v1_tiny` as the backbone.')
file_lines = file_lines.replace(
'''
'pvt_v2_b2', 'pvt_v2_b5', # 9-bs10, 10-bs5
][6]
''',
'''
'pvt_v2_b2', 'pvt_v2_b5', # 9-bs10, 10-bs5
][3]
''',
)
with open('config.py', mode="w") as fp:
fp.write(file_lines)
else:
file_lines = file_lines.replace(
'''
'pvt_v2_b2', 'pvt_v2_b5', # 9-bs10, 10-bs5
][3]
''',
'''
'pvt_v2_b2', 'pvt_v2_b5', # 9-bs10, 10-bs5
][6]
''',
)
with open('config.py', mode="w") as fp:
fp.write(file_lines)In [ ]:
from utils import check_state_dict
from models.birefnet import BiRefNet
birefnet = BiRefNet(bb_pretrained=False)
state_dict = torch.load('./{}'.format(weights_file), map_location=device)
state_dict = check_state_dict(state_dict)
birefnet.load_state_dict(state_dict)
torch.set_float32_matmul_precision(['high', 'highest'][0])
birefnet.to(device)
_ = birefnet.eval()In [ ]:
from torchvision.ops.deform_conv import DeformConv2d
import deform_conv2d_onnx_exporter
# register deform_conv2d operator
deform_conv2d_onnx_exporter.register_deform_conv2d_onnx_op()
def convert_to_onnx(net, file_name='output.onnx', input_shape=(1024, 1024), device=device):
input = torch.randn(1, 3, input_shape[0], input_shape[1]).to(device)
input_layer_names = ['input_image']
output_layer_names = ['output_image']
torch.onnx.export(
net,
input,
file_name,
verbose=False,
opset_version=17,
input_names=input_layer_names,
output_names=output_layer_names,
)
convert_to_onnx(birefnet, weights_file.replace('.pth', '.onnx'), input_shape=(1024, 1024), device=device)In [ ]:
from PIL import Image
from torchvision import transforms
transform_image = transforms.Compose([
transforms.Resize((1024, 1024)),
transforms.ToTensor(),
transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225])
])
imagepath = './Helicopter-HR.jpg'
image = Image.open(imagepath)
input_images = transform_image(image).unsqueeze(0).to(device)
input_images_numpy = input_images.cpu().numpy()In [ ]:
import onnxruntime
import matplotlib.pyplot as plt
providers = ['CPUExecutionProvider'] if device == 'cpu' else ['CUDAExecutionProvider']
onnx_session = onnxruntime.InferenceSession(
weights_file.replace('.pth', '.onnx'),
providers=providers
)
input_name = onnx_session.get_inputs()[0].name
print(onnxruntime.get_device(), onnx_session.get_providers())In [ ]:
from time import time
import matplotlib.pyplot as plt
time_st = time()
pred_onnx = torch.tensor(
onnx_session.run(None, {input_name: input_images_numpy if device == 'cpu' else input_images_numpy})[-1]
).squeeze(0).sigmoid().cpu()
print(time() - time_st)
plt.imshow(pred_onnx.squeeze(), cmap='gray'); plt.show()In [ ]:
with torch.no_grad():
preds = birefnet(input_images)[-1].sigmoid().cpu()
plt.imshow(preds.squeeze(), cmap='gray'); plt.show()In [ ]:
diff = abs(preds - pred_onnx)
print('sum(diff):', diff.sum())
plt.imshow((diff).squeeze(), cmap='gray'); plt.show()In [ ]:
%%timeit
with torch.no_grad():
preds = birefnet(input_images)[-1].sigmoid().cpu()In [ ]:
%%timeit
pred_onnx = torch.tensor(
onnx_session.run(None, {input_name: input_images_numpy})[-1]
).squeeze(0).sigmoid().cpu()In [ ]: