Views
No views yet
| Metric | Value |
|---|---|
| Val Accuracy (exact match) | 80.0% |
| Val Angular Error | 3.2 degrees |
| Angle Bins | 36 (10-degree steps) |
1from prismatic import load
2
3vlm = load("path/to/minivla-angle-selector")1from prismatic.models.materialize import (
2 get_vision_backbone_and_transform,
3 get_llm_backbone_and_tokenizer,
4 get_vlm,
5)
6import torch
7
8# Build model
9vision_backbone, _ = get_vision_backbone_and_transform(
10 "dinosiglip-vit-so-224px", "resize-naive", image_sequence_len=1
11)
12llm_backbone, tokenizer = get_llm_backbone_and_tokenizer(
13 "qwen25-0_5b-pure", llm_max_length=2048, inference_mode=True,
14)
15vlm = get_vlm(
16 model_id="minivla-angle-selector",
17 arch_specifier="no-align+fused-gelu-mlp",
18 vision_backbone=vision_backbone,
19 llm_backbone=llm_backbone,
20)
21
22# Load weights
23ckpt = torch.load("checkpoints/latest-checkpoint.pt", map_location="cpu")["model"]
24vlm.projector.load_state_dict(ckpt["projector"])
25vlm.llm_backbone.load_state_dict(ckpt["llm_backbone"])
26vlm.vision_backbone.load_state_dict(ckpt["vision_backbone"])
27vlm.to("cuda", dtype=torch.bfloat16)
28vlm.eval()
29
30# Predict angle from image
31from PIL import Image
32image = Image.open("drone_view.png")
33prompt_builder = vlm.get_prompt_builder()
34prompt_builder.add_turn("human", "Navigate the drone to the red cube")
35input_prompt = prompt_builder.get_prompt()
36
37tok = vlm.llm_backbone.tokenizer
38input_ids = tok(input_prompt, return_tensors="pt").input_ids.to("cuda")
39pixel_values = vlm.vision_backbone.get_image_transform()(image)
40pixel_values = pixel_values[None, ...].to("cuda", dtype=torch.bfloat16)
41
42with torch.no_grad():
43 output = vlm.forward(input_ids=input_ids, pixel_values=pixel_values, return_dict=True)
44 num_patches = vlm.vision_backbone.num_patches
45 action_logit = output.logits[0, num_patches:, :][-1, :]
46 token_id = action_logit.argmax().item()
47
48# Convert token to angle
49vocab_size = len(tok)
50angle_code = (vocab_size - 1 - token_id) % 36
51angle_degrees = angle_code * 10
52print(f"Predicted angle: {angle_degrees} degrees")| Code | Angle | Direction |
|---|---|---|
| 0 | 0 deg | +X (right) |
| 9 | 90 deg | +Y (forward) |
| 18 | 180 deg | -X (left) |
| 27 | 270 deg | -Y (backward) |
token_id = vocab_size - 1 - angle_code