Batch copy_ with torch._foreach_copy_ (#18558)
This commit is contained in:
@@ -23,7 +23,7 @@ import os
|
|||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
from dataclasses import dataclass
|
from dataclasses import dataclass
|
||||||
from functools import partial
|
from functools import partial
|
||||||
from typing import TYPE_CHECKING, Callable, Dict, Optional, Union
|
from typing import TYPE_CHECKING, Callable, Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
import tqdm
|
import tqdm
|
||||||
@@ -93,6 +93,19 @@ if TYPE_CHECKING:
|
|||||||
from sglang.srt.model_executor.model_runner import ModelRunner
|
from sglang.srt.model_executor.model_runner import ModelRunner
|
||||||
|
|
||||||
|
|
||||||
|
def _grouped_foreach_copy_(dsts: List[torch.Tensor], srcs: List[torch.Tensor]) -> None:
|
||||||
|
"""Call torch._foreach_copy_ grouped by (dst_dtype, src_dtype) pairs."""
|
||||||
|
groups: Dict[Tuple[torch.dtype, torch.dtype], Tuple[List, List]] = {}
|
||||||
|
for dst, src in zip(dsts, srcs):
|
||||||
|
key = (dst.dtype, src.dtype)
|
||||||
|
if key not in groups:
|
||||||
|
groups[key] = ([], [])
|
||||||
|
groups[key][0].append(dst)
|
||||||
|
groups[key][1].append(src)
|
||||||
|
for group_dsts, group_srcs in groups.values():
|
||||||
|
torch._foreach_copy_(group_dsts, group_srcs)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class DecodeInputBuffers(ForwardInputBuffers):
|
class DecodeInputBuffers(ForwardInputBuffers):
|
||||||
|
|
||||||
@@ -234,34 +247,42 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
if self.mamba_track_mask is not None:
|
if self.mamba_track_mask is not None:
|
||||||
self.mamba_track_mask.fill_(False)
|
self.mamba_track_mask.fill_(False)
|
||||||
|
|
||||||
# Common inputs
|
# Build batched copy lists for all GPU tensors.
|
||||||
self.input_ids[:raw_num_token].copy_(forward_batch.input_ids)
|
dsts = [
|
||||||
self.req_pool_indices[:raw_bs].copy_(forward_batch.req_pool_indices)
|
self.input_ids[:raw_num_token],
|
||||||
self.seq_lens[:raw_bs].copy_(forward_batch.seq_lens)
|
self.req_pool_indices[:raw_bs],
|
||||||
self.out_cache_loc[:raw_num_token].copy_(forward_batch.out_cache_loc)
|
self.seq_lens[:raw_bs],
|
||||||
self.positions[:raw_num_token].copy_(forward_batch.positions)
|
self.out_cache_loc[:raw_num_token],
|
||||||
|
self.positions[:raw_num_token],
|
||||||
|
]
|
||||||
|
srcs = [
|
||||||
|
forward_batch.input_ids,
|
||||||
|
forward_batch.req_pool_indices,
|
||||||
|
forward_batch.seq_lens,
|
||||||
|
forward_batch.out_cache_loc,
|
||||||
|
forward_batch.positions,
|
||||||
|
]
|
||||||
|
|
||||||
if (
|
if (
|
||||||
self.mamba_track_indices is not None
|
self.mamba_track_indices is not None
|
||||||
and forward_batch.mamba_track_indices is not None
|
and forward_batch.mamba_track_indices is not None
|
||||||
):
|
):
|
||||||
self.mamba_track_indices[:raw_bs].copy_(forward_batch.mamba_track_indices)
|
dsts.append(self.mamba_track_indices[:raw_bs])
|
||||||
|
srcs.append(forward_batch.mamba_track_indices)
|
||||||
if (
|
if (
|
||||||
self.mamba_track_mask is not None
|
self.mamba_track_mask is not None
|
||||||
and forward_batch.mamba_track_mask is not None
|
and forward_batch.mamba_track_mask is not None
|
||||||
):
|
):
|
||||||
self.mamba_track_mask[:raw_bs].copy_(forward_batch.mamba_track_mask)
|
dsts.append(self.mamba_track_mask[:raw_bs])
|
||||||
|
srcs.append(forward_batch.mamba_track_mask)
|
||||||
if forward_batch.seq_lens_cpu is not None:
|
|
||||||
if bs != raw_bs:
|
|
||||||
self.seq_lens_cpu.fill_(seq_len_fill_value)
|
|
||||||
self.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
|
|
||||||
|
|
||||||
if self.encoder_lens is not None and forward_batch.encoder_lens is not None:
|
if self.encoder_lens is not None and forward_batch.encoder_lens is not None:
|
||||||
self.encoder_lens[:raw_bs].copy_(forward_batch.encoder_lens)
|
dsts.append(self.encoder_lens[:raw_bs])
|
||||||
|
srcs.append(forward_batch.encoder_lens)
|
||||||
|
|
||||||
if forward_batch.mrope_positions is not None:
|
if forward_batch.mrope_positions is not None:
|
||||||
self.mrope_positions[:, :raw_num_token].copy_(forward_batch.mrope_positions)
|
dsts.append(self.mrope_positions[:, :raw_num_token])
|
||||||
|
srcs.append(forward_batch.mrope_positions)
|
||||||
|
|
||||||
if require_gathered_buffer:
|
if require_gathered_buffer:
|
||||||
self.global_num_tokens_gpu.fill_(bs * num_tokens_per_bs)
|
self.global_num_tokens_gpu.fill_(bs * num_tokens_per_bs)
|
||||||
@@ -274,16 +295,28 @@ class DecodeInputBuffers(ForwardInputBuffers):
|
|||||||
global_num_token_non_padded=forward_batch.num_token_non_padded,
|
global_num_token_non_padded=forward_batch.num_token_non_padded,
|
||||||
num_tokens_per_dp=num_tokens_per_dp,
|
num_tokens_per_dp=num_tokens_per_dp,
|
||||||
)
|
)
|
||||||
self.num_token_non_padded.copy_(local)
|
dsts.append(self.num_token_non_padded)
|
||||||
|
srcs.append(local)
|
||||||
else:
|
else:
|
||||||
self.num_token_non_padded.copy_(forward_batch.num_token_non_padded)
|
dsts.append(self.num_token_non_padded)
|
||||||
|
srcs.append(forward_batch.num_token_non_padded)
|
||||||
|
|
||||||
# Pipeline-parallel proxy tensors.
|
# Pipeline-parallel proxy tensors.
|
||||||
if pp_proxy_tensors is not None and self.pp_proxy_tensors is not None:
|
if pp_proxy_tensors is not None and self.pp_proxy_tensors is not None:
|
||||||
for key, buf in self.pp_proxy_tensors.items():
|
for key, buf in self.pp_proxy_tensors.items():
|
||||||
src = pp_proxy_tensors.tensors[key]
|
src = pp_proxy_tensors.tensors[key]
|
||||||
dim = src.shape[0]
|
dim = src.shape[0]
|
||||||
buf[:dim].copy_(src)
|
dsts.append(buf[:dim])
|
||||||
|
srcs.append(src)
|
||||||
|
|
||||||
|
# Batch all GPU copies, grouped by dtype pair.
|
||||||
|
_grouped_foreach_copy_(dsts, srcs)
|
||||||
|
|
||||||
|
# CPU tensor copy (cannot be batched with GPU tensors).
|
||||||
|
if forward_batch.seq_lens_cpu is not None:
|
||||||
|
if bs != raw_bs:
|
||||||
|
self.seq_lens_cpu.fill_(seq_len_fill_value)
|
||||||
|
self.seq_lens_cpu[:raw_bs].copy_(forward_batch.seq_lens_cpu)
|
||||||
|
|
||||||
|
|
||||||
# Detect whether the current forward pass is in capture mode
|
# Detect whether the current forward pass is in capture mode
|
||||||
|
|||||||
Reference in New Issue
Block a user