| 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 |
| |
| |
| |
| 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 = 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) |
| |
| |
| processed_features.append(_hidden_state) |
|
|
| return processed_features |