Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model
0
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 