task23-base_cnn_tuned / flock_robotics_adapter.py
flockgo's picture
Upload folder using huggingface_hub
411f238 verified
Raw
History Blame Contribute Delete
7.3 kB
from __future__ import annotations
import hashlib
import json
import math
import re
from pathlib import Path
import numpy as np
import torch
import torch.nn as nn
import torch.nn.functional as F
VARIANTS = {
"small_cnn": {"width": 32, "depth": 3, "chunk": 1, "dropout": 0.05},
"base_cnn": {"width": 64, "depth": 4, "chunk": 1, "dropout": 0.10},
"chunk_cnn": {"width": 48, "depth": 4, "chunk": 8, "dropout": 0.10},
"ensemble_cnn": {"width": 64, "depth": 5, "chunk": 8, "dropout": 0.15},
}
class VisionBackbone(nn.Module):
def __init__(self, width, depth):
super().__init__()
layers = [nn.Conv2d(3, width, 5, 2, 2), nn.GroupNorm(max(1, width // 8), width), nn.SiLU()]
channels = width
for index in range(depth - 1):
next_channels = min(width * (2 ** min(index + 1, 2)), 256)
layers += [nn.Conv2d(channels, next_channels, 3, 2, 1),
nn.GroupNorm(max(1, next_channels // 8), next_channels), nn.SiLU()]
channels = next_channels
layers += [nn.AdaptiveAvgPool2d((1, 1)), nn.Flatten(), nn.Linear(channels, 384), nn.SiLU()]
self.net = nn.Sequential(*layers)
def forward(self, images):
return self.net(images)
class PolicyNet(nn.Module):
def __init__(self, config):
super().__init__()
variant = config["variant"]
cfg = VARIANTS[variant]
self.chunk = int(cfg["chunk"])
self.vision = VisionBackbone(int(cfg["width"]), int(cfg["depth"]))
self.proprio = nn.Sequential(nn.Linear(int(config["proprio_dim"]), 128), nn.LayerNorm(128), nn.SiLU())
self.text = nn.Sequential(nn.Linear(int(config["text_dim"]), 128), nn.LayerNorm(128), nn.SiLU())
self.task = nn.Embedding(int(config["task_count"]), 32)
self.difficulty = nn.Embedding(int(config["difficulty_count"]), 16)
self.head = nn.Sequential(
nn.Linear(384 + 128 + 128 + 32 + 16, 512), nn.LayerNorm(512), nn.SiLU(),
nn.Dropout(float(cfg["dropout"])), nn.Linear(512, 256), nn.SiLU(),
nn.Linear(256, self.chunk * 7))
def forward(self, images, proprio, text, task_id, difficulty_id):
if images.ndim != 4:
raise ValueError("image batch must have four dimensions")
if images.shape[-1] == 3:
images = images.permute(0, 3, 1, 2)
images = images.float()
if images.max() > 2:
images = images / 255.0
images = F.interpolate(images, size=(96, 96), mode="bilinear", align_corners=False)
fused = torch.cat([
self.vision(images), self.proprio(proprio.float()), self.text(text.float()),
self.task(task_id.clamp(0, self.task.num_embeddings - 1)),
self.difficulty(difficulty_id.clamp(0, self.difficulty.num_embeddings - 1)),
], dim=-1)
return torch.tanh(self.head(fused)).reshape(-1, self.chunk, 7)
def _text_vector(text, dim):
result = np.zeros(dim, dtype=np.float32)
for token in re.findall(r"[a-z0-9_]+", str(text).lower()):
value = int.from_bytes(hashlib.blake2b(token.encode(), digest_size=8).digest(), "little")
result[value % dim] += 1.0 if value & 1 else -1.0
norm = float(np.linalg.norm(result))
return result / norm if norm else result
def _proprio(obs, config):
value = np.asarray(obs.get("proprio", np.zeros(25, dtype=np.float32)), dtype=np.float32).reshape(-1)
if value.size == 21:
value = np.pad(value, (0, 4))
if value.size != 25:
raise ValueError(f"proprio must have 25 values, got {value.size}")
step = float(obs.get("step", 0))
horizon = float(obs.get("horizon", 320) or 320)
value = np.concatenate([value, np.asarray([step / max(horizon, 1.0)], dtype=np.float32)])
if value.size < int(config["proprio_dim"]):
value = np.pad(value, (0, int(config["proprio_dim"]) - value.size))
value = value[: int(config["proprio_dim"])]
mean = np.asarray(config["proprio_mean"], dtype=np.float32)
std = np.asarray(config["proprio_std"], dtype=np.float32)
return ((value - mean) / np.maximum(std, 1e-4)).astype(np.float32)
class Policy:
def __init__(self, model, config, device):
self.model, self.config, self.device = model, config, device
self.chunk = int(config["chunk"])
self.last_step = None
self.chunks = {}
self.last_action = np.zeros(7, dtype=np.float32)
self.task_to_id = config["task_to_id"]
self.difficulty_to_id = config["difficulty_to_id"]
@torch.inference_mode()
def act(self, obs):
image = np.asarray(obs.get("image", np.zeros((96, 96, 3), dtype=np.uint8)))
if image.ndim == 2:
image = np.repeat(image[..., None], 3, axis=-1)
if image.shape[-1] == 4:
image = image[..., :3]
if image.shape[-1] != 3:
raise ValueError("image must have 3 or 4 channels")
step = int(obs.get("step", 0))
if self.last_step is None or step == 0 or step != self.last_step + 1:
self.chunks = {}
if step == self.last_step:
return self.last_action.copy()
task = str(obs.get("task", ""))
difficulty = str(obs.get("difficulty", "") or "")
text = f"task {task} difficulty {difficulty} instruction {obs.get('instruction', '')}"
text_feature = _text_vector(text, int(self.config["text_dim"]))
proprio = _proprio(obs, self.config)
task_id = self.task_to_id.get(task, len(self.task_to_id))
difficulty_id = self.difficulty_to_id.get(difficulty, len(self.difficulty_to_id))
raw = self.model(
torch.from_numpy(image).unsqueeze(0).to(self.device),
torch.from_numpy(proprio).unsqueeze(0).to(self.device),
torch.from_numpy(text_feature).unsqueeze(0).to(self.device),
torch.tensor([task_id], device=self.device),
torch.tensor([difficulty_id], device=self.device),
)[0].float().cpu().numpy()
self.chunks[step] = raw
candidates, weights = [], []
for start, chunk in list(self.chunks.items()):
offset = step - start
if 0 <= offset < self.chunk:
candidates.append(chunk[offset])
weights.append(math.exp(-0.55 * offset))
else:
del self.chunks[start]
action = np.average(np.asarray(candidates), axis=0, weights=np.asarray(weights))
action = np.clip(action, -1.0, 1.0).astype(np.float32)
if action[6] > 0.12:
action[6] = 1.0
elif action[6] < -0.12:
action[6] = -1.0
self.last_step, self.last_action = step, action.copy()
return action
def load_policy(model_dir: str, device: str, dtype: str):
root = Path(model_dir)
config = json.loads((root / "vla_config.json").read_text())
selected = torch.device(device if str(device).startswith("cuda") and torch.cuda.is_available() else "cpu")
model = PolicyNet(config).to(selected)
checkpoint = torch.load(root / "model.pt", map_location=selected, weights_only=True)
state = checkpoint["state_dict"] if "state_dict" in checkpoint else checkpoint
model.load_state_dict(state, strict=True)
model.eval()
return Policy(model, config, selected)