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 |