From 5fa06f2646c34eed3d8ae30f47732e83b2bf1ae3 Mon Sep 17 00:00:00 2001 From: Himanshu Janbandhu Date: Tue, 18 Aug 2026 19:16:44 +0530 Subject: [PATCH] Zero gradients between steps in the tensor parallelism examples The three training loops in distributed/tensor_parallelism call backward() and optimizer.step() without ever calling zero_grad(), and none of the files calls it anywhere. PyTorch accumulates gradients into .grad by default, so each iteration steps on the running sum of every gradient computed so far rather than on that iteration's own. With num_iters = 10 the loops still run and still print, so nothing looks wrong -- but the optimizer is not doing what the example appears to demonstrate, and these files are a common starting point for real training code. Placed after optimizer.step() to match the sibling distributed/FSDP2/example.py. --- distributed/tensor_parallelism/fsdp_tp_example.py | 1 + distributed/tensor_parallelism/sequence_parallel_example.py | 1 + distributed/tensor_parallelism/tensor_parallel_example.py | 1 + 3 files changed, 3 insertions(+) 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!")