refactor: move image processors to separate files (#4229)
This commit is contained in:
@@ -1,7 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from functools import lru_cache
|
||||
from typing import Optional
|
||||
from typing import Optional, Tuple
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
@@ -22,47 +22,29 @@ from sglang.srt.layers.quantization import QuantizationConfig
|
||||
from sglang.srt.utils import add_prefix
|
||||
|
||||
|
||||
def rotate_half(x: torch.Tensor, interleaved: bool = False) -> torch.Tensor:
|
||||
if not interleaved:
|
||||
x1, x2 = x.chunk(2, dim=-1)
|
||||
return torch.cat((-x2, x1), dim=-1)
|
||||
else:
|
||||
x1, x2 = x[..., ::2], x[..., 1::2]
|
||||
return rearrange(
|
||||
torch.stack((-x2, x1), dim=-1), "... d two -> ... (d two)", two=2
|
||||
)
|
||||
# Copied from transformers, modeling_qwen2_vl.py
|
||||
def rotate_half(x):
|
||||
"""Rotates half the hidden dims of the input."""
|
||||
x1 = x[..., : x.shape[-1] // 2]
|
||||
x2 = x[..., x.shape[-1] // 2 :]
|
||||
return torch.cat((-x2, x1), dim=-1)
|
||||
|
||||
|
||||
def apply_rotary_emb_torch(
|
||||
x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor, interleaved: bool = False
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
x: (batch_size, seqlen, nheads, headdim)
|
||||
cos, sin: (seqlen, rotary_dim / 2) or (batch_size, seqlen, rotary_dim / 2)
|
||||
"""
|
||||
ro_dim = cos.shape[-1] * 2
|
||||
assert ro_dim <= x.shape[-1]
|
||||
cos = repeat(
|
||||
cos, "... d -> ... 1 (2 d)" if not interleaved else "... d -> ... 1 (d 2)"
|
||||
)
|
||||
sin = repeat(
|
||||
sin, "... d -> ... 1 (2 d)" if not interleaved else "... d -> ... 1 (d 2)"
|
||||
)
|
||||
return torch.cat(
|
||||
[
|
||||
x[..., :ro_dim] * cos + rotate_half(x[..., :ro_dim], interleaved) * sin,
|
||||
x[..., ro_dim:],
|
||||
],
|
||||
dim=-1,
|
||||
)
|
||||
def apply_rotary_pos_emb_vision(
|
||||
q: torch.Tensor, k: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor
|
||||
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
orig_q_dtype = q.dtype
|
||||
orig_k_dtype = k.dtype
|
||||
q, k = q.float(), k.float()
|
||||
|
||||
cos, sin = cos.unsqueeze(-2).float(), sin.unsqueeze(-2).float()
|
||||
q_embed = (q * cos) + (rotate_half(q) * sin)
|
||||
k_embed = (k * cos) + (rotate_half(k) * sin)
|
||||
|
||||
def apply_rotary_pos_emb_vision(t: torch.Tensor, freqs: torch.Tensor) -> torch.Tensor:
|
||||
t_ = t.float()
|
||||
cos = freqs.cos()
|
||||
sin = freqs.sin()
|
||||
output = apply_rotary_emb_torch(t_, cos, sin).type_as(t)
|
||||
return output
|
||||
q_embed = q_embed.to(orig_q_dtype)
|
||||
k_embed = k_embed.to(orig_k_dtype)
|
||||
|
||||
return q_embed, k_embed
|
||||
|
||||
|
||||
class VisionAttention(nn.Module):
|
||||
@@ -75,8 +57,8 @@ class VisionAttention(nn.Module):
|
||||
use_context_forward (bool, default to True):
|
||||
if ``True``, a flash_attn style attention will be applied
|
||||
Otherwise, a full-sequence attention will be applied.
|
||||
use_full_precision_softmax (bool, default to False):
|
||||
if ``True``, the softmax will be performed in full-precision
|
||||
softmax_in_single_precision (bool, default to False):
|
||||
if ``True``, the softmax will be performed in single-precision
|
||||
Otherwise, it will be performed in half-precision
|
||||
|
||||
"""
|
||||
@@ -90,7 +72,7 @@ class VisionAttention(nn.Module):
|
||||
quant_config: Optional[QuantizationConfig] = None,
|
||||
dropout: float = 0.0,
|
||||
use_context_forward: bool = True,
|
||||
use_full_precision_softmax: bool = False,
|
||||
softmax_in_single_precision: bool = False,
|
||||
flatten_batch: bool = False,
|
||||
prefix: str = "",
|
||||
):
|
||||
@@ -113,7 +95,7 @@ class VisionAttention(nn.Module):
|
||||
head_size=self.head_size,
|
||||
dropout=dropout,
|
||||
flatten_batch=flatten_batch,
|
||||
use_full_precision_softmax=use_full_precision_softmax,
|
||||
softmax_in_single_precision=softmax_in_single_precision,
|
||||
)
|
||||
|
||||
self.use_qkv_parallel = use_qkv_parallel
|
||||
@@ -143,7 +125,7 @@ class VisionAttention(nn.Module):
|
||||
self,
|
||||
x: torch.Tensor,
|
||||
cu_seqlens: Optional[torch.Tensor] = None,
|
||||
rotary_pos_emb: torch.Tensor = None,
|
||||
position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||
attention_mask: Optional[torch.Tensor] = None,
|
||||
) -> torch.Tensor:
|
||||
r"""
|
||||
@@ -151,21 +133,17 @@ class VisionAttention(nn.Module):
|
||||
x: [b, s, embed_dim]
|
||||
cu_seqlens: [b]
|
||||
Returns:
|
||||
[s, b, num_heads * head]
|
||||
[s, b, head * head_size]
|
||||
"""
|
||||
bsz, s, _ = x.shape
|
||||
head = self.num_attention_heads_per_partition
|
||||
if self.use_qkv_parallel:
|
||||
# [b, s, embed_dim] --> [b, s, embed_dim]
|
||||
qkv, _ = self.qkv_proj(x)
|
||||
q, k, v = qkv.chunk(3, dim=-1)
|
||||
|
||||
# [b, s, embed_dim] --> [b * s, num_heads, head_size]
|
||||
q, k, v = [
|
||||
x.reshape(
|
||||
bsz * s, self.num_attention_heads_per_partition, -1
|
||||
).contiguous()
|
||||
for x in (q, k, v)
|
||||
]
|
||||
# [b, s, embed_dim] --> [b * s, head, head_size]
|
||||
q, k, v = [x.reshape(bsz * s, head, -1).contiguous() for x in (q, k, v)]
|
||||
else:
|
||||
# [b, s, embed_dim] --> [s, b, embed_dim]
|
||||
x = rearrange(x, "b s ... -> s b ...")
|
||||
@@ -173,7 +151,7 @@ class VisionAttention(nn.Module):
|
||||
qkv, _ = self.qkv_proj(x)
|
||||
# [s, b, head * 3 * head_size] --> [s, b, head, 3 * head_size]
|
||||
new_x_shape = qkv.size()[:-1] + (
|
||||
self.num_attention_heads_per_partition,
|
||||
head,
|
||||
3 * self.hidden_size_per_attention_head,
|
||||
)
|
||||
qkv = qkv.view(*new_x_shape)
|
||||
@@ -186,9 +164,12 @@ class VisionAttention(nn.Module):
|
||||
rearrange(x, "s b ... -> b s ...").contiguous() for x in (q, k, v)
|
||||
]
|
||||
|
||||
if rotary_pos_emb is not None:
|
||||
q = apply_rotary_pos_emb_vision(q, rotary_pos_emb)
|
||||
k = apply_rotary_pos_emb_vision(k, rotary_pos_emb)
|
||||
if position_embeddings is not None:
|
||||
cos, sin = position_embeddings
|
||||
original_shape = q.shape
|
||||
q, k = q.view(s, head, -1), k.view(s, head, -1)
|
||||
q, k = apply_rotary_pos_emb_vision(q, k, cos, sin)
|
||||
q, k = q.reshape(original_shape), k.reshape(original_shape)
|
||||
|
||||
if self.use_qkv_parallel:
|
||||
pass
|
||||
@@ -230,12 +211,12 @@ class VisionSdpaAttention(nn.Module):
|
||||
head_size: int,
|
||||
dropout: float = 0.0,
|
||||
flatten_batch: bool = False,
|
||||
use_full_precision_softmax: bool = False,
|
||||
softmax_in_single_precision: bool = False,
|
||||
):
|
||||
super().__init__()
|
||||
self.head_size = head_size
|
||||
self.flatten_batch = flatten_batch
|
||||
self.use_full_precision_softmax = use_full_precision_softmax
|
||||
self.softmax_in_single_precision = softmax_in_single_precision
|
||||
self.dropout = dropout
|
||||
|
||||
@staticmethod
|
||||
@@ -319,14 +300,14 @@ class VisionSdpaAttention(nn.Module):
|
||||
)
|
||||
|
||||
if attention_mask is None:
|
||||
if self.use_full_precision_softmax:
|
||||
if self.softmax_in_single_precision:
|
||||
raise RuntimeError("Empty attention mask")
|
||||
else:
|
||||
attention_mask = attention_mask.to(device=q.device)
|
||||
|
||||
q, k, v = [rearrange(x, "(b s) h d -> b h s d", b=bsz) for x in [q, k, v]]
|
||||
|
||||
if self.use_full_precision_softmax:
|
||||
if self.softmax_in_single_precision:
|
||||
scale = self.head_size**-0.5
|
||||
k_transposed = rearrange(k, "b h s d -> b h d s")
|
||||
attn_weights = torch.matmul(q, k_transposed) * scale
|
||||
|
||||
Reference in New Issue
Block a user