wgblas 2.1.0 → 2.2.1

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (99) hide show
  1. package/README.md +2 -0
  2. package/dist/wgblas.browser.js +1016 -43
  3. package/index.d.mts +26 -53
  4. package/index.mjs +11 -0
  5. package/package.json +132 -63
  6. package/src/classes/Complex32.d.mts +43 -0
  7. package/src/classes/Complex32.mjs +82 -0
  8. package/src/classes/Complex64.d.mts +44 -0
  9. package/src/classes/Complex64.mjs +76 -0
  10. package/src/classes/GpuMatrix.d.mts +41 -41
  11. package/src/classes/GpuMatrix.mjs +112 -10
  12. package/src/classes/GpuVector.d.mts +36 -40
  13. package/src/classes/GpuVector.mjs +39 -2
  14. package/src/cscal/cscal.d.mts +47 -0
  15. package/src/cscal/cscal.mjs +98 -0
  16. package/src/dasum/dasum.mjs +31 -15
  17. package/src/daxpy/daxpy.d.mts +56 -0
  18. package/src/daxpy/daxpy.mjs +150 -0
  19. package/src/dcopy/dcopy.d.mts +52 -0
  20. package/src/dcopy/dcopy.mjs +140 -0
  21. package/src/ddot/ddot.d.mts +62 -0
  22. package/src/ddot/ddot.mjs +184 -0
  23. package/src/dnrm2/dnrm2.d.mts +50 -0
  24. package/src/dnrm2/dnrm2.mjs +189 -0
  25. package/src/drot/drot.d.mts +67 -0
  26. package/src/drot/drot.mjs +170 -0
  27. package/src/drotm/drotm.d.mts +67 -0
  28. package/src/drotm/drotm.mjs +171 -0
  29. package/src/dscal/dscal.d.mts +52 -0
  30. package/src/dscal/dscal.mjs +119 -0
  31. package/src/dswap/dswap.d.mts +57 -0
  32. package/src/dswap/dswap.mjs +155 -0
  33. package/src/idamax/idamax.mjs +49 -19
  34. package/src/init.mjs +6 -3
  35. package/src/isamax/isamax.mjs +17 -14
  36. package/src/random/random.d.mts +37 -40
  37. package/src/random/random.mjs +39 -7
  38. package/src/sasum/sasum.mjs +13 -11
  39. package/src/saxpy/saxpy.mjs +9 -8
  40. package/src/scopy/scopy.mjs +8 -6
  41. package/src/sdot/sdot.mjs +13 -11
  42. package/src/sgemm/sgemm.d.mts +2 -2
  43. package/src/sgemm/sgemm.mjs +91 -35
  44. package/src/sgemmtr/sgemmtr.d.mts +3 -2
  45. package/src/sgemmtr/sgemmtr.mjs +92 -35
  46. package/src/sgemv/sgemv.d.mts +2 -2
  47. package/src/sgemv/sgemv.mjs +41 -25
  48. package/src/sger/sger.d.mts +2 -2
  49. package/src/sger/sger.mjs +38 -16
  50. package/src/shaders/__test_pipeline_a.wgsl +3 -0
  51. package/src/shaders/__test_pipeline_b.wgsl +2 -0
  52. package/src/shaders/cscal.wgsl +33 -0
  53. package/src/shaders/daxpy.wgsl +66 -0
  54. package/src/shaders/dcopy.wgsl +34 -0
  55. package/src/shaders/ddot.wgsl +106 -0
  56. package/src/shaders/dnrm2.wgsl +167 -0
  57. package/src/shaders/drot.wgsl +81 -0
  58. package/src/shaders/drotm.wgsl +99 -0
  59. package/src/shaders/dscal.wgsl +60 -0
  60. package/src/shaders/dswap.wgsl +38 -0
  61. package/src/shaders/f64/utils/add.wgsl +6 -0
  62. package/src/shaders/f64/utils/divide.wgsl +45 -0
  63. package/src/shaders/f64/utils/multiply.wgsl +29 -11
  64. package/src/shaders/f64/utils/sqrt.wgsl +45 -0
  65. package/src/shaders/index.mjs +69 -0
  66. package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
  67. package/src/snrm2/snrm2.mjs +20 -14
  68. package/src/srot/srot.mjs +10 -7
  69. package/src/srotm/srotm.mjs +9 -12
  70. package/src/sscal/sscal.mjs +7 -7
  71. package/src/sswap/sswap.mjs +14 -8
  72. package/src/ssymm/ssymm.d.mts +5 -4
  73. package/src/ssymm/ssymm.mjs +142 -55
  74. package/src/ssymv/ssymv.d.mts +2 -2
  75. package/src/ssymv/ssymv.mjs +42 -23
  76. package/src/ssyr/ssyr.d.mts +2 -2
  77. package/src/ssyr/ssyr.mjs +34 -15
  78. package/src/ssyr2/ssyr2.d.mts +2 -2
  79. package/src/ssyr2/ssyr2.mjs +43 -18
  80. package/src/ssyr2k/ssyr2k.d.mts +3 -2
  81. package/src/ssyr2k/ssyr2k.mjs +132 -55
  82. package/src/ssyrk/ssyrk.d.mts +3 -2
  83. package/src/ssyrk/ssyrk.mjs +84 -33
  84. package/src/strmm/strmm.d.mts +5 -4
  85. package/src/strmm/strmm.mjs +153 -54
  86. package/src/strmv/strmv.d.mts +2 -2
  87. package/src/strmv/strmv.mjs +37 -17
  88. package/src/strsm/strsm.d.mts +6 -4
  89. package/src/strsm/strsm.mjs +418 -172
  90. package/src/strsv/strsv.d.mts +5 -3
  91. package/src/strsv/strsv.mjs +82 -31
  92. package/src/util/benchmark.mjs +5 -3
  93. package/src/util/buffer.mjs +33 -12
  94. package/src/util/complex.mjs +87 -0
  95. package/src/util/compute.mjs +14 -8
  96. package/src/util/device.mjs +18 -3
  97. package/src/util/pipeline.mjs +40 -5
  98. package/src/util/workgroup.mjs +23 -6
  99. package/src/shaders/f64add.wgsl +0 -281
