hoe / int_matmul.py
Cunt1257's picture
Upload 8 files
96ba80e verified
Raw
History Blame Contribute Delete
1.36 kB
import torch
import torch.nn as nn
import torch.nn.functional as F
from quantize.quantizer import UniformAffineQuantizer
class QuantMatMul(nn.Module):
def __init__(
self,
x1_quant_params: dict = {},
x2_quant_params: dict = {},
disable_act_quant=False,
matmul_func=torch.bmm,
):
super().__init__()
# de-activate the quantized forward default
self.use_act_quant = False
self.quant_rate = 1.0
# initialize quantizer
self.i_cluster_counts = None
self.x1_quantizer = UniformAffineQuantizer(**x1_quant_params)
self.x2_quantizer = UniformAffineQuantizer(**x2_quant_params)
self.matmul_func = matmul_func
self.disable_act_quant = disable_act_quant
def set_quant_state(self, weight_quant: bool = False, act_quant: bool = False, quant_rate:float = 1.0):
self.use_weight_quant = weight_quant
self.use_act_quant = act_quant
self.quant_rate = quant_rate
def quant_x1(self, x1):
if self.use_act_quant:
x1 = self.x1_quantizer(x1,self.quant_rate)
return x1
def quant_x2(self, x2):
if self.use_act_quant:
x2 = self.x2_quantizer(x2,self.quant_rate)
return x2
def forward(self, x1, x2):
out = self.matmul_func(x1, x2)
return out