已合并
add fsdp test case #22803
zqwen创建于 2025年7月9日
add fsdp test case #22803
已合并
zqwen创建于 2025年7月9日
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+ @property
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+ @property
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+ @SupportedDevices(['Ascend910B'])
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+ @SupportedDevices(['Ascend910B'])
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+ @SupportedDevices(['Ascend910B'])
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+ @property
298+ def world_size(self) -> int:
299+ return min(4, torch.npu.device_count())
300+ 
301+ @SupportedDevices(['Ascend910B'])
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+ @SupportedDevices(['Ascend910B'])
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+ @property
391+ def world_size(self) -> int:
392+ return min(4, torch.npu.device_count())
393+ 
394+ @SupportedDevices(['Ascend910B'])
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+ @SupportedDevices(['Ascend910B'])
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+ @SupportedDevices(['Ascend910B'])
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+ @SupportedDevices(['Ascend910B'])
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+ @SupportedDevices(['Ascend910B'])
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+ @property
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+ @property
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+ @SupportedDevices(['Ascend910B'])
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+ @SupportedDevices(['Ascend910B'])
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()