diff --git a/distributed/tensor_parallelism/fsdp_tp_example.py b/distributed/tensor_parallelism/fsdp_tp_example.py index fb0d5ba1f5..748c912586 100644 --- a/distributed/tensor_parallelism/fsdp_tp_example.py +++ b/distributed/tensor_parallelism/fsdp_tp_example.py @@ -170,6 +170,7 @@ output = sharded_model(inp) output.sum().backward() optimizer.step() + optimizer.zero_grad() rank_log(_rank, logger, f"2D iter {i} complete") rank_log(_rank, logger, "2D training successfully completed!") diff --git a/distributed/tensor_parallelism/sequence_parallel_example.py b/distributed/tensor_parallelism/sequence_parallel_example.py index 73320f5bcc..2702a88eca 100644 --- a/distributed/tensor_parallelism/sequence_parallel_example.py +++ b/distributed/tensor_parallelism/sequence_parallel_example.py @@ -105,6 +105,7 @@ def forward(self, x): output = sp_model(inp) output.sum().backward() optimizer.step() + optimizer.zero_grad() rank_log(_rank, logger, f"Sequence Parallel iter {i} completed") rank_log(_rank, logger, "Sequence Parallel training completed!") diff --git a/distributed/tensor_parallelism/tensor_parallel_example.py b/distributed/tensor_parallelism/tensor_parallel_example.py index 6a4b4ea531..ad674c2cd4 100755 --- a/distributed/tensor_parallelism/tensor_parallel_example.py +++ b/distributed/tensor_parallelism/tensor_parallel_example.py @@ -119,6 +119,7 @@ def forward(self, x): output = tp_model(inp) output.sum().backward() optimizer.step() + optimizer.zero_grad() rank_log(_rank, logger, f"Tensor Parallel iter {i} completed") rank_log(_rank, logger, "Tensor Parallel training completed!")