| name | bitvla-robot-control |
| title | BitVLA: 1-bit Vision-Language-Action Models for Robotics Manipulation |
| version | 0.0.2 |
| engine | skillxiv-v0.0.2-claude-opus-4.6 |
| license | MIT |
| url | https://arxiv.org/abs/2506.07530 |
| keywords | ["quantization","robotics","vision-language","manipulation","edge-deployment","1bit"] |
| description | Build fully ternary quantized vision-language-action models for robotic manipulation, achieving 11x memory reduction and 4.4x speedup while maintaining task performance on edge devices. |
BitVLA: 1-bit Vision-Language-Action Models for Robotics Manipulation
Core Concept
BitVLA is the first fully native 1-bit quantized vision-language-action model where every parameter is ternary ({−1, 0, 1}), enabling efficient robotic policies on edge devices. Through a novel Quantize-then-Distill strategy, it compresses vision encoders to 1.58-bit weights while maintaining alignment with language understanding. The 2B parameter model achieves 11x memory reduction and 4.4x speedup compared to full-precision baselines, matching larger models on manipulation benchmarks while fitting on resource-constrained robot hardware.
Architecture Overview
- BitNet b1.58 Foundation: 1-bit LLM backbone with ternary weights
- Quantized Vision Encoder: SigLIP-L compressed to 1.58-bit via Quantize-then-Distill
- Lightweight Connector: Full-precision alignment layer (bottleneck acceptable)
- Ternary Action Head: Binary token output for robot actions
- Causal Attention Preservation: Maintains LLM-style masking for task requirements
- INT8 Activations: Per-token symmetric quantization during inference
Implementation
Step 1: Implement Ternary Quantization Functions
Create the core quantization operators:
import torch
import torch.nn as nn
import torch.nn.functional as F
class TernaryQuantizer(nn.Module):
"""Quantize weights and activations to ternary values"""
def __init__(self):
super().__init__()
@staticmethod
def quantize_weights_ternary(weight, scale=None):
"""
Quantize weights to {-1, 0, 1}.
Uses absmean scaling: scale = mean(|weight|)
"""
if scale is None:
scale = weight.abs().mean()
threshold = scale / 2.0
weight_normalized = weight / scale
weight_ternary = torch.sign(weight_normalized)
weight_ternary[weight_normalized.abs() < 0.5] = 0
return weight_ternary, scale
@staticmethod
def quantize_activations_int8(activation):
"""
Quantize activations to INT8.
Per-token symmetric quantization using absmax scaling.
"""
batch_size = activation.shape[0]
activation_reshaped = activation.reshape(batch_size, -1)
scales = activation_reshaped.abs().max(dim=)[]
activation_normalized = (activation / scales.unsqueeze(-)).clamp(-, )
activation_int8 = (activation_normalized * ).().to(torch.int8)
activation_int8, scales
():
(quantized, torch.Tensor) quantized.dtype == torch.int8:
quantized.() / * scale.unsqueeze(-)
:
quantized * scale
Step 2: Build Quantized Vision Encoder
Compress vision encoder using Quantize-then-Distill:
class QuantizedVisionEncoder(nn.Module):
"""
SigLIP vision encoder compressed to 1.58-bit weights.
Uses knowledge distillation from full-precision teacher.
"""
def __init__(self, teacher_encoder, target_bits=1.58):
super().__init__()
self.teacher = teacher_encoder
self.target_bits = target_bits
self.quantizer = TernaryQuantizer()
self.student = self.create_student_encoder()
def create_student_encoder(self):
"""Create ternary student encoder"""
config = self.teacher.config
student = self.teacher.__class__(config)
for name, param in student.named_parameters():
if 'weight' in name and len(param.shape) > 1:
ternary_weight, _ = self.quantizer.quantize_weights_ternary(param)
param.data = ternary_weight
return student
def forward(self, images):
"""Forward pass with quantized vision encoder"""
features = self.student.encoder(images)
return features
():
torch.no_grad():
teacher_features = .teacher.encoder(images)
teacher_hidden = .teacher.encoder.pool(teacher_features)
student_features = .student.encoder(images)
student_hidden = .student.encoder.pool(student_features)
teacher_norm = F.normalize(teacher_hidden, dim=-)
student_norm = F.normalize(student_hidden, dim=-)
mse_loss = F.mse_loss(student_norm, teacher_norm)
mse_loss
():
name, param .student.named_parameters():
name (param.shape) > :
param.grad :
param.grad.data = param.grad.data / (param.().mean() + )
torch.no_grad():
ternary_weight, scale = .quantizer.quantize_weights_ternary(param)
param.data = ternary_weight
Step 3: Build BitVLA Model Architecture
Integrate quantized vision and language components:
class BitVLA(nn.Module):
"""Fully ternary vision-language-action model for robotics"""
def __init__(self, llm_checkpoint="OpenBitNet/bitnet-b1.58-2B",
vision_checkpoint="google/siglip-base-patch16-512"):
super().__init__()
from transformers import AutoModelForCausalLM, AutoTokenizer
self.llm = AutoModelForCausalLM.from_pretrained(llm_checkpoint)
self.tokenizer = AutoTokenizer.from_pretrained(llm_checkpoint)
self.llm_dim = self.llm.config.hidden_size
teacher_vision = self.load_vision_encoder(vision_checkpoint)
self.vision_encoder = QuantizedVisionEncoder(teacher_vision, target_bits=1.58)
self.vision_dim = 768
self.connector = nn.Linear(self.vision_dim, self.llm_dim)
self.action_head = TernaryActionHead(self.llm_dim)
def load_vision_encoder(self, checkpoint):
"""Load SigLIP vision encoder"""
from transformers import AutoModel
return AutoModel.from_pretrained(checkpoint)
def forward():
vision_features = .vision_encoder(images)
vision_features = vision_features.mean(dim=)
aligned_features = .connector(vision_features)
prompt_ids = .tokenizer.encode(action_prompt, return_tensors=)
prompt_embeddings = .llm.get_input_embeddings()(prompt_ids)
combined_embeddings = torch.cat([
prompt_embeddings,
aligned_features.unsqueeze()
], dim=)
outputs = .llm(
inputs_embeds=combined_embeddings,
attention_mask=torch.ones(combined_embeddings.shape[:]),
use_cache=
)
logits = .action_head(outputs.hidden_states[-][:, -, :])
actions = torch.argmax(logits, dim=-)
actions
(nn.Module):
():
().__init__()
.hidden_dim = hidden_dim
.num_actions = num_actions
.action_projection = nn.Linear(hidden_dim, num_actions)
():
logits = .action_projection(hidden_state)
logits
Step 4: Implement Three-Stage Training Pipeline
Create the complete training procedure:
class BitVLATrainer:
"""Three-stage training: multimodal alignment -> quantization -> RL fine-tuning"""
def __init__(self, model, device='cuda'):
self.model = model.to(device)
self.device = device
def stage_1_multimodal_alignment(self, image_text_dataset, epochs=5):
"""
Stage 1: Align vision-language features (vision encoder frozen).
Training objective: make vision features compatible with LLM space.
"""
print("Stage 1: Multimodal alignment...")
optimizer = torch.optim.AdamW(
[p for n, p in self.model.named_parameters() if 'connector' in n or 'action_head' in n],
lr=1e-4
)
for epoch in range(epochs):
total_loss = 0
for batch in image_text_dataset:
images = batch['images'].to(self.device)
text_ids = batch['text_ids'].to(self.device)
optimizer.zero_grad()
vision_features = self.model.vision_encoder(images)
vision_features = vision_features.mean(dim=1)
aligned = self.model.connector(vision_features)
text_embeddings = .model.llm.get_input_embeddings()(text_ids)
text_features = .model.llm(input_ids=text_ids).hidden_states[-][:, , :]
loss = .contrastive_loss(aligned, text_features)
loss.backward()
torch.nn.utils.clip_grad_norm_(.model.parameters(), )
optimizer.step()
total_loss += loss.item()
()
():
()
optimizer = torch.optim.AdamW(
.model.vision_encoder.student.parameters(),
lr=
)
epoch (epochs):
total_loss =
batch image_dataset:
images = batch[].to(.device)
optimizer.zero_grad()
loss = .model.vision_encoder.distillation_loss(images)
loss.backward()
torch.nn.utils.clip_grad_norm_(.model.vision_encoder.parameters(), )
.model.vision_encoder.update_ternary_weights()
optimizer.step()
total_loss += loss.item()
()
():
()
optimizer = torch.optim.AdamW(
.model.action_head.parameters(),
lr=
)
epoch (epochs):
total_reward =
batch robot_trajectory_dataset:
images = batch[].to(.device)
actions = batch[].to(.device)
rewards = batch[].to(.device)
optimizer.zero_grad()
predicted_actions = .model(images)
action_loss = F.cross_entropy(predicted_actions, actions)
loss = (action_loss * rewards).mean()
loss.backward()
torch.nn.utils.clip_grad_norm_(.model.parameters(), )
optimizer.step()
total_reward += rewards.mean().item()
()
():
v1 = F.normalize(v1, dim=-)
v2 = F.normalize(v2, dim=-)
sim = torch.mm(v1, v2.t()) / temperature
labels = torch.arange(v1.shape[]).to(.device)
loss = F.cross_entropy(sim, labels)
loss
Practical Guidance
- Memory Footprint: 1.4GB total (vs. 15.4GB full-precision), enabling edge deployment
- Latency: 73ms per inference, 341.1 Hz throughput on typical edge hardware
- Quantization Strategy: Vision = 1.58-bit, LLM = 1-bit (BitNet), connector = full-precision
- Causal Attention: Preserve causal masking in LLM backbone for task requirements
- Distillation Temperature: Start with temperature=4.0; adjust based on convergence
- Training Data: ~1M robot trajectories from mix of sources (LIBERO, real data, simulation)
- Performance Maintenance: ternary quantization loses <1% accuracy on manipulation benchmarks
- Hardware Targets: Optimized for NVIDIA Jetson, but works on any INT8-capable device
Reference
- BitNet b1.58 proves that 1-bit quantization works for large language models
- Quantize-then-Distill separates compression from multimodal alignment concerns
- Straight-through estimators enable gradient flow through quantization operations
- Causal attention preservation is critical for sequential decision-making in robotics