Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 27 additions & 2 deletions internal/backend/webgpu/backend.go
Original file line number Diff line number Diff line change
Expand Up @@ -165,6 +165,11 @@
tensors map[*LazyGPUData]struct{}
}

// subgroupsEnabled is true when the device was created with the SubgroupOperations
// feature and subgroup shaders can be used safely. Detected once at device creation
// and used to select the subgroup MatMul path at dispatch time.
subgroupsEnabled bool

// Memory tracking
memoryStats struct {
totalAllocatedBytes uint64
Expand Down Expand Up @@ -229,7 +234,22 @@

info := adapter.Info()

device, err := adapter.RequestDevice(nil)
// Request SubgroupOperations when the adapter advertises support.
// On Vulkan the HAL ignores the feature flag at device creation (subgroup
// ops are governed by SPIR-V capabilities), but we still record whether
// the feature was granted so the rest of the backend can guard dispatch.
// On DX12 (SM 6.0+) and the Rust wgpu-native backend the flag is honored.
var desc *wgpu.DeviceDescriptor
subgroupsEnabled := false
if adapter.Features().Contains(gputypes.FeatureSubgroupOperations) {
desc = &wgpu.DeviceDescriptor{
Label: "born",
RequiredFeatures: gputypes.Features(gputypes.FeatureSubgroupOperations),

Check warning on line 247 in internal/backend/webgpu/backend.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/backend.go#L242-L247

Added lines #L242 - L247 were not covered by tests
}
subgroupsEnabled = true

Check warning on line 249 in internal/backend/webgpu/backend.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/backend.go#L249

Added line #L249 was not covered by tests
}

device, err := adapter.RequestDevice(desc)

Check warning on line 252 in internal/backend/webgpu/backend.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/backend.go#L252

Added line #L252 was not covered by tests
if err != nil {
adapter.Release()
instance.Release()
Expand All @@ -244,7 +264,12 @@
return nil, fmt.Errorf("webgpu: failed to get queue")
}

return newBackendFromDevice(instance, adapter, device, queue, &info)
b, err := newBackendFromDevice(instance, adapter, device, queue, &info)
if err != nil {
return nil, err

Check warning on line 269 in internal/backend/webgpu/backend.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/backend.go#L267-L269

Added lines #L267 - L269 were not covered by tests
}
b.subgroupsEnabled = subgroupsEnabled
return b, nil

Check warning on line 272 in internal/backend/webgpu/backend.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/backend.go#L271-L272

Added lines #L271 - L272 were not covered by tests
}

// newSoftwareBackend creates a Backend using the software HAL — CPU-based
Expand Down
106 changes: 72 additions & 34 deletions internal/backend/webgpu/compute.go
Original file line number Diff line number Diff line change
Expand Up @@ -656,9 +656,6 @@
return nil, fmt.Errorf("webgpu: matmul shape mismatch: [%d,%d] @ [%d,%d]", M, K, other.Shape()[0], N)
}

shader := b.compileShader("matmul", matmulShader)
entry := b.getOrCreatePipeline("matmul", shader, bglBinary)

bufferA := b.createBuffer(a.Data(), gputypes.BufferUsageStorage|gputypes.BufferUsageCopySrc)
defer bufferA.Release()

Expand All @@ -676,16 +673,35 @@
}
defer bufferResult.Release()

// Uniform: M, K, N as u32 (3×4 = 12 bytes, padded to 16).
params := make([]byte, 16)
binary.LittleEndian.PutUint32(params[0:4], M)
binary.LittleEndian.PutUint32(params[4:8], K)
binary.LittleEndian.PutUint32(params[8:12], N)
bufferParams := b.createUniformBuffer(params)
// Uniform buffer: M, K, N plus one u32 pad to meet std140 16-byte alignment.
paramBytes := make([]byte, 16)
binary.LittleEndian.PutUint32(paramBytes[0:4], M)
binary.LittleEndian.PutUint32(paramBytes[4:8], K)
binary.LittleEndian.PutUint32(paramBytes[8:12], N)
bufferParams := b.createUniformBuffer(paramBytes)

Check warning on line 681 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L677-L681

Added lines #L677 - L681 were not covered by tests
defer bufferParams.Release()

aSize := uint64(a.ByteSize()) //nolint:gosec // G115: integer overflow conversion int -> uint64
otherSize := uint64(other.ByteSize()) //nolint:gosec // G115: integer overflow conversion int -> uint64

var workgroupsX, workgroupsY uint32
var entry pipelineEntry

Check warning on line 688 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L687-L688

Added lines #L687 - L688 were not covered by tests

if b.subgroupsEnabled {

Check warning on line 690 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L690

Added line #L690 was not covered by tests
// Subgroup cooperative K-reduction: one workgroup(32 threads) per output element.
// Dispatch (N, M, 1) — col along X, row along Y.
shader := b.compileShader("matmul_subgroup", matmulSubgroupShader)
entry = b.getOrCreatePipeline("matmul_subgroup", shader, bglBinary)
workgroupsX = N
workgroupsY = M
} else {

Check warning on line 697 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L693-L697

Added lines #L693 - L697 were not covered by tests
// Scalar K-loop: 16×16 tile dispatch.
shader := b.compileShader("matmul", matmulShader)
entry = b.getOrCreatePipeline("matmul", shader, bglBinary)
workgroupsX = uint32(math.Ceil(float64(N) / 16.0))
workgroupsY = uint32(math.Ceil(float64(M) / 16.0))

Check warning on line 702 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L699-L702

Added lines #L699 - L702 were not covered by tests
}

bg := b.createBindGroupFromBuffers(entry.layout, []bindGroupBuffer{
bufBinding(bufferA, aSize),
bufBinding(bufferOther, otherSize),
Expand All @@ -694,9 +710,6 @@
})
defer bg.Release()

// 2D workgroup dispatch: 16×16 tiles.
workgroupsX := uint32(math.Ceil(float64(N) / 16.0))
workgroupsY := uint32(math.Ceil(float64(M) / 16.0))
resultData := b.execComputeAndRead(entry.pipeline, bg, workgroupsX, workgroupsY, 1, bufferResult, resultSize)

result, err := tensor.NewRaw(resultShape, tensor.Float32, tensor.WebGPU)
Expand Down Expand Up @@ -897,9 +910,6 @@

numClasses := uint32(input.Shape()[1]) //nolint:gosec // G115: safe, tensor dimensions are non-negative and fit in uint32

shader := b.compileShader("softmax", softmaxShader)
entry := b.getOrCreatePipeline("softmax", shader, bglUnary)

bufferInput := b.createBuffer(input.Data(), gputypes.BufferUsageStorage|gputypes.BufferUsageCopySrc)
defer bufferInput.Release()

Expand All @@ -913,22 +923,37 @@
}
defer bufferResult.Release()

// Uniform: batch_size, num_classes as u32.
params := make([]byte, 16)
binary.LittleEndian.PutUint32(params[0:4], batchSize)
binary.LittleEndian.PutUint32(params[4:8], numClasses)
bufferParams := b.createUniformBuffer(params)
// Uniform: batch_size, num_classes, two padding u32 to reach 16-byte std140 alignment.
paramBytes := make([]byte, 16)
binary.LittleEndian.PutUint32(paramBytes[0:4], batchSize)
binary.LittleEndian.PutUint32(paramBytes[4:8], numClasses)
bufferParams := b.createUniformBuffer(paramBytes)

Check warning on line 930 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L927-L930

Added lines #L927 - L930 were not covered by tests
defer bufferParams.Release()

var workgroups uint32
var entry pipelineEntry

Check warning on line 934 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L933-L934

Added lines #L933 - L934 were not covered by tests

if b.subgroupsEnabled {

Check warning on line 936 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L936

Added line #L936 was not covered by tests
// Subgroup cooperative reduction: one workgroup (32 lanes) per row.
// Dispatch (batchSize, 1, 1) — each workgroup handles exactly one row.
shader := b.compileShader("softmax_subgroup", softmaxSubgroupShader)
entry = b.getOrCreatePipeline("softmax_subgroup", shader, bglUnary)
workgroups = batchSize
} else {

Check warning on line 942 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L939-L942

Added lines #L939 - L942 were not covered by tests
// Scalar per-thread: one thread per row.
// Dispatch (ceil(batchSize/256), 1, 1).
shader := b.compileShader("softmax", softmaxShader)
entry = b.getOrCreatePipeline("softmax", shader, bglUnary)
workgroups = (batchSize + workgroupSize - 1) / workgroupSize

Check warning on line 947 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L945-L947

Added lines #L945 - L947 were not covered by tests
}

