MCPcopy Create free account
hub / github.com/evilsocket/cake / forward

Method forward

cake-core/src/models/common/attention.rs:152–357  ·  view source on GitHub ↗

Process the input tensor using the given state indexes and cache.

(
        &self,
        x: &Tensor,
        index_pos: usize,
        block_idx: usize,
        cache: &mut super::Cache,
    )

Source from the content-addressed store, hash-verified

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 {

Callers

nothing calls this directly

Calls 15

flash_attentionFunction · 0.85
apply_rotary_embMethod · 0.80
process_kv_windowedMethod · 0.80
process_kvMethod · 0.80
dtypeMethod · 0.80
maskMethod · 0.80
shapeMethod · 0.80
masked_fillFunction · 0.70
linear_forwardMethod · 0.45
rms_normMethod · 0.45
sdpaMethod · 0.45
repeat_kvMethod · 0.45

Tested by

no test coverage detected