diag eager gate fix
Browse files
app.py
CHANGED
|
@@ -480,7 +480,8 @@ def aoti_diag() -> str:
|
|
| 480 |
f"tiles {geometry.n_tiles}, topk {geometry.topk}", ""]
|
| 481 |
|
| 482 |
def eager(gate_h):
|
| 483 |
-
|
|
|
|
| 484 |
|
| 485 |
def functional(gate_h):
|
| 486 |
gate_t = None if gate_h is None else gate_h.transpose(1, 2)
|
|
|
|
| 480 |
f"tiles {geometry.n_tiles}, topk {geometry.topk}", ""]
|
| 481 |
|
| 482 |
def eager(gate_h):
|
| 483 |
+
# eager sparse_attention takes the gate as [B, S, H, D] and transposes internally
|
| 484 |
+
return vsa_h3.sparse_attention(q, k, v, gate_h, geometry)
|
| 485 |
|
| 486 |
def functional(gate_h):
|
| 487 |
gate_t = None if gate_h is None else gate_h.transpose(1, 2)
|