[vlm] fix: preserve per-request vit graph metadata for qwen-vl (#37043)
This commit is contained in:
@@ -616,6 +616,11 @@ class Qwen2_5_VisionTransformer(nn.Module, RotaryPosMixin):
|
|||||||
rotary_pos_emb = self.rot_pos_emb(grid_thw)
|
rotary_pos_emb = self.rot_pos_emb(grid_thw)
|
||||||
|
|
||||||
window_index, cu_window_seqlens = self.get_window_index(grid_thw)
|
window_index, cu_window_seqlens = self.get_window_index(grid_thw)
|
||||||
|
cu_window_layout = tuple(
|
||||||
|
value
|
||||||
|
for index, value in enumerate(cu_window_seqlens)
|
||||||
|
if index == 0 or value != cu_window_seqlens[index - 1]
|
||||||
|
)
|
||||||
cu_window_seqlens = torch.tensor(
|
cu_window_seqlens = torch.tensor(
|
||||||
cu_window_seqlens,
|
cu_window_seqlens,
|
||||||
device=x.device,
|
device=x.device,
|
||||||
@@ -658,6 +663,11 @@ class Qwen2_5_VisionTransformer(nn.Module, RotaryPosMixin):
|
|||||||
]
|
]
|
||||||
)
|
)
|
||||||
cu_seqlens = torch.cat([cu_seqlens.new_zeros(1), cu_seqlens])
|
cu_seqlens = torch.cat([cu_seqlens.new_zeros(1), cu_seqlens])
|
||||||
|
full_layout = [0, 0]
|
||||||
|
total_tokens = 0
|
||||||
|
for temporal, height, width in grid_thw.tolist():
|
||||||
|
total_tokens += temporal * height * width
|
||||||
|
full_layout.append(total_tokens)
|
||||||
|
|
||||||
return self.cuda_graph_runner.run(
|
return self.cuda_graph_runner.run(
|
||||||
x=x,
|
x=x,
|
||||||
@@ -665,6 +675,7 @@ class Qwen2_5_VisionTransformer(nn.Module, RotaryPosMixin):
|
|||||||
cu_seqlens=cu_seqlens,
|
cu_seqlens=cu_seqlens,
|
||||||
cu_window_seqlens=cu_window_seqlens,
|
cu_window_seqlens=cu_window_seqlens,
|
||||||
output_indices=reverse_indices,
|
output_indices=reverse_indices,
|
||||||
|
attention_layout_key=(tuple(full_layout), cu_window_layout),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1053,6 +1053,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
|
|||||||
rotary_pos_emb_cos,
|
rotary_pos_emb_cos,
|
||||||
rotary_pos_emb_sin,
|
rotary_pos_emb_sin,
|
||||||
) = self._prepare_graph_inputs(x, grid_thw)
|
) = self._prepare_graph_inputs(x, grid_thw)
|
||||||
|
attention_layout_key = (tuple(cu_seqlens.tolist()), None)
|
||||||
if not isinstance(cu_seqlens, torch.Tensor):
|
if not isinstance(cu_seqlens, torch.Tensor):
|
||||||
cu_seqlens = torch.tensor(cu_seqlens, device=x.device, dtype=torch.int32)
|
cu_seqlens = torch.tensor(cu_seqlens, device=x.device, dtype=torch.int32)
|
||||||
else:
|
else:
|
||||||
@@ -1067,6 +1068,7 @@ class Qwen3VLMoeVisionModel(nn.Module, RotaryPosMixin):
|
|||||||
cu_seqlens=cu_seqlens,
|
cu_seqlens=cu_seqlens,
|
||||||
cu_window_seqlens=None,
|
cu_window_seqlens=None,
|
||||||
output_indices=None,
|
output_indices=None,
|
||||||
|
attention_layout_key=attention_layout_key,
|
||||||
)
|
)
|
||||||
|
|
||||||
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
|
def load_weights(self, weights: Iterable[tuple[str, torch.Tensor]]) -> set[str]:
|
||||||
|
|||||||
@@ -60,8 +60,13 @@ class ViTCudaGraphRunner:
|
|||||||
self.cu_full_len_kk: Dict[Hashable, torch.Tensor] = {}
|
self.cu_full_len_kk: Dict[Hashable, torch.Tensor] = {}
|
||||||
self.cu_window_len_kk: Dict[Hashable, torch.Tensor] = {}
|
self.cu_window_len_kk: Dict[Hashable, torch.Tensor] = {}
|
||||||
|
|
||||||
# rotary position buffers shared across graphs
|
# Current rotary workspace plus older allocations retained by graphs
|
||||||
|
# captured before the workspace grew.
|
||||||
self.sin_cos_ws: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
|
self.sin_cos_ws: Optional[Tuple[torch.Tensor, torch.Tensor]] = None
|
||||||
|
self._retired_sin_cos_ws: List[Tuple[torch.Tensor, torch.Tensor]] = []
|
||||||
|
self._sin_cos_ws_by_graph: Dict[Hashable, Tuple[torch.Tensor, torch.Tensor]] = (
|
||||||
|
{}
|
||||||
|
)
|
||||||
self.max_context_len = getattr(vit, "max_context_len", None)
|
self.max_context_len = getattr(vit, "max_context_len", None)
|
||||||
|
|
||||||
# Qwen2.5-VL specific viarable.
|
# Qwen2.5-VL specific viarable.
|
||||||
@@ -91,31 +96,67 @@ class ViTCudaGraphRunner:
|
|||||||
def dtype(self) -> torch.dtype:
|
def dtype(self) -> torch.dtype:
|
||||||
return self.vit.dtype
|
return self.vit.dtype
|
||||||
|
|
||||||
def _ensure_sin_cos_ws(self, seq_len: int, head_dim: int):
|
def _get_sin_cos_ws(
|
||||||
if self.sin_cos_ws is None:
|
self, graph_key: Hashable, seq_len: int, head_dim: int
|
||||||
max_shape = self.max_context_len or seq_len
|
) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||||
max_shape = max(max_shape, seq_len)
|
"""Return the stable rotary buffers captured by one graph."""
|
||||||
|
graph_ws = self._sin_cos_ws_by_graph.get(graph_key)
|
||||||
|
if graph_ws is not None:
|
||||||
|
return graph_ws
|
||||||
|
|
||||||
|
needs_new_workspace = self.sin_cos_ws is None or (
|
||||||
|
self.sin_cos_ws[0].size(0) < seq_len
|
||||||
|
or self.sin_cos_ws[0].size(1) < head_dim
|
||||||
|
)
|
||||||
|
if needs_new_workspace:
|
||||||
|
previous = self.sin_cos_ws
|
||||||
|
previous_seq_len = previous[0].size(0) if previous is not None else 0
|
||||||
|
previous_head_dim = previous[0].size(1) if previous is not None else 0
|
||||||
|
max_shape = max(
|
||||||
|
self.max_context_len or 0,
|
||||||
|
previous_seq_len * 2,
|
||||||
|
seq_len,
|
||||||
|
)
|
||||||
|
max_head_dim = max(previous_head_dim, head_dim)
|
||||||
cos_ws = torch.empty(
|
cos_ws = torch.empty(
|
||||||
max_shape, head_dim, dtype=self.dtype, device=self.device
|
max_shape, max_head_dim, dtype=self.dtype, device=self.device
|
||||||
)
|
)
|
||||||
sin_ws = torch.empty(
|
sin_ws = torch.empty(
|
||||||
max_shape, head_dim, dtype=self.dtype, device=self.device
|
max_shape, max_head_dim, dtype=self.dtype, device=self.device
|
||||||
)
|
)
|
||||||
|
if previous is not None:
|
||||||
|
# CUDA graphs retain captured addresses, so an older allocation
|
||||||
|
# cannot be freed when a larger request grows the workspace.
|
||||||
|
self._retired_sin_cos_ws.append(previous)
|
||||||
self.sin_cos_ws = (cos_ws, sin_ws)
|
self.sin_cos_ws = (cos_ws, sin_ws)
|
||||||
else:
|
|
||||||
if self.sin_cos_ws[0].size(0) < seq_len:
|
|
||||||
max_shape = max(self.sin_cos_ws[0].size(0) * 2, seq_len)
|
|
||||||
cos_ws = torch.empty(
|
|
||||||
max_shape, head_dim, dtype=self.dtype, device=self.device
|
|
||||||
)
|
|
||||||
sin_ws = torch.empty(
|
|
||||||
max_shape, head_dim, dtype=self.dtype, device=self.device
|
|
||||||
)
|
|
||||||
self.sin_cos_ws = (cos_ws, sin_ws)
|
|
||||||
|
|
||||||
def _get_graph_key(self, x_3d: torch.Tensor) -> int:
|
graph_ws = (
|
||||||
# x_3d: [S, B, H], B=1, S as graph_key
|
self.sin_cos_ws[0][:seq_len, :head_dim],
|
||||||
return x_3d.shape[0]
|
self.sin_cos_ws[1][:seq_len, :head_dim],
|
||||||
|
)
|
||||||
|
self._sin_cos_ws_by_graph[graph_key] = graph_ws
|
||||||
|
return graph_ws
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _sequence_layout_key(cu_seqlens: Optional[torch.Tensor]) -> Optional[tuple]:
|
||||||
|
if cu_seqlens is None:
|
||||||
|
return None
|
||||||
|
return tuple(int(value) for value in cu_seqlens.tolist())
|
||||||
|
|
||||||
|
def _get_graph_key(
|
||||||
|
self,
|
||||||
|
x_3d: torch.Tensor,
|
||||||
|
cu_seqlens: torch.Tensor,
|
||||||
|
cu_window_seqlens: Optional[torch.Tensor],
|
||||||
|
attention_layout_key: Optional[Hashable] = None,
|
||||||
|
) -> Hashable:
|
||||||
|
"""Include attention boundaries so equal-length batches stay distinct."""
|
||||||
|
if attention_layout_key is None:
|
||||||
|
attention_layout_key = (
|
||||||
|
self._sequence_layout_key(cu_seqlens),
|
||||||
|
self._sequence_layout_key(cu_window_seqlens),
|
||||||
|
)
|
||||||
|
return (x_3d.shape[0], attention_layout_key)
|
||||||
|
|
||||||
def _capture_context(self):
|
def _capture_context(self):
|
||||||
# A DP-sharded encoder intentionally lets each rank capture only the
|
# A DP-sharded encoder intentionally lets each rank capture only the
|
||||||
@@ -265,15 +306,22 @@ class ViTCudaGraphRunner:
|
|||||||
self,
|
self,
|
||||||
x_3d: torch.Tensor, # [S, 1, H]
|
x_3d: torch.Tensor, # [S, 1, H]
|
||||||
cu_seqlens: torch.Tensor,
|
cu_seqlens: torch.Tensor,
|
||||||
cu_window_seqlens: torch.Tensor,
|
cu_window_seqlens: Optional[torch.Tensor],
|
||||||
position_embeddings: Optional[
|
position_embeddings: Optional[
|
||||||
Tuple[torch.Tensor, torch.Tensor]
|
Tuple[torch.Tensor, torch.Tensor]
|
||||||
], # (cos, sin), [S, D]
|
], # (cos, sin), [S, D]
|
||||||
rotary_pos_emb_cos: Optional[torch.Tensor] = None,
|
rotary_pos_emb_cos: Optional[torch.Tensor] = None,
|
||||||
rotary_pos_emb_sin: Optional[torch.Tensor] = None,
|
rotary_pos_emb_sin: Optional[torch.Tensor] = None,
|
||||||
) -> int:
|
attention_layout_key: Optional[Hashable] = None,
|
||||||
|
) -> Hashable:
|
||||||
vit = self.vit
|
vit = self.vit
|
||||||
graph_key = self._get_graph_key(x_3d)
|
graph_key = self._get_graph_key(
|
||||||
|
x_3d,
|
||||||
|
cu_seqlens,
|
||||||
|
cu_window_seqlens,
|
||||||
|
attention_layout_key,
|
||||||
|
)
|
||||||
|
seq_len = x_3d.shape[0]
|
||||||
|
|
||||||
if graph_key in self.block_graphs:
|
if graph_key in self.block_graphs:
|
||||||
return graph_key
|
return graph_key
|
||||||
@@ -291,7 +339,7 @@ class ViTCudaGraphRunner:
|
|||||||
x_3d, device=self.device
|
x_3d, device=self.device
|
||||||
).contiguous()
|
).contiguous()
|
||||||
self.block_ws[graph_key] = torch.empty(
|
self.block_ws[graph_key] = torch.empty(
|
||||||
graph_key,
|
seq_len,
|
||||||
num_heads,
|
num_heads,
|
||||||
attn_head_dim,
|
attn_head_dim,
|
||||||
device=self.device,
|
device=self.device,
|
||||||
@@ -315,12 +363,10 @@ class ViTCudaGraphRunner:
|
|||||||
self.block_input[graph_key].copy_(x_3d)
|
self.block_input[graph_key].copy_(x_3d)
|
||||||
|
|
||||||
if position_embeddings is not None:
|
if position_embeddings is not None:
|
||||||
# make sure rotary workspace
|
|
||||||
head_dim = position_embeddings[0].shape[1]
|
head_dim = position_embeddings[0].shape[1]
|
||||||
self._ensure_sin_cos_ws(graph_key, head_dim)
|
used_cos_ws, used_sin_ws = self._get_sin_cos_ws(
|
||||||
|
graph_key, seq_len, head_dim
|
||||||
used_cos_ws = self.sin_cos_ws[0][:graph_key, :]
|
)
|
||||||
used_sin_ws = self.sin_cos_ws[1][:graph_key, :]
|
|
||||||
used_cos_ws.copy_(position_embeddings[0])
|
used_cos_ws.copy_(position_embeddings[0])
|
||||||
used_sin_ws.copy_(position_embeddings[1])
|
used_sin_ws.copy_(position_embeddings[1])
|
||||||
persist_position_embeddings = (used_cos_ws, used_sin_ws)
|
persist_position_embeddings = (used_cos_ws, used_sin_ws)
|
||||||
@@ -328,12 +374,10 @@ class ViTCudaGraphRunner:
|
|||||||
graph_key=graph_key, position_embeddings=persist_position_embeddings
|
graph_key=graph_key, position_embeddings=persist_position_embeddings
|
||||||
)
|
)
|
||||||
elif rotary_pos_emb_cos is not None and rotary_pos_emb_sin is not None:
|
elif rotary_pos_emb_cos is not None and rotary_pos_emb_sin is not None:
|
||||||
# make sure rotary workspace
|
|
||||||
head_dim = rotary_pos_emb_cos.shape[1]
|
head_dim = rotary_pos_emb_cos.shape[1]
|
||||||
self._ensure_sin_cos_ws(graph_key, head_dim)
|
used_cos_ws, used_sin_ws = self._get_sin_cos_ws(
|
||||||
|
graph_key, seq_len, head_dim
|
||||||
used_cos_ws = self.sin_cos_ws[0][:graph_key, :]
|
)
|
||||||
used_sin_ws = self.sin_cos_ws[1][:graph_key, :]
|
|
||||||
used_cos_ws.copy_(rotary_pos_emb_cos)
|
used_cos_ws.copy_(rotary_pos_emb_cos)
|
||||||
used_sin_ws.copy_(rotary_pos_emb_sin)
|
used_sin_ws.copy_(rotary_pos_emb_sin)
|
||||||
self._create_graph(
|
self._create_graph(
|
||||||
@@ -347,7 +391,7 @@ class ViTCudaGraphRunner:
|
|||||||
|
|
||||||
def replay(
|
def replay(
|
||||||
self,
|
self,
|
||||||
graph_key: int,
|
graph_key: Hashable,
|
||||||
x_3d: torch.Tensor,
|
x_3d: torch.Tensor,
|
||||||
position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
|
||||||
rotary_pos_emb_cos: Optional[torch.Tensor] = None,
|
rotary_pos_emb_cos: Optional[torch.Tensor] = None,
|
||||||
@@ -355,20 +399,19 @@ class ViTCudaGraphRunner:
|
|||||||
output_indices: Optional[torch.Tensor] = None,
|
output_indices: Optional[torch.Tensor] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
|
|
||||||
|
seq_len = x_3d.shape[0]
|
||||||
if position_embeddings is not None:
|
if position_embeddings is not None:
|
||||||
# update rotary workspace content
|
|
||||||
head_dim = position_embeddings[0].shape[1]
|
head_dim = position_embeddings[0].shape[1]
|
||||||
self._ensure_sin_cos_ws(graph_key, head_dim)
|
used_cos_ws, used_sin_ws = self._get_sin_cos_ws(
|
||||||
used_cos_ws = self.sin_cos_ws[0][:graph_key, :]
|
graph_key, seq_len, head_dim
|
||||||
used_sin_ws = self.sin_cos_ws[1][:graph_key, :]
|
)
|
||||||
used_cos_ws.copy_(position_embeddings[0])
|
used_cos_ws.copy_(position_embeddings[0])
|
||||||
used_sin_ws.copy_(position_embeddings[1])
|
used_sin_ws.copy_(position_embeddings[1])
|
||||||
elif rotary_pos_emb_cos is not None and rotary_pos_emb_sin is not None:
|
elif rotary_pos_emb_cos is not None and rotary_pos_emb_sin is not None:
|
||||||
# update rotary workspace content
|
|
||||||
head_dim = rotary_pos_emb_cos.shape[1]
|
head_dim = rotary_pos_emb_cos.shape[1]
|
||||||
self._ensure_sin_cos_ws(graph_key, head_dim)
|
used_cos_ws, used_sin_ws = self._get_sin_cos_ws(
|
||||||
used_cos_ws = self.sin_cos_ws[0][:graph_key, :]
|
graph_key, seq_len, head_dim
|
||||||
used_sin_ws = self.sin_cos_ws[1][:graph_key, :]
|
)
|
||||||
used_cos_ws.copy_(rotary_pos_emb_cos)
|
used_cos_ws.copy_(rotary_pos_emb_cos)
|
||||||
used_sin_ws.copy_(rotary_pos_emb_sin)
|
used_sin_ws.copy_(rotary_pos_emb_sin)
|
||||||
|
|
||||||
@@ -390,15 +433,21 @@ class ViTCudaGraphRunner:
|
|||||||
self,
|
self,
|
||||||
x: torch.Tensor,
|
x: torch.Tensor,
|
||||||
cu_seqlens: torch.Tensor,
|
cu_seqlens: torch.Tensor,
|
||||||
cu_window_seqlens: torch.Tensor,
|
cu_window_seqlens: Optional[torch.Tensor],
|
||||||
position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]],
|
position_embeddings: Optional[Tuple[torch.Tensor, torch.Tensor]],
|
||||||
rotary_pos_emb_cos: Optional[torch.Tensor] = None,
|
rotary_pos_emb_cos: Optional[torch.Tensor] = None,
|
||||||
rotary_pos_emb_sin: Optional[torch.Tensor] = None,
|
rotary_pos_emb_sin: Optional[torch.Tensor] = None,
|
||||||
output_indices: Optional[torch.Tensor] = None,
|
output_indices: Optional[torch.Tensor] = None,
|
||||||
|
attention_layout_key: Optional[Hashable] = None,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
# x: [seq_len, hidden] -> [S, B=1, H]
|
# x: [seq_len, hidden] -> [S, B=1, H]
|
||||||
x_3d = x.unsqueeze(1)
|
x_3d = x.unsqueeze(1)
|
||||||
graph_key = self._get_graph_key(x_3d)
|
graph_key = self._get_graph_key(
|
||||||
|
x_3d,
|
||||||
|
cu_seqlens,
|
||||||
|
cu_window_seqlens,
|
||||||
|
attention_layout_key,
|
||||||
|
)
|
||||||
|
|
||||||
if graph_key not in self.block_graphs:
|
if graph_key not in self.block_graphs:
|
||||||
self.create_graph(
|
self.create_graph(
|
||||||
@@ -408,6 +457,7 @@ class ViTCudaGraphRunner:
|
|||||||
cu_window_seqlens=cu_window_seqlens,
|
cu_window_seqlens=cu_window_seqlens,
|
||||||
rotary_pos_emb_cos=rotary_pos_emb_cos,
|
rotary_pos_emb_cos=rotary_pos_emb_cos,
|
||||||
rotary_pos_emb_sin=rotary_pos_emb_sin,
|
rotary_pos_emb_sin=rotary_pos_emb_sin,
|
||||||
|
attention_layout_key=attention_layout_key,
|
||||||
)
|
)
|
||||||
|
|
||||||
return self.replay(
|
return self.replay(
|
||||||
|
|||||||
@@ -0,0 +1,88 @@
|
|||||||
|
import sys
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
import torch
|
||||||
|
from torch import nn
|
||||||
|
|
||||||
|
from sglang.srt.multimodal.vit_cuda_graph_runner import ViTCudaGraphRunner
|
||||||
|
from sglang.test.ci.ci_register import register_cuda_ci
|
||||||
|
|
||||||
|
register_cuda_ci(est_time=10, stage="base-b", runner_config="1-gpu-large")
|
||||||
|
|
||||||
|
|
||||||
|
class _BoundaryBlock(nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.attn = type(
|
||||||
|
"AttentionConfig",
|
||||||
|
(),
|
||||||
|
{
|
||||||
|
"num_attention_heads_per_partition": 1,
|
||||||
|
"head_size": 1,
|
||||||
|
"qkv_backend_name": "triton_attn",
|
||||||
|
},
|
||||||
|
)()
|
||||||
|
|
||||||
|
def forward(
|
||||||
|
self,
|
||||||
|
x,
|
||||||
|
*,
|
||||||
|
cu_seqlens,
|
||||||
|
position_embeddings,
|
||||||
|
output_ws=None,
|
||||||
|
):
|
||||||
|
boundary = cu_seqlens[0][1].to(x.dtype)
|
||||||
|
position = position_embeddings[0][: x.shape[0], :1].unsqueeze(1)
|
||||||
|
return x + boundary + position
|
||||||
|
|
||||||
|
|
||||||
|
class _Merger(nn.Module):
|
||||||
|
def forward(self, x):
|
||||||
|
return x.squeeze(1)
|
||||||
|
|
||||||
|
|
||||||
|
class _VisionTower(nn.Module):
|
||||||
|
def __init__(self):
|
||||||
|
super().__init__()
|
||||||
|
self.blocks = nn.ModuleList([_BoundaryBlock()])
|
||||||
|
self.merger = _Merger()
|
||||||
|
self.use_data_parallel = True
|
||||||
|
self.deepstack_visual_indexes = []
|
||||||
|
self.deepstack_merger_list = None
|
||||||
|
self.max_context_len = None
|
||||||
|
self.register_buffer("anchor", torch.empty(0, device="cuda"))
|
||||||
|
|
||||||
|
@property
|
||||||
|
def device(self):
|
||||||
|
return self.anchor.device
|
||||||
|
|
||||||
|
@property
|
||||||
|
def dtype(self):
|
||||||
|
return torch.float32
|
||||||
|
|
||||||
|
|
||||||
|
def test_vit_graph_replays_current_attention_and_position_metadata():
|
||||||
|
runner = ViTCudaGraphRunner(_VisionTower())
|
||||||
|
|
||||||
|
def run(seq_len, boundaries, position):
|
||||||
|
x = torch.zeros(seq_len, 1, device="cuda")
|
||||||
|
cu_seqlens = torch.tensor(boundaries, dtype=torch.int32, device="cuda")
|
||||||
|
positions = torch.full((seq_len, 1), position, device="cuda")
|
||||||
|
output = runner.run(x, cu_seqlens, None, (positions, positions))
|
||||||
|
torch.cuda.synchronize()
|
||||||
|
return output.cpu()
|
||||||
|
|
||||||
|
first = run(4, [0, 2, 4], 1)
|
||||||
|
different_layout = run(4, [0, 1, 4], 1)
|
||||||
|
run(8, [0, 8], 7)
|
||||||
|
small_after_growth = run(4, [0, 2, 4], 5)
|
||||||
|
|
||||||
|
torch.testing.assert_close(first, torch.full_like(first, 3))
|
||||||
|
torch.testing.assert_close(different_layout, torch.full_like(different_layout, 2))
|
||||||
|
torch.testing.assert_close(
|
||||||
|
small_after_growth, torch.full_like(small_after_growth, 7)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
sys.exit(pytest.main([__file__, "-v"]))
|
||||||
@@ -3,6 +3,7 @@ from types import SimpleNamespace
|
|||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
import torch
|
||||||
|
|
||||||
from sglang.test.ci.ci_register import register_cpu_ci
|
from sglang.test.ci.ci_register import register_cpu_ci
|
||||||
|
|
||||||
@@ -76,6 +77,39 @@ def test_vit_graph_runner_caches_resolved_backend_name():
|
|||||||
assert runner._attn_backend == "fa3"
|
assert runner._attn_backend == "fa3"
|
||||||
|
|
||||||
|
|
||||||
|
def test_vit_graph_key_includes_full_and_window_attention_boundaries():
|
||||||
|
runner = _runner(use_data_parallel=True)
|
||||||
|
x = torch.empty(8, 1, 4)
|
||||||
|
|
||||||
|
first = runner._get_graph_key(
|
||||||
|
x,
|
||||||
|
torch.tensor([0, 4, 8]),
|
||||||
|
torch.tensor([0, 2, 4, 8]),
|
||||||
|
)
|
||||||
|
second = runner._get_graph_key(
|
||||||
|
x,
|
||||||
|
torch.tensor([0, 2, 8]),
|
||||||
|
torch.tensor([0, 4, 6, 8]),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert first != second
|
||||||
|
|
||||||
|
|
||||||
|
def test_vit_graph_keeps_rotary_workspace_address_after_growth():
|
||||||
|
runner = _runner(use_data_parallel=True)
|
||||||
|
runner.vit.device = torch.device("cpu")
|
||||||
|
runner.vit.dtype = torch.float32
|
||||||
|
|
||||||
|
small = runner._get_sin_cos_ws("small", seq_len=4, head_dim=2)
|
||||||
|
small_address = small[0].data_ptr()
|
||||||
|
runner._get_sin_cos_ws("large", seq_len=16, head_dim=2)
|
||||||
|
|
||||||
|
assert runner._get_sin_cos_ws("small", seq_len=4, head_dim=2)[0].data_ptr() == (
|
||||||
|
small_address
|
||||||
|
)
|
||||||
|
assert len(runner._retired_sin_cos_ws) == 1
|
||||||
|
|
||||||
|
|
||||||
def test_internvl_graph_runner_caches_resolved_backend_name():
|
def test_internvl_graph_runner_caches_resolved_backend_name():
|
||||||
attention = SimpleNamespace(
|
attention = SimpleNamespace(
|
||||||
qkv_backend_name="triton_attn",
|
qkv_backend_name="triton_attn",
|
||||||
|
|||||||
Reference in New Issue
Block a user