converter.py 52.3 KB
Newer Older
Y
yejianwu 已提交
1 2 3 4 5 6 7 8 9 10 11 12 13
# Copyright 2018 Xiaomi, Inc.  All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
#     http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
14 15

import argparse
Y
yejianwu 已提交
16
import filelock
L
liuqi 已提交
17
import glob
18
import hashlib
19
import os
L
liuqi 已提交
20
import re
L
Liangliang He 已提交
21
import sh
22 23
import subprocess
import sys
24
import urllib
Y
yejianwu 已提交
25
import yaml
L
liuqi 已提交
26

27
from enum import Enum
28

29
import sh_commands
30
from sh_commands import BuildType
L
Liangliang He 已提交
31

32
from common import CaffeEnvType
33
from common import DeviceType
34 35 36
from common import mace_check
from common import MaceLogger
from common import StringFormatter
37

38 39 40 41
################################
# common definitions
################################
BUILD_OUTPUT_DIR = 'build'
42
PHONE_DATA_DIR = "/data/local/tmp/mace_run"
43
MODEL_OUTPUT_DIR_NAME = 'model'
L
liuqi 已提交
44
MODEL_HEADER_DIR_PATH = 'include/mace/public'
45 46
BUILD_TMP_DIR_NAME = '_tmp'
BUILD_TMP_GENERAL_OUTPUT_DIR_NAME = 'general'
Y
yejianwu 已提交
47
OUTPUT_LIBRARY_DIR_NAME = 'lib'
48
OUTPUT_OPENCL_BINARY_DIR_NAME = 'opencl'
L
liuqi 已提交
49
OUTPUT_OPENCL_BINARY_FILE_NAME = 'compiled_opencl_kernel'
50
CL_COMPILED_BINARY_FILE_NAME = "mace_cl_compiled_program.bin"
L
liuqi 已提交
51 52
CODEGEN_BASE_DIR = 'mace/codegen'
MODEL_CODEGEN_DIR = CODEGEN_BASE_DIR + '/models'
Y
yejianwu 已提交
53
LIBMACE_SO_TARGET = "//mace:libmace.so"
Y
yejianwu 已提交
54 55
MACE_RUN_STATIC_TARGET = "//mace/tools/validation:mace_run_static"
MACE_RUN_SHARED_TARGET = "//mace/tools/validation:mace_run_shared"
56
ALL_SOC_TAG = 'all'
57 58

ABITypeStrs = [
L
liuqi 已提交
59 60 61
    'armeabi-v7a',
    'arm64-v8a',
    'host',
62
]
L
liuqi 已提交
63 64 65 66 67 68


class ABIType(object):
    armeabi_v7a = 'armeabi-v7a'
    arm64_v8a = 'arm64-v8a'
    host = 'host'
69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107


PlatformTypeStrs = [
    "tensorflow",
    "caffe",
]
PlatformType = Enum('PlatformType', [(ele, ele) for ele in PlatformTypeStrs],
                    type=str)

RuntimeTypeStrs = [
    "cpu",
    "gpu",
    "dsp",
    "cpu+gpu"
]


class RuntimeType(object):
    cpu = 'cpu'
    gpu = 'gpu'
    dsp = 'dsp'
    cpu_gpu = 'cpu+gpu'


CPUDataTypeStrs = [
    "fp32",
]

CPUDataType = Enum('CPUDataType', [(ele, ele) for ele in CPUDataTypeStrs],
                   type=str)

GPUDataTypeStrs = [
    "fp16_fp32",
    "fp32_fp32",
]

GPUDataType = Enum('GPUDataType', [(ele, ele) for ele in GPUDataTypeStrs],
                   type=str)

L
liuqi 已提交
108 109 110 111 112 113 114
DSPDataTypeStrs = [
    "uint8",
]

DSPDataType = Enum('DSPDataType', [(ele, ele) for ele in DSPDataTypeStrs],
                   type=str)

115 116
WinogradParameters = [0, 2, 4]

117 118 119 120 121 122 123 124 125 126 127 128 129 130

class DefaultValues(object):
    omp_num_threads = -1,
    cpu_affinity_policy = 1,
    gpu_perf_hint = 3,
    gpu_priority_hint = 3,


class YAMLKeyword(object):
    library_name = 'library_name'
    target_abis = 'target_abis'
    target_socs = 'target_socs'
    build_type = 'build_type'
    embed_model_data = 'embed_model_data'
Y
yejianwu 已提交
131
    linkshared = 'linkshared'
132 133 134 135 136 137 138 139 140
    models = 'models'
    platform = 'platform'
    model_file_path = 'model_file_path'
    model_sha256_checksum = 'model_sha256_checksum'
    weight_file_path = 'weight_file_path'
    weight_sha256_checksum = 'weight_sha256_checksum'
    subgraphs = 'subgraphs'
    input_tensors = 'input_tensors'
    input_shapes = 'input_shapes'
李寅 已提交
141
    input_ranges = 'input_ranges'
142 143 144 145 146 147 148 149 150
    output_tensors = 'output_tensors'
    output_shapes = 'output_shapes'
    runtime = 'runtime'
    data_type = 'data_type'
    limit_opencl_kernel_time = 'limit_opencl_kernel_time'
    nnlib_graph_mode = 'nnlib_graph_mode'
    obfuscate = 'obfuscate'
    winograd = 'winograd'
    validation_inputs_data = 'validation_inputs_data'
李寅 已提交
151
    graph_optimize_options = 'graph_optimize_options'  # internal use for now
152 153 154 155 156


class ModuleName(object):
    YAML_CONFIG = 'YAML CONFIG'
    MODEL_CONVERTER = 'Model Converter'
L
liuqi 已提交
157 158
    RUN = 'RUN'
    BENCHMARK = 'Benchmark'
159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180


CPP_KEYWORDS = [
    'alignas', 'alignof', 'and', 'and_eq', 'asm', 'atomic_cancel',
    'atomic_commit', 'atomic_noexcept', 'auto', 'bitand', 'bitor',
    'bool', 'break', 'case', 'catch', 'char', 'char16_t', 'char32_t',
    'class', 'compl', 'concept', 'const', 'constexpr', 'const_cast',
    'continue', 'co_await', 'co_return', 'co_yield', 'decltype', 'default',
    'delete', 'do', 'double', 'dynamic_cast', 'else', 'enum', 'explicit',
    'export', 'extern', 'false', 'float', 'for', 'friend', 'goto', 'if',
    'import', 'inline', 'int', 'long', 'module', 'mutable', 'namespace',
    'new', 'noexcept', 'not', 'not_eq', 'nullptr', 'operator', 'or', 'or_eq',
    'private', 'protected', 'public', 'register', 'reinterpret_cast',
    'requires', 'return', 'short', 'signed', 'sizeof', 'static',
    'static_assert', 'static_cast', 'struct', 'switch', 'synchronized',
    'template', 'this', 'thread_local', 'throw', 'true', 'try', 'typedef',
    'typeid', 'typename', 'union', 'unsigned', 'using', 'virtual', 'void',
    'volatile', 'wchar_t', 'while', 'xor', 'xor_eq', 'override', 'final',
    'transaction_safe', 'transaction_safe_dynamic', 'if', 'elif', 'else',
    'endif', 'defined', 'ifdef', 'ifndef', 'define', 'undef', 'include',
    'line', 'error', 'pragma',
]
Y
yejianwu 已提交
181

182

183 184 185
################################
# common functions
################################
186
def parse_device_type(runtime):
Y
yejianwu 已提交
187
    device_type = ""
188

189
    if runtime == RuntimeType.dsp:
190
        device_type = DeviceType.HEXAGON
191
    elif runtime == RuntimeType.gpu:
192
        device_type = DeviceType.GPU
193
    elif runtime == RuntimeType.cpu:
194
        device_type = DeviceType.CPU
195

196
    return device_type
197

Y
yejianwu 已提交
198 199

def get_hexagon_mode(configs):
L
Liangliang He 已提交
200
    runtime_list = []
L
liuqi 已提交
201 202 203 204
    for model_name in configs[YAMLKeyword.models]:
        model_runtime =\
            configs[YAMLKeyword.models][model_name].get(
                YAMLKeyword.runtime, "")
L
Liangliang He 已提交
205 206
        runtime_list.append(model_runtime.lower())

L
liuqi 已提交
207
    if RuntimeType.dsp in runtime_list:
Y
yejianwu 已提交
208 209 210 211
        return True
    return False


212 213 214 215
def md5sum(str):
    md5 = hashlib.md5()
    md5.update(str)
    return md5.hexdigest()
216

Y
yejianwu 已提交
217

218 219 220 221 222 223
def sha256_checksum(fname):
    hash_func = hashlib.sha256()
    with open(fname, "rb") as f:
        for chunk in iter(lambda: f.read(4096), b""):
            hash_func.update(chunk)
    return hash_func.hexdigest()
Y
yejianwu 已提交
224

W
wuchenghui 已提交
225

226 227
def format_model_config(flags):
    with open(flags.config) as f:
228
        configs = yaml.load(f)
W
wuchenghui 已提交
229

230 231
    library_name = configs.get(YAMLKeyword.library_name, "")
    mace_check(len(library_name) > 0,
L
liuqi 已提交
232
               ModuleName.YAML_CONFIG, "library name should not be empty")
233

234 235 236 237
    if flags.target_abis:
        target_abis = flags.target_abis.split(',')
    else:
        target_abis = configs.get(YAMLKeyword.target_abis, [])
238 239
    mace_check((isinstance(target_abis, list) and len(target_abis) > 0),
               ModuleName.YAML_CONFIG, "target_abis list is needed")
240
    configs[YAMLKeyword.target_abis] = target_abis
241 242 243 244 245 246
    for abi in target_abis:
        mace_check(abi in ABITypeStrs,
                   ModuleName.YAML_CONFIG,
                   "target_abis must be in " + str(ABITypeStrs))

    target_socs = configs.get(YAMLKeyword.target_socs, "")
247 248 249 250
    if flags.target_socs:
        configs[YAMLKeyword.target_socs] = \
               [soc.lower() for soc in flags.target_socs.split(',')]
    elif not target_socs:
251 252 253 254
        configs[YAMLKeyword.target_socs] = []
    elif not isinstance(target_socs, list):
        configs[YAMLKeyword.target_socs] = [target_socs]

255 256 257
    configs[YAMLKeyword.target_socs] = \
        [soc.lower() for soc in configs[YAMLKeyword.target_socs]]

L
liuqi 已提交
258 259
    if ABIType.armeabi_v7a in target_abis \
            or ABIType.arm64_v8a in target_abis:
260
        available_socs = sh_commands.adb_get_all_socs()
261 262 263 264 265 266 267
        target_socs = configs[YAMLKeyword.target_socs]
        if ALL_SOC_TAG in target_socs:
            mace_check(available_socs,
                       ModuleName.YAML_CONFIG,
                       "Build for all SOCs plugged in computer, "
                       "you at least plug in one phone")
        else:
268 269 270 271 272 273 274
            for soc in target_socs:
                mace_check(soc in available_socs,
                           ModuleName.YAML_CONFIG,
                           "Build specified SOC library, "
                           "you must plug in a phone using the SOC")

    build_type = BuildType.code
275 276 277 278
    if flags.build_type:
        build_type_str = flags.build_type
    else:
        build_type_str = configs.get(YAMLKeyword.build_type, "")
279 280 281 282 283 284 285 286 287 288
    if build_type_str == BuildType.proto:
        build_type = BuildType.proto
    elif build_type_str == BuildType.code:
        build_type = BuildType.code
    else:
        MaceLogger.error(ModuleName.YAML_CONFIG,
                         "Invalid build type " + build_type_str
                         + ". only support [proto|code] format, "
                         + "proto for converting model to ProtoBuf file, "
                         + "code for converting model to c++ code.")
289
    configs[YAMLKeyword.build_type] = build_type
290 291 292 293 294 295 296 297 298
    embed_model_data = configs.get(YAMLKeyword.embed_model_data, "")
    if embed_model_data == "" or not isinstance(embed_model_data, int) or \
       embed_model_data < 0 or embed_model_data > 1:
        MaceLogger.error(ModuleName.YAML_CONFIG,
                         "embed_model_data must be 0 or 1. "
                         "0 for embed model data to code, 1 not.")
    if build_type == BuildType.proto:
        configs[YAMLKeyword.embed_model_data] = 0

Y
yejianwu 已提交
299 300 301 302 303 304
    linkshared = configs.get(YAMLKeyword.linkshared, "")
    if linkshared == "":
        configs[YAMLKeyword.linkshared] = 0
        linkshared = 0
    if not isinstance(linkshared, int) or linkshared < 0 or \
       linkshared > 1:
Y
yejianwu 已提交
305
        MaceLogger.error(ModuleName.YAML_CONFIG,
Y
yejianwu 已提交
306
                         "linkshared must be 0 or 1. "
Y
yejianwu 已提交
307
                         "default is 0, for link mace lib statically, "
Y
yejianwu 已提交
308 309
                         "1 for dynamic linking.")
    if build_type == BuildType.code and linkshared == 1:
Y
yejianwu 已提交
310
        MaceLogger.error(ModuleName.YAML_CONFIG,
Y
yejianwu 已提交
311
                         "'linkshared == 1' only support when "
Y
yejianwu 已提交
312 313
                         "'build_type == proto'")

314 315 316 317
    model_names = configs.get(YAMLKeyword.models, [])
    mace_check(len(model_names) > 0, ModuleName.YAML_CONFIG,
               "no model found in config file")

L
liuqi 已提交
318
    model_name_reg = re.compile(r'^[a-zA-Z0-9_]+$')
319 320 321 322 323 324 325 326
    for model_name in model_names:
        # check model_name legality
        mace_check(model_name not in CPP_KEYWORDS,
                   ModuleName.YAML_CONFIG,
                   "model name should not be c++ keyword.")
        mace_check((model_name[0] == '_' or model_name[0].isalpha())
                   and bool(model_name_reg.match(model_name)),
                   ModuleName.YAML_CONFIG,
L
liuqi 已提交
327
                   "model name should Meet the c++ naming convention"
328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383
                   " which start with '_' or alpha"
                   " and only contain alpha, number and '_'")

        model_config = configs[YAMLKeyword.models][model_name]
        platform = model_config.get(YAMLKeyword.platform, "")
        mace_check(platform in PlatformTypeStrs,
                   ModuleName.YAML_CONFIG,
                   "'platform' must be in " + str(PlatformTypeStrs))

        for key in [YAMLKeyword.model_file_path,
                    YAMLKeyword.model_sha256_checksum]:
            value = model_config.get(key, "")
            mace_check(value != "", ModuleName.YAML_CONFIG,
                       "'%s' is necessary" % key)

        weight_file_path = model_config.get(YAMLKeyword.weight_file_path, "")
        if weight_file_path:
            weight_checksum =\
                model_config.get(YAMLKeyword.weight_sha256_checksum, "")
            mace_check(weight_checksum != "", ModuleName.YAML_CONFIG,
                       "'%s' is necessary" %
                       YAMLKeyword.weight_sha256_checksum)
        else:
            model_config[YAMLKeyword.weight_sha256_checksum] = ""

        runtime = model_config.get(YAMLKeyword.runtime, "")
        mace_check(runtime in RuntimeTypeStrs,
                   ModuleName.YAML_CONFIG,
                   "'runtime' must be in " + str(RuntimeTypeStrs))
        if ABIType.host in target_abis:
            mace_check(runtime == RuntimeType.cpu,
                       ModuleName.YAML_CONFIG,
                       "host only support cpu runtime now.")

        data_type = model_config.get(YAMLKeyword.data_type, "")
        if runtime == RuntimeType.cpu_gpu and data_type not in GPUDataTypeStrs:
            model_config[YAMLKeyword.data_type] = \
                GPUDataType.fp16_fp32.value
        elif runtime == RuntimeType.cpu:
            if len(data_type) > 0:
                mace_check(data_type in CPUDataTypeStrs,
                           ModuleName.YAML_CONFIG,
                           "'data_type' must be in " + str(CPUDataTypeStrs)
                           + " for cpu runtime")
            else:
                model_config[YAMLKeyword.data_type] = \
                    CPUDataType.fp32.value
        elif runtime == RuntimeType.gpu:
            if len(data_type) > 0:
                mace_check(data_type in GPUDataTypeStrs,
                           ModuleName.YAML_CONFIG,
                           "'data_type' must be in " + str(GPUDataTypeStrs)
                           + " for gpu runtime")
            else:
                model_config[YAMLKeyword.data_type] =\
                    GPUDataType.fp16_fp32.value
