wgblas 2.0.0 → 2.2.0

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 (124) hide show
  1. package/README.md +20 -18
  2. package/dist/wgblas.browser.js +2172 -1174
  3. package/index.d.mts +49 -44
  4. package/index.mjs +11 -0
  5. package/package.json +133 -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 +126 -17
  12. package/src/classes/GpuVector.d.mts +36 -40
  13. package/src/classes/GpuVector.mjs +66 -11
  14. package/src/cscal/cscal.d.mts +47 -0
  15. package/src/cscal/cscal.mjs +98 -0
  16. package/src/dasum/dasum.d.mts +4 -4
  17. package/src/dasum/dasum.mjs +38 -20
  18. package/src/daxpy/daxpy.d.mts +56 -0
  19. package/src/daxpy/daxpy.mjs +150 -0
  20. package/src/dcopy/dcopy.d.mts +52 -0
  21. package/src/dcopy/dcopy.mjs +140 -0
  22. package/src/ddot/ddot.d.mts +62 -0
  23. package/src/ddot/ddot.mjs +184 -0
  24. package/src/devdocs.mjs +13 -0
  25. package/src/dnrm2/dnrm2.d.mts +50 -0
  26. package/src/dnrm2/dnrm2.mjs +189 -0
  27. package/src/drot/drot.d.mts +67 -0
  28. package/src/drot/drot.mjs +170 -0
  29. package/src/drotm/drotm.d.mts +67 -0
  30. package/src/drotm/drotm.mjs +171 -0
  31. package/src/dscal/dscal.d.mts +52 -0
  32. package/src/dscal/dscal.mjs +119 -0
  33. package/src/dswap/dswap.d.mts +57 -0
  34. package/src/dswap/dswap.mjs +155 -0
  35. package/src/idamax/idamax.d.mts +20 -2
  36. package/src/idamax/idamax.mjs +56 -24
  37. package/src/init.mjs +117 -56
  38. package/src/isamax/isamax.d.mts +20 -2
  39. package/src/isamax/isamax.mjs +21 -16
  40. package/src/random/random.d.mts +37 -39
  41. package/src/random/random.mjs +39 -7
  42. package/src/sasum/sasum.d.mts +2 -2
  43. package/src/sasum/sasum.mjs +20 -16
  44. package/src/saxpy/saxpy.d.mts +2 -2
  45. package/src/saxpy/saxpy.mjs +14 -11
  46. package/src/scopy/scopy.d.mts +2 -2
  47. package/src/scopy/scopy.mjs +13 -9
  48. package/src/sdot/sdot.d.mts +2 -2
  49. package/src/sdot/sdot.mjs +21 -17
  50. package/src/sgemm/sgemm.d.mts +2 -2
  51. package/src/sgemm/sgemm.mjs +109 -40
  52. package/src/sgemmtr/sgemmtr.d.mts +3 -2
  53. package/src/sgemmtr/sgemmtr.mjs +98 -40
  54. package/src/sgemv/sgemv.d.mts +2 -2
  55. package/src/sgemv/sgemv.mjs +69 -41
  56. package/src/sger/sger.d.mts +2 -2
  57. package/src/sger/sger.mjs +43 -19
  58. package/src/shaders/__test_pipeline_a.wgsl +3 -0
  59. package/src/shaders/__test_pipeline_b.wgsl +2 -0
  60. package/src/shaders/cscal.wgsl +33 -0
  61. package/src/shaders/daxpy.wgsl +66 -0
  62. package/src/shaders/dcopy.wgsl +34 -0
  63. package/src/shaders/ddot.wgsl +106 -0
  64. package/src/shaders/dnrm2.wgsl +167 -0
  65. package/src/shaders/drot.wgsl +81 -0
  66. package/src/shaders/drotm.wgsl +99 -0
  67. package/src/shaders/dscal.wgsl +60 -0
  68. package/src/shaders/dswap.wgsl +38 -0
  69. package/src/shaders/f64/utils/add.wgsl +6 -0
  70. package/src/shaders/f64/utils/divide.wgsl +45 -0
  71. package/src/shaders/f64/utils/multiply.wgsl +19 -10
  72. package/src/shaders/f64/utils/sqrt.wgsl +45 -0
  73. package/src/shaders/index.mjs +233 -14
  74. package/src/shaders/reduction/scaledSum.wgsl +65 -0
  75. package/src/shaders/reduction/scaledSumF64.wgsl +93 -0
  76. package/src/shaders/sgemm_large.wgsl +107 -18
  77. package/src/shaders/sgemm_small.wgsl +115 -15
  78. package/src/shaders/sgemmtr_large.wgsl +4 -1
  79. package/src/shaders/sgemmtr_small.wgsl +4 -1
  80. package/src/shaders/sgemv_n.wgsl +3 -1
  81. package/src/shaders/sgemv_t.wgsl +3 -1
  82. package/src/shaders/snrm2.wgsl +72 -23
  83. package/src/shaders/ssymv.wgsl +3 -1
  84. package/src/snrm2/snrm2.d.mts +2 -2
  85. package/src/snrm2/snrm2.mjs +41 -23
  86. package/src/srot/srot.d.mts +2 -4
  87. package/src/srot/srot.mjs +16 -11
  88. package/src/srotm/srotm.d.mts +2 -4
  89. package/src/srotm/srotm.mjs +17 -11
  90. package/src/sscal/sscal.d.mts +3 -3
  91. package/src/sscal/sscal.mjs +14 -12
  92. package/src/sswap/sswap.d.mts +2 -2
  93. package/src/sswap/sswap.mjs +18 -10
  94. package/src/ssymm/ssymm.d.mts +5 -4
  95. package/src/ssymm/ssymm.mjs +150 -54
  96. package/src/ssymv/ssymv.d.mts +2 -2
  97. package/src/ssymv/ssymv.mjs +47 -26
  98. package/src/ssyr/ssyr.d.mts +2 -2
  99. package/src/ssyr/ssyr.mjs +38 -17
  100. package/src/ssyr2/ssyr2.d.mts +2 -2
  101. package/src/ssyr2/ssyr2.mjs +48 -21
  102. package/src/ssyr2k/ssyr2k.d.mts +3 -2
  103. package/src/ssyr2k/ssyr2k.mjs +140 -62
  104. package/src/ssyrk/ssyrk.d.mts +3 -2
  105. package/src/ssyrk/ssyrk.mjs +91 -39
  106. package/src/strmm/strmm.d.mts +5 -4
  107. package/src/strmm/strmm.mjs +174 -60
  108. package/src/strmv/strmv.d.mts +2 -2
  109. package/src/strmv/strmv.mjs +42 -20
  110. package/src/strsm/strsm.d.mts +6 -4
  111. package/src/strsm/strsm.mjs +438 -174
  112. package/src/strsv/strsv.d.mts +5 -3
  113. package/src/strsv/strsv.mjs +89 -34
  114. package/src/util/benchmark.mjs +9 -9
  115. package/src/util/bindgroup.mjs +1 -3
  116. package/src/util/buffer.mjs +139 -24
  117. package/src/util/complex.mjs +87 -0
  118. package/src/util/compute.mjs +19 -16
  119. package/src/util/constants.mjs +57 -0
  120. package/src/util/device.mjs +49 -0
  121. package/src/util/pipeline.mjs +44 -10
  122. package/src/util/workgroup.mjs +72 -7
  123. package/src/shaders/browser-shaders.mjs +0 -81
  124. package/src/shaders/f64add.wgsl +0 -281
@@ -4,33 +4,54 @@ import {
4
4
  createStorageBuffer,
5
5
  stageReadback,
6
6
  destroyBuffers,
7
+ vec4ViewBinding,
7
8
  } from "../util/buffer.mjs";
8
9
  import { createBindGroup } from "../util/bindgroup.mjs";
9
10
  import { beginTimedEncoder, encodePass, submit } from "../util/compute.mjs";
10
11
  import { extractResult } from "../util/result.mjs";
11
12
  import { resolveTimestamp, extractTimestamp } from "../util/benchmark.mjs";
12
13
  import { getPipeline } from "../util/pipeline.mjs";
13
- import { calcWorkgroups } from "../util/workgroup.mjs";
14
+ import {
15
+ calcWorkgroups,
16
+ requireWorkgroups,
17
+ requireWorkgroupCount,
18
+ } from "../util/workgroup.mjs";
14
19
  import { GpuMatrix } from "../classes/GpuMatrix.mjs";
15
-
16
- const BLOCK_SIZE = 64; // must match strsv_invert_block.wgsl's own constant
17
- const BM_SMALL = 32, BN_SMALL = 32; // sgemm_small.wgsl's block tile
18
- const BM_LARGE = 64, BN_LARGE = 64; // sgemm_large.wgsl's block tile
19
- const LARGE_TILE_WORKGROUP_THRESHOLD = 36; // same threshold sgemm/ssymm/strmm use
20
+ import {
21
+ BM_SMALL,
22
+ BN_SMALL,
23
+ BM_LARGE,
24
+ BN_LARGE,
25
+ LARGE_TILE_WORKGROUP_THRESHOLD,
26
+ } from "../util/constants.mjs";
27
+ import { BLOCK_SIZE } from "../util/constants.mjs";
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);
54
+ requireSameDevice(device, "strsm", { A, B });
34
55
  if (side !== "left" && side !== "right")
