File size: 2,222 Bytes
60ba429
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from typing import List, Optional, Sequence, Tuple

import numpy as np


def seq_len_from_output(output: np.ndarray) -> Optional[int]:
    if output.ndim < 2:
        return None
    if output.ndim == 2:
        return int(output.shape[0])
    return int(output.shape[-2])


def normalize_vit_output(
    output: np.ndarray,
    target_hidden_size: int,
    expected_tokens: Optional[int] = None,
) -> np.ndarray:
    normalized = output
    if expected_tokens is not None:
        if normalized.ndim == 3 and normalized.shape[1] == target_hidden_size and normalized.shape[2] == expected_tokens:
            normalized = np.transpose(normalized, (0, 2, 1))
        elif normalized.ndim == 2 and normalized.shape[0] == target_hidden_size and normalized.shape[1] == expected_tokens:
            normalized = np.transpose(normalized, (1, 0))
    return normalized


def describe_output_shapes(outputs: Sequence[np.ndarray]) -> List[Tuple[int, ...]]:
    return [tuple(int(v) for v in output.shape) for output in outputs]


def select_vit_output(
    outputs: Sequence[np.ndarray],
    target_hidden_size: int,
    expected_tokens: Optional[int] = None,
) -> np.ndarray:
    normalized_outputs = [
        normalize_vit_output(output, target_hidden_size, expected_tokens=expected_tokens) for output in outputs
    ]

    image_embeds = None
    if expected_tokens is not None:
        for output in normalized_outputs:
            if output.ndim >= 2 and seq_len_from_output(output) == expected_tokens and output.shape[-1] == target_hidden_size:
                image_embeds = output
                break
        if image_embeds is None:
            for output in normalized_outputs:
                if output.ndim >= 2 and seq_len_from_output(output) == expected_tokens:
                    image_embeds = output
                    break

    if image_embeds is None:
        for output in normalized_outputs:
            if output.ndim >= 2 and output.shape[-1] == target_hidden_size:
                image_embeds = output
                break

    if image_embeds is None:
        image_embeds = normalized_outputs[0]
    if image_embeds.ndim == 2:
        image_embeds = image_embeds[None, ...]
    return image_embeds