gitarist commited on
Commit
655386b
·
verified ·
1 Parent(s): f72bfc4

Dequant-free Apple matmul; copy-free GQA decode attention

Browse files
Files changed (3) hide show
  1. README.md +17 -14
  2. bpdq.py +200 -3
  3. demo.ipynb +1 -1
README.md CHANGED
@@ -97,14 +97,14 @@ on this machine.
97
 
98
  | batch | decode | prefill |
99
  |---|---|---|
100
- | 1 | 28.6 tok/s | 166 tok/s |
101
- | 4 | 61.2 | 176 |
102
- | 8 | 74.4 | 175 |
103
- | 16 | 77.0 | 175 |
104
 
105
- Per matmul against dense fp16 of the same shape: 5.8–6.5x at batch 1, 1.3–1.7x at 16,
106
- 0.7–1.0x from 32 up. At batch 1 the matmuls read 81–99 GB/s; a device-to-device copy
107
- reaches 98 GB/s.
108
 
109
  ### Quality
110
 
@@ -161,17 +161,20 @@ While a forward runs at most 8 tokens under vLLM, a second stream loads each wei
161
  matrix into L2 two matmuls ahead of its use, paced by a progress flag the matmuls write.
162
  GPUs below compute capability 8.0 get a Triton kernel instead.
163
 
164
- **Apple** is three Metal kernels. Decode gives each output row several threads, each
165
- taking an interleaved share of the 256-column groups against a 4-entry value table; the
166
- partial sums meet in threadgroup memory. The 8-bit `lm_head` decodes the same way.
167
- Prefill feeds `simdgroup_half8x8` matrices from a weight slab rebuilt in threadgroup
168
- memory. None writes anything dense to device memory. The one-token decode step runs
169
- under `torch.compile`; the first step compiles for about 30 s.
 
 
 
170
 
171
  ## Limits
172
 
173
  - **TP = 1.** No tensor parallelism.
174
- - On Apple, batches of 32 tokens and up run at 0.68–1.0x of a dense fp16 matmul. The win
175
  is at decode batch sizes.
176
  - At 2 bits the model invents confident detail about anything obscure, and past roughly
177
  250 tokens long open-ended answers drift. Use it for short exchanges and bounded tasks.
 
97
 
98
  | batch | decode | prefill |
99
  |---|---|---|
100
+ | 1 | 31.3 tok/s | 169 tok/s |
101
+ | 4 | 110.6 | 173 |
102
+ | 8 | 144.1 | 168 |
103
+ | 16 | 137.7 | 168 |
104
 
105
+ Per matmul against dense fp16 of the same shape: 4.7–5.6x at 1 token, 3.2–3.9x at 8,
106
+ 1.1–1.4x at 32, 0.7–0.9x from 64 up. From 1 to 4 tokens a matmul takes at most 1.25x
107
+ the time to read its weights at the 98 GB/s a device-to-device copy reaches.
108
 
109
  ### Quality
110
 
 
161
  matrix into L2 two matmuls ahead of its use, paced by a progress flag the matmuls write.
162
  GPUs below compute capability 8.0 get a Triton kernel instead.
163
 
164
+ **Apple** multiplies without dequantising. Per 256-column group a row's output is
165
+ `c0 * S0 + c1 * S1 + bias * sum(x)`, where `S0` and `S1` are the sums of the activations
166
+ each bit-plane selects. For every 8 columns a threadgroup builds a 256-entry table of all
167
+ subset sums of those 8 activations, 4 tokens per entry, and one byte of a bit-plane
168
+ indexes it; the planes are the checkpoint's own, transposed at load. From 8 tokens on,
169
+ the kernel is bound by random reads of those tables. The 8-bit `lm_head` has its own
170
+ Metal kernel. One-token decode attention folds each KV head's query heads into one pass
171
+ instead of repeating the KV cache, and the decode step runs under `torch.compile`; the
172
+ first step compiles for about 30 s.
173
 
174
  ## Limits
175
 
176
  - **TP = 1.** No tensor parallelism.
177
+ - On Apple, batches of 64 tokens and up run at 0.7–0.9x of a dense fp16 matmul. The win
178
  is at decode batch sizes.
179
  - At 2 bits the model invents confident detail about anything obscure, and past roughly
180
  250 tokens long open-ended answers drift. Use it for short exchanges and bounded tasks.
