已合并
[master] 优化 DSA 长序列上下文并行显存 #1282
lzy0920232创建于 8月29日
[master] 优化 DSA 长序列上下文并行显存 #1282
已合并
共 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 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 | + | ||
| 56 | def _is_tensor_or_dtensor(value: Any) -> bool: | 184 | def _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_name | 374 | style.query_rope_kwarg_name = query_rope_kwarg_name |
| 245 | style.key_rope_kwarg_name = key_rope_kwarg_name | 375 | style.key_rope_kwarg_name = key_rope_kwarg_name |
| 246 | style.use_local_output = use_local_output | 376 | 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 | ||
| 249 | def _apply_sparse_attention_boundary( | 381 | def _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_name | 672 | self.query_rope_kwarg_name = query_rope_kwarg_name |
| 527 | self.key_rope_kwarg_name = key_rope_kwarg_name | 673 | self.key_rope_kwarg_name = key_rope_kwarg_name |
| 528 | self.use_local_output = use_local_output | 674 | 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_kwargs | 803 | 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 | ) |
| 40 | from hyper_parallel.core.context_parallel.async_dsa_context_parallel import _AsyncSequenceReplicateSlot | 40 | from 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 | +) | ||
| 41 | from hyper_parallel.core.dtensor.device_mesh import init_device_mesh, _DEVICE_MESH_MAP | 45 | from hyper_parallel.core.dtensor.device_mesh import init_device_mesh, _DEVICE_MESH_MAP |
| 42 | from hyper_parallel.core.dtensor.dtensor import DTensor | 46 | from hyper_parallel.core.dtensor.dtensor import DTensor |
| 43 | from hyper_parallel.core.dtensor.placement_types import Replicate, Shard, StridedShard | 47 | from 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 | + | ||
| 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 | + | ||
| 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 | 442 | ||
| 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.""" |