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

Triton矩阵乘法结果与PyTorch不符的异常问题求助

Triton矩阵乘法结果与PyTorch不符的异常问题求助

我最近在尝试用Triton实现矩阵点积运算,但遇到了一个很奇怪的问题——Triton计算出的结果和PyTorch的torch.matmul完全对不上,折腾了半天也没找到问题所在,想请大家帮忙看看!

先说说我的测试用例:我定义了两个矩阵P和V,其中P是32x32的矩阵,除了最后一列全为1之外,其他元素都是0;V则是从32*64到64*64的连续整数,reshape成32x64的矩阵。按照矩阵乘法的逻辑,P和V的点积结果应该等于V的最后一行,因为P只有最后一列有非零值,相当于取V的最后一行。

P和V的构造代码如下:

P = torch.zeros((32,32), device = 'cuda',  dtype = torch.float32)
P[:,-1] = 1
V = torch.arange(32*64, 64 * 64, device = 'cuda', dtype = torch.float32).reshape(32, 64)

现在问题来了:当我用自己写的Triton内核调用tl.dot(P, V)时,看起来P和V都加载正确了,但输出结果却是重复的成对数值:

[4032., 4032., 4034., 4034., 4036., 4036., 4038., 4038., 4040., 4040.,
         4042., 4042., 4044., 4044., 4046., 4046., 4048., 4048., 4050., 4050.,
         4052., 4052., 4054., 4054., 4056., 4056., 4058., 4058., 4060., 4060.,
         4062., 4062., 4064., 4064., 4066., 4066., 4068., 4068., 4070., 4070.,
         4072., 4072., 4074., 4074., 4076., 4076., 4078., 4078., 4080., 4080.,
         4082., 4082., 4084., 4084., 4086., 4086., 4088., 4088., 4090., 4090.,
         4092., 4092., 4094., 4094.]

而用PyTorch的torch.matmul(P, V)得到的是正确的连续递增结果,也就是V的最后一行:

[4032., 4033., 4034., 4035., 4036., 4037., 4038., 4039., 4040., 4041.,
         4042., 4043., 4044., 4045., 4046., 4047., 4048., 4049., 4050., 4051.,
         4052., 4053., 4054., 4055., 4056., 4057., 4058., 4059., 4060., 4061.,
         4062., 4063., 4064., 4065., 4066., 4067., 4068., 4069., 4070., 4071.,
         4072., 4073., 4074., 4075., 4076., 4077., 4078., 4079., 4080., 4081.,
         4082., 4083., 4084., 4085., 4086., 4087., 4088., 4089., 4090., 4091.,
         4092., 4093., 4094., 4095.]

下面是我写的Triton内核和辅助函数代码:

import triton
import triton.language as tl
import torch
torch.cuda.is_available()
torch.set_printoptions(profile="full")

@triton.jit
def test_kernel(x_ptr,y_ptr,output_ptr,
                M, K, N,
                stride_xm, stride_xk,
                stride_yk, stride_yn,
                stride_om, stride_on,
                BLOCK_SIZE_M: tl.constexpr,
                BLOCK_SIZE_K: tl.constexpr,
                BLOCK_SIZE_N: tl.constexpr):
    pid_m = tl.program_id(axis = 0) * BLOCK_SIZE_M
    pid_n = tl.program_id(axis = 1) * BLOCK_SIZE_N
    x_ptr += (pid_m + tl.arange(0, BLOCK_SIZE_M))[:,None] * stride_xm +  (pid_n + tl.arange(0, BLOCK_SIZE_K))[None,:]*stride_xk
    y_ptr += (pid_m + tl.arange(0, BLOCK_SIZE_K))[:,None] * stride_yk +  (pid_n + tl.arange(0, BLOCK_SIZE_N))[None,:]*stride_yn
    x = tl.load(x_ptr) 
    y = tl.load(y_ptr) 
    output_offset = (pid_m + tl.arange(0, BLOCK_SIZE_M))[:,None] *stride_om + (pid_n + tl.arange(0, BLOCK_SIZE_N))[None, :] *stride_on
    tl.store(output_ptr + output_offset,  tl.dot(x,y))

def helper(x: torch.Tensor, y: torch.Tensor):
    M , K = x.shape
    K1, N = y.shape
    assert K == K1
    output = torch.empty((M, N), device = 'cuda', dtype = torch.float32)
    assert x.is_cuda and y.is_cuda and output.is_cuda

    grid = lambda meta: (triton.cdiv(M, meta['BLOCK_SIZE_M']), triton.cdiv(N, meta['BLOCK_SIZE_N']),)
    test_kernel[grid](x, y, output, 
                        M, K, N, 
                        x.stride(0), x.stride(1),
                        y.stride(0), y.stride(1),
                        output.stride(0), output.stride(1),
                        BLOCK_SIZE_N = 64, 
                        BLOCK_SIZE_K = 32, 
                        BLOCK_SIZE_M = 32, 
                        )
    return output

最让我困惑的是,如果我把V的构造改成从0开始的连续整数(torch.arange(0, 32*64, device = 'cuda', dtype = torch.float32).reshape(32, 64)),Triton计算出来的结果就和PyTorch完全一致了!这是不是说明我在处理指针偏移的时候哪里写错了?

麻烦大家帮我看看代码里有没有什么明显的错误,或者有没有什么Triton的细节我没注意到导致了这个问题?


备注:内容来源于stack exchange,提问作者Div

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.13 19:18:05