diff --git a/bin/cnocr_models/azur_lane/model.onnx b/bin/cnocr_models/azur_lane/model.onnx new file mode 100644 index 000000000..b9f70a1c1 Binary files /dev/null and b/bin/cnocr_models/azur_lane/model.onnx differ diff --git a/bin/cnocr_models/azur_lane_jp/model.onnx b/bin/cnocr_models/azur_lane_jp/model.onnx new file mode 100644 index 000000000..bb2f287ac Binary files /dev/null and b/bin/cnocr_models/azur_lane_jp/model.onnx differ diff --git a/bin/cnocr_models/cnocr/model.onnx b/bin/cnocr_models/cnocr/model.onnx new file mode 100644 index 000000000..37b9c1b6c Binary files /dev/null and b/bin/cnocr_models/cnocr/model.onnx differ diff --git a/bin/cnocr_models/jp/model.onnx b/bin/cnocr_models/jp/model.onnx new file mode 100644 index 000000000..db5bf0250 Binary files /dev/null and b/bin/cnocr_models/jp/model.onnx differ diff --git a/bin/cnocr_models/tw/model.onnx b/bin/cnocr_models/tw/model.onnx new file mode 100644 index 000000000..574cd2546 Binary files /dev/null and b/bin/cnocr_models/tw/model.onnx differ diff --git a/config/deploy.template.yaml b/config/deploy.template.yaml index a3b8dbce0..1cd32ce72 100644 --- a/config/deploy.template.yaml +++ b/config/deploy.template.yaml @@ -79,6 +79,10 @@ Deploy: # Address of ocr server for alas instance to connect # [Default] 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: # Use auto update and builtin updater feature diff --git a/deploy/config.py b/deploy/config.py index ea56c2299..e1ebbd233 100644 --- a/deploy/config.py +++ b/deploy/config.py @@ -35,6 +35,7 @@ class ConfigModel: StartOcrServer: bool = False OcrServerPort: int = 22268 OcrClientAddress: str = "127.0.0.1:22268" + UseOcrOnnx: bool = False # Update EnableReload: bool = True diff --git a/dev_tools/convert_mxnet_to_onnx.py b/dev_tools/convert_mxnet_to_onnx.py new file mode 100644 index 000000000..ed7893bf9 --- /dev/null +++ b/dev_tools/convert_mxnet_to_onnx.py @@ -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.") diff --git a/module/ocr/al_ocr.py b/module/ocr/al_ocr.py index 2827c5dc7..40ea1a7b9 100644 --- a/module/ocr/al_ocr.py +++ b/module/ocr/al_ocr.py @@ -176,7 +176,7 @@ class AlOcr(CnOcr): prefix = os.path.join(self._model_dir, self._model_file_prefix) data_names = ['data'] 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( prefix, self._model_epoch, @@ -235,3 +235,62 @@ class AlOcr(CnOcr): img_list, img_widths = self._pad_arrays(img_list) image = cv2.hconcat(img_list)[0, :, :] 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) diff --git a/module/ocr/models.py b/module/ocr/models.py index 07c389631..ed11bf390 100644 --- a/module/ocr/models.py +++ b/module/ocr/models.py @@ -1,65 +1,76 @@ 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: @cached_property def azur_lane(self): # Folder: ./bin/cnocr_models/azur_lane - # Size: 3.25MB + # Size: 3.25MB (MXNet) / 3.34MB (ONNX) # Model: densenet-lite-gru # Epoch: 15 # Validation accuracy: 99.43% # Font: Impact, AgencyFB-Regular, MStiffHeiHK-UltraBold # Charset: 0123456789ABCDEFGHIJKLMNPQRSTUVWXYZ:/- (Letter 'O' and is not included) # _num_classes: 39 - from module.ocr.al_ocr import AlOcr - return AlOcr(model_name='densenet-lite-gru', model_epoch=15, root='./bin/cnocr_models/azur_lane', - name='azur_lane') + return _ocr_class()(model_name='densenet-lite-gru', model_epoch=15, + root='./bin/cnocr_models/azur_lane', name='azur_lane') @cached_property def azur_lane_jp(self): # Folder: ./bin/cnocr_models/azur_lane_jp - # Size: 3.25MB + # Size: 3.25MB (MXNet) / 3.34MB (ONNX) # Model: densenet-lite-gru # Epoch: 20 # Validation accuracy: 99.01% # Font: Impact, VibeMO Compressed Pro Thin, Folk R, Source Han Serif JP # Charset: 0123456789ABCDEFGHIJKLMNPQRSTUVWXYZ:/- (Letter 'O' and is not included) # _num_classes: 39 - from module.ocr.al_ocr import AlOcr - return AlOcr(model_name='densenet-lite-gru', model_epoch=20, root='./bin/cnocr_models/azur_lane_jp', - name='azur_lane_jp') + return _ocr_class()(model_name='densenet-lite-gru', model_epoch=20, + root='./bin/cnocr_models/azur_lane_jp', name='azur_lane_jp') @cached_property def cnocr(self): # Folder: ./bin/cnocr_models/cnocr - # Size: 9.51MB + # Size: 9.51MB (MXNet) / 9.75MB (ONNX) # Model: densenet-lite-gru # Epoch: 39 # Validation accuracy: 99.04% # Font: Various # Charset: Number, English character, Chinese character, symbols, # _num_classes: 6426 - from module.ocr.al_ocr import AlOcr - return AlOcr(model_name='densenet-lite-gru', model_epoch=39, root='./bin/cnocr_models/cnocr', name='cnocr') + return _ocr_class()(model_name='densenet-lite-gru', model_epoch=39, + root='./bin/cnocr_models/cnocr', name='cnocr') @cached_property def jp(self): - from module.ocr.al_ocr import AlOcr - return AlOcr(model_name='densenet-lite-gru', model_epoch=125, root='./bin/cnocr_models/jp', name='jp') + # Folder: ./bin/cnocr_models/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 def tw(self): # Folder: ./bin/cnocr_models/tw - # Size: 8.43MB + # Size: 8.43MB (MXNet) / 8.64MB (ONNX) # Model: densenet-lite-gru # Epoch: 63 # Validation accuracy: 99.24% # Font: Various, 6 kinds # Charset: Numbers, Upper english characters, Chinese traditional characters # _num_classes: 5322 - from module.ocr.al_ocr import AlOcr - return AlOcr(model_name='densenet-lite-gru', model_epoch=63, root='./bin/cnocr_models/tw', name='tw') + return _ocr_class()(model_name='densenet-lite-gru', model_epoch=63, + root='./bin/cnocr_models/tw', name='tw') OCR_MODEL = OcrModel()