Skip to content
This repository was archived by the owner on Nov 17, 2023. It is now read-only.
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 28 additions & 0 deletions include/mxnet/c_api.h
Original file line number Diff line number Diff line change
Expand Up @@ -1623,6 +1623,34 @@ MXNET_DLL int MXQuantizeSymbol(SymbolHandle sym_handle, SymbolHandle *ret_sym_ha
const mx_uint num_offline, const char **offline_params,
const char *quantized_dtype, const bool calib_quantize);

/*!
* \brief Convert a symbol into a mixed precision symbol with cast operators for target dtype casting
* \param sym_handle symbol to be converted
* \param ret_sym_handle mixed precision symbol result
* \param target_dtype target_dtype for the mixed precision model
* \param num_target_dtype_op_names number of ops to be casted to target_dtype
* \param num_fp32_op_names number of ops to be casted to FP32
* \param num_widest_dtype_op_names number of ops to be casted to widest dtype
* \param num_conditional_fp32_op_names number of ops to be cast to fp32 based on condition
* \param target_dtype_op_names op names to be casted to target_dtype
* \param fp32_op_names op names to be casted to FP32
* \param widest_dtype_op_names names to be casted to widest dtype
* \param conditional_fp32_op_names names to be casted to FP32 conditionally
*/
MXNET_DLL int MXReducePrecisionSymbol(SymbolHandle sym_handle,
SymbolHandle *ret_sym_handle,
const int* target_dtype,
const mx_uint num_target_dtype_op_names,
const mx_uint num_fp32_op_names,
const mx_uint num_widest_dtype_op_names,
const mx_uint num_conditional_fp32_op_names,
const mx_uint num_excluded_symbols,
const char **target_dtype_op_names,
const char **fp32_op_names,
const char **widest_dtype_op_names,
const char **conditional_fp32_op_names,
const char **excluded_symbols);

