MCPcopy Create free account
hub / github.com/XiaoMi/mace / MultiNetDefInfo

Class MultiNetDefInfo

tools/layers_validate.py:53–164  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

51
52
53class MultiNetDefInfo:
54 def __init__(self, multi_net_def, layers):
55 self.init_multi_net_def_info(multi_net_def)
56 self.EndIndexDecrement()
57 self.handle_index(layers)
58
59 def init_multi_net_def_info(self, multi_net_def):
60 netdefs = multi_net_def.net_def
61 self.net_num = len(netdefs)
62 self.net_defs = [None] * self.net_num
63 self.net_op_nums = [0] * self.net_num
64 self.quantizes = [False] * self.net_num
65 self.hexagons = [False] * self.net_num
66 for net_def in netdefs:
67 order = net_def.infer_order
68 self.net_defs[order] = net_def
69 self.net_op_nums[order] = len(net_def.op)
70 is_quantize = ConverterUtil.get_arg(
71 net_def, MaceKeyword.mace_quantize_flag_arg_str)
72 self.quantizes[order] = \
73 False if is_quantize is None else is_quantize.i == 1
74 self.hexagons[order] = self.quantizes[order] and \
75 (net_def.op[-1].type == HexagonOp.DequantizeOUTPUT_8tof.name or
76 net_def.op[-1].type == HexagonOp.OUTPUT.name)
77
78 self.end_index = self.start_index = 0
79 for op_num in self.net_op_nums:
80 self.end_index = self.end_index + op_num
81 self.start_net_idx = 0
82 self.start_op_idx = 0
83 self.end_net_idx = self.net_num
84 self.end_op_idx = self.net_op_nums[self.end_net_idx - 1]
85
86 def handle_index(self, layers):
87 num_layers = self.end_index - self.start_index + 1
88 if ':' in layers:
89 start_index, end_index = layers.split(':')
90 start_index = int(start_index) if start_index else 0
91 end_index = int(end_index) if end_index else num_layers - 1
92 else:
93 start_index = int(layers)
94 end_index = start_index + 1
95 if start_index < 0:
96 start_index += num_layers
97 if end_index < 0:
98 end_index += num_layers
99 start_index += self.start_index
100 end_index += self.start_index
101 start_index = \
102 max(self.start_index, min(self.end_index - 1, start_index))
103 end_index = max(self.start_index + 1, min(self.end_index, end_index))
104
105 for i in range(self.net_num):
106 start_index = start_index - self.net_op_nums[i]
107 if start_index < 0:
108 self.start_net_idx = i
109 self.start_op_idx = start_index + self.net_op_nums[i]
110 break

Callers 1

convertFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected