| 162 | |
| 163 | |
| 164 | def get_input_shapes(sample_case, param_names): |
| 165 | param_shape_map = {} |
| 166 | name_array = [] |
| 167 | shape_array = [] |
| 168 | for i in range(len(param_names)): |
| 169 | file_name = sample_case + '/input_' + str(i) + '.pb' |
| 170 | name, data = read_pb_file(file_name) |
| 171 | param_shape_map[name] = data.shape |
| 172 | shape_array.append(data.shape) |
| 173 | if name: |
| 174 | name_array.append(name) |
| 175 | |
| 176 | if len(name_array) < len(shape_array): |
| 177 | param_shape_map = {} |
| 178 | for i in range(len(param_names)): |
| 179 | param_shape_map[param_names[i]] = shape_array[i] |
| 180 | |
| 181 | return param_shape_map |
| 182 | |
| 183 | for name in param_names: |
| 184 | if not name in param_shape_map: |
| 185 | print("Input {} does not exist!".format(name)) |
| 186 | sys.exit() |
| 187 | |
| 188 | return param_shape_map |
| 189 | |
| 190 | |
| 191 | def run_one_case(model, param_map): |