bg := b.createBindGroupFromBuffers(entry.layout, []bindGroupBuffer{
bufBinding(bufferInput, resultSize),
bufBinding(bufferResult, resultSize),
bufBinding(bufferParams, 16),
})
defer bg.Release()

// Each workgroup handles one row (batch sample).
workgroups := (batchSize + workgroupSize - 1) / workgroupSize
resultData := b.execComputeAndRead(entry.pipeline, bg, workgroups, 1, 1, bufferResult, resultSize)

result, err := tensor.NewRaw(input.Shape(), tensor.Float32, tensor.WebGPU)
Expand Down Expand Up @@ -984,9 +1009,6 @@
resultShape = tensor.Shape{shapeA[0], shapeA[1], int(M), int(N)}
}

shader := b.compileShader("batchMatMul", batchMatMulShader)
entry := b.getOrCreatePipeline("batchMatMul", shader, bglBinary)

bufferA := b.createBuffer(a.Data(), gputypes.BufferUsageStorage|gputypes.BufferUsageCopySrc)
defer bufferA.Release()

Expand All @@ -1003,16 +1025,35 @@
}
defer bufferResult.Release()

params := make([]byte, 16)
binary.LittleEndian.PutUint32(params[0:4], batch)
binary.LittleEndian.PutUint32(params[4:8], M)
binary.LittleEndian.PutUint32(params[8:12], K)
binary.LittleEndian.PutUint32(params[12:16], N)
bufferParams := b.createUniformBuffer(params)
paramBytes := make([]byte, 16)
binary.LittleEndian.PutUint32(paramBytes[0:4], batch)
binary.LittleEndian.PutUint32(paramBytes[4:8], M)
binary.LittleEndian.PutUint32(paramBytes[8:12], K)
binary.LittleEndian.PutUint32(paramBytes[12:16], N)
bufferParams := b.createUniformBuffer(paramBytes)

Check warning on line 1033 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L1028-L1033

Added lines #L1028 - L1033 were not covered by tests
defer bufferParams.Release()

aSize := uint64(a.ByteSize()) //nolint:gosec // G115: integer overflow conversion int -> uint64
otherSize := uint64(other.ByteSize()) //nolint:gosec // G115: integer overflow conversion int -> uint64

var workgroupsX, workgroupsY uint32
var entry pipelineEntry

Check warning on line 1040 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L1039-L1040

Added lines #L1039 - L1040 were not covered by tests

if b.subgroupsEnabled {

Check warning on line 1042 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L1042

Added line #L1042 was not covered by tests
// Subgroup cooperative K-reduction: one workgroup(32 threads) per output element.
// Dispatch (N, M, batch) — col along X, row along Y, batch along Z.
shader := b.compileShader("batchMatMul_subgroup", batchMatMulSubgroupShader)
entry = b.getOrCreatePipeline("batchMatMul_subgroup", shader, bglBinary)
workgroupsX = N
workgroupsY = M
} else {

Check warning on line 1049 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L1045-L1049

Added lines #L1045 - L1049 were not covered by tests
// Scalar K-loop: (N+7)/8 × (M+7)/8 × batch dispatch.
shader := b.compileShader("batchMatMul", batchMatMulShader)
entry = b.getOrCreatePipeline("batchMatMul", shader, bglBinary)
workgroupsX = (N + 7) / 8
workgroupsY = (M + 7) / 8

Check warning on line 1054 in internal/backend/webgpu/compute.go

View check run for this annotation

Codecov / codecov/patch

internal/backend/webgpu/compute.go#L1051-L1054

Added lines #L1051 - L1054 were not covered by tests
}

bg := b.createBindGroupFromBuffers(entry.layout, []bindGroupBuffer{
bufBinding(bufferA, aSize),
bufBinding(bufferB, otherSize),
Expand All @@ -1021,9 +1062,6 @@
})
defer bg.Release()

// Dispatch: (N+7)/8 × (M+7)/8 × batch.
workgroupsX := (N + 7) / 8
workgroupsY := (M + 7) / 8
resultData := b.execComputeAndRead(entry.pipeline, bg, workgroupsX, workgroupsY, batch, bufferResult, resultSize)

result, err := tensor.NewRaw(resultShape, tensor.Float32, tensor.WebGPU)
Expand Down
Loading
Loading