LLaVA-UHD-v4-8M-Baseline / downsample_mlp.py
PhoenixGS's picture
Add config and tokenizer files
9465d36
Raw
History Blame Contribute Delete
2.67 kB
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