diff --git a/transformer_engine/pytorch/module/grouped_linear.py b/transformer_engine/pytorch/module/grouped_linear.py index f534da5c3..ae636a398 100644 --- a/transformer_engine/pytorch/module/grouped_linear.py +++ b/transformer_engine/pytorch/module/grouped_linear.py @@ -1617,7 +1617,8 @@ def set_tensor_parallel_attributes(self, defer_init=False) -> None: if grouped_weight is not None: set_tensor_model_parallel_attributes( tensor=grouped_weight, - is_parallel=True, + # Replicated when parallel_mode is None; see Linear.reset_parameters. + is_parallel=self.parallel_mode is not None, dim=1 if self.parallel_mode == "row" else 0, stride=1, ) @@ -1625,7 +1626,8 @@ def set_tensor_parallel_attributes(self, defer_init=False) -> None: for i in range(self.num_gemms): set_tensor_model_parallel_attributes( tensor=getattr(self, f"weight{i}"), - is_parallel=True, + # Replicated when parallel_mode is None. + is_parallel=self.parallel_mode is not None, dim=1 if self.parallel_mode == "row" else 0, stride=1, ) diff --git a/transformer_engine/pytorch/module/layernorm_linear.py b/transformer_engine/pytorch/module/layernorm_linear.py index 479b346bf..81c978459 100644 --- a/transformer_engine/pytorch/module/layernorm_linear.py +++ b/transformer_engine/pytorch/module/layernorm_linear.py @@ -1703,7 +1703,8 @@ def reset_parameters(self, defer_init=False): for weight in self.weight_names: set_tensor_model_parallel_attributes( tensor=getattr(self, weight), - is_parallel=True, + # Replicated when parallel_mode is None; see Linear.reset_parameters. + is_parallel=self.parallel_mode is not None, dim=1 if self.parallel_mode == "row" else 0, stride=1, ) diff --git a/transformer_engine/pytorch/module/linear.py b/transformer_engine/pytorch/module/linear.py index c4c9318b7..84df0f4a8 100644 --- a/transformer_engine/pytorch/module/linear.py +++ b/transformer_engine/pytorch/module/linear.py @@ -1843,7 +1843,12 @@ def reset_parameters(self, defer_init=False): for weight in self.weight_names: set_tensor_model_parallel_attributes( tensor=getattr(self, weight), - is_parallel=True, + # A weight is only tensor-model-parallel when the layer is. For + # parallel_mode=None the weight is replicated on every TP rank, and + # marking it parallel makes downstream consumers (e.g. Megatron's + # param_is_not_tensor_parallel_duplicate) admit it to the global + # gradient norm once per rank instead of once. + is_parallel=self.parallel_mode is not None, dim=1 if self.parallel_mode == "row" else 0, stride=1, )