L
liuqi 已提交
384 385 386 387 388 389 390 391 392
        elif runtime == RuntimeType.dsp:
            if len(data_type) > 0:
                mace_check(data_type in DSPDataTypeStrs,
                           ModuleName.YAML_CONFIG,
                           "'data_type' must be in " + str(DSPDataTypeStrs)
                           + " for dsp runtime")
            else:
                model_config[YAMLKeyword.data_type] = \
                    DSPDataType.uint8.value
393 394 395 396 397 398 399 400 401 402 403 404 405 406 407

        subgraphs = model_config.get(YAMLKeyword.subgraphs, "")
        mace_check(len(subgraphs) > 0, ModuleName.YAML_CONFIG,
                   "at least one subgraph is needed")

        for subgraph in subgraphs:
            for key in [YAMLKeyword.input_tensors,
                        YAMLKeyword.input_shapes,
                        YAMLKeyword.output_tensors,
                        YAMLKeyword.output_shapes]:
                value = subgraph.get(key, "")
                mace_check(value != "", ModuleName.YAML_CONFIG,
                           "'%s' is necessary in subgraph" % key)
                if not isinstance(value, list):
                    subgraph[key] = [value]
L
liuqi 已提交
408 409 410 411 412 413 414 415
            validation_inputs_data = subgraph.get(
                YAMLKeyword.validation_inputs_data, [])
            if not isinstance(validation_inputs_data, list):
                subgraph[YAMLKeyword.validation_inputs_data] = [
                    validation_inputs_data]
            else:
                subgraph[YAMLKeyword.validation_inputs_data] = \
                    validation_inputs_data
416 417 418 419 420 421
            input_ranges = subgraph.get(
                YAMLKeyword.input_ranges, [])
            if not isinstance(input_ranges, list):
                subgraph[YAMLKeyword.input_ranges] = [input_ranges]
            else:
                subgraph[YAMLKeyword.input_ranges] = input_ranges
W
wuchenghui 已提交
422

423 424 425 426 427 428 429
        for key in [YAMLKeyword.limit_opencl_kernel_time,
                    YAMLKeyword.nnlib_graph_mode,
                    YAMLKeyword.obfuscate,
                    YAMLKeyword.winograd]:
            value = model_config.get(key, "")
            if value == "":
                model_config[key] = 0
L
Liangliang He 已提交
430

431 432 433 434 435 436
        mace_check(model_config[YAMLKeyword.winograd] in WinogradParameters,
                   ModuleName.YAML_CONFIG,
                   "'winograd' parameters must be in "
                   + str(WinogradParameters) +
                   ". 0 for disable winograd convolution")

L
liuqi 已提交
437 438
        weight_file_path = model_config.get(YAMLKeyword.weight_file_path, "")
        model_config[YAMLKeyword.weight_file_path] = weight_file_path
Y
yejianwu 已提交
439

440
    return configs
Y
yejianwu 已提交
441

W
wuchenghui 已提交
442

443 444 445 446
def get_build_binary_dir(library_name, target_abi, target_soc,
                         serial_num):
    if not target_soc or not serial_num:
        binary_path_digest = md5sum(target_abi)
L
liuqi 已提交
447
        binary_path_digest = "%s_%s" % (target_abi, binary_path_digest)
Y
yejianwu 已提交
448
    else:
L
liuqi 已提交
449
        device_name = sh_commands.adb_get_device_name_by_serialno(serial_num)
450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473
        binary_path_digest = md5sum(target_abi + target_soc + serial_num)
        binary_path_digest = "%s_%s_%s" % \
                             (device_name, target_soc, binary_path_digest)
    return "%s/%s/%s/%s" % (
        BUILD_OUTPUT_DIR, library_name, BUILD_TMP_DIR_NAME, binary_path_digest)


def get_build_model_dirs(library_name, model_name, target_abi, target_soc,
                         serial_num, model_file_path):
    model_path_digest = md5sum(model_file_path)
    model_output_base_dir = "%s/%s/%s/%s/%s" % (
        BUILD_OUTPUT_DIR, library_name, BUILD_TMP_DIR_NAME,
        model_name, model_path_digest)

    if target_abi == ABIType.host:
        model_output_dir = "%s/%s" % (model_output_base_dir, target_abi)
    elif not target_soc or not serial_num:
        model_output_dir = "%s/%s/%s" % (
            model_output_base_dir, BUILD_TMP_GENERAL_OUTPUT_DIR_NAME,
            target_abi)
    else:
        device_name = \
            sh_commands.adb_get_device_name_by_serialno(serial_num)
        model_output_dir = "%s/%s_%s/%s" % (
L
liuqi 已提交
474
            model_output_base_dir, device_name,
475
            target_soc, target_abi)
Y
yejianwu 已提交
476

477 478 479 480 481 482
    mace_model_dir = \
        '%s/%s/%s' % (BUILD_OUTPUT_DIR, library_name, MODEL_OUTPUT_DIR_NAME)

    return model_output_base_dir, model_output_dir, mace_model_dir


L
liuqi 已提交
483 484 485 486 487
def get_opencl_binary_output_path(library_name, target_abi,
                                  target_soc, serial_num):
    device_name = \
        sh_commands.adb_get_device_name_by_serialno(serial_num)
    return '%s/%s/%s/%s/%s_%s.%s.%s.bin' % \
488 489 490
           (BUILD_OUTPUT_DIR,
            library_name,
            OUTPUT_OPENCL_BINARY_DIR_NAME,
L
liuqi 已提交
491 492 493 494 495
            target_abi,
            library_name,
            OUTPUT_OPENCL_BINARY_FILE_NAME,
            device_name,
            target_soc)
496 497


Y
yejianwu 已提交
498 499 500 501 502 503 504
def get_shared_library_dir(library_name, abi):
    return '%s/%s/%s/%s' % (BUILD_OUTPUT_DIR,
                            library_name,
                            OUTPUT_LIBRARY_DIR_NAME,
                            abi)


505 506 507 508 509 510 511
################################
# build
################################
def pull_opencl_binary_and_tuning_param(target_abi,
                                        serialno,
                                        model_output_dirs):
    sh_commands.pull_binaries(target_abi, serialno, model_output_dirs,
512
                              CL_COMPILED_BINARY_FILE_NAME)
513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528


def print_configuration(flags, configs):
    title = "Common Configuration"
    header = ["key", "value"]
    data = list()
    data.append([YAMLKeyword.library_name,
                 configs[YAMLKeyword.library_name]])
    data.append([YAMLKeyword.target_abis,
                 configs[YAMLKeyword.target_abis]])
    data.append([YAMLKeyword.target_socs,
                 configs[YAMLKeyword.target_socs]])
    data.append([YAMLKeyword.build_type,
                 configs[YAMLKeyword.build_type]])
    data.append([YAMLKeyword.embed_model_data,
                 configs[YAMLKeyword.embed_model_data]])
