From 2efe2ef6248ae7bf82c4aca402511e0ee642cacc Mon Sep 17 00:00:00 2001 From: ruixiangw Date: Wed, 18 Feb 2026 16:22:11 +0000 Subject: [PATCH 1/5] packing optimziation with cache to reduce D2H copy --- unsloth/utils/packing.py | 44 ++++++++++++++++++++++++++++++++++++---- 1 file changed, 40 insertions(+), 4 deletions(-) diff --git a/unsloth/utils/packing.py b/unsloth/utils/packing.py index 63a57c04da1..a8ca51997e2 100644 --- a/unsloth/utils/packing.py +++ b/unsloth/utils/packing.py @@ -36,6 +36,15 @@ _XFORMERS_MASK_CACHE_MAXSIZE = 32 _XFORMERS_MASK_CACHE: OrderedDict[Tuple[Tuple[int, ...], int], Any] = OrderedDict() +# Cache for get_packed_info_from_kwargs to avoid repeated D2H sync across layers +_PACKED_INFO_CACHE: dict = {"seq_lengths": None, "result": None} + +# Cache for build_sdpa_packed_attention_mask to avoid repeated D2H sync across layers +_SDPA_MASK_CACHE: dict = {"seq_lengths": None, "params": None, "mask": None} + +# Cache for build_xformers_block_causal_mask to avoid repeated D2H sync across layers +_XFORMERS_BLOCK_MASK_CACHE: dict = {"seq_lengths": None, "params": None, "mask": None} + 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 + if _PACKED_INFO_CACHE["seq_lengths"] is seq_lengths: + return _PACKED_INFO_CACHE["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["seq_lengths"] = seq_lengths + _PACKED_INFO_CACHE["result"] = result + return result def build_xformers_block_causal_mask( @@ -243,11 +257,23 @@ 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 + params = (sliding_window,) + if ( + _XFORMERS_BLOCK_MASK_CACHE["seq_lengths"] is seq_lengths + and _XFORMERS_BLOCK_MASK_CACHE["params"] == params + ): + return _XFORMERS_BLOCK_MASK_CACHE["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["seq_lengths"] = seq_lengths + _XFORMERS_BLOCK_MASK_CACHE["params"] = params + _XFORMERS_BLOCK_MASK_CACHE["mask"] = mask else: mask = base_mask @@ -269,6 +295,11 @@ def build_sdpa_packed_attention_mask( sliding_window: Optional[int] = None, ) -> torch.Tensor: seq_lengths, _, _ = seq_info + + params = (dtype, device, sliding_window) + if _SDPA_MASK_CACHE["seq_lengths"] is seq_lengths and _SDPA_MASK_CACHE["params"] == params: + return _SDPA_MASK_CACHE["mask"] + total_tokens = int(seq_lengths.sum().item()) mask = torch.full( (total_tokens, total_tokens), @@ -297,7 +328,12 @@ 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["seq_lengths"] = seq_lengths + _SDPA_MASK_CACHE["params"] = params + _SDPA_MASK_CACHE["mask"] = result + return result def _normalize_packed_lengths( From 837b5f8e67c02165f0bb1760e02d6f5750555913 Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 6 Mar 2026 21:24:02 +0000 Subject: [PATCH 2/5] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- unsloth/utils/packing.py | 9 ++++++--- 1 file changed, 6 insertions(+), 3 deletions(-) diff --git a/unsloth/utils/packing.py b/unsloth/utils/packing.py index a8ca51997e2..04ce05b9868 100644 --- a/unsloth/utils/packing.py +++ b/unsloth/utils/packing.py @@ -264,7 +264,7 @@ def build_xformers_block_causal_mask( and _XFORMERS_BLOCK_MASK_CACHE["params"] == params ): return _XFORMERS_BLOCK_MASK_CACHE["mask"] - + lengths_tensor = seq_lengths.to("cpu", torch.int32) if lengths_tensor.numel() == 0: return None @@ -297,7 +297,10 @@ def build_sdpa_packed_attention_mask( seq_lengths, _, _ = seq_info params = (dtype, device, sliding_window) - if _SDPA_MASK_CACHE["seq_lengths"] is seq_lengths and _SDPA_MASK_CACHE["params"] == params: + if ( + _SDPA_MASK_CACHE["seq_lengths"] is seq_lengths + and _SDPA_MASK_CACHE["params"] == params + ): return _SDPA_MASK_CACHE["mask"] total_tokens = int(seq_lengths.sum().item()) @@ -328,7 +331,7 @@ def build_sdpa_packed_attention_mask( block = block.masked_fill(window_mask, float("-inf")) mask[offset : offset + length, offset : offset + length] = block offset += length - + result = mask.unsqueeze(0).unsqueeze(0) _SDPA_MASK_CACHE["seq_lengths"] = seq_lengths _SDPA_MASK_CACHE["params"] = params From 67930ba5ccace801850a0014c807a2ec6822a535 Mon Sep 17 00:00:00 2001 From: ruixiang Date: Thu, 12 Mar 2026 02:54:47 +0800 Subject: [PATCH 3/5] cache per device to avoid race condition for multi-gpu --- unsloth/utils/packing.py | 47 ++++++++++++++++++++-------------------- 1 file changed, 24 insertions(+), 23 deletions(-) diff --git a/unsloth/utils/packing.py b/unsloth/utils/packing.py index 04ce05b9868..4cbccea4cd8 100644 --- a/unsloth/utils/packing.py +++ b/unsloth/utils/packing.py @@ -36,14 +36,14 @@ _XFORMERS_MASK_CACHE_MAXSIZE = 32 _XFORMERS_MASK_CACHE: OrderedDict[Tuple[Tuple[int, ...], int], Any] = OrderedDict() -# Cache for get_packed_info_from_kwargs to avoid repeated D2H sync across layers -_PACKED_INFO_CACHE: dict = {"seq_lengths": None, "result": None} +# Cache per device for get_packed_info_from_kwargs to avoid repeated D2H sync across layers +_PACKED_INFO_CACHE: dict = {} -# Cache for build_sdpa_packed_attention_mask to avoid repeated D2H sync across layers -_SDPA_MASK_CACHE: dict = {"seq_lengths": None, "params": None, "mask": None} +# Cache per device for build_sdpa_packed_attention_mask to avoid repeated D2H sync across layers +_SDPA_MASK_CACHE: dict = {} -# Cache for build_xformers_block_causal_mask to avoid repeated D2H sync across layers -_XFORMERS_BLOCK_MASK_CACHE: dict = {"seq_lengths": None, "params": None, "mask": None} +# 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: @@ -233,8 +233,9 @@ def get_packed_info_from_kwargs( if seq_lengths is None: return None - if _PACKED_INFO_CACHE["seq_lengths"] is seq_lengths: - return _PACKED_INFO_CACHE["result"] + 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.zeros(lengths.numel() + 1, dtype = torch.int32, device = device) @@ -242,8 +243,7 @@ def get_packed_info_from_kwargs( max_seqlen = int(lengths.max().item()) result = (lengths, cu_seqlens, max_seqlen) - _PACKED_INFO_CACHE["seq_lengths"] = seq_lengths - _PACKED_INFO_CACHE["result"] = result + _PACKED_INFO_CACHE[device] = {"seq_lengths": seq_lengths, "result": result} return result @@ -258,12 +258,15 @@ def build_xformers_block_causal_mask( 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 ( - _XFORMERS_BLOCK_MASK_CACHE["seq_lengths"] is seq_lengths - and _XFORMERS_BLOCK_MASK_CACHE["params"] == params + entry is not None + and entry["seq_lengths"] is seq_lengths + and entry["params"] == params ): - return _XFORMERS_BLOCK_MASK_CACHE["mask"] + return entry["mask"] lengths_tensor = seq_lengths.to("cpu", torch.int32) if lengths_tensor.numel() == 0: @@ -271,9 +274,7 @@ def build_xformers_block_causal_mask( lengths = tuple(int(x) for x in lengths_tensor.tolist()) mask = _get_cached_block_mask(lengths, sliding_window) - _XFORMERS_BLOCK_MASK_CACHE["seq_lengths"] = seq_lengths - _XFORMERS_BLOCK_MASK_CACHE["params"] = params - _XFORMERS_BLOCK_MASK_CACHE["mask"] = mask + _XFORMERS_BLOCK_MASK_CACHE[device] = {"seq_lengths": seq_lengths, "params": params, "mask": mask} else: mask = base_mask @@ -296,12 +297,14 @@ def build_sdpa_packed_attention_mask( ) -> torch.Tensor: seq_lengths, _, _ = seq_info - params = (dtype, device, sliding_window) + params = (dtype, sliding_window) + entry = _SDPA_MASK_CACHE.get(device) if ( - _SDPA_MASK_CACHE["seq_lengths"] is seq_lengths - and _SDPA_MASK_CACHE["params"] == params + entry is not None + and entry["seq_lengths"] is seq_lengths + and entry["params"] == params ): - return _SDPA_MASK_CACHE["mask"] + return entry["mask"] total_tokens = int(seq_lengths.sum().item()) mask = torch.full( @@ -333,9 +336,7 @@ def build_sdpa_packed_attention_mask( offset += length result = mask.unsqueeze(0).unsqueeze(0) - _SDPA_MASK_CACHE["seq_lengths"] = seq_lengths - _SDPA_MASK_CACHE["params"] = params - _SDPA_MASK_CACHE["mask"] = result + _SDPA_MASK_CACHE[device] = {"seq_lengths": seq_lengths, "params": params, "mask": result} return result From ecc5106f23f4be3d66b2ccd245b7d067488f5c7c Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Wed, 11 Mar 2026 19:07:46 +0000 Subject: [PATCH 4/5] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- unsloth/utils/packing.py | 12 ++++++++++-- 1 file changed, 10 insertions(+), 2 deletions(-) diff --git a/unsloth/utils/packing.py b/unsloth/utils/packing.py index 4cbccea4cd8..227896013ae 100644 --- a/unsloth/utils/packing.py +++ b/unsloth/utils/packing.py @@ -274,7 +274,11 @@ def build_xformers_block_causal_mask( 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} + _XFORMERS_BLOCK_MASK_CACHE[device] = { + "seq_lengths": seq_lengths, + "params": params, + "mask": mask, + } else: mask = base_mask @@ -336,7 +340,11 @@ def build_sdpa_packed_attention_mask( offset += length result = mask.unsqueeze(0).unsqueeze(0) - _SDPA_MASK_CACHE[device] = {"seq_lengths": seq_lengths, "params": params, "mask": result} + _SDPA_MASK_CACHE[device] = { + "seq_lengths": seq_lengths, + "params": params, + "mask": result, + } return result From 60da93df798e2d549a58b6ca5b38c53749722deb Mon Sep 17 00:00:00 2001 From: ruixiang Date: Thu, 12 Mar 2026 03:23:36 +0800 Subject: [PATCH 5/5] add cache freeing up func --- unsloth/utils/packing.py | 8 ++++++++ 1 file changed, 8 insertions(+) diff --git a/unsloth/utils/packing.py b/unsloth/utils/packing.py index 227896013ae..bd6631d4618 100644 --- a/unsloth/utils/packing.py +++ b/unsloth/utils/packing.py @@ -389,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", @@ -399,4 +406,5 @@ def mask_packed_sequence_boundaries( "build_xformers_block_causal_mask", "build_sdpa_packed_attention_mask", "mask_packed_sequence_boundaries", + "clear_packed_caches", ]