已合并
add fsdp test case #22803
zqwen创建于 2025年7月9日
add fsdp test case #22803
已合并
从refs/pull/22803/head合入到v2.7.1
共 1 个文件变更+1127-0
| @@ -0,0 +1,1127 @@ | |||
| 1 | +import copy | ||
| 2 | +import functools | ||
| 3 | +import itertools | ||
| 4 | +import unittest | ||
| 5 | +from typing import Callable, List, Optional, Tuple, Union | ||
| 6 | + | ||
| 7 | +import torch | ||
| 8 | +import torch.distributed as dist | ||
| 9 | +import torch.nn as nn | ||
| 10 | +import torch.nn.functional as F | ||
| 11 | +from torch.distributed._composable import checkpoint, replicate | ||
| 12 | +from torch.distributed.device_mesh import DeviceMesh, init_device_mesh | ||
| 13 | +from torch.distributed.fsdp import ( | ||
| 14 | + FSDPModule, | ||
| 15 | + fully_shard, | ||
| 16 | + MixedPrecisionPolicy, | ||
| 17 | + OffloadPolicy, | ||
| 18 | +) | ||
| 19 | +from torch.distributed.fsdp._fully_shard._fsdp_collectives import ( | ||
| 20 | + _div_if_needed, | ||
| 21 | + _get_gradient_divide_factors, | ||
| 22 | + foreach_all_gather, | ||
| 23 | + foreach_all_gather_copy_out, | ||
| 24 | + foreach_reduce, | ||
| 25 | +) | ||
| 26 | +from torch.distributed.fsdp._fully_shard._fsdp_common import FSDPMeshInfo, TrainingState | ||
| 27 | +from torch.distributed.fsdp._fully_shard._fsdp_init import ( | ||
| 28 | + _get_post_forward_mesh_info, | ||
| 29 | + _init_default_fully_shard_mesh, | ||
| 30 | +) | ||
| 31 | +from torch.distributed.fsdp._fully_shard._fsdp_param import ShardedState | ||
| 32 | +from torch.distributed.fsdp._fully_shard._fsdp_param_group import FSDPParamGroup | ||
| 33 | +from torch.distributed.tensor import DTensor | ||
| 34 | +from torch.distributed.tensor.debug import CommDebugMode | ||
| 35 | +from torch.distributed.tensor.experimental import implicit_replication | ||
| 36 | +from torch.testing._internal.common_fsdp import ( | ||
| 37 | + check_sharded_parity, | ||
| 38 | + DoubleLinear, | ||
| 39 | + FSDPTest, | ||
| 40 | + FSDPTestMultiThread, | ||
| 41 | + MLP, | ||
| 42 | + patch_post_backward, | ||
| 43 | + patch_reshard, | ||
| 44 | + patch_unshard, | ||
| 45 | +) | ||
| 46 | +from torch.testing._internal.common_utils import run_tests | ||
| 47 | +from torch.testing._internal.distributed._tensor.common_dtensor import ( | ||
| 48 | + ModelArgs, | ||
| 49 | + Transformer, | ||
| 50 | + TransformerBlock, | ||
| 51 | +) | ||
| 52 | + | ||
| 53 | +import torch_npu | ||
| 54 | +from torch_npu.testing.common_utils import SupportedDevices | ||
| 55 | +from torch_npu.testing._internal.common_fsdp import FSDPNPUTest | ||
| 56 | + | ||
| 57 | + | ||
| 58 | +c10d_ops = torch.ops.c10d | ||
| 59 | + | ||
| 60 | +# For recording FSDP events like unshard or post-backward | ||
| 61 | +EventType = Tuple[str, str, TrainingState] | ||
| 62 | + | ||
| 63 | + | ||
| 64 | +class TestFullyShardCollectiveOps(FSDPTestMultiThread): | ||
| 65 | + | ||
| 66 | + def world_size(self) -> int: | ||
| 67 | + return 128 | ||
| 68 | + | ||
| 69 | + def perThreadSetUp(self): | ||
| 70 | + super().perThreadSetUp() | ||
| 71 | + torch.npu.set_device(0) | ||
| 72 | + | ||
| 73 | + | ||
| 74 | + def device(self) -> torch.device: | ||
| 75 | + return torch.device("npu:0") | ||
| 76 | + | ||
| 77 | + def _get_param_sizes(self) -> List[torch.Size]: | ||
| 78 | + # For world size 128, the fp32 all-gather and reduce-scatter testing | ||
| 79 | + # requires ~0.22 GB | ||
| 80 | + return [ | ||
| 81 | + torch.Size([17, 257]), | ||
| 82 | + torch.Size([17]), | ||
| 83 | + torch.Size([64, 312]), | ||
| 84 | + torch.Size([64]), | ||
| 85 | + torch.Size([64, 64]), | ||
| 86 | + torch.Size([512, 64]), | ||
| 87 | + torch.Size([256]), | ||
| 88 | + torch.Size([64, 297]), | ||
| 89 | + ] | ||
| 90 | + | ||
| 91 | + def _init_params(self, param_sizes: List[torch.Size]) -> List[nn.Parameter]: | ||
| 92 | + torch.manual_seed(42) | ||
| 93 | + orig_params = [nn.Parameter(torch.randn(size, device=self.device)) for size in param_sizes] | ||
| 94 | + # Since seed is per process, not per thread, we broadcast to ensure the | ||
| 95 | + # same original parameters across ranks | ||
| 96 | + for orig_param in orig_params: | ||
| 97 | + dist.broadcast(orig_param, src=0) | ||
| 98 | + return orig_params | ||
| 99 | + | ||
| 100 | + def _init_fsdp_param_group( | ||
| 101 | + self, params: List[nn.Parameter], reshard_after_forward: Union[bool, int] | ||
| 102 | + ): | ||
| 103 | + module = nn.ParameterList([param.detach().clone() for param in params]) | ||
| 104 | + mesh_info = FSDPMeshInfo(_init_default_fully_shard_mesh(), shard_mesh_dim=0) | ||
| 105 | + post_forward_mesh_info = _get_post_forward_mesh_info( | ||
| 106 | + reshard_after_forward, mesh_info | ||
| 107 | + ) | ||
| 108 | + fsdp_param_group = FSDPParamGroup( | ||
| 109 | + list(module.parameters()), | ||
| 110 | + (module,), | ||
| 111 | + mesh_info, | ||
| 112 | + post_forward_mesh_info, | ||
| 113 | + self.device, | ||
| 114 | + None, # shard_placement_fn | ||
| 115 | + MixedPrecisionPolicy(), | ||
| 116 | + OffloadPolicy(), | ||
| 117 | + ) | ||
| 118 | + fsdp_param_group.lazy_init() | ||
| 119 | + return fsdp_param_group | ||
| 120 | + | ||
| 121 | + | ||
| 122 | + def test_all_gather_fp32(self): | ||
| 123 | + param_sizes = self._get_param_sizes() | ||
| 124 | + default_stream = torch.npu.current_stream() | ||
| 125 | + stream1, stream2 = torch.npu.Stream(), torch.npu.Stream() | ||
| 126 | + for async_op, streams, reshard_after_forward in itertools.product( | ||
| 127 | + (False, True), | ||
| 128 | + ((default_stream, default_stream), (stream1, stream2)), | ||
| 129 | + (True, 8), | ||
| 130 | + ): | ||
| 131 | + all_gather_copy_in_stream, all_gather_stream = streams | ||
| 132 | + # Save test time by only testing reshard after forward as an int | ||
| 133 | + # for non-async and non-default streams (like in pre-backward) | ||
| 134 | + if type(reshard_after_forward) is int and ( | ||
| 135 | + async_op or all_gather_stream is default_stream | ||
| 136 | + ): | ||
| 137 | + continue | ||
| 138 | + self._test_all_gather( | ||
| 139 | + param_sizes, | ||
| 140 | + reshard_after_forward=reshard_after_forward, | ||
| 141 | + async_op=async_op, | ||
| 142 | + all_gather_copy_in_stream=all_gather_copy_in_stream, | ||
| 143 | + all_gather_stream=all_gather_stream, | ||
| 144 | + ) | ||
| 145 | + | ||
| 146 | + def _test_all_gather( | ||
| 147 | + self, | ||
| 148 | + param_sizes: List[torch.Size], | ||
| 149 | + reshard_after_forward: Union[bool, int], | ||
| 150 | + async_op: bool, | ||
| 151 | + all_gather_copy_in_stream: torch.npu.Stream, | ||
| 152 | + all_gather_stream: torch.npu.Stream, | ||
| 153 | + ): | ||
| 154 | + def all_gather(fsdp_param_group: FSDPParamGroup, group: dist.ProcessGroup): | ||
| 155 | + all_gather_result = foreach_all_gather( | ||
| 156 | + fsdp_param_group.fsdp_params, | ||
| 157 | + group, | ||
| 158 | + async_op=async_op, | ||
| 159 | + all_gather_copy_in_stream=all_gather_copy_in_stream, | ||
| 160 | + all_gather_stream=all_gather_stream, | ||
| 161 | + device=self.device, | ||
| 162 | + ) | ||
| 163 | + foreach_all_gather_copy_out(all_gather_result, fsdp_params, group) | ||
| 164 | + # Transition to unsharded state to register unsharded parameters | ||
| 165 | + for fsdp_param in fsdp_param_group.fsdp_params: | ||
| 166 | + fsdp_param.init_unsharded_param() | ||
| 167 | + fsdp_param_group._to_unsharded() | ||
| 168 | + | ||
| 169 | + def check_all_gathered_params( | ||
| 170 | + orig_params: List[nn.Parameter], module: nn.Module | ||
| 171 | + ): | ||
| 172 | + for orig_param, param in zip(orig_params, module.parameters()): | ||
| 173 | + self.assertIsInstance(param, torch.Tensor) | ||
| 174 | + self.assertIsInstance(param, nn.Parameter) | ||
| 175 | + self.assertEqual(param, orig_param.to(param.dtype)) | ||
| 176 | + | ||
| 177 | + # Set up the reference parameters and construct the FSDP group | ||
| 178 | + orig_params = self._init_params(param_sizes) | ||
| 179 | + fsdp_param_group = self._init_fsdp_param_group( | ||
| 180 | + orig_params, reshard_after_forward | ||
| 181 | + ) | ||
| 182 | + fsdp_params = fsdp_param_group.fsdp_params | ||
| 183 | + module = fsdp_param_group.modules[0] | ||
| 184 | + | ||
| 185 | + # Sanity check that the parameter sharding is as expected | ||
| 186 | + for orig_param, param in zip(orig_params, module.parameters()): | ||
| 187 | + self.assertTrue(isinstance(param, DTensor)) | ||
| 188 | + self.assertEqual(param.full_tensor(), orig_param) | ||
| 189 | + | ||
| 190 | + # Run the foreach all-gather (including copy-in and copy-out) | ||
| 191 | + all_gather(fsdp_param_group, fsdp_param_group.mesh_info.shard_process_group) | ||
| 192 | + | ||
| 193 | + # Check all-gather correctness | ||
| 194 | + check_all_gathered_params(orig_params, module) | ||
| 195 | + | ||
| 196 | + # For reshard after after forward as an int, further test emulating the | ||
| 197 | + # pre-backward all-gather | ||
| 198 | + if type(reshard_after_forward) is not int: | ||
| 199 | + return | ||
| 200 | + fsdp_param_group._to_sharded_post_forward() | ||
| 201 | + all_gather( | ||
| 202 | + fsdp_param_group, | ||
| 203 | + fsdp_param_group.post_forward_mesh_info.shard_process_group, | ||
| 204 | + ) | ||
| 205 | + check_all_gathered_params(orig_params, module) | ||
| 206 | + | ||
| 207 | + | ||
| 208 | + def test_reduce_scatter_fp32(self): | ||
| 209 | + param_sizes = self._get_param_sizes() | ||
| 210 | + default_stream = torch.npu.current_stream() | ||
| 211 | + stream = torch.npu.Stream() | ||
| 212 | + for reduce_scatter_stream in (default_stream, stream): | ||
| 213 | + self._test_reduce_scatter( | ||
| 214 | + param_sizes, | ||
| 215 | + reduce_scatter_stream=reduce_scatter_stream, | ||
| 216 | + reduce_scatter_dtype=torch.float32, | ||
| 217 | + ) | ||
| 218 | + | ||
| 219 | + | ||
| 220 | + def test_reduce_scatter_fp16(self): | ||
| 221 | + param_sizes = self._get_param_sizes() | ||
| 222 | + default_stream = torch.npu.current_stream() | ||
| 223 | + stream = torch.npu.Stream() | ||
| 224 | + for reduce_scatter_stream in (default_stream, stream): | ||
| 225 | + self._test_reduce_scatter( | ||
| 226 | + param_sizes, | ||
| 227 | + reduce_scatter_stream=reduce_scatter_stream, | ||
| 228 | + reduce_scatter_dtype=torch.float16, | ||
| 229 | + ) | ||
| 230 | + | ||
| 231 | + def _test_reduce_scatter( | ||
| 232 | + self, | ||
| 233 | + param_sizes: List[torch.Size], | ||
| 234 | + reduce_scatter_stream: torch.npu.Stream, | ||
| 235 | + reduce_scatter_dtype: torch.dtype, | ||
| 236 | + ): | ||
| 237 | + # Set up the reference parameters and construct the FSDP group | ||
| 238 | + orig_params = self._init_params(param_sizes) | ||
| 239 | + fsdp_param_group = self._init_fsdp_param_group(orig_params, True) | ||
| 240 | + fsdp_params = fsdp_param_group.fsdp_params | ||
| 241 | + fsdp_param_group.comm_ctx.lazy_init(self.device) | ||
| 242 | + | ||
| 243 | + # Run one unshard to initialize metadata | ||
| 244 | + fsdp_param_group.unshard() | ||
| 245 | + fsdp_param_group.wait_for_unshard() | ||
| 246 | + fsdp_param_group.reshard() | ||
| 247 | + | ||
| 248 | + # Run the foreach reduce-scatter (including copy-in and view-out) | ||
| 249 | + torch.manual_seed(42) | ||
| 250 | + unsharded_grads = [torch.ones_like(param) * self.rank for param in orig_params] | ||
| 251 | + group = fsdp_param_group.mesh_info.shard_process_group | ||
| 252 | + self.assertEqual(group.size(), self.world_size) | ||
| 253 | + all_reduce_stream = torch.npu.Stream() | ||
| 254 | + ( | ||
| 255 | + reduce_scatter_input, | ||
| 256 | + reduce_scatter_event, | ||
| 257 | + post_reduce_event, | ||
| 258 | + _, | ||
| 259 | + _, | ||
| 260 | + _, | ||
| 261 | + ) = foreach_reduce( | ||
| 262 | + fsdp_params, | ||
| 263 | + unsharded_grads, | ||
| 264 | + group, | ||
| 265 | + reduce_scatter_stream, | ||
| 266 | + orig_dtype=orig_params[0].dtype, | ||
| 267 | + reduce_dtype=reduce_scatter_dtype, | ||
| 268 | + device=self.device, | ||
| 269 | + reduce_scatter_reduce_op=None, | ||
| 270 | + all_reduce_group=None, | ||
| 271 | + all_reduce_stream=all_reduce_stream, | ||
| 272 | + all_reduce_grads=True, | ||
| 273 | + partial_reduce_output=None, | ||
| 274 | + ) | ||
| 275 | + torch.npu.current_stream().wait_event(post_reduce_event) | ||
| 276 | + | ||
| 277 | + # Check reduce-scatter correctness | ||
| 278 | + predivide_factor, postdivide_factor = _get_gradient_divide_factors( | ||
| 279 | + group, None, reduce_scatter_dtype | ||
| 280 | + ) | ||
| 281 | + reduced_grads = [grad.detach().clone() for grad in unsharded_grads] | ||
| 282 | + for grad in reduced_grads: | ||
| 283 | + _div_if_needed(grad, predivide_factor) | ||
| 284 | + dist.all_reduce( | ||
| 285 | + grad, | ||
| 286 | + group=group, | ||
| 287 | + op=dist.ReduceOp.AVG if predivide_factor is None else dist.ReduceOp.SUM, | ||
| 288 | + ) | ||
| 289 | + _div_if_needed(grad, postdivide_factor) | ||
| 290 | + for fsdp_param, reduced_grad in zip(fsdp_params, reduced_grads): | ||
| 291 | + sharded_grad = fsdp_param.sharded_param.grad | ||
| 292 | + self.assertIsInstance(sharded_grad, DTensor) | ||
| 293 | + self.assertEqual(sharded_grad.full_tensor(), reduced_grad) | ||
| 294 | + | ||
| 295 | + | ||
| 296 | +class TestFullyShardCommunication(FSDPNPUTest): | ||
| 297 | + | ||
| 298 | + def world_size(self) -> int: | ||
| 299 | + return min(4, torch.npu.device_count()) | ||
| 300 | + | ||
| 301 | + | ||
| 302 | + def test_fully_shard_communication_count(self): | ||
| 303 | + """ | ||
| 304 | + Tests that FSDP issues the expected number of all-gathers and | ||
| 305 | + reduce-scatters during forward and backward. | ||
| 306 | + """ | ||
| 307 | + self.run_subtests( | ||
| 308 | + {"reshard_after_forward": [True, False, 2]}, | ||
| 309 | + self._test_communication_count, | ||
| 310 | + ) | ||
| 311 | + | ||
| 312 | + def _test_communication_count( | ||
| 313 | + self, | ||
| 314 | + reshard_after_forward: Union[bool, int], | ||
| 315 | + ): | ||
| 316 | + torch.manual_seed(42) | ||
| 317 | + model_args = ModelArgs() | ||
| 318 | + model = Transformer(model_args) | ||
| 319 | + fully_shard_fn = functools.partial( | ||
| 320 | + fully_shard, reshard_after_forward=reshard_after_forward | ||
| 321 | + ) | ||
| 322 | + num_blocks = 0 | ||
| 323 | + for module in model.modules(): | ||
| 324 | + if isinstance(module, TransformerBlock): | ||
| 325 | + fully_shard_fn(module) | ||
| 326 | + num_blocks += 1 | ||
| 327 | + fully_shard_fn(model) | ||
| 328 | + # We construct `num_blocks` plus 1 FSDP states/communication groups | ||
| 329 | + | ||
| 330 | + torch.manual_seed(42 + self.rank) | ||
| 331 | + inp = torch.randint(0, model_args.vocab_size, (2, 16), device="npu") | ||
| 332 | + with CommDebugMode() as fwd_comm_mode: | ||
| 333 | + loss = model(inp) | ||
| 334 | + fwd_comm_counts = fwd_comm_mode.get_comm_counts() | ||
| 335 | + self.assertEqual(len(fwd_comm_counts), 1) | ||
| 336 | + self.assertEqual(fwd_comm_counts[c10d_ops._allgather_base_], num_blocks + 1) | ||
| 337 | + with CommDebugMode() as bwd_comm_mode: | ||
| 338 | + loss.sum().backward() | ||
| 339 | + bwd_comm_counts = bwd_comm_mode.get_comm_counts() | ||
| 340 | + if reshard_after_forward is False: | ||
| 341 | + self.assertEqual(len(bwd_comm_counts), 1) | ||
| 342 | + else: | ||
| 343 | + # The root always does not reshard after forward | ||
| 344 | + self.assertEqual(len(bwd_comm_counts), 2) | ||
| 345 | + self.assertEqual(bwd_comm_counts[c10d_ops._allgather_base_], num_blocks) | ||
| 346 | + self.assertEqual( | ||
| 347 | + bwd_comm_counts[c10d_ops._reduce_scatter_base_], num_blocks + 1 | ||
| 348 | + ) | ||
| 349 | + | ||
| 350 | + | ||
| 351 | + def test_manual_reshard_with_reshard_after_forward_false(self): | ||
| 352 | + """ | ||
| 353 | + Tests that we can manually call ``reshard`` on FSDP modules that were | ||
| 354 | + initialized with ``reshard_after_forward=False`` and still run unshard. | ||
| 355 | + """ | ||
| 356 | + torch.manual_seed(42) | ||
| 357 | + model_args = ModelArgs() | ||
| 358 | + model = Transformer(model_args) | ||
| 359 | + for module in model.modules(): | ||
| 360 | + if isinstance(module, TransformerBlock): | ||
| 361 | + fully_shard(module, reshard_after_forward=False) | ||
| 362 | + model = fully_shard(model, reshard_after_forward=False) | ||
| 363 | + num_fsdp_modules = sum( | ||
| 364 | + isinstance(module, FSDPModule) for module in model.modules() | ||
| 365 | + ) | ||
| 366 | + | ||
| 367 | + torch.manual_seed(42 + self.rank) | ||
| 368 | + inp = torch.randint(0, model_args.vocab_size, (2, 16), device="npu") | ||
| 369 | + with CommDebugMode() as fwd_comm_mode: | ||
| 370 | + loss = model(inp) | ||
| 371 | + fwd_comm_counts = fwd_comm_mode.get_comm_counts() | ||
| 372 | + self.assertEqual(len(fwd_comm_counts), 1) | ||
| 373 | + self.assertEqual(fwd_comm_counts[c10d_ops._allgather_base_], num_fsdp_modules) | ||
| 374 | + | ||
| 375 | + for module in model.modules(): | ||
| 376 | + if isinstance(module, FSDPModule): | ||
| 377 | + module.reshard() | ||
| 378 | + | ||
| 379 | + with CommDebugMode() as bwd_comm_mode: | ||
| 380 | + loss.sum().backward() | ||
| 381 | + bwd_comm_counts = bwd_comm_mode.get_comm_counts() | ||
| 382 | + self.assertEqual(len(bwd_comm_counts), 2) | ||
| 383 | + self.assertEqual(bwd_comm_counts[c10d_ops._allgather_base_], num_fsdp_modules) | ||
| 384 | + self.assertEqual( | ||
| 385 | + bwd_comm_counts[c10d_ops._reduce_scatter_base_], num_fsdp_modules | ||
| 386 | + ) | ||
| 387 | + | ||
| 388 | + | ||
| 389 | +class TestFullyShardPrefetch(FSDPNPUTest): | ||
| 390 | + | ||
| 391 | + def world_size(self) -> int: | ||
| 392 | + return min(4, torch.npu.device_count()) | ||
| 393 | + | ||
| 394 | + | ||
| 395 | + def test_fully_shard_backward_prefetch(self): | ||
| 396 | + # Activation checkpointing should not affect the expected FSDP events | ||
| 397 | + self.run_subtests( | ||
| 398 | + { | ||
| 399 | + "reshard_after_forward": [True, False, 2], | ||
| 400 | + "checkpoint_impl": [None, "utils", "composable"], | ||
| 401 | + }, | ||
| 402 | + self._test_backward_prefetch_forward_backward, | ||
| 403 | + ) | ||
| 404 | + self.run_subtests( | ||
| 405 | + { | ||
| 406 | + "reshard_after_forward": [True, False, 2], | ||
| 407 | + "checkpoint_impl": [None, "utils", "composable"], | ||
| 408 | + }, | ||
| 409 | + self._test_backward_prefetch_multi_forward, | ||
| 410 | + ) | ||
| 411 | + self._test_backward_prefetch_unused_in_backward(True) | ||
| 412 | + | ||
| 413 | + def _test_backward_prefetch_forward_backward( | ||
| 414 | + self, reshard_after_forward: Union[bool, int], checkpoint_impl: Optional[str] | ||
| 415 | + ): | ||
| 416 | + n_layers = 3 | ||
| 417 | + model, optim, inp = self._init_transformer( | ||
| 418 | + n_layers, reshard_after_forward, checkpoint_impl | ||
| 419 | + ) | ||
| 420 | + events: List[EventType] = [] | ||
| 421 | + unshard_with_record = self._get_unshard_with_record( | ||
| 422 | + FSDPParamGroup.unshard, events | ||
| 423 | + ) | ||
| 424 | + post_backward_with_record = self._get_post_backward_with_record( | ||
| 425 | + FSDPParamGroup.post_backward, events | ||
| 426 | + ) | ||
| 427 | + # Check the order for normal 1 forward, 1 backward, 1 optimizer step | ||
| 428 | + with patch_unshard(unshard_with_record), patch_post_backward( | ||
| 429 | + post_backward_with_record | ||
| 430 | + ): | ||
| 431 | + for iter_idx in range(3): | ||
| 432 | + loss = model(inp) | ||
| 433 | + expected_events = [ | ||
| 434 | + ("unshard", "", TrainingState.FORWARD), # root | ||
| 435 | + ("unshard", "layers.0", TrainingState.FORWARD), | ||
| 436 | + ("unshard", "layers.1", TrainingState.FORWARD), | ||
| 437 | + ("unshard", "layers.2", TrainingState.FORWARD), | ||
| 438 | + ] | ||
| 439 | + self.assertEqual(events, expected_events) | ||
| 440 | + events.clear() | ||
| 441 | + loss.sum().backward() | ||
| 442 | + expected_events = [ | ||
| 443 | + # Root does not reshard after forward so there is no | ||
| 444 | + # unshard event for it in backward | ||
| 445 | + ("unshard", "layers.2", TrainingState.PRE_BACKWARD), | ||
| 446 | + # Explicit backward prefetching moves the unshards early | ||
| 447 | + # by one module (note how swapping each unshard down one | ||
| 448 | + # event would give the natural event order) | ||
| 449 | + ("unshard", "layers.1", TrainingState.PRE_BACKWARD), | ||
| 450 | + ("post_backward", "layers.2", TrainingState.POST_BACKWARD), | ||
| 451 | + ("unshard", "layers.0", TrainingState.PRE_BACKWARD), | ||
| 452 | + ("post_backward", "layers.1", TrainingState.POST_BACKWARD), | ||
| 453 | + ("post_backward", "layers.0", TrainingState.POST_BACKWARD), | ||
| 454 | + ("post_backward", "", TrainingState.POST_BACKWARD), | ||
| 455 | + ] | ||
| 456 | + if reshard_after_forward is False: | ||
| 457 | + # No reshard after forward means no backward unshards | ||
| 458 | + expected_events = [e for e in expected_events if e[0] != "unshard"] | ||
| 459 | + self.assertEqual(events, expected_events) | ||
| 460 | + events.clear() | ||
| 461 | + optim.step() | ||
| 462 | + optim.zero_grad(set_to_none=(iter_idx % 2 == 0)) | ||
| 463 | + | ||
| 464 | + def _test_backward_prefetch_multi_forward( | ||
| 465 | + self, reshard_after_forward: Union[bool, int], checkpoint_impl: Optional[str] | ||
| 466 | + ): | ||
| 467 | + n_layers = 3 | ||
| 468 | + model, optim, inp = self._init_transformer( | ||
| 469 | + n_layers, reshard_after_forward, checkpoint_impl | ||
| 470 | + ) | ||
| 471 | + events: List[EventType] = [] | ||
| 472 | + unshard_with_record = self._get_unshard_with_record( | ||
| 473 | + FSDPParamGroup.unshard, events | ||
| 474 | + ) | ||
| 475 | + post_backward_with_record = self._get_post_backward_with_record( | ||
| 476 | + FSDPParamGroup.post_backward, events | ||
| 477 | + ) | ||
| 478 | + # Check the order for multiple forwards before 1 backward | ||
| 479 | + with patch_unshard(unshard_with_record), patch_post_backward( | ||
| 480 | + post_backward_with_record | ||
| 481 | + ): | ||
| 482 | + loss1 = model(inp) | ||
| 483 | + loss2 = model(inp) | ||
| 484 | + expected_events = [ | ||
| 485 | + ("unshard", "", TrainingState.FORWARD), # root | ||
| 486 | + ("unshard", "layers.0", TrainingState.FORWARD), | ||
| 487 | + ("unshard", "layers.1", TrainingState.FORWARD), | ||
| 488 | + ("unshard", "layers.2", TrainingState.FORWARD), | ||
| 489 | + # Root does not reshard after forward so there is not another | ||
| 490 | + # unshard event for it | ||
| 491 | + ("unshard", "layers.0", TrainingState.FORWARD), | ||
| 492 | + ("unshard", "layers.1", TrainingState.FORWARD), | ||
| 493 | + ("unshard", "layers.2", TrainingState.FORWARD), | ||
| 494 | + ] | ||
| 495 | + if reshard_after_forward is False: | ||
| 496 | + # No reshard after forward means no second set of unshards | ||
| 497 | + expected_events = expected_events[:-3] | ||
| 498 | + self.assertEqual(events, expected_events) | ||
| 499 | + events.clear() | ||
| 500 | + (loss1 + loss2).sum().backward() | ||
| 501 | + expected_events = [ | ||
| 502 | + # Same as the single forward/backward case except the root's | ||
| 503 | + # post-backward does not run until the end of backward in the | ||
| 504 | + # final callback (since the input not requiring gradient means | ||
| 505 | + # that we do not have a tensor on which to hook for | ||
| 506 | + # post-backward) | ||
| 507 | + ("unshard", "layers.2", TrainingState.PRE_BACKWARD), | ||
| 508 | + ("unshard", "layers.1", TrainingState.PRE_BACKWARD), | ||
| 509 | + ("post_backward", "layers.2", TrainingState.POST_BACKWARD), | ||
| 510 | + ("unshard", "layers.0", TrainingState.PRE_BACKWARD), | ||
| 511 | + ("post_backward", "layers.1", TrainingState.POST_BACKWARD), | ||
| 512 | + ("post_backward", "layers.0", TrainingState.POST_BACKWARD), | ||
| 513 | + ] | ||
| 514 | + if reshard_after_forward is False: | ||
| 515 | + # No reshard after forward means no backward unshards | ||
| 516 | + expected_events = [e for e in expected_events if e[0] != "unshard"] | ||
| 517 | + # However, the post-backward reshards, so the second set of | ||
| 518 | + # unshards will run as real ops | ||
| 519 | + expected_events += [ | ||
| 520 | + # Repeat the same pattern except with the root's post-backward | ||
| 521 | + # at the end since the final callback runs | ||
| 522 | + ("unshard", "layers.2", TrainingState.PRE_BACKWARD), | ||
| 523 | + ("unshard", "layers.1", TrainingState.PRE_BACKWARD), | ||
| 524 | + ("post_backward", "layers.2", TrainingState.POST_BACKWARD), | ||
| 525 | + ("unshard", "layers.0", TrainingState.PRE_BACKWARD), | ||
| 526 | + ("post_backward", "layers.1", TrainingState.POST_BACKWARD), | ||
| 527 | + ("post_backward", "layers.0", TrainingState.POST_BACKWARD), | ||
| 528 | + ("post_backward", "", TrainingState.POST_BACKWARD), | ||
| 529 | + ] | ||
| 530 | + self.assertEqual(events, expected_events) | ||
| 531 | + events.clear() | ||
| 532 | + | ||
| 533 | + def _test_backward_prefetch_unused_in_backward( | ||
| 534 | + self, reshard_after_forward: Union[bool, int] | ||
| 535 | + ): | ||
| 536 | + """ | ||
| 537 | + Test a model with a linear module then a split into two linear modules, | ||
| 538 | + where we run backward through one path first before the other, meaning | ||
| 539 | + that (1) only one linear of the two split is used per backward and (2) | ||
| 540 | + the initial shared linear is used in both backwards. | ||
| 541 | + """ | ||
| 542 | + dim = 8 | ||
| 543 | + model = nn.Sequential(nn.Linear(dim, dim), DoubleLinear(dim)) | ||
| 544 | + fully_shard(model[0], reshard_after_forward=reshard_after_forward) | ||
| 545 | + fully_shard(model[1].lin1, reshard_after_forward=reshard_after_forward) | ||
| 546 | + fully_shard(model[1].lin2, reshard_after_forward=reshard_after_forward) | ||
| 547 | + fully_shard(model, reshard_after_forward=reshard_after_forward) | ||
| 548 | + inp = torch.randn((4, dim), device="npu") | ||
| 549 | + events: List[EventType] = [] | ||
| 550 | + unshard_with_record = self._get_unshard_with_record( | ||
| 551 | + FSDPParamGroup.unshard, events | ||
| 552 | + ) | ||
| 553 | + post_backward_with_record = self._get_post_backward_with_record( | ||
| 554 | + FSDPParamGroup.post_backward, events | ||
| 555 | + ) | ||
| 556 | + with patch_unshard(unshard_with_record), patch_post_backward( | ||
| 557 | + post_backward_with_record | ||
| 558 | + ): | ||
| 559 | + loss1, loss2 = model(inp) | ||
| 560 | + expected_events = [ | ||
| 561 | + # Root has no parameters, so it does not have an unshard | ||
| 562 | + ("unshard", "0", TrainingState.FORWARD), | ||
| 563 | + ("unshard", "1.lin1", TrainingState.FORWARD), | ||
| 564 | + ("unshard", "1.lin2", TrainingState.FORWARD), | ||
| 565 | + ] | ||
| 566 | + self.assertEqual(events, expected_events) | ||
| 567 | + events.clear() | ||
| 568 | + | ||
| 569 | + model.set_is_last_backward(False) | ||
| 570 | + loss2.sum().backward(retain_graph=True) | ||
| 571 | + expected_events = [ | ||
| 572 | + ("unshard", "1.lin2", TrainingState.PRE_BACKWARD), | ||
| 573 | + # NOTE: This `1.lin1` unshard is a mistargeted prefetch. | ||
| 574 | + ("unshard", "1.lin1", TrainingState.PRE_BACKWARD), | ||
| 575 | + ("post_backward", "1.lin2", TrainingState.POST_BACKWARD), | ||
| 576 | + ("unshard", "0", TrainingState.PRE_BACKWARD), | ||
| 577 | + ("post_backward", "0", TrainingState.POST_BACKWARD), | ||
| 578 | + ] | ||
| 579 | + self.assertEqual(events, expected_events) | ||
| 580 | + events.clear() | ||
| 581 | + | ||
| 582 | + model.set_is_last_backward(True) | ||
| 583 | + loss1.sum().backward() | ||
| 584 | + expected_events = [ | ||
| 585 | + # NOTE: `1.lin1` is already unsharded from the mistargeted | ||
| 586 | + # prefetch in the first backward. | ||
| 587 | + # Prefetch `0` | ||
| 588 | + ("unshard", "0", TrainingState.PRE_BACKWARD), | ||
| 589 | + ("post_backward", "1.lin1", TrainingState.POST_BACKWARD), | ||
| 590 | + ("post_backward", "0", TrainingState.POST_BACKWARD), | ||
| 591 | + ] | ||
| 592 | + self.assertEqual(events, expected_events) | ||
| 593 | + events.clear() | ||
| 594 | + | ||
| 595 | + | ||
| 596 | + def test_set_modules_to_forward_prefetch(self): | ||
| 597 | + n_layers = 4 | ||
| 598 | + reshard_after_forward = True | ||
| 599 | + checkpoint_impl = "utils" | ||
| 600 | + model, _, inp = self._init_transformer( | ||
| 601 | + n_layers, reshard_after_forward, checkpoint_impl | ||
| 602 | + ) | ||
| 603 | + | ||
| 604 | + def set_forward_prefetch(model: Transformer, num_to_prefetch: int) -> None: | ||
| 605 | + # Use model-specific knowledge to configure forward prefetching: | ||
| 606 | + # each transformer block (layer) prefetches for the next few | ||
| 607 | + for i, layer in enumerate(model.layers): | ||
| 608 | + if i >= len(model.layers) - num_to_prefetch: | ||
| 609 | + break | ||
| 610 | + layers_to_prefetch = [model.layers[i + j] for j in range(1, num_to_prefetch + 1)] | ||
| 611 | + layer.set_modules_to_forward_prefetch(layers_to_prefetch) | ||
| 612 | + | ||
| 613 | + events: List[EventType] = [] | ||
| 614 | + unshard_with_record = self._get_unshard_with_record( | ||
| 615 | + FSDPParamGroup.unshard, events | ||
| 616 | + ) | ||
| 617 | + reshard_with_record = self._get_reshard_with_record( | ||
| 618 | + FSDPParamGroup.reshard, events | ||
| 619 | + ) | ||
| 620 | + post_backward_with_record = self._get_post_backward_with_record( | ||
| 621 | + FSDPParamGroup.post_backward, events | ||
| 622 | + ) | ||
| 623 | + expected_backward_events = [ | ||
| 624 | + # Default backward prefetching | ||
| 625 | + ("unshard", "layers.3", TrainingState.PRE_BACKWARD), | ||
| 626 | + ("unshard", "layers.2", TrainingState.PRE_BACKWARD), | ||
| 627 | + ("reshard", "layers.3", TrainingState.POST_BACKWARD), | ||
| 628 | + ("post_backward", "layers.3", TrainingState.POST_BACKWARD), | ||
| 629 | + ("unshard", "layers.1", TrainingState.PRE_BACKWARD), | ||
| 630 | + ("reshard", "layers.2", TrainingState.POST_BACKWARD), | ||
| 631 | + ("post_backward", "layers.2", TrainingState.POST_BACKWARD), | ||
| 632 | + ("unshard", "layers.0", TrainingState.PRE_BACKWARD), | ||
| 633 | + ("reshard", "layers.1", TrainingState.POST_BACKWARD), | ||
| 634 | + ("post_backward", "layers.1", TrainingState.POST_BACKWARD), | ||
| 635 | + ("reshard", "layers.0", TrainingState.POST_BACKWARD), | ||
| 636 | + ("post_backward", "layers.0", TrainingState.POST_BACKWARD), | ||
| 637 | + ("reshard", "", TrainingState.POST_BACKWARD), | ||
| 638 | + ("post_backward", "", TrainingState.POST_BACKWARD), | ||
| 639 | + ] | ||
| 640 | + with patch_unshard(unshard_with_record), patch_reshard( | ||
| 641 | + reshard_with_record | ||
| 642 | + ), patch_post_backward(post_backward_with_record): | ||
| 643 | + set_forward_prefetch(model, num_to_prefetch=1) | ||
| 644 | + loss = model(inp) | ||
| 645 | + expected_forward_events = [ | ||
| 646 | + ("unshard", "", TrainingState.FORWARD), | ||
| 647 | + # `layers.i` prefetches `layers.i+1` | ||
| 648 | + ("unshard", "layers.0", TrainingState.FORWARD), | ||
| 649 | + ("unshard", "layers.1", TrainingState.FORWARD), | ||
| 650 | + ("reshard", "layers.0", TrainingState.FORWARD), | ||
| 651 | + ("unshard", "layers.2", TrainingState.FORWARD), | ||
| 652 | + ("reshard", "layers.1", TrainingState.FORWARD), | ||
| 653 | + ("unshard", "layers.3", TrainingState.FORWARD), | ||
| 654 | + ("reshard", "layers.2", TrainingState.FORWARD), | ||
| 655 | + ("reshard", "layers.3", TrainingState.FORWARD), | ||
| 656 | + ] | ||
| 657 | + self.assertEqual(events, expected_forward_events) | ||
| 658 | + events.clear() | ||
| 659 | + loss.sum().backward() | ||
| 660 | + self.assertEqual(events, expected_backward_events) | ||
| 661 | + events.clear() | ||
| 662 | + | ||
| 663 | + set_forward_prefetch(model, num_to_prefetch=2) | ||
| 664 | + loss = model(inp) | ||
| 665 | + expected_forward_events = [ | ||
| 666 | + ("unshard", "", TrainingState.FORWARD), | ||
| 667 | + # `layers.i` prefetches `layers.i+1` and `layers.i+2` | ||
| 668 | + ("unshard", "layers.0", TrainingState.FORWARD), | ||
| 669 | + ("unshard", "layers.1", TrainingState.FORWARD), | ||
| 670 | + ("unshard", "layers.2", TrainingState.FORWARD), | ||
| 671 | + ("reshard", "layers.0", TrainingState.FORWARD), | ||
| 672 | + ("unshard", "layers.3", TrainingState.FORWARD), | ||
| 673 | + ("reshard", "layers.1", TrainingState.FORWARD), | ||
| 674 | + ("reshard", "layers.2", TrainingState.FORWARD), | ||
| 675 | + ("reshard", "layers.3", TrainingState.FORWARD), | ||
| 676 | + ] | ||
| 677 | + self.assertEqual(events, expected_forward_events) | ||
| 678 | + events.clear() | ||
| 679 | + loss.sum().backward() | ||
| 680 | + self.assertEqual(events, expected_backward_events) | ||
| 681 | + events.clear() | ||
| 682 | + | ||
| 683 | + | ||
| 684 | + def test_set_modules_to_backward_prefetch(self): | ||
| 685 | + n_layers = 4 | ||
| 686 | + reshard_after_forward = True | ||
| 687 | + checkpoint_impl = "utils" | ||
| 688 | + model, _, inp = self._init_transformer( | ||
| 689 | + n_layers, reshard_after_forward, checkpoint_impl | ||
| 690 | + ) | ||
| 691 | + | ||
| 692 | + def set_backward_prefetch(model: Transformer, num_to_prefetch: int) -> None: | ||
| 693 | + # Use model-specific knowledge to configure backward prefetching: | ||
| 694 | + # each transformer block (layer) prefetches for the previous few | ||
| 695 | + for i, layer in enumerate(model.layers): | ||
| 696 | + if i < num_to_prefetch: | ||
| 697 | + continue | ||
| 698 | + layers_to_prefetch = [model.layers[i - j] for j in range(1, num_to_prefetch + 1)] | ||
| 699 | + layer.set_modules_to_backward_prefetch(layers_to_prefetch) | ||
| 700 | + | ||
| 701 | + events: List[EventType] = [] | ||
| 702 | + unshard_with_record = self._get_unshard_with_record( | ||
| 703 | + FSDPParamGroup.unshard, events | ||
| 704 | + ) | ||
| 705 | + reshard_with_record = self._get_reshard_with_record( | ||
| 706 | + FSDPParamGroup.reshard, events | ||
| 707 | + ) | ||
| 708 | + post_backward_with_record = self._get_post_backward_with_record( | ||
| 709 | + FSDPParamGroup.post_backward, events | ||
| 710 | + ) | ||
| 711 | + expected_forward_events = [ | ||
| 712 | + # Default forward prefetching | ||
| 713 | + ("unshard", "", TrainingState.FORWARD), # root | ||
| 714 | + ("unshard", "layers.0", TrainingState.FORWARD), | ||
| 715 | + ("reshard", "layers.0", TrainingState.FORWARD), | ||
| 716 | + ("unshard", "layers.1", TrainingState.FORWARD), | ||
| 717 | + ("reshard", "layers.1", TrainingState.FORWARD), | ||
| 718 | + ("unshard", "layers.2", TrainingState.FORWARD), | ||
| 719 | + ("reshard", "layers.2", TrainingState.FORWARD), | ||
| 720 | + ("unshard", "layers.3", TrainingState.FORWARD), | ||
| 721 | + ("reshard", "layers.3", TrainingState.FORWARD), | ||
| 722 | + ] | ||
| 723 | + with patch_unshard(unshard_with_record), patch_reshard( | ||
| 724 | + reshard_with_record | ||
| 725 | + ), patch_post_backward(post_backward_with_record): | ||
| 726 | + set_backward_prefetch(model, num_to_prefetch=1) | ||
| 727 | + loss = model(inp) | ||
| 728 | + self.assertEqual(events, expected_forward_events) | ||
| 729 | + events.clear() | ||
| 730 | + loss.sum().backward() | ||
| 731 | + expected_backward_events = [ | ||
| 732 | + # Root prefetches `layers.3` per default | ||
| 733 | + ("unshard", "layers.3", TrainingState.PRE_BACKWARD), | ||
| 734 | + # `layers.i` prefetches for `layers.i-1` (same as default) | ||
| 735 | + ("unshard", "layers.2", TrainingState.PRE_BACKWARD), | ||
| 736 | + ("reshard", "layers.3", TrainingState.POST_BACKWARD), | ||
| 737 | + ("post_backward", "layers.3", TrainingState.POST_BACKWARD), | ||
| 738 | + ("unshard", "layers.1", TrainingState.PRE_BACKWARD), | ||
| 739 | + ("reshard", "layers.2", TrainingState.POST_BACKWARD), | ||
| 740 | + ("post_backward", "layers.2", TrainingState.POST_BACKWARD), | ||
| 741 | + ("unshard", "layers.0", TrainingState.PRE_BACKWARD), | ||
| 742 | + ("reshard", "layers.1", TrainingState.POST_BACKWARD), | ||
| 743 | + ("post_backward", "layers.1", TrainingState.POST_BACKWARD), | ||
| 744 | + ("reshard", "layers.0", TrainingState.POST_BACKWARD), | ||
| 745 | + ("post_backward", "layers.0", TrainingState.POST_BACKWARD), | ||
| 746 | + ("reshard", "", TrainingState.POST_BACKWARD), | ||
| 747 | + ("post_backward", "", TrainingState.POST_BACKWARD), | ||
| 748 | + ] | ||
| 749 | + self.assertEqual(events, expected_backward_events) | ||
| 750 | + events.clear() | ||
| 751 | + | ||
| 752 | + set_backward_prefetch(model, num_to_prefetch=2) | ||
| 753 | + loss = model(inp) | ||
| 754 | + self.assertEqual(events, expected_forward_events) | ||
| 755 | + events.clear() | ||
| 756 | + loss.sum().backward() | ||
| 757 | + expected_backward_events = [ | ||
| 758 | + # Root prefetches `layers.3` per default | ||
| 759 | + ("unshard", "layers.3", TrainingState.PRE_BACKWARD), | ||
| 760 | + # `layers.i` prefetches for `layers.i-1` and `layers.i-2` | ||
| 761 | + ("unshard", "layers.2", TrainingState.PRE_BACKWARD), | ||
| 762 | + ("unshard", "layers.1", TrainingState.PRE_BACKWARD), | ||
| 763 | + ("reshard", "layers.3", TrainingState.POST_BACKWARD), | ||
| 764 | + ("post_backward", "layers.3", TrainingState.POST_BACKWARD), | ||
| 765 | + ("unshard", "layers.0", TrainingState.PRE_BACKWARD), | ||
| 766 | + ("reshard", "layers.2", TrainingState.POST_BACKWARD), | ||
| 767 | + ("post_backward", "layers.2", TrainingState.POST_BACKWARD), | ||
| 768 | + ("reshard", "layers.1", TrainingState.POST_BACKWARD), | ||
| 769 | + ("post_backward", "layers.1", TrainingState.POST_BACKWARD), | ||
| 770 | + ("reshard", "layers.0", TrainingState.POST_BACKWARD), | ||
| 771 | + ("post_backward", "layers.0", TrainingState.POST_BACKWARD), | ||
| 772 | + ("reshard", "", TrainingState.POST_BACKWARD), | ||
| 773 | + ("post_backward", "", TrainingState.POST_BACKWARD), | ||
| 774 | + ] | ||
| 775 | + self.assertEqual(events, expected_backward_events) | ||
| 776 | + events.clear() | ||
| 777 | + | ||
| 778 | + | ||
| 779 | + def test_fully_shard_multi_module_backward_prefetch(self): | ||
| 780 | + n_layers = 5 | ||
| 781 | + model_args = ModelArgs(n_layers=n_layers, checkpoint_activations=True) | ||
| 782 | + model = Transformer(model_args) | ||
| 783 | + for i in range(n_layers): | ||
| 784 | + if i == 0: | ||
| 785 | + fully_shard(model.layers[i]) | ||
| 786 | + elif i % 2 == 1: | ||
| 787 | + fully_shard([model.layers[i], model.layers[i + 1]]) | ||
| 788 | + fully_shard([model.tok_embeddings, model.pos_embeddings]) | ||
| 789 | + fully_shard([model.norm, model.output], reshard_after_forward=False) | ||
| 790 | + fully_shard(model) | ||
| 791 | + optim = torch.optim.AdamW(model.parameters(), lr=1e-2) | ||
| 792 | + | ||
| 793 | + events: List[EventType] = [] | ||
| 794 | + unshard_with_record = self._get_unshard_with_record( | ||
| 795 | + FSDPParamGroup.unshard, events | ||
| 796 | + ) | ||
| 797 | + post_backward_with_record = self._get_post_backward_with_record( | ||
| 798 | + FSDPParamGroup.post_backward, events | ||
| 799 | + ) | ||
| 800 | + inp = torch.randint( | ||
| 801 | + 0, model_args.vocab_size, (2, model_args.max_seq_len), device="npu" | ||
| 802 | + ) | ||
| 803 | + with patch_unshard(unshard_with_record), patch_post_backward( | ||
| 804 | + post_backward_with_record | ||
| 805 | + ): | ||
| 806 | + for _ in range(3): | ||
| 807 | + loss = model(inp) | ||
| 808 | + expected_events = [ | ||
| 809 | + ( | ||
| 810 | + "unshard", | ||
| 811 | + "tok_embeddings, pos_embeddings", | ||
| 812 | + TrainingState.FORWARD, | ||
| 813 | + ), | ||
| 814 | + ("unshard", "layers.0", TrainingState.FORWARD), | ||
| 815 | + ("unshard", "layers.1, layers.2", TrainingState.FORWARD), | ||
| 816 | + ("unshard", "layers.3, layers.4", TrainingState.FORWARD), | ||
| 817 | + ("unshard", "norm, output", TrainingState.FORWARD), | ||
| 818 | + ] | ||
| 819 | + self.assertEqual(events, expected_events) | ||
| 820 | + events.clear() | ||
| 821 | + loss.sum().backward() | ||
| 822 | + expected_events = [ | ||
| 823 | + # (norm, output) does not reshard after forward, so there is | ||
| 824 | + # no unshard to begin backward | ||
| 825 | + ("unshard", "layers.3, layers.4", TrainingState.PRE_BACKWARD), | ||
| 826 | + ("post_backward", "norm, output", TrainingState.POST_BACKWARD), | ||
| 827 | + ("unshard", "layers.1, layers.2", TrainingState.PRE_BACKWARD), | ||
| 828 | + ( | ||
| 829 | + "post_backward", | ||
| 830 | + "layers.3, layers.4", | ||
| 831 | + TrainingState.POST_BACKWARD, | ||
| 832 | + ), | ||
| 833 | + ("unshard", "layers.0", TrainingState.PRE_BACKWARD), | ||
| 834 | + ( | ||
| 835 | + "post_backward", | ||
| 836 | + "layers.1, layers.2", | ||
| 837 | + TrainingState.POST_BACKWARD, | ||
| 838 | + ), | ||
| 839 | + ( | ||
| 840 | + "unshard", | ||
| 841 | + "tok_embeddings, pos_embeddings", | ||
| 842 | + TrainingState.PRE_BACKWARD, | ||
| 843 | + ), | ||
| 844 | + ("post_backward", "layers.0", TrainingState.POST_BACKWARD), | ||
| 845 | + ( | ||
| 846 | + "post_backward", | ||
| 847 | + "tok_embeddings, pos_embeddings", | ||
| 848 | + TrainingState.POST_BACKWARD, | ||
| 849 | + ), | ||
| 850 | + ] | ||
| 851 | + events.clear() | ||
| 852 | + optim.step() | ||
| 853 | + optim.zero_grad() | ||
| 854 | + | ||
| 855 | + | ||
| 856 | + def test_fully_shard_multi_module_unused_module(self): | ||
| 857 | + class ModuleWithUnusedLinear(nn.Module): | ||
| 858 | + def __init__(self) -> None: | ||
| 859 | + super().__init__() | ||
| 860 | + self.unused_lin = nn.Linear(1, 1) | ||
| 861 | + self.lin = nn.Linear(16, 16) | ||
| 862 | + | ||
| 863 | + def forward(self, x: torch.Tensor) -> torch.Tensor: | ||
| 864 | + return nn.functional.relu(self.lin(x)) | ||
| 865 | + | ||
| 866 | + model = nn.Sequential( | ||
| 867 | + ModuleWithUnusedLinear(), ModuleWithUnusedLinear(), nn.Linear(16, 16) | ||
| 868 | + ) | ||
| 869 | + fully_shard([model[0].unused_lin, model[0].lin], reshard_after_forward=True) | ||
| 870 | + fully_shard([model[1].unused_lin, model[1].lin], reshard_after_forward=True) | ||
| 871 | + fully_shard(model) | ||
| 872 | + optim = torch.optim.AdamW(model.parameters(), lr=1e-2) | ||
| 873 | + | ||
| 874 | + events: List[EventType] = [] | ||
| 875 | + unshard_with_record = self._get_unshard_with_record( | ||
| 876 | + FSDPParamGroup.unshard, events | ||
| 877 | + ) | ||
| 878 | + post_backward_with_record = self._get_post_backward_with_record( | ||
| 879 | + FSDPParamGroup.post_backward, events | ||
| 880 | + ) | ||
| 881 | + inp = torch.randn((2, 16), device="npu") | ||
| 882 | + with patch_unshard(unshard_with_record), patch_post_backward( | ||
| 883 | + post_backward_with_record | ||
| 884 | + ): | ||
| 885 | + for _ in range(3): | ||
| 886 | + loss = model(inp) | ||
| 887 | + expected_events = [ | ||
| 888 | + ("unshard", "", TrainingState.FORWARD), | ||
| 889 | + ("unshard", "0.unused_lin, 0.lin", TrainingState.FORWARD), | ||
| 890 | + ("unshard", "1.unused_lin, 1.lin", TrainingState.FORWARD), | ||
| 891 | + ] | ||
| 892 | + self.assertEqual(events, expected_events) | ||
| 893 | + events.clear() | ||
| 894 | + loss.sum().backward() | ||
| 895 | + expected_events = [ | ||
| 896 | + # Since both `model[0]` and `model[1]` have unused modules | ||
| 897 | + # that never ran forward, they do not reshard after forward | ||
| 898 | + # despite setting it to `True`. Check that there are no | ||
| 899 | + # unshards in backward. | ||
| 900 | + ( | ||
| 901 | + "post_backward", | ||
| 902 | + "1.unused_lin, 1.lin", | ||
| 903 | + TrainingState.POST_BACKWARD, | ||
| 904 | + ), | ||
| 905 | + ( | ||
| 906 | + "post_backward", | ||
| 907 | + "0.unused_lin, 0.lin", | ||
| 908 | + TrainingState.POST_BACKWARD, | ||
| 909 | + ), | ||
| 910 | + ("post_backward", "", TrainingState.POST_BACKWARD), | ||
| 911 | + ] | ||
| 912 | + events.clear() | ||
| 913 | + optim.step() | ||
| 914 | + optim.zero_grad() | ||
| 915 | + | ||
| 916 | + def test_backward_misprefetch(self): | ||
| 917 | + torch.manual_seed(42) | ||
| 918 | + model = MLP(dim=16, device="npu") | ||
| 919 | + ref_model = copy.deepcopy(model) | ||
| 920 | + ref_optim = torch.optim.Adam(ref_model.parameters(), lr=1e-2) | ||
| 921 | + fully_shard(model.in_proj) | ||
| 922 | + fully_shard(model.out_proj) | ||
| 923 | + fully_shard(model) | ||
| 924 | + optim = torch.optim.Adam(model.parameters(), lr=1e-2) | ||
| 925 | + | ||
| 926 | + # Backward should run through `out_proj` -> `in_proj`, so if `in_proj` | ||
| 927 | + # prefetches for `out_proj`, then this is a misprefetch, as `out_proj` | ||
| 928 | + # should not be needed anymore for backward. | ||
| 929 | + model.in_proj.set_modules_to_backward_prefetch([model.out_proj]) | ||
| 930 | + | ||
| 931 | + torch.manual_seed(self.rank + 1) | ||
| 932 | + inp = torch.randn((2, 16), device="npu") | ||
| 933 | + for _ in range(3): | ||
| 934 | + ref_optim.zero_grad() | ||
| 935 | + ref_loss = ref_model(inp).sum() | ||
| 936 | + ref_loss.backward() | ||
| 937 | + for param in ref_model.parameters(): | ||
| 938 | + dist.all_reduce(param.grad, op=dist.ReduceOp.AVG) | ||
| 939 | + ref_optim.step() | ||
| 940 | + optim.zero_grad() | ||
| 941 | + loss = model(inp).sum() | ||
| 942 | + loss.backward() | ||
| 943 | + optim.step() | ||
| 944 | + self.assertEqual(ref_loss, loss) | ||
| 945 | + | ||
| 946 | + def _init_transformer( | ||
| 947 | + self, | ||
| 948 | + n_layers: int, | ||
| 949 | + reshard_after_forward: Union[bool, int], | ||
| 950 | + checkpoint_impl: Optional[str], | ||
| 951 | + ): | ||
| 952 | + model_args = ModelArgs( | ||
| 953 | + n_layers=n_layers, checkpoint_activations=(checkpoint_impl == "utils") | ||
| 954 | + ) | ||
| 955 | + model = Transformer(model_args) | ||
| 956 | + for module in model.modules(): | ||
| 957 | + if isinstance(module, TransformerBlock): | ||
| 958 | + if checkpoint_impl == "composable": | ||
| 959 | + checkpoint(module) | ||
| 960 | + fully_shard(module, reshard_after_forward=reshard_after_forward) | ||
| 961 | + fully_shard(model, reshard_after_forward=reshard_after_forward) | ||
| 962 | + optim = torch.optim.Adam(model.parameters(), lr=1e-2) | ||
| 963 | + inp = torch.randint( | ||
| 964 | + 0, model_args.vocab_size, (2, model_args.max_seq_len), device="npu" | ||
| 965 | + ) | ||
| 966 | + return model, optim, inp | ||
| 967 | + | ||
| 968 | + def _get_unshard_with_record( | ||
| 969 | + self, orig_unshard: Callable, events: List[EventType] | ||
| 970 | + ) -> Callable: | ||
| 971 | + def unshard_with_record(self, *args, **kwargs): | ||
| 972 | + nonlocal events | ||
| 973 | + if ( | ||
| 974 | + self._all_gather_result is None | ||
| 975 | + and self._sharded_state != ShardedState.UNSHARDED | ||
| 976 | + ): # skip no-ops | ||
| 977 | + events.append(("unshard", self._module_fqn, self._training_state)) | ||
| 978 | + return orig_unshard(self, *args, **kwargs) | ||
| 979 | + | ||
| 980 | + return unshard_with_record | ||
| 981 | + | ||
| 982 | + def _get_reshard_with_record( | ||
| 983 | + self, orig_reshard: Callable, events: List[EventType] | ||
| 984 | + ) -> Callable: | ||
| 985 | + def reshard_with_record(self, *args, **kwargs): | ||
| 986 | + nonlocal events | ||
| 987 | + if ( | ||
| 988 | + self._training_state == TrainingState.FORWARD | ||
| 989 | + and not self._reshard_after_forward | ||
| 990 | + ): # skip no-ops | ||
| 991 | + return None | ||
| 992 | + events.append(("reshard", self._module_fqn, self._training_state)) | ||
| 993 | + return orig_reshard(self, *args, **kwargs) | ||
| 994 | + | ||
| 995 | + return reshard_with_record | ||
| 996 | + | ||
| 997 | + def _get_post_backward_with_record( | ||
| 998 | + self, orig_post_backward: Callable, events: List[EventType] | ||
| 999 | + ) -> Callable: | ||
| 1000 | + def post_backward_with_record(self, *args, **kwargs): | ||
| 1001 | + nonlocal events | ||
| 1002 | + ret = orig_post_backward(self, *args, **kwargs) | ||
| 1003 | + # Use training state after running post-backward to check that the | ||
| 1004 | + # state is transitioned to `POST_BACKWARD` as expected | ||
| 1005 | + events.append(("post_backward", self._module_fqn, self._training_state)) | ||
| 1006 | + return ret | ||
| 1007 | + | ||
| 1008 | + return post_backward_with_record | ||
| 1009 | + | ||
| 1010 | + | ||
| 1011 | +class TestFullyShardUnshardMultiProcess(FSDPNPUTest): | ||
| 1012 | + | ||
| 1013 | + def world_size(self) -> int: | ||
| 1014 | + return min(torch.npu.device_count(), 2) | ||
| 1015 | + | ||
| 1016 | + def test_unshard_async(self): | ||
| 1017 | + class ReduceModule(nn.Module): | ||
| 1018 | + def __init__(self, dim: int, mesh: DeviceMesh): | ||
| 1019 | + super().__init__() | ||
| 1020 | + self.mesh = mesh | ||
| 1021 | + self.weight = nn.Parameter(torch.randn(dim, dim)) | ||
| 1022 | + | ||
| 1023 | + def forward(self, x: torch.Tensor): | ||
| 1024 | + y = F.relu(x @ self.weight) | ||
| 1025 | + # NOTE: This all-reduce is not differentiable and is included | ||
| 1026 | + # to exercise the overlap. | ||
| 1027 | + work = dist.all_reduce(y, group=self.mesh.get_group(), async_op=True) | ||
| 1028 | + return y, work | ||
| 1029 | + | ||
| 1030 | + class MLPs(nn.Module): | ||
| 1031 | + def __init__(self, dim: int): | ||
| 1032 | + super().__init__() | ||
| 1033 | + self.mlp1 = MLP(dim) | ||
| 1034 | + self.mlp2 = MLP(dim) | ||
| 1035 | + self.mlp3 = MLP(dim) | ||
| 1036 | + | ||
| 1037 | + def forward(self, ys: List[torch.Tensor], works: List[dist.Work]): | ||
| 1038 | + (y1, y2, y3), (work1, work2, work3) = ys, works | ||
| 1039 | + work1.wait() | ||
| 1040 | + z1 = self.mlp1(y1) | ||
| 1041 | + work2.wait() | ||
| 1042 | + z2 = self.mlp2(y2) | ||
| 1043 | + work3.wait() | ||
| 1044 | + z3 = self.mlp3(y3) | ||
| 1045 | + return z1 + z2 + z3 | ||
| 1046 | + | ||
| 1047 | + class ReduceModel(nn.Module): | ||
| 1048 | + def __init__(self, dim: int, mesh: DeviceMesh): | ||
| 1049 | + super().__init__() | ||
| 1050 | + self.reduce_module1 = ReduceModule(dim, mesh) | ||
| 1051 | + self.reduce_module2 = ReduceModule(dim, mesh) | ||
| 1052 | + self.reduce_module3 = ReduceModule(dim, mesh) | ||
| 1053 | + self.mlps = MLPs(dim) | ||
| 1054 | + | ||
| 1055 | + def forward(self, x: torch.Tensor): | ||
| 1056 | + y1, work1 = self.reduce_module1(x) | ||
| 1057 | + if isinstance(self.mlps.mlp1, FSDPModule): | ||
| 1058 | + self.mlps.mlp1.unshard(async_op=True) | ||
| 1059 | + y2, work2 = self.reduce_module2(x) | ||
| 1060 | + if isinstance(self.mlps.mlp2, FSDPModule): | ||
| 1061 | + self.mlps.mlp2.unshard(async_op=True) | ||
| 1062 | + y3, work3 = self.reduce_module3(x) | ||
| 1063 | + if isinstance(self.mlps.mlp3, FSDPModule): | ||
| 1064 | + self.mlps.mlp3.unshard(async_op=True) | ||
| 1065 | + return self.mlps([y1, y2, y3], [work1, work2, work3]) | ||
| 1066 | + | ||
| 1067 | + mesh = init_device_mesh("npu", (self.world_size,)) | ||
| 1068 | + batch_size, dim = 2, 8 | ||
| 1069 | + torch.manual_seed(42) | ||
| 1070 | + ref_model = replicate(ReduceModel(dim, mesh).npu()) | ||
| 1071 | + ref_optim = torch.optim.Adam(ref_model.parameters(), lr=1e-2) | ||
| 1072 | + torch.manual_seed(42) | ||
| 1073 | + model = ReduceModel(dim, mesh) | ||
| 1074 | + fully_shard(model.mlps.mlp1, reshard_after_forward=False) | ||
| 1075 | + fully_shard(model.mlps.mlp2, reshard_after_forward=False) | ||
| 1076 | + fully_shard(model.mlps.mlp3, reshard_after_forward=False) | ||
| 1077 | + fully_shard(model.mlps) | ||
| 1078 | + replicate(model.npu()) | ||
| 1079 | + optim = torch.optim.Adam(model.parameters(), lr=1e-2, foreach=True) | ||
| 1080 | + torch.manual_seed(42 + self.rank + 1) | ||
| 1081 | + inp = torch.randn((batch_size, dim), device="npu") | ||
| 1082 | + for _ in range(10): | ||
| 1083 | + losses: List[torch.Tensor] = [] | ||
| 1084 | + for _model, _optim in ((ref_model, ref_optim), (model, optim)): | ||
| 1085 | + losses.append(_model(inp).sum()) | ||
| 1086 | + losses[-1].backward() | ||
| 1087 | + with implicit_replication(): | ||
| 1088 | + _optim.step() | ||
| 1089 | + _optim.zero_grad() | ||
| 1090 | + self.assertEqual(losses[0], losses[1]) | ||
| 1091 | + | ||
| 1092 | + | ||
| 1093 | +class TestFullyShardUnshardMultiThread(FSDPTestMultiThread): | ||
| 1094 | + | ||
| 1095 | + def world_size(self) -> int: | ||
| 1096 | + return 2 | ||
| 1097 | + | ||
| 1098 | + def perThreadSetUp(self): | ||
| 1099 | + super().perThreadSetUp() | ||
| 1100 | + torch.npu.set_device(0) | ||
| 1101 | + | ||
| 1102 | + | ||
| 1103 | + def test_unshard_no_param_group(self): | ||
| 1104 | + # Check that we can call `unshard()` on a module with no parameter | ||
| 1105 | + # group / no managed parameters without erroring | ||
| 1106 | + model = nn.Sequential(nn.Linear(4, 4), nn.Linear(4, 4)) | ||
| 1107 | + for lin in model: | ||
| 1108 | + fully_shard(lin) | ||
| 1109 | + fully_shard(model) | ||
| 1110 | + handle = model.unshard(async_op=True) | ||
| 1111 | + handle.wait() | ||
| 1112 | + | ||
| 1113 | + | ||
| 1114 | + def test_unshard_without_lazy_init(self): | ||
| 1115 | + torch.manual_seed(42) | ||
| 1116 | + model = MLP(4) | ||
| 1117 | + for param in model.parameters(): | ||
| 1118 | + dist.broadcast(param, src=0) | ||
| 1119 | + ref_model = copy.deepcopy(model) | ||
| 1120 | + fully_shard(model) | ||
| 1121 | + model.unshard() # no lazy init yet | ||
| 1122 | + for ref_param, param in zip(ref_model.parameters(), model.parameters()): | ||
| 1123 | + self.assertEqual(ref_param, param) | ||
| 1124 | + | ||
| 1125 | + | ||
| 1126 | +if __name__ == "__main__": | ||
| 1127 | + run_tests() | ||