File size: 4,508 Bytes
b66f552 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 | # Copyright (c) 2023-2025, Songlin Yang, Yu Zhang
import math
import torch
def tril_softmax(scores: torch.Tensor, strict: bool = True) -> torch.Tensor:
"""
Row-wise causal softmax over strictly lower-triangular (j < i) positions.
Args:
scores: [B, H, T, T] raw attention scores (q @ k^T).
strict: if True, mask out diagonal as well (strictly causal). Otherwise include diagonal.
Returns:
probs: [B, H, T, T] with probabilities on j < i (or j <= i if strict=False), zeros elsewhere.
"""
T = scores.size(-1)
device = scores.device
i = torch.arange(T, device=device).view(1, 1, T, 1)
j = torch.arange(T, device=device).view(1, 1, 1, T)
if strict:
mask = (j < i)
else:
mask = (j <= i)
masked = scores.masked_fill(~mask, float('-inf'))
max_per_row = masked.max(dim=-1, keepdim=True).values
exp = (masked - max_per_row).exp()
exp = exp.masked_fill(~mask, 0.0)
denom = exp.sum(dim=-1, keepdim=True).clamp_min_(1e-20)
probs = exp / denom
return probs
def naive_causal_attention_bhtd(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
) -> torch.Tensor:
B, H, T, D = q.shape
qk_scale = 1.0 / math.sqrt(D)
scores = torch.matmul(q, k.transpose(-1, -2)) * qk_scale # [B, H, T, T]
causal_mask = torch.triu(torch.ones(T, T, device=q.device), diagonal=1).bool()
scores = scores.masked_fill(causal_mask, float('-inf'))
attn_weights = torch.softmax(scores, dim=-1) # [B, H, T, T]
o = torch.matmul(attn_weights, v) # [B, H, T, D]
return o
def naive_deltaformer_attn_head_first(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor | None = None,
) -> torch.Tensor:
"""
Naive reference implementation of DeltaFormer attention for head-first format.
Two-stage process:
1. Computes u[i] = v[i] - beta[i] * sum_{j<i} softmax(q[i] @ k[:i]^T) @ u[:i]
2. Applies causal attention: o = causal_attn(q, k, u)
Args:
q: [B, H, T, D]
k: [B, H, T, D]
v: [B, H, T, D]
beta: [B, H, T] or None (defaults to ones)
Returns:
o: [B, H, T, D]
"""
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "q,k,v must be [B,H,T,D]"
B, H, T, D = q.shape
assert k.shape == (B, H, T, D) and v.shape == (B, H, T, D)
orig_dtype = q.dtype
qf = q.float()
kf = k.float()
vf = v.float()
if beta is None:
betaf = torch.ones((B, H, T), dtype=torch.float32, device=q.device)
else:
assert beta.shape == (B, H, T)
betaf = beta.float()
qk_scale = 1.0 / math.sqrt(D)
scores = torch.matmul(qf, kf.transpose(-1, -2)) * qk_scale
probs = tril_softmax(scores, strict=True) # [B,H,T,T] float32
u_list = []
for t in range(T):
if t == 0:
u_t = vf[:, :, t, :]
else:
w = probs[:, :, t, :t] # [B,H,t]
u_prev = torch.stack(u_list, dim=-2) # [B,H,t,D]
weighted_sum = (w.unsqueeze(-1) * u_prev).sum(dim=-2) # [B,H,D]
u_t = vf[:, :, t, :] - betaf[:, :, t].unsqueeze(-1) * weighted_sum
u_list.append(u_t)
u = torch.stack(u_list, dim=2) # [B,H,T,D]
o = naive_causal_attention_bhtd(q, k, u.to(orig_dtype))
return o.to(orig_dtype)
def naive_deltaformer_attn(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
beta: torch.Tensor | None = None,
) -> torch.Tensor:
"""
Naive reference implementation of DeltaFormer attention for sequence-first format.
Args:
q: [B, T, H, D]
k: [B, T, H, D]
v: [B, T, H, D]
beta: [B, T, H] or None (defaults to ones)
Returns:
o: [B, T, H, D]
"""
assert q.dim() == 4 and k.dim() == 4 and v.dim() == 4, "q,k,v must be [B,T,H,D]"
B, T, H, D = q.shape
assert k.shape == (B, T, H, D) and v.shape == (B, T, H, D)
q_bhtd = q.transpose(1, 2) # [B, T, H, D] -> [B, H, T, D]
k_bhtd = k.transpose(1, 2) # [B, T, H, D] -> [B, H, T, D]
v_bhtd = v.transpose(1, 2) # [B, T, H, D] -> [B, H, T, D]
if beta is not None:
assert beta.shape == (B, T, H)
beta_bhtd = beta.transpose(1, 2) # [B, T, H] -> [B, H, T]
else:
beta_bhtd = None
o_bhtd = naive_deltaformer_attn_head_first(q_bhtd, k_bhtd, v_bhtd, beta_bhtd)
o_bthd = o_bhtd.transpose(1, 2) # [B, H, T, D] -> [B, T, H, D]
return o_bthd
__all__ = [
'naive_deltaformer_attn',
'tril_softmax',
]
|