| 860 | raise NotImplementedError("Saving Embedding8bit module is not implemented") |
| 861 | |
| 862 | def forward(self, input: Tensor) -> Tensor: |
| 863 | if not hasattr(self.weight, "SCB"): |
| 864 | raise RuntimeError("Embedding layer is not quantized. Please call .cuda() or .to(device) first.") |
| 865 | |
| 866 | rows = self.weight.data |
| 867 | row_stats = self.weight.SCB |
| 868 | |
| 869 | assert rows.shape == (self.num_embeddings, self.embedding_dim) |
| 870 | assert row_stats.shape == (self.num_embeddings,) |
| 871 | |
| 872 | compressed_output = F.embedding(input, rows) |
| 873 | compressed_output_stats = F.embedding(input, row_stats.view(self.num_embeddings, 1)) |
| 874 | |
| 875 | output = compressed_output * (compressed_output_stats / 127.0) |
| 876 | |
| 877 | return output.to(self.dtype) |
| 878 | |
| 879 | |
| 880 | class Embedding4bit(nn.Embedding): |