[GLM-OCR] Support GLM-OCR Model (#17582)
Signed-off-by: Xinyuan Tong <xinyuantong.cs@gmail.com> Co-authored-by: Xinyuan Tong <115166877+JustinTong0323@users.noreply.github.com> Co-authored-by: Xinyuan Tong <xinyuantong.cs@gmail.com>
This commit is contained in:
@@ -590,6 +590,7 @@ class VisionAttention(nn.Module):
|
||||
num_dummy_heads: int = 0,
|
||||
qkv_bias: bool = True,
|
||||
qk_normalization: bool = False,
|
||||
qk_normalization_by_head_size: bool = False,
|
||||
layer_norm_eps: float = 1e-06,
|
||||
customized_position_embedding_applier: Callable[
|
||||
[torch.Tensor, torch.Tensor, Any, Any], Tuple[torch.Tensor, torch.Tensor]
|
||||
@@ -617,30 +618,19 @@ class VisionAttention(nn.Module):
|
||||
self.kv_size = self.num_attention_kv_heads_per_partition * self.head_size
|
||||
|
||||
self.qk_normalization = qk_normalization
|
||||
self.qk_normalization_by_head_size = qk_normalization_by_head_size
|
||||
|
||||
# Additional dummy heads are used to enable TP for common GPU counts.
|
||||
self.dummy_dim = (num_dummy_heads + num_heads) * self.head_size
|
||||
|
||||
if self.qk_normalization:
|
||||
norm_kwargs = (
|
||||
dict(
|
||||
weight_dtype=torch.float32,
|
||||
cast_x_before_out_mul=True,
|
||||
)
|
||||
if get_global_server_args().rl_on_policy_target is not None
|
||||
else {}
|
||||
self.q_norm, self.k_norm = self._init_qk_norm(
|
||||
self.dummy_dim, layer_norm_eps, embed_dim
|
||||
)
|
||||
self.q_norm = RMSNorm(
|
||||
self.dummy_dim,
|
||||
eps=layer_norm_eps,
|
||||
var_hidden_size=embed_dim,
|
||||
**norm_kwargs,
|
||||
)
|
||||
self.k_norm = RMSNorm(
|
||||
self.dummy_dim,
|
||||
eps=layer_norm_eps,
|
||||
var_hidden_size=embed_dim,
|
||||
**norm_kwargs,
|
||||
|
||||
elif self.qk_normalization_by_head_size:
|
||||
self.q_norm, self.k_norm = self._init_qk_norm(
|
||||
self.head_size, layer_norm_eps
|
||||
)
|
||||
|
||||
# Select attention backend via a unified method
|
||||
@@ -702,6 +692,31 @@ class VisionAttention(nn.Module):
|
||||
self.aux_stream = aux_stream
|
||||
self.ln_events = [torch.cuda.Event(), torch.cuda.Event()] if aux_stream else []
|
||||
|
||||
def _init_qk_norm(
|
||||
self, norm_dim: int, eps: float, var_hidden_size: Optional[int] = None
|
||||
):
|
||||
norm_kwargs = (
|
||||
dict(
|
||||
weight_dtype=torch.float32,
|
||||
cast_x_before_out_mul=True,
|
||||
)
|
||||
if get_global_server_args().rl_on_policy_target is not None
|
||||
else {}
|
||||
)
|
||||
q_norm = RMSNorm(
|
||||
norm_dim,
|
||||
eps=eps,
|
||||
var_hidden_size=var_hidden_size,
|
||||
**norm_kwargs,
|
||||
)
|
||||
k_norm = RMSNorm(
|
||||
norm_dim,
|
||||
eps=eps,
|
||||
var_hidden_size=var_hidden_size,
|
||||
**norm_kwargs,
|
||||
)
|
||||
return q_norm, k_norm
|
||||
|
||||
def _determine_attention_backend(self, passed_backend: Optional[str]) -> str:
|
||||
"""Decide the multimodal attention backend string.
|
||||
|
||||
@@ -734,6 +749,16 @@ class VisionAttention(nn.Module):
|
||||
|
||||
return backend
|
||||
|
||||
def _apply_qk_norm_head_size(self, q: torch.Tensor, k: torch.Tensor):
|
||||
"""apply qk norm for GLM-OCR vit attn"""
|
||||
q_by_head = q.reshape(-1, self.head_size)
|
||||
q_by_head = self.q_norm(q_by_head)
|
||||
k_by_head = k.reshape(-1, self.head_size)
|
||||
k_by_head = self.k_norm(k_by_head)
|
||||
q = q_by_head.view(q.shape)
|
||||
k = k_by_head.view(k.shape)
|
||||
return q, k
|
||||
|
||||
def _apply_qk_norm(self, q: torch.Tensor, k: torch.Tensor):
|
||||
"""apply qk norm for internvl vit attn"""
|
||||
|
||||
@@ -816,6 +841,8 @@ class VisionAttention(nn.Module):
|
||||
q = q.reshape(bsz * s, head, -1).contiguous()
|
||||
k = k.reshape(bsz * s, kv_head, -1).contiguous()
|
||||
v = v.reshape(bsz * s, kv_head, -1).contiguous()
|
||||
if self.qk_normalization_by_head_size:
|
||||
q, k = self._apply_qk_norm_head_size(q, k)
|
||||
else:
|
||||
# [b, s, embed_dim] --> [s, b, embed_dim]
|
||||
x = rearrange(x, "b s ... -> s b ...")
|
||||
@@ -837,6 +864,9 @@ class VisionAttention(nn.Module):
|
||||
rearrange(x, "s b ... -> b s ...").contiguous() for x in (q, k, v)
|
||||
]
|
||||
|
||||
if self.qk_normalization_by_head_size:
|
||||
q, k = self._apply_qk_norm_head_size(q, k)
|
||||
|
||||
cos = None
|
||||
sin = None
|
||||
|
||||
@@ -881,7 +911,7 @@ class VisionAttention(nn.Module):
|
||||
assert v.dim() == 3, v.dim()
|
||||
|
||||
# internvl
|
||||
if self.qk_normalization:
|
||||
if self.qk_normalization and not self.qk_normalization_by_head_size:
|
||||
# jit kernel
|
||||
if can_use_jit_qk_norm(self.head_size, q.dtype):
|
||||
|
||||
|
||||
Reference in New Issue
Block a user