35
56
  throw new Error("side must be 'left' or 'right'.");
36
57
  if (uplo !== "lower" && uplo !== "upper")
@@ -41,11 +62,15 @@ export async function strsm(
41
62
  throw new Error("diag must be 'unit' or 'non-unit'.");
42
63
  if (layout !== "row-major" && layout !== "column-major")
43
64
  throw new Error("layout must be 'row-major' or 'column-major'.");
44
- if (typeof alpha !== "number")
45
- throw new Error("alpha must be a number.");
65
+ if (typeof alpha !== "number") throw new Error("alpha must be a number.");
46
66
  if (Number.isNaN(alpha)) throw new Error("alpha must not be NaN.");
47
67
  if (!Number.isFinite(alpha)) throw new Error("alpha must be finite.");
48
- 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
+ )
49
74
  throw new Error("m, n, lda, and ldb must be integers.");
50
75
  if (!AIsGpu && !(A instanceof Float32Array))
51
76
  throw new Error("A must be a Float32Array or GpuMatrix.");
@@ -61,30 +86,51 @@ export async function strsm(
61
86
 
62
87
  // A: triangular, order = m (side='left') or n (side='right').
63
88
  const aOrder = side === "left" ? m : n;
64
- 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") + ".");
65
91
  if (AIsGpu) {
66
- if (lda !== A.lda) throw new Error("lda must match A.lda when A is a GpuMatrix.");
67
- 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.");
68
96
  } else if (A.length < (aOrder - 1) * lda + aOrder) {
69
- 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
+ );
70
100
  }
71
101
 
72
102
  // B: always m x n, overwritten in place with the same ldb.
73
103
  const bOuter = effLayoutB === "column-major" ? n : m;
74
104
  const bInner = effLayoutB === "column-major" ? m : n;
75
105
  if (ldb < bInner)
76
- 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
+ );
77
109
  if (BIsGpu) {
78
- if (ldb !== B.lda) throw new Error("ldb must match B.lda when B is a GpuMatrix.");
79
- 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.");
80
114
  } else if (B.length < (bOuter - 1) * ldb + bInner) {
81
- 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
+ );
82
118
  }
83
119
 
84
120
  // A isn't symmetric: column-major = genuine transpose, so flip transA;
85
121
  // transposing also swaps which triangle looks stored, so flip uplo too.
86
- const uploEffA = effLayoutA === "column-major" ? (uplo === "lower" ? "upper" : "lower") : uplo;
87
- 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;
88
134
 
89
135
  const otherLen = side === "left" ? n : m;
90
136
  const blockIsRow = side === "left";
