Skip to content
Merged
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
131 changes: 102 additions & 29 deletions python/tvm/relay/backend/contrib/ethosu/legalize.py
Original file line number Diff line number Diff line change
Expand Up @@ -135,21 +135,67 @@ def get_lut_from_func(
ofm_scale: float,
ofm_zp: int,
func: Callable[[float], float],
dtype,
) -> List[int]:
"""Calculates the values of the lookup table based on the calculation function"""

lut_values = list()
# Only int8 is currently supported
dtype = np.int8
qmin, qmax = np.iinfo(dtype).min, np.iinfo(dtype).max
for x in range(qmin, qmax + 1):
x_real = ifm_scale * (x - ifm_zp)
out_real = func(x_real)
lut_result = int(util.round_away_zero(ofm_zp + out_real / ofm_scale))
lut_result = min(qmax, max(qmin, lut_result))
lut_values.append(lut_result)
assert dtype in ["int8", "int16"]

return lut_values
if dtype == "int8":
lut_values = list()
qmin, qmax = np.iinfo(dtype).min, np.iinfo(dtype).max
for x in range(qmin, qmax + 1):
x_real = ifm_scale * (x - ifm_zp)
out_real = func(x_real)
lut_result = int(util.round_away_zero(ofm_zp + out_real / ofm_scale))
lut_result = min(qmax, max(qmin, lut_result))
lut_values.append(lut_result)

return lut_values
else:
# dtype == "int16"
table_min = np.iinfo(np.int16).min
table_max = np.iinfo(np.int16).max

input_min = ifm_scale * (table_min - ifm_zp)
input_max = ifm_scale * (table_max - ifm_zp)

output_min = ofm_scale * (table_min - ofm_zp)
output_max = ofm_scale * (table_max - ofm_zp)
# Create 16 bit lut following the reference
nbr_steps = 512
step = (input_max - input_min) / nbr_steps
half_step = step / 2
output_scaling_inv = (table_max - table_min + 1) / (output_max - output_min)

values = []
for i in range(nbr_steps):
val = func(input_min + i * step)
val_midpoint = func(input_min + i * step + half_step)
val_next = func(input_min + (i + 1) * step)

sample_val = util.round_away_zero(val * output_scaling_inv)
midpoint_interp_val = util.round_away_zero(
(val_next * output_scaling_inv + util.round_away_zero(val * output_scaling_inv)) / 2
)
midpoint_val = util.round_away_zero(val_midpoint * output_scaling_inv)
midpoint_err = midpoint_interp_val - midpoint_val
bias = util.round_away_zero(midpoint_err / 2)

lut_result = min(max(sample_val - bias, table_min), table_max)
values.append(lut_result)

val = util.round_away_zero(func(input_max) * output_scaling_inv)
lut_result = min(max(val, table_min), table_max)
values.append(lut_result)
# Convert to hardware 16bit lut with base and slope
lut = [0] * nbr_steps
for i in range(nbr_steps):
slope = (int(values[i + 1]) - int(values[i])) << 16
base = int(values[i])
lut[i] = slope + base

return lut


class LutActivationRewriter(DFPatternCallback):
Expand All @@ -176,25 +222,40 @@ def callback(self, pre: tvm.relay.Expr, post: tvm.relay.Expr, node_map: tvm.ir.c
output_scale = float(params.ofm.q_params.scale_f32)
output_zp = int(params.ofm.q_params.zero_point)

lut_values = get_lut_from_func(
input_scale,
input_zp,
output_scale,
output_zp,
self.calc_func,
)
lut = relay.const(lut_values, dtype=params.ifm.dtype)
# Validation function from pattern matching checks that the input type can be int8 or int16
ifm_dtype = params.ifm.dtype
if ifm_dtype == "int8":
lut_values = get_lut_from_func(
input_scale, input_zp, output_scale, output_zp, self.calc_func, ifm_dtype
)
lut = relay.const(lut_values, dtype=ifm_dtype)

# We baked the requantization into the LUT, so we don't requantize the identity operator
identity = ethosu_ops.ethosu_identity(
ifm=params.ifm.tensor,
lut=lut,
ifm_scale=input_scale,
ifm_zero_point=input_zp,
ofm_scale=input_scale,
ofm_zero_point=input_zp,
activation=self.activation_type,
)
# We baked the requantization into the LUT, so we don't requantize the identity operator
identity = ethosu_ops.ethosu_identity(
ifm=params.ifm.tensor,
lut=lut,
ifm_scale=input_scale,
ifm_zero_point=input_zp,
ofm_scale=input_scale,
ofm_zero_point=input_zp,
activation=self.activation_type,
)

else:
# ifm_dtype == "int16"
lut = get_lut_from_func(
input_scale, input_zp, output_scale, output_zp, self.calc_func, ifm_dtype
)
lut = relay.const(lut, dtype="int32")
identity = ethosu_ops.ethosu_identity(
ifm=params.ifm.tensor,
lut=lut,
ifm_scale=input_scale,
ifm_zero_point=0,
ofm_scale=output_scale,
ofm_zero_point=0,
activation=self.activation_type,
)

return identity

Expand All @@ -208,6 +269,17 @@ def __init__(self):
)


