Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
shared.py1040 linesDownload Raw Back to export
1# Copyright (c) Facebook, Inc. and its affiliates.2 3import collections4import copy5import functools6import logging7import numpy as np8import os9from typing import Any, Callable, Dict, List, Optional, Tuple, Union10from unittest import mock11import caffe2.python.utils as putils12import torch13import torch.nn.functional as F14from caffe2.proto import caffe2_pb215from caffe2.python import core, net_drawer, workspace16from torch.nn.functional import interpolate as interp17 18logger = logging.getLogger(__name__)19 20 21# ==== torch/utils_toffee/cast.py =======================================22 23 24def to_device(t, device_str):25    """26    This function is a replacement of .to(another_device) such that it allows the27    casting to be traced properly by explicitly calling the underlying copy ops.28    It also avoids introducing unncessary op when casting to the same device.29    """30    src = t.device31    dst = torch.device(device_str)32 33    if src == dst:34        return t35    elif src.type == "cuda" and dst.type == "cpu":36        return torch.ops._caffe2.CopyGPUToCPU(t)37    elif src.type == "cpu" and dst.type == "cuda":38        return torch.ops._caffe2.CopyCPUToGPU(t)39    else:40        raise RuntimeError("Can't cast tensor from device {} to device {}".format(src, dst))41 42 43# ==== torch/utils_toffee/interpolate.py =======================================44 45 46# Note: borrowed from vision/detection/fair/detectron/detectron/modeling/detector.py47def BilinearInterpolation(tensor_in, up_scale):48    assert up_scale % 2 == 0, "Scale should be even"49 50    def upsample_filt(size):51        factor = (size + 1) // 252        if size % 2 == 1:53            center = factor - 154        else:55            center = factor - 0.556 57        og = np.ogrid[:size, :size]58        return (1 - abs(og[0] - center) / factor) * (1 - abs(og[1] - center) / factor)59 60    kernel_size = int(up_scale) * 261    bil_filt = upsample_filt(kernel_size)62 63    dim = int(tensor_in.shape[1])64    kernel = np.zeros((dim, dim, kernel_size, kernel_size), dtype=np.float32)65    kernel[range(dim), range(dim), :, :] = bil_filt66 67    tensor_out = F.conv_transpose2d(68        tensor_in,69        weight=to_device(torch.Tensor(kernel), tensor_in.device),70        bias=None,71        stride=int(up_scale),72        padding=int(up_scale / 2),73    )74 75    return tensor_out76 77 78# NOTE: ONNX is incompatible with traced torch.nn.functional.interpolate if79# using dynamic `scale_factor` rather than static `size`. (T43166860)80# NOTE: Caffe2 Int8 conversion might not be able to quantize `size` properly.81def onnx_compatibale_interpolate(82    input, size=None, scale_factor=None, mode="nearest", align_corners=None83):84    # NOTE: The input dimensions are interpreted in the form:85    # `mini-batch x channels x [optional depth] x [optional height] x width`.86    if size is None and scale_factor is not None:87        if input.dim() == 4:88            if isinstance(scale_factor, (int, float)):89                height_scale, width_scale = (scale_factor, scale_factor)90            else:91                assert isinstance(scale_factor, (tuple, list))92                assert len(scale_factor) == 293                height_scale, width_scale = scale_factor94 95            assert not align_corners, "No matching C2 op for align_corners == True"96            if mode == "nearest":97                return torch.ops._caffe2.ResizeNearest(98                    input, order="NCHW", width_scale=width_scale, height_scale=height_scale99                )100            elif mode == "bilinear":101                logger.warning(102                    "Use F.conv_transpose2d for bilinear interpolate"103                    " because there's no such C2 op, this may cause significant"104                    " slowdown and the boundary pixels won't be as same as"105                    " using F.interpolate due to padding."106                )107                assert height_scale == width_scale108                return BilinearInterpolation(input, up_scale=height_scale)109        logger.warning("Output size is not static, it might cause ONNX conversion issue")110 111    return interp(input, size, scale_factor, mode, align_corners)112 113 114def mock_torch_nn_functional_interpolate():115    def decorator(func):116        @functools.wraps(func)117        def _mock_torch_nn_functional_interpolate(*args, **kwargs):118            if torch.onnx.is_in_onnx_export():119                with mock.patch(120                    "torch.nn.functional.interpolate", side_effect=onnx_compatibale_interpolate121                ):122                    return func(*args, **kwargs)123            else:124                return func(*args, **kwargs)125 126        return _mock_torch_nn_functional_interpolate127 128    return decorator129 130 131# ==== torch/utils_caffe2/ws_utils.py ==========================================132 133 134class ScopedWS(object):135    def __init__(self, ws_name, is_reset, is_cleanup=False):136        self.ws_name = ws_name137        self.is_reset = is_reset138        self.is_cleanup = is_cleanup139        self.org_ws = ""140 141    def __enter__(self):142        self.org_ws = workspace.CurrentWorkspace()143        if self.ws_name is not None:144            workspace.SwitchWorkspace(self.ws_name, True)145        if self.is_reset:146            workspace.ResetWorkspace()147 148        return workspace149 150    def __exit__(self, *args):151        if self.is_cleanup:152            workspace.ResetWorkspace()153        if self.ws_name is not None:154            workspace.SwitchWorkspace(self.org_ws)155 156 157def fetch_any_blob(name):158    bb = None159    try:160        bb = workspace.FetchBlob(name)161    except TypeError:162        bb = workspace.FetchInt8Blob(name)163    except Exception as e:164        logger.error("Get blob {} error: {}".format(name, e))165 166    return bb167 168 169# ==== torch/utils_caffe2/protobuf.py ==========================================170 171 172def get_pb_arg(pb, arg_name):173    for x in pb.arg:174        if x.name == arg_name:175            return x176    return None177 178 179def get_pb_arg_valf(pb, arg_name, default_val):180    arg = get_pb_arg(pb, arg_name)181    return arg.f if arg is not None else default_val182 183 184def get_pb_arg_floats(pb, arg_name, default_val):185    arg = get_pb_arg(pb, arg_name)186    return list(map(float, arg.floats)) if arg is not None else default_val187 188 189def get_pb_arg_ints(pb, arg_name, default_val):190    arg = get_pb_arg(pb, arg_name)191    return list(map(int, arg.ints)) if arg is not None else default_val192 193 194def get_pb_arg_vali(pb, arg_name, default_val):195    arg = get_pb_arg(pb, arg_name)196    return arg.i if arg is not None else default_val197 198 199def get_pb_arg_vals(pb, arg_name, default_val):200    arg = get_pb_arg(pb, arg_name)201    return arg.s if arg is not None else default_val202 203 204def get_pb_arg_valstrings(pb, arg_name, default_val):205    arg = get_pb_arg(pb, arg_name)206    return list(arg.strings) if arg is not None else default_val207 208 209def check_set_pb_arg(pb, arg_name, arg_attr, arg_value, allow_override=False):210    arg = get_pb_arg(pb, arg_name)211    if arg is None:212        arg = putils.MakeArgument(arg_name, arg_value)213        assert hasattr(arg, arg_attr)214        pb.arg.extend([arg])215    if allow_override and getattr(arg, arg_attr) != arg_value:216        logger.warning(217            "Override argument {}: {} -> {}".format(arg_name, getattr(arg, arg_attr), arg_value)218        )219        setattr(arg, arg_attr, arg_value)220    else:221        assert arg is not None222        assert getattr(arg, arg_attr) == arg_value, "Existing value {}, new value {}".format(223            getattr(arg, arg_attr), arg_value224        )225 226 227def _create_const_fill_op_from_numpy(name, tensor, device_option=None):228    assert type(tensor) == np.ndarray229    kTypeNameMapper = {230        np.dtype("float32"): "GivenTensorFill",231        np.dtype("int32"): "GivenTensorIntFill",232        np.dtype("int64"): "GivenTensorInt64Fill",233        np.dtype("uint8"): "GivenTensorStringFill",234    }235 236    args_dict = {}237    if tensor.dtype == np.dtype("uint8"):238        args_dict.update({"values": [str(tensor.data)], "shape": [1]})239    else:240        args_dict.update({"values": tensor, "shape": tensor.shape})241 242    if device_option is not None:243        args_dict["device_option"] = device_option244 245    return core.CreateOperator(kTypeNameMapper[tensor.dtype], [], [name], **args_dict)246 247 248def _create_const_fill_op_from_c2_int8_tensor(name, int8_tensor):249    assert type(int8_tensor) == workspace.Int8Tensor250    kTypeNameMapper = {251        np.dtype("int32"): "Int8GivenIntTensorFill",252        np.dtype("uint8"): "Int8GivenTensorFill",253    }254 255    tensor = int8_tensor.data256    assert tensor.dtype in [np.dtype("uint8"), np.dtype("int32")]257    values = tensor.tobytes() if tensor.dtype == np.dtype("uint8") else tensor258 259    return core.CreateOperator(260        kTypeNameMapper[tensor.dtype],261        [],262        [name],263        values=values,264        shape=tensor.shape,265        Y_scale=int8_tensor.scale,266        Y_zero_point=int8_tensor.zero_point,267    )268 269 270def create_const_fill_op(271    name: str,272    blob: Union[np.ndarray, workspace.Int8Tensor],273    device_option: Optional[caffe2_pb2.DeviceOption] = None,274) -> caffe2_pb2.OperatorDef:275    """276    Given a blob object, return the Caffe2 operator that creates this blob277    as constant. Currently support NumPy tensor and Caffe2 Int8Tensor.278    """279 280    tensor_type = type(blob)281    assert tensor_type in [282        np.ndarray,283        workspace.Int8Tensor,284    ], 'Error when creating const fill op for "{}", unsupported blob type: {}'.format(285        name, type(blob)286    )287 288    if tensor_type == np.ndarray:289        return _create_const_fill_op_from_numpy(name, blob, device_option)290    elif tensor_type == workspace.Int8Tensor:291        assert device_option is None292        return _create_const_fill_op_from_c2_int8_tensor(name, blob)293 294 295def construct_init_net_from_params(296    params: Dict[str, Any], device_options: Optional[Dict[str, caffe2_pb2.DeviceOption]] = None297) -> caffe2_pb2.NetDef:298    """299    Construct the init_net from params dictionary300    """301    init_net = caffe2_pb2.NetDef()302    device_options = device_options or {}303    for name, blob in params.items():304        if isinstance(blob, str):305            logger.warning(306                (307                    "Blob {} with type {} is not supported in generating init net,"308                    " skipped.".format(name, type(blob))309                )310            )311            continue312        init_net.op.extend(313            [create_const_fill_op(name, blob, device_option=device_options.get(name, None))]314        )315        init_net.external_output.append(name)316    return init_net317 318 319def get_producer_map(ssa):320    """321    Return dict from versioned blob to (i, j),322        where i is index of producer op, j is the index of output of that op.323    """324    producer_map = {}325    for i in range(len(ssa)):326        outputs = ssa[i][1]327        for j, outp in enumerate(outputs):328            producer_map[outp] = (i, j)329    return producer_map330 331 332def get_consumer_map(ssa):333    """334    Return dict from versioned blob to list of (i, j),335        where i is index of consumer op, j is the index of input of that op.336    """337    consumer_map = collections.defaultdict(list)338    for i in range(len(ssa)):339        inputs = ssa[i][0]340        for j, inp in enumerate(inputs):341            consumer_map[inp].append((i, j))342    return consumer_map343 344 345def get_params_from_init_net(346    init_net: caffe2_pb2.NetDef,347) -> [Dict[str, Any], Dict[str, caffe2_pb2.DeviceOption]]:348    """349    Take the output blobs from init_net by running it.350    Outputs:351        params: dict from blob name to numpy array352        device_options: dict from blob name to the device option of its creating op353    """354    # NOTE: this assumes that the params is determined by producer op with the355    # only exception be CopyGPUToCPU which is CUDA op but returns CPU tensor.356    def _get_device_option(producer_op):357        if producer_op.type == "CopyGPUToCPU":358            return caffe2_pb2.DeviceOption()359        else:360            return producer_op.device_option361 362    with ScopedWS("__get_params_from_init_net__", is_reset=True, is_cleanup=True) as ws:363        ws.RunNetOnce(init_net)364        params = {b: fetch_any_blob(b) for b in init_net.external_output}365    ssa, versions = core.get_ssa(init_net)366    producer_map = get_producer_map(ssa)367    device_options = {368        b: _get_device_option(init_net.op[producer_map[(b, versions[b])][0]])369        for b in init_net.external_output370    }371    return params, device_options372 373 374def _updater_raise(op, input_types, output_types):375    raise RuntimeError(376        "Failed to apply updater for op {} given input_types {} and"377        " output_types {}".format(op, input_types, output_types)378    )379 380 381def _generic_status_identifier(382    predict_net: caffe2_pb2.NetDef,383    status_updater: Callable,384    known_status: Dict[Tuple[str, int], Any],385) -> Dict[Tuple[str, int], Any]:386    """387    Statically infer the status of each blob, the status can be such as device type388        (CPU/GPU), layout (NCHW/NHWC), data type (float32/int8), etc. "Blob" here389        is versioned blob (Tuple[str, int]) in the format compatible with ssa.390    Inputs:391        predict_net: the caffe2 network392        status_updater: a callable, given an op and the status of its input/output,393            it returns the updated status of input/output. `None` is used for394            representing unknown status.395        known_status: a dict containing known status, used as initialization.396    Outputs:397        A dict mapping from versioned blob to its status398    """399    ssa, versions = core.get_ssa(predict_net)400    versioned_ext_input = [(b, 0) for b in predict_net.external_input]401    versioned_ext_output = [(b, versions[b]) for b in predict_net.external_output]402    all_versioned_blobs = set().union(*[set(x[0] + x[1]) for x in ssa])403 404    allowed_vbs = all_versioned_blobs.union(versioned_ext_input).union(versioned_ext_output)405    assert all(k in allowed_vbs for k in known_status)406    assert all(v is not None for v in known_status.values())407    _known_status = copy.deepcopy(known_status)408 409    def _check_and_update(key, value):410        assert value is not None411        if key in _known_status:412            if not _known_status[key] == value:413                raise RuntimeError(414                    "Confilict status for {}, existing status {}, new status {}".format(415                        key, _known_status[key], value416                    )417                )418        _known_status[key] = value419 420    def _update_i(op, ssa_i):421        versioned_inputs = ssa_i[0]422        versioned_outputs = ssa_i[1]423 424        inputs_status = [_known_status.get(b, None) for b in versioned_inputs]425        outputs_status = [_known_status.get(b, None) for b in versioned_outputs]426 427        new_inputs_status, new_outputs_status = status_updater(op, inputs_status, outputs_status)428 429        for versioned_blob, status in zip(430            versioned_inputs + versioned_outputs, new_inputs_status + new_outputs_status431        ):432            if status is not None:433                _check_and_update(versioned_blob, status)434 435    for op, ssa_i in zip(predict_net.op, ssa):436        _update_i(op, ssa_i)437    for op, ssa_i in zip(reversed(predict_net.op), reversed(ssa)):438        _update_i(op, ssa_i)439 440    # NOTE: This strictly checks all the blob from predict_net must be assgined441    # a known status. However sometimes it's impossible (eg. having deadend op),442    # we may relax this constraint if443    for k in all_versioned_blobs:444        if k not in _known_status:445            raise NotImplementedError(446                "Can not infer the status for {}. Currently only support the case where"447                " a single forward and backward pass can identify status for all blobs.".format(k)448            )449 450    return _known_status451 452 453def infer_device_type(454    predict_net: caffe2_pb2.NetDef,455    known_status: Dict[Tuple[str, int], Any],456    device_name_style: str = "caffe2",457) -> Dict[Tuple[str, int], str]:458    """Return the device type ("cpu" or "gpu"/"cuda") of each (versioned) blob"""459 460    assert device_name_style in ["caffe2", "pytorch"]461    _CPU_STR = "cpu"462    _GPU_STR = "gpu" if device_name_style == "caffe2" else "cuda"463 464    def _copy_cpu_to_gpu_updater(op, input_types, output_types):465        if input_types[0] == _GPU_STR or output_types[0] == _CPU_STR:466            _updater_raise(op, input_types, output_types)467        return ([_CPU_STR], [_GPU_STR])468 469    def _copy_gpu_to_cpu_updater(op, input_types, output_types):470        if input_types[0] == _CPU_STR or output_types[0] == _GPU_STR:471            _updater_raise(op, input_types, output_types)472        return ([_GPU_STR], [_CPU_STR])473 474    def _other_ops_updater(op, input_types, output_types):475        non_none_types = [x for x in input_types + output_types if x is not None]476        if len(non_none_types) > 0:477            the_type = non_none_types[0]478            if not all(x == the_type for x in non_none_types):479                _updater_raise(op, input_types, output_types)480        else:481            the_type = None482        return ([the_type for _ in op.input], [the_type for _ in op.output])483 484    def _device_updater(op, *args, **kwargs):485        return {486            "CopyCPUToGPU": _copy_cpu_to_gpu_updater,487            "CopyGPUToCPU": _copy_gpu_to_cpu_updater,488        }.get(op.type, _other_ops_updater)(op, *args, **kwargs)489 490    return _generic_status_identifier(predict_net, _device_updater, known_status)491 492 493# ==== torch/utils_caffe2/vis.py ===============================================494 495 496def _modify_blob_names(ops, blob_rename_f):497    ret = []498 499    def _replace_list(blob_list, replaced_list):500        del blob_list[:]501        blob_list.extend(replaced_list)502 503    for x in ops:504        cur = copy.deepcopy(x)505        _replace_list(cur.input, list(map(blob_rename_f, cur.input)))506        _replace_list(cur.output, list(map(blob_rename_f, cur.output)))507        ret.append(cur)508 509    return ret510 511 512def _rename_blob(name, blob_sizes, blob_ranges):513    def _list_to_str(bsize):514        ret = ", ".join([str(x) for x in bsize])515        ret = "[" + ret + "]"516        return ret517 518    ret = name519    if blob_sizes is not None and name in blob_sizes:520        ret += "\n" + _list_to_str(blob_sizes[name])521    if blob_ranges is not None and name in blob_ranges:522        ret += "\n" + _list_to_str(blob_ranges[name])523 524    return ret525 526 527# graph_name could not contain word 'graph'528def save_graph(net, file_name, graph_name="net", op_only=True, blob_sizes=None, blob_ranges=None):529    blob_rename_f = functools.partial(_rename_blob, blob_sizes=blob_sizes, blob_ranges=blob_ranges)530    return save_graph_base(net, file_name, graph_name, op_only, blob_rename_f)531 532 533def save_graph_base(net, file_name, graph_name="net", op_only=True, blob_rename_func=None):534    graph = None535    ops = net.op536    if blob_rename_func is not None:537        ops = _modify_blob_names(ops, blob_rename_func)538    if not op_only:539        graph = net_drawer.GetPydotGraph(ops, graph_name, rankdir="TB")540    else:541        graph = net_drawer.GetPydotGraphMinimal(542            ops, graph_name, rankdir="TB", minimal_dependency=True543        )544 545    try:546        par_dir = os.path.dirname(file_name)547        if not os.path.exists(par_dir):548            os.makedirs(par_dir)549 550        format = os.path.splitext(os.path.basename(file_name))[-1]551        if format == ".png":552            graph.write_png(file_name)553        elif format == ".pdf":554            graph.write_pdf(file_name)555        elif format == ".svg":556            graph.write_svg(file_name)557        else:558            print("Incorrect format {}".format(format))559    except Exception as e:560        print("Error when writing graph to image {}".format(e))561 562    return graph563 564 565# ==== torch/utils_toffee/aten_to_caffe2.py ====================================566 567 568def group_norm_replace_aten_with_caffe2(predict_net: caffe2_pb2.NetDef):569    """570    For ONNX exported model, GroupNorm will be represented as ATen op,571        this can be a drop in replacement from ATen to GroupNorm572    """573    count = 0574    for op in predict_net.op:575        if op.type == "ATen":576            op_name = get_pb_arg_vals(op, "operator", None)  # return byte in py3577            if op_name and op_name.decode() == "group_norm":578                op.arg.remove(get_pb_arg(op, "operator"))579 580                if get_pb_arg_vali(op, "cudnn_enabled", None):581                    op.arg.remove(get_pb_arg(op, "cudnn_enabled"))582 583                num_groups = get_pb_arg_vali(op, "num_groups", None)584                if num_groups is not None:585                    op.arg.remove(get_pb_arg(op, "num_groups"))586                    check_set_pb_arg(op, "group", "i", num_groups)587 588                op.type = "GroupNorm"589                count += 1590    if count > 1:591        logger.info("Replaced {} ATen operator to GroupNormOp".format(count))592 593 594# ==== torch/utils_toffee/alias.py =============================================595 596 597def alias(x, name, is_backward=False):598    if not torch.onnx.is_in_onnx_export():599        return x600    assert isinstance(x, torch.Tensor)601    return torch.ops._caffe2.AliasWithName(x, name, is_backward=is_backward)602 603 604def fuse_alias_placeholder(predict_net, init_net):605    """Remove AliasWithName placeholder and rename the input/output of it"""606    # First we finish all the re-naming607    for i, op in enumerate(predict_net.op):608        if op.type == "AliasWithName":609            assert len(op.input) == 1610            assert len(op.output) == 1611            name = get_pb_arg_vals(op, "name", None).decode()612            is_backward = bool(get_pb_arg_vali(op, "is_backward", 0))613            rename_op_input(predict_net, init_net, i, 0, name, from_producer=is_backward)614            rename_op_output(predict_net, i, 0, name)615 616    # Remove AliasWithName, should be very safe since it's a non-op617    new_ops = []618    for op in predict_net.op:619        if op.type != "AliasWithName":620            new_ops.append(op)621        else:622            # safety check623            assert op.input == op.output624            assert op.input[0] == op.arg[0].s.decode()625    del predict_net.op[:]626    predict_net.op.extend(new_ops)627 628 629# ==== torch/utils_caffe2/graph_transform.py ===================================630 631 632class IllegalGraphTransformError(ValueError):633    """When a graph transform function call can't be executed."""634 635 636def _rename_versioned_blob_in_proto(637    proto: caffe2_pb2.NetDef,638    old_name: str,639    new_name: str,640    version: int,641    ssa: List[Tuple[List[Tuple[str, int]], List[Tuple[str, int]]]],642    start_versions: Dict[str, int],643    end_versions: Dict[str, int],644):645    """In given proto, rename all blobs with matched version"""646    # Operater list647    for op, i_th_ssa in zip(proto.op, ssa):648        versioned_inputs, versioned_outputs = i_th_ssa649        for i in range(len(op.input)):650            if versioned_inputs[i] == (old_name, version):651                op.input[i] = new_name652        for i in range(len(op.output)):653            if versioned_outputs[i] == (old_name, version):654                op.output[i] = new_name655    # external_input656    if start_versions.get(old_name, 0) == version:657        for i in range(len(proto.external_input)):658            if proto.external_input[i] == old_name:659                proto.external_input[i] = new_name660    # external_output661    if end_versions.get(old_name, 0) == version:662        for i in range(len(proto.external_output)):663            if proto.external_output[i] == old_name:664                proto.external_output[i] = new_name665 666 667def rename_op_input(668    predict_net: caffe2_pb2.NetDef,669    init_net: caffe2_pb2.NetDef,670    op_id: int,671    input_id: int,672    new_name: str,673    from_producer: bool = False,674):675    """676    Rename the op_id-th operator in predict_net, change it's input_id-th input's677        name to the new_name. It also does automatic re-route and change678        external_input and init_net if necessary.679    - It requires the input is only consumed by this op.680    - This function modifies predict_net and init_net in-place.681    - When from_producer is enable, this also updates other operators that consumes682        the same input. Be cautious because may trigger unintended behavior.683    """684    assert isinstance(predict_net, caffe2_pb2.NetDef)685    assert isinstance(init_net, caffe2_pb2.NetDef)686 687    init_net_ssa, init_net_versions = core.get_ssa(init_net)688    predict_net_ssa, predict_net_versions = core.get_ssa(689        predict_net, copy.deepcopy(init_net_versions)690    )691 692    versioned_inputs, versioned_outputs = predict_net_ssa[op_id]693    old_name, version = versioned_inputs[input_id]694 695    if from_producer:696        producer_map = get_producer_map(predict_net_ssa)697        if not (old_name, version) in producer_map:698            raise NotImplementedError(699                "Can't find producer, the input {} is probably from"700                " init_net, this is not supported yet.".format(old_name)701            )702        producer = producer_map[(old_name, version)]703        rename_op_output(predict_net, producer[0], producer[1], new_name)704        return705 706    def contain_targets(op_ssa):707        return (old_name, version) in op_ssa[0]708 709    is_consumer = [contain_targets(op_ssa) for op_ssa in predict_net_ssa]710    if sum(is_consumer) > 1:711        raise IllegalGraphTransformError(712            (713                "Input '{}' of operator(#{}) are consumed by other ops, please use"714                + " rename_op_output on the producer instead. Offending op: \n{}"715            ).format(old_name, op_id, predict_net.op[op_id])716        )717 718    # update init_net719    _rename_versioned_blob_in_proto(720        init_net, old_name, new_name, version, init_net_ssa, {}, init_net_versions721    )722    # update predict_net723    _rename_versioned_blob_in_proto(724        predict_net,725        old_name,726        new_name,727        version,728        predict_net_ssa,729        init_net_versions,730        predict_net_versions,731    )732 733 734def rename_op_output(predict_net: caffe2_pb2.NetDef, op_id: int, output_id: int, new_name: str):735    """736    Rename the op_id-th operator in predict_net, change it's output_id-th input's737        name to the new_name. It also does automatic re-route and change738        external_output and if necessary.739    - It allows multiple consumers of its output.740    - This function modifies predict_net in-place, doesn't need init_net.741    """742    assert isinstance(predict_net, caffe2_pb2.NetDef)743 744    ssa, blob_versions = core.get_ssa(predict_net)745 746    versioned_inputs, versioned_outputs = ssa[op_id]747    old_name, version = versioned_outputs[output_id]748 749    # update predict_net750    _rename_versioned_blob_in_proto(751        predict_net, old_name, new_name, version, ssa, {}, blob_versions752    )753 754 755def get_sub_graph_external_input_output(756    predict_net: caffe2_pb2.NetDef, sub_graph_op_indices: List[int]757) -> Tuple[List[Tuple[str, int]], List[Tuple[str, int]]]:758    """759    Return the list of external input/output of sub-graph,760    each element is tuple of the name and corresponding version in predict_net.761 762    external input/output is defined the same way as caffe2 NetDef.763    """764    ssa, versions = core.get_ssa(predict_net)765 766    all_inputs = []767    all_outputs = []768    for op_id in sub_graph_op_indices:769        all_inputs += [inp for inp in ssa[op_id][0] if inp not in all_inputs]770        all_outputs += list(ssa[op_id][1])  # ssa output won't repeat771 772    # for versioned blobs, external inputs are just those blob in all_inputs773    # but not in all_outputs774    ext_inputs = [inp for inp in all_inputs if inp not in all_outputs]775 776    # external outputs are essentially outputs of this subgraph that are used777    # outside of this sub-graph (including predict_net.external_output)778    all_other_inputs = sum(779        (ssa[i][0] for i in range(len(ssa)) if i not in sub_graph_op_indices),780        [(outp, versions[outp]) for outp in predict_net.external_output],781    )782    ext_outputs = [outp for outp in all_outputs if outp in set(all_other_inputs)]783 784    return ext_inputs, ext_outputs785 786 787class DiGraph:788    """A DAG representation of caffe2 graph, each vertice is a versioned blob."""789 790    def __init__(self):791        self.vertices = set()792        self.graph = collections.defaultdict(list)793 794    def add_edge(self, u, v):795        self.graph[u].append(v)796        self.vertices.add(u)797        self.vertices.add(v)798 799    # grab from https://www.geeksforgeeks.org/find-paths-given-source-destination/800    def get_all_paths(self, s, d):801        visited = {k: False for k in self.vertices}802        path = []803        all_paths = []804 805        def _get_all_paths_util(graph, u, d, visited, path):806            visited[u] = True807            path.append(u)808            if u == d:809                all_paths.append(copy.deepcopy(path))810            else:811                for i in graph[u]:812                    if not visited[i]:813                        _get_all_paths_util(graph, i, d, visited, path)814            path.pop()815            visited[u] = False816 817        _get_all_paths_util(self.graph, s, d, visited, path)818        return all_paths819 820    @staticmethod821    def from_ssa(ssa):822        graph = DiGraph()823        for op_id in range(len(ssa)):824            for inp in ssa[op_id][0]:825                for outp in ssa[op_id][1]:826                    graph.add_edge(inp, outp)827        return graph828 829 830def _get_dependency_chain(ssa, versioned_target, versioned_source):831    """832    Return the index list of relevant operator to produce target blob from source blob,833        if there's no dependency, return empty list.834    """835 836    # finding all paths between nodes can be O(N!), thus we can only search837    # in the subgraph using the op starting from the first consumer of source blob838    # to the producer of the target blob.839    consumer_map = get_consumer_map(ssa)840    producer_map = get_producer_map(ssa)841    start_op = min(x[0] for x in consumer_map[versioned_source]) - 15842    end_op = (843        producer_map[versioned_target][0] + 15 if versioned_target in producer_map else start_op844    )845    sub_graph_ssa = ssa[start_op : end_op + 1]846    if len(sub_graph_ssa) > 30:847        logger.warning(848            "Subgraph bebetween {} and {} is large (from op#{} to op#{}), it"849            " might take non-trival time to find all paths between them.".format(850                versioned_source, versioned_target, start_op, end_op851            )852        )853 854    dag = DiGraph.from_ssa(sub_graph_ssa)855    paths = dag.get_all_paths(versioned_source, versioned_target)  # include two ends856    ops_in_paths = [[producer_map[blob][0] for blob in path[1:]] for path in paths]857    return sorted(set().union(*[set(ops) for ops in ops_in_paths]))858 859 860def identify_reshape_sub_graph(predict_net: caffe2_pb2.NetDef) -> List[List[int]]:861    """862    Idenfity the reshape sub-graph in a protobuf.863    The reshape sub-graph is defined as matching the following pattern:864 865    (input_blob) -> Op_1 -> ... -> Op_N -> (new_shape) -โ”€โ”866        โ””-------------------------------------------> Reshape -> (output_blob)867 868    Return:869        List of sub-graphs, each sub-graph is represented as a list of indices870        of the relavent ops, [Op_1, Op_2, ..., Op_N, Reshape]871    """872 873    ssa, _ = core.get_ssa(predict_net)874 875    ret = []876    for i, op in enumerate(predict_net.op):877        if op.type == "Reshape":878            assert len(op.input) == 2879            input_ssa = ssa[i][0]880            data_source = input_ssa[0]881            shape_source = input_ssa[1]882            op_indices = _get_dependency_chain(ssa, shape_source, data_source)883            ret.append(op_indices + [i])884    return ret885 886 887def remove_reshape_for_fc(predict_net, params):888    """889    In PyTorch nn.Linear has to take 2D tensor, this often leads to reshape890        a 4D tensor to 2D by calling .view(). However this (dynamic) reshaping891        doesn't work well with ONNX and Int8 tools, and cause using extra892        ops (eg. ExpandDims) that might not be available on mobile.893    Luckily Caffe2 supports 4D tensor for FC, so we can remove those reshape894        after exporting ONNX model.895    """896    from caffe2.python import core897 898    # find all reshape sub-graph that can be removed, which is now all Reshape899    # sub-graph whose output is only consumed by FC.900    # TODO: to make it safer, we may need the actually value to better determine901    # if a Reshape before FC is removable.902    reshape_sub_graphs = identify_reshape_sub_graph(predict_net)903    sub_graphs_to_remove = []904    for reshape_sub_graph in reshape_sub_graphs:905        reshape_op_id = reshape_sub_graph[-1]906        assert predict_net.op[reshape_op_id].type == "Reshape"907        ssa, _ = core.get_ssa(predict_net)908        reshape_output = ssa[reshape_op_id][1][0]909        consumers = [i for i in range(len(ssa)) if reshape_output in ssa[i][0]]910        if all(predict_net.op[consumer].type == "FC" for consumer in consumers):911            # safety check if the sub-graph is isolated, for this reshape sub-graph,912            # it means it has one non-param external input and one external output.913            ext_inputs, ext_outputs = get_sub_graph_external_input_output(914                predict_net, reshape_sub_graph915            )916            non_params_ext_inputs = [inp for inp in ext_inputs if inp[1] != 0]917            if len(non_params_ext_inputs) == 1 and len(ext_outputs) == 1:918                sub_graphs_to_remove.append(reshape_sub_graph)919 920    # perform removing subgraph by:921    # 1: rename the Reshape's output to its input, then the graph can be922    #   seen as in-place itentify, meaning whose external input/output are the same.923    # 2: simply remove those ops.924    remove_op_ids = []925    params_to_remove = []926    for sub_graph in sub_graphs_to_remove:927        logger.info(928            "Remove Reshape sub-graph:\n{}".format(929                "".join(["(#{:>4})\n{}".format(i, predict_net.op[i]) for i in sub_graph])930            )931        )932        reshape_op_id = sub_graph[-1]933        new_reshap_output = predict_net.op[reshape_op_id].input[0]934        rename_op_output(predict_net, reshape_op_id, 0, new_reshap_output)935        ext_inputs, ext_outputs = get_sub_graph_external_input_output(predict_net, sub_graph)936        non_params_ext_inputs = [inp for inp in ext_inputs if inp[1] != 0]937        params_ext_inputs = [inp for inp in ext_inputs if inp[1] == 0]938        assert len(non_params_ext_inputs) == 1 and len(ext_outputs) == 1939        assert ext_outputs[0][0] == non_params_ext_inputs[0][0]940        assert ext_outputs[0][1] == non_params_ext_inputs[0][1] + 1941        remove_op_ids.extend(sub_graph)942        params_to_remove.extend(params_ext_inputs)943 944    predict_net = copy.deepcopy(predict_net)945    new_ops = [op for i, op in enumerate(predict_net.op) if i not in remove_op_ids]946    del predict_net.op[:]947    predict_net.op.extend(new_ops)948    for versioned_params in params_to_remove:949        name = versioned_params[0]950        logger.info("Remove params: {} from init_net and predict_net.external_input".format(name))951        del params[name]952        predict_net.external_input.remove(name)953 954    return predict_net, params955 956 957def fuse_copy_between_cpu_and_gpu(predict_net: caffe2_pb2.NetDef):958    """959    In-place fuse extra copy ops between cpu/gpu for the following case:960        a -CopyAToB-> b -CopyBToA> c1 -NextOp1-> d1961                        -CopyBToA> c2 -NextOp2-> d2962    The fused network will look like:963        a -NextOp1-> d1964          -NextOp2-> d2965    """966 967    _COPY_OPS = ["CopyCPUToGPU", "CopyGPUToCPU"]968 969    def _fuse_once(predict_net):970        ssa, blob_versions = core.get_ssa(predict_net)971        consumer_map = get_consumer_map(ssa)972        versioned_external_output = [973            (name, blob_versions[name]) for name in predict_net.external_output974        ]975 976        for op_id, op in enumerate(predict_net.op):977            if op.type in _COPY_OPS:978                fw_copy_versioned_output = ssa[op_id][1][0]979                consumer_ids = [x[0] for x in consumer_map[fw_copy_versioned_output]]980                reverse_op_type = _COPY_OPS[1 - _COPY_OPS.index(op.type)]981 982                is_fusable = (983                    len(consumer_ids) > 0984                    and fw_copy_versioned_output not in versioned_external_output985                    and all(986                        predict_net.op[_op_id].type == reverse_op_type987                        and ssa[_op_id][1][0] not in versioned_external_output988                        for _op_id in consumer_ids989                    )990                )991 992                if is_fusable:993                    for rv_copy_op_id in consumer_ids:994                        # making each NextOp uses "a" directly and removing Copy ops995                        rs_copy_versioned_output = ssa[rv_copy_op_id][1][0]996                        next_op_id, inp_id = consumer_map[rs_copy_versioned_output][0]997                        predict_net.op[next_op_id].input[inp_id] = op.input[0]998                    # remove CopyOps999                    new_ops = [1000                        op1001                        for i, op in enumerate(predict_net.op)1002                        if i != op_id and i not in consumer_ids1003                    ]1004                    del predict_net.op[:]1005                    predict_net.op.extend(new_ops)1006                    return True1007 1008        return False1009 1010    # _fuse_once returns False is nothing can be fused1011    while _fuse_once(predict_net):1012        pass1013 1014 1015def remove_dead_end_ops(net_def: caffe2_pb2.NetDef):1016    """remove ops if its output is not used or not in external_output"""1017    ssa, versions = core.get_ssa(net_def)1018    versioned_external_output = [(name, versions[name]) for name in net_def.external_output]1019    consumer_map = get_consumer_map(ssa)1020    removed_op_ids = set()1021 1022    def _is_dead_end(versioned_blob):1023        return not (1024            versioned_blob in versioned_external_output1025            or (1026                len(consumer_map[versioned_blob]) > 01027                and all(x[0] not in removed_op_ids for x in consumer_map[versioned_blob])1028            )1029        )1030 1031    for i, ssa_i in reversed(list(enumerate(ssa))):1032        versioned_outputs = ssa_i[1]1033        if all(_is_dead_end(outp) for outp in versioned_outputs):1034            removed_op_ids.add(i)1035 1036    # simply removing those deadend ops should have no effect to external_output1037    new_ops = [op for i, op in enumerate(net_def.op) if i not in removed_op_ids]1038    del net_def.op[:]1039    net_def.op.extend(new_ops)1040