| |
| |
| |
| |
| import * as WG from './wgsl.js'; |
|
|
| const ROLES_NEEDED = new Set([ |
| 'wqkv', 'wo', 'w13', 'w2', 'lm_head', 'gate_bias', 'final_norm', 'rope_golden', |
| 'embed', 'img_proj', |
| ]); |
| |
| const HEAD_ROLES = { |
| 'coord_dec.w1': 'coord_decoder.w1', 'coord_dec.w2': 'coord_decoder.w2', |
| 'size_dec.w1': 'size_decoder.w1', 'size_dec.w2': 'size_decoder.w2', |
| 'coord_enc.embed': 'coord_encoder.embed', 'coord_enc.transform': 'coord_encoder.transform', |
| 'size_enc.embed': 'size_encoder.embed', 'size_enc.transform': 'size_encoder.transform', |
| }; |
|
|
| function u32(...vals) { return new Uint32Array(vals); } |
| function f16ToF32(u16) { |
| const out = new Float32Array(u16.length); |
| for (let i = 0; i < u16.length; i++) { |
| const h = u16[i], sg = (h & 0x8000) ? -1 : 1, e = (h >> 10) & 0x1F, m = h & 0x3FF; |
| out[i] = e === 0 ? sg * m * 2 ** -24 |
| : e === 31 ? (m ? NaN : sg * Infinity) |
| : sg * (1 + m / 1024) * 2 ** (e - 15); |
| } |
| return out; |
| } |
| function f32u32(f) { return new Float32Array([f])[0]; } |
|
|
| export class Engine { |
| static async create(device, weightsBuf, |
| { maxS = 640, maxT = 1024, maxK = 32, headsBuf = null, sgmat = 'auto' } = {}) { |
| const e = new Engine(); |
| e.device = device; |
| e.maxS = maxS; |
| e.maxT = maxT; |
| e.maxK = maxK; |
| e._parseWeights(weightsBuf); |
| if (headsBuf && !e.heads) e._parseHeads(headsBuf); |
| await e._buildPipelines(sgmat !== false); |
| e._allocBuffers(); |
| |
| e.useSgmat = false; |
| if (sgmat !== false && e.p.gemmSg8) { |
| const err = await e._smokeSgmat(); |
| e.useSgmat = err < 5e-2; |
| console.log(`[engine] sgmat smoke: normErr=${err.toExponential(2)} -> ${e.useSgmat ? 'ВКЛ' : 'ВЫКЛ'}`); |
| } |
| if (sgmat === true && !e.useSgmat) console.warn('[engine] sgmat запрошен, но недоступен/не прошёл smoke'); |
| e._buildBindGroups(); |
| return e; |
| } |
|
|
| |
| async _smokeSgmat() { |
| const d = this.device, M = 35, N = 64, K = 256, nb = K / 128; |
| const rnd = (n, f) => { const a = new f(n); let s = 777; |
| for (let i = 0; i < n; i++) { s = (Math.imul(s, 1664525) + 1013904223) >>> 0; |
| a[i] = f === Float32Array ? ((s / 4294967296) - 0.5) : (s & 0xFF); } return a; }; |
| const up = (data, usage = GPUBufferUsage.STORAGE) => { |
| const b = d.createBuffer({ size: Math.ceil(data.byteLength / 4) * 4, |
| usage: usage | GPUBufferUsage.COPY_DST }); |
| d.queue.writeBuffer(b, 0, data.buffer ? data.buffer : data); return b; |
| }; |
| const packed = up(rnd(N * K, Uint8Array)); |
| const scaleBits = new Uint16Array(N * nb); |
| for (let i = 0; i < scaleBits.length; i++) scaleBits[i] = 0x2400 + (i % 7); |
| const scales = up(scaleBits); |
| const zeros = up(rnd(N * nb, Uint8Array)); |
| const x = up(rnd(M * K, Float32Array)); |
| const mk = () => d.createBuffer({ size: M * N * 4, |
| usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC }); |
| const y1 = mk(), y2 = mk(); |
| const params = up(new Uint32Array([M, N, K, 0]), GPUBufferUsage.UNIFORM); |
| const run = (pipe, y, gx, gy) => { |
| const bg = this._bg(pipe, [params, packed, scales, zeros, x, y]); |
| const enc = d.createCommandEncoder(); |
| this._pass(enc, pipe, bg, gx, gy); |
| d.queue.submit([enc.finish()]); |
| }; |
| run(this.p.gemm8, y1, Math.ceil(N / 32), Math.ceil(M / 32)); |
| run(this.p.gemmSg8, y2, Math.ceil(N / 64), Math.ceil(M / 32)); |
| const [a, b] = await Promise.all([this.readF32(y1, M * N), this.readF32(y2, M * N)]); |
| let maxAbs = 0, maxRef = 1e-9; |
| for (let i = 0; i < M * N; i++) { |
| maxAbs = Math.max(maxAbs, Math.abs(a[i] - b[i])); |
| maxRef = Math.max(maxRef, Math.abs(a[i])); |
| } |
| for (const buf of [packed, scales, zeros, x, y1, y2, params]) buf.destroy(); |
| return maxAbs / maxRef; |
| } |
|
|
| |
| _parseHeads(buf) { |
| const dv = new DataView(buf); |
| if (dv.getUint32(0, true) !== 0x34515046) throw new Error('heads: не FPQ4'); |
| const mlen = dv.getUint32(4, true); |
| const man = JSON.parse(new TextDecoder().decode(new Uint8Array(buf, 8, mlen))); |
| this.heads = {}; |
| for (const t of man.tensors) { |
| const d = t.data; |
| if (t.dtype === 'f16') { |
| const b = this.device.createBuffer({ |
| size: d.length, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST }); |
| this.device.queue.writeBuffer(b, 0, buf, d.offset, d.length); |
| this.heads[t.name] = { buffer: b, shape: t.shape }; |
| } else { |
| this.heads[t.name] = { cpu: new Float32Array(buf.slice(d.offset, d.offset + d.length)), |
| shape: t.shape }; |
| } |
| } |
| } |
|
|
| |
|
|
| _parseWeights(buf) { |
| const dv = new DataView(buf); |
| if (dv.getUint32(0, true) !== 0x34515046) throw new Error('не FPQ4'); |
| const mlen = dv.getUint32(4, true); |
| const man = JSON.parse(new TextDecoder().decode(new Uint8Array(buf, 8, mlen))); |
| this.cfg = man.model; |
| const upload = (o, l) => { |
| const b = this.device.createBuffer({ |
| size: Math.ceil(l / 4) * 4, |
| usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST, |
| }); |
| this.device.queue.writeBuffer(b, 0, buf, o, l); |
| return b; |
| }; |
| this.w = { layers: Array.from({ length: 28 }, () => ({})) }; |
| const gateBias = new Float32Array(28 * 16); |
| for (const t of man.tensors) { |
| if (HEAD_ROLES[t.role]) { |
| this.heads = this.heads || {}; |
| const d = t.data; |
| const isEnc = t.role.includes('_enc.'); |
| if (isEnc) { |
| |
| const cpu = t.dtype === 'f16' |
| ? f16ToF32(new Uint16Array(buf.slice(d.offset, d.offset + d.length))) |
| : new Float32Array(buf.slice(d.offset, d.offset + d.length)); |
| this.heads[HEAD_ROLES[t.role]] = { cpu, shape: t.shape }; |
| } else { |
| this.heads[HEAD_ROLES[t.role]] = { buffer: upload(d.offset, d.length), shape: t.shape }; |
| } |
| continue; |
| } |
| if (!ROLES_NEEDED.has(t.role)) continue; |
| if (t.dtype === 'q4' || t.dtype === 'q8') { |
| const m = { |
| n: t.shape[0], k: t.shape[1], |
| bits: (t.quant && t.quant.bits) || 4, |
| packed: upload(t.packed.offset, t.packed.length), |
| scales: upload(t.scales.offset, t.scales.length), |
| zeros: upload(t.zeros.offset, t.zeros.length), |
| }; |
| if (t.role === 'lm_head') this.w.lm_head = m; |
| else this.w.layers[t.layer][t.role] = m; |
| } else if (t.role === 'gate_bias') { |
| gateBias.set(new Float32Array(buf, t.data.offset, 16), t.layer * 16); |
| } else if (t.role === 'final_norm') { |
| this.w.final_norm = upload(t.data.offset, t.data.length); |
| } else if (t.role === 'embed') { |
| this.w.embed = upload(t.data.offset, t.data.length); |
| } else if (t.role === 'img_proj') { |
| this.w.img_proj = upload(t.data.offset, t.data.length); |
| } else if (t.role === 'rope_golden') { |
| this.golden = new Float32Array(buf.slice(t.data.offset, t.data.offset + t.data.length)); |
| } |
| } |
| this.w.gate_bias = this.device.createBuffer({ |
| size: gateBias.byteLength, usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST }); |
| this.device.queue.writeBuffer(this.w.gate_bias, 0, gateBias); |
| } |
|
|
| |
|
|
| async _pipe(key, code) { |
| const device = this.device; |
| device.pushErrorScope('validation'); |
| const module = device.createShaderModule({ code, label: key }); |
| const info = await module.getCompilationInfo(); |
| const errs = info.messages.filter((m) => m.type === 'error'); |
| if (errs.length) { |
| await device.popErrorScope().catch(() => null); |
| throw new Error(`${key}: ${errs.map((m) => `${m.lineNum}:${m.linePos} ${m.message}`).join('; ')}`); |
| } |
| let p; |
| try { |
| p = await device.createComputePipelineAsync({ |
| label: key, layout: 'auto', compute: { module, entryPoint: 'main' } }); |
| } catch (err) { |
| await device.popErrorScope().catch(() => null); |
| throw new Error(`${key}: ${err.message}`); |
| } |
| const e2 = await device.popErrorScope(); |
| if (e2) throw new Error(`${key}: ${e2.message}`); |
| return p; |
| } |
|
|
| async _buildPipelines(trySgmat = true) { |
| this.p = {}; |
| const defs = { |
| rmsnorm: WG.rmsnorm(1024, false), |
| rmsnormW: WG.rmsnorm(1024, true), |
| gemv4: WG.gemvQ4(4), |
| gemv8: WG.gemvQ4(8), |
| gemm4: WG.gemmQ4(true, 4), |
| gemm8: WG.gemmQ4(true, 8), |
| gemm4F32x: WG.gemmQ4(false, 4), |
| gemm8F32x: WG.gemmQ4(false, 8), |
| qkvPost: WG.qkvPost(), |
| scores: WG.attnScores(), |
| softmax: WG.attnSoftmaxGate(), |
| pv: WG.attnPV(), |
| decodeAttn: WG.attnDecode(), |
| decodeAttnMq: WG.attnDecodeMq(), |
| mlpAct: WG.mlpAct(), |
| argmax1: WG.argmaxStage1(), |
| argmax2: WG.argmaxStage2(), |
| gather: WG.gatherEmbed(), |
| gemvF16: WG.gemvF16W(), |
| gemmF16: WG.gemmF16W(), |
| relu2: WG.relu2(), |
| }; |
| for (const [k, code] of Object.entries(defs)) this.p[k] = await this._pipe(k, code); |
| if (trySgmat && this.device.features.has('chromium-experimental-subgroup-matrix')) { |
| try { |
| this.p.gemmSg4 = await this._pipe('gemmSg4', WG.gemmQ4Sg(4)); |
| this.p.gemmSg8 = await this._pipe('gemmSg8', WG.gemmQ4Sg(8)); |
| } catch (e) { |
| console.warn('[engine] sgmat-кернелы не скомпилировались:', e.message); |
| delete this.p.gemmSg4; delete this.p.gemmSg8; |
| } |
| } |
| |
| this.pv = (m) => (m.bits === 8 ? this.p.gemv8 : this.p.gemv4); |
| this.pg = (m, xf32) => (m.bits === 8 |
| ? (xf32 ? this.p.gemm8F32x : this.p.gemm8) |
| : (xf32 ? this.p.gemm4F32x : this.p.gemm4)); |
| } |
|
|
| |
|
|
| _sb(bytes, label) { |
| return this.device.createBuffer({ |
| label, size: bytes, |
| usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_DST | GPUBufferUsage.COPY_SRC }); |
| } |
| _ub(vals) { |
| const b = this.device.createBuffer({ |
| size: 16, usage: GPUBufferUsage.UNIFORM | GPUBufferUsage.COPY_DST }); |
| this.device.queue.writeBuffer(b, 0, vals); |
| return b; |
| } |
|
|
| _allocBuffers() { |
| const { maxS, maxT } = this; |
| this.b = { |
| x: this._sb(maxS * 1024 * 4, 'x'), |
| h: this._sb(maxS * 1024 * 4, 'h'), |
| qkv: this._sb(maxS * 4096 * 4, 'qkv'), |
| qbuf: this._sb(16 * maxS * 128 * 4, 'qbuf'), |
| krot16: this._sb(16 * maxS * 128 * 2, 'krot16'), |
| v16: this._sb(16 * maxS * 128 * 2, 'v16'), |
| kcache: this._sb(28 * 8 * maxT * 128 * 2, 'kcache'), |
| vcache: this._sb(28 * 8 * maxT * 128 * 2, 'vcache'), |
| scores: this._sb(16 * maxS * maxS * 4, 'scores'), |
| lse: this._sb(16 * maxS * 4, 'lse'), |
| gate: this._sb(16 * maxS * 4, 'gate'), |
| o: this._sb(maxS * 2048 * 4, 'o'), |
| u: this._sb(maxS * 6144 * 4, 'u'), |
| ymlp: this._sb(maxS * 3072 * 4, 'ymlp'), |
| hidden: this._sb(maxS * 1024 * 4, 'hidden'), |
| logits: this._sb(65536 * 4, 'logits'), |
| rope1: this._sb(maxT * 32 * 2 * 4, 'rope1'), |
| rope2: this._sb(maxT * 16 * 32 * 2 * 4, 'rope2'), |
| |
| pval: this._sb(this.maxK * 256 * 4, 'pval'), |
| pidx: this._sb(this.maxK * 256 * 4, 'pidx'), |
| tokId: this._sb(4, 'tokId'), |
| tokIds: this._sb(this.maxK * 4, 'tokIds'), |
| ids: this._sb(maxS * 4, 'ids'), |
| logitsK: this._sb(this.maxK * 65536 * 4, 'logitsK'), |
| |
| headU: this._sb(2 * this.maxK * 8192 * 4, 'headU'), |
| headU2: this._sb(2 * this.maxK * 8192 * 4, 'headU2'), |
| headLogits: this._sb(2048 * 4, 'headLogits'), |
| headLogitsK: this._sb(2 * this.maxK * 2048 * 4, 'headLogitsK'), |
| patches: this._sb(maxS * 768 * 4, 'patches'), |
| }; |
| |
| const eIn = this.cfg.eps_inner, eF = this.cfg.eps_final; |
| const F = (v) => new Float32Array([0, 0, v, 0]); |
| const normIn = new Uint32Array(F(eIn).buffer); |
| const normF = new Uint32Array(F(eF).buffer); |
| this.u = { |
| normIn: this._ub(normIn), |
| normF: this._ub(normF), |
| |
| gWqkv: this._ub(u32(0, 4096, 1024, 0)), |
| gWo: this._ub(u32(0, 1024, 2048, 1)), |
| gW13: this._ub(u32(0, 6144, 1024, 0)), |
| gW2: this._ub(u32(0, 1024, 3072, 1)), |
| |
| vWqkv: this._ub(u32(4096, 1024, 0, 0)), |
| vWo: this._ub(u32(1024, 2048, 0, 1)), |
| vW13: this._ub(u32(6144, 1024, 0, 0)), |
| vW2: this._ub(u32(1024, 3072, 0, 1)), |
| vLm: this._ub(u32(65536, 1024, 0, 0)), |
| qkvPost: this._ub(u32(0, 0, this.maxT, 1)), |
| sc: this._ub(u32(0, 0, 0, 0)), |
| mlp: this._ub(u32(0, 0, 0, 0)), |
| soft: Array.from({ length: 28 }, (_, l) => this._ub(u32(0, 0, 0, l))), |
| dec: Array.from({ length: 28 }, (_, l) => this._ub(u32(0, this.maxT, l, 0))), |
| |
| decMq: Array.from({ length: 28 }, (_, l) => this._ub(u32(0, this.maxT, l, 0))), |
| gLm: this._ub(u32(0, 65536, 1024, 0)), |
| hW1g: this._ub(u32(0, 8192, 1024, 0)), |
| hReluK: this._ub(u32(0, 0, 0, 0)), |
| hW2gC: this._ub(u32(0, 2048, 8192, 0)), |
| hW2gS: this._ub(u32(0, 2048, 8192, this.maxK)), |
| am1: this._ub(u32(65536, 0, 0, 0)), |
| am2: this._ub(u32(256, 0, 0, 0)), |
| gatherDec: this._ub(u32(1, 0, 0, 0)), |
| gatherPre: this._ub(u32(0, 0, 0, 0)), |
| proj: this._ub(u32(0, 1024, 768, 0)), |
| hW1: this._ub(u32(8192, 1024, 0, 0)), |
| hRelu: this._ub(u32(8192, 0, 0, 0)), |
| hW2: this._ub(u32(2048, 8192, 0, 0)), |
| }; |
| } |
|
|
| _bg(pipeline, buffers) { |
| return this.device.createBindGroup({ |
| layout: pipeline.getBindGroupLayout(0), |
| entries: buffers.map((r, i) => ({ |
| binding: i, resource: r.buffer ? r : { buffer: r } })), |
| }); |
| } |
|
|
| _buildBindGroups() { |
| const { b, u, p, w } = this; |
| const layerCacheK = (l) => ({ buffer: b.kcache, offset: l * 8 * this.maxT * 128 * 2, |
| size: 8 * this.maxT * 128 * 2 }); |
| const layerCacheV = (l) => ({ buffer: b.vcache, offset: l * 8 * this.maxT * 128 * 2, |
| size: 8 * this.maxT * 128 * 2 }); |
| |
| |
| const mm = (pipe, unif, m, xbuf, ybuf, tileN = 32, tileM = 32) => |
| ({ pipe, bg: this._bg(pipe, [unif, m.packed, m.scales, m.zeros, xbuf, ybuf]), |
| tileN, tileM }); |
| const mmSg = (unif, m, xbuf, ybuf) => (this.useSgmat |
| ? mm(m.bits === 8 ? this.p.gemmSg8 : this.p.gemmSg4, unif, m, xbuf, ybuf, 64, 32) |
| : mm(this.pg(m), unif, m, xbuf, ybuf)); |
| this.bgs = { |
| normAttn: this._bg(p.rmsnorm, [u.normIn, b.x, b.h]), |
| normMlp: this._bg(p.rmsnorm, [u.normIn, b.x, b.h]), |
| normFinal: this._bg(p.rmsnormW, [u.normF, b.x, b.hidden, w.final_norm]), |
| scores: this._bg(p.scores, [u.sc, b.qbuf, b.krot16, b.scores]), |
| pv: this._bg(p.pv, [u.sc, b.scores, b.v16, b.o]), |
| mlpAct: this._bg(p.mlpAct, [u.mlp, b.u, b.ymlp]), |
| lm: mm(this.pv(w.lm_head), u.vLm, w.lm_head, b.hidden, b.logits), |
| layers: this.w.layers.map((lw, l) => ({ |
| gemmQkv: mmSg(u.gWqkv, lw.wqkv, b.h, b.qkv), |
| gemmWo: mmSg(u.gWo, lw.wo, b.o, b.x), |
| gemmW13: mmSg(u.gW13, lw.w13, b.h, b.u), |
| gemmW2: mm(this.pg(lw.w2, true), u.gW2, lw.w2, b.ymlp, b.x), |
| gemvQkv: mm(this.pv(lw.wqkv), u.vWqkv, lw.wqkv, b.h, b.qkv), |
| gemvWo: mm(this.pv(lw.wo), u.vWo, lw.wo, b.o, b.x), |
| gemvW13: mm(this.pv(lw.w13), u.vW13, lw.w13, b.h, b.u), |
| gemvW2: mm(this.pv(lw.w2), u.vW2, lw.w2, b.ymlp, b.x), |
| qkvPost: this._bg(p.qkvPost, [u.qkvPost, b.qkv, b.rope1, b.rope2, b.qbuf, |
| layerCacheK(l), layerCacheV(l), b.krot16, b.v16]), |
| softmax: this._bg(p.softmax, [u.soft[l], b.scores, w.gate_bias, b.lse, b.gate]), |
| decodeAttn: this._bg(p.decodeAttn, [u.dec[l], b.qbuf, layerCacheK(l), |
| layerCacheV(l), b.rope2, w.gate_bias, b.o]), |
| decodeAttnMq: this._bg(p.decodeAttnMq, [u.decMq[l], b.qbuf, layerCacheK(l), |
| layerCacheV(l), b.rope2, w.gate_bias, b.o]), |
| })), |
| |
| lmK: mm(this.pg(w.lm_head), u.gLm, w.lm_head, b.hidden, b.logitsK), |
| }; |
| |
| if (w.embed) { |
| this.bgs.gatherDec = this._bg(p.gather, [u.gatherDec, b.tokId, w.embed, b.x]); |
| this.bgs.gatherPre = this._bg(p.gather, [u.gatherPre, b.ids, w.embed, b.x]); |
| } |
| if (w.img_proj) { |
| this.bgs.proj = this._bg(p.gemmF16, [u.proj, w.img_proj, b.patches, b.x]); |
| } |
| this.bgs.argmax1 = this._bg(p.argmax1, [u.am1, b.logits, b.pval, b.pidx]); |
| this.bgs.argmax2 = this._bg(p.argmax2, [u.am2, b.pval, b.pidx, b.tokId]); |
| this.bgs.argmax1K = this._bg(p.argmax1, [u.am1, b.logitsK, b.pval, b.pidx]); |
| this.bgs.argmax2K = this._bg(p.argmax2, [u.am2, b.pval, b.pidx, b.tokIds]); |
| if (this.heads) { |
| const hb = (n) => this.heads[n].buffer; |
| this.bgs.heads = {}; |
| for (const kind of ['coord', 'size']) { |
| this.bgs.heads[kind] = { |
| w1: this._bg(p.gemvF16, [u.hW1, hb(`${kind}_decoder.w1`), b.hidden, b.headU]), |
| act: this._bg(p.relu2, [u.hRelu, b.headU, b.headU2]), |
| w2: this._bg(p.gemvF16, [u.hW2, hb(`${kind}_decoder.w2`), b.headU2, b.headLogits]), |
| }; |
| } |
| |
| |
| const regB = (buf) => ({ buffer: buf, offset: this.maxK * 8192 * 4, |
| size: this.maxK * 8192 * 4 }); |
| this.bgs.headsK = { |
| coordW1: this._bg(p.gemmF16, [u.hW1g, hb('coord_decoder.w1'), b.hidden, b.headU]), |
| coordAct: this._bg(p.relu2, [u.hReluK, b.headU, b.headU2]), |
| coordW2: this._bg(p.gemmF16, [u.hW2gC, hb('coord_decoder.w2'), b.headU2, b.headLogitsK]), |
| sizeW1: this._bg(p.gemmF16, [u.hW1g, hb('size_decoder.w1'), b.hidden, regB(b.headU)]), |
| sizeAct: this._bg(p.relu2, [u.hReluK, regB(b.headU), regB(b.headU2)]), |
| sizeW2: this._bg(p.gemmF16, [u.hW2gS, hb('size_decoder.w2'), regB(b.headU2), b.headLogitsK]), |
| }; |
| } |
| } |
|
|
| |
|
|
| setRopeTables(rope1F32, rope2F32) { |
| this.device.queue.writeBuffer(this.b.rope1, 0, rope1F32); |
| this.device.queue.writeBuffer(this.b.rope2, 0, rope2F32); |
| } |
|
|
| |
|
|
| _pass(enc, pipeline, bg, gx, gy = 1, gz = 1) { |
| const pass = enc.beginComputePass(); |
| pass.setPipeline(pipeline); |
| pass.setBindGroup(0, bg); |
| pass.dispatchWorkgroups(gx, gy, gz); |
| pass.end(); |
| } |
| |
| |
| _d(pass, pipeline, bg, gx, gy = 1, gz = 1) { |
| pass.setPipeline(pipeline); |
| pass.setBindGroup(0, bg); |
| pass.dispatchWorkgroups(gx, gy, gz); |
| } |
| _gemvGrid(n) { return [Math.min(n, WG.GRID_X), Math.ceil(n / WG.GRID_X)]; } |
|
|
| async readF32(buffer, count, offsetBytes = 0) { |
| const st = this.device.createBuffer({ |
| size: count * 4, usage: GPUBufferUsage.MAP_READ | GPUBufferUsage.COPY_DST }); |
| const enc = this.device.createCommandEncoder(); |
| enc.copyBufferToBuffer(buffer, offsetBytes, st, 0, count * 4); |
| this.device.queue.submit([enc.finish()]); |
| await st.mapAsync(GPUMapMode.READ); |
| const out = new Float32Array(st.getMappedRange().slice(0)); |
| st.destroy(); |
| return out; |
| } |
| async readF16(buffer, count, offsetBytes = 0) { |
| const st = this.device.createBuffer({ |
| size: count * 2, usage: GPUBufferUsage.MAP_READ | GPUBufferUsage.COPY_DST }); |
| const enc = this.device.createCommandEncoder(); |
| enc.copyBufferToBuffer(buffer, offsetBytes, st, 0, count * 2); |
| this.device.queue.submit([enc.finish()]); |
| await st.mapAsync(GPUMapMode.READ); |
| const out = new Uint16Array(st.getMappedRange().slice(0)); |
| st.destroy(); |
| return out; |
| } |
|
|
| |
|
|
| |
| |
| async prefill(emb, bidirEnd, hooks = {}, sLen = 0) { |
| const S = emb ? emb.length / 1024 : sLen; |
| if (S > this.maxS) throw new Error(`S=${S} > maxS=${this.maxS}`); |
| if (!S) throw new Error('prefill: нет ни emb, ни sLen'); |
| const q = this.device.queue; |
| if (emb) q.writeBuffer(this.b.x, 0, emb); |
| |
| q.writeBuffer(this.u.gWqkv, 0, u32(S)); |
| q.writeBuffer(this.u.gWo, 0, u32(S)); |
| q.writeBuffer(this.u.gW13, 0, u32(S)); |
| q.writeBuffer(this.u.gW2, 0, u32(S)); |
| q.writeBuffer(this.u.qkvPost, 0, u32(S, 0, this.maxT, 1)); |
| q.writeBuffer(this.u.sc, 0, u32(S, S)); |
| q.writeBuffer(this.u.mlp, 0, u32(S * 3072)); |
| q.writeBuffer(this.u.vLm, 0, u32(65536, 1024, (S - 1) * 256, 0)); |
| for (let l = 0; l < 28; l++) q.writeBuffer(this.u.soft[l], 0, u32(S, S, bidirEnd)); |
|
|
| const g16 = Math.ceil(S / 16), g32 = Math.ceil(S / 32); |
| |
| |
| const single = !hooks.afterAttn && !hooks.afterMlp; |
| let enc = single ? this.device.createCommandEncoder() : null; |
| let pass = single ? enc.beginComputePass() : null; |
| const D = (pipe, bg, gx, gy = 1, gz = 1) => { |
| if (single) this._d(pass, pipe, bg, gx, gy, gz); |
| else this._pass(enc, pipe, bg, gx, gy, gz); |
| }; |
| for (let l = 0; l < 28; l++) { |
| const lb = this.bgs.layers[l]; |
| if (!single) enc = this.device.createCommandEncoder(); |
| D(this.p.rmsnorm, this.bgs.normAttn, S); |
| D(lb.gemmQkv.pipe, lb.gemmQkv.bg, |
| Math.ceil(4096 / lb.gemmQkv.tileN), Math.ceil(S / lb.gemmQkv.tileM)); |
| D(this.p.qkvPost, lb.qkvPost, S, 24); |
| D(this.p.scores, this.bgs.scores, g16, g16, 16); |
| D(this.p.softmax, lb.softmax, S, 16); |
| D(this.p.pv, this.bgs.pv, 8, g16, 16); |
| D(lb.gemmWo.pipe, lb.gemmWo.bg, |
| Math.ceil(1024 / lb.gemmWo.tileN), Math.ceil(S / lb.gemmWo.tileM)); |
| if (!single) { |
| q.submit([enc.finish()]); |
| if (hooks.afterAttn) { await q.onSubmittedWorkDone(); await hooks.afterAttn(l); } |
| enc = this.device.createCommandEncoder(); |
| } |
| D(this.p.rmsnorm, this.bgs.normMlp, S); |
| D(lb.gemmW13.pipe, lb.gemmW13.bg, |
| Math.ceil(6144 / lb.gemmW13.tileN), Math.ceil(S / lb.gemmW13.tileM)); |
| D(this.p.mlpAct, this.bgs.mlpAct, Math.ceil(S * 3072 / 256)); |
| D(lb.gemmW2.pipe, lb.gemmW2.bg, Math.ceil(1024 / 32), g32); |
| if (!single) { |
| q.submit([enc.finish()]); |
| if (hooks.afterMlp) { await q.onSubmittedWorkDone(); await hooks.afterMlp(l); } |
| else if (l % 4 === 3) await q.onSubmittedWorkDone(); |
| } |
| } |
| if (!single) enc = this.device.createCommandEncoder(); |
| else { } |
| D(this.p.rmsnormW, this.bgs.normFinal, S); |
| D(this.bgs.lm.pipe, this.bgs.lm.bg, ...this._gemvGrid(65536)); |
| if (single) pass.end(); |
| q.submit([enc.finish()]); |
| await q.onSubmittedWorkDone(); |
| this.T = S; |
| this._decUniformsDirty = true; |
| return { S }; |
| } |
|
|
| |
| async decodeStep(emb) { |
| const q = this.device.queue; |
| const T = this.T + 1; |
| q.writeBuffer(this.b.x, 0, emb); |
| q.writeBuffer(this.u.qkvPost, 0, u32(1, this.T, this.maxT, 0)); |
| q.writeBuffer(this.u.mlp, 0, u32(3072)); |
| q.writeBuffer(this.u.vLm, 0, u32(65536, 1024, 0, 0)); |
| for (let l = 0; l < 28; l++) q.writeBuffer(this.u.dec[l], 0, u32(T)); |
|
|
| const enc = this.device.createCommandEncoder(); |
| for (let l = 0; l < 28; l++) { |
| const lb = this.bgs.layers[l]; |
| this._pass(enc, this.p.rmsnorm, this.bgs.normAttn, 1); |
| this._pass(enc, lb.gemvQkv.pipe, lb.gemvQkv.bg, ...this._gemvGrid(4096)); |
| this._pass(enc, this.p.qkvPost, lb.qkvPost, 1, 24); |
| this._pass(enc, this.p.decodeAttn, lb.decodeAttn, 8); |
| this._pass(enc, lb.gemvWo.pipe, lb.gemvWo.bg, ...this._gemvGrid(1024)); |
| this._pass(enc, this.p.rmsnorm, this.bgs.normMlp, 1); |
| this._pass(enc, lb.gemvW13.pipe, lb.gemvW13.bg, ...this._gemvGrid(6144)); |
| this._pass(enc, this.p.mlpAct, this.bgs.mlpAct, Math.ceil(3072 / 256)); |
| this._pass(enc, lb.gemvW2.pipe, lb.gemvW2.bg, ...this._gemvGrid(1024)); |
| } |
| this._pass(enc, this.p.rmsnormW, this.bgs.normFinal, 1); |
| this._pass(enc, this.bgs.lm.pipe, this.bgs.lm.bg, ...this._gemvGrid(65536)); |
| q.submit([enc.finish()]); |
| await q.onSubmittedWorkDone(); |
| this.T = T; |
| } |
|
|
| async readU32(buffer, count, offsetBytes = 0) { |
| const f = await this.readF32(buffer, count, offsetBytes); |
| return new Uint32Array(f.buffer); |
| } |
|
|
| |
| |
| async buildPrefillEmbeddings(idsU32, patchesF32, imgStart) { |
| const q = this.device.queue; |
| const C = idsU32.length; |
| q.writeBuffer(this.b.ids, 0, idsU32); |
| q.writeBuffer(this.u.gatherPre, 0, u32(C)); |
| const enc = this.device.createCommandEncoder(); |
| this._pass(enc, this.p.gather, this.bgs.gatherPre, C); |
| if (patchesF32 && patchesF32.length) { |
| const m = patchesF32.length / 768; |
| q.writeBuffer(this.b.patches, 0, patchesF32); |
| q.writeBuffer(this.u.proj, 0, u32(m, 1024, 768, imgStart)); |
| this._pass(enc, this.p.gemmF16, this.bgs.proj, 1024 / 16, Math.ceil(m / 16)); |
| } |
| q.submit([enc.finish()]); |
| } |
|
|
| |
| writeX(embF32) { this.device.queue.writeBuffer(this.b.x, 0, embF32); } |
|
|
| |
| |
| |
| |
| async decodeStepAuto(useGather) { |
| const q = this.device.queue; |
| const T = this.T + 1; |
| q.writeBuffer(this.u.qkvPost, 0, u32(1, this.T, this.maxT, 0)); |
| if (this._decUniformsDirty) { |
| q.writeBuffer(this.u.mlp, 0, u32(3072)); |
| q.writeBuffer(this.u.vLm, 0, u32(65536, 1024, 0, 0)); |
| this._decUniformsDirty = false; |
| } |
| for (let l = 0; l < 28; l++) q.writeBuffer(this.u.dec[l], 0, u32(T)); |
|
|
| const enc = this.device.createCommandEncoder(); |
| const pass = enc.beginComputePass(); |
| if (useGather) this._d(pass, this.p.gather, this.bgs.gatherDec, 1); |
| for (let l = 0; l < 28; l++) { |
| const lb = this.bgs.layers[l]; |
| this._d(pass, this.p.rmsnorm, this.bgs.normAttn, 1); |
| this._d(pass, lb.gemvQkv.pipe, lb.gemvQkv.bg, ...this._gemvGrid(4096)); |
| this._d(pass, this.p.qkvPost, lb.qkvPost, 1, 24); |
| this._d(pass, this.p.decodeAttn, lb.decodeAttn, 8); |
| this._d(pass, lb.gemvWo.pipe, lb.gemvWo.bg, ...this._gemvGrid(1024)); |
| this._d(pass, this.p.rmsnorm, this.bgs.normMlp, 1); |
| this._d(pass, lb.gemvW13.pipe, lb.gemvW13.bg, ...this._gemvGrid(6144)); |
| this._d(pass, this.p.mlpAct, this.bgs.mlpAct, Math.ceil(3072 / 256)); |
| this._d(pass, lb.gemvW2.pipe, lb.gemvW2.bg, ...this._gemvGrid(1024)); |
| } |
| this._d(pass, this.p.rmsnormW, this.bgs.normFinal, 1); |
| this._d(pass, this.bgs.lm.pipe, this.bgs.lm.bg, ...this._gemvGrid(65536)); |
| this._d(pass, this.p.argmax1, this.bgs.argmax1, 256); |
| this._d(pass, this.p.argmax2, this.bgs.argmax2, 1); |
| pass.end(); |
| q.submit([enc.finish()]); |
| this.T = T; |
| return (await this.readU32(this.b.tokId, 1))[0]; |
| } |
|
|
| |
| |
| |
| |
| |
| |
| |
| |
| |
| |
| async verifyStep(items) { |
| const k = items.length, q = this.device.queue, T0 = this.T; |
| if (k < 1 || k > this.maxK) throw new Error(`verifyStep: k=${k} вне [1, ${this.maxK}]`); |
| if (T0 + k > this.maxT) throw new Error(`verifyStep: T0+k=${T0 + k} > maxT`); |
| |
| const ids = new Uint32Array(k).fill(0xFFFFFFFF); |
| for (let i = 0; i < k; i++) if (!items[i].emb) ids[i] = items[i].id; |
| q.writeBuffer(this.b.ids, 0, ids); |
| q.writeBuffer(this.u.gatherPre, 0, u32(k)); |
| for (let i = 0; i < k; i++) |
| if (items[i].emb) q.writeBuffer(this.b.x, i * 4096, items[i].emb); |
| |
| q.writeBuffer(this.u.gWqkv, 0, u32(k)); |
| q.writeBuffer(this.u.gWo, 0, u32(k)); |
| q.writeBuffer(this.u.gW13, 0, u32(k)); |
| q.writeBuffer(this.u.gW2, 0, u32(k)); |
| q.writeBuffer(this.u.qkvPost, 0, u32(k, T0, this.maxT, 0)); |
| q.writeBuffer(this.u.mlp, 0, u32(k * 3072)); |
| q.writeBuffer(this.u.gLm, 0, u32(k)); |
| for (let l = 0; l < 28; l++) |
| this.device.queue.writeBuffer(this.u.decMq[l], 0, u32(T0, this.maxT, l, k)); |
| const hasHeads = !!(this.bgs.headsK); |
| if (hasHeads) { |
| q.writeBuffer(this.u.hW1g, 0, u32(k)); |
| q.writeBuffer(this.u.hReluK, 0, u32(k * 8192)); |
| q.writeBuffer(this.u.hW2gC, 0, u32(k)); |
| q.writeBuffer(this.u.hW2gS, 0, u32(k)); |
| } |
|
|
| const g16 = Math.ceil(k / 16), g32 = Math.ceil(k / 32); |
| const enc = this.device.createCommandEncoder(); |
| const pass = enc.beginComputePass(); |
| this._d(pass, this.p.gather, this.bgs.gatherPre, k); |
| for (let l = 0; l < 28; l++) { |
| const lb = this.bgs.layers[l]; |
| this._d(pass, this.p.rmsnorm, this.bgs.normAttn, k); |
| this._d(pass, lb.gemmQkv.pipe, lb.gemmQkv.bg, |
| Math.ceil(4096 / lb.gemmQkv.tileN), Math.ceil(k / lb.gemmQkv.tileM)); |
| this._d(pass, this.p.qkvPost, lb.qkvPost, k, 24); |
| this._d(pass, this.p.decodeAttnMq, lb.decodeAttnMq, 8, k); |
| this._d(pass, lb.gemmWo.pipe, lb.gemmWo.bg, |
| Math.ceil(1024 / lb.gemmWo.tileN), Math.ceil(k / lb.gemmWo.tileM)); |
| this._d(pass, this.p.rmsnorm, this.bgs.normMlp, k); |
| this._d(pass, lb.gemmW13.pipe, lb.gemmW13.bg, |
| Math.ceil(6144 / lb.gemmW13.tileN), Math.ceil(k / lb.gemmW13.tileM)); |
| this._d(pass, this.p.mlpAct, this.bgs.mlpAct, Math.ceil(k * 3072 / 256)); |
| this._d(pass, lb.gemmW2.pipe, lb.gemmW2.bg, Math.ceil(1024 / 32), g32); |
| } |
| this._d(pass, this.p.rmsnormW, this.bgs.normFinal, k); |
| this._d(pass, this.bgs.lmK.pipe, this.bgs.lmK.bg, Math.ceil(65536 / 32), g32); |
| this._d(pass, this.p.argmax1, this.bgs.argmax1K, 256, k); |
| this._d(pass, this.p.argmax2, this.bgs.argmax2K, k); |
| if (hasHeads) { |
| const hk = this.bgs.headsK; |
| this._d(pass, this.p.gemmF16, hk.coordW1, 8192 / 16, g16); |
| this._d(pass, this.p.gemmF16, hk.sizeW1, 8192 / 16, g16); |
| this._d(pass, this.p.relu2, hk.coordAct, Math.ceil(k * 8192 / 256)); |
| this._d(pass, this.p.relu2, hk.sizeAct, Math.ceil(k * 8192 / 256)); |
| this._d(pass, this.p.gemmF16, hk.coordW2, 2048 / 16, g16); |
| this._d(pass, this.p.gemmF16, hk.sizeW2, 2048 / 16, g16); |
| } |
| pass.end(); |
| |
| const cBytes = k * 2048 * 4; |
| const st = this.device.createBuffer({ |
| size: k * 4 + (hasHeads ? 2 * cBytes : 0), |
| usage: GPUBufferUsage.MAP_READ | GPUBufferUsage.COPY_DST }); |
| enc.copyBufferToBuffer(this.b.tokIds, 0, st, 0, k * 4); |
| if (hasHeads) { |
| enc.copyBufferToBuffer(this.b.headLogitsK, 0, st, k * 4, cBytes); |
| enc.copyBufferToBuffer(this.b.headLogitsK, this.maxK * 2048 * 4, st, k * 4 + cBytes, cBytes); |
| } |
| q.submit([enc.finish()]); |
| this._decUniformsDirty = true; |
| await st.mapAsync(GPUMapMode.READ); |
| const raw = st.getMappedRange().slice(0); |
| st.destroy(); |
| return { |
| T0, |
| tokens: new Uint32Array(raw, 0, k), |
| coordLogits: hasHeads ? new Float32Array(raw, k * 4, k * 2048) : null, |
| sizeLogits: hasHeads ? new Float32Array(raw, k * 4 + cBytes, k * 2048) : null, |
| }; |
| } |
|
|
| |
| async runHead(kind) { |
| const h = this.bgs.heads[kind]; |
| const enc = this.device.createCommandEncoder(); |
| const pass = enc.beginComputePass(); |
| this._d(pass, this.p.gemvF16, h.w1, ...this._gemvGrid(8192)); |
| this._d(pass, this.p.relu2, h.act, Math.ceil(8192 / 256)); |
| this._d(pass, this.p.gemvF16, h.w2, ...this._gemvGrid(2048)); |
| pass.end(); |
| this.device.queue.submit([enc.finish()]); |
| return this.readF32(this.b.headLogits, 2048); |
| } |
|
|
| |
| encodeCoords(kind, xy) { |
| const eE = this.heads[`${kind}_encoder.embed`]; |
| const eT = this.heads[`${kind}_encoder.transform`]; |
| if (!eE || !eE.cpu || !eT || !eT.cpu) { |
| throw new Error(`encodeCoords(${kind}): нет весов энкодера; ` + |
| `eE=${eE && JSON.stringify({keys: Object.keys(eE), shape: eE.shape, cpuLen: eE.cpu && eE.cpu.length})}`); |
| } |
| const E = eE.cpu; |
| const Tw = eT.cpu; |
| const feat = new Float32Array(512); |
| for (let i = 0; i < 256; i++) { |
| const a = 6.2831855 * (E[i * 2] * xy[0] + E[i * 2 + 1] * xy[1]); |
| feat[i] = Math.cos(a); |
| feat[256 + i] = Math.sin(a); |
| } |
| const out = new Float32Array(1024); |
| for (let n = 0; n < 1024; n++) { |
| let s = 0; |
| const base = n * 512; |
| for (let k = 0; k < 512; k++) s += Tw[base + k] * feat[k]; |
| out[n] = s; |
| } |
| return out; |
| } |
| } |
|
|