dq_w = dequant_mxfp4(layer.weight, layer.weight_scale, x.dtype) qdq_x = self.quant_dequant_func(x) return F.linear(qdq_x, dq_w, bias)