@@ -11,26 +11,46 @@ import { beginTimedEncoder, encodePass, submit } from "../util/compute.mjs";
11
11
  import { extractResult } from "../util/result.mjs";
12
12
  import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
13
13
  import { getPipeline } from "../util/pipeline.mjs";
14
- import { calcWorkgroups, requireWorkgroups, requireWorkgroupCount } from "../util/workgroup.mjs";
14
+ import {
15
+ calcWorkgroups,
16
+ requireWorkgroups,
17
+ requireWorkgroupCount,
18
+ } from "../util/workgroup.mjs";
15
19
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
16
- import { BM_SMALL, BN_SMALL, BM_LARGE, BN_LARGE, LARGE_TILE_WORKGROUP_THRESHOLD } from "../util/constants.mjs";
20
+ import {
21
+ BM_SMALL,
22
+ BN_SMALL,
23
+ BM_LARGE,
24
+ BN_LARGE,
25
+ LARGE_TILE_WORKGROUP_THRESHOLD,
26
+ } from "../util/constants.mjs";
17
27
  import { BLOCK_SIZE } from "../util/constants.mjs";
18
- import { requireSameDevice } from "../util/device.mjs";
19
-
28
+ import { requireGpuDevice, requireSameDevice } from "../util/device.mjs";
20
29
 
21
30
  // strsm: B := alpha*op(A)^-1*B (side='left') or alpha*B*op(A)^-1 (side='right'),
22
31
  // A triangular. Blocked substitution (strsv's own technique, generalized to
23
32
  // a matrix RHS): strsv_invert_block + sgemm, unchanged; every per-block B/A
24
33
  // access goes through block_transfer.wgsl (see that shader for why).
