已合并
[master] 优化 DSA 长序列上下文并行显存 #1282
[master] 优化 DSA 长序列上下文并行显存 #1282
已合并
lzy0920232创建于 8月29日
共 2 个文件变更+258-6
@@ -53,6 +53,134 @@ _SUPPORTED_LOSS_VARIANTS = ("sparse", "dense")
53_DEFAULT_ARG_INDEX = object()53_DEFAULT_ARG_INDEX = object()
54 54 
55 55 
56+def _move_dim_to_front(value: Any, dim: int) -> tuple[Any, list[int]]:
57+ """Move ``dim`` to the leading position and return the inverse permutation."""
58+ order = [dim] + [index for index in range(len(value.shape)) if index != dim]
59+ inverse = [0] * len(order)
60+ for new_index, old_index in enumerate(order):
61+ inverse[old_index] = new_index
62+ return value.permute(order).contiguous(), inverse
63+ 
64+ 
65+class _DSASequenceReplicateGradientBridge(platform.Function):
66+ """Reuse one gathered forward value while preserving per-consumer backward."""
67+ 
68+ @staticmethod
69+ def forward(
70+ ctx: Any,
71+ local_value: Any,
72+ replicated_value: Any,
73+ group: Any,
74+ world_size: int,
75+ seq_dim: int,
76+ ) -> Any:
77+ """Return the detached shared value and retain only collective metadata."""
78+ del local_value
79+ ctx.group = group
80+ ctx.world_size = world_size
81+ ctx.seq_dim = seq_dim
82+ return replicated_value
83+ 
84+ @staticmethod
85+ def backward( # pylint: disable=arguments-differ
86+ ctx: Any, grad_output: Any
87+ ) -> tuple[Any, None, None, None, None]:
88+ """Reduce-scatter this consumer's gradient independently to its local input."""
89+ if ctx.world_size == 1:
90+ return grad_output, None, None, None, None
91+ grad_front, inverse = _move_dim_to_front(grad_output, ctx.seq_dim)
92+ output_shape = list(grad_front.shape)
93+ if output_shape[0] % ctx.world_size != 0:
94+ raise ValueError(
95+ "DSA shared replicate backward requires a divisible sequence dimension, "
96+ f"got {output_shape[0]} and CP size {ctx.world_size}."
97+ )
98+ output_shape[0] //= ctx.world_size
99+ local_grad, work = platform.reduce_scatter_single(
100+ grad_front, output_shape, ctx.group, async_op=False
101+ )
102+ if work is not None:
103+ work.wait()
104+ local_grad = local_grad.permute(inverse).contiguous()
105+ return local_grad, None, None, None, None
106+ 
107+ 
108+class DSASequenceReplicateCache:
109+ """Share semantically identical CP AllGather results within one DSA layer.
110+ 
111+ Callers assign stable semantic slot names only to local tensors that are
112+ mathematically identical views of the same activation. The first consumer
113+ performs the sequence AllGather and later consumers reuse that storage.
114+ Each consumer receives an independent gradient bridge so its reverse
115+ ReduceScatter remains separate, matching the pre-cache accumulation order.
116+ 
117+ The cache is deliberately scoped to one sparse-attention/indexer-loss
118+ forward interval. :meth:`begin` drops stale entries before sparse attention
119+ and :meth:`clear` releases references after indexer-loss inputs are built.
120+ """
121+ 
122+ def __init__(self) -> None:
123+ """Initialize an empty per-layer communication cache."""
124+ self._values: dict[str, Any] = {}
125+ self._signatures: dict[str, tuple] = {}
126+ 
127+ @staticmethod
128+ def _signature(value: Any, device_mesh: DeviceMesh, seq_dim: int) -> tuple:
129+ """Build a metadata-only signature without synchronizing the device."""
130+ if isinstance(value, DTensor):
131+ shape = tuple(value.local_shape)
132+ dtype = value.dtype
133+ else:
134+ shape = tuple(value.shape)
135+ dtype = value.dtype
136+ return shape, dtype, tuple(device_mesh.rank_list), seq_dim
137+ 
138+ def begin(self) -> None:
139+ """Start a new sparse-attention interval and discard stale references."""
140+ self.clear()
141+ 
142+ def replicate(self, slot_name: str, value: Any, device_mesh: DeviceMesh, seq_dim: int) -> Any:
143+ """Return the shared replicated DTensor for one semantic activation."""
144+ if not _is_tensor_or_dtensor(value):
145+ return value
146+ signature = self._signature(value, device_mesh, seq_dim)
147+ if slot_name in self._values:
148+ if self._signatures[slot_name] != signature:
149+ raise ValueError(
150+ f"DSA shared replicate slot {slot_name!r} received incompatible "
151+ f"metadata: expected {self._signatures[slot_name]}, got {signature}."
152+ )
153+ replicated = self._values[slot_name]
154+ else:
155+ replicated = _to_sequence_replicate(value, device_mesh, seq_dim)
156+ if isinstance(replicated, DTensor):
157+ replicated = DTensor.from_local_with_layout(
158+ replicated.to_local().detach(), replicated.layout
159+ )
160+ else:
161+ replicated = replicated.detach()
162+ self._values[slot_name] = replicated
163+ self._signatures[slot_name] = signature
164+ 
165+ local_value = value.to_local() if isinstance(value, DTensor) else value
166+ replicated_value = replicated.to_local() if isinstance(replicated, DTensor) else replicated
167+ bridged_value = _DSASequenceReplicateGradientBridge.apply(
168+ local_value,
169+ replicated_value,
170+ device_mesh.get_group(),
171+ device_mesh.size(),
172+ seq_dim,
173+ )
174+ if isinstance(replicated, DTensor):
175+ return DTensor.from_local_with_layout(bridged_value, replicated.layout)
176+ return bridged_value
177+ 
178+ def clear(self) -> None:
179+ """Release cache-owned references after all forward consumers are wired."""
180+ self._values.clear()
181+ self._signatures.clear()
182+ 
183+ 
56def _is_tensor_or_dtensor(value: Any) -> bool:184def _is_tensor_or_dtensor(value: Any) -> bool:
57 """Return True for framework tensors and HyperParallel DTensors."""185 """Return True for framework tensors and HyperParallel DTensors."""
58 return isinstance(value, DTensor) or platform.is_tensor(value)186 return isinstance(value, DTensor) or platform.is_tensor(value)
@@ -225,6 +353,8 @@ def _configure_sparse_attention_boundary( # pylint: disable=too-many-arguments
225 query_rope_kwarg_name: Optional[str],353 query_rope_kwarg_name: Optional[str],
226 key_rope_kwarg_name: Optional[str],354 key_rope_kwarg_name: Optional[str],
227 use_local_output: bool,355 use_local_output: bool,
356+ shared_replicate_cache: Optional[DSASequenceReplicateCache],
357+ share_key_value: bool,
228) -> None:358) -> None:
229 """Store sparse-attention boundary configuration on ``style``."""359 """Store sparse-attention boundary configuration on ``style``."""
230 layout, seq_dim = _validate_layout_and_mode(style.__class__.__name__, layout, mode)360 layout, seq_dim = _validate_layout_and_mode(style.__class__.__name__, layout, mode)
@@ -244,6 +374,8 @@ def _configure_sparse_attention_boundary( # pylint: disable=too-many-arguments
244 style.query_rope_kwarg_name = query_rope_kwarg_name374 style.query_rope_kwarg_name = query_rope_kwarg_name
245 style.key_rope_kwarg_name = key_rope_kwarg_name375 style.key_rope_kwarg_name = key_rope_kwarg_name
246 style.use_local_output = use_local_output376 style.use_local_output = use_local_output
377+ style.shared_replicate_cache = shared_replicate_cache
378+ style.share_key_value = share_key_value
247 379 
248 380 
249def _apply_sparse_attention_boundary(381def _apply_sparse_attention_boundary(
@@ -262,12 +394,19 @@ def _apply_sparse_attention_boundary(
262 def _replicate(slot_name: str):394 def _replicate(slot_name: str):
263 if async_state is not None:395 if async_state is not None:
264 return lambda value: async_state.wait(slot_name, value)396 return lambda value: async_state.wait(slot_name, value)
397+ if style.shared_replicate_cache is not None:
398+ return lambda value: style.shared_replicate_cache.replicate(
399+ slot_name, value, cp_mesh, style.seq_dim
400+ )
265 return lambda value: _to_sequence_replicate(value, cp_mesh, style.seq_dim)401 return lambda value: _to_sequence_replicate(value, cp_mesh, style.seq_dim)
266 402 
403+ key_slot = "main_kv" if style.share_key_value else "key"
404+ value_slot = "main_kv" if style.share_key_value else "value"
405+ 
267 specs = [406 specs = [
268 _ParamSpec(style.query_index, style.query_kwarg_name, _shard),407 _ParamSpec(style.query_index, style.query_kwarg_name, _shard),
269- _ParamSpec(style.key_index, style.key_kwarg_name, _replicate("key")),408+ _ParamSpec(style.key_index, style.key_kwarg_name, _replicate(key_slot)),
270- _ParamSpec(style.value_index, style.value_kwarg_name, _replicate("value")),409+ _ParamSpec(style.value_index, style.value_kwarg_name, _replicate(value_slot)),
271 _ParamSpec(style.topk_index, style.topk_kwarg_name, _shard),410 _ParamSpec(style.topk_index, style.topk_kwarg_name, _shard),
272 _ParamSpec(style.query_rope_index, style.query_rope_kwarg_name, _shard),411 _ParamSpec(style.query_rope_index, style.query_rope_kwarg_name, _shard),
273 _ParamSpec(style.key_rope_index, style.key_rope_kwarg_name, _replicate("key_rope")),412 _ParamSpec(style.key_rope_index, style.key_rope_kwarg_name, _replicate("key_rope")),
@@ -276,6 +415,8 @@ def _apply_sparse_attention_boundary(
276 def _pre_hook(hook_module, args, kwargs):415 def _pre_hook(hook_module, args, kwargs):
277 new_args = list(args)416 new_args = list(args)
278 new_kwargs = dict(kwargs)417 new_kwargs = dict(kwargs)
418+ if style.shared_replicate_cache is not None:
419+ style.shared_replicate_cache.begin()
279 _record_query_output_layout(420 _record_query_output_layout(
280 hook_module,421 hook_module,
281 _read_value(new_args, new_kwargs, style.query_index, style.query_kwarg_name),422 _read_value(new_args, new_kwargs, style.query_index, style.query_kwarg_name),
@@ -405,6 +546,8 @@ class DSASparseAttentionContextParallel(ParallelStyle):
405 query_rope_kwarg_name: Optional[str] = "query_rope",546 query_rope_kwarg_name: Optional[str] = "query_rope",
406 key_rope_kwarg_name: Optional[str] = "key_rope",547 key_rope_kwarg_name: Optional[str] = "key_rope",
407 use_local_output: bool = False,548 use_local_output: bool = False,
549+ shared_replicate_cache: Optional[DSASequenceReplicateCache] = None,
550+ share_key_value: bool = False,
408 ) -> None:551 ) -> None:
409 super().__init__()552 super().__init__()
410 _configure_sparse_attention_boundary(553 _configure_sparse_attention_boundary(
@@ -424,6 +567,8 @@ class DSASparseAttentionContextParallel(ParallelStyle):
424 query_rope_kwarg_name=query_rope_kwarg_name,567 query_rope_kwarg_name=query_rope_kwarg_name,
425 key_rope_kwarg_name=key_rope_kwarg_name,568 key_rope_kwarg_name=key_rope_kwarg_name,
426 use_local_output=use_local_output,569 use_local_output=use_local_output,
570+ shared_replicate_cache=shared_replicate_cache,
571+ share_key_value=share_key_value,
427 )572 )
428 573 
429 def __repr__(self) -> str:574 def __repr__(self) -> str:
@@ -491,6 +636,7 @@ class DSAIndexerLossContextParallel(ParallelStyle):
491 query_rope_kwarg_name: Optional[str] = None,636 query_rope_kwarg_name: Optional[str] = None,
492 key_rope_kwarg_name: Optional[str] = None,637 key_rope_kwarg_name: Optional[str] = None,
493 use_local_output: bool = False,638 use_local_output: bool = False,
639+ shared_replicate_cache: Optional[DSASequenceReplicateCache] = None,
494 ) -> None:640 ) -> None:
495 super().__init__()641 super().__init__()
496 layout, seq_dim = _validate_layout_and_mode(self.__class__.__name__, layout, mode)642 layout, seq_dim = _validate_layout_and_mode(self.__class__.__name__, layout, mode)
@@ -526,6 +672,7 @@ class DSAIndexerLossContextParallel(ParallelStyle):
526 self.query_rope_kwarg_name = query_rope_kwarg_name672 self.query_rope_kwarg_name = query_rope_kwarg_name
527 self.key_rope_kwarg_name = key_rope_kwarg_name673 self.key_rope_kwarg_name = key_rope_kwarg_name
528 self.use_local_output = use_local_output674 self.use_local_output = use_local_output
675+ self.shared_replicate_cache = shared_replicate_cache
529 676 
530 def __repr__(self) -> str:677 def __repr__(self) -> str:
531 return (678 return (
@@ -633,6 +780,7 @@ class DSAIndexerLossContextParallel(ParallelStyle):
633 specs: list[_ParamSpec],780 specs: list[_ParamSpec],
634 local_idx: int,781 local_idx: int,
635 cp_mesh: DeviceMesh,782 cp_mesh: DeviceMesh,
783+ completion_fn: Optional[Callable[[], None]] = None,
636 ) -> Module:784 ) -> Module:
637 """Register indexer-loss hooks driven by parameter specs."""785 """Register indexer-loss hooks driven by parameter specs."""
638 def _pre_hook(hook_module, args, kwargs):786 def _pre_hook(hook_module, args, kwargs):
@@ -647,7 +795,11 @@ class DSAIndexerLossContextParallel(ParallelStyle):
647 key_shape = self._read_key_indexer_shape(new_args, new_kwargs)795 key_shape = self._read_key_indexer_shape(new_args, new_kwargs)
648 setattr(hook_module, "_hp_dsa_loss_key_index_local_shape", key_shape)796 setattr(hook_module, "_hp_dsa_loss_key_index_local_shape", key_shape)
649 setattr(hook_module, "_hp_dsa_loss_local_idx", local_idx)797 setattr(hook_module, "_hp_dsa_loss_local_idx", local_idx)
650- _apply_param_specs(new_args, new_kwargs, specs)798+ try:
799+ _apply_param_specs(new_args, new_kwargs, specs)
800+ finally:
801+ if completion_fn is not None:
802+ completion_fn()
651 return tuple(new_args), new_kwargs803 return tuple(new_args), new_kwargs
652 804 
653 platform.register_forward_pre_hook(module, _pre_hook, with_kwargs=True)805 platform.register_forward_pre_hook(module, _pre_hook, with_kwargs=True)
@@ -659,14 +811,28 @@ class DSAIndexerLossContextParallel(ParallelStyle):
659 cp_mesh = _ensure_1d(device_mesh)811 cp_mesh = _ensure_1d(device_mesh)
660 812 
661 def replicate(value: Any) -> Any:813 def replicate(value: Any) -> Any:
814+ """Replicate a key tensor without sharing its gradient path."""
662 return self._replicate_key_side(value, cp_mesh)815 return self._replicate_key_side(value, cp_mesh)
663 816 
817+ def shared_replicate(slot_name: str, value: Any) -> Any:
818+ """Reuse a gradient-safe main-attention activation when enabled."""
819+ if self.shared_replicate_cache is None:
820+ return replicate(value)
821+ return self.shared_replicate_cache.replicate(slot_name, value, cp_mesh, self.seq_dim)
822+ 
664 specs = self._build_loss_specs(823 specs = self._build_loss_specs(
665 cp_mesh,824 cp_mesh,
666 replicate_fn_map={825 replicate_fn_map={
667- "key": replicate,826+ "key": lambda value: shared_replicate("main_kv", value),
668 "key_indexer": replicate,827 "key_indexer": replicate,
669- "key_rope": replicate,828+ "key_rope": lambda value: shared_replicate("key_rope", value),
670 },829 },
671 )830 )
672- return self._apply_with_loss_specs(module, specs, self._get_local_idx(cp_mesh), cp_mesh)831+ completion_fn = self.shared_replicate_cache.clear if self.shared_replicate_cache is not None else None
832+ return self._apply_with_loss_specs(
833+ module,
834+ specs,
835+ self._get_local_idx(cp_mesh),
836+ cp_mesh,
837+ completion_fn=completion_fn,
838+ )
@@ -38,6 +38,10 @@ from hyper_parallel.core.context_parallel.context_parallel import (
38 _to_cp_dtensor,38 _to_cp_dtensor,
39)39)
40from hyper_parallel.core.context_parallel.async_dsa_context_parallel import _AsyncSequenceReplicateSlot40from hyper_parallel.core.context_parallel.async_dsa_context_parallel import _AsyncSequenceReplicateSlot
41+from hyper_parallel.core.context_parallel.dsa_context_parallel import (
42+ DSASequenceReplicateCache,
43+ _to_sequence_replicate,
44+)
41from hyper_parallel.core.dtensor.device_mesh import init_device_mesh, _DEVICE_MESH_MAP45from hyper_parallel.core.dtensor.device_mesh import init_device_mesh, _DEVICE_MESH_MAP
42from hyper_parallel.core.dtensor.dtensor import DTensor46from hyper_parallel.core.dtensor.dtensor import DTensor
43from hyper_parallel.core.dtensor.placement_types import Replicate, Shard, StridedShard47from hyper_parallel.core.dtensor.placement_types import Replicate, Shard, StridedShard
@@ -353,6 +357,88 @@ class TestDsaContextParallel(unittest.TestCase):
353 self.assertEqual(out[4].placements, (Shard(1),))357 self.assertEqual(out[4].placements, (Shard(1),))
354 self.assertEqual(out[5].placements, (Replicate(),))358 self.assertEqual(out[5].placements, (Replicate(),))
355 359 
360+ @patch("hyper_parallel.core.dtensor.device_mesh.platform")
361+ def test_sparse_attention_and_loss_reuse_main_kv_gathers(self, mock_mesh_platform):
362+ """Sparse DSA should share only gradient-safe global K/V activations."""
363+ mesh = self._make_cp_mesh(mock_mesh_platform)
364+ shared_cache = DSASequenceReplicateCache()
365+ attention_style = DSASparseAttentionContextParallel(
366+ layout="BSND",
367+ use_local_output=False,
368+ shared_replicate_cache=shared_cache,
369+ share_key_value=True,
370+ )
371+ loss_style = DSAIndexerLossContextParallel(
372+ layout="BSND",
373+ use_local_output=False,
374+ shared_replicate_cache=shared_cache,
375+ )
376+ attention_module = _IdentityModule()
377+ loss_module = _IdentityModule()
378+ attention_style.apply(attention_module, mesh)
379+ 
380+ query = torch.randn(2, 4, 8, 16)
381+ latent = torch.randn(2, 4, 1, 16)
382+ topk = torch.randint(0, 4, (2, 4, 1, 2), dtype=torch.int32)
383+ query_rope = torch.randn(2, 4, 8, 8)
384+ key_rope = torch.randn(2, 4, 1, 8)
385+ query_index = torch.randn(2, 4, 8, 16)
386+ key_index = torch.randn(2, 4, 1, 16)
387+ weights = torch.randn(2, 4, 8)
388+ softmax_max = torch.randn(2, 4, 8, 1)
389+ softmax_sum = torch.randn(2, 4, 8, 1)
390+ 
391+ with _patch_torch_dist_rank(), patch.object(mesh, "get_group", return_value=None), patch(
392+ "hyper_parallel.core.context_parallel.dsa_context_parallel._to_sequence_replicate",
393+ wraps=_to_sequence_replicate,
394+ ) as mock_replicate:
395+ loss_style.apply(loss_module, mesh)
396+ attention_out = attention_module(query, latent, latent.view_as(latent), topk, query_rope, key_rope)
397+ loss_out = loss_module(
398+ query,
399+ latent,
400+ query_index,
401+ key_index,
402+ weights,
403+ topk,
404+ softmax_max,
405+ softmax_sum,
406+ query_rope,
407+ key_rope,
408+ )
409+ 
410+ self.assertEqual(attention_out[1].to_local().data_ptr(), attention_out[2].to_local().data_ptr())
411+ self.assertEqual(loss_out[1].to_local().data_ptr(), attention_out[1].to_local().data_ptr())
412+ self.assertEqual(loss_out[9].to_local().data_ptr(), attention_out[5].to_local().data_ptr())
413+ self.assertNotEqual(loss_out[3].to_local().data_ptr(), attention_out[1].to_local().data_ptr())
414+ self.assertEqual(mock_replicate.call_count, 3)
415+ 
416+ @patch("hyper_parallel.core.dtensor.device_mesh.platform")
417+ def test_shared_main_kv_preserves_parent_gradient(self, mock_mesh_platform):
418+ """Aliased K/V gather should preserve the gradient of their common latent."""
419+ mesh = self._make_cp_mesh(mock_mesh_platform)
420+ shared_cache = DSASequenceReplicateCache()
421+ style = DSASparseAttentionContextParallel(
422+ layout="BSND",
423+ use_local_output=False,
424+ shared_replicate_cache=shared_cache,
425+ share_key_value=True,
426+ )
427+ module = _IdentityModule()
428+ style.apply(module, mesh)
429+ 
430+ latent = torch.randn(2, 4, 1, 16, requires_grad=True)
431+ query = torch.randn(2, 4, 8, 16)
432+ topk = torch.randint(0, 4, (2, 4, 1, 2), dtype=torch.int32)
433+ query_rope = torch.randn(2, 4, 8, 8)
434+ key_rope = torch.randn(2, 4, 1, 8)
435+ 
436+ with _patch_torch_dist_rank(), patch.object(mesh, "get_group", return_value=None):
437+ out = module(query, latent * 1.0, latent.view_as(latent), topk, query_rope, key_rope)
438+ (out[1].to_local().sum() + out[2].to_local().sum()).backward()
439+ 
440+ self.assertTrue(torch.equal(latent.grad, torch.full_like(latent, 2.0)))
441+ 
356 @patch("hyper_parallel.core.dtensor.device_mesh.platform")442 @patch("hyper_parallel.core.dtensor.device_mesh.platform")
357 def test_sparse_attention_boundary_kwargs_are_rewritten(self, mock_mesh_platform):443 def test_sparse_attention_boundary_kwargs_are_rewritten(self, mock_mesh_platform):
358 """Sparse FA boundary should also rewrite configured keyword arguments."""444 """Sparse FA boundary should also rewrite configured keyword arguments."""