diff --git a/megatron/core/transformer/dot_product_attention.py b/megatron/core/transformer/dot_product_attention.py index 2a958722e46..2a6ac65a685 100644 --- a/megatron/core/transformer/dot_product_attention.py +++ b/megatron/core/transformer/dot_product_attention.py @@ -126,6 +126,8 @@ def __init__( ) ), ) + if config.perform_initialization: + self.softmax_offset = config.init_method(self.softmax_offset) else: raise ValueError("Softmax type not supported") diff --git a/megatron/core/transformer/moe/router.py b/megatron/core/transformer/moe/router.py index 068d680c798..7fa4692ef2f 100644 --- a/megatron/core/transformer/moe/router.py +++ b/megatron/core/transformer/moe/router.py @@ -66,6 +66,8 @@ def reset_parameters(self): """Reset the router parameters.""" if self.config.perform_initialization: self.config.init_method(self.weight) + if self.bias is not None: + self.config.init_method(self.bias) self.weight.data = self.weight.data.to(dtype=self.config.params_dtype) setattr(self.weight, 'sequence_parallel', self.config.sequence_parallel) if self.bias is not None: