File size: 7,282 Bytes
561d0f9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
"""Tactus inference: open-vocabulary object recognition from pressure-array data.

Tactus embeds a short window of 32x32 tactile pressure frames into the fusion-embedding
canonical 2048-d text space, so recognition is a cosine ranking against text queries rather
than a fixed classifier head.

Two halves, loaded from two places:
  * PRESSURE side (this repository): a trained CNN trunk + projector (``tactus_head.pt``),
    defined in the vendored ``tactile.py`` (byte-identical to the training-time module, so the
    checkpoint loads with ``strict=True``).
  * TEXT side (the base model repository): the canonical whitened text readout of
    ``EximiusLabs/fusion-embedding-2-2b-preview`` via the ``fusion-embedding`` package.
    Tactus was trained against that exact readout; embedding text any other way (including the
    plain base model without whitening) will misrank.

Input contract: frames are STAG-normalized pressure maps. If your data is raw sensor counts,
convert with ``stag_raw_to_unit`` (the dataset's own calibration: ``clip((raw-500)/150, 0, 1)``)
before embedding, or pass ``raw="stag"``. uint8 0-255 maps produced by the STAG caches are
accepted directly.

Usage:
    from inference import TactusEmbedder

    ta = TactusEmbedder.from_pretrained("EximiusLabs/fusion-embedding-2-tactus",
                                        revision="v0.1-preview")
    window = np.load("grasp.npy")            # [F,32,32] uint8/float, or [32,32]
    for text, score in ta.rank(window, ["a mug", "scissors", "a full soda can"]):
        print(f"{score:+.3f}  {text}")
"""
from __future__ import annotations

import os

import numpy as np
import torch
import torch.nn.functional as F

from tactile import TactileEncoder, build_projector, preprocess_pressure  # vendored, verbatim

REPO_DEFAULT = "EximiusLabs/fusion-embedding-2-tactus"
BASE_REPO = "EximiusLabs/fusion-embedding-2-2b-preview"
SAFETENSORS_FILE = "model.safetensors"      # head weights; config in the file's metadata
CKPT_FILE = "tactus_head.pt"                # legacy .pt with full training provenance

# STAG's own sensor calibration (classification/TouchDataset.py::transformPressure in the
# reference code): resting level raw ~500 counts, informative band [500, 650].
STAG_RAW_BASELINE = 500.0
STAG_RAW_SCALE = 150.0


def stag_raw_to_unit(raw) -> np.ndarray:
    """Raw pressure counts -> [0,1] floats via STAG's calibration affine."""
    a = np.asarray(raw, dtype=np.float32)
    return np.clip((a - STAG_RAW_BASELINE) / STAG_RAW_SCALE, 0.0, 1.0)


class TactusEmbedder:
    """Pressure -> 2048-d canonical embeddings, plus the matching text side."""

    def __init__(self, head_path: str, device: str = "cuda", dtype=torch.bfloat16,
                 base_repo: str = BASE_REPO, load_text: bool = True):
        self.device = device
        if head_path.endswith(".safetensors"):
            import json as _json

            import safetensors.torch as st
            from safetensors import safe_open
            flat = st.load_file(head_path)
            with safe_open(head_path, framework="pt") as f:
                cfg = _json.loads((f.metadata() or {})["config"])
            enc_sd = {k[len("encoder."):]: v for k, v in flat.items() if k.startswith("encoder.")}
            proj_sd = {k[len("proj."):]: v for k, v in flat.items() if k.startswith("proj.")}
        else:                                          # legacy .pt checkpoint
            blob = torch.load(head_path, map_location="cpu", weights_only=False)
            cfg = blob["config"]
            enc_sd = {k: v.float() for k, v in blob["encoder"].items()}
            proj_sd = {k: v.float() for k, v in blob["proj"].items()}
        self.cfg = cfg
        self.window_frames = int(cfg.get("window_frames", 8))

        self.enc = TactileEncoder(temporal=cfg.get("temporal", "conv"),
                                  depth=cfg.get("depth", "resnet18"),
                                  frames=self.window_frames,
                                  flow=bool(cfg.get("flow", False)))
        self.proj = build_projector()
        self.enc.load_state_dict(enc_sd, strict=True)
        self.proj.load_state_dict(proj_sd, strict=True)
        self.enc.eval().to(device)
        self.proj.eval().to(device)

        # The text side is the canonical whitened readout Tactus was trained against.
        self.ue = None
        if load_text:
            from fusion_embedding import UnifiedEmbedder   # pip install fusion-embedding[hf]
            self.ue = UnifiedEmbedder.from_pretrained(base_repo, device=device, dtype=dtype)

    @classmethod
    def from_pretrained(cls, repo: str = REPO_DEFAULT, revision: str | None = None,
                        device: str = "cuda", dtype=torch.bfloat16, load_text: bool = True):
        if os.path.isdir(repo):
            head = None
            for name in (SAFETENSORS_FILE, CKPT_FILE):
                for cand in (os.path.join(repo, name), os.path.join(repo, "out", name)):
                    if os.path.exists(cand):
                        head = cand
                        break
                if head:
                    break
            assert head, f"no {SAFETENSORS_FILE} or {CKPT_FILE} under {repo}"
        else:
            from huggingface_hub import hf_hub_download
            try:
                head = hf_hub_download(repo, SAFETENSORS_FILE, revision=revision)
            except Exception:  # noqa: BLE001 -- older revisions ship only the .pt
                head = hf_hub_download(repo, CKPT_FILE, revision=revision)
        return cls(head, device=device, dtype=dtype, load_text=load_text)

    @torch.no_grad()
    def embed_pressure(self, frames, raw: str | None = None) -> torch.Tensor:
        """One grasp window ``[F,32,32]`` (or a single ``[32,32]`` frame) -> 2048-d, L2-normalized.

        ``raw="stag"`` applies the sensor calibration affine first (input is raw counts).
        Otherwise frames must already be [0,1] floats or uint8 0-255 maps (the cache format)."""
        if raw == "stag":
            frames = stag_raw_to_unit(frames)
        elif raw is not None:
            raise ValueError(f"unknown raw mode {raw!r}; use raw='stag' or pre-normalize")
        x = preprocess_pressure(frames).to(self.device)        # [F,32,32] in [0,1]
        feat = self.enc(x)                                     # [512]
        vec = self.proj(feat)                                  # [2048] raw
        return F.normalize(vec.float(), dim=-1).cpu()

    @torch.no_grad()
    def embed_text(self, texts) -> torch.Tensor:
        """Text -> canonical 2048-d embeddings (whitened readout). Requires load_text=True."""
        assert self.ue is not None, "constructed with load_text=False"
        if isinstance(texts, str):
            texts = [texts]
        return torch.stack([self.ue.embed_text(t).float() for t in texts])

    @torch.no_grad()
    def rank(self, frames, texts, raw: str | None = None):
        """Rank candidate texts against one pressure window; returns [(text, cosine)] sorted."""
        p = self.embed_pressure(frames, raw=raw)               # [2048]
        t = F.normalize(self.embed_text(texts), dim=-1)        # [N,2048]
        sims = (t @ p).tolist()
        return sorted(zip(list(texts), sims), key=lambda kv: -kv[1])