bpdq.py CHANGED
@@ -1027,6 +1027,172 @@ def _matmul_mps8(x, codes, tbl, N, K):
1027
  return y[:T]
1028
 
1029
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1030
  def _select_expr(msbits, word, table):
1031
  """Nested ternary over the register table: msbits selects, no conversions."""
1032
  def rec(prefix, depth):
@@ -1138,6 +1304,10 @@ class PackedBPDQ:
1138
  if mma_supported(msbits, group_size, in_f, out_f, device):
1139
  codes, tbl = mma_repack(planes, coeffs, in_f, lut_bias, device)
1140
  return cls(codes, tbl, dead, in_f, out_f, msbits, coeffs.shape[2], group_size)
 
 
 
 
1141
  cpw = 32 // msbits
1142
  cwords = (in_f + cpw - 1) // cpw
1143
  codes = torch.zeros(cwords, out_f, dtype=torch.int32, device=device)
@@ -1195,6 +1365,8 @@ class PackedBPDQ:
1195
  return _mma_dense(self.codes, self.coeffs, self.in_f)
1196
  if self.msbits == 8:
1197
  return _gptq8_rows(self.codes, self.coeffs, 0, self.out_f).float()
 
 
1198
  return _dequant_rt(self.codes, self.coeffs, self.in_f, self.msbits,
1199
  self.n_coeff, self.group_size, self.cpw)
1200
 
@@ -1210,8 +1382,8 @@ class PackedBPDQ:
1210
  def packed_matmul(x, codes, coeffs, dead, in_f, msbits, n_coeff, group_size,
1211
  scratch=None):
1212
  """x [T, in_f] -> [T, out_f] in x.dtype, straight off the packed weight.
1213
- 4-D codes are the mma layout (CUDA); otherwise the row-major runtime layout."""
1214
- out_f = codes.shape[0] * 16 if codes.dim() == 4 else codes.shape[1]
1215
  cpw = 32 // msbits
1216
  wpg = group_size // cpw
1217
  lut_bias = n_coeff > msbits
