move more files under srt/utils (#11285)
This commit is contained in:
@@ -1,2 +1,2 @@
|
||||
# Temporarily do this to avoid changing all imports in the repo
|
||||
from .common import *
|
||||
from sglang.srt.utils.common import *
|
||||
|
||||
@@ -0,0 +1,100 @@
|
||||
import asyncio
|
||||
|
||||
|
||||
class RWLock:
|
||||
def __init__(self):
|
||||
# Protects internal state
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
# Condition variable used to wait for state changes
|
||||
self._cond = asyncio.Condition(self._lock)
|
||||
|
||||
# Number of readers currently holding the lock
|
||||
self._readers = 0
|
||||
|
||||
# Whether a writer is currently holding the lock
|
||||
self._writer_active = False
|
||||
|
||||
# How many writers are queued waiting for a turn
|
||||
self._waiting_writers = 0
|
||||
|
||||
@property
|
||||
def reader_lock(self):
|
||||
"""
|
||||
A context manager for acquiring a shared (reader) lock.
|
||||
|
||||
Example:
|
||||
async with rwlock.reader_lock:
|
||||
# read-only access
|
||||
"""
|
||||
return _ReaderLock(self)
|
||||
|
||||
@property
|
||||
def writer_lock(self):
|
||||
"""
|
||||
A context manager for acquiring an exclusive (writer) lock.
|
||||
|
||||
Example:
|
||||
async with rwlock.writer_lock:
|
||||
# exclusive access
|
||||
"""
|
||||
return _WriterLock(self)
|
||||
|
||||
async def acquire_reader(self):
|
||||
async with self._lock:
|
||||
# Wait until there is no active writer or waiting writer
|
||||
# to ensure fairness.
|
||||
while self._writer_active or self._waiting_writers > 0:
|
||||
await self._cond.wait()
|
||||
self._readers += 1
|
||||
|
||||
async def release_reader(self):
|
||||
async with self._lock:
|
||||
self._readers -= 1
|
||||
# If this was the last reader, wake up anyone waiting
|
||||
# (potentially a writer or new readers).
|
||||
if self._readers == 0:
|
||||
self._cond.notify_all()
|
||||
|
||||
async def acquire_writer(self):
|
||||
async with self._lock:
|
||||
# Increment the count of writers waiting
|
||||
self._waiting_writers += 1
|
||||
try:
|
||||
# Wait while either a writer is active or readers are present
|
||||
while self._writer_active or self._readers > 0:
|
||||
await self._cond.wait()
|
||||
self._writer_active = True
|
||||
finally:
|
||||
# Decrement waiting writers only after we've acquired the writer lock
|
||||
self._waiting_writers -= 1
|
||||
|
||||
async def release_writer(self):
|
||||
async with self._lock:
|
||||
self._writer_active = False
|
||||
# Wake up anyone waiting (readers or writers)
|
||||
self._cond.notify_all()
|
||||
|
||||
|
||||
class _ReaderLock:
|
||||
def __init__(self, rwlock: RWLock):
|
||||
self._rwlock = rwlock
|
||||
|
||||
async def __aenter__(self):
|
||||
await self._rwlock.acquire_reader()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
await self._rwlock.release_reader()
|
||||
|
||||
|
||||
class _WriterLock:
|
||||
def __init__(self, rwlock: RWLock):
|
||||
self._rwlock = rwlock
|
||||
|
||||
async def __aenter__(self):
|
||||
await self._rwlock.acquire_writer()
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
await self._rwlock.release_writer()
|
||||
@@ -0,0 +1,137 @@
|
||||
import os
|
||||
import sys
|
||||
from contextlib import nullcontext
|
||||
|
||||
import torch
|
||||
|
||||
|
||||
# NOTE copied and modified from DeepGEMM
|
||||
class suppress_stdout_stderr:
|
||||
def __enter__(self):
|
||||
self.outnull_file = open(os.devnull, "w")
|
||||
self.errnull_file = open(os.devnull, "w")
|
||||
|
||||
self.old_stdout_fileno_undup = sys.stdout.fileno()
|
||||
self.old_stderr_fileno_undup = sys.stderr.fileno()
|
||||
|
||||
self.old_stdout_fileno = os.dup(sys.stdout.fileno())
|
||||
self.old_stderr_fileno = os.dup(sys.stderr.fileno())
|
||||
|
||||
self.old_stdout = sys.stdout
|
||||
self.old_stderr = sys.stderr
|
||||
|
||||
os.dup2(self.outnull_file.fileno(), self.old_stdout_fileno_undup)
|
||||
os.dup2(self.errnull_file.fileno(), self.old_stderr_fileno_undup)
|
||||
|
||||
sys.stdout = self.outnull_file
|
||||
sys.stderr = self.errnull_file
|
||||
return self
|
||||
|
||||
def __exit__(self, *_):
|
||||
sys.stdout = self.old_stdout
|
||||
sys.stderr = self.old_stderr
|
||||
|
||||
os.dup2(self.old_stdout_fileno, self.old_stdout_fileno_undup)
|
||||
os.dup2(self.old_stderr_fileno, self.old_stderr_fileno_undup)
|
||||
|
||||
os.close(self.old_stdout_fileno)
|
||||
os.close(self.old_stderr_fileno)
|
||||
|
||||
self.outnull_file.close()
|
||||
self.errnull_file.close()
|
||||
|
||||
|
||||
# NOTE copied and modified from DeepGEMM
|
||||
def bench_kineto(
|
||||
fn,
|
||||
kernel_names,
|
||||
num_tests: int = 30,
|
||||
suppress_kineto_output: bool = False,
|
||||
trace_path: str = None,
|
||||
flush_l2: bool = True,
|
||||
with_multiple_kernels: bool = False,
|
||||
):
|
||||
# Conflict with Nsight Systems
|
||||
using_nsys = int(os.environ.get("SGLANG_NSYS_PROFILING", 0))
|
||||
|
||||
# By default, flush L2 with an excessive 8GB memset to give the GPU some (literal) chill time without full idle
|
||||
flush_l2_size = int(8e9 // 4)
|
||||
|
||||
# For some auto-tuning kernels with prints
|
||||
fn()
|
||||
|
||||
# Profile
|
||||
suppress = (
|
||||
suppress_stdout_stderr
|
||||
if suppress_kineto_output and not using_nsys
|
||||
else nullcontext
|
||||
)
|
||||
with suppress():
|
||||
schedule = (
|
||||
torch.profiler.schedule(wait=0, warmup=1, active=1, repeat=1)
|
||||
if not using_nsys
|
||||
else None
|
||||
)
|
||||
profiler = (
|
||||
torch.profiler.profile(
|
||||
activities=[torch.profiler.ProfilerActivity.CUDA], schedule=schedule
|
||||
)
|
||||
if not using_nsys
|
||||
else nullcontext()
|
||||
)
|
||||
with profiler:
|
||||
for i in range(2):
|
||||
for _ in range(num_tests):
|
||||
if flush_l2:
|
||||
torch.empty(
|
||||
flush_l2_size, dtype=torch.int, device="cuda"
|
||||
).zero_()
|
||||
fn()
|
||||
|
||||
if not using_nsys:
|
||||
profiler.step()
|
||||
|
||||
# Return 1 if using Nsight Systems
|
||||
if using_nsys:
|
||||
return 1
|
||||
|
||||
# Parse the profiling table
|
||||
assert isinstance(kernel_names, str) or isinstance(kernel_names, tuple)
|
||||
is_tuple = isinstance(kernel_names, tuple)
|
||||
prof_lines = (
|
||||
profiler.key_averages()
|
||||
.table(sort_by="cuda_time_total", max_name_column_width=100)
|
||||
.split("\n")
|
||||
)
|
||||
kernel_names = (kernel_names,) if isinstance(kernel_names, str) else kernel_names
|
||||
assert all([isinstance(name, str) for name in kernel_names])
|
||||
if not with_multiple_kernels:
|
||||
for name in kernel_names:
|
||||
assert (
|
||||
sum([name in line for line in prof_lines]) == 1
|
||||
), f"Errors of the kernel {name} in the profiling table (table: {prof_lines})"
|
||||
|
||||
# Save chrome traces
|
||||
if trace_path is not None:
|
||||
profiler.export_chrome_trace(trace_path)
|
||||
|
||||
# Return average kernel times
|
||||
units = {"ms": 1e3, "us": 1e6}
|
||||
kernel_times = []
|
||||
for name in kernel_names:
|
||||
total_time = 0
|
||||
total_num = 0
|
||||
for line in prof_lines:
|
||||
if name in line:
|
||||
time_str = line.split()[-2]
|
||||
num_str = line.split()[-1]
|
||||
for unit, scale in units.items():
|
||||
if unit in time_str:
|
||||
total_time += (
|
||||
float(time_str.replace(unit, "")) / scale * int(num_str)
|
||||
)
|
||||
total_num += int(num_str)
|
||||
break
|
||||
kernel_times.append(total_time / total_num)
|
||||
|
||||
return tuple(kernel_times) if is_tuple else kernel_times[0]
|
||||
@@ -487,7 +487,7 @@ def make_layers(
|
||||
# circula imports
|
||||
from sglang.srt.distributed import get_pp_indices
|
||||
from sglang.srt.layers.utils import PPMissingLayer
|
||||
from sglang.srt.offloader import get_offloader
|
||||
from sglang.srt.utils.offloader import get_offloader
|
||||
|
||||
assert not pp_size or num_hidden_layers >= pp_size
|
||||
start_layer, end_layer = (
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
import logging
|
||||
import os
|
||||
from dataclasses import dataclass
|
||||
from multiprocessing import shared_memory
|
||||
from pathlib import Path
|
||||
from typing import List, Optional
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
|
||||
from sglang.srt.distributed.naive_distributed import get_naive_distributed
|
||||
from sglang.srt.utils import check_cuda_result
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class HostSharedMemoryManager:
|
||||
def __init__(self, base_name: str):
|
||||
self._base_name = Path(base_name)
|
||||
self._operation_index = 0
|
||||
self._records: List[_Record] = []
|
||||
|
||||
def malloc(self, *, shape, dtype):
|
||||
meta_tensor = torch.empty(size=shape, dtype=dtype, device="meta")
|
||||
raw = self._malloc_raw(num_bytes=meta_tensor.nbytes)
|
||||
return raw.view(dtype).view(*shape)
|
||||
|
||||
def _malloc_raw(self, *, num_bytes: int) -> torch.Tensor:
|
||||
import cuda.bindings.runtime as cuda_rt
|
||||
|
||||
self._operation_index += 1
|
||||
shm_name = f"{self._base_name}_op{self._operation_index}"
|
||||
|
||||
# TODO handle dispose
|
||||
if get_naive_distributed().get_rank() == 0:
|
||||
shm = shared_memory.SharedMemory(name=shm_name, create=True, size=num_bytes)
|
||||
|
||||
get_naive_distributed().barrier()
|
||||
|
||||
if get_naive_distributed().get_rank() != 0:
|
||||
shm = shared_memory.SharedMemory(name=shm_name)
|
||||
|
||||
np_array = np.ndarray((num_bytes,), dtype=np.uint8, buffer=shm.buf)
|
||||
tensor = torch.from_numpy(np_array)
|
||||
|
||||
check_cuda_result(
|
||||
cuda_rt.cudaHostRegister(
|
||||
tensor.data_ptr(), num_bytes, cuda_rt.cudaHostRegisterPortable
|
||||
)
|
||||
)
|
||||
|
||||
get_naive_distributed().barrier()
|
||||
|
||||
self._records.append(
|
||||
_Record(
|
||||
shm=shm,
|
||||
np_array=np_array,
|
||||
tensor=tensor,
|
||||
)
|
||||
)
|
||||
return tensor
|
||||
|
||||
|
||||
@dataclass
|
||||
class _Record:
|
||||
shm: shared_memory.SharedMemory
|
||||
np_array: np.ndarray
|
||||
tensor: torch.Tensor
|
||||
|
||||
|
||||
# Can have multi instances if needed
|
||||
_instance: Optional[HostSharedMemoryManager] = None
|
||||
|
||||
|
||||
def get_host_shared_memory_manager():
|
||||
assert _instance is not None
|
||||
return _instance
|
||||
|
||||
|
||||
def set_host_shared_memory_manager(instance: HostSharedMemoryManager):
|
||||
global _instance
|
||||
assert _instance is None
|
||||
_instance = instance
|
||||
@@ -0,0 +1,572 @@
|
||||
import logging
|
||||
import os
|
||||
from abc import ABC
|
||||
from typing import Callable, Generator, List, Optional
|
||||
|
||||
import torch
|
||||
from torch.func import functional_call
|
||||
|
||||
from sglang.srt.distributed.naive_distributed import (
|
||||
NaiveDistributed,
|
||||
get_naive_distributed,
|
||||
set_naive_distributed,
|
||||
)
|
||||
from sglang.srt.layers.parameter import ModelWeightParameter
|
||||
from sglang.srt.server_args import ServerArgs
|
||||
from sglang.srt.utils import MultiprocessingSerializer, is_pin_memory_available
|
||||
from sglang.srt.utils.host_shared_memory import (
|
||||
HostSharedMemoryManager,
|
||||
get_host_shared_memory_manager,
|
||||
set_host_shared_memory_manager,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_SubmoduleAccessor = Callable[[torch.nn.Module], torch.nn.Module]
|
||||
_WhitelistParamNamesCreator = Callable[[torch.nn.Module], List[str]]
|
||||
|
||||
|
||||
class BaseOffloader(ABC):
|
||||
def wrap_modules(
|
||||
self,
|
||||
all_modules_generator: Generator[torch.nn.Module, None, None],
|
||||
submodule_accessor: Optional[_SubmoduleAccessor] = None,
|
||||
whitelist_param_names_creator: Optional[_WhitelistParamNamesCreator] = None,
|
||||
):
|
||||
return list(all_modules_generator)
|
||||
|
||||
def post_init(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
def forbid_copy_engine_usage(self):
|
||||
return False
|
||||
|
||||
|
||||
class NoopOffloader(BaseOffloader):
|
||||
pass
|
||||
|
||||
|
||||
# For simplicity use singleton, but can surely support multi instance
|
||||
_instance: Optional[BaseOffloader] = NoopOffloader()
|
||||
|
||||
|
||||
def get_offloader():
|
||||
assert _instance is not None
|
||||
return _instance
|
||||
|
||||
|
||||
def set_offloader(instance: BaseOffloader):
|
||||
global _instance
|
||||
_instance = instance
|
||||
|
||||
|
||||
def create_offloader_from_server_args(server_args: ServerArgs, dp_rank: int):
|
||||
if server_args.cpu_offload_gb > 0:
|
||||
return OffloaderV1(
|
||||
cpu_offload_max_bytes=int(server_args.cpu_offload_gb * 1024**3)
|
||||
)
|
||||
if server_args.offload_group_size > 0:
|
||||
assert (
|
||||
server_args.cpu_offload_gb == 0
|
||||
), "V2 offload does not support cpu_offload_gb yet"
|
||||
return OffloaderV2(
|
||||
group_size=server_args.offload_group_size,
|
||||
num_in_group=server_args.offload_num_in_group,
|
||||
prefetch_step=server_args.offload_prefetch_step,
|
||||
mode=server_args.offload_mode,
|
||||
dp_rank=dp_rank,
|
||||
dp_size=server_args.dp_size,
|
||||
)
|
||||
return NoopOffloader()
|
||||
|
||||
|
||||
class OffloaderV1(BaseOffloader):
|
||||
def __init__(self, cpu_offload_max_bytes: int):
|
||||
self._cpu_offload_bytes = 0
|
||||
self._cpu_offload_max_bytes = cpu_offload_max_bytes
|
||||
|
||||
def wrap_modules(
|
||||
self,
|
||||
all_modules_generator: Generator[torch.nn.Module, None, None],
|
||||
submodule_accessor: Optional[_SubmoduleAccessor] = None,
|
||||
whitelist_param_names_creator: Optional[_WhitelistParamNamesCreator] = None,
|
||||
):
|
||||
return [self.maybe_offload_to_cpu(module) for module in all_modules_generator]
|
||||
|
||||
def maybe_offload_to_cpu(self, module: torch.nn.Module) -> torch.nn.Module:
|
||||
if (params := next(module.parameters(), None)) is None:
|
||||
return module
|
||||
|
||||
device = params.device
|
||||
|
||||
if device == torch.device("cpu"):
|
||||
return module
|
||||
|
||||
if self._cpu_offload_bytes >= self._cpu_offload_max_bytes:
|
||||
return module
|
||||
|
||||
pin_memory = is_pin_memory_available()
|
||||
# offload parameters to CPU
|
||||
# use pin_memory if possible, which helps cudagraph capture speed
|
||||
offloaded_parameters = False
|
||||
for p in module.parameters():
|
||||
if self._cpu_offload_bytes >= self._cpu_offload_max_bytes:
|
||||
# we use per-parameter offloading
|
||||
# one module might have some parameters offloaded and some not
|
||||
break
|
||||
|
||||
# `torch.empty_like` does not support `pin_memory` argument
|
||||
cpu_data = torch.empty_strided(
|
||||
size=p.data.size(),
|
||||
stride=p.data.stride(),
|
||||
dtype=p.data.dtype,
|
||||
layout=p.data.layout,
|
||||
device="cpu",
|
||||
pin_memory=pin_memory,
|
||||
)
|
||||
cpu_data.copy_(p.data)
|
||||
p.data = cpu_data
|
||||
self._cpu_offload_bytes += p.data.numel() * p.data.element_size()
|
||||
offloaded_parameters = True
|
||||
|
||||
if offloaded_parameters:
|
||||
original_forward = module.forward
|
||||
|
||||
def forward(*args, **kwargs):
|
||||
module.forward = original_forward
|
||||
device_state = {
|
||||
# here we blindly call `to(device)`
|
||||
# if the parameter is already on the device, it will be a no-op
|
||||
k: v.to(device, non_blocking=True)
|
||||
for k, v in module.state_dict().items()
|
||||
}
|
||||
output = functional_call(module, device_state, args=args, kwargs=kwargs)
|
||||
module.forward = forward
|
||||
return output
|
||||
|
||||
module.forward = forward
|
||||
|
||||
return module
|
||||
|
||||
|
||||
class OffloaderV2(BaseOffloader):
|
||||
def __init__(
|
||||
self,
|
||||
group_size: int,
|
||||
num_in_group: int,
|
||||
prefetch_step: int,
|
||||
mode: str,
|
||||
dp_rank: int,
|
||||
dp_size: int,
|
||||
):
|
||||
self.group_size = group_size
|
||||
self.num_in_group = num_in_group
|
||||
self.prefetch_step = prefetch_step
|
||||
self.mode = mode
|
||||
|
||||
run_id = os.environ["SGLANG_RUN_ID"]
|
||||
|
||||
# Temporarily init inside Offloader, can move if other modules also need this
|
||||
if self.mode in {"sharded_gpu", "shm_cpu"}:
|
||||
from sglang.srt.distributed import get_tensor_model_parallel_world_size
|
||||
|
||||
assert (
|
||||
get_tensor_model_parallel_world_size() == 1
|
||||
), "not yet support tp_size!=1"
|
||||
set_naive_distributed(
|
||||
NaiveDistributed(
|
||||
rank=dp_rank,
|
||||
world_size=dp_size,
|
||||
rendezvous=f"/tmp/{run_id}",
|
||||
)
|
||||
)
|
||||
if self.mode in {"shm_cpu"}:
|
||||
set_host_shared_memory_manager(
|
||||
HostSharedMemoryManager(
|
||||
base_name=run_id,
|
||||
)
|
||||
)
|
||||
|
||||
self.offloaders = []
|
||||
|
||||
def wrap_modules(
|
||||
self,
|
||||
all_modules_generator: Generator[torch.nn.Module, None, None],
|
||||
submodule_accessor: Optional[_SubmoduleAccessor] = None,
|
||||
whitelist_param_names_creator: Optional[_WhitelistParamNamesCreator] = None,
|
||||
):
|
||||
assert len(self.offloaders) == 0, "should only call wrap_modules once"
|
||||
|
||||
alt_stream = torch.cuda.Stream()
|
||||
|
||||
all_modules = []
|
||||
offload_submodules = []
|
||||
for module_index, module in enumerate(all_modules_generator):
|
||||
all_modules.append(module)
|
||||
if module_index % self.group_size >= self.group_size - self.num_in_group:
|
||||
submodule = submodule_accessor(module)
|
||||
whitelist_param_names = whitelist_param_names_creator(submodule)
|
||||
logger.info(
|
||||
f"[offloader] offload {module_index=} submodule={type(submodule)} params={whitelist_param_names} memory_allocated={torch.cuda.memory_allocated()}"
|
||||
)
|
||||
offload_submodules.append(submodule)
|
||||
self.offloaders.append(
|
||||
_ModuleOffloader(
|
||||
mode=self.mode,
|
||||
module=submodule,
|
||||
alt_stream=alt_stream,
|
||||
whitelist_param_names=whitelist_param_names,
|
||||
)
|
||||
)
|
||||
|
||||
for index, module in enumerate(offload_submodules):
|
||||
_hook_module_forward_for_offloader(
|
||||
index=index,
|
||||
module=module,
|
||||
offloaders=self.offloaders,
|
||||
prefetch_step=self.prefetch_step,
|
||||
)
|
||||
|
||||
return all_modules
|
||||
|
||||
def post_init(self):
|
||||
for offloader in self.offloaders:
|
||||
offloader.post_init()
|
||||
|
||||
for i in range(self.prefetch_step):
|
||||
self.offloaders[i].start_onload()
|
||||
|
||||
@property
|
||||
def forbid_copy_engine_usage(self):
|
||||
return self.mode == "cpu"
|
||||
|
||||
|
||||
def _hook_module_forward_for_offloader(index, module, offloaders, prefetch_step):
|
||||
def _on_forward_end():
|
||||
offloaders[(index + prefetch_step) % len(offloaders)].start_onload()
|
||||
offloaders[index].offload()
|
||||
|
||||
_hook_module_forward_raw(
|
||||
module,
|
||||
on_forward_end=_on_forward_end,
|
||||
get_parameter_and_buffer_dicts=lambda: offloaders[
|
||||
index
|
||||
].wait_and_get_device_tensors(),
|
||||
)
|
||||
|
||||
|
||||
def _hook_module_forward_raw(module, on_forward_end, get_parameter_and_buffer_dicts):
|
||||
original_forward = module.forward
|
||||
|
||||
def forward(*args, **kwargs):
|
||||
module.forward = original_forward
|
||||
output = functional_call(
|
||||
module, get_parameter_and_buffer_dicts(), args=args, kwargs=kwargs
|
||||
)
|
||||
on_forward_end()
|
||||
module.forward = forward
|
||||
return output
|
||||
|
||||
module.forward = forward
|
||||
|
||||
|
||||
class _ModuleOffloader(ABC):
|
||||
def __init__(
|
||||
self,
|
||||
mode: str,
|
||||
module: torch.nn.Module,
|
||||
alt_stream: torch.cuda.Stream,
|
||||
whitelist_param_names: List[str],
|
||||
):
|
||||
self.mode = mode
|
||||
self.module = module
|
||||
self.device = next(module.parameters()).device
|
||||
self.alt_stream = alt_stream
|
||||
|
||||
assert self.device != torch.device(
|
||||
"cpu"
|
||||
), "not handled device=cpu case yet (should skip this tensor)"
|
||||
|
||||
self._device_tensors = None
|
||||
self._load_event = None
|
||||
|
||||
param_dict = dict(self.module.named_parameters())
|
||||
assert all(
|
||||
name in param_dict for name in whitelist_param_names
|
||||
), f"{whitelist_param_names=} {list(param_dict.keys())=}"
|
||||
|
||||
self._param_offloaders = {
|
||||
name: _BaseParamOffloader.create(mode, module=module, param_name=name)
|
||||
for name in whitelist_param_names
|
||||
}
|
||||
|
||||
def post_init(self):
|
||||
for name, param_offloader in self._param_offloaders.items():
|
||||
param_offloader.post_init()
|
||||
|
||||
def start_onload(self):
|
||||
self.alt_stream.wait_stream(torch.cuda.current_stream())
|
||||
with torch.cuda.stream(self.alt_stream):
|
||||
self._device_tensors = self._create_device_tensors()
|
||||
self._load_event = torch.cuda.Event()
|
||||
self._load_event.record()
|
||||
|
||||
def offload(self):
|
||||
self._device_tensors = None
|
||||
self._load_event = None
|
||||
|
||||
def wait_and_get_device_tensors(self):
|
||||
assert self._device_tensors is not None
|
||||
self._load_event.wait()
|
||||
return self._device_tensors
|
||||
|
||||
def _create_device_tensors(self):
|
||||
return {k: v.create_device_tensor() for k, v in self._param_offloaders.items()}
|
||||
|
||||
|
||||
class _BaseParamOffloader(ABC):
|
||||
@staticmethod
|
||||
def create(mode: str, **kwargs) -> "_BaseParamOffloader":
|
||||
return {
|
||||
"meta": _MetaParamOffloader,
|
||||
"cpu": _CpuParamOffloader,
|
||||
"shm_cpu": _ShmCpuParamOffloader,
|
||||
"sharded_gpu": _ShardedGpuParamOffloader,
|
||||
}[mode](**kwargs)
|
||||
|
||||
def __init__(self, module, param_name):
|
||||
self._module = module
|
||||
self._param_name = param_name
|
||||
|
||||
@property
|
||||
def _param(self):
|
||||
return getattr(self._module, self._param_name)
|
||||
|
||||
def post_init(self):
|
||||
pass
|
||||
|
||||
def create_device_tensor(self):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class _MetaParamOffloader(_BaseParamOffloader):
|
||||
"""Usually used for debugging."""
|
||||
|
||||
def __init__(self, module, param_name):
|
||||
super().__init__(module, param_name)
|
||||
_move_param_to_meta(module, param_name)
|
||||
|
||||
def create_device_tensor(self):
|
||||
return torch.empty_like(self._param.data, device="cuda")
|
||||
|
||||
|
||||
class _CpuParamOffloader(_BaseParamOffloader):
|
||||
def __init__(self, module, param_name):
|
||||
super().__init__(module, param_name)
|
||||
_move_param_to_cpu(self._param, pin_memory=True)
|
||||
|
||||
def create_device_tensor(self):
|
||||
return self._param.to("cuda", non_blocking=True)
|
||||
|
||||
|
||||
class _ShmCpuParamOffloader(_BaseParamOffloader):
|
||||
def __init__(self, module, param_name):
|
||||
super().__init__(module, param_name)
|
||||
self._rank = get_naive_distributed().get_rank()
|
||||
self._world_size = get_naive_distributed().get_world_size()
|
||||
|
||||
from sglang.srt.distributed import get_tensor_model_parallel_world_size
|
||||
|
||||
assert get_tensor_model_parallel_world_size() == 1, "not yet support tp_size!=1"
|
||||
assert (
|
||||
self._param.data.is_contiguous()
|
||||
), f"not yet support non-contiguous tensor {self._param.shape=} {self._param.stride()=}"
|
||||
|
||||
self.shm_cpu_data = get_host_shared_memory_manager().malloc(
|
||||
shape=self._param.shape, dtype=self._param.dtype
|
||||
)
|
||||
|
||||
if self._rank == 0:
|
||||
self.shm_cpu_data.copy_(self._param.data.to("cpu"))
|
||||
self._param.data = self.shm_cpu_data
|
||||
else:
|
||||
_move_param_to_meta(self._module, self._param_name)
|
||||
get_naive_distributed().barrier()
|
||||
|
||||
def post_init(self):
|
||||
if self._rank == 0:
|
||||
assert (
|
||||
self.shm_cpu_data.data_ptr() == self._param.data.data_ptr()
|
||||
), f"{self.shm_cpu_data.data_ptr()=} {self._param.data.data_ptr()=} {self.shm_cpu_data=} {self._param.data=}"
|
||||
|
||||
_move_param_to_meta(self._module, self._param_name)
|
||||
|
||||
def create_device_tensor(self):
|
||||
return self.shm_cpu_data.to("cuda", non_blocking=True)
|
||||
|
||||
|
||||
def update_param(param, new_tensor):
|
||||
"""Update parameter while keeping properties needed by Offloader (e.g. pinned host memory)."""
|
||||
|
||||
if param.device == new_tensor.device:
|
||||
param.data = new_tensor
|
||||
else:
|
||||
assert param.device == torch.device(
|
||||
"cpu"
|
||||
), f"{param.device=} {new_tensor.device=}"
|
||||
param.data = _create_cpu_data(new_tensor, pin_memory=True)
|
||||
|
||||
|
||||
def _move_param_to_cpu(param, pin_memory: bool):
|
||||
param.data = _create_cpu_data(param.data, pin_memory=pin_memory)
|
||||
|
||||
|
||||
def _create_cpu_data(data, pin_memory: bool):
|
||||
cpu_data = _empty_strided_like(
|
||||
data,
|
||||
device="cpu",
|
||||
pin_memory=pin_memory,
|
||||
)
|
||||
cpu_data.copy_(data)
|
||||
return cpu_data
|
||||
|
||||
|
||||
def _move_param_to_meta(module, param_name):
|
||||
old_param = getattr(module, param_name)
|
||||
old_param_type = type(old_param)
|
||||
|
||||
new_data = old_param.data.to("meta")
|
||||
|
||||
if old_param_type == ModelWeightParameter:
|
||||
# manually checked how `w13_weight` and `w2_weight` are constructed
|
||||
new_param = ModelWeightParameter(
|
||||
data=new_data,
|
||||
**{
|
||||
k: getattr(old_param, k)
|
||||
for k in ["input_dim", "output_dim", "weight_loader"]
|
||||
},
|
||||
)
|
||||
elif old_param_type == torch.nn.Parameter:
|
||||
new_param = torch.nn.Parameter(
|
||||
data=new_data,
|
||||
requires_grad=False,
|
||||
)
|
||||
else:
|
||||
raise ValueError(f"Unknown {old_param_type=} {old_param=}")
|
||||
|
||||
setattr(module, param_name, new_param)
|
||||
|
||||
|
||||
def _empty_strided_like(x: torch.Tensor, device, pin_memory=False):
|
||||
return torch.empty_strided(
|
||||
size=x.size(),
|
||||
stride=x.stride(),
|
||||
dtype=x.dtype,
|
||||
layout=x.layout,
|
||||
device=device,
|
||||
pin_memory=pin_memory,
|
||||
)
|
||||
|
||||
|
||||
# ----------------------------------------- ShardedGpu ------------------------------------------------------
|
||||
|
||||
|
||||
# TODO unify with ShmCpu mode
|
||||
class _ShardedGpuParamOffloader(_BaseParamOffloader):
|
||||
def __init__(self, module, param_name):
|
||||
super().__init__(module, param_name)
|
||||
self._rank = get_naive_distributed().get_rank()
|
||||
self._world_size = get_naive_distributed().get_world_size()
|
||||
|
||||
from sglang.srt.distributed import get_tensor_model_parallel_world_size
|
||||
|
||||
assert get_tensor_model_parallel_world_size() == 1, "not yet support tp_size!=1"
|
||||
assert (
|
||||
self._param.data.is_contiguous()
|
||||
), f"not yet support non-contiguous tensor {self._param.shape=} {self._param.stride()=}"
|
||||
|
||||
if self._rank == 0:
|
||||
_move_param_to_cpu(self._param, pin_memory=True)
|
||||
else:
|
||||
_move_param_to_meta(self._module, self._param_name)
|
||||
|
||||
self.sharded_param_handles = None
|
||||
|
||||
def post_init(self):
|
||||
# check again since it may be changed
|
||||
assert (
|
||||
self._param.data.is_contiguous()
|
||||
), f"not yet support non-contiguous tensor {self._param.shape=} {self._param.stride()=}"
|
||||
|
||||
scatter_src = self._param.data
|
||||
|
||||
logger.info(
|
||||
f"[offloader] post_init {scatter_src.nbytes=} {scatter_src.dtype=} {scatter_src.shape=} {torch.cuda.memory_allocated()=}"
|
||||
)
|
||||
|
||||
if self._rank == 0:
|
||||
scatter_src = scatter_src.to("cuda")
|
||||
scatter_list = _even_chunk(scatter_src, self._world_size)
|
||||
|
||||
sharded_param = torch.empty(
|
||||
scatter_list[0].shape, dtype=scatter_list[0].dtype, device="cuda"
|
||||
)
|
||||
self.sharded_param_handles = _create_shared_buffer_tensors(
|
||||
local_tensor=sharded_param
|
||||
)
|
||||
|
||||
get_naive_distributed().scatter(
|
||||
sharded_param, scatter_list if self._rank == 0 else None
|
||||
)
|
||||
|
||||
_move_param_to_meta(self._module, self._param_name)
|
||||
|
||||
def create_device_tensor(self):
|
||||
output = _empty_strided_like(self._param, device="cuda")
|
||||
output_chunks = output.chunk(self._world_size)
|
||||
|
||||
for index in range(self._world_size):
|
||||
src_rank = (self._rank + index) % self._world_size
|
||||
src_buf = self.sharded_param_handles[src_rank]
|
||||
output_chunks[src_rank].copy_(src_buf)
|
||||
|
||||
return output
|
||||
|
||||
|
||||
def _even_chunk(x: torch.Tensor, chunks: int):
|
||||
assert x.shape[0] % chunks == 0, f"{x.shape=} {chunks=}"
|
||||
return list(x.chunk(chunks))
|
||||
|
||||
|
||||
def _create_shared_buffer_tensors(local_tensor: torch.Tensor) -> List[torch.Tensor]:
|
||||
self_rank = get_naive_distributed().get_rank()
|
||||
world_size = get_naive_distributed().get_world_size()
|
||||
|
||||
object_list = get_naive_distributed().all_gather_object(
|
||||
dict(
|
||||
dup_serialized_local_tensor=[
|
||||
(
|
||||
None
|
||||
if interesting_rank == self_rank
|
||||
else MultiprocessingSerializer.serialize(local_tensor)
|
||||
)
|
||||
for interesting_rank in range(world_size)
|
||||
]
|
||||
)
|
||||
)
|
||||
|
||||
output_tensors = []
|
||||
for output_rank in range(world_size):
|
||||
remote_serialized_tensor = object_list[output_rank][
|
||||
"dup_serialized_local_tensor"
|
||||
][self_rank]
|
||||
if output_rank == self_rank:
|
||||
assert remote_serialized_tensor is None
|
||||
output_tensors.append(local_tensor)
|
||||
else:
|
||||
output_tensors.append(
|
||||
MultiprocessingSerializer.deserialize(remote_serialized_tensor)
|
||||
)
|
||||
|
||||
return output_tensors
|
||||
@@ -0,0 +1,92 @@
|
||||
import logging
|
||||
from abc import ABC
|
||||
from contextlib import contextmanager
|
||||
|
||||
try:
|
||||
import torch_memory_saver
|
||||
|
||||
_memory_saver = torch_memory_saver.torch_memory_saver
|
||||
import_error = None
|
||||
except ImportError as e:
|
||||
import_error = e
|
||||
pass
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class TorchMemorySaverAdapter(ABC):
|
||||
@staticmethod
|
||||
def create(enable: bool):
|
||||
if enable and import_error is not None:
|
||||
logger.warning(
|
||||
"enable_memory_saver is enabled, but "
|
||||
"torch-memory-saver is not installed. Please install it "
|
||||
"via `pip3 install torch-memory-saver`. "
|
||||
)
|
||||
raise import_error
|
||||
return (
|
||||
_TorchMemorySaverAdapterReal() if enable else _TorchMemorySaverAdapterNoop()
|
||||
)
|
||||
|
||||
def check_validity(self, caller_name):
|
||||
if not self.enabled:
|
||||
logger.warning(
|
||||
f"`{caller_name}` will not save memory because torch_memory_saver is not enabled. "
|
||||
f"Potential causes: `enable_memory_saver` is false, or torch_memory_saver has installation issues."
|
||||
)
|
||||
|
||||
def configure_subprocess(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def region(self, tag: str, enable_cpu_backup: bool = False):
|
||||
raise NotImplementedError
|
||||
|
||||
def pause(self, tag: str):
|
||||
raise NotImplementedError
|
||||
|
||||
def resume(self, tag: str):
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
def enabled(self):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class _TorchMemorySaverAdapterReal(TorchMemorySaverAdapter):
|
||||
"""Adapter for TorchMemorySaver with tag-based control"""
|
||||
|
||||
def configure_subprocess(self):
|
||||
return torch_memory_saver.configure_subprocess()
|
||||
|
||||
def region(self, tag: str, enable_cpu_backup: bool = False):
|
||||
return _memory_saver.region(tag=tag, enable_cpu_backup=enable_cpu_backup)
|
||||
|
||||
def pause(self, tag: str):
|
||||
return _memory_saver.pause(tag=tag)
|
||||
|
||||
def resume(self, tag: str):
|
||||
return _memory_saver.resume(tag=tag)
|
||||
|
||||
@property
|
||||
def enabled(self):
|
||||
return _memory_saver is not None and _memory_saver.enabled
|
||||
|
||||
|
||||
class _TorchMemorySaverAdapterNoop(TorchMemorySaverAdapter):
|
||||
@contextmanager
|
||||
def configure_subprocess(self):
|
||||
yield
|
||||
|
||||
@contextmanager
|
||||
def region(self, tag: str, enable_cpu_backup: bool = False):
|
||||
yield
|
||||
|
||||
def pause(self, tag: str):
|
||||
pass
|
||||
|
||||
def resume(self, tag: str):
|
||||
pass
|
||||
|
||||
@property
|
||||
def enabled(self):
|
||||
return False
|
||||
Reference in New Issue
Block a user