| 51 | |
| 52 | |
| 53 | class 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 |