mirror of
https://github.com/sui-feng-cb/AzurLaneAutoScript1.git
synced 2026-08-06 01:22:59 +08:00
Opt: use ONNX model converted from MXNet for faster OCR
This commit is contained in:
BIN
bin/cnocr_models/azur_lane/model.onnx
Normal file
BIN
bin/cnocr_models/azur_lane/model.onnx
Normal file
Binary file not shown.
BIN
bin/cnocr_models/azur_lane_jp/model.onnx
Normal file
BIN
bin/cnocr_models/azur_lane_jp/model.onnx
Normal file
Binary file not shown.
BIN
bin/cnocr_models/cnocr/model.onnx
Normal file
BIN
bin/cnocr_models/cnocr/model.onnx
Normal file
Binary file not shown.
BIN
bin/cnocr_models/jp/model.onnx
Normal file
BIN
bin/cnocr_models/jp/model.onnx
Normal file
Binary file not shown.
BIN
bin/cnocr_models/tw/model.onnx
Normal file
BIN
bin/cnocr_models/tw/model.onnx
Normal file
Binary file not shown.
@@ -79,6 +79,10 @@ Deploy:
|
|||||||
# Address of ocr server for alas instance to connect
|
# Address of ocr server for alas instance to connect
|
||||||
# [Default] 127.0.0.1:22268
|
# [Default] 127.0.0.1:22268
|
||||||
OcrClientAddress: 127.0.0.1:22268
|
OcrClientAddress: 127.0.0.1:22268
|
||||||
|
# Use ONNX Runtime for OCR inference instead of MXNet
|
||||||
|
# Requires onnxruntime package and model.onnx in each model directory
|
||||||
|
# [Default] false
|
||||||
|
UseOcrOnnx: false
|
||||||
|
|
||||||
Update:
|
Update:
|
||||||
# Use auto update and builtin updater feature
|
# Use auto update and builtin updater feature
|
||||||
|
|||||||
@@ -35,6 +35,7 @@ class ConfigModel:
|
|||||||
StartOcrServer: bool = False
|
StartOcrServer: bool = False
|
||||||
OcrServerPort: int = 22268
|
OcrServerPort: int = 22268
|
||||||
OcrClientAddress: str = "127.0.0.1:22268"
|
OcrClientAddress: str = "127.0.0.1:22268"
|
||||||
|
UseOcrOnnx: bool = False
|
||||||
|
|
||||||
# Update
|
# Update
|
||||||
EnableReload: bool = True
|
EnableReload: bool = True
|
||||||
|
|||||||
607
dev_tools/convert_mxnet_to_onnx.py
Normal file
607
dev_tools/convert_mxnet_to_onnx.py
Normal file
@@ -0,0 +1,607 @@
|
|||||||
|
"""
|
||||||
|
Convert MXNet cnocr checkpoints to ONNX format.
|
||||||
|
|
||||||
|
Builds equivalent ONNX graphs directly from the MXNet symbol JSON and
|
||||||
|
parameter files. MXNet's built-in ONNX exporter cannot handle the
|
||||||
|
hybridized bidirectional GRU (internal ops _rnn_param_concat + RNN),
|
||||||
|
so the graph is reconstructed node-by-node.
|
||||||
|
|
||||||
|
Output is written alongside the source checkpoint as model.onnx.
|
||||||
|
|
||||||
|
The MODELS list at the bottom of this file defines which models to
|
||||||
|
convert. Each entry specifies the model name, directory, checkpoint
|
||||||
|
prefix, and epoch. Edit this list to add or remove models.
|
||||||
|
|
||||||
|
Notes:
|
||||||
|
- MXNet GRU gate order is [r, z, h]; ONNX expects [z, r, h].
|
||||||
|
- MXNet Reshape shape values -2, -3 (relative dimensions) are rewritten
|
||||||
|
because ONNX only supports -1.
|
||||||
|
"""
|
||||||
|
import os, sys, json
|
||||||
|
sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||||
|
|
||||||
|
import numpy as np
|
||||||
|
import onnx
|
||||||
|
from onnx import helper, TensorProto, numpy_helper
|
||||||
|
|
||||||
|
import mxnet as mx
|
||||||
|
|
||||||
|
OPSET = 13
|
||||||
|
FLOAT = TensorProto.FLOAT
|
||||||
|
INT64 = TensorProto.INT64
|
||||||
|
|
||||||
|
|
||||||
|
def parse_tuple(s):
|
||||||
|
"""Parse '(1, 2)' or '()' string into tuple of ints."""
|
||||||
|
s = s.strip().strip('()')
|
||||||
|
if not s:
|
||||||
|
return ()
|
||||||
|
return tuple(int(x.strip()) for x in s.split(','))
|
||||||
|
|
||||||
|
|
||||||
|
def ndarr(arr):
|
||||||
|
"""numpy array -> ONNX initializer."""
|
||||||
|
return numpy_helper.from_array(arr.astype(np.float32))
|
||||||
|
|
||||||
|
|
||||||
|
def shape_init(name, vals):
|
||||||
|
"""Create an INT64 initializer for shape/axes."""
|
||||||
|
tensor = numpy_helper.from_array(np.array(vals, dtype=np.int64), name=name)
|
||||||
|
return tensor
|
||||||
|
|
||||||
|
|
||||||
|
class ONNXBuilder:
|
||||||
|
def __init__(self, model_name):
|
||||||
|
self.model_name = model_name
|
||||||
|
self.nodes = []
|
||||||
|
self.initializers = []
|
||||||
|
self.name_counter = [0]
|
||||||
|
# Maps MXNet node name -> ONNX tensor name
|
||||||
|
self.name_map = {}
|
||||||
|
self.eps = 1e-5
|
||||||
|
|
||||||
|
def new_name(self, base):
|
||||||
|
self.name_counter[0] += 1
|
||||||
|
return f'{base}_{self.name_counter[0]}'
|
||||||
|
|
||||||
|
def resolve(self, mxnet_ref):
|
||||||
|
"""
|
||||||
|
Args:
|
||||||
|
mxnet_ref is either a node index (int) or a node name (str).
|
||||||
|
|
||||||
|
Return ONNX tensor name for an MXNet output, or None.
|
||||||
|
"""
|
||||||
|
if isinstance(mxnet_ref, int):
|
||||||
|
# Look up node name by index
|
||||||
|
if 0 <= mxnet_ref < len(self.nodes_json):
|
||||||
|
mxnet_ref = self.nodes_json[mxnet_ref]['name']
|
||||||
|
else:
|
||||||
|
return None
|
||||||
|
return self.name_map.get(mxnet_ref)
|
||||||
|
|
||||||
|
def make_conv(self, node, arg_params):
|
||||||
|
attrs = node['attrs']
|
||||||
|
name = node['name']
|
||||||
|
mx_in = self.resolve(node['inputs'][0][0])
|
||||||
|
if mx_in is None:
|
||||||
|
return False
|
||||||
|
# Weight is always input[1] (a param node)
|
||||||
|
w_node_idx = node['inputs'][1][0]
|
||||||
|
w_name = self.nodes_json[w_node_idx]['name']
|
||||||
|
weight = arg_params[w_name]
|
||||||
|
w_tensor = self.new_name(w_name)
|
||||||
|
self.initializers.append(ndarr(weight.asnumpy()))
|
||||||
|
# Note: ndarr above loses name; rebuild with name
|
||||||
|
self.initializers[-1] = numpy_helper.from_array(weight.asnumpy(), name=w_tensor)
|
||||||
|
|
||||||
|
kernel = parse_tuple(attrs['kernel'])
|
||||||
|
stride = parse_tuple(attrs['stride'])
|
||||||
|
pad = parse_tuple(attrs['pad'])
|
||||||
|
dilate = parse_tuple(attrs['dilate'])
|
||||||
|
groups = int(attrs['num_group'])
|
||||||
|
|
||||||
|
out = self.new_name(name)
|
||||||
|
# ONNX pads: [top, left, bottom, right] = [ph, pw, ph, pw]
|
||||||
|
pads = [pad[0], pad[1], pad[0], pad[1]] if len(pad) == 2 else list(pad) * 2
|
||||||
|
conv = helper.make_node(
|
||||||
|
'Conv', [mx_in, w_tensor], [out],
|
||||||
|
name=name,
|
||||||
|
kernel_shape=list(kernel),
|
||||||
|
strides=list(stride),
|
||||||
|
pads=pads,
|
||||||
|
dilations=list(dilate),
|
||||||
|
group=groups,
|
||||||
|
)
|
||||||
|
self.nodes.append(conv)
|
||||||
|
self.name_map[name] = out
|
||||||
|
return True
|
||||||
|
|
||||||
|
def make_batchnorm(self, node, arg_params, aux_params):
|
||||||
|
attrs = node['attrs']
|
||||||
|
name = node['name']
|
||||||
|
mx_in = self.resolve(node['inputs'][0][0])
|
||||||
|
if mx_in is None:
|
||||||
|
return False
|
||||||
|
eps = float(attrs['eps'])
|
||||||
|
|
||||||
|
gamma_name = self.nodes_json[node['inputs'][1][0]]['name']
|
||||||
|
beta_name = self.nodes_json[node['inputs'][2][0]]['name']
|
||||||
|
mean_name = self.nodes_json[node['inputs'][3][0]]['name']
|
||||||
|
var_name = self.nodes_json[node['inputs'][4][0]]['name']
|
||||||
|
|
||||||
|
gamma_t = self.new_name(gamma_name)
|
||||||
|
beta_t = self.new_name(beta_name)
|
||||||
|
mean_t = self.new_name(mean_name)
|
||||||
|
var_t = self.new_name(var_name)
|
||||||
|
for t, key in [(gamma_t, gamma_name), (beta_t, beta_name),
|
||||||
|
(mean_t, mean_name), (var_t, var_name)]:
|
||||||
|
src = arg_params.get(key)
|
||||||
|
if src is None:
|
||||||
|
src = aux_params.get(key)
|
||||||
|
self.initializers.append(numpy_helper.from_array(src.asnumpy().astype(np.float32), name=t))
|
||||||
|
|
||||||
|
out = self.new_name(name)
|
||||||
|
bn = helper.make_node(
|
||||||
|
'BatchNormalization',
|
||||||
|
[mx_in, gamma_t, beta_t, mean_t, var_t], [out],
|
||||||
|
name=name, epsilon=eps,
|
||||||
|
)
|
||||||
|
self.nodes.append(bn)
|
||||||
|
self.name_map[name] = out
|
||||||
|
return True
|
||||||
|
|
||||||
|
def make_relu(self, node):
|
||||||
|
name = node['name']
|
||||||
|
mx_in = self.resolve(node['inputs'][0][0])
|
||||||
|
if mx_in is None:
|
||||||
|
return False
|
||||||
|
out = self.new_name(name)
|
||||||
|
relu = helper.make_node('Relu', [mx_in], [out], name=name)
|
||||||
|
self.nodes.append(relu)
|
||||||
|
self.name_map[name] = out
|
||||||
|
return True
|
||||||
|
|
||||||
|
def make_concat(self, node):
|
||||||
|
name = node['name']
|
||||||
|
axis = int(node['attrs']['dim'])
|
||||||
|
mx_inputs = [self.resolve(i[0]) for i in node['inputs']]
|
||||||
|
if any(x is None for x in mx_inputs):
|
||||||
|
return False
|
||||||
|
out = self.new_name(name)
|
||||||
|
concat = helper.make_node('Concat', mx_inputs, [out], name=name, axis=axis)
|
||||||
|
self.nodes.append(concat)
|
||||||
|
self.name_map[name] = out
|
||||||
|
return True
|
||||||
|
|
||||||
|
def make_pooling(self, node):
|
||||||
|
attrs = node['attrs']
|
||||||
|
name = node['name']
|
||||||
|
mx_in = self.resolve(node['inputs'][0][0])
|
||||||
|
if mx_in is None:
|
||||||
|
return False
|
||||||
|
pool_type = attrs['pool_type']
|
||||||
|
kernel = parse_tuple(attrs['kernel'])
|
||||||
|
stride = parse_tuple(attrs['stride'])
|
||||||
|
pad = parse_tuple(attrs['pad'])
|
||||||
|
|
||||||
|
out = self.new_name(name)
|
||||||
|
pads = [pad[0], pad[1], pad[0], pad[1]] if len(pad) == 2 else list(pad) * 2
|
||||||
|
if pool_type == 'max':
|
||||||
|
pool = helper.make_node(
|
||||||
|
'MaxPool', [mx_in], [out], name=name,
|
||||||
|
kernel_shape=list(kernel), strides=list(stride), pads=pads,
|
||||||
|
)
|
||||||
|
elif pool_type == 'avg':
|
||||||
|
pool = helper.make_node(
|
||||||
|
'AveragePool', [mx_in], [out], name=name,
|
||||||
|
kernel_shape=list(kernel), strides=list(stride), pads=pads,
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
raise ValueError(f'Unknown pool_type: {pool_type}')
|
||||||
|
self.nodes.append(pool)
|
||||||
|
self.name_map[name] = out
|
||||||
|
return True
|
||||||
|
|
||||||
|
def make_reshape(self, node):
|
||||||
|
name = node['name']
|
||||||
|
shape_str = node['attrs']['shape']
|
||||||
|
shape = parse_tuple(shape_str)
|
||||||
|
mx_in = self.resolve(node['inputs'][0][0])
|
||||||
|
if mx_in is None:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Convert MXNet shape notation to ONNX-compatible
|
||||||
|
# -2, -3, -4 (relative dims) are not supported by ONNX; handled by caller
|
||||||
|
onnx_shape = []
|
||||||
|
for s in shape:
|
||||||
|
if s >= 0 or s == -1:
|
||||||
|
onnx_shape.append(s)
|
||||||
|
else:
|
||||||
|
# Unsupported relative dim — caller should have rewritten these
|
||||||
|
raise ValueError(f'Unsupported shape value {s} in {name}')
|
||||||
|
|
||||||
|
out = self.new_name(name)
|
||||||
|
shape_t = self.new_name(f'{name}_shape')
|
||||||
|
self.initializers.append(shape_init(shape_t, onnx_shape))
|
||||||
|
reshape = helper.make_node('Reshape', [mx_in, shape_t], [out], name=name)
|
||||||
|
self.nodes.append(reshape)
|
||||||
|
self.name_map[name] = out
|
||||||
|
return True
|
||||||
|
|
||||||
|
def make_unsqueeze(self, node):
|
||||||
|
name = node['name']
|
||||||
|
axis = int(node['attrs']['axis'])
|
||||||
|
mx_in = self.resolve(node['inputs'][0][0])
|
||||||
|
if mx_in is None:
|
||||||
|
return False
|
||||||
|
out = self.new_name(name)
|
||||||
|
axes_t = self.new_name(f'{name}_axes')
|
||||||
|
self.initializers.append(shape_init(axes_t, [axis]))
|
||||||
|
unsq = helper.make_node('Unsqueeze', [mx_in, axes_t], [out], name=name)
|
||||||
|
self.nodes.append(unsq)
|
||||||
|
self.name_map[name] = out
|
||||||
|
return True
|
||||||
|
|
||||||
|
def make_squeeze(self, node):
|
||||||
|
name = node['name']
|
||||||
|
axis = int(node['attrs']['axis'])
|
||||||
|
mx_in = self.resolve(node['inputs'][0][0])
|
||||||
|
if mx_in is None:
|
||||||
|
return False
|
||||||
|
out = self.new_name(name)
|
||||||
|
axes_t = self.new_name(f'{name}_axes')
|
||||||
|
self.initializers.append(shape_init(axes_t, [axis]))
|
||||||
|
sq = helper.make_node('Squeeze', [mx_in, axes_t], [out], name=name)
|
||||||
|
self.nodes.append(sq)
|
||||||
|
self.name_map[name] = out
|
||||||
|
return True
|
||||||
|
|
||||||
|
def make_transpose(self, node):
|
||||||
|
name = node['name']
|
||||||
|
perm = parse_tuple(node['attrs']['axes'])
|
||||||
|
mx_in = self.resolve(node['inputs'][0][0])
|
||||||
|
if mx_in is None:
|
||||||
|
return False
|
||||||
|
out = self.new_name(name)
|
||||||
|
tr = helper.make_node('Transpose', [mx_in], [out], name=name, perm=list(perm))
|
||||||
|
self.nodes.append(tr)
|
||||||
|
self.name_map[name] = out
|
||||||
|
return True
|
||||||
|
|
||||||
|
def make_dropout(self, node):
|
||||||
|
# Inference: identity
|
||||||
|
name = node['name']
|
||||||
|
mx_in = self.resolve(node['inputs'][0][0])
|
||||||
|
if mx_in is None:
|
||||||
|
return False
|
||||||
|
self.name_map[name] = mx_in
|
||||||
|
return True
|
||||||
|
|
||||||
|
def make_gemm(self, node, arg_params):
|
||||||
|
name = node['name']
|
||||||
|
mx_in = self.resolve(node['inputs'][0][0])
|
||||||
|
if mx_in is None:
|
||||||
|
return False
|
||||||
|
w_node_idx = node['inputs'][1][0]
|
||||||
|
w_name = self.nodes_json[w_node_idx]['name']
|
||||||
|
b_node_idx = node['inputs'][2][0]
|
||||||
|
b_name = self.nodes_json[b_node_idx]['name']
|
||||||
|
|
||||||
|
w_t = self.new_name(w_name)
|
||||||
|
b_t = self.new_name(b_name)
|
||||||
|
self.initializers.append(numpy_helper.from_array(arg_params[w_name].asnumpy().astype(np.float32), name=w_t))
|
||||||
|
self.initializers.append(numpy_helper.from_array(arg_params[b_name].asnumpy().astype(np.float32), name=b_t))
|
||||||
|
|
||||||
|
out = self.new_name(name)
|
||||||
|
# Y = X @ W^T + b (transB=1)
|
||||||
|
gemm = helper.make_node('Gemm', [mx_in, w_t, b_t], [out], name=name, transB=1, alpha=1.0, beta=1.0)
|
||||||
|
self.nodes.append(gemm)
|
||||||
|
self.name_map[name] = out
|
||||||
|
return True
|
||||||
|
|
||||||
|
def build(self, sym_dict, arg_params, aux_params, seq_len, hidden_dim, img_width):
|
||||||
|
self.nodes_json = sym_dict['nodes']
|
||||||
|
node_map = {node['name']: node for node in self.nodes_json}
|
||||||
|
|
||||||
|
# Find input node (data)
|
||||||
|
data_node = self.nodes_json[0]
|
||||||
|
# Dynamic width: batch and width are free dimensions
|
||||||
|
graph_input = helper.make_tensor_value_info('data', FLOAT, [None, 1, 32, None])
|
||||||
|
self.name_map['data'] = 'data'
|
||||||
|
|
||||||
|
# Walk nodes in order, handling each op
|
||||||
|
for idx, node in enumerate(self.nodes_json):
|
||||||
|
op = node.get('op', 'null')
|
||||||
|
name = node.get('name', '')
|
||||||
|
|
||||||
|
if op == 'null':
|
||||||
|
# Parameter node (weight/data) — no-op
|
||||||
|
continue
|
||||||
|
if op == '_zeros':
|
||||||
|
# Initial hidden state for RNN — handled in GRU section
|
||||||
|
continue
|
||||||
|
# Skip MXNet nodes that we replace with a dynamic-width-friendly
|
||||||
|
# equivalent. The MXNet reshape(0,-3,0) + expand_dims + squeeze +
|
||||||
|
# transpose chain has different 0-index semantics than ONNX and
|
||||||
|
# doesn't support dynamic width.
|
||||||
|
if name in ('densenet0_reshape0', 'densenet0_expand_dims0',
|
||||||
|
'dropout0_fwd', 'crnn0_squeeze0', 'crnn0_transpose0'):
|
||||||
|
if name == 'densenet0_reshape0':
|
||||||
|
self.make_post_densenet(node)
|
||||||
|
continue
|
||||||
|
if op == 'Reshape':
|
||||||
|
# For GRU param reshapes (gru0_reshape*), skip — they're for
|
||||||
|
# MXNet's internal RNN param packing which we bypass.
|
||||||
|
if name.startswith('gru0_reshape'):
|
||||||
|
continue
|
||||||
|
# For reshape0 (the final reshape before FC): (-3, -2) -> (-1, hidden_dim)
|
||||||
|
if name == 'reshape0':
|
||||||
|
mx_in = self.resolve(node['inputs'][0][0])
|
||||||
|
out = self.new_name(name)
|
||||||
|
shape_t = self.new_name(f'{name}_shape')
|
||||||
|
self.initializers.append(shape_init(shape_t, [-1, hidden_dim]))
|
||||||
|
self.nodes.append(helper.make_node('Reshape', [mx_in, shape_t], [out], name=name))
|
||||||
|
self.name_map[name] = out
|
||||||
|
continue
|
||||||
|
self.make_reshape(node)
|
||||||
|
elif op == 'Convolution':
|
||||||
|
self.make_conv(node, arg_params)
|
||||||
|
elif op == 'BatchNorm':
|
||||||
|
self.make_batchnorm(node, arg_params, aux_params)
|
||||||
|
elif op == 'Activation':
|
||||||
|
self.make_relu(node)
|
||||||
|
elif op == 'Concat':
|
||||||
|
self.make_concat(node)
|
||||||
|
elif op == 'Pooling':
|
||||||
|
self.make_pooling(node)
|
||||||
|
elif op == 'expand_dims':
|
||||||
|
self.make_unsqueeze(node)
|
||||||
|
elif op == 'squeeze':
|
||||||
|
self.make_squeeze(node)
|
||||||
|
elif op == 'transpose':
|
||||||
|
self.make_transpose(node)
|
||||||
|
elif op == 'Dropout':
|
||||||
|
self.make_dropout(node)
|
||||||
|
elif op == 'FullyConnected':
|
||||||
|
self.make_gemm(node, arg_params)
|
||||||
|
elif op == '_rnn_param_concat':
|
||||||
|
# Bypassed — GRU built from raw params in make_gru()
|
||||||
|
continue
|
||||||
|
elif op == 'RNN':
|
||||||
|
self.make_gru(node, arg_params)
|
||||||
|
else:
|
||||||
|
print(f' WARNING: unhandled op {op} ({name})')
|
||||||
|
|
||||||
|
# Find the pred_fc output tensor name
|
||||||
|
pred_fc_node = node_map['pred_fc']
|
||||||
|
pred_fc_out = self.name_map.get('pred_fc')
|
||||||
|
|
||||||
|
# Add softmax
|
||||||
|
softmax_out = 'softmax_output'
|
||||||
|
sm = helper.make_node('Softmax', [pred_fc_out], [softmax_out], name='softmax', axis=-1)
|
||||||
|
self.nodes.append(sm)
|
||||||
|
|
||||||
|
graph_output = helper.make_tensor_value_info(softmax_out, FLOAT, [None, None])
|
||||||
|
|
||||||
|
graph = helper.make_graph(
|
||||||
|
self.nodes,
|
||||||
|
f'{self.model_name}_graph',
|
||||||
|
[graph_input],
|
||||||
|
[graph_output],
|
||||||
|
initializer=self.initializers,
|
||||||
|
)
|
||||||
|
model = helper.make_model(graph, opset_imports=[helper.make_opsetid('', OPSET)])
|
||||||
|
model.ir_version = 7
|
||||||
|
return model
|
||||||
|
|
||||||
|
def make_post_densenet(self, node):
|
||||||
|
"""
|
||||||
|
Replace MXNet's reshape+expand_dims+dropout+squeeze+transpose chain
|
||||||
|
with a dynamic-width-friendly equivalent.
|
||||||
|
|
||||||
|
MXNet DenseNet output: (batch, 256, 2, seq) where seq = W/4
|
||||||
|
Target GRU input: (seq, batch, 512)
|
||||||
|
|
||||||
|
Sequence:
|
||||||
|
1. transpose (0, 3, 1, 2): (batch, seq, 256, 2)
|
||||||
|
2. reshape (0, 0, -1): (batch, seq, 512) [flatten 256*2]
|
||||||
|
3. transpose (1, 0, 2): (seq, batch, 512)
|
||||||
|
"""
|
||||||
|
name = node['name']
|
||||||
|
mx_in = self.resolve(node['inputs'][0][0]) # DenseNet pool output
|
||||||
|
if mx_in is None:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# 1. transpose (batch, 256, 2, seq) -> (batch, seq, 256, 2)
|
||||||
|
t1 = self.new_name(f'{name}_t1')
|
||||||
|
self.nodes.append(helper.make_node('Transpose', [mx_in], [t1],
|
||||||
|
name=f'{name}_t1', perm=[0, 3, 1, 2]))
|
||||||
|
|
||||||
|
# 2. reshape (batch, seq, 256*2=512)
|
||||||
|
t2 = self.new_name(f'{name}_t2')
|
||||||
|
shape_t = self.new_name(f'{name}_shape')
|
||||||
|
self.initializers.append(shape_init(shape_t, [0, 0, -1]))
|
||||||
|
self.nodes.append(helper.make_node('Reshape', [t1, shape_t], [t2],
|
||||||
|
name=f'{name}_t2'))
|
||||||
|
|
||||||
|
# 3. transpose (seq, batch, 512) — the GRU input layout
|
||||||
|
t3 = self.new_name(f'{name}_t3')
|
||||||
|
self.nodes.append(helper.make_node('Transpose', [t2], [t3],
|
||||||
|
name=f'{name}_t3', perm=[1, 0, 2]))
|
||||||
|
|
||||||
|
# Store as crnn0_transpose0 output so the RNN node can consume it
|
||||||
|
self.name_map['crnn0_transpose0'] = t3
|
||||||
|
return True
|
||||||
|
|
||||||
|
def make_gru(self, node, arg_params):
|
||||||
|
"""
|
||||||
|
Build ONNX bidirectional GRU from raw forward/backward params.
|
||||||
|
|
||||||
|
MXNet RNN op consumes (data, params_concat, initial_h).
|
||||||
|
Bypass the internal param packing and build ONNX GRU directly:
|
||||||
|
X: (seq_len, batch, 512)
|
||||||
|
W: (2, 384, 512) [fwd_i2h, bwd_i2h]
|
||||||
|
R: (2, 384, 128) [fwd_h2h, bwd_h2h]
|
||||||
|
B: (2, 768) [fwd[i2h_bias,h2h_bias], bwd[...]]
|
||||||
|
direction=bidirectional, linear_before_reset=0
|
||||||
|
"""
|
||||||
|
name = node['name']
|
||||||
|
mx_in = self.resolve(node['inputs'][0][0]) # transposed DenseNet output
|
||||||
|
if mx_in is None:
|
||||||
|
return False
|
||||||
|
|
||||||
|
# Derive GRU dims from actual params (architecture is densenet-lite-gru)
|
||||||
|
i2h0 = arg_params['gru0_l0_i2h_weight'].asnumpy()
|
||||||
|
hidden = i2h0.shape[0] // 3 # 3*hidden, hidden = 128
|
||||||
|
input_size = i2h0.shape[1]
|
||||||
|
# MXNet GRU gate order is [r, z, h]; ONNX GRU expects [z, r, h].
|
||||||
|
# Empirical finding: reorder (0<->1) + linear_before_reset=1 matches MXNet.
|
||||||
|
idx = np.array([1, 0, 2])
|
||||||
|
|
||||||
|
def reorder_gates(mat):
|
||||||
|
# mat shape (3*hidden, ...) or (3*hidden,)
|
||||||
|
if mat.ndim == 1:
|
||||||
|
return mat.reshape(3, hidden)[idx].reshape(-1)
|
||||||
|
else:
|
||||||
|
return mat.reshape(3, hidden, -1)[idx].reshape(-1, mat.shape[-1])
|
||||||
|
|
||||||
|
# Collect forward/backward GRU params
|
||||||
|
fwd_i2h = reorder_gates(arg_params['gru0_l0_i2h_weight'].asnumpy()) # (384, 512)
|
||||||
|
fwd_h2h = reorder_gates(arg_params['gru0_l0_h2h_weight'].asnumpy()) # (384, 128)
|
||||||
|
fwd_i2h_b = reorder_gates(arg_params['gru0_l0_i2h_bias'].asnumpy()) # (384,)
|
||||||
|
fwd_h2h_b = reorder_gates(arg_params['gru0_l0_h2h_bias'].asnumpy()) # (384,)
|
||||||
|
bwd_i2h = reorder_gates(arg_params['gru0_r0_i2h_weight'].asnumpy())
|
||||||
|
bwd_h2h = reorder_gates(arg_params['gru0_r0_h2h_weight'].asnumpy())
|
||||||
|
bwd_i2h_b = reorder_gates(arg_params['gru0_r0_i2h_bias'].asnumpy())
|
||||||
|
bwd_h2h_b = reorder_gates(arg_params['gru0_r0_h2h_bias'].asnumpy())
|
||||||
|
|
||||||
|
# ONNX GRU weights: W (2, 3h, input), R (2, 3h, h), B (2, 6h)
|
||||||
|
W = np.stack([fwd_i2h, bwd_i2h]).astype(np.float32) # (2, 384, 512)
|
||||||
|
R = np.stack([fwd_h2h, bwd_h2h]).astype(np.float32) # (2, 384, 128)
|
||||||
|
B = np.stack([np.concatenate([fwd_i2h_b, fwd_h2h_b]),
|
||||||
|
np.concatenate([bwd_i2h_b, bwd_h2h_b])]).astype(np.float32) # (2, 768)
|
||||||
|
|
||||||
|
w_t = self.new_name('gru_W')
|
||||||
|
r_t = self.new_name('gru_R')
|
||||||
|
b_t = self.new_name('gru_B')
|
||||||
|
self.initializers.append(numpy_helper.from_array(W, name=w_t))
|
||||||
|
self.initializers.append(numpy_helper.from_array(R, name=r_t))
|
||||||
|
self.initializers.append(numpy_helper.from_array(B, name=b_t))
|
||||||
|
|
||||||
|
# GRU output Y: (seq_len, num_directions, batch, hidden)
|
||||||
|
y_t = self.new_name('gru_Y')
|
||||||
|
gru = helper.make_node(
|
||||||
|
'GRU', [mx_in, w_t, r_t, b_t], [y_t],
|
||||||
|
name=name,
|
||||||
|
hidden_size=hidden,
|
||||||
|
direction='bidirectional',
|
||||||
|
linear_before_reset=1,
|
||||||
|
)
|
||||||
|
self.nodes.append(gru)
|
||||||
|
|
||||||
|
# Transform Y (seq_len, 2, batch, 128) -> (seq_len, batch, 2, 128) -> (seq_len, batch, 256)
|
||||||
|
y_tr_t = self.new_name('gru_Y_tr')
|
||||||
|
self.nodes.append(helper.make_node('Transpose', [y_t], [y_tr_t], name='gru_Y_tr', perm=[0, 2, 1, 3]))
|
||||||
|
y_rs_t = self.new_name('gru_Y_rs')
|
||||||
|
shape_t = self.new_name('gru_Y_shape')
|
||||||
|
self.initializers.append(shape_init(shape_t, [0, 0, -1]))
|
||||||
|
self.nodes.append(helper.make_node('Reshape', [y_tr_t, shape_t], [y_rs_t], name='gru_Y_reshape'))
|
||||||
|
|
||||||
|
# Store result for the final reshape0 (which expects (seq_len, batch, 256))
|
||||||
|
self.name_map['gru0_rnn0'] = y_rs_t
|
||||||
|
return True
|
||||||
|
|
||||||
|
|
||||||
|
def convert(model_info):
|
||||||
|
name = model_info['name']
|
||||||
|
model_dir = model_info['dir']
|
||||||
|
prefix_path = os.path.join(model_dir, model_info['prefix'])
|
||||||
|
epoch = model_info['epoch']
|
||||||
|
img_width = model_info['img_width']
|
||||||
|
|
||||||
|
seq_len = img_width // 4
|
||||||
|
hidden_dim = model_info['hidden'] * 2
|
||||||
|
|
||||||
|
print(f"\n{'='*60}")
|
||||||
|
print(f"Building ONNX model: {name}")
|
||||||
|
|
||||||
|
sym, arg_params, aux_params = mx.model.load_checkpoint(prefix_path, epoch)
|
||||||
|
print(f" Loaded MXNet checkpoint (epoch {epoch})")
|
||||||
|
|
||||||
|
pred_fc = sym.get_internals()['pred_fc_output']
|
||||||
|
sym_dict = json.loads(pred_fc.tojson())
|
||||||
|
|
||||||
|
builder = ONNXBuilder(name)
|
||||||
|
model = builder.build(sym_dict, arg_params, aux_params, seq_len, hidden_dim, img_width)
|
||||||
|
|
||||||
|
output_path = os.path.join(model_dir, 'model.onnx')
|
||||||
|
onnx.save(model, output_path)
|
||||||
|
size_kb = os.path.getsize(output_path) / 1024
|
||||||
|
print(f" SAVED: {output_path} ({size_kb:.1f} KB)")
|
||||||
|
|
||||||
|
# Verify
|
||||||
|
try:
|
||||||
|
onnx.checker.check_model(model)
|
||||||
|
print(f" ONNX structural check: PASSED")
|
||||||
|
except Exception as e:
|
||||||
|
print(f" ONNX structural check: FAILED - {e}")
|
||||||
|
|
||||||
|
# Print graph summary
|
||||||
|
ops = {}
|
||||||
|
for n in model.graph.node:
|
||||||
|
ops[n.op_type] = ops.get(n.op_type, 0) + 1
|
||||||
|
print(f" Ops ({len(ops)} types): {dict(sorted(ops.items()))}")
|
||||||
|
return output_path
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == '__main__':
|
||||||
|
print("ONNX Model Builder for ALAS cnocr Models")
|
||||||
|
print(f"MXNet: {mx.__version__}, onnx: {onnx.__version__}")
|
||||||
|
|
||||||
|
models = [
|
||||||
|
{
|
||||||
|
'name': 'azur_lane',
|
||||||
|
'dir': './bin/cnocr_models/azur_lane',
|
||||||
|
'prefix': 'cnocr-v1.2.0-densenet-lite-gru',
|
||||||
|
'epoch': 15,
|
||||||
|
'img_width': 280,
|
||||||
|
'hidden': 128,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
'name': 'azur_lane_jp',
|
||||||
|
'dir': './bin/cnocr_models/azur_lane_jp',
|
||||||
|
'prefix': 'cnocr-v1.2.0-densenet-lite-gru',
|
||||||
|
'epoch': 20,
|
||||||
|
'img_width': 280,
|
||||||
|
'hidden': 128,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
'name': 'cnocr',
|
||||||
|
'dir': './bin/cnocr_models/cnocr',
|
||||||
|
'prefix': 'cnocr-v1.2.0-densenet-lite-gru',
|
||||||
|
'epoch': 39,
|
||||||
|
'img_width': 280,
|
||||||
|
'hidden': 128,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
'name': 'jp',
|
||||||
|
'dir': './bin/cnocr_models/jp',
|
||||||
|
'prefix': 'cnocr-v1.2.0-densenet-lite-gru',
|
||||||
|
'epoch': 125,
|
||||||
|
'img_width': 280,
|
||||||
|
'hidden': 128,
|
||||||
|
},
|
||||||
|
{
|
||||||
|
'name': 'tw',
|
||||||
|
'dir': './bin/cnocr_models/tw',
|
||||||
|
'prefix': 'cnocr-v1.2.0-densenet-lite-gru',
|
||||||
|
'epoch': 63,
|
||||||
|
'img_width': 280,
|
||||||
|
'hidden': 128,
|
||||||
|
},
|
||||||
|
]
|
||||||
|
|
||||||
|
for m in models:
|
||||||
|
convert(m)
|
||||||
|
|
||||||
|
print(f"\nDone.")
|
||||||
@@ -176,7 +176,7 @@ class AlOcr(CnOcr):
|
|||||||
prefix = os.path.join(self._model_dir, self._model_file_prefix)
|
prefix = os.path.join(self._model_dir, self._model_file_prefix)
|
||||||
data_names = ['data']
|
data_names = ['data']
|
||||||
data_shapes = [(data_names[0], (hp.batch_size, 1, hp.img_height, hp.img_width))]
|
data_shapes = [(data_names[0], (hp.batch_size, 1, hp.img_height, hp.img_width))]
|
||||||
logger.info('Loading OCR model: %s' % self._model_dir) # Change log appearance.
|
logger.info('Loading OCR model (MXNET): %s' % self._model_dir) # Change log appearance.
|
||||||
mod = load_module(
|
mod = load_module(
|
||||||
prefix,
|
prefix,
|
||||||
self._model_epoch,
|
self._model_epoch,
|
||||||
@@ -235,3 +235,62 @@ class AlOcr(CnOcr):
|
|||||||
img_list, img_widths = self._pad_arrays(img_list)
|
img_list, img_widths = self._pad_arrays(img_list)
|
||||||
image = cv2.hconcat(img_list)[0, :, :]
|
image = cv2.hconcat(img_list)[0, :, :]
|
||||||
Image.fromarray(image).show()
|
Image.fromarray(image).show()
|
||||||
|
|
||||||
|
|
||||||
|
class _OnnxModule:
|
||||||
|
"""
|
||||||
|
Thin wrapper that presents the same `predict(sample) -> NDArray`
|
||||||
|
interface as an MXNet Module, so that CnOcr._predict works unchanged.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, session):
|
||||||
|
self._session = session
|
||||||
|
self._input_name = session.get_inputs()[0].name
|
||||||
|
|
||||||
|
def predict(self, sample):
|
||||||
|
import mxnet as mx
|
||||||
|
|
||||||
|
if isinstance(sample, mx.nd.NDArray):
|
||||||
|
sample = sample.asnumpy()
|
||||||
|
# sample: (batch, 1, 32, width) -> output: (seq_len * batch, num_classes)
|
||||||
|
out = self._session.run(None, {self._input_name: np.asarray(sample, dtype=np.float32)})[0]
|
||||||
|
return mx.nd.array(out)
|
||||||
|
|
||||||
|
|
||||||
|
class AlOcrOnnx(AlOcr):
|
||||||
|
"""
|
||||||
|
Subclass of AlOcr that runs inference through ONNX Runtime instead of MXNet,
|
||||||
|
providing faster CPU inference with numerically identical output.
|
||||||
|
All pre/post-processing is inherited unchanged.
|
||||||
|
|
||||||
|
Requires onnxruntime installed and model.onnx in the model directory
|
||||||
|
(generated by dev_tools/convert_mxnet_to_onnx.py from the MXNet checkpoint).
|
||||||
|
|
||||||
|
Parameters are identical to AlOcr.
|
||||||
|
|
||||||
|
Enable by UseOcrOnnx in deploy config.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def _get_module(self, context):
|
||||||
|
_network, self._hp = gen_network(self._model_name, self._hp, self._net_prefix)
|
||||||
|
onnx_path = os.path.join(self._model_dir, 'model.onnx')
|
||||||
|
|
||||||
|
if not os.path.exists(onnx_path):
|
||||||
|
logger.warning(f'ONNX model not found: {onnx_path}, '
|
||||||
|
'falling back to MXNet')
|
||||||
|
return AlOcr._get_module(self, context)
|
||||||
|
|
||||||
|
try:
|
||||||
|
import onnxruntime as ort
|
||||||
|
except ImportError:
|
||||||
|
logger.warning('onnxruntime not installed, falling back to MXNet')
|
||||||
|
return AlOcr._get_module(self, context)
|
||||||
|
|
||||||
|
logger.info(f'Loading OCR model (ONNX): {onnx_path}')
|
||||||
|
so = ort.SessionOptions()
|
||||||
|
# intra_op_num_threads=1 avoids thread-scheduling overhead on the small GRU model (hidden=128).
|
||||||
|
# Using all cores (default) is measurably slower
|
||||||
|
# because the per-op parallelisation gain is outweighed by thread-pool synchronisation cost.
|
||||||
|
so.intra_op_num_threads = 1
|
||||||
|
session = ort.InferenceSession(onnx_path, so, providers=['CPUExecutionProvider'])
|
||||||
|
return _OnnxModule(session)
|
||||||
|
|||||||
@@ -1,65 +1,76 @@
|
|||||||
from module.base.decorator import cached_property
|
from module.base.decorator import cached_property
|
||||||
|
from module.webui.setting import State
|
||||||
|
|
||||||
|
|
||||||
|
def _ocr_class():
|
||||||
|
if not State.deploy_config.UseOcrOnnx:
|
||||||
|
from module.ocr.al_ocr import AlOcr
|
||||||
|
return AlOcr
|
||||||
|
from module.ocr.al_ocr import AlOcrOnnx
|
||||||
|
return AlOcrOnnx
|
||||||
|
|
||||||
|
|
||||||
class OcrModel:
|
class OcrModel:
|
||||||
@cached_property
|
@cached_property
|
||||||
def azur_lane(self):
|
def azur_lane(self):
|
||||||
# Folder: ./bin/cnocr_models/azur_lane
|
# Folder: ./bin/cnocr_models/azur_lane
|
||||||
# Size: 3.25MB
|
# Size: 3.25MB (MXNet) / 3.34MB (ONNX)
|
||||||
# Model: densenet-lite-gru
|
# Model: densenet-lite-gru
|
||||||
# Epoch: 15
|
# Epoch: 15
|
||||||
# Validation accuracy: 99.43%
|
# Validation accuracy: 99.43%
|
||||||
# Font: Impact, AgencyFB-Regular, MStiffHeiHK-UltraBold
|
# Font: Impact, AgencyFB-Regular, MStiffHeiHK-UltraBold
|
||||||
# Charset: 0123456789ABCDEFGHIJKLMNPQRSTUVWXYZ:/- (Letter 'O' and <space> is not included)
|
# Charset: 0123456789ABCDEFGHIJKLMNPQRSTUVWXYZ:/- (Letter 'O' and <space> is not included)
|
||||||
# _num_classes: 39
|
# _num_classes: 39
|
||||||
from module.ocr.al_ocr import AlOcr
|
return _ocr_class()(model_name='densenet-lite-gru', model_epoch=15,
|
||||||
return AlOcr(model_name='densenet-lite-gru', model_epoch=15, root='./bin/cnocr_models/azur_lane',
|
root='./bin/cnocr_models/azur_lane', name='azur_lane')
|
||||||
name='azur_lane')
|
|
||||||
|
|
||||||
@cached_property
|
@cached_property
|
||||||
def azur_lane_jp(self):
|
def azur_lane_jp(self):
|
||||||
# Folder: ./bin/cnocr_models/azur_lane_jp
|
# Folder: ./bin/cnocr_models/azur_lane_jp
|
||||||
# Size: 3.25MB
|
# Size: 3.25MB (MXNet) / 3.34MB (ONNX)
|
||||||
# Model: densenet-lite-gru
|
# Model: densenet-lite-gru
|
||||||
# Epoch: 20
|
# Epoch: 20
|
||||||
# Validation accuracy: 99.01%
|
# Validation accuracy: 99.01%
|
||||||
# Font: Impact, VibeMO Compressed Pro Thin, Folk R, Source Han Serif JP
|
# Font: Impact, VibeMO Compressed Pro Thin, Folk R, Source Han Serif JP
|
||||||
# Charset: 0123456789ABCDEFGHIJKLMNPQRSTUVWXYZ:/- (Letter 'O' and <space> is not included)
|
# Charset: 0123456789ABCDEFGHIJKLMNPQRSTUVWXYZ:/- (Letter 'O' and <space> is not included)
|
||||||
# _num_classes: 39
|
# _num_classes: 39
|
||||||
from module.ocr.al_ocr import AlOcr
|
return _ocr_class()(model_name='densenet-lite-gru', model_epoch=20,
|
||||||
return AlOcr(model_name='densenet-lite-gru', model_epoch=20, root='./bin/cnocr_models/azur_lane_jp',
|
root='./bin/cnocr_models/azur_lane_jp', name='azur_lane_jp')
|
||||||
name='azur_lane_jp')
|
|
||||||
|
|
||||||
@cached_property
|
@cached_property
|
||||||
def cnocr(self):
|
def cnocr(self):
|
||||||
# Folder: ./bin/cnocr_models/cnocr
|
# Folder: ./bin/cnocr_models/cnocr
|
||||||
# Size: 9.51MB
|
# Size: 9.51MB (MXNet) / 9.75MB (ONNX)
|
||||||
# Model: densenet-lite-gru
|
# Model: densenet-lite-gru
|
||||||
# Epoch: 39
|
# Epoch: 39
|
||||||
# Validation accuracy: 99.04%
|
# Validation accuracy: 99.04%
|
||||||
# Font: Various
|
# Font: Various
|
||||||
# Charset: Number, English character, Chinese character, symbols, <space>
|
# Charset: Number, English character, Chinese character, symbols, <space>
|
||||||
# _num_classes: 6426
|
# _num_classes: 6426
|
||||||
from module.ocr.al_ocr import AlOcr
|
return _ocr_class()(model_name='densenet-lite-gru', model_epoch=39,
|
||||||
return AlOcr(model_name='densenet-lite-gru', model_epoch=39, root='./bin/cnocr_models/cnocr', name='cnocr')
|
root='./bin/cnocr_models/cnocr', name='cnocr')
|
||||||
|
|
||||||
@cached_property
|
@cached_property
|
||||||
def jp(self):
|
def jp(self):
|
||||||
from module.ocr.al_ocr import AlOcr
|
# Folder: ./bin/cnocr_models/jp
|
||||||
return AlOcr(model_name='densenet-lite-gru', model_epoch=125, root='./bin/cnocr_models/jp', name='jp')
|
# Size: 6.36MB (ONNX)
|
||||||
|
# Model: densenet-lite-gru
|
||||||
|
# Epoch: 125
|
||||||
|
return _ocr_class()(model_name='densenet-lite-gru', model_epoch=125,
|
||||||
|
root='./bin/cnocr_models/jp', name='jp')
|
||||||
|
|
||||||
@cached_property
|
@cached_property
|
||||||
def tw(self):
|
def tw(self):
|
||||||
# Folder: ./bin/cnocr_models/tw
|
# Folder: ./bin/cnocr_models/tw
|
||||||
# Size: 8.43MB
|
# Size: 8.43MB (MXNet) / 8.64MB (ONNX)
|
||||||
# Model: densenet-lite-gru
|
# Model: densenet-lite-gru
|
||||||
# Epoch: 63
|
# Epoch: 63
|
||||||
# Validation accuracy: 99.24%
|
# Validation accuracy: 99.24%
|
||||||
# Font: Various, 6 kinds
|
# Font: Various, 6 kinds
|
||||||
# Charset: Numbers, Upper english characters, Chinese traditional characters
|
# Charset: Numbers, Upper english characters, Chinese traditional characters
|
||||||
# _num_classes: 5322
|
# _num_classes: 5322
|
||||||
from module.ocr.al_ocr import AlOcr
|
return _ocr_class()(model_name='densenet-lite-gru', model_epoch=63,
|
||||||
return AlOcr(model_name='densenet-lite-gru', model_epoch=63, root='./bin/cnocr_models/tw', name='tw')
|
root='./bin/cnocr_models/tw', name='tw')
|
||||||
|
|
||||||
|
|
||||||
OCR_MODEL = OcrModel()
|
OCR_MODEL = OcrModel()
|
||||||
|
|||||||
Reference in New Issue
Block a user