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

Class MultiNetDefInfo

tools/python/layers_validate.py:53–165  ·  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] = \
75 self.quantizes[order] and \
76 (net_def.op[-1].type == HexagonOp.DequantizeOUTPUT_8tof.name or
77 net_def.op[-1].type == HexagonOp.OUTPUT.name)
78
79 self.end_index = self.start_index = 0
80 for op_num in self.net_op_nums:
81 self.end_index = self.end_index + op_num
82 self.start_net_idx = 0
83 self.start_op_idx = 0
84 self.end_net_idx = self.net_num
85 self.end_op_idx = self.net_op_nums[self.end_net_idx - 1]
86
87 def handle_index(self, layers):
88 num_layers = self.end_index - self.start_index + 1
89 if ':' in layers:
90 start_index, end_index = layers.split(':')
91 start_index = int(start_index) if start_index else 0
92 end_index = int(end_index) if end_index else num_layers - 1
93 else:
94 start_index = int(layers)
95 end_index = start_index + 1
96 if start_index < 0:
97 start_index += num_layers
98 if end_index < 0:
99 end_index += num_layers
100 start_index += self.start_index
101 end_index += self.start_index
102 start_index = \
103 max(self.start_index, min(self.end_index - 1, start_index))
104 end_index = max(self.start_index + 1, min(self.end_index, end_index))
105
106 for i in range(self.net_num):
107 start_index = start_index - self.net_op_nums[i]
108 if start_index < 0:
109 self.start_net_idx = i
110 self.start_op_idx = start_index + self.net_op_nums[i]

Callers 1

convertFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected