@@ -166,15 +166,15 @@ def forward(self, x, mask, angles, input_pos=None):
166166 x_flat = x .reshape (1 , - 1 ) # Shape: (1, d_in)
167167 input_dtype = x .dtype
168168
169- queries_flat = self .aie_query_gemv (None , x_flat )
169+ queries_flat = self .aie_query_gemv (x_flat )
170170 queries = queries_flat .reshape (b , num_tokens , self .d_out ).to (input_dtype )
171171
172- keys_flat = self .aie_key_gemv (None , x_flat )
172+ keys_flat = self .aie_key_gemv (x_flat )
173173 keys = keys_flat .reshape (
174174 b , num_tokens , self .num_kv_groups * self .head_dim
175175 ).to (input_dtype )
176176
177- values_flat = self .aie_value_gemv (None , x_flat )
177+ values_flat = self .aie_value_gemv (x_flat )
178178 values = values_flat .reshape (
179179 b , num_tokens , self .num_kv_groups * self .head_dim
180180 ).to (input_dtype )
@@ -384,7 +384,7 @@ def my_mha(queries, keys, values):
384384 # Choose output projection based on phase
385385 if self .cfg ["use_kv_cache" ] and is_decode and self .cfg ["use_aie_gemv" ]:
386386 context_vec_flat = context_vec .reshape (1 , - 1 )
387- output_flat = self .aie_out_proj_gemv (None , context_vec_flat )
387+ output_flat = self .aie_out_proj_gemv (context_vec_flat )
388388 context_vec = output_flat .reshape (b , num_tokens , self .d_out ).to (input_dtype )
389389 elif self .cfg ["use_aie_attn_projection_gemm" ]:
390390 context_vec_flat = context_vec .reshape (- 1 , self .d_out )
0 commit comments