Y
yejianwu 已提交
529 530
    data.append([YAMLKeyword.linkshared,
                 configs[YAMLKeyword.linkshared]])
531
    data.append(["Tuning", flags.disable_tuning])
532
    MaceLogger.summary(StringFormatter.table(header, data, title))
L
Liangliang He 已提交
533

Y
yejianwu 已提交
534

535 536 537
def download_model_files(model_file_path,
                         model_output_dir,
                         weight_file_path=""):
538
    MaceLogger.info("Downloading model, please wait ...")
L
Liangliang He 已提交
539 540
    if model_file_path.startswith("http://") or \
            model_file_path.startswith("https://"):
Y
yejianwu 已提交
541 542
        model_file = model_output_dir + "/model.pb"
        urllib.urlretrieve(model_file_path, model_file)
543 544 545 546 547

    if weight_file_path.startswith("http://") or \
            weight_file_path.startswith("https://"):
        weight_file = model_output_dir + "/model.caffemodel"
        urllib.urlretrieve(weight_file_path, weight_file)
548
    MaceLogger.info("Model downloaded successfully.")
549 550 551 552 553 554 555 556


def get_model_files_path(model_file_path,
                         model_output_dir,
                         weight_file_path=""):
    if model_file_path.startswith("http://") or \
            model_file_path.startswith("https://"):
        model_file = model_output_dir + "/model.pb"
Y
yejianwu 已提交
557 558
    else:
        model_file = model_file_path
L
Liangliang He 已提交
559 560 561

    if weight_file_path.startswith("http://") or \
            weight_file_path.startswith("https://"):
Y
yejianwu 已提交
562
        weight_file = model_output_dir + "/model.caffemodel"
Y
yejianwu 已提交
563 564
    else:
        weight_file = weight_file_path
Y
yejianwu 已提交
565 566

    return model_file, weight_file
L
Liangliang He 已提交
567

L
liuqi 已提交
568

569 570 571 572 573
def convert_model(configs):
    # Remove previous output dirs
    library_name = configs[YAMLKeyword.library_name]
    if not os.path.exists(BUILD_OUTPUT_DIR):
        os.makedirs(BUILD_OUTPUT_DIR)
L
liuqi 已提交
574 575 576
    elif os.path.exists(os.path.join(BUILD_OUTPUT_DIR, library_name)):
        sh.rm("-rf", os.path.join(BUILD_OUTPUT_DIR, library_name))
    os.makedirs(os.path.join(BUILD_OUTPUT_DIR, library_name))
577 578 579

    model_output_dir = \
        '%s/%s/%s' % (BUILD_OUTPUT_DIR, library_name, MODEL_OUTPUT_DIR_NAME)
L
liuqi 已提交
580 581
    model_header_dir = \
        '%s/%s/%s' % (BUILD_OUTPUT_DIR, library_name, MODEL_HEADER_DIR_PATH)
582 583 584
    if os.path.exists(model_output_dir):
        sh.rm("-rf", model_output_dir)
    os.makedirs(model_output_dir)
L
liuqi 已提交
585 586 587 588 589
    if os.path.exists(model_header_dir):
        sh.rm("-rf", model_header_dir)
    os.makedirs(model_header_dir)
    # copy header files
    sh.cp("-f", glob.glob("mace/public/*.h"), model_header_dir)
590

L
liuqi 已提交
591 592 593 594 595 596 597
    sh_commands.gen_mace_engine_factory_source(
        configs[YAMLKeyword.models].keys(),
        configs[YAMLKeyword.build_type])
    if configs[YAMLKeyword.build_type] == BuildType.code:
        sh.cp("-f", glob.glob("mace/codegen/engine/*.h"),
              model_header_dir)

598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643
    embed_model_data = configs[YAMLKeyword.embed_model_data]

    sh_commands.clear_model_codegen()
    for model_name in configs[YAMLKeyword.models]:
        MaceLogger.header(
            StringFormatter.block("Convert %s model" % model_name))
        model_config = configs[YAMLKeyword.models][model_name]
        runtime = model_config[YAMLKeyword.runtime]

        # Create model build directory
        model_path_digest = md5sum(
            model_config[YAMLKeyword.model_file_path])

        model_output_base_dir = "%s/%s/%s/%s/%s" % (
            BUILD_OUTPUT_DIR, library_name, BUILD_TMP_DIR_NAME,
            model_name, model_path_digest)

        if os.path.exists(model_output_base_dir):
            sh.rm("-rf", model_output_base_dir)
        os.makedirs(model_output_base_dir)

        download_model_files(
            model_config[YAMLKeyword.model_file_path],
            model_output_base_dir,
            model_config[YAMLKeyword.weight_file_path])

        model_file_path, weight_file_path = get_model_files_path(
            model_config[YAMLKeyword.model_file_path],
            model_output_base_dir,
            model_config[YAMLKeyword.weight_file_path])

        if sha256_checksum(model_file_path) != \
                model_config[YAMLKeyword.model_sha256_checksum]:
            MaceLogger.error(ModuleName.MODEL_CONVERTER,
                             "model file sha256checksum not match")

        if weight_file_path:
            if sha256_checksum(weight_file_path) != \
                    model_config[YAMLKeyword.weight_sha256_checksum]:
                MaceLogger.error(ModuleName.MODEL_CONVERTER,
                                 "weight file sha256checksum not match")

        data_type = model_config[YAMLKeyword.data_type]
        # TODO(liuqi): support multiple subgraphs
        subgraphs = model_config[YAMLKeyword.subgraphs]

L
liuqi 已提交
644
        model_codegen_dir = "%s/%s" % (MODEL_CODEGEN_DIR, model_name)
645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661
        sh_commands.gen_model_code(
            model_codegen_dir,
            model_config[YAMLKeyword.platform],
            model_file_path,
            weight_file_path,
            model_config[YAMLKeyword.model_sha256_checksum],
            model_config[YAMLKeyword.weight_sha256_checksum],
            ",".join(subgraphs[0][YAMLKeyword.input_tensors]),
            ",".join(subgraphs[0][YAMLKeyword.output_tensors]),
            runtime,
            model_name,
            ":".join(subgraphs[0][YAMLKeyword.input_shapes]),
            model_config[YAMLKeyword.nnlib_graph_mode],
            embed_model_data,
            model_config[YAMLKeyword.winograd],
            model_config[YAMLKeyword.obfuscate],
            configs[YAMLKeyword.build_type],
李寅 已提交
662
            data_type,
李寅 已提交
663
            ",".join(model_config.get(YAMLKeyword.graph_optimize_options, [])))
664

L
liuqi 已提交
665 666 667 668
        if configs[YAMLKeyword.build_type] == BuildType.proto:
            sh.mv("-f",
                  '%s/%s.pb' % (model_codegen_dir, model_name),
                  model_output_dir)
L
liuqi 已提交
669 670 671
            sh.mv("-f",
                  '%s/%s.data' % (model_codegen_dir, model_name),
                  model_output_dir)
L
liuqi 已提交
672 673
        else:
            if not embed_model_data:
L
liuqi 已提交
674
                sh.mv("-f",
L
liuqi 已提交
675
                      '%s/%s.data' % (model_codegen_dir, model_name),
L
liuqi 已提交
676
                      model_output_dir)
L
liuqi 已提交
677 678
            sh.cp("-f", glob.glob("mace/codegen/models/*/*.h"),
                  model_header_dir)
679

L
liuqi 已提交
680
        MaceLogger.summary(
681 682 683 684 685 686 687 688 689
            StringFormatter.block("Model %s converted" % model_name))


def build_specific_lib(target_abi, target_soc, serial_num,
                       configs, tuning, enable_openmp,
                       address_sanitizer):
    library_name = configs[YAMLKeyword.library_name]
    build_type = configs[YAMLKeyword.build_type]
    embed_model_data = configs[YAMLKeyword.embed_model_data]
Y
yejianwu 已提交
690
    linkshared = configs[YAMLKeyword.linkshared]
691 692 693 694 695 696 697 698 699
    hexagon_mode = get_hexagon_mode(configs)
    model_output_dirs = []

    build_tmp_binary_dir = get_build_binary_dir(library_name, target_abi,
                                                target_soc, serial_num)
    if os.path.exists(build_tmp_binary_dir):
        sh.rm("-rf", build_tmp_binary_dir)
    os.makedirs(build_tmp_binary_dir)

700
    sh_commands.gen_tuning_param_code(model_output_dirs)
Y
yejianwu 已提交
701 702
    if linkshared == 0:
        mace_run_target = MACE_RUN_STATIC_TARGET
Y
yejianwu 已提交
703
    else:
Y
yejianwu 已提交
704
        mace_run_target = MACE_RUN_SHARED_TARGET
Y
yejianwu 已提交
705 706 707 708 709 710 711 712 713 714 715 716 717
        sh_commands.bazel_build(
            LIBMACE_SO_TARGET,
            abi=target_abi,
            hexagon_mode=hexagon_mode,
            enable_openmp=enable_openmp,
            address_sanitizer=address_sanitizer
        )
        sh_commands.update_libmace_shared_library(serial_num,
                                                  target_abi,
                                                  library_name,
                                                  BUILD_OUTPUT_DIR,
                                                  OUTPUT_LIBRARY_DIR_NAME)

718
    sh_commands.bazel_build(
Y
yejianwu 已提交
719
        mace_run_target,
720 721 722 723 724
        abi=target_abi,
        hexagon_mode=hexagon_mode,
        enable_openmp=enable_openmp,
        address_sanitizer=address_sanitizer
    )
Y
yejianwu 已提交
725
    sh_commands.update_mace_run_lib(build_tmp_binary_dir, linkshared)
726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743
    binary_changed = False

    for model_name in configs[YAMLKeyword.models]:
        model_config = configs[YAMLKeyword.models][model_name]
        model_runtime = model_config[YAMLKeyword.runtime]
        # Create model build directory
        model_output_base_dir, model_output_dir, mace_model_dir = \
            get_build_model_dirs(library_name, model_name, target_abi,
                                 target_soc, serial_num,
                                 model_config[YAMLKeyword.model_file_path])

        model_output_dirs.append(model_output_dir)

        if os.path.exists(model_output_dir):
            sh.rm("-rf", model_output_dir)
        os.makedirs(model_output_dir)

        # build for specified soc
744
        if not address_sanitizer and target_abi != ABIType.host \
745 746 747 748 749 750 751 752 753 754
                and target_soc is not None and \
                model_runtime in [RuntimeType.gpu, RuntimeType.cpu_gpu]:
            sh_commands.clear_phone_data_dir(serial_num, PHONE_DATA_DIR)

            subgraphs = model_config[YAMLKeyword.subgraphs]
            # generate input data
            sh_commands.gen_random_input(
                model_output_dir,
                subgraphs[0][YAMLKeyword.input_tensors],
                subgraphs[0][YAMLKeyword.input_shapes],
李寅 已提交
755
                subgraphs[0][YAMLKeyword.validation_inputs_data],
756
                input_ranges=subgraphs[0][YAMLKeyword.input_ranges])
757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778

            device_type = parse_device_type(RuntimeType.gpu)
            sh_commands.tuning_run(
                abi=target_abi,
                serialno=serial_num,
                mace_run_dir=build_tmp_binary_dir,
                vlog_level=0,
                embed_model_data=embed_model_data,
                model_output_dir=model_output_dir,
                input_nodes=subgraphs[0][YAMLKeyword.input_tensors],
                output_nodes=subgraphs[0][YAMLKeyword.output_tensors],
                input_shapes=subgraphs[0][YAMLKeyword.input_shapes],
                output_shapes=subgraphs[0][YAMLKeyword.output_shapes],
                mace_model_dir=mace_model_dir,
                model_tag=model_name,
                device_type=device_type,
                running_round=0,
                restart_round=1,
                limit_opencl_kernel_time=model_config[YAMLKeyword.limit_opencl_kernel_time],  # noqa
                tuning=tuning,
                out_of_range_check=False,
                phone_data_dir=PHONE_DATA_DIR,
779 780
                build_type=build_type,
                opencl_binary_file="",
Y
yejianwu 已提交
781
                shared_library_dir=get_shared_library_dir(library_name, target_abi),  # noqa
Y
yejianwu 已提交
782
                linkshared=linkshared,
783 784 785 786
            )

            pull_opencl_binary_and_tuning_param(target_abi, serial_num,
                                                [model_output_dir])
787
            sh_commands.touch_tuned_file_flag(build_tmp_binary_dir)
788 789 790
            binary_changed = True

    if binary_changed:
L
liuqi 已提交
791 792 793
        opencl_output_bin_path = get_opencl_binary_output_path(
            library_name, target_abi, target_soc, serial_num
        )
794 795
        sh_commands.merge_opencl_binaries(
            model_output_dirs, CL_COMPILED_BINARY_FILE_NAME,
L
liuqi 已提交
796
            opencl_output_bin_path)
797
        sh_commands.gen_tuning_param_code(model_output_dirs)
798
        sh_commands.bazel_build(
Y
yejianwu 已提交
799
            mace_run_target,
800 801 802 803 804
            abi=target_abi,
            hexagon_mode=hexagon_mode,
            enable_openmp=enable_openmp,
            address_sanitizer=address_sanitizer
        )
Y
yejianwu 已提交
805
        sh_commands.update_mace_run_lib(build_tmp_binary_dir, linkshared)
806 807 808 809 810 811 812

    if target_abi == ABIType.host:
        sh_commands.build_host_libraries(build_type, target_abi)

    # build benchmark_model binary
    sh_commands.build_benchmark_model(target_abi,
                                      build_tmp_binary_dir,
Y
yejianwu 已提交
813
                                      hexagon_mode,
L
liuqi 已提交
814
                                      enable_openmp,
Y
yejianwu 已提交
815
                                      linkshared)
816 817

    # generate library
Y
yejianwu 已提交
818
    if linkshared == 0:
Y
yejianwu 已提交
819 820 821 822 823 824 825 826
        sh_commands.merge_libs(target_soc,
                               serial_num,
                               target_abi,
                               library_name,
                               BUILD_OUTPUT_DIR,
                               OUTPUT_LIBRARY_DIR_NAME,
                               build_type,
                               hexagon_mode)
827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851


def generate_library(configs, tuning, enable_openmp, address_sanitizer):
    MaceLogger.header(StringFormatter.block("Building library"))
    # generate source
    MaceLogger.info('* generate common source files...')
    sh_commands.gen_mace_version()
    sh_commands.gen_encrypted_opencl_source()
    MaceLogger.info('generate common source files done')

    # create build dirs
    library_name = configs[YAMLKeyword.library_name]
    if not os.path.exists(BUILD_OUTPUT_DIR):
        os.makedirs(BUILD_OUTPUT_DIR)
    tmp_build_dir = os.path.join(BUILD_OUTPUT_DIR, library_name,
                                 BUILD_TMP_DIR_NAME)
    if not os.path.exists(tmp_build_dir):
        os.makedirs(tmp_build_dir)
    library_out_dir = os.path.join(BUILD_OUTPUT_DIR, library_name,
                                   OUTPUT_LIBRARY_DIR_NAME)
    if os.path.exists(library_out_dir):
        sh.rm('-rf', library_out_dir)

    target_socs = configs[YAMLKeyword.target_socs]
    for target_abi in configs[YAMLKeyword.target_abis]:
L
liuqi 已提交
852
        if not target_socs or target_abi == ABIType.host:
853 854 855
            build_specific_lib(target_abi, None, None, configs,
                               tuning, enable_openmp, address_sanitizer)
        else:
856 857
            if ALL_SOC_TAG in target_socs:
                target_socs = sh_commands.adb_get_all_socs()
858
            for target_soc in target_socs:
L
liuqi 已提交
859 860 861 862 863 864 865
                serial_nums = \
                    sh_commands.get_target_socs_serialnos([target_soc])
                for serial_num in serial_nums:
                    with sh_commands.device_lock(serial_num):
                        build_specific_lib(target_abi, target_soc, serial_num,
                                           configs, tuning, enable_openmp,
                                           address_sanitizer)
866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883

    # package library
    sh_commands.packaging_lib(BUILD_OUTPUT_DIR,
                              configs[YAMLKeyword.library_name])


def print_library_summary(configs):
    library_name = configs[YAMLKeyword.library_name]
    title = "Library"
    header = ["key", "value"]
    data = list()
    data.append(["library package",
                 "%s/%s/libmace_%s.tar.gz"
                 % (BUILD_OUTPUT_DIR, library_name, library_name)])
    MaceLogger.summary(StringFormatter.table(header, data, title))


def build_library(flags):
884
    configs = format_model_config(flags)
885 886 887 888 889

    print_configuration(flags, configs)

    convert_model(configs)

890 891
    generate_library(configs, flags.disable_tuning,
                     flags.disable_openmp, flags.address_sanitizer)
892 893 894 895 896 897 898 899 900 901 902 903

    print_library_summary(configs)


################################
# run
################################
def report_run_statistics(stdout,
                          abi,
                          serialno,
                          model_name,
                          device_type,
904 905
                          output_dir,
                          tuned):
906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926
    metrics = [0] * 3
    for line in stdout.split('\n'):
        line = line.strip()
        parts = line.split()
        if len(parts) == 4 and parts[0].startswith("time"):
            metrics[0] = str(float(parts[1]))
            metrics[1] = str(float(parts[2]))
            metrics[2] = str(float(parts[3]))
            break

    device_name = ""
    target_soc = ""
    if abi != "host":
        props = sh_commands.adb_getprop_by_serialno(serialno)
        device_name = props.get("ro.product.model", "")
        target_soc = props.get("ro.board.platform", "")

    report_filename = output_dir + "/report.csv"
    if not os.path.exists(report_filename):
        with open(report_filename, 'w') as f:
            f.write("model_name,device_name,soc,abi,runtime,"
927
                    "init,warmup,run_avg,tuned\n")
928 929

    data_str = "{model_name},{device_name},{soc},{abi},{device_type}," \
930
               "{init},{warmup},{run_avg},{tuned}\n" \
931 932 933 934 935 936 937
        .format(model_name=model_name,
                device_name=device_name,
                soc=target_soc,
                abi=abi,
                device_type=device_type,
                init=metrics[0],
                warmup=metrics[1],
938 939
                run_avg=metrics[2],
                tuned=tuned,
940 941 942 943 944 945 946 947 948 949
                )
    with open(report_filename, 'a') as f:
        f.write(data_str)


def run_specific_target(flags, configs, target_abi,
                        target_soc, serial_num):
    library_name = configs[YAMLKeyword.library_name]
    build_type = configs[YAMLKeyword.build_type]
    embed_model_data = configs[YAMLKeyword.embed_model_data]
L
liuqi 已提交
950
    opencl_output_bin_path = ""
Y
yejianwu 已提交
951
    linkshared = configs[YAMLKeyword.linkshared]
952 953 954 955 956 957
    if not configs[YAMLKeyword.target_socs]:
        build_tmp_binary_dir = get_build_binary_dir(library_name, target_abi,
                                                    None, None)
    else:
        build_tmp_binary_dir = get_build_binary_dir(library_name, target_abi,
                                                    target_soc, serial_num)
L
liuqi 已提交
958 959 960
        opencl_output_bin_path = get_opencl_binary_output_path(
            library_name, target_abi, target_soc, serial_num
        )
L
liuqi 已提交
961 962 963
    mace_check(os.path.exists(build_tmp_binary_dir),
               ModuleName.RUN,
               'You should build before run.')
964 965

    for model_name in configs[YAMLKeyword.models]:
L
liuqi 已提交
966 967 968 969 970 971 972 973
        if target_abi == ABIType.host:
            device_name = ABIType.host
        else:
            device_name =\
                sh_commands.adb_get_device_name_by_serialno(serial_num)
        MaceLogger.header(
            StringFormatter.block(
                "Run model %s on %s" % (model_name, device_name)))
974 975 976 977 978 979 980 981 982 983 984 985 986 987
        model_config = configs[YAMLKeyword.models][model_name]
        model_runtime = model_config[YAMLKeyword.runtime]
        subgraphs = model_config[YAMLKeyword.subgraphs]

        if not configs[YAMLKeyword.target_socs]:
            model_output_base_dir, model_output_dir, mace_model_dir = \
                get_build_model_dirs(library_name, model_name, target_abi,
                                     None, None,
                                     model_config[YAMLKeyword.model_file_path])
        else:
            model_output_base_dir, model_output_dir, mace_model_dir = \
                get_build_model_dirs(library_name, model_name, target_abi,
                                     target_soc, serial_num,
                                     model_config[YAMLKeyword.model_file_path])
L
liuqi 已提交
988 989 990 991
        mace_check(os.path.exists(model_output_dir)
                   and os.path.exists(mace_model_dir),
                   ModuleName.RUN,
                   'You should build before run.')
992 993 994 995 996 997 998 999
        if target_abi != ABIType.host:
            sh_commands.clear_phone_data_dir(serial_num, PHONE_DATA_DIR)

        # generate input data
        sh_commands.gen_random_input(
            model_output_dir,
            subgraphs[0][YAMLKeyword.input_tensors],
            subgraphs[0][YAMLKeyword.input_shapes],
李寅 已提交
1000
            subgraphs[0][YAMLKeyword.validation_inputs_data],
1001
            input_ranges=subgraphs[0][YAMLKeyword.input_ranges])
1002 1003 1004 1005 1006 1007 1008 1009 1010
        runtime_list = []
        if target_abi == ABIType.host:
            runtime_list.extend([RuntimeType.cpu])
        elif model_runtime == RuntimeType.cpu_gpu:
            runtime_list.extend([RuntimeType.cpu, RuntimeType.gpu])
        else:
            runtime_list.extend([model_runtime])
        for runtime in runtime_list:
            device_type = parse_device_type(runtime)
L
liuqi 已提交
1011

1012 1013 1014 1015 1016 1017 1018 1019 1020 1021 1022 1023 1024 1025 1026 1027 1028 1029
            run_output = sh_commands.tuning_run(
                abi=target_abi,
                serialno=serial_num,
                mace_run_dir=build_tmp_binary_dir,
                vlog_level=flags.vlog_level,
                embed_model_data=embed_model_data,
                model_output_dir=model_output_dir,
                input_nodes=subgraphs[0][YAMLKeyword.input_tensors],
                output_nodes=subgraphs[0][YAMLKeyword.output_tensors],
                input_shapes=subgraphs[0][YAMLKeyword.input_shapes],
                output_shapes=subgraphs[0][YAMLKeyword.output_shapes],
                mace_model_dir=mace_model_dir,
                model_tag=model_name,
                device_type=device_type,
                running_round=flags.round,
                restart_round=flags.restart_round,
                limit_opencl_kernel_time=model_config[YAMLKeyword.limit_opencl_kernel_time],  # noqa
                tuning=False,
L
Liangliang He 已提交
1030
                out_of_range_check=flags.gpu_out_of_range_check,
1031 1032 1033 1034 1035 1036 1037 1038
                phone_data_dir=PHONE_DATA_DIR,
                build_type=build_type,
                omp_num_threads=flags.omp_num_threads,
                cpu_affinity_policy=flags.cpu_affinity_policy,
                gpu_perf_hint=flags.gpu_perf_hint,
                gpu_priority_hint=flags.gpu_priority_hint,
                runtime_failure_ratio=flags.runtime_failure_ratio,
                address_sanitizer=flags.address_sanitizer,
L
liuqi 已提交
1039
                opencl_binary_file=opencl_output_bin_path,
Y
yejianwu 已提交
1040
                shared_library_dir=get_shared_library_dir(library_name, target_abi),  # noqa
Y
yejianwu 已提交
1041
                linkshared=linkshared,
1042 1043 1044
            )
            if flags.validate:
                model_file_path, weight_file_path = get_model_files_path(
L
liuqi 已提交
1045
                    model_config[YAMLKeyword.model_file_path],
1046
                    model_output_base_dir,
L
liuqi 已提交
1047
                    model_config[YAMLKeyword.weight_file_path])
