Files
go-mxl-pattern-generator/cmd/wgpu-sample/main.go
T
Dmitry Sergeev f48195a8bf WGPU is coming
2026-09-07 17:18:03 +03:00

283 lines
7.3 KiB
Go

// Command wgpu-sample renders 75% SMPTE color bars (Rec.709, 10-bit v210)
// with a wgpu compute shader, verifies the output against the reference
// table, and reports per-frame timings.
//
// Run from the repo root: go run ./cmd/wgpu-sample
package main
import (
"context"
"encoding/binary"
"fmt"
"log"
"os"
"time"
"github.com/gogpu/gputypes"
"github.com/gogpu/wgpu"
// Register all available GPU backends (Vulkan, DX12, GLES, Metal, etc.)
_ "github.com/gogpu/wgpu/hal/allbackends"
)
const (
width, height = 1920, 1080
wgSize = 64
blocks = width * height / 6
frameSize = uint64(blocks * 16)
benchFrames = 100
)
// 75% SMPTE bars, Rec.709, 10-bit {Y, Cb, Cr}
var bars = [7][3]uint32{
{721, 512, 512}, {674, 176, 543}, {581, 589, 176},
{534, 253, 207}, {251, 771, 817}, {204, 435, 848}, {111, 848, 481},
}
func main() {
if err := run(); err != nil {
log.Fatalf("FATAL: %v", err)
}
}
func run() error {
instance, err := wgpu.CreateInstance(nil)
if err != nil {
return fmt.Errorf("CreateInstance: %w", err)
}
defer instance.Release()
adapter, err := instance.RequestAdapter(nil)
if err != nil {
return fmt.Errorf("RequestAdapter: %w", err)
}
defer adapter.Release()
info := adapter.Info()
fmt.Printf("adapter: %s (%s, %s, driver %s)\n", info.Name, info.Vendor, info.DeviceType, info.Driver)
device, err := adapter.RequestDevice(nil)
if err != nil {
return fmt.Errorf("RequestDevice: %w", err)
}
defer device.Release()
queue := device.Queue()
wgsl, err := os.ReadFile("kernels/v210_bars.wgsl")
if err != nil {
return fmt.Errorf("read kernel: %w", err)
}
out, err := device.CreateBuffer(&wgpu.BufferDescriptor{
Label: "v210-out", Size: frameSize,
Usage: wgpu.BufferUsageStorage | wgpu.BufferUsageCopySrc,
})
if err != nil {
return err
}
defer out.Release()
staging, err := device.CreateBuffer(&wgpu.BufferDescriptor{
Label: "v210-staging", Size: frameSize,
Usage: wgpu.BufferUsageCopyDst | wgpu.BufferUsageMapRead,
})
if err != nil {
return err
}
defer staging.Release()
params := make([]byte, 16)
binary.LittleEndian.PutUint32(params[0:], width)
binary.LittleEndian.PutUint32(params[4:], height)
uniform, err := device.CreateBuffer(&wgpu.BufferDescriptor{
Label: "params", Size: uint64(len(params)),
Usage: wgpu.BufferUsageUniform | wgpu.BufferUsageCopyDst,
})
if err != nil {
return err
}
defer uniform.Release()
if err := queue.WriteBuffer(uniform, 0, params); err != nil {
return fmt.Errorf("write params: %w", err)
}
shader, err := device.CreateShaderModule(&wgpu.ShaderModuleDescriptor{
Label: "bars-shader", WGSL: string(wgsl),
})
if err != nil {
return fmt.Errorf("shader: %w", err)
}
defer shader.Release()
bgl, err := device.CreateBindGroupLayout(&wgpu.BindGroupLayoutDescriptor{
Label: "bars-bgl",
Entries: []gputypes.BindGroupLayoutEntry{
{Binding: 0, Visibility: wgpu.ShaderStageCompute, Buffer: &gputypes.BufferBindingLayout{Type: gputypes.BufferBindingTypeStorage}},
{Binding: 1, Visibility: wgpu.ShaderStageCompute, Buffer: &gputypes.BufferBindingLayout{Type: gputypes.BufferBindingTypeUniform}},
},
})
if err != nil {
return err
}
defer bgl.Release()
bg, err := device.CreateBindGroup(&wgpu.BindGroupDescriptor{
Label: "bars-bg", Layout: bgl,
Entries: []wgpu.BindGroupEntry{
{Binding: 0, Buffer: out, Size: frameSize},
{Binding: 1, Buffer: uniform, Size: uint64(len(params))},
},
})
if err != nil {
return err
}
defer bg.Release()
pl, err := device.CreatePipelineLayout(&wgpu.PipelineLayoutDescriptor{
Label: "bars-pl", BindGroupLayouts: []*wgpu.BindGroupLayout{bgl},
})
if err != nil {
return err
}
defer pl.Release()
pipeline, err := device.CreateComputePipeline(&wgpu.ComputePipelineDescriptor{
Label: "bars-pipeline", Layout: pl, Module: shader, EntryPoint: "main",
})
if err != nil {
return fmt.Errorf("pipeline: %w", err)
}
defer pipeline.Release()
render := func(frame uint32) error {
binary.LittleEndian.PutUint32(params[8:], frame)
if err := queue.WriteBuffer(uniform, 0, params); err != nil {
return err
}
encoder, err := device.CreateCommandEncoder(nil)
if err != nil {
return err
}
pass, err := encoder.BeginComputePass(nil)
if err != nil {
return err
}
pass.SetPipeline(pipeline)
pass.SetBindGroup(0, bg, nil)
pass.Dispatch((blocks+wgSize-1)/wgSize, 1, 1)
if err := pass.End(); err != nil {
return err
}
encoder.CopyBufferToBuffer(out, 0, staging, 0, frameSize)
cmd, err := encoder.Finish()
if err != nil {
return err
}
if _, err := queue.Submit(cmd); err != nil {
return err
}
return nil
}
readback := func() ([]byte, error) {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
if err := staging.Map(ctx, wgpu.MapModeRead, 0, frameSize); err != nil {
return nil, fmt.Errorf("map: %w", err)
}
rng, err := staging.MappedRange(0, frameSize)
if err != nil {
_ = staging.Unmap()
return nil, fmt.Errorf("mapped range: %w", err)
}
data := make([]byte, frameSize)
copy(data, rng.Bytes())
if err := staging.Unmap(); err != nil {
return nil, err
}
return data, nil
}
// Correctness: render one frame and verify every pixel.
fmt.Printf("frame: %dx%d v210, %d bytes\n", width, height, frameSize)
if err := render(0); err != nil {
return fmt.Errorf("render: %w", err)
}
data, err := readback()
if err != nil {
return fmt.Errorf("readback: %w", err)
}
if err := verify(data); err != nil {
return err
}
fmt.Println("verify: PASS (all 1920x1080 pixels match the 75% Rec.709 table)")
// Rough timing: full dispatch + copy + map/unmap per frame.
for i := 0; i < 10; i++ {
if err := render(uint32(i)); err != nil {
return err
}
if _, err := readback(); err != nil {
return err
}
}
start := time.Now()
for i := 0; i < benchFrames; i++ {
if err := render(uint32(i)); err != nil {
return err
}
if _, err := readback(); err != nil {
return err
}
}
perFrame := time.Since(start) / benchFrames
fmt.Printf("timing: %v/frame (dispatch+copy+readback), %.1f fps equivalent\n",
perFrame, float64(time.Second)/float64(perFrame))
return nil
}
// verify decodes the v210 frame and checks Y at every pixel and the
// co-sited chroma (even columns) against the reference table.
func verify(data []byte) error {
barOf := func(x int) int {
if b := x * 7 / width; b < 7 {
return b
}
return 6
}
for y := 0; y < height; y++ {
for x := 0; x < width; x++ {
p := y*width + x
off := (p / 6) * 16
w0 := binary.LittleEndian.Uint32(data[off:])
w1 := binary.LittleEndian.Uint32(data[off+4:])
w2 := binary.LittleEndian.Uint32(data[off+8:])
w3 := binary.LittleEndian.Uint32(data[off+12:])
var yv, cb, cr uint32
switch p % 6 {
case 0:
yv, cb, cr = (w0>>10)&0x3FF, w0&0x3FF, (w0>>20)&0x3FF
case 1:
yv = w1 & 0x3FF
case 2:
yv, cb, cr = (w1>>20)&0x3FF, (w1>>10)&0x3FF, w2&0x3FF
case 3:
yv = (w2 >> 10) & 0x3FF
case 4:
yv, cb, cr = w3&0x3FF, (w2>>20)&0x3FF, (w3>>10)&0x3FF
case 5:
yv = (w3 >> 20) & 0x3FF
}
b := bars[barOf(x)]
if yv != b[0] {
return fmt.Errorf("pixel (%d,%d): Y=%d want %d", x, y, yv, b[0])
}
if x%2 == 0 && (cb != b[1] || cr != b[2]) {
return fmt.Errorf("pixel (%d,%d): Cb=%d Cr=%d want %d/%d", x, y, cb, cr, b[1], b[2])
}
}
}
return nil
}