| 447 | |
| 448 | @register_lower_rule(mops.FakeQuant) |
| 449 | def fakequant_lower(ctx, *args: Union[HLOTensor, Sequence[HLOTensor]]): |
| 450 | assert ( |
| 451 | len(args) == 3 and len(ctx.vars_in) == 3 |
| 452 | ), f"{len(args)}, {len(ctx.vars_in)}, {len(ctx.vars_out)}" |
| 453 | assert args[1].ndim == args[2].ndim, f"{args[1].shape}, {args[2].shape}" |
| 454 | qmax = np.array(ctx.param["qmax"], dtype=args[0].dtype) |
| 455 | qmin = np.array(ctx.param["qmin"], dtype=args[0].dtype) |
| 456 | inp, scale, zero_point = args[:3] |
| 457 | res = round(inp / scale) + zero_point |
| 458 | res = minimum(maximum(res, qmin), qmax) |
| 459 | res = (res - zero_point) * scale |
| 460 | return res |
| 461 | |
| 462 | |
| 463 | @register_lower_rule("FakeQuantBackward") |