class TanhFixedPointRewriter(LutActivationRewriter):
"""This pass adds tanh with fixed point as a LUT to the identity operator"""

def __init__(self):
super().__init__(
params_class=ethosu_patterns.TanhFixedPointParams,
activation_type="TANH",
calc_func=math.tanh,
)


def sigmoid_calc_func(x: float) -> float:
"""Function to calculate the values for sigmoid"""
# These limits are inherited from TFLite
Expand Down Expand Up @@ -1690,6 +1762,7 @@ def transform_npu_function(self, _, func: relay.Function) -> relay.Function:
ShlRewriter(),
AbsRewriter(),
TanhRewriter(),
TanhFixedPointRewriter(),
HardSwishRewriter(),
LeakyReLURewriter(),
MeanRewriter(),
Expand Down
61 changes: 60 additions & 1 deletion python/tvm/relay/op/contrib/ethosu.py
Original file line number Diff line number Diff line change
Expand Up @@ -1251,7 +1251,7 @@ def is_valid(self):
"""
This function checks whether activation has compatible attributes with the NPU
"""
if not check_valid_dtypes([self.ifm, self.ofm], supported_dtypes=[np.int8]):
if not check_valid_dtypes([self.ifm, self.ofm], supported_dtypes=[np.int8, np.int16]):
return False
return True

Expand All @@ -1269,6 +1269,60 @@ def tanh_pattern():
return quant


class TanhFixedPointParams:
"""
This class will parse a call to a ethos-u.tanh_fixed_point composite function
and extract the parameter information.
"""

composite_name = "ethos-u.tanh_fixed_point"

@requires_vela
def __init__(self, func_body):
layout = "NHWC"

tanh_fixed_point = func_body.args[0]
tanh = tanh_fixed_point.args[0]
# fixed_point_multiply relay operation uses multiplier with 31 fractional bits
# so to determine the size of the fraction use the formula: 31 - shift
self.fraction_size = 31 - tanh_fixed_point.attrs.shift
fract_scale = tvm.relay.Constant(tvm.nd.array(np.array(1 / 2**self.fraction_size)))
fract_zero_point = tvm.relay.Constant(tvm.nd.array(np.array(0, dtype="int32")))

self.ifm = TensorParams(
tanh.args[0].args[0].args[0],
layout=layout,
scale=fract_scale,
zero_point=fract_zero_point,
)
self.ofm = TensorParams(
func_body,
layout=layout,
scale=fract_scale,
zero_point=fract_zero_point,
)

def is_valid(self) -> bool:
"""
This function checks whether activation has compatible attributes with the NPU
"""

if self.fraction_size < 0 or self.fraction_size > 16:
return False
if not check_valid_dtypes([self.ifm, self.ofm], supported_dtypes=[np.int8, np.int16]):
return False
return True


def tanh_fixed_point_pattern():
"""Create pattern for fixed point tanh"""
ifm = is_op("cast")(wildcard())
ifm = is_op("fixed_point_multiply")(ifm)
tanh = is_op("tanh")(ifm)
tanh = is_op("fixed_point_multiply")(tanh)
return is_op("cast")(tanh)


