(ctx, *args: Union[HLOTensor, Sequence[HLOTensor]])
| 462 | |
| 463 | @register_lower_rule("FakeQuantBackward") |
| 464 | def fakequant_grad_rule(ctx, *args: Union[HLOTensor, Sequence[HLOTensor]]): |
| 465 | assert ( |
| 466 | len(args) == 4 and len(ctx.vars_in) == 4 |
| 467 | ), f"{len(args)}, {len(ctx.vars_in)}, {len(ctx.vars_out)}" |
| 468 | assert args[0].ndim == args[1].ndim, f"{args[0].shape}, {args[1].shape}" |
| 469 | assert args[2].ndim == args[3].ndim, f"{args[2].shape}, {args[3].shape}" |
| 470 | diff, inp, scale, zere_point = args[:4] |
| 471 | x = round(inp / scale) + zere_point |
| 472 | qmax = np.array(ctx.param["qmax"], dtype=x.dtype) |
| 473 | qmin = np.array(ctx.param["qmin"], dtype=x.dtype) |
| 474 | mask1 = x <= qmax |
| 475 | mask2 = x >= qmin |
| 476 | mask = logical_and(mask1, mask2).astype(diff.dtype) |
| 477 | return diff * mask |
| 478 | |
| 479 | |
| 480 | @register_lower_rule(mops.TQT) |
nothing calls this directly
no test coverage detected