YellowCab / model.py
1ouai's picture
Release YellowCab under Apache-2.0
fcb9b70
Raw
History Blame Contribute Delete
8.42 kB
"""Exact architecture and integrity-checked loader for YellowCab v0."""
from __future__ import annotations
import hashlib
import json
from collections import Counter
from pathlib import Path
from typing import Any, Mapping
import torch
from safetensors.torch import load_file
from torch import nn
from torchvision.models import efficientnet_b0
class CheckpointIntegrityError(RuntimeError):
"""Raised when a release checkpoint does not match its signed-off contract."""
class TemporalFusionHead(nn.Module):
"""The trained temporal and telemetry fusion head."""
def __init__(
self,
*,
image_feature_dim: int,
telemetry_dim: int,
num_classes: int,
hidden_dim: int = 256,
dropout: float = 0.25,
) -> None:
super().__init__()
self.image_feature_dim = image_feature_dim
self.image_norm = nn.LayerNorm(image_feature_dim)
self.temporal = nn.GRU(
input_size=image_feature_dim,
hidden_size=hidden_dim,
batch_first=True,
)
self.telemetry = nn.Sequential(
nn.LayerNorm(telemetry_dim),
nn.Linear(telemetry_dim, 96),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(96, 96),
nn.GELU(),
)
self.classifier = nn.Sequential(
nn.LayerNorm(hidden_dim + 96),
nn.Linear(hidden_dim + 96, hidden_dim),
nn.GELU(),
nn.Dropout(dropout),
nn.Linear(hidden_dim, num_classes),
)
def forward(
self,
image_features: torch.Tensor,
telemetry: torch.Tensor,
) -> torch.Tensor:
temporal_output, _ = self.temporal(self.image_norm(image_features))
telemetry_output = self.telemetry(telemetry)
fused = torch.cat((temporal_output[:, -1], telemetry_output), dim=-1)
return self.classifier(fused)
class TaxiManeuverModel(nn.Module):
"""Frozen EfficientNet-B0 plus the trained v0 fusion head."""
def __init__(self, head_config: Mapping[str, Any]) -> None:
super().__init__()
self.encoder = efficientnet_b0(weights=None)
self.encoder.classifier = nn.Identity()
self.head = TemporalFusionHead(**dict(head_config))
def forward(
self,
images: torch.Tensor,
telemetry: torch.Tensor,
) -> torch.Tensor:
if images.ndim != 5 or images.shape[1] != 3 or images.shape[2] != 3:
raise ValueError("images must have shape [batch, 3, 3, height, width]")
batch_size, frame_count, channels, height, width = images.shape
flat_images = images.reshape(
batch_size * frame_count,
channels,
height,
width,
)
image_features = self.encoder(flat_images).reshape(
batch_size,
frame_count,
-1,
)
return self.head(image_features, telemetry)
def load_release_config(repository_path: str | Path) -> dict[str, Any]:
"""Load and minimally validate the one canonical release configuration."""
root = Path(repository_path)
path = root / "config.json"
try:
config = json.loads(path.read_text(encoding="utf-8"))
except (OSError, json.JSONDecodeError) as exc:
raise CheckpointIntegrityError("release configuration is unavailable") from exc
if config.get("schema") != "gdc_taxi_maneuver_release_config_v1":
raise CheckpointIntegrityError("unsupported release configuration schema")
for identity_field in ("model_id", "model_name", "model_version"):
identity_value = config.get(identity_field)
if not isinstance(identity_value, str) or not identity_value.strip():
raise CheckpointIntegrityError(
f"release configuration {identity_field} is invalid"
)
labels = config.get("labels")
feature_order = config.get("telemetry", {}).get("feature_order")
if not isinstance(labels, list) or len(labels) != 5 or len(set(labels)) != 5:
raise CheckpointIntegrityError("canonical label configuration is invalid")
if (
not isinstance(feature_order, list)
or len(feature_order) != 12
or len(set(feature_order)) != 12
):
raise CheckpointIntegrityError("canonical telemetry configuration is invalid")
return config
def sha256_file(path: str | Path) -> str:
digest = hashlib.sha256()
with Path(path).open("rb") as handle:
for chunk in iter(lambda: handle.read(1024 * 1024), b""):
digest.update(chunk)
return digest.hexdigest()
def _verify_checkpoint_metadata(
*,
checkpoint_path: Path,
weights: Mapping[str, torch.Tensor],
model: nn.Module,
checkpoint_config: Mapping[str, Any],
) -> None:
expected_bytes = int(checkpoint_config["bytes"])
expected_sha256 = str(checkpoint_config["sha256"])
if checkpoint_path.stat().st_size != expected_bytes:
raise CheckpointIntegrityError("checkpoint byte size does not match config.json")
if sha256_file(checkpoint_path) != expected_sha256:
raise CheckpointIntegrityError("checkpoint SHA-256 does not match config.json")
expected_state = model.state_dict()
actual_keys = set(weights)
expected_keys = set(expected_state)
missing = sorted(expected_keys - actual_keys)
unexpected = sorted(actual_keys - expected_keys)
if missing or unexpected:
raise CheckpointIntegrityError(
"checkpoint tensor keys do not exactly match the model architecture"
)
for name, expected_tensor in expected_state.items():
actual_tensor = weights[name]
if actual_tensor.shape != expected_tensor.shape:
raise CheckpointIntegrityError(
f"checkpoint tensor shape mismatch for {name}"
)
if actual_tensor.dtype != expected_tensor.dtype:
raise CheckpointIntegrityError(
f"checkpoint tensor dtype mismatch for {name}"
)
actual_tensor_count = len(weights)
actual_state_numel = sum(tensor.numel() for tensor in weights.values())
actual_dtype_counts = Counter(str(tensor.dtype) for tensor in weights.values())
expected_dtype_counts = {
str(name): int(count)
for name, count in checkpoint_config["state_dtypes"].items()
}
if actual_tensor_count != int(checkpoint_config["state_tensor_count"]):
raise CheckpointIntegrityError("checkpoint tensor count does not match config.json")
if actual_state_numel != int(checkpoint_config["state_numel"]):
raise CheckpointIntegrityError("checkpoint state size does not match config.json")
if dict(actual_dtype_counts) != expected_dtype_counts:
raise CheckpointIntegrityError("checkpoint dtype inventory does not match config.json")
parameter_count = sum(parameter.numel() for parameter in model.parameters())
if parameter_count != int(checkpoint_config["parameter_count"]):
raise CheckpointIntegrityError(
"model parameter count does not match config.json"
)
def load_model(
repository_path: str | Path,
*,
device: torch.device | str = "cpu",
) -> tuple[TaxiManeuverModel, dict[str, Any]]:
"""Load the checkpoint only after hash, keys, shapes and dtypes all match."""
root = Path(repository_path)
config = load_release_config(root)
checkpoint_config = config["checkpoint"]
checkpoint_path = root / str(checkpoint_config["file"])
try:
weights = load_file(str(checkpoint_path), device="cpu")
except (OSError, ValueError) as exc:
raise CheckpointIntegrityError("checkpoint cannot be read") from exc
model = TaxiManeuverModel(config["architecture"]["head_config"])
_verify_checkpoint_metadata(
checkpoint_path=checkpoint_path,
weights=weights,
model=model,
checkpoint_config=checkpoint_config,
)
try:
model.load_state_dict(weights, strict=True)
except RuntimeError as exc:
raise CheckpointIntegrityError("strict checkpoint loading failed") from exc
model.eval()
model.requires_grad_(False)
model.to(torch.device(device))
return model, config
__all__ = [
"CheckpointIntegrityError",
"TaxiManeuverModel",
"TemporalFusionHead",
"load_model",
"load_release_config",
"sha256_file",
]