/*!
* \brief Set calibration table to node attributes in the sym
* \param sym_handle symbol whose node attributes are to be set by calibration table
Expand Down
1 change: 1 addition & 0 deletions python/mxnet/contrib/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,3 +33,4 @@
from . import quantization
from . import quantization as quant
from . import tensorrt
from . import amp
174 changes: 174 additions & 0 deletions python/mxnet/contrib/amp.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,174 @@
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you 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.
"""Contrib AMP API to demonstrate conversion to mixed precision model"""

import ctypes
import numpy as np
from ..base import _LIB, check_call
from ..base import mx_uint, c_str_array
from ..base import SymbolHandle
from ..symbol import Symbol
from ..ndarray import _DTYPE_NP_TO_MX


def _convert_symbol(sym, target_dtype="float16", target_dtype_ops=None,
fp32_ops=None, widest_dtype_ops=None, conditional_fp32_ops=None,
excluded_sym_names=None):
"""Given a symbol object representing a neural network of data type FP32 and target_dtype,
add cast layers according to the op lists (target_precision_ops, fp32_ops,
widest_precision_ops, conditional_fp32_ops) if provided, otherwise use the default
lists provided by the framework.

Parameters
----------
sym : Symbol
FP32 neural network symbol
target_dtype : str or numpy
currently only supports float16. The target dtype indicates to add cast layers
when possible so that lower precision computation can be leveraged.
target_precision_ops : list of strs
Override the list of operator names casted to target_dtype.
If None, uses the framework's default list to be casted to target dtype.
fp32_ops : list of strs
Override the lists of operator names casted to FP32.
If None, uses the framework's default list to be casted to FP32.
widest_precision_ops : list of strs
Override the list of operator names which should run in widest precision among its
input arguments.
If None, uses the framework's default list of widest_precision_ops.
conditional_fp32_ops : list of (string, string, list of string)
Override the list of functions casted to FP32.
The format of the list is
(name of the function, name of the parameter,
list of values of the parameter that make the operator to be casted to
fp32)
excluded_sym_names : list of strs
A list of strings that represent the names of symbols that users want to exclude
from being quantized.
"""
if target_dtype != "float16":
raise ValueError("Only target_dtype float16 is supported currently")
num_target_dtype_ops = 0
num_fp32_ops = 0
num_widest_dtype_ops = 0
num_conditional_fp32_ops = 0
num_excluded_syms = 0

if target_dtype_ops is not None:
assert isinstance(target_dtype_ops, list)
num_target_dtype_ops = len(target_dtype_ops)
else:
target_dtype_ops = []

if fp32_ops is not None:
assert isinstance(fp32_ops, list)
num_fp32_ops = len(fp32_ops)
else:
fp32_ops = []

if widest_dtype_ops is not None:
assert isinstance(widest_dtype_ops, list)
num_widest_dtype_ops = len(widest_dtype_ops)
else:
widest_dtype_ops = []

if conditional_fp32_ops is not None:
assert isinstance(conditional_fp32_ops, list)
num_conditional_fp32_ops = len(conditional_fp32_ops)
else:
conditional_fp32_ops = []

if excluded_sym_names is not None:
assert isinstance(excluded_sym_names, list)
num_excluded_syms = len(excluded_sym_names)
else:
excluded_sym_names = []

target_dtype = _DTYPE_NP_TO_MX[np.dtype(target_dtype).type]

out = SymbolHandle()
# currently this passes str for conditional_fp32_ops for PoC, this will change
check_call(_LIB.MXReducePrecisionSymbol(sym.handle,
ctypes.byref(out),
ctypes.byref(ctypes.c_int(target_dtype)),
mx_uint(num_target_dtype_ops),
mx_uint(num_fp32_ops),
mx_uint(num_widest_dtype_ops),
mx_uint(num_conditional_fp32_ops),
mx_uint(num_excluded_syms),
c_str_array(target_dtype_ops),
c_str_array(fp32_ops),
c_str_array(widest_dtype_ops),
c_str_array(conditional_fp32_ops),
c_str_array(excluded_sym_names)))
return Symbol(out)


def convert_model(sym, arg_params, aux_params, target_dtype="float16", target_precision_ops=None,
fp32_ops=None, widest_precision_ops=None,
conditional_fp32_ops=None, excluded_sym_names=None):
"""API for converting a model from FP32 model to a mixed precision model.
MXNet tries to convert the FP32 model to mixed precision model by adding
cast layers using amp_cast and amp_multicast operators. The decision on
which cast layer to add is based on hardcoded lists for Automatic Mixed Precision
in MXNet. These lists can be overridden by the user by providing their own lists
using : targe_precision_ops, fp32_ops, widest_precision_ops, conditional_fp32_ops

Parameters
----------
sym : str or Symbol
Defines the structure of a neural network for FP32 types.
arg_params : dict
Dictionary of name to `NDArray`.
aux_params : dict
Dictionary of name to `NDArray`.
target_dtype : str
Currently only supports float16. The target dtype indicates to add cast layers
when possible so that lower precision computation can be leveraged.
target_precision_ops : list of strs
Override the list of operator names casted to target_dtype.
If None, uses the framework's default list to be casted to target dtype.
fp32_ops : list of strs
Override the lists of operator names casted to FP32.
If None, uses the framework's default list to be casted to FP32.
widest_precision_ops : list of strs
A list of op names provided by user which should run in widest precision among its inputs.
If None, uses the framework's default list of widest_precision_ops.
conditional_fp32_ops : list of (string, string, list of string)
Override the list of operators to be casted to FP32.
The format of the list is
(name of the function, name of the parameter,
list of values of the parameter that make the operator to be casted to
fp32)
excluded_sym_names : list of strs
A list of strings that represent the names of symbols that users want to exclude
from being quantized.
"""
if excluded_sym_names is None:
excluded_sym_names = []
if not isinstance(excluded_sym_names, list):
raise ValueError('excluded_sym_names must be a list of strings representing'
' the names of the symbols that should not be casted,'
' while received type %s' % str(type(excluded_sym_names)))

if target_dtype != "float16":
raise ValueError("Only target_dtype float16 is supported currently")

sym = _convert_symbol(sym, target_dtype, target_precision_ops,
fp32_ops, widest_precision_ops, conditional_fp32_ops,
excluded_sym_names)
return sym, arg_params, aux_params
4 changes: 4 additions & 0 deletions python/mxnet/visualization.py
Original file line number Diff line number Diff line change
Expand Up @@ -369,6 +369,10 @@ def looks_like_weight(name):
attr["fillcolor"] = cm[5]
elif op == "Softmax":
attr["fillcolor"] = cm[6]
elif op == "amp_multicast":
label = "amp_multicast"
elif op == "amp_cast":
label = "amp_cast"
else:
attr["fillcolor"] = cm[7]
if op == "Custom":
Expand Down
55 changes: 53 additions & 2 deletions src/c_api/c_api_symbolic.cc
Original file line number Diff line number Diff line change
Expand Up @@ -695,6 +695,57 @@ int MXQuantizeSymbol(SymbolHandle sym_handle,
API_END_HANDLE_ERROR(delete s);
}

int MXReducePrecisionSymbol(SymbolHandle sym_handle,
SymbolHandle *ret_sym_handle,
const int* target_dtype,
const mx_uint num_target_dtype_op_names,
const mx_uint num_fp32_op_names,
const mx_uint num_widest_dtype_op_names,
const mx_uint num_conditional_fp32_op_names,
const mx_uint num_excluded_symbols,
const char **target_dtype_op_names,
const char **fp32_op_names,
const char **widest_dtype_op_names,
const char **conditional_fp32_op_names,
const char **excluded_symbols) {
nnvm::Symbol *s = new nnvm::Symbol();
API_BEGIN();
nnvm::Symbol *sym = static_cast<nnvm::Symbol*>(sym_handle);
nnvm::Graph g = Symbol2Graph(*sym);
std::unordered_set<std::string> target_dtype_ops;
std::unordered_set<std::string> fp32_ops;
std::unordered_set<std::string> widest_dtype_ops;
std::unordered_set<std::string> conditional_fp32_ops;
std::unordered_set<std::string> excluded_syms;
int target_dt = *target_dtype;
for (size_t i = 0; i < num_target_dtype_op_names; ++i) {
target_dtype_ops.emplace(target_dtype_op_names[i]);
}
for (size_t i = 0; i < num_fp32_op_names; ++i) {
fp32_ops.emplace(fp32_op_names[i]);
}
for (size_t i = 0; i < num_widest_dtype_op_names; ++i) {
widest_dtype_ops.emplace(widest_dtype_op_names[i]);
}
for (size_t i = 0; i < num_conditional_fp32_op_names; ++i) {
conditional_fp32_ops.emplace(conditional_fp32_op_names[i]);
}
for (size_t i = 0; i < num_excluded_symbols; ++i) {
excluded_syms.emplace(excluded_symbols[i]);
}
g.attrs["target_dtype_ops"] = std::make_shared<nnvm::any>(std::move(target_dtype_ops));
g.attrs["fp32_ops"] = std::make_shared<nnvm::any>(std::move(fp32_ops));
g.attrs["widest_dtype_ops"] = std::make_shared<nnvm::any>(std::move(widest_dtype_ops));
g.attrs["conditional_fp32_ops"] = std::make_shared<nnvm::any>(std::move(conditional_fp32_ops));
g.attrs["excluded_syms"] = std::make_shared<nnvm::any>(std::move(excluded_syms));
g.attrs["target_dtype"] = std::make_shared<nnvm::any>(target_dt);
g = ApplyPass(std::move(g), "ReducePrecision");
s->outputs = g.outputs;
*ret_sym_handle = s;
API_END_HANDLE_ERROR(delete s);
}


int MXSetCalibTableToQuantizedSymbol(SymbolHandle qsym_handle,
const mx_uint num_layers,
const char** layer_names,
Expand Down Expand Up @@ -725,11 +776,11 @@ int MXGenBackendSubgraph(SymbolHandle sym_handle, const char *backend,
std::vector<mxnet::op::SubgraphPropertyPtr> properties =
mxnet::op::SubgraphPropertyRegistry::Get()->CreateSubgraphProperty(backend);
for (auto property : properties) {
nnvm::Graph g = Symbol2Graph(*s);
nnvm::Graph g = Symbol2Graph(*s);
property->SetAttr("graph", g);
g.attrs["subgraph_property"] = std::make_shared<nnvm::any>(std::move(property));
g = ApplyPass(std::move(g), "BuildSubgraph");
s->outputs = g.outputs;
s->outputs = g.outputs;
}
*ret_sym_handle = s;
API_END_HANDLE_ERROR(delete s);
Expand Down
Loading