You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

Metal计算出现意外结果,请求排查代码问题

编辑说明

在下方评论中,pmdj指出了我代码中的2个Bug。
为清晰起见,我不会复制所有代码和输出缓冲区内容,而是对两者进行了修改,并将原始版本保留为注释。


原问题

我正在学习Metal计算,但即使是我能想到的最简单的代码,也得到了完全不符合预期的结果。以下是我的测试环境:

我编写了如下单元测试函数:

func test_test() throws {
    let metalImageProcessor = MetalImageProcessor()!
    metalImageProcessor.test()
}

我的想法是测试完成后,可通过调试器查看日志。

我的MetalImageProcessor实现如下:

func test() {
    let device = MTLCreateSystemDefaultDevice()!
    let commandQueue = device.makeCommandQueue()!
    let library = device.makeDefaultLibrary()!
    let testFunction = library.makeFunction(name: "testFunction")!
    let pipelineState = try! device.makeComputePipelineState(function: testFunction)
    let commandBuffer = commandQueue.makeCommandBuffer()!
    let commandEncoder = commandBuffer.makeComputeCommandEncoder()!

    let threadExecutionWidth  = makeTexturePipelineState.threadExecutionWidth
    let threadExecutionHeight = makeTexturePipelineState.maxTotalThreadsPerThreadgroup / threadExecutionWidth
    let threadsPerThreadgroup = MTLSize(width: threadExecutionWidth, height: threadExecutionHeight, depth: 1) 
    let threadsPerGrid = MTLSize(width: threadExecutionWidth, height: threadExecutionHeight, depth: 1)
    
    let outputBufferSize = threadsPerGrid.width * threadsPerGrid.height * MemoryLayout<Float>.size * 10 // shader writes 10 floats per thread  
// WAS: let outputBufferSize = threadsPerGrid.width * threadsPerGrid.height * MemoryLayout<Float>.size // shader writes 10 floats
    var buffer1 = [UInt8](repeating: 0, count: outputBufferSize)
    let outputBuffer1 = device.makeBuffer(bytes: &buffer1, length: outputBufferSize, options: .storageModeShared)
    
    commandEncoder.setComputePipelineState(pipelineState)
    commandEncoder.setBuffer(outputBuffer1, offset: 0, index: 0)

    commandEncoder.dispatchThreadgroups(threadsPerGrid, threadsPerThreadgroup: threadsPerThreadgroup)
    commandEncoder.endEncoding()
    commandBuffer.commit()
    commandBuffer.waitUntilCompleted()
    
    debugBufferValues(name: "outputBuffer1", buffer: outputBuffer1!, width: threadExecutionWidth, height: threadExecutionHeight)
}

func debugBufferValues(name: String, buffer: MTLBuffer, width: Int, height: Int) {
    let debugData = buffer.contents().bindMemory(to: Float.self, capacity: width * height)
    print("\n")
    for i in 0 ..< min(width * height, 8320) {
        print(name, "[\(i)]: \(debugData[i])")
    }
}

它调用了如下着色器内核函数:

kernel void testFunction(device float *outputBuffer [[buffer(0)]],
                         uint2 gridSize [[grid_size]],
                         uint2 threadGroupIDinGrid [[threadgroup_position_in_grid]],
                         uint2 threadIDinThreadGroup [[thread_position_in_threadgroup]],
                         uint2 threadGroupSize [[threads_per_threadgroup]],
                         uint2 gid [[thread_position_in_grid]]) {
    uint nrFloatsPerThread = 10; // see below
    // uint testValueIndex = 0; // start value // out commented 
    uint outputBufferIndex = (gridSize.x * gid.y + gid.x) * nrFloatsPerThread; // WAS: uint outputBufferIndex = (gridSize.x * gid.y + gid.x) * nrFloatsPerThread + testValueIndex;

    outputBuffer[outputBufferIndex] = gridSize.x;               outputBufferIndex += 1; // WAS: testValueIndex += 1;
    outputBuffer[outputBufferIndex] = gridSize.y;               outputBufferIndex += 1; // WAS: testValueIndex += 1;
    outputBuffer[outputBufferIndex] = threadGroupIDinGrid.x;    outputBufferIndex += 1; // WAS: testValueIndex += 1;
    outputBuffer[outputBufferIndex] = threadGroupIDinGrid.y;    outputBufferIndex += 1; // WAS: testValueIndex += 1;
    outputBuffer[outputBufferIndex] = threadIDinThreadGroup.x;  outputBufferIndex += 1; // WAS: testValueIndex += 1;
    outputBuffer[outputBufferIndex] = threadIDinThreadGroup.y;  outputBufferIndex += 1; // WAS: testValueIndex += 1;
    outputBuffer[outputBufferIndex] = threadGroupSize.x;        outputBufferIndex += 1; // WAS: testValueIndex += 1;
    outputBuffer[outputBufferIndex] = threadGroupSize.y;        outputBufferIndex += 1; // WAS: testValueIndex += 1;
    outputBuffer[outputBufferIndex] = gid.x;                    outputBufferIndex += 1; // WAS: testValueIndex += 1;
    outputBuffer[outputBufferIndex] = gid.y;                    outputBufferIndex += 1; // WAS: testValueIndex += 1;
}

