Initialize Linear4bit class. Args: input_features (`str`): Number of input features of the linear layer. output_features (`str`): Number of output features of the linear layer. bias (`bool`, defaults to `True`):
(
self,
input_features,
output_features,
bias=True,
compute_dtype=None,
compress_statistics=True,
quant_type="fp4",
quant_storage=torch.uint8,
device=None,
)
| 535 | """ |
| 536 | |
| 537 | def __init__( |
| 538 | self, |
| 539 | input_features, |
| 540 | output_features, |
| 541 | bias=True, |
| 542 | compute_dtype=None, |
| 543 | compress_statistics=True, |
| 544 | quant_type="fp4", |
| 545 | quant_storage=torch.uint8, |
| 546 | device=None, |
| 547 | ): |
| 548 | """ |
| 549 | Initialize Linear4bit class. |
| 550 | |
| 551 | Args: |
| 552 | input_features (`str`): |
| 553 | Number of input features of the linear layer. |
| 554 | output_features (`str`): |
| 555 | Number of output features of the linear layer. |
| 556 | bias (`bool`, defaults to `True`): |
| 557 | Whether the linear class uses the bias term as well. |
| 558 | """ |
| 559 | super().__init__(input_features, output_features, bias, device) |
| 560 | self.weight = Params4bit( |
| 561 | self.weight.data, |
| 562 | requires_grad=False, |
| 563 | compress_statistics=compress_statistics, |
| 564 | quant_type=quant_type, |
| 565 | quant_storage=quant_storage, |
| 566 | module=self, |
| 567 | ) |
| 568 | # self.persistent_buffers = [] # TODO consider as way to save quant state |
| 569 | self.compute_dtype = compute_dtype |
| 570 | self.compute_type_is_set = compute_dtype is not None |
| 571 | self.quant_state = None |
| 572 | self.quant_storage = quant_storage |
| 573 | self.support_avx512bf16_for_cpu = has_avx512bf16() |
| 574 | |
| 575 | def set_compute_type(self, x): |
| 576 | if x.dtype in [torch.float32, torch.bfloat16]: |
nothing calls this directly
no test coverage detected