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

Function merge_params

tools/python/utils/convert_util.py:33–101  ·  view source on GitHub ↗
(net_def, data_type)

Source from the content-addressed store, hash-verified

31
32
33def merge_params(net_def, data_type):
34 def tensor_to_bytes(tensor):
35 if tensor.data_type == mace_pb2.DT_HALF:
36 data = bytearray(
37 np.array(tensor.float_data).astype(np.float16).tobytes())
38 tensor.data_size = len(tensor.float_data)
39 elif tensor.data_type == mace_pb2.DT_FLOAT:
40 data = bytearray(
41 np.array(tensor.float_data).astype(np.float32).tobytes())
42 tensor.data_size = len(tensor.float_data)
43 elif tensor.data_type == mace_pb2.DT_INT32:
44 data = bytearray(
45 np.array(tensor.int32_data).astype(np.int32).tobytes())
46 tensor.data_size = len(tensor.int32_data)
47 elif tensor.data_type == mace_pb2.DT_UINT8:
48 data = bytearray(
49 np.array(tensor.int32_data).astype(np.uint8).tolist())
50 tensor.data_size = len(tensor.int32_data)
51 elif tensor.data_type == mace_pb2.DT_INT8:
52 data = bytearray(
53 np.array(tensor.int32_data).astype(np.uint8).tolist())
54 tensor.data_size = len(tensor.int32_data)
55 elif tensor.data_type == mace_pb2.DT_FLOAT16:
56 data = bytearray(
57 np.array(tensor.float_data).astype(np.float16).tobytes())
58 tensor.data_size = len(tensor.float_data)
59 elif tensor.data_type == mace_pb2.DT_BFLOAT16:
60 data = Float2BFloat16Bytes(tensor.float_data)
61 tensor.data_size = len(tensor.float_data)
62 elif tensor.data_type == mace_pb2.DT_INT16:
63 data = bytearray(
64 np.array(tensor.int32_data).astype(np.int16).tobytes())
65 tensor.data_size = len(tensor.int32_data)
66 else:
67 raise Exception('Tensor data type %s not supported' %
68 tensor.data_type)
69 return data
70
71 model_data = []
72 offset = 0
73 for tensor in net_def.tensors:
74 if tensor.data_type == mace_pb2.DT_FLOAT:
75 tensor.data_type = data_type
76 raw_data = tensor_to_bytes(tensor)
77 if tensor.data_type != mace_pb2.DT_UINT8 and offset % 4 != 0:
78 padding = 4 - offset % 4
79 model_data.extend(bytearray([0] * padding))
80 offset += padding
81
82 tensor.offset = offset
83 model_data.extend(raw_data)
84 offset += len(raw_data)
85
86 for tensor in net_def.tensors:
87 if tensor.data_type == mace_pb2.DT_FLOAT \
88 or tensor.data_type == mace_pb2.DT_HALF \
89 or tensor.data_type == mace_pb2.DT_FLOAT16 \
90 or tensor.data_type == mace_pb2.DT_BFLOAT16:

Callers 1

convertFunction · 0.90

Calls 1

tensor_to_bytesFunction · 0.85

Tested by

no test coverage detected