Spaces:
Running
Running
| from __future__ import annotations | |
| import copy | |
| import json | |
| from pathlib import Path | |
| import numpy as np | |
| import pandas as pd | |
| import torch | |
| import trackio | |
| from model import BitMLP, parameter_count | |
| from packing import pack_binary_model | |
| from safetensors.torch import save_file | |
| from sklearn.datasets import load_digits | |
| from sklearn.model_selection import train_test_split | |
| from torch import nn | |
| from torch.nn import functional as F | |
| from torch.utils.data import DataLoader, TensorDataset | |
| PROJECT_DIR = Path(__file__).resolve().parent | |
| ARTIFACT_DIR = PROJECT_DIR / "artifacts" / "bitforge-1bit" | |
| DATA_DIR = PROJECT_DIR / "data" | |
| def seed_everything(seed: int) -> None: | |
| np.random.seed(seed) | |
| torch.manual_seed(seed) | |
| torch.set_num_threads(1) | |
| def make_loader( | |
| images: np.ndarray, | |
| labels: np.ndarray, | |
| *, | |
| shuffle: bool, | |
| seed: int, | |
| ) -> DataLoader: | |
| return DataLoader( | |
| TensorDataset( | |
| torch.from_numpy(images.astype(np.float32)), | |
| torch.from_numpy(labels.astype(np.int64)), | |
| ), | |
| batch_size=256, | |
| shuffle=shuffle, | |
| generator=torch.Generator().manual_seed(seed), | |
| ) | |
| def evaluate(model: nn.Module, loader: DataLoader) -> dict: | |
| model.eval() | |
| correct = 0 | |
| total = 0 | |
| losses = [] | |
| for images, labels in loader: | |
| logits = model(images) | |
| losses.append(float(F.cross_entropy(logits, labels))) | |
| correct += int((logits.argmax(1) == labels).sum()) | |
| total += len(labels) | |
| return {"accuracy": correct / total, "cross_entropy": float(np.mean(losses))} | |
| def train_teacher( | |
| model: BitMLP, train_loader: DataLoader, validation_loader: DataLoader | |
| ) -> BitMLP: | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=2e-3, weight_decay=1e-4) | |
| best = copy.deepcopy(model.state_dict()) | |
| best_accuracy = 0.0 | |
| for epoch in range(1, 101): | |
| model.train() | |
| for images, labels in train_loader: | |
| loss = F.cross_entropy(model(images), labels) | |
| optimizer.zero_grad() | |
| loss.backward() | |
| optimizer.step() | |
| if epoch % 5 == 0: | |
| validation = evaluate(model, validation_loader) | |
| trackio.log( | |
| { | |
| "teacher_epoch": epoch, | |
| "teacher_validation_accuracy": validation["accuracy"], | |
| } | |
| ) | |
| if validation["accuracy"] > best_accuracy: | |
| best_accuracy = validation["accuracy"] | |
| best = copy.deepcopy(model.state_dict()) | |
| model.load_state_dict(best) | |
| return model | |
| def train_student( | |
| name: str, | |
| model: BitMLP, | |
| teacher: BitMLP, | |
| train_loader: DataLoader, | |
| validation_loader: DataLoader, | |
| ) -> BitMLP: | |
| optimizer = torch.optim.AdamW(model.parameters(), lr=1.5e-3, weight_decay=2e-5) | |
| best = copy.deepcopy(model.state_dict()) | |
| best_accuracy = 0.0 | |
| temperature = 2.5 | |
| teacher.eval() | |
| for epoch in range(1, 151): | |
| model.train() | |
| for images, labels in train_loader: | |
| logits = model(images) | |
| with torch.no_grad(): | |
| teacher_logits = teacher(images) | |
| hard = F.cross_entropy(logits, labels) | |
| soft = F.kl_div( | |
| F.log_softmax(logits / temperature, dim=1), | |
| F.softmax(teacher_logits / temperature, dim=1), | |
| reduction="batchmean", | |
| ) * temperature**2 | |
| loss = 0.45 * hard + 0.55 * soft | |
| optimizer.zero_grad() | |
| loss.backward() | |
| torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) | |
| optimizer.step() | |
| if epoch % 5 == 0: | |
| validation = evaluate(model, validation_loader) | |
| trackio.log( | |
| { | |
| f"{name}_epoch": epoch, | |
| f"{name}_validation_accuracy": validation["accuracy"], | |
| } | |
| ) | |
| if validation["accuracy"] > best_accuracy: | |
| best_accuracy = validation["accuracy"] | |
| best = copy.deepcopy(model.state_dict()) | |
| model.load_state_dict(best) | |
| return model | |
| def main() -> None: | |
| seed_everything(2043) | |
| digits = load_digits() | |
| images = (digits.images / 16.0).astype(np.float32) | |
| labels = digits.target.astype(np.int64) | |
| indices = np.arange(len(images)) | |
| train_indices, test_indices = train_test_split( | |
| indices, test_size=0.25, random_state=2043, stratify=labels | |
| ) | |
| train_indices, validation_indices = train_test_split( | |
| train_indices, | |
| test_size=0.18, | |
| random_state=3043, | |
| stratify=labels[train_indices], | |
| ) | |
| train_loader = make_loader( | |
| images[train_indices], labels[train_indices], shuffle=True, seed=2043 | |
| ) | |
| validation_loader = make_loader( | |
| images[validation_indices], | |
| labels[validation_indices], | |
| shuffle=False, | |
| seed=3043, | |
| ) | |
| test_loader = make_loader( | |
| images[test_indices], labels[test_indices], shuffle=False, seed=4043 | |
| ) | |
| teacher = BitMLP("fp32") | |
| binary = BitMLP("binary") | |
| ternary = BitMLP("ternary") | |
| trackio.init( | |
| project="bitforge-1bit", | |
| name="binary-ternary-distillation-v1", | |
| config={ | |
| "parameters_per_variant": parameter_count(teacher), | |
| "teacher_epochs": 100, | |
| "student_epochs": 150, | |
| "binary_matrix_weight_bits": 1, | |
| }, | |
| ) | |
| teacher = train_teacher(teacher, train_loader, validation_loader) | |
| binary = train_student( | |
| "binary", binary, teacher, train_loader, validation_loader | |
| ) | |
| ternary = train_student( | |
| "ternary", ternary, teacher, train_loader, validation_loader | |
| ) | |
| ARTIFACT_DIR.mkdir(parents=True, exist_ok=True) | |
| save_file(teacher.state_dict(), ARTIFACT_DIR / "fp32.safetensors") | |
| save_file(binary.state_dict(), ARTIFACT_DIR / "binary_qat.safetensors") | |
| save_file(ternary.state_dict(), ARTIFACT_DIR / "ternary_qat.safetensors") | |
| packing = pack_binary_model(binary, ARTIFACT_DIR / "binary_weights.npz") | |
| fp32_payload_bytes = parameter_count(teacher) * 4 | |
| results = { | |
| "benchmark": "BitForge 1-bit", | |
| "parameters_per_variant": parameter_count(teacher), | |
| "matrix_weight_count": int( | |
| teacher.hidden.weight.numel() + teacher.output.weight.numel() | |
| ), | |
| "test": { | |
| "fp32": evaluate(teacher, test_loader), | |
| "binary_weight": evaluate(binary, test_loader), | |
| "ternary_weight": evaluate(ternary, test_loader), | |
| }, | |
| "storage": { | |
| "fp32_parameter_payload_bytes": fp32_payload_bytes, | |
| **packing, | |
| "measured_payload_compression": fp32_payload_bytes | |
| / packing["packed_payload_bytes"], | |
| }, | |
| "precision_boundary": { | |
| "matrix_weights": "one packed bit in binary variant", | |
| "scales": "float32 per output channel", | |
| "biases": "float32", | |
| "activations": "float32", | |
| }, | |
| } | |
| (ARTIFACT_DIR / "evaluation.json").write_text( | |
| json.dumps(results, indent=2), encoding="utf-8" | |
| ) | |
| DATA_DIR.mkdir(parents=True, exist_ok=True) | |
| pd.DataFrame( | |
| { | |
| "source_index": indices, | |
| "label": labels, | |
| "split": np.select( | |
| [ | |
| np.isin(indices, train_indices), | |
| np.isin(indices, validation_indices), | |
| ], | |
| ["train", "validation"], | |
| default="test", | |
| ), | |
| } | |
| ).to_parquet(DATA_DIR / "split_manifest.parquet", index=False) | |
| trackio.log( | |
| { | |
| "fp32_test_accuracy": results["test"]["fp32"]["accuracy"], | |
| "binary_test_accuracy": results["test"]["binary_weight"]["accuracy"], | |
| "ternary_test_accuracy": results["test"]["ternary_weight"]["accuracy"], | |
| "binary_payload_compression": results["storage"][ | |
| "measured_payload_compression" | |
| ], | |
| } | |
| ) | |
| trackio.finish() | |
| print(json.dumps(results, indent=2)) | |
| if __name__ == "__main__": | |
| main() | |