Skip to content

Commit 2fab3f1

Browse files
authored
[SM6.10] LinAlg: Update vector accumulate to match spec (#8644)
microsoft/hlsl-specs#892 updates the spec for vector accumulate. This PR updates the implementation to match those changes. Fixes #8639
1 parent bb42209 commit 2fab3f1

22 files changed

Lines changed: 73 additions & 73 deletions

File tree

include/dxc/DXIL/DxilInstructions.h

Lines changed: 12 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -10981,20 +10981,20 @@ struct DxilInst_LinAlgVectorAccumulateToDescriptor {
1098110981
bool requiresUniformInputs() const { return false; }
1098210982
// Operand indexes
1098310983
enum OperandIdx {
10984-
arg_vector = 1,
10985-
arg_handle = 2,
10986-
arg_offset = 3,
10987-
arg_align = 4,
10984+
arg_handle = 1,
10985+
arg_offset = 2,
10986+
arg_align = 3,
10987+
arg_vector = 4,
1098810988
};
1098910989
// Accessors
10990-
llvm::Value *get_vector() const { return Instr->getOperand(1); }
10991-
void set_vector(llvm::Value *val) { Instr->setOperand(1, val); }
10992-
llvm::Value *get_handle() const { return Instr->getOperand(2); }
10993-
void set_handle(llvm::Value *val) { Instr->setOperand(2, val); }
10994-
llvm::Value *get_offset() const { return Instr->getOperand(3); }
10995-
void set_offset(llvm::Value *val) { Instr->setOperand(3, val); }
10996-
llvm::Value *get_align() const { return Instr->getOperand(4); }
10997-
void set_align(llvm::Value *val) { Instr->setOperand(4, val); }
10990+
llvm::Value *get_handle() const { return Instr->getOperand(1); }
10991+
void set_handle(llvm::Value *val) { Instr->setOperand(1, val); }
10992+
llvm::Value *get_offset() const { return Instr->getOperand(2); }
10993+
void set_offset(llvm::Value *val) { Instr->setOperand(2, val); }
10994+
llvm::Value *get_align() const { return Instr->getOperand(3); }
10995+
void set_align(llvm::Value *val) { Instr->setOperand(3, val); }
10996+
llvm::Value *get_vector() const { return Instr->getOperand(4); }
10997+
void set_vector(llvm::Value *val) { Instr->setOperand(4, val); }
1099810998
};
1099910999

1100011000
/// This instruction triggers a breakpoint if debugging is enabled

lib/DXIL/DxilOperations.cpp

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6686,10 +6686,10 @@ Function *OP::GetOpFunc(OpCode opCode, Type *pOverloadType) {
66866686
case OpCode::LinAlgVectorAccumulateToDescriptor:
66876687
A(pV);
66886688
A(pI32);
6689-
A(pETy);
66906689
A(pRes);
66916690
A(pI32);
66926691
A(pI32);
6692+
A(pETy);
66936693
break;
66946694

66956695
//
@@ -6858,6 +6858,7 @@ llvm::Type *OP::GetOverloadType(OpCode opCode, llvm::Function *F) {
68586858
case OpCode::StorePrimitiveOutput:
68596859
case OpCode::DispatchMesh:
68606860
case OpCode::RawBufferVectorStore:
6861+
case OpCode::LinAlgVectorAccumulateToDescriptor:
68616862
if (FT->getNumParams() <= 4)
68626863
return nullptr;
68636864
return FT->getParamType(4);
@@ -6886,7 +6887,6 @@ llvm::Type *OP::GetOverloadType(OpCode opCode, llvm::Function *F) {
68866887
case OpCode::LinAlgMatrixGetCoordinate:
68876888
case OpCode::LinAlgMatrixStoreToDescriptor:
68886889
case OpCode::LinAlgMatrixAccumulateToDescriptor:
6889-
case OpCode::LinAlgVectorAccumulateToDescriptor:
68906890
if (FT->getNumParams() <= 1)
68916891
return nullptr;
68926892
return FT->getParamType(1);

lib/HLSL/HLOperationLower.cpp

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -7173,16 +7173,16 @@ Value *TranslateLinAlgVectorAccumulateToDescriptor(
71737173

71747174
Constant *OpArg = HlslOp->GetU32Const(static_cast<unsigned>(OpCode));
71757175

7176-
Value *Vector = CI->getArgOperand(1);
7177-
Value *ResHandle = CI->getArgOperand(2);
7178-
Value *Offset = CI->getArgOperand(3);
7179-
Value *Align = CI->getArgOperand(4);
7176+
Value *ResHandle = CI->getArgOperand(1);
7177+
Value *Offset = CI->getArgOperand(2);
7178+
Value *Align = CI->getArgOperand(3);
7179+
Value *Vector = CI->getArgOperand(4);
71807180

71817181
// Get the DXIL function for the operation
71827182
Function *DxilFunc = HlslOp->GetOpFunc(OpCode, Vector->getType());
71837183

71847184
return Builder.CreateCall(DxilFunc,
7185-
{OpArg, Vector, ResHandle, Offset, Align});
7185+
{OpArg, ResHandle, Offset, Align, Vector});
71867186
}
71877187

71887188
} // namespace

tools/clang/lib/Headers/hlsl/dx/linalg.h

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -629,11 +629,11 @@ OuterProduct(vector<InputElTy, M> VecA, vector<InputElTy, N> VecB) {
629629
return Result;
630630
}
631631

632-
template <typename InputElTy, SIZE_TYPE M>
632+
template <uint Align = 64, typename InputElTy, SIZE_TYPE M>
633633
typename hlsl::enable_if<hlsl::is_arithmetic<InputElTy>::value, void>::type
634-
InterlockedAccumulate(vector<InputElTy, M> Vec, RWByteAddressBuffer Res,
635-
uint StartOffset, uint Align = 64) {
636-
__builtin_LinAlg_VectorAccumulateToDescriptor(Vec, Res, StartOffset, Align);
634+
InterlockedAccumulate(RWByteAddressBuffer Res, uint StartOffset,
635+
vector<InputElTy, M> Vec) {
636+
__builtin_LinAlg_VectorAccumulateToDescriptor(Res, StartOffset, Align, Vec);
637637
}
638638

639639
} // namespace linalg

tools/clang/test/CodeGenDXIL/hlsl/linalg/api/vectors.hlsl

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -196,12 +196,12 @@ void main(uint ID : SV_GroupID) {
196196
// CHECK-SAME: ; LinAlgMatVecMulAdd(matrix,isOutputSigned,inputVector,inputInterpretation,biasVector,biasInterpretation)
197197
vector<half, 7> vec24 = MultiplyAdd<half>(Mat_7_15_Packed, interpVecH15Packed, memBias7Packed);
198198

199-
// CHECK: call void @dx.op.linAlgVectorAccumulateToDescriptor.v4f16(i32 -2147483617, <4 x half>
200-
// CHECK-SAME: <half 0xH4926, half 0xH4926, half 0xH4926, half 0xH4926>, %dx.types.Handle %{{[0-9]+}}, i32 0, i32 64)
201-
// CHECK-SAME: ; LinAlgVectorAccumulateToDescriptor(vector,handle,offset,align)
202-
InterlockedAccumulate(vec1, RWBAB, 0);
203-
204-
// CHECK: call void @dx.op.linAlgVectorAccumulateToDescriptor.v8f16(i32 -2147483617, <8 x half> %{{[0-9]+}},
205-
// CHECK-SAME: %dx.types.Handle %{{[0-9]+}}, i32 8, i32 64) ; LinAlgVectorAccumulateToDescriptor(vector,handle,offset,align)
206-
InterlockedAccumulate(vec2, RWBAB, 8);
199+
// CHECK: call void @dx.op.linAlgVectorAccumulateToDescriptor.v4f16(i32 -2147483617, %dx.types.Handle %{{[0-9]+}},
200+
// CHECK-SAME: i32 0, i32 64, <4 x half> <half 0xH4926, half 0xH4926, half 0xH4926, half 0xH4926>)
201+
// CHECK-SAME: ; LinAlgVectorAccumulateToDescriptor(handle,offset,align,vector)
202+
InterlockedAccumulate(RWBAB, 0, vec1);
203+
204+
// CHECK: call void @dx.op.linAlgVectorAccumulateToDescriptor.v8f16(i32 -2147483617, %dx.types.Handle %{{[0-9]+}},
205+
// CHECK-SAME: i32 8, i32 64, <8 x half> %{{[0-9]+}}) ; LinAlgVectorAccumulateToDescriptor(handle,offset,align,vector)
206+
InterlockedAccumulate(RWBAB, 8, vec2);
207207
}

tools/clang/test/CodeGenDXIL/hlsl/linalg/builtins/vectoraccumulatetodescriptor/nominal.hlsl

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -8,12 +8,12 @@ RWByteAddressBuffer outbuf;
88
void main() {
99
// CHECK-LABEL: define void @main()
1010

11-
// CHECK: call void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32 -2147483617, <4 x float>
12-
// CHECK-SAME: <float 9.000000e+00, float 8.000000e+00, float 7.000000e+00, float 6.000000e+00>, %dx.types.Handle %{{.*}}, i32 16, i32 64)
13-
// CHECK-SAME: ; LinAlgVectorAccumulateToDescriptor(vector,handle,offset,align)
11+
// CHECK: call void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32 -2147483617, %dx.types.Handle %{{.*}}, i32 16,
12+
// CHECK-SAME: i32 64, <4 x float> <float 9.000000e+00, float 8.000000e+00, float 7.000000e+00, float 6.000000e+00>)
13+
// CHECK-SAME: ; LinAlgVectorAccumulateToDescriptor(handle,offset,align,vector)
1414

15-
// CHECK2: call void @"dx.hl.op..void (i32, <4 x float>, %dx.types.Handle, i32, i32)"
16-
// CHECK2-SAME: (i32 423, <4 x float> %{{.*}}, %dx.types.Handle %{{.*}}, i32 16, i32 64)
15+
// CHECK2: call void @"dx.hl.op..void (i32, %dx.types.Handle, i32, i32, <4 x float>)"
16+
// CHECK2-SAME: (i32 423, %dx.types.Handle %{{.*}}, i32 16, i32 64, <4 x float> %{{.*}})
1717
float4 vec = {9.0, 8.0, 7.0, 6.0};
18-
__builtin_LinAlg_VectorAccumulateToDescriptor(vec, outbuf, 16, 64);
18+
__builtin_LinAlg_VectorAccumulateToDescriptor(outbuf, 16, 64, vec);
1919
}

tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-as.ll

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -85,7 +85,7 @@ define void @mainAS() {
8585
%v16 = call <4 x float> @dx.op.linAlgConvert.v4f32.v4i32(i32 -2147483618, <4 x i32> zeroinitializer, i32 1, i32 2) ; LinAlgConvert(inputVector,inputInterpretation,outputInterpretation)
8686

8787
; dx.op.linAlgVectorAccumulateToDescriptor
88-
call void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32 -2147483617, <4 x float> zeroinitializer, %dx.types.Handle %handle, i32 0, i32 64) ; LinAlgVectorAccumulateToDescriptor(vector,handle,offset,align)
88+
call void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32 -2147483617, %dx.types.Handle %handle, i32 0, i32 64, <4 x float> zeroinitializer) ; LinAlgVectorAccumulateToDescriptor(handle,offset,align,vector)
8989

9090
;
9191
; Built-ins restricted to compute, mesh and amplification shaders
@@ -167,7 +167,7 @@ declare <4 x i32> @dx.op.linAlgMatVecMulAdd.v4i32.mC4M5N4U0S2.v4i32.v4i32(i32, %
167167
declare <4 x float> @dx.op.linAlgConvert.v4f32.v4i32(i32, <4 x i32>, i32, i32) #0
168168

169169
; Function Attrs: nounwind
170-
declare void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32, <4 x float>, %dx.types.Handle, i32, i32) #0
170+
declare void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32, %dx.types.Handle, i32, i32, <4 x float>) #0
171171

172172
; Function Attrs: nounwind
173173
declare %dx.types.LinAlgMatrixC4M4N5U1S2 @dx.op.linAlgCopyConvertMatrix.mC4M4N5U1S2.mC4M5N4U0S2(i32, %dx.types.LinAlgMatrixC4M5N4U0S2, i1) #0

tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-cs.ll

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -58,7 +58,7 @@ define void @mainCS() {
5858
%v16 = call <4 x float> @dx.op.linAlgConvert.v4f32.v4i32(i32 -2147483618, <4 x i32> zeroinitializer, i32 1, i32 2) ; LinAlgConvert(inputVector,inputInterpretation,outputInterpretation)
5959

6060
; dx.op.linAlgVectorAccumulateToDescriptor
61-
call void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32 -2147483617, <4 x float> zeroinitializer, %dx.types.Handle %handle, i32 0, i32 64) ; LinAlgVectorAccumulateToDescriptor(vector,handle,offset,align)
61+
call void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32 -2147483617, %dx.types.Handle %handle, i32 0, i32 64, <4 x float> zeroinitializer) ; LinAlgVectorAccumulateToDescriptor(handle,offset,align,vector)
6262

6363
;
6464
; Built-ins restricted to compute, mesh and amplification shaders
@@ -140,7 +140,7 @@ declare <4 x i32> @dx.op.linAlgMatVecMulAdd.v4i32.mC4M5N4U0S2.v4i32.v4i32(i32, %
140140
declare <4 x float> @dx.op.linAlgConvert.v4f32.v4i32(i32, <4 x i32>, i32, i32) #0
141141

142142
; Function Attrs: nounwind
143-
declare void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32, <4 x float>, %dx.types.Handle, i32, i32) #0
143+
declare void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32, %dx.types.Handle, i32, i32, <4 x float>) #0
144144

145145
; Function Attrs: nounwind
146146
declare %dx.types.LinAlgMatrixC4M4N5U1S2 @dx.op.linAlgCopyConvertMatrix.mC4M4N5U1S2.mC4M5N4U0S2(i32, %dx.types.LinAlgMatrixC4M5N4U0S2, i1) #0

tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-ds.ll

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -86,7 +86,7 @@ define void @MainDS() {
8686
%v16 = call <4 x float> @dx.op.linAlgConvert.v4f32.v4i32(i32 -2147483618, <4 x i32> zeroinitializer, i32 1, i32 2) ; LinAlgConvert(inputVector,inputInterpretation,outputInterpretation)
8787

8888
; dx.op.linAlgVectorAccumulateToDescriptor
89-
call void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32 -2147483617, <4 x float> zeroinitializer, %dx.types.Handle %handle, i32 0, i32 64) ; LinAlgVectorAccumulateToDescriptor(vector,handle,offset,align)
89+
call void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32 -2147483617, %dx.types.Handle %handle, i32 0, i32 64, <4 x float> zeroinitializer) ; LinAlgVectorAccumulateToDescriptor(handle,offset,align,vector)
9090

9191
;
9292
; Built-ins restricted to compute, mesh and amplification shaders
@@ -173,7 +173,7 @@ declare <4 x i32> @dx.op.linAlgMatVecMulAdd.v4i32.mC4M5N4U0S2.v4i32.v4i32(i32, %
173173
declare <4 x float> @dx.op.linAlgConvert.v4f32.v4i32(i32, <4 x i32>, i32, i32) #0
174174

175175
; Function Attrs: nounwind
176-
declare void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32, <4 x float>, %dx.types.Handle, i32, i32) #0
176+
declare void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32, %dx.types.Handle, i32, i32, <4 x float>) #0
177177

178178
; Function Attrs: nounwind
179179
declare %dx.types.LinAlgMatrixC4M4N5U1S2 @dx.op.linAlgCopyConvertMatrix.mC4M4N5U1S2.mC4M5N4U0S2(i32, %dx.types.LinAlgMatrixC4M5N4U0S2, i1) #0

tools/clang/test/LitDXILValidation/LinAlgMatrix/linalgmatrix-gs.ll

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -86,7 +86,7 @@ define void @MainGS() {
8686
%v16 = call <4 x float> @dx.op.linAlgConvert.v4f32.v4i32(i32 -2147483618, <4 x i32> zeroinitializer, i32 1, i32 2) ; LinAlgConvert(inputVector,inputInterpretation,outputInterpretation)
8787

8888
; dx.op.linAlgVectorAccumulateToDescriptor
89-
call void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32 -2147483617, <4 x float> zeroinitializer, %dx.types.Handle %handle, i32 0, i32 64) ; LinAlgVectorAccumulateToDescriptor(vector,handle,offset,align)
89+
call void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32 -2147483617, %dx.types.Handle %handle, i32 0, i32 64, <4 x float> zeroinitializer) ; LinAlgVectorAccumulateToDescriptor(handle,offset,align,vector)
9090

9191
;
9292
; Built-ins restricted to compute, mesh and amplification shaders
@@ -172,7 +172,7 @@ declare <4 x i32> @dx.op.linAlgMatVecMulAdd.v4i32.mC4M5N4U0S2.v4i32.v4i32(i32, %
172172
declare <4 x float> @dx.op.linAlgConvert.v4f32.v4i32(i32, <4 x i32>, i32, i32) #0
173173

174174
; Function Attrs: nounwind
175-
declare void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32, <4 x float>, %dx.types.Handle, i32, i32) #0
175+
declare void @dx.op.linAlgVectorAccumulateToDescriptor.v4f32(i32, %dx.types.Handle, i32, i32, <4 x float>) #0
176176

177177
; Function Attrs: nounwind
178178
declare %dx.types.LinAlgMatrixC4M4N5U1S2 @dx.op.linAlgCopyConvertMatrix.mC4M4N5U1S2.mC4M5N4U0S2(i32, %dx.types.LinAlgMatrixC4M5N4U0S2, i1) #0

0 commit comments

Comments
 (0)