diff --git a/src/mcore_bridge/model/gpt_model.py b/src/mcore_bridge/model/gpt_model.py index 8ed9994..d6d1c10 100644 --- a/src/mcore_bridge/model/gpt_model.py +++ b/src/mcore_bridge/model/gpt_model.py @@ -300,7 +300,7 @@ def forward( full_attention_mask = full_attention_mask['full_attention'] if mcore_016 and full_attention_mask is not None: assert packed_seq_params is None - padding_mask = ~((~full_attention_mask).sum(dim=(1, 2)) > 0) + padding_mask = full_attention_mask.all(dim=(1, 2)) if self.config.context_parallel_size > 1: padding_mask = split_cp_inputs(padding_mask, None, 1) tp_size = self.config.tensor_model_parallel_size diff --git a/src/mcore_bridge/model/hybrid_model.py b/src/mcore_bridge/model/hybrid_model.py index 9e6424d..f9f9322 100644 --- a/src/mcore_bridge/model/hybrid_model.py +++ b/src/mcore_bridge/model/hybrid_model.py @@ -59,7 +59,7 @@ def _get_padding_mask(self, attention_mask) -> Optional[torch.Tensor]: attention_mask = attention_mask['full_attention'] if attention_mask is None: return None - padding_mask = ~((~attention_mask).sum(dim=(1, 2)) > 0) + padding_mask = attention_mask.all(dim=(1, 2)) if self.config.context_parallel_size > 1: padding_mask = split_cp_inputs(padding_mask, None, 1) tp_size = self.config.tensor_model_parallel_size