Dequant-free Apple matmul; copy-free GQA decode attention
Browse files- README.md +17 -14
- bpdq.py +200 -3
- demo.ipynb +1 -1
README.md
CHANGED
|
@@ -97,14 +97,14 @@ on this machine.
|
|
| 97 |
|
| 98 |
| batch | decode | prefill |
|
| 99 |
|---|---|---|
|
| 100 |
-
| 1 |
|
| 101 |
-
| 4 |
|
| 102 |
-
| 8 |
|
| 103 |
-
| 16 |
|
| 104 |
|
| 105 |
-
Per matmul against dense fp16 of the same shape:
|
| 106 |
-
0.7–
|
| 107 |
-
|
| 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**
|
| 165 |
-
|
| 166 |
-
|
| 167 |
-
|
| 168 |
-
|
| 169 |
-
|
|
|
|
|
|
|
|
|
|
| 170 |
|
| 171 |
## Limits
|
| 172 |
|
| 173 |
- **TP = 1.** No tensor parallelism.
|
| 174 |
-
- On Apple, batches of
|
| 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
|
| 1214 |
-
out_f = codes.shape[0] * 16
|
| 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
|
| 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 | **
|
| 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 |
]
|