|
| 1 | +import { GUI } from 'dat.gui'; |
| 2 | +import packedWGSL from './packed.wgsl'; |
| 3 | +import { quitIfWebGPUNotAvailableOrMissingFeatures } from '../util'; |
| 4 | + |
| 5 | +type vec4i = readonly [number, number, number, number]; |
| 6 | +// Pack four signed 8-bit components into a u32, low byte first. |
| 7 | +function pack4xI8([x, y, z, w]: vec4i): number { |
| 8 | + // `&` operator applies sign extension to i32 before operating. |
| 9 | + // `>>> 0` converts the final i32 to u32. |
| 10 | + /*prettier-ignore*/ |
| 11 | + return ((x & 0xff) | |
| 12 | + ((y & 0xff) << 8) | |
| 13 | + ((z & 0xff) << 16) | |
| 14 | + ((w & 0xff) << 24)) >>> 0; |
| 15 | +} |
| 16 | + |
| 17 | +const outputElement = document.querySelector('#output') as HTMLElement; |
| 18 | +if ( |
| 19 | + !navigator.gpu?.wgslLanguageFeatures.has('packed_4x8_integer_dot_product') |
| 20 | +) { |
| 21 | + result.textContent = |
| 22 | + "This sample requires the WGSL language feature 'packed_4x8_integer_dot_product'."; |
| 23 | +} else { |
| 24 | + const adapter = await navigator.gpu.requestAdapter({ |
| 25 | + featureLevel: 'compatibility', |
| 26 | + }); |
| 27 | + const device = await adapter?.requestDevice(); |
| 28 | + quitIfWebGPUNotAvailableOrMissingFeatures(adapter, device); |
| 29 | + |
| 30 | + const kInputSize = 2 * Uint32Array.BYTES_PER_ELEMENT; |
| 31 | + const inputBuffer = device.createBuffer({ |
| 32 | + size: kInputSize, |
| 33 | + usage: GPUBufferUsage.COPY_DST | GPUBufferUsage.STORAGE, |
| 34 | + }); |
| 35 | + |
| 36 | + const kOutputSize = Int32Array.BYTES_PER_ELEMENT; |
| 37 | + const outputBuffer = device.createBuffer({ |
| 38 | + size: kOutputSize, |
| 39 | + usage: GPUBufferUsage.STORAGE | GPUBufferUsage.COPY_SRC, |
| 40 | + }); |
| 41 | + const readbackBuffer = device.createBuffer({ |
| 42 | + size: kOutputSize, |
| 43 | + usage: GPUBufferUsage.COPY_DST | GPUBufferUsage.MAP_READ, |
| 44 | + }); |
| 45 | + |
| 46 | + const pipeline = await device.createComputePipelineAsync({ |
| 47 | + layout: 'auto', |
| 48 | + compute: { module: device.createShaderModule({ code: packedWGSL }) }, |
| 49 | + }); |
| 50 | + const bindGroup = device.createBindGroup({ |
| 51 | + layout: pipeline.getBindGroupLayout(0), |
| 52 | + entries: [ |
| 53 | + { binding: 0, resource: { buffer: inputBuffer } }, |
| 54 | + { binding: 1, resource: { buffer: outputBuffer } }, |
| 55 | + ], |
| 56 | + }); |
| 57 | + |
| 58 | + async function updateResult() { |
| 59 | + // If an update is still in progress just wait until it's done. |
| 60 | + if (readbackBuffer.mapState !== 'unmapped') { |
| 61 | + setTimeout(updateResult, 0); |
| 62 | + return; |
| 63 | + } |
| 64 | + |
| 65 | + const lhs = [settings.lhs0, settings.lhs1, settings.lhs2, settings.lhs3]; |
| 66 | + const rhs = [settings.rhs0, settings.rhs1, settings.rhs2, settings.rhs3]; |
| 67 | + |
| 68 | + device.queue.writeBuffer( |
| 69 | + inputBuffer, |
| 70 | + 0, |
| 71 | + new Uint32Array([lhs, rhs].map(pack4xI8)) |
| 72 | + ); |
| 73 | + const encoder = device.createCommandEncoder(); |
| 74 | + const pass = encoder.beginComputePass(); |
| 75 | + pass.setPipeline(pipeline); |
| 76 | + pass.setBindGroup(0, bindGroup); |
| 77 | + pass.dispatchWorkgroups(1); |
| 78 | + pass.end(); |
| 79 | + encoder.copyBufferToBuffer(outputBuffer, 0, readbackBuffer, 0, kOutputSize); |
| 80 | + device.queue.submit([encoder.finish()]); |
| 81 | + |
| 82 | + await readbackBuffer.mapAsync(GPUMapMode.READ); |
| 83 | + const result = new Int32Array(readbackBuffer.getMappedRange())[0]; |
| 84 | + |
| 85 | + // Result should be the same in JS, show that for comparison. |
| 86 | + const expected = |
| 87 | + lhs[0] * rhs[0] + lhs[1] * rhs[1] + lhs[2] * rhs[2] + lhs[3] * rhs[3]; |
| 88 | + |
| 89 | + const lhsStr = `[${lhs |
| 90 | + .map((x) => x.toString().padStart(4)) |
| 91 | + .join(', ')}] (0x${pack4xI8(lhs).toString(16).padStart(8, '0')})`; |
| 92 | + const rhsStr = `[${rhs |
| 93 | + .map((x) => x.toString().padStart(4)) |
| 94 | + .join(', ')}] (0x${pack4xI8(rhs).toString(16).padStart(8, '0')})`; |
| 95 | + const outStr = result.toString().padStart(6); |
| 96 | + const expStr = expected.toString().padStart(6); |
| 97 | + outputElement.textContent = ` |
| 98 | +
|
| 99 | +WGSL dot4I8Packed of ${lhsStr} |
| 100 | + by ${rhsStr} gave ${outStr} (JS gave ${expStr})`; |
| 101 | + |
| 102 | + readbackBuffer.unmap(); |
| 103 | + } |
| 104 | + |
| 105 | + const settings = { |
| 106 | + lhs0: 1, |
| 107 | + lhs1: -2, |
| 108 | + lhs2: 3, |
| 109 | + lhs3: -4, |
| 110 | + rhs0: -5, |
| 111 | + rhs1: 6, |
| 112 | + rhs2: -7, |
| 113 | + rhs3: 8, |
| 114 | + }; |
| 115 | + const gui = new GUI(); |
| 116 | + gui.add(settings, 'lhs0', -127, 128, 1).onChange(updateResult); |
| 117 | + gui.add(settings, 'lhs1', -127, 128, 1).onChange(updateResult); |
| 118 | + gui.add(settings, 'lhs2', -127, 128, 1).onChange(updateResult); |
| 119 | + gui.add(settings, 'lhs3', -127, 128, 1).onChange(updateResult); |
| 120 | + gui.add(settings, 'rhs0', -127, 128, 1).onChange(updateResult); |
| 121 | + gui.add(settings, 'rhs1', -127, 128, 1).onChange(updateResult); |
| 122 | + gui.add(settings, 'rhs2', -127, 128, 1).onChange(updateResult); |
| 123 | + gui.add(settings, 'rhs3', -127, 128, 1).onChange(updateResult); |
| 124 | + updateResult(); |
| 125 | +} |
0 commit comments