@@ -103,19 +149,22 @@ export async function strsm(
103
149
  const transferPipeline = await getPipeline(device, "block_transfer");
104
150
  const scalarPipeline = await getPipeline(device, "sscal");
105
151
 
106
- const ABuffer = AIsGpu ? A._buf : uploadBuffer(A, "strsm-A", false);
107
- const BBuffer = BIsGpu ? B._buf : uploadBuffer(B, "strsm-B", true);
108
- const AinvBuffer = createStorageBuffer(numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4, "strsm-Ainv");
152
+ // Null-init here and allocate inside the try below, so a throw partway
153
+ // through the sequence still reaches finally with every handle visible
154
+ // (strsv.mjs is the reference for this pattern).
155
+ let ABuffer = null;
156
+ let BBuffer = null;
157
+ let AinvBuffer = null;
109
158
 
110
159
  const paramsBuffers = [];
111
160
  const scratchBuffers = [];
112
161
  function scratch(size, label) {
113
- const buf = createStorageBuffer(size, label);
162
+ const buf = createStorageBuffer(device, size, label);
114
163
  scratchBuffers.push(buf);
115
164
  return buf;
116
165
  }
117
166
  function params(entries, label) {
118
- const buf = createParamsBuffer(entries, label);
167
+ const buf = createParamsBuffer(device, entries, label);
119
168
  paramsBuffers.push(buf);
120
169
  return buf;
121
170
  }
@@ -126,9 +175,19 @@ export async function strsm(
126
175
  const bScaleLen = (bOuter - 1) * ldb + bInner;
127
176
 
128
177
  try {
178
+ ABuffer = AIsGpu ? A._buf : uploadBuffer(device, A, "strsm-A", false);
179
+ BBuffer = BIsGpu ? B._buf : uploadBuffer(device, B, "strsm-B", true);
180
+ AinvBuffer = createStorageBuffer(
181
+ device,
182
+ numBlocks * BLOCK_SIZE * BLOCK_SIZE * 4,
183
+ "strsm-Ainv",
184
+ );
185
+
129
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.
130
189
  let preScaleBindGroup = null;
131
- if (alpha !== 1.0) {
190
+ if (alpha !== 1.0 && alpha !== 0) {
132
191
  const scaleParams = params(
133
192
  [
134
193
  { value: bScaleLen, type: "u32" },
@@ -137,7 +196,11 @@ export async function strsm(
137
196
  ],
138
197
  "strsm-scale-params",
139
198
  );
140
- preScaleBindGroup = createBindGroup(scalarPipeline.getBindGroupLayout(0), [BBuffer, scaleParams]);
199
+ preScaleBindGroup = createBindGroup(
200
+ device,
201
+ scalarPipeline.getBindGroupLayout(0),
202
+ [BBuffer, scaleParams],
203
+ );
141
204
  }
142
205
 
143
206
  // Every diagonal block's inverse, fully parallel, one dispatch, unchanged.
@@ -151,7 +214,11 @@ export async function strsm(
151
214
  ],
152
215
  "strsm-invert-params",
153
216
  );
154
- const invertBindGroup = createBindGroup(invertPipeline.getBindGroupLayout(0), [ABuffer, AinvBuffer, invertParams]);
217
+ const invertBindGroup = createBindGroup(
218
+ device,
219
+ invertPipeline.getBindGroupLayout(0),
220
+ [ABuffer, AinvBuffer, invertParams],
221
+ );
155
222
 
156
223
  // Reusable scratch buffers, sized for the worst case, bound at offset 0.
157
224
  const Bblock = scratch(BLOCK_SIZE * otherLen * 4, "strsm-Bblock");
@@ -159,168 +226,365 @@ export async function strsm(
159
226
  const Aoff = scratch(aOrder * BLOCK_SIZE * 4, "strsm-Aoff");
160
227
  const delta = scratch(aOrder * otherLen * 4, "strsm-delta");
161
228
 
162
- const { commandEncoder, querySet } = beginTimedEncoder();
229
+ const { commandEncoder, querySet } = beginTimedEncoder(device);
163
230
 
164
231
  if (alpha === 0) {
165
- // BLAS: alpha=0 means A is not referenced — skip straight to B:=0.
166
- const zeroDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0, endOfPassWriteIndex: 1 } } : undefined;
167
- encodePass(commandEncoder, scalarPipeline, preScaleBindGroup, calcWorkgroups(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
+ );
168
305
  } else {
169
306
  if (preScaleBindGroup) {
170
- encodePass(commandEncoder, scalarPipeline, preScaleBindGroup, calcWorkgroups(bScaleLen));
307
+ encodePass(
308
+ commandEncoder,
309
+ scalarPipeline,
310
+ preScaleBindGroup,
311
+ calcWorkgroups(device, bScaleLen),
312
+ );
171
313
  }
172
- const invertDesc = querySet ? { timestampWrites: { querySet, beginningOfPassWriteIndex: 0 } } : undefined;
173
- 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
+ );
174
324
 
175
325
  for (let bi = 0; bi < blockStarts.length; bi++) {
176
- const blockStart = blockStarts[bi];
177
- const blockEnd = Math.min(blockStart + BLOCK_SIZE, aOrder);
178
- const blockLen = blockEnd - blockStart;
179
- const blockIndex = blockStart / BLOCK_SIZE;
180
- const isLastPass = bi === blockStarts.length - 1;
181
-
182
- // 1) gather B's current block into a tight scratch buffer.
183
- const gatherBParams = params(
184
- [
185
- { value: blockStart, type: "u32" },
186
- { value: blockLen, type: "u32" },
187
- { value: 0, type: "u32" },
188
- { value: otherLen, type: "u32" },
189
- { value: ldb, type: "u32" },
190
- { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
191
- { value: blockIsRow ? 1 : 0, type: "u32" },
192
- { value: 2, type: "u32" }, // gather
193
- ],
194
- "strsm-gather-B-params",
195
- );
196
- const gatherBBindGroup = createBindGroup(transferPipeline.getBindGroupLayout(0), [Bblock, BBuffer, gatherBParams]);
197
- encodePass(commandEncoder, transferPipeline, gatherBBindGroup, calcWorkgroups(blockLen, otherLen));
198
-
199
- // 2) apply: Xblock := op(Ainv_block) @ Bblock (side='left'), or the
200
- // transpose-trick equivalent for side='right' (same trick strmm uses).
201
- {
202
- const mg = blockLen, ng = otherLen, kg = blockLen;
203
- const largeWgX = Math.ceil(ng / BN_LARGE), largeWgY = Math.ceil(mg / BM_LARGE);
204
- const useLarge = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
205
- const gemmPipeline = await getPipeline(device, useLarge ? "sgemm_large" : "sgemm_small");
206
- 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(
207
334
  [
208
- { value: mg, type: "u32" },
209
- { value: ng, type: "u32" },
210
- { value: kg, type: "u32" },
211
- { value: 1.0, type: "f32" }, // alpha already applied to B up front
212
- { value: 0.0, type: "f32" }, // beta — fresh output, no accumulation
213
- { value: BLOCK_SIZE, type: "u32" }, // ldX = Ainv's own dense stride
214
- { value: otherLen, type: "u32" }, // ldY = Bblock's own tight stride
215
- { value: otherLen, type: "u32" }, // ldc = Xblock's own tight stride
216
- { value: side === "right" ? 1 : 0, type: "u32" }, // transX: side='right' needs Ainv^T
217
- { 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
218
343
  ],
219
- "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),
220
356
  );
221
- const applyBindGroup = createBindGroup(gemmPipeline.getBindGroupLayout(0), [
222
- { buffer: AinvBuffer, offset: blockIndex * BLOCK_SIZE * BLOCK_SIZE * 4, size: BLOCK_SIZE * BLOCK_SIZE * 4 },
223
- Bblock,
224
- Xblock,
225
- applyParams,
226
- ]);
227
- const wg = useLarge
228
- ? { x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension), y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension) }
229
- : { x: Math.min(Math.ceil(ng / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension), y: Math.min(Math.ceil(mg / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension) };
230
- encodePass(commandEncoder, gemmPipeline, applyBindGroup, wg);
231
- }
232
357
 
233
- // 3) scatter the solved block back into B.
234
- const rangeStart = forward ? blockEnd : 0;
235
- const rangeEnd = forward ? aOrder : blockStart;
236
- const hasRemaining = rangeStart < rangeEnd;
237
- const scatterParams = params(
238
- [
239
- { value: blockStart, type: "u32" },
240
- { value: blockLen, type: "u32" },
241
- { value: 0, type: "u32" },
242
- { value: otherLen, type: "u32" },
243
- { value: ldb, type: "u32" },
244
- { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
245
- { value: blockIsRow ? 1 : 0, type: "u32" },
246
- { value: 0, type: "u32" }, // overwrite
247
- ],
248
- "strsm-scatter-params",
249
- );
250
- const scatterBindGroup = createBindGroup(transferPipeline.getBindGroupLayout(0), [Xblock, BBuffer, scatterParams]);
251
- const scatterDesc = isLastPass && !hasRemaining && querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
252
- encodePass(commandEncoder, transferPipeline, scatterBindGroup, calcWorkgroups(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
+ );
253
459
 
254
- // 4) trailing update: subtract this block's contribution from B.
255
- if (!hasRemaining) continue;
256
- const remCount = rangeEnd - rangeStart;
460
+ // 4) trailing update: subtract this block's contribution from B.
461
+ if (!hasRemaining) continue;
462
+ const remCount = rangeEnd - rangeStart;
257
463
 
258
- const gatherAParams = params(
259
- [
260
- { value: rangeStart, type: "u32" },
261
- { value: remCount, type: "u32" },
262
- { value: blockStart, type: "u32" },
263
- { value: blockLen, type: "u32" },
264
- { value: lda, type: "u32" },
265
- { value: transEffA === "transpose" ? 1 : 0, type: "u32" },
266
- { value: blockIsRow ? 1 : 0, type: "u32" },
267
- { value: 2, type: "u32" }, // gather
268
- ],
269
- "strsm-gather-A-params",
270
- );
271
- const gatherABindGroup = createBindGroup(transferPipeline.getBindGroupLayout(0), [Aoff, ABuffer, gatherAParams]);
272
- encodePass(commandEncoder, transferPipeline, gatherABindGroup, calcWorkgroups(remCount, blockLen));
273
-
274
- {
275
- const mg = remCount, ng = otherLen, kg = blockLen;
276
- const largeWgX = Math.ceil(ng / BN_LARGE), largeWgY = Math.ceil(mg / BM_LARGE);
277
- const useLarge = largeWgX * largeWgY >= LARGE_TILE_WORKGROUP_THRESHOLD;
278
- const gemmPipeline = await getPipeline(device, useLarge ? "sgemm_large" : "sgemm_small");
279
- const updateParams = params(
464
+ const gatherAParams = params(
280
465
  [
281
- { value: mg, type: "u32" },
282
- { value: ng, type: "u32" },
283
- { value: kg, type: "u32" },
284
- { value: 1.0, type: "f32" },
285
- { value: 0.0, type: "f32" },
286
- { value: blockLen, type: "u32" }, // ldX = Aoff's own tight stride
287
- { value: otherLen, type: "u32" }, // ldY = Xblock's own tight stride
288
- { value: otherLen, type: "u32" }, // ldc = delta's own tight stride
289
- { value: 0, type: "u32" }, // transX: Aoff already read in the right orientation
290
- { 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
291
474
  ],
292
- "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),
293
487
  );
294
- const updateBindGroup = createBindGroup(gemmPipeline.getBindGroupLayout(0), [Aoff, Xblock, delta, updateParams]);
295
- const wg = useLarge
296
- ? { x: Math.min(largeWgX, device.limits.maxComputeWorkgroupsPerDimension), y: Math.min(largeWgY, device.limits.maxComputeWorkgroupsPerDimension) }
297
- : { x: Math.min(Math.ceil(ng / BN_SMALL), device.limits.maxComputeWorkgroupsPerDimension), y: Math.min(Math.ceil(mg / BM_SMALL), device.limits.maxComputeWorkgroupsPerDimension) };
298
- encodePass(commandEncoder, gemmPipeline, updateBindGroup, wg);
299
- }
300
488
 
301
- const scatterSubParams = params(
302
- [
303
- { value: rangeStart, type: "u32" },
304
- { value: remCount, type: "u32" },
305
- { value: 0, type: "u32" },
306
- { value: otherLen, type: "u32" },
307
- { value: ldb, type: "u32" },
308
- { value: effLayoutB === "column-major" ? 1 : 0, type: "u32" },
309
- { value: blockIsRow ? 1 : 0, type: "u32" },
310
- { value: 1, type: "u32" }, // subtract
311
- ],
312
- "strsm-scatter-sub-params",
313
- );
314
- const scatterSubBindGroup = createBindGroup(transferPipeline.getBindGroupLayout(0), [delta, BBuffer, scatterSubParams]);
315
- const subDesc = isLastPass && querySet ? { timestampWrites: { querySet, endOfPassWriteIndex: 1 } } : undefined;
316
- encodePass(commandEncoder, transferPipeline, scatterSubBindGroup, calcWorkgroups(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
+ );
317
579
  }
318
580
  }
319
581
 
320
- const ts = resolveTimestamp(commandEncoder, querySet);
321
- const readBuffer = BIsGpu ? null : stageReadback(commandEncoder, BBuffer);
582
+ const ts = resolveTimestamp(device, commandEncoder, querySet);
583
+ const readBuffer = BIsGpu
584
+ ? null
585
+ : stageReadback(device, commandEncoder, BBuffer);
322
586
 
323
- submit(commandEncoder);
587
+ submit(device, commandEncoder);
324
588
 
325
589
  const gpuTimeMs = await extractTimestamp(ts);
326
590
 
@@ -333,9 +597,9 @@ export async function strsm(
333
597
  if (gpuTimeMs !== undefined) return { B: result, gpuTimeMs };
334
598
  return { B: result };
335
599
  } finally {
336
- if (!AIsGpu) destroyBuffers(ABuffer);
337
- if (!BIsGpu) destroyBuffers(BBuffer);
338
- destroyBuffers(AinvBuffer);
600
+ if (!AIsGpu && ABuffer) destroyBuffers(ABuffer);
601
+ if (!BIsGpu && BBuffer) destroyBuffers(BBuffer);
602
+ if (AinvBuffer) destroyBuffers(AinvBuffer);
339
603
  destroyBuffers(scratchBuffers);
340
604
  destroyBuffers(paramsBuffers);
341
605
  }