(ctx, input: Tensor, lower: float, upper: float, bitwidth: int)
| 80 | class _RoundSTE(torch.autograd.Function): |
| 81 | @staticmethod |
| 82 | def forward(ctx, input: Tensor, lower: float, upper: float, bitwidth: int) -> Tensor: |
| 83 | clamped = input.clamp(lower, upper) |
| 84 | denom = (2 ** bitwidth) - 1 |
| 85 | q_step = (upper - lower) / denom |
| 86 | levels = torch.round((clamped - lower) / q_step) |
| 87 | return levels * q_step + lower |
| 88 | |
| 89 | @staticmethod |
| 90 | def backward(ctx, grad_output: Tensor) -> Tuple[Tensor, None, None, None]: |
nothing calls this directly
no outgoing calls
no test coverage detected