25
34
  export async function strsm(
26
- device, side, uplo, transA, diag, m, n, alpha, A, lda, B, ldb, layout = "row-major",
35
+ device,
36
+ side,
37
+ uplo,
38
+ transA,
39
+ diag,
40
+ m,
41
+ n,
42
+ alpha,
43
+ A,
44
+ lda,
45
+ B,
46
+ ldb,
47
+ layout = "row-major",
27
48
  ) {
28
49
  const AIsGpu = A instanceof GpuMatrix;
29
50
  const BIsGpu = B instanceof GpuMatrix;
30
51
  const isUnit = diag === "unit";
31
52
 
32
- if (!(device instanceof GPUDevice))
33
- throw new Error("device must be a GPUDevice.");
53
+ requireGpuDevice(device);
34
54
  requireSameDevice(device, "strsm", { A, B });
35
55
  if (side !== "left" && side !== "right")
36
56
  throw new Error("side must be 'left' or 'right'.");
@@ -42,11 +62,15 @@ export async function strsm(
42
62
  throw new Error("diag must be 'unit' or 'non-unit'.");
43
63
  if (layout !== "row-major" && layout !== "column-major")
44
64
  throw new Error("layout must be 'row-major' or 'column-major'.");
45
- if (typeof alpha !== "number")
46
- throw new Error("alpha must be a number.");
65
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
47
66
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
48
67
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
49
- if (!Number.isInteger(m) || !Number.isInteger(n) || !Number.isInteger(lda) || !Number.isInteger(ldb))
68
+ if (
69
+ !Number.isInteger(m) ||
70
+ !Number.isInteger(n) ||
71
+ !Number.isInteger(lda) ||
72
+ !Number.isInteger(ldb)
73
+ )
50
74
  throw new Error("m, n, lda, and ldb must be integers.");
51
75
  if (!AIsGpu && !(A instanceof Float32Array))
52
76
  throw new Error("A must be a Float32Array or GpuMatrix.");
@@ -62,30 +86,51 @@ export async function strsm(
62
86
 
63
87
  // A: triangular, order = m (side='left') or n (side='right').
64
88
  const aOrder = side === "left" ? m : n;
65
- if (lda < aOrder) throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
89
+ if (lda < aOrder)
90
+ throw new Error("lda must be >= " + (side === "left" ? "m" : "n") + ".");
66
91
  if (AIsGpu) {
67
- if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
68
- if (A.rows < aOrder || A.cols < aOrder) throw new Error("A is too small for the given m/n and side.");
92
+ if (lda !== A.lda)
93
+ throw new Error("lda must match A.lda when A is a GpuMatrix.");
94
+ if (A.rows < aOrder || A.cols < aOrder)
95
+ throw new Error("A is too small for the given m/n and side.");
69
96
  } else if (A.length < (aOrder - 1) * lda + aOrder) {
70
- throw new Error("A does not have enough elements for the given dimensions and lda.");
97
+ throw new Error(
98
+ "A does not have enough elements for the given dimensions and lda.",
99
+ );
71
100
  }
72
101
 
73
102
  // B: always m x n, overwritten in place with the same ldb.
74
103
  const bOuter = effLayoutB === "column-major" ? n : m;
75
104
  const bInner = effLayoutB === "column-major" ? m : n;
76
105
  if (ldb < bInner)
77
- throw new Error(`ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`);
106
+ throw new Error(
107
+ `ldb must be >= ${effLayoutB === "column-major" ? "rows" : "cols"} of B as stored.`,
108
+ );
78
109
  if (BIsGpu) {
79
- if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
80
- if (B.rows < m || B.cols < n) throw new Error("B is too small for the given m and n.");
110
+ if (ldb !== B.lda)
111
+ throw new Error("ldb must match B.lda when B is a GpuMatrix.");
112
+ if (B.rows < m || B.cols < n)
113
+ throw new Error("B is too small for the given m and n.");
81
114
  } else if (B.length < (bOuter - 1) * ldb + bInner) {
82
- throw new Error("B does not have enough elements for the given dimensions and ldb.");
115
+ throw new Error(
116
+ "B does not have enough elements for the given dimensions and ldb.",
117
+ );
83
118
  }
84
119
 
85
120
  // A isn't symmetric: column-major = genuine transpose, so flip transA;
86
121
  // transposing also swaps which triangle looks stored, so flip uplo too.
87
- const uploEffA = effLayoutA === "column-major" ? (uplo === "lower" ? "upper" : "lower") : uplo;
88
- const transEffA = effLayoutA === "column-major" ? (transA === "no-transpose" ? "transpose" : "no-transpose") : transA;
122
+ const uploEffA =
123
+ effLayoutA === "column-major"
124
+ ? uplo === "lower"
125
+ ? "upper"
126
+ : "lower"
127
+ : uplo;
128
+ const transEffA =
129
+ effLayoutA === "column-major"
130
+ ? transA === "no-transpose"
131
+ ? "transpose"
132
+ : "no-transpose"
133
+ : transA;
89
134
 
90
135
  const otherLen = side === "left" ? n : m;
91
136
  const blockIsRow = side === "left";
@@ -132,11 +177,17 @@ export async function strsm(
132
177
  try {
133
178
  ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strsm-A", false);
134
179
  BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "strsm-B", true);
135
- AinvBuffer = createStorageBuffer(device, numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4, "strsm-Ainv");
180
+ AinvBuffer = createStorageBuffer(
181
+ device,
182
+ numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4,
183
+ "strsm-Ainv",
184
+ );
136
185
 
137
186
  // Pre-scale B by alpha once (reuses sscal, so no per-block alpha handling).
187
+ // Skipped for alpha=0 — sscal computes 0*B[i], which leaks NaN/Inf from a
188
+ // poisoned B; see the literal zero-write dispatched below instead.
138
189
  let preScaleBindGroup = null;
139
- if (alpha !== 1.0) {
190
+ if (alpha !== 1.0 && alpha !== 0) {
140
191
  const scaleParams = params(
141
192
  [
142
193
  { value: bScaleLen, type: "u32" },
@@ -145,7 +196,11 @@ export async function strsm(
145
196
  ],
146
197
  "strsm-scale-params",
147
198
  );
148
- preScaleBindGroup = createBindGroup(device, scalarPipeline.getBindGroupLayout(0), [BBuffer, scaleParams]);
199
+ preScaleBindGroup = createBindGroup(
200
+ device,
201
+ scalarPipeline.getBindGroupLayout(0),
202
+ [BBuffer, scaleParams],
203
+ );
149
204
  }
150
205
 
151
206
  // Every diagonal block's inverse, fully parallel, one dispatch, unchanged.
@@ -159,7 +214,11 @@ export async function strsm(
159
214
  ],
160
215
  "strsm-invert-params",
161
216
  );
162
- const invertBindGroup = createBindGroup(device, invertPipeline.getBindGroupLayout(0), [ABuffer, AinvBuffer, invertParams]);
217
+ const invertBindGroup = createBindGroup(
218
+ device,
219
+ invertPipeline.getBindGroupLayout(0),
220
+ [ABuffer, AinvBuffer, invertParams],
221
+ );
163
222
 
164
223
  // Reusable scratch buffers, sized for the worst case, bound at offset 0.
165
224
  const Bblock = scratch(BLOCK_SIZE * otherLen * 4, "strsm-Bblock");
@@ -170,173 +229,360 @@ export async function strsm(
170
229
  const { commandEncoder, querySet } = beginTimedEncoder(device);
171
230
 
172
231
  if (alpha === 0) {
173
- // BLAS: alpha=0 means A is not referenced — skip straight to B:=0.
174
- const zeroDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0, endOfPassWriteIndex: 1 } } : undefined;
175
- encodePass(commandEncoder, scalarPipeline, preScaleBindGroup, calcWorkgroups(device, bScaleLen), zeroDesc);
232
+ // BLAS: alpha=0 means B must become a literal zero, not 0*B (leaks
233
+ // NaN/Inf via sscal). Reuse sgemm with k=0 (X/Y never read, so
234
+ // AinvBuffer is a safe dummy for both) and beta=0.
235
+ const zeroLargeWgX = Math.ceil(bInner / BN_LARGE);
236
+ const zeroLargeWgY = Math.ceil(bOuter / BM_LARGE);
237
+ const zeroUseLarge =
238
+ zeroLargeWgX * zeroLargeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
239
+ const zeroPipeline = await getPipeline(
240
+ device,
241
+ zeroUseLarge ? "sgemm_large" : "sgemm_small",
242
+ );
243
+ const zeroParams = params(
244
+ [
245
+ { value: bOuter, type: "u32" },
246
+ { value: bInner, type: "u32" },
247
+ { value: 0, type: "u32" }, // k forced to 0 — X/Y never read
248
+ { value: 0.0, type: "f32" }, // alpha
249
+ { value: 0.0, type: "f32" }, // beta
250
+ { value: 1, type: "u32" }, // ldX (dummy, unread)
251
+ { value: 1, type: "u32" }, // ldY (dummy, unread)
252
+ { value: ldb, type: "u32" }, // ldc — B's own real stride
253
+ { value: 0, type: "u32" }, // transX (dummy)
254
+ { value: 0, type: "u32" }, // transY (dummy)
255
+ ],
256
+ "strsm-zero-params",
257
+ );
258
+ const zeroBindGroup = createBindGroup(
259
+ device,
260
+ zeroPipeline.getBindGroupLayout(0),
261
+ [
262
+ AinvBuffer,
263
+ vec4ViewBinding(device, AinvBuffer),
264
+ AinvBuffer,
265
+ vec4ViewBinding(device, AinvBuffer),
266
+ BBuffer,
267
+ zeroParams,
268
+ ],
269
+ );
270
+ const zeroWgCount = zeroUseLarge
271
+ ? {
272
+ x: requireWorkgroupCount(device, zeroLargeWgX, "strsm", "x"),
273
+ y: requireWorkgroupCount(device, zeroLargeWgY, "strsm", "y"),
274
+ }
275
+ : {
276
+ x: requireWorkgroupCount(
277
+ device,
278
+ Math.ceil(bInner / BN_SMALL),
279
+ "strsm",
280
+ "x",
281
+ ),
282
+ y: requireWorkgroupCount(
283
+ device,
284
+ Math.ceil(bOuter / BM_SMALL),
285
+ "strsm",
286
+ "y",
287
+ ),
288
+ };
289
+ const zeroDesc = querySet
290
+ ? {
291
+ timestampWrites: {
292
+ querySet,
293
+ beginningOfPassWriteIndex: 0,
294
+ endOfPassWriteIndex: 1,
295
+ },
296
+ }
297
+ : undefined;
298
+ encodePass(
299
+ commandEncoder,
300
+ zeroPipeline,
301
+ zeroBindGroup,
302
+ zeroWgCount,
303
+ zeroDesc,
304
+ );
176
305
  } else {
177
306
  if (preScaleBindGroup) {
178
- encodePass(commandEncoder, scalarPipeline, preScaleBindGroup, calcWorkgroups(device, bScaleLen));
307
+ encodePass(
308
+ commandEncoder,
309
+ scalarPipeline,
310
+ preScaleBindGroup,
311
+ calcWorkgroups(device, bScaleLen),
312
+ );
179
313
  }
180
- const invertDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
181
- encodePass(commandEncoder, invertPipeline, invertBindGroup, { x: BLOCK_SIZE, y: numBlocks }, invertDesc);
314
+ const invertDesc = querySet
315
+ ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } }
316
+ : undefined;
317
+ encodePass(
318
+ commandEncoder,
319
+ invertPipeline,
320
+ invertBindGroup,
321
+ { x: BLOCK_SIZE, y: numBlocks },
322
+ invertDesc,
323
+ );
182
324
 
183
325
  for (let bi = 0; bi < blockStarts.length; bi++) {
184
- const blockStart = blockStarts[bi];
185
- const blockEnd = Math.min(blockStart + BLOCK_SIZE, aOrder);
186
- const blockLen = blockEnd - blockStart;
187
- const blockIndex = blockStart / BLOCK_SIZE;
188
- const isLastPass = bi === blockStarts.length - 1;
189
-
190
- // 1) gather B's current block into a tight scratch buffer.
191
- const gatherBParams = params(
192
- [
193
- { value: blockStart, type: "u32" },
194
- { value: blockLen, type: "u32" },
195
- { value: 0, type: "u32" },
196
- { value: otherLen, type: "u32" },
197
- { value: ldb, type: "u32" },
198
- { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
199
- { value: blockIsRow ? 1 : 0, type: "u32" },
200
- { value: 2, type: "u32" }, // gather
201
- ],
202
- "strsm-gather-B-params",
203
- );
204
- const gatherBBindGroup = createBindGroup(device, transferPipeline.getBindGroupLayout(0), [Bblock, BBuffer, gatherBParams]);
205
- encodePass(commandEncoder, transferPipeline, gatherBBindGroup, requireWorkgroups(device, "strsm", blockLen, otherLen));
206
-
207
- // 2) apply: Xblock := op(Ainv_block) @ Bblock (side='left'), or the
208
- // transpose-trick equivalent for side='right' (same trick strmm uses).
209
- {
210
- const mg = blockLen, ng = otherLen, kg = blockLen;
211
- const largeWgX = Math.ceil(ng / BN_LARGE), largeWgY = Math.ceil(mg / BM_LARGE);
212
- const useLarge = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
213
- const gemmPipeline = await getPipeline(device, useLarge ? "sgemm_large" : "sgemm_small");
214
- const applyParams = params(
326
+ const blockStart = blockStarts[bi];
327
+ const blockEnd = Math.min(blockStart + BLOCK_SIZE, aOrder);
328
+ const blockLen = blockEnd - blockStart;
329
+ const blockIndex = blockStart / BLOCK_SIZE;
330
+ const isLastPass = bi === blockStarts.length - 1;
331
+
332
+ // 1) gather B's current block into a tight scratch buffer.
333
+ const gatherBParams = params(
215
334
  [
216
- { value: mg, type: "u32" },
217
- { value: ng, type: "u32" },
218
- { value: kg, type: "u32" },
219
- { value: 1.0, type: "f32" }, // alpha already applied to B up front
220
- { value: 0.0, type: "f32" }, // beta — fresh output, no accumulation
221
- { value: BLOCK_SIZE, type: "u32" }, // ldX = Ainv's own dense stride
222
- { value: otherLen, type: "u32" }, // ldY = Bblock's own tight stride
223
- { value: otherLen, type: "u32" }, // ldc = Xblock's own tight stride
224
- { value: side === "right" ? 1 : 0, type: "u32" }, // transX: side='right' needs Ainv^T
225
- { value: 0, type: "u32" }, // transY: Bblock is always read as-is
335
+ { value: blockStart, type: "u32" },
336
+ { value: blockLen, type: "u32" },
337
+ { value: 0, type: "u32" },
338
+ { value: otherLen, type: "u32" },
339
+ { value: ldb, type: "u32" },
340
+ { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
341
+ { value: blockIsRow ? 1 : 0, type: "u32" },
342
+ { value: 2, type: "u32" }, // gather
226
343
  ],
227
- "strsm-apply-params",
344
+ "strsm-gather-B-params",
345
+ );
346
+ const gatherBBindGroup = createBindGroup(
347
+ device,
348
+ transferPipeline.getBindGroupLayout(0),
349
+ [Bblock, BBuffer, gatherBParams],
350
+ );
351
+ encodePass(
352
+ commandEncoder,
353
+ transferPipeline,
354
+ gatherBBindGroup,
355
+ requireWorkgroups(device, "strsm", blockLen, otherLen),
228
356
  );
229
- const ainvBlock = { buffer: AinvBuffer, offset: blockIndex * BLOCK_SIZE * BLOCK_SIZE * 4, size: BLOCK_SIZE * BLOCK_SIZE * 4 };
230
- const applyBindGroup = createBindGroup(device, gemmPipeline.getBindGroupLayout(0), [
231
- ainvBlock,
232
- vec4ViewBinding(device, ainvBlock),
233
- Bblock,
234
- vec4ViewBinding(device, Bblock),
235
- Xblock,
236
- applyParams,
237
- ]);
238
- const wg = useLarge
239
- ? { x: requireWorkgroupCount(device, largeWgX, "strsm", "x"), y: requireWorkgroupCount(device, largeWgY, "strsm", "y") }
240
- : { x: requireWorkgroupCount(device, Math.ceil(ng / BN_SMALL), "strsm", "x"), y: requireWorkgroupCount(device, Math.ceil(mg / BM_SMALL), "strsm", "y") };
241
- encodePass(commandEncoder, gemmPipeline, applyBindGroup, wg);
242
- }
243
357
 
244
- // 3) scatter the solved block back into B.
245
- const rangeStart = forward ? blockEnd : 0;
246
- const rangeEnd = forward ? aOrder : blockStart;
247
- const hasRemaining = rangeStart < rangeEnd;
248
- const scatterParams = params(
249
- [
250
- { value: blockStart, type: "u32" },
251
- { value: blockLen, type: "u32" },
252
- { value: 0, type: "u32" },
253
- { value: otherLen, type: "u32" },
254
- { value: ldb, type: "u32" },
255
- { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
256
- { value: blockIsRow ? 1 : 0, type: "u32" },
257
- { value: 0, type: "u32" }, // overwrite
258
- ],
259
- "strsm-scatter-params",
260
- );
261
- const scatterBindGroup = createBindGroup(device, transferPipeline.getBindGroupLayout(0), [Xblock, BBuffer, scatterParams]);
262
- const scatterDesc = isLastPass && !hasRemaining && querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
263
- encodePass(commandEncoder, transferPipeline, scatterBindGroup, requireWorkgroups(device, "strsm", blockLen, otherLen), scatterDesc);
358
+ // 2) apply: Xblock := op(Ainv_block) @ Bblock (side='left'), or the
359
+ // transpose-trick equivalent for side='right' (same trick strmm uses).
360
+ {
361
+ const mg = blockLen,
362
+ ng = otherLen,
363
+ kg = blockLen;
364
+ const largeWgX = Math.ceil(ng / BN_LARGE),
365
+ largeWgY = Math.ceil(mg / BM_LARGE);
366
+ const useLarge =
367
+ largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
368
+ const gemmPipeline = await getPipeline(
369
+ device,
370
+ useLarge ? "sgemm_large" : "sgemm_small",
371
+ );
372
+ const applyParams = params(
373
+ [
374
+ { value: mg, type: "u32" },
375
+ { value: ng, type: "u32" },
376
+ { value: kg, type: "u32" },
377
+ { value: 1.0, type: "f32" }, // alpha already applied to B up front
378
+ { value: 0.0, type: "f32" }, // beta — fresh output, no accumulation
379
+ { value: BLOCK_SIZE, type: "u32" }, // ldX = Ainv's own dense stride
380
+ { value: otherLen, type: "u32" }, // ldY = Bblock's own tight stride
381
+ { value: otherLen, type: "u32" }, // ldc = Xblock's own tight stride
382
+ { value: side === "right" ? 1 : 0, type: "u32" }, // transX: side='right' needs Ainv^T
383
+ { value: 0, type: "u32" }, // transY: Bblock is always read as-is
384
+ ],
385
+ "strsm-apply-params",
386
+ );
387
+ const ainvBlock = {
388
+ buffer: AinvBuffer,
389
+ offset: blockIndex * BLOCK_SIZE * BLOCK_SIZE * 4,
390
+ size: BLOCK_SIZE * BLOCK_SIZE * 4,
391
+ };
392
+ const applyBindGroup = createBindGroup(
393
+ device,
394
+ gemmPipeline.getBindGroupLayout(0),
395
+ [
396
+ ainvBlock,
397
+ vec4ViewBinding(device, ainvBlock),
398
+ Bblock,
399
+ vec4ViewBinding(device, Bblock),
400
+ Xblock,
401
+ applyParams,
402
+ ],
403
+ );
404
+ const wg = useLarge
405
+ ? {
406
+ x: requireWorkgroupCount(device, largeWgX, "strsm", "x"),
407
+ y: requireWorkgroupCount(device, largeWgY, "strsm", "y"),
408
+ }
409
+ : {
410
+ x: requireWorkgroupCount(
411
+ device,
412
+ Math.ceil(ng / BN_SMALL),
413
+ "strsm",
414
+ "x",
415
+ ),
416
+ y: requireWorkgroupCount(
417
+ device,
418
+ Math.ceil(mg / BM_SMALL),
419
+ "strsm",
420
+ "y",
421
+ ),
422
+ };
423
+ encodePass(commandEncoder, gemmPipeline, applyBindGroup, wg);
424
+ }
425
+
426
+ // 3) scatter the solved block back into B.
427
+ const rangeStart = forward ? blockEnd : 0;
428
+ const rangeEnd = forward ? aOrder : blockStart;
429
+ const hasRemaining = rangeStart < rangeEnd;
430
+ const scatterParams = params(
431
+ [
432
+ { value: blockStart, type: "u32" },
433
+ { value: blockLen, type: "u32" },
434
+ { value: 0, type: "u32" },
435
+ { value: otherLen, type: "u32" },
436
+ { value: ldb, type: "u32" },
437
+ { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
438
+ { value: blockIsRow ? 1 : 0, type: "u32" },
439
+ { value: 0, type: "u32" }, // overwrite
440
+ ],
441
+ "strsm-scatter-params",
442
+ );
443
+ const scatterBindGroup = createBindGroup(
444
+ device,
445
+ transferPipeline.getBindGroupLayout(0),
446
+ [Xblock, BBuffer, scatterParams],
447
+ );
448
+ const scatterDesc =
449
+ isLastPass && !hasRemaining && querySet
450
+ ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
451
+ : undefined;
452
+ encodePass(
453
+ commandEncoder,
454
+ transferPipeline,
455
+ scatterBindGroup,
456
+ requireWorkgroups(device, "strsm", blockLen, otherLen),
457
+ scatterDesc,
458
+ );
264
459
 
265
- // 4) trailing update: subtract this block's contribution from B.
266
- if (!hasRemaining) continue;
267
- const remCount = rangeEnd - rangeStart;
460
+ // 4) trailing update: subtract this block's contribution from B.
461
+ if (!hasRemaining) continue;
462
+ const remCount = rangeEnd - rangeStart;
268
463
 
269
- const gatherAParams = params(
270
- [
271
- { value: rangeStart, type: "u32" },
272
- { value: remCount, type: "u32" },
273
- { value: blockStart, type: "u32" },
274
- { value: blockLen, type: "u32" },
275
- { value: lda, type: "u32" },
276
- { value: transEffA === "transpose" ? 1 : 0, type: "u32" },
277
- { value: blockIsRow ? 1 : 0, type: "u32" },
278
- { value: 2, type: "u32" }, // gather
279
- ],
280
- "strsm-gather-A-params",
281
- );
282
- const gatherABindGroup = createBindGroup(device, transferPipeline.getBindGroupLayout(0), [Aoff, ABuffer, gatherAParams]);
283
- encodePass(commandEncoder, transferPipeline, gatherABindGroup, requireWorkgroups(device, "strsm", remCount, blockLen));
284
-
285
- {
286
- const mg = remCount, ng = otherLen, kg = blockLen;
287
- const largeWgX = Math.ceil(ng / BN_LARGE), largeWgY = Math.ceil(mg / BM_LARGE);
288
- const useLarge = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
289
- const gemmPipeline = await getPipeline(device, useLarge ? "sgemm_large" : "sgemm_small");
290
- const updateParams = params(
464
+ const gatherAParams = params(
291
465
  [
292
- { value: mg, type: "u32" },
293
- { value: ng, type: "u32" },
294
- { value: kg, type: "u32" },
295
- { value: 1.0, type: "f32" },
296
- { value: 0.0, type: "f32" },
297
- { value: blockLen, type: "u32" }, // ldX = Aoff's own tight stride
298
- { value: otherLen, type: "u32" }, // ldY = Xblock's own tight stride
299
- { value: otherLen, type: "u32" }, // ldc = delta's own tight stride
300
- { value: 0, type: "u32" }, // transX: Aoff already read in the right orientation
301
- { value: 0, type: "u32" }, // transY: Xblock read as-is
466
+ { value: rangeStart, type: "u32" },
467
+ { value: remCount, type: "u32" },
468
+ { value: blockStart, type: "u32" },
469
+ { value: blockLen, type: "u32" },
470
+ { value: lda, type: "u32" },
471
+ { value: transEffA === "transpose" ? 1 : 0, type: "u32" },
472
+ { value: blockIsRow ? 1 : 0, type: "u32" },
473
+ { value: 2, type: "u32" }, // gather
302
474
  ],
303
- "strsm-update-params",
475
+ "strsm-gather-A-params",
476
+ );
477
+ const gatherABindGroup = createBindGroup(
478
+ device,
479
+ transferPipeline.getBindGroupLayout(0),
480
+ [Aoff, ABuffer, gatherAParams],
481
+ );
482
+ encodePass(
483
+ commandEncoder,
484
+ transferPipeline,
485
+ gatherABindGroup,
486
+ requireWorkgroups(device, "strsm", remCount, blockLen),
304
487
  );
305
- const updateBindGroup = createBindGroup(device, gemmPipeline.getBindGroupLayout(0), [
306
- Aoff,
307
- vec4ViewBinding(device, Aoff),
308
- Xblock,
309
- vec4ViewBinding(device, Xblock),
310
- delta,
311
- updateParams,
312
- ]);
313
- const wg = useLarge
314
- ? { x: requireWorkgroupCount(device, largeWgX, "strsm", "x"), y: requireWorkgroupCount(device, largeWgY, "strsm", "y") }
315
- : { x: requireWorkgroupCount(device, Math.ceil(ng / BN_SMALL), "strsm", "x"), y: requireWorkgroupCount(device, Math.ceil(mg / BM_SMALL), "strsm", "y") };
316
- encodePass(commandEncoder, gemmPipeline, updateBindGroup, wg);
317
- }
318
488
 
319
- const scatterSubParams = params(
320
- [
321
- { value: rangeStart, type: "u32" },
322
- { value: remCount, type: "u32" },
323
- { value: 0, type: "u32" },
324
- { value: otherLen, type: "u32" },
325
- { value: ldb, type: "u32" },
326
- { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
327
- { value: blockIsRow ? 1 : 0, type: "u32" },
328
- { value: 1, type: "u32" }, // subtract
329
- ],
330
- "strsm-scatter-sub-params",
331
- );
332
- const scatterSubBindGroup = createBindGroup(device, transferPipeline.getBindGroupLayout(0), [delta, BBuffer, scatterSubParams]);
333
- const subDesc = isLastPass && querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
334
- encodePass(commandEncoder, transferPipeline, scatterSubBindGroup, requireWorkgroups(device, "strsm", remCount, otherLen), subDesc);
489
+ {
490
+ const mg = remCount,
491
+ ng = otherLen,
492
+ kg = blockLen;
493
+ const largeWgX = Math.ceil(ng / BN_LARGE),
494
+ largeWgY = Math.ceil(mg / BM_LARGE);
495
+ const useLarge =
496
+ largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
497
+ const gemmPipeline = await getPipeline(
498
+ device,
499
+ useLarge ? "sgemm_large" : "sgemm_small",
500
+ );
501
+ const updateParams = params(
502
+ [
503
+ { value: mg, type: "u32" },
504
+ { value: ng, type: "u32" },
505
+ { value: kg, type: "u32" },
506
+ { value: 1.0, type: "f32" },
507
+ { value: 0.0, type: "f32" },
508
+ { value: blockLen, type: "u32" }, // ldX = Aoff's own tight stride
509
+ { value: otherLen, type: "u32" }, // ldY = Xblock's own tight stride
510
+ { value: otherLen, type: "u32" }, // ldc = delta's own tight stride
511
+ { value: 0, type: "u32" }, // transX: Aoff already read in the right orientation
512
+ { value: 0, type: "u32" }, // transY: Xblock read as-is
513
+ ],
514
+ "strsm-update-params",
515
+ );
516
+ const updateBindGroup = createBindGroup(
517
+ device,
518
+ gemmPipeline.getBindGroupLayout(0),
519
+ [
520
+ Aoff,
521
+ vec4ViewBinding(device, Aoff),
522
+ Xblock,
523
+ vec4ViewBinding(device, Xblock),
524
+ delta,
525
+ updateParams,
526
+ ],
527
+ );
528
+ const wg = useLarge
529
+ ? {
530
+ x: requireWorkgroupCount(device, largeWgX, "strsm", "x"),
531
+ y: requireWorkgroupCount(device, largeWgY, "strsm", "y"),
532
+ }
533
+ : {
534
+ x: requireWorkgroupCount(
535
+ device,
536
+ Math.ceil(ng / BN_SMALL),
537
+ "strsm",
538
+ "x",
539
+ ),
540
+ y: requireWorkgroupCount(
541
+ device,
542
+ Math.ceil(mg / BM_SMALL),
543
+ "strsm",
544
+ "y",
545
+ ),
546
+ };
547
+ encodePass(commandEncoder, gemmPipeline, updateBindGroup, wg);
548
+ }
549
+
550
+ const scatterSubParams = params(
551
+ [
552
+ { value: rangeStart, type: "u32" },
553
+ { value: remCount, type: "u32" },
554
+ { value: 0, type: "u32" },
555
+ { value: otherLen, type: "u32" },
556
+ { value: ldb, type: "u32" },
557
+ { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
558
+ { value: blockIsRow ? 1 : 0, type: "u32" },
559
+ { value: 1, type: "u32" }, // subtract
560
+ ],
561
+ "strsm-scatter-sub-params",
562
+ );
563
+ const scatterSubBindGroup = createBindGroup(
564
+ device,
565
+ transferPipeline.getBindGroupLayout(0),
566
+ [delta, BBuffer, scatterSubParams],
567
+ );
568
+ const subDesc =
569
+ isLastPass && querySet
570
+ ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } }
571
+ : undefined;
572
+ encodePass(
573
+ commandEncoder,
574
+ transferPipeline,
575
+ scatterSubBindGroup,
576
+ requireWorkgroups(device, "strsm", remCount, otherLen),
577
+ subDesc,
578
+ );
335
579
  }
336
580
  }
337
581
 
338
582
  const ts = resolveTimestamp(device, commandEncoder, querySet);
339
- const readBuffer = BIsGpu ? null : stageReadback(device, commandEncoder, BBuffer);
583
+ const readBuffer = BIsGpu
584
+ ? null
585
+ : stageReadback(device, commandEncoder, BBuffer);
340
586
 
341
587
  submit(device, commandEncoder);
342
588