预期结果

由于网格与线程组维度相同,因此只有一个线程组,在我的测试中大小为(32,16,1)。
该组中的每个线程都会执行着色器函数,并在输出缓冲区中写入10个浮点数,预期值如下:

outputBuffer1 [0]: 32.0     // gridSize.x
outputBuffer1 [1]: 16.0     // gridSize.y
outputBuffer1 [2]: 0.0      // threadGroupIDinGrid.x
outputBuffer1 [3]: 0.0      // threadGroupIDinGrid.y
outputBuffer1 [4]: 0.0      // threadIDinThreadGroup.x
outputBuffer1 [5]: 0.0      // threadIDinThreadGroup.y
outputBuffer1 [6]: 32.0     // threadGroupSize.x
outputBuffer1 [7]: 16.0     // threadGroupSize.y
outputBuffer1 [8]: 0.0      // gid.x
outputBuffer1 [9]: 0.0      // gid.y
outputBuffer1 [0]: 32.0     // gridSize.x
outputBuffer1 [1]: 16.0     // gridSize.y
outputBuffer1 [2]: 0.0      // threadGroupIDinGrid.x
outputBuffer1 [3]: 0.0      // threadGroupIDinGrid.y
outputBuffer1 [4]: 1.0      // threadIDinThreadGroup.x
outputBuffer1 [5]: 0.0      // threadIDinThreadGroup.y
outputBuffer1 [6]: 32.0     // threadGroupSize.x
outputBuffer1 [7]: 16.0     // threadGroupSize.y
outputBuffer1 [8]: 1.0      // gid.x
outputBuffer1 [9]: 0.0      // gid.y  

以此类推。

实际结果

outputBuffer1 [0]: 0.0 // WAS: 250.0
outputBuffer1 [1]: 0.0
outputBuffer1 [2]: 0.0
outputBuffer1 [3]: 15.0 // WAS: 0.0
outputBuffer1 [4]: 0.0
outputBuffer1 [5]: 12.0 // WAS: 0.0
outputBuffer1 [6]: 32.0 // WAS: 0.0
outputBuffer1 [7]: 16.0 // WAS: 0.0
outputBuffer1 [8]: 0.0
outputBuffer1 [9]: 255.0 // WAS: 0.0
outputBuffer1 [10]: 0.0 // WAS: 250.0
outputBuffer1 [11]: 0.0
outputBuffer1 [12]: 0.0
outputBuffer1 [13]: 15.0 // WAS: 0.0
outputBuffer1 [14]: 1.0 // WAS: 0.0
outputBuffer1 [15]: 12.0 // WAS: 0.0
outputBuffer1 [16]: 32.0 // WAS: 0.0
outputBuffer1 [17]: 16.0 // WAS: 0.0
outputBuffer1 [18]: 1.0 // WAS: 0.0
outputBuffer1 [19]: 255.0 // WAS: 0.0  

以此类推。

输出缓冲区的前10个值就已经错误。我既无法理解outputBuffer1 [0]中的250.0,也无法理解后续的零值。
显然我的代码存在问题,但找不到Bug所在。
欢迎任何帮助!


编辑说明

代码修正后的变化:

  • gridSize.x和gridSize.y均为0,无法理解。
  • threadGroupIDinGrid.y现在为15,可能是线程组中最后一个线程留下的值。
  • gid.y现在为255,无法理解,因为线程组中有32*16个线程,最后一个索引应为511。

内容的提问来源于stack exchange,提问作者Reinhard Männer

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.22 14:14:50