Process the input tensor using the given state indexes and cache.
(
&self,
x: &Tensor,
index_pos: usize,
block_idx: usize,
cache: &mut super::Cache,
)
| 150 | |
| 151 | /// Process the input tensor using the given state indexes and cache. |
| 152 | pub fn forward( |
| 153 | &self, |
| 154 | x: &Tensor, |
| 155 | index_pos: usize, |
| 156 | block_idx: usize, |
| 157 | cache: &mut super::Cache, |
| 158 | ) -> anyhow::Result<Tensor> { |
| 159 | let (b_sz, seq_len, _hidden_size) = x.dims3().map_err(|e| anyhow!("x.dims3 -> {e}"))?; |
| 160 | |
| 161 | // Single fused QKV projection (routed through backend for GPU acceleration) |
| 162 | let qkv = self |
| 163 | .backend.linear_forward(x, &self.qkv_proj_weight, self.qkv_proj_bias.as_ref()) |
| 164 | .map_err(|e| anyhow!("qkv.forward -> {e}"))?; |
| 165 | |
| 166 | let q = qkv |
| 167 | .narrow(D::Minus1, 0, self.size_q) |
| 168 | .map_err(|e| anyhow!("q split -> {e}"))?; |
| 169 | let k = qkv |
| 170 | .narrow(D::Minus1, self.size_q, self.size_kv) |
| 171 | .map_err(|e| anyhow!("k split -> {e}"))?; |
| 172 | let v = qkv |
| 173 | .narrow(D::Minus1, self.size_q + self.size_kv, self.size_kv) |
| 174 | .map_err(|e| anyhow!("v split -> {e}"))?; |
| 175 | |
| 176 | // OLMo2-style: apply QK-norm BEFORE head reshape (norm dim = size_q/size_kv). |
| 177 | let (q, k) = if self.pre_reshape_qk_norm { |
| 178 | let q = if let Some(w) = &self.q_norm_weight { |
| 179 | self.backend.rms_norm(&q.contiguous() |
| 180 | .map_err(|e| anyhow!("pre_reshape q contiguous -> {e}"))?, w, self.qk_norm_eps) |
| 181 | .map_err(|e| anyhow!("pre_reshape q_norm -> {e}"))? |
| 182 | } else { q }; |
| 183 | let k = if let Some(w) = &self.k_norm_weight { |
| 184 | self.backend.rms_norm(&k.contiguous() |
| 185 | .map_err(|e| anyhow!("pre_reshape k contiguous -> {e}"))?, w, self.qk_norm_eps) |
| 186 | .map_err(|e| anyhow!("pre_reshape k_norm -> {e}"))? |
| 187 | } else { k }; |
| 188 | (q, k) |
| 189 | } else { |
| 190 | (q, k) |
| 191 | }; |
| 192 | |
| 193 | // Reshape: (b, seq, heads, head_dim) |
| 194 | let q = q |
| 195 | .reshape((b_sz, seq_len, self.num_attention_heads, self.head_dim))?; |
| 196 | let k = k |
| 197 | .reshape((b_sz, seq_len, self.num_key_value_heads, self.head_dim))?; |
| 198 | let v = v |
| 199 | .reshape((b_sz, seq_len, self.num_key_value_heads, self.head_dim))?; |
| 200 | |
| 201 | // Standard QK-norm: applied after reshape (on head_dim, last dim) before transpose. |
| 202 | let q = if !self.pre_reshape_qk_norm { |
| 203 | if let Some(w) = &self.q_norm_weight { |
| 204 | self.backend.rms_norm(&q.contiguous() |
| 205 | .map_err(|e| anyhow!("q contiguous -> {e}"))?, w, self.qk_norm_eps) |
| 206 | .map_err(|e| anyhow!("q_norm -> {e}"))? |
| 207 | } else { q } |
| 208 | } else { q }; |
| 209 | let k = if !self.pre_reshape_qk_norm { |
nothing calls this directly
no test coverage detected