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()