[Fix] Fix llava on multi images (#1247)
This commit is contained in:
@@ -16,7 +16,7 @@ limitations under the License.
|
||||
"""ModelRunner runs the forward passes of the models."""
|
||||
from dataclasses import dataclass
|
||||
from enum import IntEnum, auto
|
||||
from typing import TYPE_CHECKING, List, Optional
|
||||
from typing import TYPE_CHECKING, List
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
@@ -58,6 +58,7 @@ class InputMetadata:
|
||||
|
||||
# For extend
|
||||
extend_seq_lens: torch.Tensor = None
|
||||
extend_prefix_lens: torch.Tensor = None
|
||||
extend_start_loc: torch.Tensor = None
|
||||
extend_no_prefix: bool = None
|
||||
|
||||
@@ -69,8 +70,8 @@ class InputMetadata:
|
||||
|
||||
# For multimodal
|
||||
pixel_values: List[torch.Tensor] = None
|
||||
image_sizes: List[List[int]] = None
|
||||
image_offsets: List[int] = None
|
||||
image_sizes: List[List[List[int]]] = None
|
||||
image_offsets: List[List[int]] = None
|
||||
|
||||
# Trition attention backend
|
||||
triton_max_seq_len: int = 0
|
||||
@@ -87,20 +88,8 @@ class InputMetadata:
|
||||
def init_multimuldal_info(self, batch: ScheduleBatch):
|
||||
reqs = batch.reqs
|
||||
self.pixel_values = [r.pixel_values for r in reqs]
|
||||
self.image_sizes = [r.image_size for r in reqs]
|
||||
self.image_offsets = []
|
||||
for r in reqs:
|
||||
if isinstance(r.image_offset, list):
|
||||
self.image_offsets.append(
|
||||
[
|
||||
(image_offset - len(r.prefix_indices))
|
||||
for image_offset in r.image_offset
|
||||
]
|
||||
)
|
||||
elif isinstance(r.image_offset, int):
|
||||
self.image_offsets.append(r.image_offset - len(r.prefix_indices))
|
||||
elif r.image_offset is None:
|
||||
self.image_offsets.append(0)
|
||||
self.image_sizes = [r.image_sizes for r in reqs]
|
||||
self.image_offsets = [r.image_offsets for r in reqs]
|
||||
|
||||
def compute_positions(self, batch: ScheduleBatch):
|
||||
position_ids_offsets = batch.position_ids_offsets
|
||||
@@ -153,6 +142,7 @@ class InputMetadata:
|
||||
for i, r in enumerate(batch.reqs)
|
||||
]
|
||||
self.extend_seq_lens = torch.tensor(extend_lens_cpu, device="cuda")
|
||||
self.extend_prefix_lens = torch.tensor(batch.prefix_lens_cpu, device="cuda")
|
||||
self.extend_start_loc = torch.zeros_like(self.seq_lens)
|
||||
self.extend_start_loc[1:] = torch.cumsum(self.extend_seq_lens[:-1], dim=0)
|
||||
self.extend_no_prefix = all(l == 0 for l in batch.prefix_lens_cpu)
|
||||
@@ -238,10 +228,10 @@ class InputMetadata:
|
||||
prefix_lens_cpu,
|
||||
flashinfer_use_ragged,
|
||||
):
|
||||
if self.forward_mode != ForwardMode.DECODE:
|
||||
prefix_lens = torch.tensor(prefix_lens_cpu, device="cuda")
|
||||
else:
|
||||
if self.forward_mode == ForwardMode.DECODE:
|
||||
prefix_lens = None
|
||||
else:
|
||||
prefix_lens = self.extend_prefix_lens
|
||||
|
||||
update_flashinfer_indices(
|
||||
self.forward_mode,
|
||||
|
||||
@@ -50,7 +50,7 @@ from sglang.srt.mem_cache.memory_pool import (
|
||||
MLATokenToKVPool,
|
||||
ReqToTokenPool,
|
||||
)
|
||||
from sglang.srt.model_config import AttentionArch
|
||||
from sglang.srt.model_config import AttentionArch, ModelConfig
|
||||
from sglang.srt.model_executor.forward_batch_info import ForwardMode, InputMetadata
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import (
|
||||
@@ -69,7 +69,7 @@ logger = logging.getLogger(__name__)
|
||||
class ModelRunner:
|
||||
def __init__(
|
||||
self,
|
||||
model_config,
|
||||
model_config: ModelConfig,
|
||||
mem_fraction_static: float,
|
||||
gpu_id: int,
|
||||
tp_rank: int,
|
||||
@@ -85,7 +85,9 @@ class ModelRunner:
|
||||
self.tp_size = tp_size
|
||||
self.nccl_port = nccl_port
|
||||
self.server_args = server_args
|
||||
self.is_multimodal_model = is_multimodal_model(self.model_config)
|
||||
self.is_multimodal_model = is_multimodal_model(
|
||||
self.model_config.hf_config.architectures
|
||||
)
|
||||
global_server_args_dict.update(
|
||||
{
|
||||
"disable_flashinfer": server_args.disable_flashinfer,
|
||||
@@ -95,6 +97,13 @@ class ModelRunner:
|
||||
}
|
||||
)
|
||||
|
||||
if self.is_multimodal_model:
|
||||
logger.info(
|
||||
"Automatically turn off --chunked-prefill-size and adjust --mem-fraction-static for multimodal models."
|
||||
)
|
||||
server_args.chunked_prefill_size = None
|
||||
server_args.mem_fraction_static *= 0.95
|
||||
|
||||
min_per_gpu_memory = self.init_torch_distributed()
|
||||
self.load_model()
|
||||
self.init_memory_pool(
|
||||
@@ -507,9 +516,9 @@ class ModelRunner:
|
||||
raise Exception(
|
||||
f"Capture cuda graph failed: {e}\n"
|
||||
"Possible solutions:\n"
|
||||
"1. disable torch compile by not using --enable-torch-compile\n"
|
||||
"2. disable cuda graph by --disable-cuda-graph\n"
|
||||
"3. set --mem-fraction-static to a smaller value\n"
|
||||
"1. disable cuda graph by --disable-cuda-graph\n"
|
||||
"2. set --mem-fraction-static to a smaller value\n"
|
||||
"3. disable torch compile by not using --enable-torch-compile\n"
|
||||
"Open an issue on GitHub https://github.com/sgl-project/sglang/issues/new/choose \n"
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user