class SigmoidParams(LutActivationParams):
"""
This class will parse a call to a ethos-u.sigmoid composite function
Expand Down Expand Up @@ -2373,6 +2427,11 @@ def pattern_table() -> List[Tuple[str, tvm.relay.dataflow_pattern.DFPattern, Cal
lambda pat: AbsParams(pat).is_valid(),
),
(TanhParams.composite_name, tanh_pattern(), lambda pat: TanhParams(pat).is_valid()),
(
TanhFixedPointParams.composite_name,
tanh_fixed_point_pattern(),
lambda pat: TanhFixedPointParams(pat).is_valid(),
),
(
MeanParams.composite_name,
mean_pattern(),
Expand Down
45 changes: 45 additions & 0 deletions tests/python/contrib/test_ethosu/test_codegen.py
Original file line number Diff line number Diff line change
Expand Up @@ -1675,5 +1675,50 @@ def convert_to_fixed_point(arr, fract_size):
)


@pytest.mark.parametrize("accel_type", ["ethos-u55-256", "ethos-u65-256"])
@pytest.mark.parametrize(
"ifm_shape,fract_size,tolerance",
[[(1, 2, 8, 4), 15, 0.001], [(1, 8), 12, 0.001], [(1, 1, 4, 8), 10, 0.002]],
)
def test_ethosu_tanh_fixed_point(accel_type, ifm_shape, fract_size, tolerance):
np.random.seed(0)
dtype = "int16"

def create_model():
ifm = relay.var("ifm", shape=ifm_shape, dtype=dtype)
ifm_fixed_point = relay.cast(ifm, "int32")
ifm_fixed_point = relay.fixed_point_multiply(ifm_fixed_point, 2**31 - 1, 0)
tanh = relay.tanh(ifm_fixed_point)
tanh = relay.fixed_point_multiply(tanh, 1, 31 - fract_size)
tanh = relay.cast(tanh, dtype)
return tvm.IRModule.from_expr(relay.Function([ifm], tanh))

def generate_ref(input_data):
return np.tanh(input_data)

def convert_to_fixed_point(arr, fract_size):
fract_fact = 0b1 << fract_size
return np.array(arr * fract_fact, dtype=np.int16)

cpu_mod = create_model()
ethosu_mod = partition_for_ethosu(cpu_mod)

input_data = {"ifm": np.random.uniform(-1, 1, size=ifm_shape)}
output_data = generate_ref(input_data["ifm"])

input_data = {"ifm": convert_to_fixed_point(input_data["ifm"], fract_size)}
output_data = {"output": convert_to_fixed_point(output_data, fract_size)}
tolerance = convert_to_fixed_point(tolerance, fract_size)

infra.compare_ethosu_with_reference(
ethosu_mod,
input_data,
output_data,
accel_type,
enable_cascader=is_u55_accel_type(accel_type),
output_tolerance=tolerance,
)


if __name__ == "__main__":
tvm.testing.main()
45 changes: 45 additions & 0 deletions tests/python/contrib/test_ethosu/test_legalize.py
Original file line number Diff line number Diff line change
Expand Up @@ -3923,5 +3923,50 @@ def _visit(stmt):
verify(mod["tvmgen_default_ethos_u_main_0"])


@pytest.mark.parametrize(
"ifm_shape,fract_size",
[[(1, 2, 8, 4), 15], [(1, 8), 12], [(1, 1, 4, 8), 10]],
)
def test_relay_tanh_fixed_point_legalize(ifm_shape, fract_size):
dtype = "int16"

def create_model():
ifm = relay.var("ifm", shape=ifm_shape, dtype=dtype)
ifm_fixed_point = relay.cast(ifm, "int32")
ifm_fixed_point = relay.fixed_point_multiply(ifm_fixed_point, 2**31 - 1, 0)
tanh = relay.tanh(ifm_fixed_point)
tanh = relay.fixed_point_multiply(tanh, 1, 31 - fract_size)
tanh = relay.cast(tanh, dtype)
return tvm.IRModule.from_expr(relay.Function([ifm], tanh))

mod = create_model()

tanh_pattern_table = [
(
ethosu.TanhFixedPointParams.composite_name,
ethosu.tanh_fixed_point_pattern(),
lambda pat: ethosu.TanhFixedPointParams(pat).is_valid(),
),
]

mod = partition_ethosu_by_table(mod, tanh_pattern_table)
mod["tvmgen_default_ethos_u_main_0"] = dataflow_pattern.rewrite(
legalize.TanhFixedPointRewriter(), mod["tvmgen_default_ethos_u_main_0"]
)
mod = relay.transform.InferType()(mod)

func = mod["tvmgen_default_ethos_u_main_0"]

identity = func.body
assert identity.op.name == "contrib.ethosu.identity"
assert identity.attrs.activation == "TANH"
assert identity.args[0].checked_type.dtype == dtype
assert tuple(identity.args[0].checked_type.shape) == ifm_shape
# check LUT size
assert tuple(identity.args[1].checked_type.shape) == (512,)
assert identity.attrs.ifm_scale == 1 / 2**fract_size
assert identity.attrs.ifm_scale == identity.attrs.ofm_scale


if __name__ == "__main__":
tvm.testing.main()