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
相关产品推荐
相关产品推荐

