fix tp capture in vit cuda graph (#17255)

This commit is contained in:
narutolhy
2026-03-27 22:38:18 +00:00
committed by GitHub
parent ec29bbb286
commit 9b29131961
@@ -17,11 +17,13 @@
from __future__ import annotations from __future__ import annotations
import inspect import inspect
from contextlib import nullcontext
from typing import Dict, Hashable, List, Optional, Tuple from typing import Dict, Hashable, List, Optional, Tuple
import torch import torch
import torch.nn as nn import torch.nn as nn
from sglang.srt.distributed.parallel_state import get_tp_group
from sglang.srt.layers.attention.vision import VisionAttention from sglang.srt.layers.attention.vision import VisionAttention
from sglang.srt.server_args import get_global_server_args from sglang.srt.server_args import get_global_server_args
@@ -139,7 +141,11 @@ class ViTCudaGraphRunner:
override_backend = get_global_server_args().mm_attention_backend override_backend = get_global_server_args().mm_attention_backend
with torch.cuda.graph(graph): tp_group = get_tp_group()
ca_comm = tp_group.ca_comm
capture_ctx = ca_comm.capture() if ca_comm is not None else nullcontext()
with capture_ctx, torch.cuda.graph(graph):
y = None y = None
deepstack_outs: List[torch.Tensor] = [] deepstack_outs: List[torch.Tensor] = []
deepstack_capture_idx = 0 deepstack_capture_idx = 0