[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:
Yuxuan Zhang
2026-01-27 14:24:00 +08:00
committed by GitHub
parent 81c0f5c5ad
commit 7106f6c8e1
9 changed files with 679 additions and 29 deletions

View File

@@ -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):