ARotting's picture
Publish Interactive FP32, binary, and ternary comparison
2eec02e verified
Raw
History Blame Contribute Delete
1.78 kB
from __future__ import annotations
from pathlib import Path
import numpy as np
import torch
from model import BitMLP
def pack_binary_model(model: BitMLP, destination: Path) -> dict:
payload = {}
payload_bytes = 0
for name in ["hidden", "output"]:
layer = getattr(model, name)
weight = layer.weight.detach().cpu().numpy()
scales = np.mean(np.abs(weight), axis=1).astype(np.float32)
signs = np.packbits((weight >= 0).reshape(-1), bitorder="little")
shape = np.asarray(weight.shape, dtype=np.int32)
bias = layer.bias.detach().cpu().numpy().astype(np.float32)
payload[f"{name}_signs"] = signs
payload[f"{name}_scales"] = scales
payload[f"{name}_shape"] = shape
payload[f"{name}_bias"] = bias
payload_bytes += signs.nbytes + scales.nbytes + bias.nbytes
np.savez(destination, **payload)
return {
"packed_payload_bytes": int(payload_bytes),
"container_bytes": int(destination.stat().st_size),
}
def load_binary_model(source: Path) -> BitMLP:
packed = np.load(source)
model = BitMLP(bits="binary")
with torch.no_grad():
for name in ["hidden", "output"]:
layer = getattr(model, name)
shape = tuple(packed[f"{name}_shape"].astype(int))
total = int(np.prod(shape))
bits = np.unpackbits(
packed[f"{name}_signs"], bitorder="little"
)[:total]
signs = np.where(bits.reshape(shape) == 1, 1.0, -1.0)
scales = packed[f"{name}_scales"][:, None]
layer.weight.copy_(
torch.from_numpy((signs * scales).astype(np.float32))
)
layer.bias.copy_(torch.from_numpy(packed[f"{name}_bias"]))
return model