File size: 2,671 Bytes
9465d36
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
from functools import partial
import numpy as np

import torch
from torch import nn
from torch.nn.init import trunc_normal_
from typing import Tuple
from einops import rearrange


class DownsampleMLP(nn.Module):
    def __init__(self, hidden_size, llm_embed_dim, merge_kernel_size=(2, 2)):
        super().__init__()
        self.merge_kernel_size = merge_kernel_size

        self.hidden_size = (
            hidden_size
            * self.merge_kernel_size[0]
            * self.merge_kernel_size[1]
        )

        self.pre_norm = torch.nn.LayerNorm(self.hidden_size, eps=1e-6)

        self.mlp = nn.Sequential(
            nn.Linear(self.hidden_size, self.hidden_size, bias=True),
            nn.GELU(),
            nn.Linear(self.hidden_size, llm_embed_dim, bias=True)
        )

    def forward(self, x: torch.Tensor) -> torch.Tensor:
        x = self.mlp(self.pre_norm(x).view(-1, self.hidden_size))
        return x


class Merger(nn.Module):
    def __init__(self, hidden_size, llm_embed_dim, merge_kernel_size=(2, 2), times=1):
        super().__init__()
        self.merge_kernel_size = merge_kernel_size
        self.times = times
        self.mlp = nn.ModuleList([DownsampleMLP(hidden_size, llm_embed_dim if i==times-1 else hidden_size, merge_kernel_size) for i in range(times)])


    def forward(self, hidden_states: torch.Tensor, tgt_sizes: torch.IntTensor,) -> torch.Tensor:
        m1, m2 = self.merge_kernel_size
        
        # batch_size, _, _ = hidden_states.shape
        # start = 0
        processed_features = []
        for batch_idx in range(len(tgt_sizes)):
            h, w = tgt_sizes[batch_idx]
            num_patches = h * w

            _hidden_state = rearrange(hidden_states[batch_idx, 0: num_patches, :], "(h p1 w p2) d -> (h w) (p1 p2 d)", h=h // m1, p1=m1, w=w // m2, p2=m2)
            # _hidden_state = rearrange(hidden_states[0, start: start+num_patches, :], "(h p1 w p2) d -> (h w) (p1 p2 d)", h=h // m1, p1=m1, w=w // m2, p2=m2)
            _hidden_state = self.mlp[0](_hidden_state)

            if self.times > 1:
                for i in range(1, self.times):
                    assert h % self.merge_kernel_size[0] == 0 and w % self.merge_kernel_size[1] == 0, "patch尺寸不能被2整除,无法拼接4个相邻patch"
                    h = h // 2
                    w = w // 2

                    _hidden_state = rearrange(_hidden_state, "(h p1 w p2) d -> (h w) (p1 p2 d)", h=h // m1, p1=m1, w=w // m2, p2=m2)
                    _hidden_state = self.mlp[i](_hidden_state)
            
            # start += num_patches
            processed_features.append(_hidden_state)

        return processed_features