(weight)
| 3 | import torch |
| 4 | |
| 5 | def quant_weight_fp16(weight): |
| 6 | weight = weight.to(torch.float) |
| 7 | s = 1.0 / weight.abs().mean().clamp_(min=1e-5) |
| 8 | new_weight = (weight * s).round().clamp(-1, 1) / s |
| 9 | return new_weight |
| 10 | |
| 11 | def quant_model(input, output): |
| 12 | tensors = {} |