[Feature] support fastsafetensors (#15091)
Signed-off-by: Xuchun Shang <xuchun.shang@gmail.com> Co-authored-by: Xuchun Shang <xuchun.shang@gmail.com>
This commit is contained in:
@@ -29,6 +29,7 @@ class LoadFormat(str, enum.Enum):
|
|||||||
REMOTE_INSTANCE = "remote_instance"
|
REMOTE_INSTANCE = "remote_instance"
|
||||||
RDMA = "rdma"
|
RDMA = "rdma"
|
||||||
LOCAL_CACHED = "local_cached"
|
LOCAL_CACHED = "local_cached"
|
||||||
|
FASTSAFETENSORS = "fastsafetensors"
|
||||||
PRIVATE = "private"
|
PRIVATE = "private"
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -89,6 +89,7 @@ from sglang.srt.environ import envs
|
|||||||
from sglang.srt.model_loader.weight_utils import (
|
from sglang.srt.model_loader.weight_utils import (
|
||||||
download_safetensors_index_file_from_hf,
|
download_safetensors_index_file_from_hf,
|
||||||
download_weights_from_hf,
|
download_weights_from_hf,
|
||||||
|
fastsafetensors_weights_iterator,
|
||||||
filter_duplicate_safetensors_files,
|
filter_duplicate_safetensors_files,
|
||||||
filter_files_not_needed_for_inference,
|
filter_files_not_needed_for_inference,
|
||||||
get_gguf_extra_tensor_names,
|
get_gguf_extra_tensor_names,
|
||||||
@@ -386,7 +387,10 @@ class DefaultModelLoader(BaseModelLoader):
|
|||||||
# Some quantized models use .pt files for storing the weights.
|
# Some quantized models use .pt files for storing the weights.
|
||||||
if load_format == LoadFormat.AUTO:
|
if load_format == LoadFormat.AUTO:
|
||||||
allow_patterns = ["*.safetensors", "*.bin"]
|
allow_patterns = ["*.safetensors", "*.bin"]
|
||||||
elif load_format == LoadFormat.SAFETENSORS:
|
elif (
|
||||||
|
load_format == LoadFormat.SAFETENSORS
|
||||||
|
or load_format == LoadFormat.FASTSAFETENSORS
|
||||||
|
):
|
||||||
use_safetensors = True
|
use_safetensors = True
|
||||||
allow_patterns = ["*.safetensors"]
|
allow_patterns = ["*.safetensors"]
|
||||||
elif load_format == LoadFormat.MISTRAL:
|
elif load_format == LoadFormat.MISTRAL:
|
||||||
@@ -474,7 +478,11 @@ class DefaultModelLoader(BaseModelLoader):
|
|||||||
get_global_server_args().weight_loader_disable_mmap
|
get_global_server_args().weight_loader_disable_mmap
|
||||||
)
|
)
|
||||||
|
|
||||||
if extra_config.get("enable_multithread_load"):
|
if self.load_config.load_format == LoadFormat.FASTSAFETENSORS:
|
||||||
|
weights_iterator = fastsafetensors_weights_iterator(
|
||||||
|
hf_weights_files,
|
||||||
|
)
|
||||||
|
elif extra_config.get("enable_multithread_load"):
|
||||||
weights_iterator = multi_thread_safetensors_weights_iterator(
|
weights_iterator = multi_thread_safetensors_weights_iterator(
|
||||||
hf_weights_files,
|
hf_weights_files,
|
||||||
max_workers=extra_config.get(
|
max_workers=extra_config.get(
|
||||||
|
|||||||
@@ -49,6 +49,21 @@ from sglang.srt.model_loader.weight_validation import (
|
|||||||
from sglang.srt.utils import find_local_repo_dir, log_info_on_rank0, print_warning_once
|
from sglang.srt.utils import find_local_repo_dir, log_info_on_rank0, print_warning_once
|
||||||
from sglang.utils import is_in_ci
|
from sglang.utils import is_in_ci
|
||||||
|
|
||||||
|
try:
|
||||||
|
from fastsafetensors import SafeTensorsFileLoader, SingleGroup
|
||||||
|
except ImportError:
|
||||||
|
|
||||||
|
class PlaceholderModule:
|
||||||
|
def __init__(self, name):
|
||||||
|
self.name = name
|
||||||
|
|
||||||
|
def __getattr__(self, name):
|
||||||
|
raise ImportError(f"Please install {self.name}")
|
||||||
|
|
||||||
|
fastsafetensors = PlaceholderModule("fastsafetensors")
|
||||||
|
SafeTensorsFileLoader = None
|
||||||
|
SingleGroup = None
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
# use system-level temp directory for file locks, so that multiple users
|
# use system-level temp directory for file locks, so that multiple users
|
||||||
@@ -826,6 +841,61 @@ def safetensors_weights_iterator(
|
|||||||
yield name, f.get_tensor(name)
|
yield name, f.get_tensor(name)
|
||||||
|
|
||||||
|
|
||||||
|
def fastsafetensors_weights_iterator(
|
||||||
|
hf_weights_files: List[str],
|
||||||
|
) -> Generator[Tuple[str, torch.Tensor], None, None]:
|
||||||
|
"""
|
||||||
|
Iterate over the weights in the model safetensor files
|
||||||
|
using fastsafetensor library to accelerate loading via GPU Direct Storage (if available).
|
||||||
|
"""
|
||||||
|
if SafeTensorsFileLoader is None:
|
||||||
|
raise ImportError(
|
||||||
|
"Please install fastsafetensors via `pip install fastsafetensors`"
|
||||||
|
)
|
||||||
|
|
||||||
|
if torch.distributed.is_initialized():
|
||||||
|
pg = torch.distributed.group.WORLD
|
||||||
|
else:
|
||||||
|
pg = SingleGroup()
|
||||||
|
|
||||||
|
try:
|
||||||
|
rank = pg.rank()
|
||||||
|
except Exception:
|
||||||
|
rank = 0
|
||||||
|
|
||||||
|
device = torch.device(f"cuda:{rank}")
|
||||||
|
|
||||||
|
weight_files_sub_lists = [
|
||||||
|
hf_weights_files[i : i + pg.size()]
|
||||||
|
for i in range(0, len(hf_weights_files), pg.size())
|
||||||
|
]
|
||||||
|
|
||||||
|
_BAR_FORMAT = (
|
||||||
|
"{l_bar}{bar}| {n_fmt}/{total_fmt} [{elapsed}<{remaining}, {rate_fmt}]"
|
||||||
|
)
|
||||||
|
|
||||||
|
for f_list in tqdm(
|
||||||
|
weight_files_sub_lists,
|
||||||
|
desc="Loading safetensors using Fastsafetensor loader",
|
||||||
|
disable=False,
|
||||||
|
bar_format=_BAR_FORMAT,
|
||||||
|
):
|
||||||
|
loader = SafeTensorsFileLoader(pg, device)
|
||||||
|
rank_file_map = {i: [f] for i, f in enumerate(f_list)}
|
||||||
|
loader.add_filenames(rank_file_map)
|
||||||
|
try:
|
||||||
|
fb = loader.copy_files_to_device()
|
||||||
|
try:
|
||||||
|
keys = list(fb.key_to_rank_lidx.keys())
|
||||||
|
for k in keys:
|
||||||
|
t = fb.get_tensor(k)
|
||||||
|
yield k, t
|
||||||
|
finally:
|
||||||
|
pass
|
||||||
|
finally:
|
||||||
|
loader.close()
|
||||||
|
|
||||||
|
|
||||||
def multi_thread_safetensors_weights_iterator(
|
def multi_thread_safetensors_weights_iterator(
|
||||||
hf_weights_files: List[str],
|
hf_weights_files: List[str],
|
||||||
is_all_weights_sharded: bool = False,
|
is_all_weights_sharded: bool = False,
|
||||||
|
|||||||
@@ -84,6 +84,7 @@ LOAD_FORMAT_CHOICES = [
|
|||||||
"flash_rl",
|
"flash_rl",
|
||||||
"remote",
|
"remote",
|
||||||
"remote_instance",
|
"remote_instance",
|
||||||
|
"fastsafetensors",
|
||||||
"private",
|
"private",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user