ARotting's picture
Publish Interactive FP32, binary, and ternary comparison
2eec02e verified
Raw
History Blame Contribute Delete
1.63 kB
from __future__ import annotations
import torch
from torch import nn
from torch.nn import functional as F
def binary_weight(weight: torch.Tensor) -> torch.Tensor:
scale = weight.abs().mean(dim=1, keepdim=True).clamp_min(1e-6)
quantized = scale * torch.where(weight >= 0, 1.0, -1.0)
return weight + (quantized - weight).detach()
def ternary_weight(weight: torch.Tensor) -> torch.Tensor:
scale = weight.abs().mean(dim=1, keepdim=True).clamp_min(1e-6)
normalized = weight / scale
quantized = scale * normalized.round().clamp(-1, 1)
return weight + (quantized - weight).detach()
class QuantizedLinear(nn.Linear):
def __init__(self, *args, bits: str = "fp32", **kwargs) -> None:
super().__init__(*args, **kwargs)
self.bits = bits
def forward(self, inputs: torch.Tensor) -> torch.Tensor:
if self.bits == "binary":
weight = binary_weight(self.weight)
elif self.bits == "ternary":
weight = ternary_weight(self.weight)
else:
weight = self.weight
return F.linear(inputs, weight, self.bias)
class BitMLP(nn.Module):
def __init__(self, bits: str = "fp32") -> None:
super().__init__()
self.bits = bits
self.hidden = QuantizedLinear(64, 64, bits=bits)
self.output = QuantizedLinear(64, 10, bits=bits)
def forward(self, images: torch.Tensor) -> torch.Tensor:
flattened = images.flatten(1)
return self.output(F.gelu(self.hidden(flattened)))
def parameter_count(module: nn.Module) -> int:
return sum(parameter.numel() for parameter in module.parameters())