Cache packed sequence metadata to reduce D2H syncs across layers (#4243)
* packing optimziation with cache to reduce D2H copy * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * cache per device to avoid race condition for multi-gpu * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * add cache freeing up func --------- Co-authored-by: ruixiangw <ruixiangw@nvidia.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> Co-authored-by: ruixiang <wangruixiang07@outlook.com>
This commit is contained in:
parent
3e1469a55a
commit
12f8525bc6
1 changed files with 60 additions and 4 deletions
|
|
@ -36,6 +36,15 @@ except Exception:
|
|||
_XFORMERS_MASK_CACHE_MAXSIZE = 32
|
||||
_XFORMERS_MASK_CACHE: OrderedDict[Tuple[Tuple[int, ...], int], Any] = OrderedDict()
|
||||
|
||||
# Cache per device for get_packed_info_from_kwargs to avoid repeated D2H sync across layers
|
||||
_PACKED_INFO_CACHE: dict = {}
|
||||
|
||||
# Cache per device for build_sdpa_packed_attention_mask to avoid repeated D2H sync across layers
|
||||
_SDPA_MASK_CACHE: dict = {}
|
||||
|
||||
# Cache per device for build_xformers_block_causal_mask to avoid repeated D2H sync across layers
|
||||
_XFORMERS_BLOCK_MASK_CACHE: dict = {}
|
||||
|
||||
|
||||
def _window_cache_key(sliding_window: Optional[int]) -> int:
|
||||
if sliding_window is None or sliding_window <= 0:
|
||||
|
|
@ -224,13 +233,18 @@ def get_packed_info_from_kwargs(
|
|||
if seq_lengths is None:
|
||||
return None
|
||||
|
||||
entry = _PACKED_INFO_CACHE.get(device)
|
||||
if entry is not None and entry["seq_lengths"] is seq_lengths:
|
||||
return entry["result"]
|
||||
|
||||
lengths = seq_lengths.to(device = device, dtype = torch.int32, non_blocking = True)
|
||||
cu_seqlens = torch.empty(lengths.numel() + 1, dtype = torch.int32, device = device)
|
||||
cu_seqlens[0] = 0
|
||||
cu_seqlens = torch.zeros(lengths.numel() + 1, dtype = torch.int32, device = device)
|
||||
torch.cumsum(lengths, dim = 0, dtype = torch.int32, out = cu_seqlens[1:])
|
||||
|
||||
max_seqlen = int(lengths.max().item())
|
||||
return lengths, cu_seqlens, max_seqlen
|
||||
result = (lengths, cu_seqlens, max_seqlen)
|
||||
_PACKED_INFO_CACHE[device] = {"seq_lengths": seq_lengths, "result": result}
|
||||
return result
|
||||
|
||||
|
||||
def build_xformers_block_causal_mask(
|
||||
|
|
@ -243,11 +257,28 @@ def build_xformers_block_causal_mask(
|
|||
return None
|
||||
if seq_info is not None:
|
||||
seq_lengths, _, _ = seq_info
|
||||
# Cache the mask to avoid repeated D2H sync across layers
|
||||
device = seq_lengths.device
|
||||
params = (sliding_window,)
|
||||
entry = _XFORMERS_BLOCK_MASK_CACHE.get(device)
|
||||
if (
|
||||
entry is not None
|
||||
and entry["seq_lengths"] is seq_lengths
|
||||
and entry["params"] == params
|
||||
):
|
||||
return entry["mask"]
|
||||
|
||||
lengths_tensor = seq_lengths.to("cpu", torch.int32)
|
||||
if lengths_tensor.numel() == 0:
|
||||
return None
|
||||
lengths = tuple(int(x) for x in lengths_tensor.tolist())
|
||||
mask = _get_cached_block_mask(lengths, sliding_window)
|
||||
|
||||
_XFORMERS_BLOCK_MASK_CACHE[device] = {
|
||||
"seq_lengths": seq_lengths,
|
||||
"params": params,
|
||||
"mask": mask,
|
||||
}
|
||||
else:
|
||||
mask = base_mask
|
||||
|
||||
|
|
@ -269,6 +300,16 @@ def build_sdpa_packed_attention_mask(
|
|||
sliding_window: Optional[int] = None,
|
||||
) -> torch.Tensor:
|
||||
seq_lengths, _, _ = seq_info
|
||||
|
||||
params = (dtype, sliding_window)
|
||||
entry = _SDPA_MASK_CACHE.get(device)
|
||||
if (
|
||||
entry is not None
|
||||
and entry["seq_lengths"] is seq_lengths
|
||||
and entry["params"] == params
|
||||
):
|
||||
return entry["mask"]
|
||||
|
||||
total_tokens = int(seq_lengths.sum().item())
|
||||
mask = torch.full(
|
||||
(total_tokens, total_tokens),
|
||||
|
|
@ -297,7 +338,14 @@ def build_sdpa_packed_attention_mask(
|
|||
block = block.masked_fill(window_mask, float("-inf"))
|
||||
mask[offset : offset + length, offset : offset + length] = block
|
||||
offset += length
|
||||
return mask.unsqueeze(0).unsqueeze(0)
|
||||
|
||||
result = mask.unsqueeze(0).unsqueeze(0)
|
||||
_SDPA_MASK_CACHE[device] = {
|
||||
"seq_lengths": seq_lengths,
|
||||
"params": params,
|
||||
"mask": result,
|
||||
}
|
||||
return result
|
||||
|
||||
|
||||
def _normalize_packed_lengths(
|
||||
|
|
@ -341,6 +389,13 @@ def mask_packed_sequence_boundaries(
|
|||
return True
|
||||
|
||||
|
||||
def clear_packed_caches():
|
||||
"""Release cached masks/metadata to free device memory."""
|
||||
_PACKED_INFO_CACHE.clear()
|
||||
_SDPA_MASK_CACHE.clear()
|
||||
_XFORMERS_BLOCK_MASK_CACHE.clear()
|
||||
|
||||
|
||||
__all__ = [
|
||||
"configure_sample_packing",
|
||||
"configure_padding_free",
|
||||
|
|
@ -351,4 +406,5 @@ __all__ = [
|
|||
"build_xformers_block_causal_mask",
|
||||
"build_sdpa_packed_attention_mask",
|
||||
"mask_packed_sequence_boundaries",
|
||||
"clear_packed_caches",
|
||||
]
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue