fix: bound CUDA memory for fast image preprocessing (#36295)
This commit is contained in:
@@ -6,6 +6,7 @@ import multiprocessing as mp
|
|||||||
import os
|
import os
|
||||||
import re
|
import re
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
|
from contextlib import contextmanager
|
||||||
from typing import (
|
from typing import (
|
||||||
Any,
|
Any,
|
||||||
Dict,
|
Dict,
|
||||||
@@ -642,6 +643,24 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
return "npu"
|
return "npu"
|
||||||
return None
|
return None
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def _temporary_fast_processor_cuda_pool(self, device: Optional[str]):
|
||||||
|
"""Release fast-processor CUDA temporaries after CPU feature transport."""
|
||||||
|
can_release = (
|
||||||
|
device is not None
|
||||||
|
and torch.device(device).type == "cuda"
|
||||||
|
and not self.keep_mm_features_on_device
|
||||||
|
and not self.precompute_hash_before_cpu_transfer
|
||||||
|
)
|
||||||
|
if not can_release:
|
||||||
|
yield
|
||||||
|
return
|
||||||
|
|
||||||
|
with torch.cuda.device(device):
|
||||||
|
pool = torch.cuda.MemPool()
|
||||||
|
with torch.cuda.use_mem_pool(pool, device=device):
|
||||||
|
yield
|
||||||
|
|
||||||
def process_mm_data(
|
def process_mm_data(
|
||||||
self,
|
self,
|
||||||
input_text,
|
input_text,
|
||||||
@@ -690,14 +709,15 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
if self.audio_config:
|
if self.audio_config:
|
||||||
kwargs.setdefault("audio_kwargs", {}).update(self.audio_config)
|
kwargs.setdefault("audio_kwargs", {}).update(self.audio_config)
|
||||||
|
|
||||||
|
processor_device = None
|
||||||
if (
|
if (
|
||||||
hasattr(processor, "image_processor")
|
hasattr(processor, "image_processor")
|
||||||
and isinstance(processor.image_processor, BaseImageProcessor)
|
and isinstance(processor.image_processor, BaseImageProcessor)
|
||||||
and not self.disable_fast_image_processor
|
and not self.disable_fast_image_processor
|
||||||
):
|
):
|
||||||
device = self._fast_image_processor_device(processor)
|
processor_device = self._fast_image_processor_device(processor)
|
||||||
if device is not None:
|
if processor_device is not None:
|
||||||
kwargs["device"] = device
|
kwargs["device"] = processor_device
|
||||||
|
|
||||||
# Avoid double BOS when the chat template already wrote one.
|
# Avoid double BOS when the chat template already wrote one.
|
||||||
if self._tokenizer_auto_adds_specials and isinstance(input_text, str):
|
if self._tokenizer_auto_adds_specials and isinstance(input_text, str):
|
||||||
@@ -705,24 +725,25 @@ class BaseMultimodalProcessor(ABC):
|
|||||||
if bos and input_text.startswith(bos):
|
if bos and input_text.startswith(bos):
|
||||||
kwargs.setdefault("add_special_tokens", False)
|
kwargs.setdefault("add_special_tokens", False)
|
||||||
|
|
||||||
result = processor.__call__(
|
with self._temporary_fast_processor_cuda_pool(processor_device):
|
||||||
text=[input_text],
|
result = processor.__call__(
|
||||||
padding=True,
|
text=[input_text],
|
||||||
return_tensors="pt",
|
padding=True,
|
||||||
**kwargs,
|
return_tensors="pt",
|
||||||
)
|
**kwargs,
|
||||||
# Deferred: the hash is computed on the GPU tensor first, and
|
)
|
||||||
# _precompute_hashes_before_cpu_transfer moves it down afterwards.
|
# Deferred: the hash is computed on the GPU tensor first, and
|
||||||
if (
|
# _precompute_hashes_before_cpu_transfer moves it down afterwards.
|
||||||
not self.keep_mm_features_on_device
|
if (
|
||||||
and not self.precompute_hash_before_cpu_transfer
|
not self.keep_mm_features_on_device
|
||||||
):
|
and not self.precompute_hash_before_cpu_transfer
|
||||||
# move feature tensors to cpu
|
):
|
||||||
for feature_name in self.FEATURE_NAMES:
|
# move feature tensors to cpu
|
||||||
if feature_name in result and isinstance(
|
for feature_name in self.FEATURE_NAMES:
|
||||||
result[feature_name], torch.Tensor
|
if feature_name in result and isinstance(
|
||||||
):
|
result[feature_name], torch.Tensor
|
||||||
result[feature_name] = result[feature_name].to("cpu")
|
):
|
||||||
|
result[feature_name] = result[feature_name].to("cpu")
|
||||||
|
|
||||||
return result
|
return result
|
||||||
|
|
||||||
|
|||||||
@@ -7,6 +7,8 @@ device has to come from what the worker was handed.
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import unittest
|
import unittest
|
||||||
|
from contextlib import nullcontext
|
||||||
|
from types import SimpleNamespace
|
||||||
from unittest.mock import patch
|
from unittest.mock import patch
|
||||||
|
|
||||||
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
|
from sglang.srt.multimodal.processors.base_processor import BaseMultimodalProcessor
|
||||||
@@ -77,5 +79,83 @@ class TestFastImageProcessorDevice(CustomTestCase):
|
|||||||
self.assertIsNone(device)
|
self.assertIsNone(device)
|
||||||
|
|
||||||
|
|
||||||
|
class TestFastImageProcessorMemoryPool(CustomTestCase):
|
||||||
|
def _processor(self, *, transport="cpu", precompute_hash=False):
|
||||||
|
processor = _make(base_gpu_id=0)
|
||||||
|
processor.mm_feature_transport = transport
|
||||||
|
processor.precompute_hash_before_cpu_transfer = precompute_hash
|
||||||
|
return processor
|
||||||
|
|
||||||
|
def test_pool_is_limited_to_immediate_cpu_transport(self):
|
||||||
|
cases = (
|
||||||
|
(self._processor(), "cuda:0", True),
|
||||||
|
(self._processor(transport="cuda_ipc"), "cuda:0", False),
|
||||||
|
(self._processor(transport="cuda_vmm"), "cuda:0", False),
|
||||||
|
(self._processor(precompute_hash=True), "cuda:0", False),
|
||||||
|
(self._processor(), "cpu", False),
|
||||||
|
(self._processor(), None, False),
|
||||||
|
)
|
||||||
|
for processor, device, expected in cases:
|
||||||
|
with (
|
||||||
|
self.subTest(device=device, transport=processor.mm_feature_transport),
|
||||||
|
patch(f"{BASE}.torch.cuda.device", return_value=nullcontext()),
|
||||||
|
patch(f"{BASE}.torch.cuda.MemPool", return_value="pool") as mem_pool,
|
||||||
|
patch(f"{BASE}.torch.cuda.use_mem_pool", return_value=nullcontext()),
|
||||||
|
):
|
||||||
|
with processor._temporary_fast_processor_cuda_pool(device):
|
||||||
|
pass
|
||||||
|
self.assertEqual(mem_pool.called, expected)
|
||||||
|
|
||||||
|
def test_processor_call_uses_private_pool_until_cpu_copy_finishes(self):
|
||||||
|
class ImageProcessor:
|
||||||
|
pass
|
||||||
|
|
||||||
|
class Feature:
|
||||||
|
def to(self, device):
|
||||||
|
events.append(("copy", device))
|
||||||
|
|
||||||
|
feature = Feature()
|
||||||
|
|
||||||
|
class Processor:
|
||||||
|
image_processor = ImageProcessor()
|
||||||
|
tokenizer = SimpleNamespace(bos_token=None)
|
||||||
|
|
||||||
|
def __call__(self, **kwargs):
|
||||||
|
events.append(("call", kwargs["device"]))
|
||||||
|
return {"pixel_values": feature}
|
||||||
|
|
||||||
|
events = []
|
||||||
|
processor = self._processor()
|
||||||
|
processor._processor = Processor()
|
||||||
|
processor._tokenizer = processor._processor.tokenizer
|
||||||
|
processor._tokenizer_auto_adds_specials = False
|
||||||
|
processor.disable_fast_image_processor = False
|
||||||
|
processor.image_config = {}
|
||||||
|
processor.video_config = {}
|
||||||
|
processor.audio_config = {}
|
||||||
|
processor.FEATURE_NAMES = ["pixel_values"]
|
||||||
|
|
||||||
|
class PoolContext:
|
||||||
|
def __enter__(self):
|
||||||
|
events.append("enter")
|
||||||
|
|
||||||
|
def __exit__(self, *args):
|
||||||
|
events.append("exit")
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(f"{BASE}.BaseImageProcessor", ImageProcessor),
|
||||||
|
patch(f"{BASE}.torch.cuda.device", return_value=nullcontext()),
|
||||||
|
patch(f"{BASE}.torch.cuda.MemPool", return_value="pool"),
|
||||||
|
patch(f"{BASE}.torch.cuda.use_mem_pool", return_value=PoolContext()),
|
||||||
|
patch(f"{BASE}.torch.Tensor", Feature),
|
||||||
|
):
|
||||||
|
processor.process_mm_data("test", images=["image"])
|
||||||
|
|
||||||
|
self.assertEqual(
|
||||||
|
events,
|
||||||
|
["enter", ("call", "cuda:0"), ("copy", "cpu"), "exit"],
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
if __name__ == "__main__":
|
if __name__ == "__main__":
|
||||||
unittest.main()
|
unittest.main()
|
||||||
|
|||||||
Reference in New Issue
Block a user