An XLA version of CpuNudge().
| 46 | |
| 47 | // An XLA version of CpuNudge(). |
| 48 | void XlaNudge(xla::XlaBuilder* b, const DataType data_type, |
| 49 | const xla::XlaOp& min, const xla::XlaOp& max, |
| 50 | const float quant_min_value, const float quant_max_value, |
| 51 | xla::XlaOp* nudged_min, xla::XlaOp* nudged_max, |
| 52 | xla::XlaOp* scale) { |
| 53 | *scale = xla::Div(xla::Sub(max, min), |
| 54 | XlaHelpers::FloatLiteral( |
| 55 | b, data_type, quant_max_value - quant_min_value)); |
| 56 | xla::XlaOp quant_min = |
| 57 | XlaHelpers::FloatLiteral(b, data_type, quant_min_value); |
| 58 | xla::XlaOp zero_point_from_min = xla::Sub(quant_min, xla::Div(min, *scale)); |
| 59 | xla::XlaOp quant_max = |
| 60 | XlaHelpers::FloatLiteral(b, data_type, quant_max_value); |
| 61 | xla::XlaOp nudged_zero_point = |
| 62 | xla::Select(xla::Le(zero_point_from_min, quant_min), quant_min, |
| 63 | xla::Select(xla::Ge(zero_point_from_min, quant_max), |
| 64 | quant_max, xla::Round(zero_point_from_min))); |
| 65 | *nudged_min = xla::Mul(xla::Sub(quant_min, nudged_zero_point), *scale); |
| 66 | *nudged_max = xla::Mul(xla::Sub(quant_max, nudged_zero_point), *scale); |
| 67 | } |
| 68 | |
| 69 | xla::XlaOp Quantize(xla::XlaBuilder* b, const xla::XlaOp& input, |
| 70 | const DataType data_type, |