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