| 26 | |
| 27 | |
| 28 | class MemComputer: |
| 29 | def __init__(self, net_def, np_data_type): |
| 30 | self.net_def = net_def |
| 31 | self.np_data_type = np_data_type |
| 32 | self.const_tensor_names = [] |
| 33 | for const_tensor in net_def.tensors: |
| 34 | self.const_tensor_names.append(const_tensor.name) |
| 35 | self.input_names = [] |
| 36 | for input_info in net_def.input_info: |
| 37 | self.input_names.append(input_info.name) |
| 38 | |
| 39 | def init_computer(self): |
| 40 | self.free_mem_list = [] |
| 41 | self.used_mem_list = [] |
| 42 | self.buffer_size = 0 |
| 43 | self.ref_counts = {} |
| 44 | for op in self.net_def.op: |
| 45 | for tensor_name in op.input: |
| 46 | if tensor_name in self.const_tensor_names or \ |
| 47 | tensor_name in self.input_names: |
| 48 | continue |
| 49 | if tensor_name not in self.ref_counts: |
| 50 | self.ref_counts[tensor_name] = 0 |
| 51 | self.ref_counts[tensor_name] += 1 |
| 52 | |
| 53 | def get_mem_size(self, op, output_shape): |
| 54 | np_data_type = self.np_data_type |
| 55 | if len(op.output_type) > 0: |
| 56 | np_data_type = \ |
| 57 | data_type_to_np_dt(op.output_type[0], self.np_data_type) |
| 58 | data_type_bytes = np.dtype(np_data_type).itemsize |
| 59 | if op.type == 'WinogradTransform' or op.type == 'GEMM': |
| 60 | mace_check(len(output_shape) == 4, |
| 61 | "WinogradTransform and GEMM only support 4-dim") |
| 62 | mem_size = output_shape[2] * output_shape[3] * output_shape[0] \ |
| 63 | * int((output_shape[1] + 3) / 4) * 4 |
| 64 | else: |
| 65 | dim_size = len(output_shape) |
| 66 | if dim_size > 0: |
| 67 | mem_size = int((output_shape[dim_size - 1] + 3) / 4) * 4 |
| 68 | for i in range(dim_size - 1): |
| 69 | mem_size *= output_shape[i] |
| 70 | else: |
| 71 | print("the op %s's output dim size is 0" % op.type) |
| 72 | mem_size = 0 |
| 73 | return mem_size * data_type_bytes |
| 74 | |
| 75 | def remove_mem_block_by_name(self, mem_list, tensor_name): |
| 76 | return_mem_block = None |
| 77 | for mem_block in mem_list: |
| 78 | if tensor_name == mem_block.tensor_name: |
| 79 | return_mem_block = mem_block |
| 80 | mem_list.remove(mem_block) |
| 81 | break |
| 82 | return return_mem_block |
| 83 | |
| 84 | def fake_new(self, op): |
| 85 | output_size = len(op.output) |