File size: 3,740 Bytes
fa5242f 2457010 fa5242f 2457010 fa5242f 2457010 fa5242f | 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 | """
Minimal loader for the STP release on HuggingFace.
import torch
from loader import STPEncoder
enc = STPEncoder.from_pretrained('yangjiange/STP')
feats = enc.encode(images) # (N, 3, 224, 224) -> (N, 768)
Or straight from the hub without this file:
import torch, models_stp
from huggingface_hub import hf_hub_download
path = hf_hub_download('yangjiange/STP', 'stp_vitb.pth')
model = models_stp.mae_vit_base_patch16()
from util.ckpt_io import load_model_state
state, _ = load_model_state(path)
model.load_state_dict(state, strict=True)
model.eval()
"""
import os
import numpy as np
import torch
from PIL import Image
IMAGENET_MEAN = (0.485, 0.456, 0.406)
IMAGENET_STD = (0.229, 0.224, 0.225)
class STPEncoder:
"""Frozen STP encoder: images in, [CLS] representations out.
Only the encoder is used. The spatial and temporal decoders exist for
pre-training and are irrelevant once you are training a policy.
"""
def __init__(self, model, device='cpu', image_size=224):
self.model = model
self.device = device
self.image_size = image_size
# ------------------------------------------------------------------
@classmethod
def from_pretrained(cls, repo_id_or_path='yangjiange/STP', filename='stp_vitb.pth',
device=None, repo_root=None):
import models_stp
device = device or ('cuda' if torch.cuda.is_available() else 'cpu')
if os.path.isfile(repo_id_or_path):
path = repo_id_or_path
else:
from huggingface_hub import hf_hub_download
path = hf_hub_download(repo_id=repo_id_or_path, filename=filename,
cache_dir=None)
import sys as _sys, os as _os
_root = repo_root or _os.path.dirname(_os.path.dirname(_os.path.abspath(__file__)))
if _root not in _sys.path:
_sys.path.insert(0, _root)
from util.ckpt_io import load_model_state
state, _ = load_model_state(path)
if repo_root and repo_root not in os.sys.path:
os.sys.path.insert(0, repo_root)
model = models_stp.mae_vit_base_patch16()
model.load_state_dict(state, strict=True)
model.eval().to(device)
return cls(model, device)
# ------------------------------------------------------------------
def preprocess(self, img):
if not isinstance(img, Image.Image):
img = Image.fromarray(np.asarray(img))
if img.size != (self.image_size, self.image_size):
img = img.resize((self.image_size, self.image_size), Image.BICUBIC)
x = torch.from_numpy(np.asarray(img.convert('RGB'), dtype=np.float32) / 255.0)
x = x.permute(2, 0, 1)
mean = torch.tensor(IMAGENET_MEAN).view(3, 1, 1)
std = torch.tensor(IMAGENET_STD).view(3, 1, 1)
return (x - mean) / std
@torch.no_grad()
def encode(self, images, batch_size=32):
"""PIL images / numpy arrays / a tensor -> (N, 768) numpy array."""
if isinstance(images, torch.Tensor):
out = []
for i in range(0, len(images), batch_size):
out.append(self.model.forward_features(
images[i:i + batch_size].to(self.device)).float().cpu())
return torch.cat(out, 0).numpy()
if isinstance(images, Image.Image):
images = [images]
feats = []
for i in range(0, len(images), batch_size):
batch = torch.stack([self.preprocess(im) for im in images[i:i + batch_size]])
feats.append(self.model.forward_features(batch.to(self.device)).float().cpu())
return torch.cat(feats, 0).numpy()
|