@@ -1224,6 +1396,8 @@ def packed_matmul(x, codes, coeffs, dead, in_f, msbits, n_coeff, group_size,
1224
  y = torch.empty(x.shape[0], out_f, device=x.device, dtype=torch.float32)
1225
  _matmul_cuda(x.contiguous(), codes, coeffs, y, out_f, in_f, msbits,
1226
  n_coeff, lut_bias, cpw, wpg)
 
 
1227
  elif dev == "mps" and msbits == 8:
1228
  y = _matmul_mps8(x.contiguous().half(), codes, coeffs, out_f, in_f)
1229
  elif dev == "mps":
@@ -1242,6 +1416,24 @@ _PROJ = ("self_attn.q_proj", "self_attn.k_proj", "self_attn.v_proj", "self_attn.
1242
  "mlp.gate_proj", "mlp.up_proj", "mlp.down_proj")
1243
 
1244
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1245
  @torch.library.custom_op("bpdq::matmul", mutates_args=())
1246
  def _op_matmul(x: torch.Tensor, codes: torch.Tensor, coeffs: torch.Tensor,
1247
  dead: Optional[torch.Tensor], in_f: int, msbits: int,
@@ -1252,7 +1444,7 @@ def _op_matmul(x: torch.Tensor, codes: torch.Tensor, coeffs: torch.Tensor,
1252
 
1253
  @_op_matmul.register_fake
1254
  def _(x, codes, coeffs, dead, in_f, msbits, n_coeff, group_size):
1255
- out_f = codes.shape[0] * 16 if codes.dim() == 4 else codes.shape[1]
1256
  return x.new_empty((x.shape[0], out_f))
1257
 
1258
 
@@ -1372,6 +1564,11 @@ def load(path, device=None, dtype=torch.float16, verbose=False):
1372
 
1373
  model.to(device)
1374
  if torch.device(device).type == "mps":
 
 
 
 
 
1375
  model._bpdq_step = torch.compile(model.forward, dynamic=False)
1376
  tok = AutoTokenizer.from_pretrained(path)
1377
  return model, tok
 
1027
  return y[:T]
1028
 
1029
 
1030
+ _METAL_LUT_SRC = r"""
1031
+ #include <metal_stdlib>
1032
+ using namespace metal;
1033
+ typedef {VT} vt;
1034
+ typedef {FT} ft;
1035
+
1036
+ // 2-bit, dequant-free: y = sum_g c0 * S0 + c1 * S1 + bias * Sx, where S0 / S1 are the
1037
+ // activation sums selected by each bit-plane. Per 8 columns, a 256-entry table holds
1038
+ // every subset sum of the 8 activations for V tokens; one byte of a plane indexes it.
1039
+ // planes [K/32, 2, M]: bit i of a word is column 32w + i. coeffs [G, 3, M] = (c0, c1, bias).
1040
+ // A thread owns R rows; tables are double-buffered per KC columns; gridDim.y splits K.
1041
+ kernel void bpdq_lut(
1042
+ device float* part [[buffer(0)]], // [split, TT, M]
1043
+ device const uint* planes [[buffer(1)]],
1044
+ device const float* coeffs [[buffer(2)]],
1045
+ device const half* x [[buffer(3)]], // [T, K]
1046
+ constant uint& M [[buffer(4)]],
1047
+ constant uint& K [[buffer(5)]],
1048
+ constant uint& KS [[buffer(6)]], // columns per split
1049
+ constant uint& T [[buffer(7)]], // real tokens; the rest read as zero
1050
+ uint2 tid2 [[thread_position_in_threadgroup]],
1051
+ uint2 gid [[threadgroup_position_in_grid]])
1052
+ {{
1053
+ const uint tid = tid2.x;
1054
+ const uint TT = {TT}u, V = {V}u, NV = {TT}u / {V}u, TG = {TG}u, KC = {KC}u, R = {R}u;
1055
+ const uint NTB = KC / 8u, TS = NV * NTB * 256u;
1056
+ threadgroup vt tab2[2 * ({TT} / {V}) * ({KC} / 8) * 256];
1057
+
1058
+ uint mr[{R}];
1059
+ for (uint r = 0; r < R; ++r) mr[r] = min(gid.x * TG * R + r * TG + tid, M - 1u);
1060
+ const uint k0 = gid.y * KS, k1 = min(k0 + KS, K);
1061
+ ft acc[{R}][{TT} / {V}];
1062
+ vt s0[{R}][{TT} / {V}], s1[{R}][{TT} / {V}], sx[{TT} / {V}];
1063
+ for (uint r = 0; r < R; ++r)
1064
+ for (uint n = 0; n < NV; ++n) {{ acc[r][n] = ft(0); s0[r][n] = vt(0); s1[r][n] = vt(0); }}
1065
+ for (uint n = 0; n < NV; ++n) sx[n] = vt(0);
1066
+
1067
+ // one thread per (token group, table, high nibble): 16 entries, one add each
1068
+ auto build = [&](uint kc, threadgroup vt* tb) {{
1069
+ for (uint e = tid; e < NV * NTB * 16u; e += TG) {{
1070
+ const uint n = e / (NTB * 16u), b = (e / 16u) % NTB, hi = e % 16u;
1071
+ vt xv[8];
1072
+ for (uint u = 0; u < V; ++u) {{
1073
+ const uint t = min(n * V + u, T - 1u);
1074
+ const device half4* xp = (const device half4*)(x + t * K + kc + b * 8u);
1075
+ const half4 a = n * V + u < T ? xp[0] : half4(0), c = n * V + u < T ? xp[1] : half4(0);
1076
+ xv[0][u] = a[0]; xv[1][u] = a[1]; xv[2][u] = a[2]; xv[3][u] = a[3];
1077
+ xv[4][u] = c[0]; xv[5][u] = c[1]; xv[6][u] = c[2]; xv[7][u] = c[3];
1078
+ }}
1079
+ vt cur = (hi & 1u ? xv[4] : vt(0)) + (hi & 2u ? xv[5] : vt(0))
1080
+ + (hi & 4u ? xv[6] : vt(0)) + (hi & 8u ? xv[7] : vt(0));
1081
+ threadgroup vt* tp = tb + (n * NTB + b) * 256u + hi * 16u;
1082
+ tp[0] = cur;
1083
+ cur = cur + xv[0]; tp[1] = cur;
1084
+ cur = cur + xv[1]; tp[3] = cur;
1085
+ cur = cur - xv[0]; tp[2] = cur;
1086
+ cur = cur + xv[2]; tp[6] = cur;
1087
+ cur = cur + xv[0]; tp[7] = cur;
1088
+ cur = cur - xv[1]; tp[5] = cur;
1089
+ cur = cur - xv[0]; tp[4] = cur;
1090
+ cur = cur + xv[3]; tp[12] = cur;
1091
+ cur = cur + xv[0]; tp[13] = cur;
1092
+ cur = cur + xv[1]; tp[15] = cur;
1093
+ cur = cur - xv[0]; tp[14] = cur;
1094
+ cur = cur - xv[2]; tp[10] = cur;
1095
+ cur = cur + xv[0]; tp[11] = cur;
1096
+ cur = cur - xv[1]; tp[9] = cur;
1097
+ cur = cur - xv[0]; tp[8] = cur;
1098
+ }}
1099
+ }};
1100
+
1101
+ build(k0, tab2);
1102
+ uint buf = 0;
1103
+ for (uint kc = k0; kc < k1; kc += KC) {{
1104
+ threadgroup_barrier(mem_flags::mem_threadgroup);
1105
+ const threadgroup vt* tab = tab2 + buf * TS;
1106
+ if (kc + KC < k1) build(kc + KC, tab2 + (buf ^ 1u) * TS);
1107
+ buf ^= 1u;
1108
+ #pragma clang loop unroll(full)
1109
+ for (uint wq = 0; wq < KC / 32u; ++wq) {{
1110
+ const uint w = kc / 32u + wq;
1111
+ #pragma clang loop unroll(full)
1112
+ for (uint j = 0; j < 4u; ++j)
1113
+ for (uint n = 0; n < NV; ++n) sx[n] += tab[(n * NTB + wq * 4u + j) * 256u + 255u];
1114
+ #pragma clang loop unroll(full)
1115
+ for (uint r = 0; r < R; ++r) {{
1116
+ const uchar4 p0 = as_type<uchar4>(planes[(w * 2u) * M + mr[r]]);
1117
+ const uchar4 p1 = as_type<uchar4>(planes[(w * 2u + 1u) * M + mr[r]]);
1118
+ #pragma clang loop unroll(full)
1119
+ for (uint j = 0; j < 4u; ++j) {{
1120
+ #pragma clang loop unroll(full)
1121
+ for (uint n = 0; n < NV; ++n) {{
1122
+ const threadgroup vt* tp = tab + (n * NTB + wq * 4u + j) * 256u;
1123
+ s0[r][n] += tp[p0[j]];
1124
+ s1[r][n] += tp[p1[j]];
1125
+ }}
1126
+ }}
1127
+ }}
1128
+ }}
1129
+ if (((kc + KC) & 255u) == 0u || kc + KC >= k1) {{ // end of a 256-column group
1130
+ const uint g = kc / 256u;
1131
+ for (uint r = 0; r < R; ++r) {{
1132
+ const uint m = mr[r];
1133
+ const float c0 = coeffs[(g * 3u) * M + m], c1 = coeffs[(g * 3u + 1u) * M + m];
1134
+ const float cb = coeffs[(g * 3u + 2u) * M + m];
1135
+ for (uint n = 0; n < NV; ++n) {{
1136
+ acc[r][n] += c0 * ft(s0[r][n]) + c1 * ft(s1[r][n]) + cb * ft(sx[n]);
1137
+ s0[r][n] = vt(0); s1[r][n] = vt(0);
1138
+ }}
1139
+ }}
1140
+ for (uint n = 0; n < NV; ++n) sx[n] = vt(0);
1141
+ }}
1142
+ }}
1143
+ for (uint r = 0; r < R; ++r) {{
1144
+ const uint row = gid.x * TG * R + r * TG + tid;
1145
+ if (row < M)
1146
+ for (uint t = 0; t < min(TT, T); ++t) part[(gid.y * T + t) * M + row] = acc[r][t / V][t % V];
1147
+ }}
1148
+ }}
1149
+ """
1150
+
1151
+
1152
+ @functools.lru_cache(maxsize=32)
1153
+ def _metal_lut_lib(tt, tg, kc, r):
1154
+ v = min(tt, 4)
1155
+ vt, ft = {2: ("half2", "float2"), 4: ("half4", "float4")}[v]
1156
+ return torch.mps.compile_shader(_METAL_LUT_SRC.format(VT=vt, FT=ft, TT=tt, V=v, TG=tg, KC=kc, R=r))
1157
+
1158
+
1159
+ def mps_lut_supported(msbits, lut_bias, group_size, in_f, device):
1160
+ return (torch.device(device).type == "mps" and msbits == 2 and lut_bias
1161
+ and group_size == 256 and in_f % 256 == 0)
1162
+
1163
+
1164
+ def mps_lut_config(T, M, K):
1165
+ """(padded tokens, threadgroup, KC, rows per thread, split) for T <= 8."""
1166
+ tt = max(2, _next_pow2(T))
1167
+ tg, r = (128, 4) if tt <= 4 else (256, 4)
1168
+ groups, split = -(-M // (tg * r)), 1
1169
+ while groups * split < 32 and split < 16 and K % (split * 512) == 0:
1170
+ split *= 2
1171
+ return tt, tg, 32, r, split
1172
+
1173
+
1174
+ def _matmul_mps_lut(x, planes, coeffs, M, K):
1175
+ T = x.shape[0]
1176
+ if T > 8:
1177
+ return torch.cat([_matmul_mps_lut(c, planes, coeffs, M, K) for c in x.split(8)], 0)
1178
+ tt, tg, kc, r, split = mps_lut_config(T, M, K)
1179
+ if x.storage_offset() % 4: # half4 loads
1180
+ x = x.clone()
1181
+ part = torch.empty(split, T, M, device=x.device, dtype=torch.float32)
1182
+ _metal_lut_lib(tt, tg, kc, r).bpdq_lut(
1183
+ part, planes, coeffs, x, M, K, K // split, T,
1184
+ threads=(-(-M // (tg * r)) * tg, split), group_size=(tg, 1))
1185
+ return (part[0] if split == 1 else part.sum(0)).half()
1186
+
1187
+
1188
+ def _planes_dense(planes, coeffs, in_f):
1189
+ """[K/32, 2, M] planes + [G, 3, M] coeffs -> dense fp32 [M, K]."""
1190
+ sh = torch.arange(32, device=planes.device, dtype=torch.int32)
1191
+ bits = ((planes[..., None] >> sh) & 1).permute(1, 2, 0, 3).reshape(2, planes.shape[2], -1)
1192
+ c = coeffs.repeat_interleave(256, 0)[:in_f] # [K, 3, M]
1193
+ return (c[:, 2].T + bits[0] * c[:, 0].T + bits[1] * c[:, 1].T).float()
1194
+
1195
+
1196
  def _select_expr(msbits, word, table):
1197
  """Nested ternary over the register table: msbits selects, no conversions."""
1198
  def rec(prefix, depth):
 
1304
  if mma_supported(msbits, group_size, in_f, out_f, device):
1305
  codes, tbl = mma_repack(planes, coeffs, in_f, lut_bias, device)
1306
  return cls(codes, tbl, dead, in_f, out_f, msbits, coeffs.shape[2], group_size)
1307
+ if mps_lut_supported(msbits, lut_bias, group_size, in_f, device):
1308
+ codes = planes.permute(2, 0, 1).contiguous().to(device) # [K/32, 2, M]
1309
+ rt = coeffs.to(device).float().permute(0, 2, 1).contiguous()
1310
+ return cls(codes, rt, dead, in_f, out_f, msbits, coeffs.shape[2], group_size)
1311
  cpw = 32 // msbits
1312
  cwords = (in_f + cpw - 1) // cpw
1313
  codes = torch.zeros(cwords, out_f, dtype=torch.int32, device=device)
 
1365
  return _mma_dense(self.codes, self.coeffs, self.in_f)
1366
  if self.msbits == 8:
1367
  return _gptq8_rows(self.codes, self.coeffs, 0, self.out_f).float()
1368
+ if self.codes.dim() == 3:
1369
+ return _planes_dense(self.codes, self.coeffs, self.in_f)
1370
  return _dequant_rt(self.codes, self.coeffs, self.in_f, self.msbits,
1371
  self.n_coeff, self.group_size, self.cpw)
1372
 
 
1382
  def packed_matmul(x, codes, coeffs, dead, in_f, msbits, n_coeff, group_size,
1383
  scratch=None):
1384
  """x [T, in_f] -> [T, out_f] in x.dtype, straight off the packed weight.
1385
+ 4-D codes: mma layout (CUDA). 3-D: bit-planes (Apple, 2-bit). 2-D: row-major runtime layout."""
1386
+ out_f = {4: codes.shape[0] * 16, 3: codes.shape[-1]}.get(codes.dim(), codes.shape[1])
1387
  cpw = 32 // msbits
1388
  wpg = group_size // cpw
1389
  lut_bias = n_coeff > msbits
 
1396
  y = torch.empty(x.shape[0], out_f, device=x.device, dtype=torch.float32)
1397
  _matmul_cuda(x.contiguous(), codes, coeffs, y, out_f, in_f, msbits,
1398
  n_coeff, lut_bias, cpw, wpg)
1399
+ elif dev == "mps" and codes.dim() == 3:
1400
+ y = _matmul_mps_lut(x.contiguous().half(), codes, coeffs, out_f, in_f)
1401
  elif dev == "mps" and msbits == 8:
1402
  y = _matmul_mps8(x.contiguous().half(), codes, coeffs, out_f, in_f)
1403
  elif dev == "mps":
 
1416
  "mlp.gate_proj", "mlp.up_proj", "mlp.down_proj")
1417
 
1418
 
1419
+ def _gqa_decode_attention(module, query, key, value, attention_mask, dropout=0.0,
1420
+ scaling=None, is_causal=None, **kw):
1421
+ """SDPA; one-token queries fold each KV head's query heads into the query axis
1422
+ instead of repeating the KV cache."""
1423
+ from transformers.integrations.sdpa_attention import sdpa_attention_forward
1424
+ B, H, Tq, D = query.shape
1425
+ Hkv = key.shape[1]
1426
+ if Tq != 1 or H == Hkv:
1427
+ return sdpa_attention_forward(module, query, key, value, attention_mask, dropout=dropout,
1428
+ scaling=scaling, is_causal=is_causal, **kw)
1429
+ mask = attention_mask
1430
+ if mask is not None:
1431
+ mask = mask[..., : key.shape[-2]].expand(B, 1, H // Hkv, key.shape[-2])
1432
+ out = torch.nn.functional.scaled_dot_product_attention(
1433
+ query.view(B, Hkv, H // Hkv, D), key, value, attn_mask=mask, scale=scaling)
1434
+ return out.view(B, H, 1, D).transpose(1, 2).contiguous(), None
1435
+
1436
+
1437
  @torch.library.custom_op("bpdq::matmul", mutates_args=())
1438
  def _op_matmul(x: torch.Tensor, codes: torch.Tensor, coeffs: torch.Tensor,
1439
  dead: Optional[torch.Tensor], in_f: int, msbits: int,
 
1444
 
1445
  @_op_matmul.register_fake
1446
  def _(x, codes, coeffs, dead, in_f, msbits, n_coeff, group_size):
1447
+ out_f = {4: codes.shape[0] * 16, 3: codes.shape[-1]}.get(codes.dim(), codes.shape[1])
1448
  return x.new_empty((x.shape[0], out_f))
1449
 
1450
 
 
1564
 
1565
  model.to(device)
1566
  if torch.device(device).type == "mps":
1567
+ from transformers import AttentionInterface
1568
+ from transformers.masking_utils import AttentionMaskInterface, sdpa_mask
1569
+ AttentionInterface.register("bpdq_gqa", _gqa_decode_attention)
1570
+ AttentionMaskInterface.register("bpdq_gqa", sdpa_mask)
1571
+ model.set_attn_implementation("bpdq_gqa")
1572
  model._bpdq_step = torch.compile(model.forward, dynamic=False)
1573
  tok = AutoTokenizer.from_pretrained(path)
1574
  return model, tok
demo.ipynb CHANGED
@@ -13,7 +13,7 @@
13
  "| weights resident | **3.65 GiB** | 15.3 GiB |\n",
14
  "| 1 sequence, RTX 4090 | **322 tok/s** | 59 |\n",
15
  "| 16 sequences, RTX 4090 | **3452 tok/s** | 890 |\n",
16
- "| Apple M4, 1 sequence | **29 tok/s** | *does not fit* |\n",
17
  "\n",
18
  "The notebook works out which machine you are on and picks the kernels."
19
  ]
 
13
  "| weights resident | **3.65 GiB** | 15.3 GiB |\n",
14
  "| 1 sequence, RTX 4090 | **322 tok/s** | 59 |\n",
15
  "| 16 sequences, RTX 4090 | **3452 tok/s** | 890 |\n",
16
+ "| Apple M4, 1 sequence | **31 tok/s** | *does not fit* |\n",
17
  "\n",
18
  "The notebook works out which machine you are on and picks the kernels."
19
  ]