broken attention

This commit is contained in:
2026-02-25 21:12:44 +08:00
parent e1b9c4b2b1
commit a23c3cd96c
+6 -5
View File
@@ -136,13 +136,14 @@ class TransformerBlock:
self.cache_kv[:, :, :, start_pos:start_pos+T, :].assign(Tensor.stack(k, v))
k = self.cache_kv[0, :, :, 0:start_pos+T, :]
v = self.cache_kv[1, :, :, 0:start_pos+T, :]
# NOTE: this mask is causal_lower_right, not the causal_upper_left generated by is_casual = True
mask = Tensor.full((1, 1, T, start_pos+T), float("-inf"), dtype=x.dtype, device=x.device).triu(int(start_pos)+1) if T > 1 else None
attn = q.scaled_dot_product_attention(k, v, attn_mask=mask, enable_gqa=True) # (B,H,T,Hd)
attn = attn.transpose(1, 2).reshape(B, T, -1) # back to (B,T,D)
attn = function(self.attn_output)(attn)
return x + attn
@function
def attention(q:Tensor, k:Tensor, v:Tensor, mask:Tensor) -> Tensor:
attn = q.scaled_dot_product_attention(k, v, attn_mask=mask, enable_gqa=True) # (B,H,T,Hd)
attn = attn.transpose(1, 2).reshape(B, T, -1) # back to (B,T,D)
return self.attn_output(attn)
return x + attention(q, k, v, mask)
@function
def _feed_forward(self, h: Tensor) -> Tensor: