yum-mirror/slang
Making it easier to work with shaders
git clone https://git.yummers.dev/yum-mirror/slang
73f9aeb83
master
1//TEST(compute):COMPARE_COMPUTE(filecheck-buffer=CHECK): -cpu -output-using-type 2//TEST(compute):COMPARE_COMPUTE(filecheck-buffer=CHECK): -vk -output-using-type 3//TEST(compute):COMPARE_COMPUTE(filecheck-buffer=CHECK): -d3d11 -output-using-type 4 5// This test verifies that row_major and column_major matrices don't create 6// duplicate DiffPair structs when used together in autodiff code. 7// Before the fix, this would generate compilation errors due to mismatched 8// DiffPair_matrixx3Cfloatx2C3x2C3x3E_0 and DiffPair_matrixx3Cfloatx2C3x2C3x3E_1 types. 9 10//TEST_INPUT:ubuffer(data=[0 0 0 0 0 0 0 0 0], stride=4):out,name=outputBuffer 11RWStructuredBuffer<float> outputBuffer; 12 13[Differentiable] 14float3 matmul33_row(no_diff float3 v, row_major float3x3 w) { 15 return mul(w, v); 16} 17 18[Differentiable] 19float3 matmul33_col(no_diff float3 v, column_major float3x3 w) { 20 return mul(w, v); 21} 22 23[numthreads(1, 1, 1)] 24void computeMain(uint3 dispatchThreadID : SV_DispatchThreadID) { 25 // Test row_major matrix with meaningful values 26 row_major float3x3 w_row = float3x3(1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0); 27 float3 v = float3(1.0, 2.0, 3.0); 28 29 DifferentialPair<row_major float3x3> dpW_row = diffPair(w_row); 30 __bwd_diff(matmul33_row)(v, dpW_row, float3(4.0, 5.0, 6.0)); 31 32 // Write gradients to output buffer to prevent dead code elimination 33 // Expected gradient matrix is dResult ⊗ v = [4,5,6]^T ⊗ [1,2,3] = [[4,8,12],[5,10,15],[6,12,18]] 34 outputBuffer[0] = dpW_row.d[0][0]; // CHECK: 4 35 outputBuffer[1] = dpW_row.d[0][1]; // CHECK: 8 36 outputBuffer[2] = dpW_row.d[0][2]; // CHECK: 12 37 38 // Test column_major matrix to ensure they share the same DiffPair struct 39 column_major float3x3 w_col = float3x3(1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0); 40 DifferentialPair<column_major float3x3> dpW_col = diffPair(w_col); 41 __bwd_diff(matmul33_col)(v, dpW_col, float3(4.0, 5.0, 6.0)); 42 43 outputBuffer[3] = dpW_col.d[1][0]; // CHECK: 5 44 outputBuffer[4] = dpW_col.d[1][1]; // CHECK: 10 45 outputBuffer[5] = dpW_col.d[1][2]; // CHECK: 15 46 47 // Additional test values from different matrix positions 48 outputBuffer[6] = dpW_row.d[2][0]; // CHECK: 6 49 outputBuffer[7] = dpW_col.d[2][1]; // CHECK: 12 50 outputBuffer[8] = dpW_row.d[2][2]; // CHECK: 18 51}