| ... | ... | @@ -23,11 +23,10 @@ pub fn addCases(ctx: *TestContext) !void { |
| 23 | 23 | var case = addPtx(ctx, "nvptx: read special registers"); |
| 24 | 24 | |
| 25 | 25 | case.compiles( |
| 26 | | \\fn threadIdX() usize { |
| 27 | | \\ var tid = asm volatile ("mov.u32 \t$0, %tid.x;" |
| 28 | | \\ : [ret] "=r" (-> u32), |
| 29 | | \\ ); |
| 30 | | \\ return @as(usize, tid); |
| 26 | \\fn threadIdX() u32 { |
| 27 | \\ return asm ("mov.u32 \t%[r], %tid.x;" |
| 28 | \\ : [r] "=r" (-> utid), |
| 29 | \\ ); |
| 31 | 30 | \\} |
| 32 | 31 | \\ |
| 33 | 32 | \\pub export fn special_reg(a: []const i32, out: []i32) callconv(.PtxKernel) void { |
| ... | ... | @@ -49,6 +48,38 @@ pub fn addCases(ctx: *TestContext) !void { |
| 49 | 48 | \\} |
| 50 | 49 | ); |
| 51 | 50 | } |
| 51 | |
| 52 | { |
| 53 | var case = addPtx(ctx, "nvptx: reduce in shared mem"); |
| 54 | case.compiles( |
| 55 | \\fn threadIdX() u32 { |
| 56 | \\ return asm ("mov.u32 \t%[r], %tid.x;" |
| 57 | \\ : [r] "=r" (-> utid), |
| 58 | \\ ); |
| 59 | \\} |
| 60 | \\ |
| 61 | \\ var _sdata: [1024]f32 addrspace(.shared) = undefined; |
| 62 | \\ pub export fn reduceSum(d_x: []const f32, out: *f32) callconv(ptx.Kernel) void { |
| 63 | \\ var sdata = @addrSpaceCast(.generic, &_sdata); |
| 64 | \\ const tid: u32 = threadIdX(); |
| 65 | \\ var sum = d_x[tid]; |
| 66 | \\ sdata[tid] = sum; |
| 67 | \\ asm volatile ("bar.sync \t0;"); |
| 68 | \\ var s: u32 = 512; |
| 69 | \\ while (s > 0) : (s = s >> 1) { |
| 70 | \\ if (tid < s) { |
| 71 | \\ sum += sdata[tid + s]; |
| 72 | \\ sdata[tid] = sum; |
| 73 | \\ } |
| 74 | \\ asm volatile ("bar.sync \t0;"); |
| 75 | \\ } |
| 76 | \\ |
| 77 | \\ if (tid == 0) { |
| 78 | \\ out.* = sum; |
| 79 | \\ } |
| 80 | \\ } |
| 81 | ); |
| 82 | } |
| 52 | 83 | } |
| 53 | 84 | |
| 54 | 85 | const nvptx_target = std.zig.CrossTarget{ |