From 0f268680c8ce84d3c9df4e3556e8a23598f981fa Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?G=C3=B6kdeniz=20G=C3=BClmez?= <60228478+Goekdeniz-Guelmez@users.noreply.github.com> Date: Thu, 4 Sep 2025 18:03:00 +0200 Subject: [PATCH] Fix Nemotron H loading error (#426) * fix * format --- mlx_lm/models/nemotron_h.py | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/mlx_lm/models/nemotron_h.py b/mlx_lm/models/nemotron_h.py index a7b3e06..12753b9 100644 --- a/mlx_lm/models/nemotron_h.py +++ b/mlx_lm/models/nemotron_h.py @@ -11,7 +11,7 @@ from .base import BaseModelArgs, create_attention_mask, scaled_dot_product_atten from .cache import KVCache, MambaCache -@dataclass(kw_only=True) +@dataclass() class ModelArgs(BaseModelArgs): model_type: str vocab_size: int @@ -35,6 +35,7 @@ class ModelArgs(BaseModelArgs): use_bias: bool use_conv_bias: bool residual_in_fp32: bool + head_dim: Optional[int] = None hybrid_override_pattern: Optional[List[str]] = None @@ -195,7 +196,11 @@ class NemotronHAttention(nn.Module): super().__init__() self.hidden_size = args.hidden_size self.num_heads = args.num_attention_heads - self.head_dim = self.hidden_size // self.num_heads + self.head_dim = ( + args.head_dim + if args.head_dim is not None + else (args.hidden_size // args.num_attention_heads) + ) self.num_key_value_heads = args.num_key_value_heads self.scale = self.head_dim**-0.5