1048 1049 1050 1051 1052 1053 1054 1055 1056 1057 1058 1059 1060 1061 1062 1063 1064 1065

                sh_commands.validate_model(
                    abi=target_abi,
                    serialno=serial_num,
                    model_file_path=model_file_path,
                    weight_file_path=weight_file_path,
                    platform=model_config[YAMLKeyword.platform],
                    device_type=device_type,
                    input_nodes=subgraphs[0][YAMLKeyword.input_tensors],
                    output_nodes=subgraphs[0][YAMLKeyword.output_tensors],
                    input_shapes=subgraphs[0][YAMLKeyword.input_shapes],
                    output_shapes=subgraphs[0][YAMLKeyword.output_shapes],
                    model_output_dir=model_output_dir,
                    phone_data_dir=PHONE_DATA_DIR,
                    caffe_env=flags.caffe_env)
            if flags.report and flags.round > 0:
                report_run_statistics(
                    run_output, target_abi, serial_num,
1066 1067
                    model_name, device_type, flags.report_dir,
                    sh_commands.is_binary_tuned(build_tmp_binary_dir))
1068 1069 1070


def run_mace(flags):
1071
    configs = format_model_config(flags)
1072 1073

    target_socs = configs[YAMLKeyword.target_socs]
1074
    if not target_socs or ALL_SOC_TAG in target_socs:
1075 1076 1077 1078 1079 1080 1081
        target_socs = sh_commands.adb_get_all_socs()

    for target_abi in configs[YAMLKeyword.target_abis]:
        if target_abi == ABIType.host:
            run_specific_target(flags, configs, target_abi, None, None)
        else:
            for target_soc in target_socs:
L
liuqi 已提交
1082 1083 1084 1085 1086 1087 1088 1089 1090
                serial_nums = \
                    sh_commands.get_target_socs_serialnos([target_soc])
                mace_check(serial_nums,
                           ModuleName.RUN,
                           'There is no device with soc: ' + target_soc)
                for serial_num in serial_nums:
                    with sh_commands.device_lock(serial_num):
                        run_specific_target(flags, configs, target_abi,
                                            target_soc, serial_num)
1091 1092 1093 1094 1095 1096 1097 1098 1099


################################
#  benchmark model
################################
def bm_specific_target(flags, configs, target_abi, target_soc, serial_num):
    library_name = configs[YAMLKeyword.library_name]
    build_type = configs[YAMLKeyword.build_type]
    embed_model_data = configs[YAMLKeyword.embed_model_data]
L
liuqi 已提交
1100
    opencl_output_bin_path = ""
Y
yejianwu 已提交
1101
    linkshared = configs[YAMLKeyword.linkshared]
1102 1103 1104 1105 1106 1107
    if not configs[YAMLKeyword.target_socs]:
        build_tmp_binary_dir = get_build_binary_dir(library_name, target_abi,
                                                    None, None)
    else:
        build_tmp_binary_dir = get_build_binary_dir(library_name, target_abi,
                                                    target_soc, serial_num)
L
liuqi 已提交
1108 1109 1110
        opencl_output_bin_path = get_opencl_binary_output_path(
            library_name, target_abi, target_soc, serial_num
        )
L
liuqi 已提交
1111 1112 1113
    mace_check(os.path.exists(build_tmp_binary_dir),
               ModuleName.BENCHMARK,
               'You should build before benchmark.')
1114 1115

    for model_name in configs[YAMLKeyword.models]:
L
liuqi 已提交
1116 1117 1118 1119 1120 1121 1122 1123
        if target_abi == ABIType.host:
            device_name = ABIType.host
        else:
            device_name = \
                sh_commands.adb_get_device_name_by_serialno(serial_num)
        MaceLogger.header(
            StringFormatter.block(
                "Benchmark model %s on %s" % (model_name, device_name)))
1124 1125 1126 1127 1128 1129 1130 1131 1132 1133 1134 1135 1136 1137
        model_config = configs[YAMLKeyword.models][model_name]
        model_runtime = model_config[YAMLKeyword.runtime]
        subgraphs = model_config[YAMLKeyword.subgraphs]

        if not configs[YAMLKeyword.target_socs]:
            model_output_base_dir, model_output_dir, mace_model_dir = \
                get_build_model_dirs(library_name, model_name, target_abi,
                                     None, None,
                                     model_config[YAMLKeyword.model_file_path])
        else:
            model_output_base_dir, model_output_dir, mace_model_dir = \
                get_build_model_dirs(library_name, model_name, target_abi,
                                     target_soc, serial_num,
                                     model_config[YAMLKeyword.model_file_path])
L
liuqi 已提交
1138 1139 1140 1141
        mace_check(os.path.exists(model_output_dir)
                   and os.path.exists(mace_model_dir),
                   ModuleName.BENCHMARK,
                   'You should build before benchmark.')
1142 1143 1144 1145 1146 1147 1148
        if target_abi != ABIType.host:
            sh_commands.clear_phone_data_dir(serial_num, PHONE_DATA_DIR)

        sh_commands.gen_random_input(
            model_output_dir,
            subgraphs[0][YAMLKeyword.input_tensors],
            subgraphs[0][YAMLKeyword.input_shapes],
李寅 已提交
1149
            subgraphs[0][YAMLKeyword.validation_inputs_data],
1150
            input_ranges=subgraphs[0][YAMLKeyword.input_ranges])
1151 1152 1153 1154 1155 1156 1157 1158 1159 1160 1161 1162 1163 1164 1165 1166 1167 1168 1169 1170 1171 1172 1173 1174 1175 1176 1177 1178
        runtime_list = []
        if target_abi == ABIType.host:
            runtime_list.extend([RuntimeType.cpu])
        elif model_runtime == RuntimeType.cpu_gpu:
            runtime_list.extend([RuntimeType.cpu, RuntimeType.gpu])
        else:
            runtime_list.extend([model_runtime])
        for runtime in runtime_list:
            device_type = parse_device_type(runtime)
            sh_commands.benchmark_model(
                abi=target_abi,
                serialno=serial_num,
                benchmark_binary_dir=build_tmp_binary_dir,
                vlog_level=0,
                embed_model_data=embed_model_data,
                model_output_dir=model_output_dir,
                input_nodes=subgraphs[0][YAMLKeyword.input_tensors],
                output_nodes=subgraphs[0][YAMLKeyword.output_tensors],
                input_shapes=subgraphs[0][YAMLKeyword.input_shapes],
                output_shapes=subgraphs[0][YAMLKeyword.output_shapes],
                mace_model_dir=mace_model_dir,
                model_tag=model_name,
                device_type=device_type,
                phone_data_dir=PHONE_DATA_DIR,
                build_type=build_type,
                omp_num_threads=flags.omp_num_threads,
                cpu_affinity_policy=flags.cpu_affinity_policy,
                gpu_perf_hint=flags.gpu_perf_hint,
1179
                gpu_priority_hint=flags.gpu_priority_hint,
Y
yejianwu 已提交
1180 1181
                opencl_binary_file=opencl_output_bin_path,
                shared_library_dir=get_shared_library_dir(library_name, target_abi),  # noqa
Y
yejianwu 已提交
1182
                linkshared=linkshared)
1183 1184 1185


def benchmark_model(flags):
1186
    configs = format_model_config(flags)
1187 1188

    target_socs = configs[YAMLKeyword.target_socs]
1189
    if not target_socs or ALL_SOC_TAG in target_socs:
1190 1191 1192
        target_socs = sh_commands.adb_get_all_socs()

    for target_abi in configs[YAMLKeyword.target_abis]:
L
liuqi 已提交
1193
        if target_abi == ABIType.host:
1194 1195 1196
            bm_specific_target(flags, configs, target_abi, None, None)
        else:
            for target_soc in target_socs:
L
liuqi 已提交
1197 1198 1199 1200 1201 1202 1203 1204 1205
                serial_nums = \
                    sh_commands.get_target_socs_serialnos([target_soc])
                mace_check(serial_nums,
                           ModuleName.BENCHMARK,
                           'There is no device with soc: ' + target_soc)
                for serial_num in serial_nums:
                    with sh_commands.device_lock(serial_num):
                        bm_specific_target(flags, configs, target_abi,
                                           target_soc, serial_num)
L
liuqi 已提交
1206

1207

L
liuqi 已提交
1208 1209 1210 1211 1212 1213 1214 1215 1216 1217 1218 1219 1220 1221
################################
# Parsing arguments
################################
def str2bool(v):
    if v.lower() in ('yes', 'true', 't', 'y', '1'):
        return True
    elif v.lower() in ('no', 'false', 'f', 'n', '0'):
        return False
    else:
        raise argparse.ArgumentTypeError('Boolean value expected.')


def str_to_caffe_env_type(v):
    if v.lower() == 'docker':
1222
        return CaffeEnvType.DOCKER
L
liuqi 已提交
1223
    elif v.lower() == 'local':
1224
        return CaffeEnvType.LOCAL
L
liuqi 已提交
1225 1226 1227 1228
    else:
        raise argparse.ArgumentTypeError('[docker | local] expected.')


1229
def parse_args():
L
Liangliang He 已提交
1230
    """Parses command line arguments."""
1231 1232 1233
    all_type_parent_parser = argparse.ArgumentParser(add_help=False)
    all_type_parent_parser.add_argument(
        '--config',
L
Liangliang He 已提交
1234
        type=str,
1235
        default="",
L
liuqi 已提交
1236
        required=True,
1237
        help="the path of model yaml configuration file.")
1238 1239 1240 1241 1242 1243 1244 1245 1246 1247
    all_type_parent_parser.add_argument(
        "--build_type",
        type=str,
        default="",
        help="Model build type, can be ['proto', 'code'].")
    all_type_parent_parser.add_argument(
        "--target_abis",
        type=str,
        default="",
        help="Target ABIs, comma seperated list.")
1248 1249 1250 1251 1252
    all_type_parent_parser.add_argument(
        "--target_socs",
        type=str,
        default="",
        help="Target SOCs, comma seperated list.")
1253 1254 1255 1256
    build_run_parent_parser = argparse.ArgumentParser(add_help=False)
    build_run_parent_parser.add_argument(
        '--address_sanitizer',
        action="store_true",
L
liuqi 已提交
1257
        help="Whether to use address sanitizer to check memory error")
1258 1259
    run_bm_parent_parser = argparse.ArgumentParser(add_help=False)
    run_bm_parent_parser.add_argument(
W
wuchenghui 已提交
1260 1261
        "--omp_num_threads",
        type=int,
1262
        default=DefaultValues.omp_num_threads,
W
wuchenghui 已提交
1263
        help="num of openmp threads")
1264
    run_bm_parent_parser.add_argument(
W
wuchenghui 已提交
1265 1266
        "--cpu_affinity_policy",
        type=int,
1267
        default=DefaultValues.cpu_affinity_policy,
W
wuchenghui 已提交
1268
        help="0:AFFINITY_NONE/1:AFFINITY_BIG_ONLY/2:AFFINITY_LITTLE_ONLY")
1269
    run_bm_parent_parser.add_argument(
W
wuchenghui 已提交
1270 1271
        "--gpu_perf_hint",
        type=int,
1272
        default=DefaultValues.gpu_perf_hint,
W
wuchenghui 已提交
1273
        help="0:DEFAULT/1:LOW/2:NORMAL/3:HIGH")
1274
    run_bm_parent_parser.add_argument(
W
wuchenghui 已提交
1275 1276
        "--gpu_priority_hint",
        type=int,
1277
        default=DefaultValues.gpu_priority_hint,
W
wuchenghui 已提交
1278
        help="0:DEFAULT/1:LOW/2:NORMAL/3:HIGH")
1279 1280 1281 1282 1283 1284 1285 1286 1287

    parser = argparse.ArgumentParser()
    subparsers = parser.add_subparsers()
    build = subparsers.add_parser(
        'build',
        parents=[all_type_parent_parser, build_run_parent_parser],
        help='build model library and test tools')
    build.set_defaults(func=build_library)
    build.add_argument(
1288 1289 1290
        '--disable_tuning',
        action="store_false",
        help="Disable tuning the parameters for the GPU of specified SoC.")
1291
    build.add_argument(
1292
        "--disable_openmp",
1293
        action="store_false",
1294
        help="Disable openmp for multiple thread.")
1295 1296 1297 1298 1299 1300 1301 1302
    run = subparsers.add_parser(
        'run',
        parents=[all_type_parent_parser, run_bm_parent_parser,
                 build_run_parent_parser],
        help='run model in command line')
    run.set_defaults(func=run_mace)
    run.add_argument(
        "--round",
L
Liangliang He 已提交
1303
        type=int,
1304 1305 1306 1307 1308
        default=1,
        help="The model running round.")
    run.add_argument(
        "--validate",
        action="store_true",
1309 1310
        help="whether to verify the results are consistent with "
             "the frameworks.")
1311
    run.add_argument(
L
liuqi 已提交
1312 1313 1314
        "--caffe_env",
        type=str_to_caffe_env_type,
        default='docker',
1315 1316
        help="[docker | local] you can specific caffe environment for"
             " validation. local environment or caffe docker image.")
1317 1318 1319 1320
    run.add_argument(
        "--vlog_level",
        type=int,
        default=0,
1321
        help="[1~5]. Verbose log level for debug.")
1322
    run.add_argument(
L
Liangliang He 已提交
1323
        "--gpu_out_of_range_check",
1324 1325 1326 1327 1328 1329
        action="store_true",
        help="Enable out of memory check for gpu.")
    run.add_argument(
        "--restart_round",
        type=int,
        default=1,
1330
        help="restart round between run.")
1331 1332 1333 1334 1335 1336
    run.add_argument(
        "--report",
        action="store_true",
        help="print run statistics report.")
    run.add_argument(
        "--report_dir",
1337 1338
        type=str,
        default="",
1339 1340
        help="print run statistics report.")
    run.add_argument(
李寅 已提交
1341 1342 1343 1344
        "--runtime_failure_ratio",
        type=float,
        default=0.0,
        help="[mock runtime failure ratio].")
1345 1346 1347 1348 1349 1350
    benchmark = subparsers.add_parser(
        'benchmark',
        parents=[all_type_parent_parser, run_bm_parent_parser,
                 build_run_parent_parser],
        help='benchmark model for detail information')
    benchmark.set_defaults(func=benchmark_model)
L
Liangliang He 已提交
1351 1352
    return parser.parse_known_args()

1353

Y
yejianwu 已提交
1354
if __name__ == "__main__":
1355 1356
    flags, unparsed = parse_args()
    flags.func(flags)