From 46bb9632b5a6c47abd17ea35fea77c1215a0e495 Mon Sep 17 00:00:00 2001 From: Achuth Reddy Date: Thu, 9 Jul 2026 10:23:12 -0500 Subject: [PATCH] fix: respect explicit config.head_dim instead of assuming hidden_size/num_heads (fixes #315) --- eagle/model/cnets.py | 13 +++++-------- 1 file changed, 5 insertions(+), 8 deletions(-) diff --git a/eagle/model/cnets.py b/eagle/model/cnets.py index a8e13ca1..c3ec7478 100644 --- a/eagle/model/cnets.py +++ b/eagle/model/cnets.py @@ -196,16 +196,13 @@ def __init__(self, config): self.config = config self.hidden_size = config.hidden_size self.num_heads = config.num_attention_heads - self.head_dim = self.hidden_size // self.num_heads + if hasattr(config, "head_dim") and config.head_dim is not None: + self.head_dim = config.head_dim + else: + self.head_dim = self.hidden_size // self.num_heads self.num_key_value_heads = config.num_key_value_heads self.num_key_value_groups = self.num_heads // self.num_key_value_heads self.max_position_embeddings = config.max_position_embeddings - - if (self.head_dim * self.num_heads) != self.hidden_size: - raise ValueError( - f"hidden_size must be divisible by num_heads (got `hidden_size`: {self.hidden_size}" - f" and `num_heads`: {self.num_heads})." - ) self.q_proj = nn.Linear(self.hidden_size * 2, self.num_heads * self.head_dim, bias=False) self.k_proj = nn.Linear(self.hidden_size * 2, self.num_key_value_heads * self.head_dim, bias=False) self.v_proj = nn.Linear(self.hidden_size * 2, self.num_key_value_heads * self.head_dim, bias=False) @@ -318,7 +315,7 @@ def forward( ) attn_output = attn_output.transpose(1, 2).contiguous() - attn_output = attn_output.reshape(bsz, q_len, self.hidden_size) + attn_output = attn_output.reshape(bsz, q_len, self.num_heads * self.head_dim) if self.config.pretraining_tp > 1: attn_output = attn_output.split(self.hidden_size // self.config.pretraining_tp, dim=2)