diff options
| author | Anders Leino <aleino@nvidia.com> | 2025-01-10 21:05:05 +0200 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2025-01-10 11:05:05 -0800 |
| commit | 803e0c9f9a9dc4b01e29ebbf3b37a5bba782ac83 (patch) | |
| tree | 4996c9f415c64692e8381ae8c9ab1ab914ee86ea /source/slang/slang-ir-legalize-binary-operator.cpp | |
| parent | 6437f2d37b08972db5e4515bd124639c2903dda1 (diff) | |
WGSL: Convert signed vector shift amounts to unsigned (#6023)
* WGSL: Fixes for signed shift amounts
- Handle the case of vector shift amounts
- Closes #5985
- Move handling of scalar case from emit to legalization
- Add tests for bitshifts.
* Move the binary operator legalization function to a common place
* Metal: Legalize binary operations
Closes #6029.
* Fix Metal filecheck test
The int shift amounts are now converted to unsigned.
* format code
---------
Co-authored-by: slangbot <186143334+slangbot@users.noreply.github.com>
Co-authored-by: Yong He <yonghe@outlook.com>
Diffstat (limited to 'source/slang/slang-ir-legalize-binary-operator.cpp')
| -rw-r--r-- | source/slang/slang-ir-legalize-binary-operator.cpp | 121 |
1 files changed, 121 insertions, 0 deletions
diff --git a/source/slang/slang-ir-legalize-binary-operator.cpp b/source/slang/slang-ir-legalize-binary-operator.cpp new file mode 100644 index 000000000..a1affb7e9 --- /dev/null +++ b/source/slang/slang-ir-legalize-binary-operator.cpp @@ -0,0 +1,121 @@ +#include "slang-ir-legalize-binary-operator.h" + +#include "slang-ir-insts.h" + +namespace Slang +{ + +void legalizeBinaryOp(IRInst* inst) +{ + // For shifts, ensure that the shift amount is unsigned, as required by + // https://www.w3.org/TR/WGSL/#bit-expr. + if (inst->getOp() == kIROp_Lsh || inst->getOp() == kIROp_Rsh) + { + IRInst* shiftAmount = inst->getOperand(1); + IRType* shiftAmountType = shiftAmount->getDataType(); + if (auto shiftAmountVectorType = as<IRVectorType>(shiftAmountType)) + { + IRType* shiftAmountElementType = shiftAmountVectorType->getElementType(); + IntInfo opIntInfo = getIntTypeInfo(shiftAmountElementType); + if (opIntInfo.isSigned) + { + IRBuilder builder(inst); + builder.setInsertBefore(inst); + opIntInfo.isSigned = false; + shiftAmountElementType = builder.getType(getIntTypeOpFromInfo(opIntInfo)); + shiftAmountVectorType = builder.getVectorType( + shiftAmountElementType, + shiftAmountVectorType->getElementCount()); + IRInst* newShiftAmount = builder.emitCast(shiftAmountVectorType, shiftAmount); + builder.replaceOperand(inst->getOperands() + 1, newShiftAmount); + } + } + else if (isIntegralType(shiftAmountType)) + { + IntInfo opIntInfo = getIntTypeInfo(shiftAmountType); + if (opIntInfo.isSigned) + { + IRBuilder builder(inst); + builder.setInsertBefore(inst); + opIntInfo.isSigned = false; + shiftAmountType = builder.getType(getIntTypeOpFromInfo(opIntInfo)); + IRInst* newShiftAmount = builder.emitCast(shiftAmountType, shiftAmount); + builder.replaceOperand(inst->getOperands() + 1, newShiftAmount); + } + } + } + + auto isVectorOrMatrix = [](IRType* type) + { + switch (type->getOp()) + { + case kIROp_VectorType: + case kIROp_MatrixType: + return true; + default: + return false; + } + }; + if (isVectorOrMatrix(inst->getOperand(0)->getDataType()) && + as<IRBasicType>(inst->getOperand(1)->getDataType())) + { + IRBuilder builder(inst); + builder.setInsertBefore(inst); + IRType* compositeType = inst->getOperand(0)->getDataType(); + IRInst* scalarValue = inst->getOperand(1); + // Retain the scalar type for shifts + if (inst->getOp() == kIROp_Lsh || inst->getOp() == kIROp_Rsh) + { + auto vectorType = as<IRVectorType>(compositeType); + compositeType = + builder.getVectorType(scalarValue->getDataType(), vectorType->getElementCount()); + } + auto newRhs = builder.emitMakeCompositeFromScalar(compositeType, scalarValue); + builder.replaceOperand(inst->getOperands() + 1, newRhs); + } + else if ( + as<IRBasicType>(inst->getOperand(0)->getDataType()) && + isVectorOrMatrix(inst->getOperand(1)->getDataType())) + { + IRBuilder builder(inst); + builder.setInsertBefore(inst); + IRType* compositeType = inst->getOperand(1)->getDataType(); + IRInst* scalarValue = inst->getOperand(0); + // Retain the scalar type for shifts + if (inst->getOp() == kIROp_Lsh || inst->getOp() == kIROp_Rsh) + { + auto vectorType = as<IRVectorType>(compositeType); + compositeType = + builder.getVectorType(scalarValue->getDataType(), vectorType->getElementCount()); + } + auto newLhs = builder.emitMakeCompositeFromScalar(compositeType, scalarValue); + builder.replaceOperand(inst->getOperands(), newLhs); + } + else if ( + isIntegralType(inst->getOperand(0)->getDataType()) && + isIntegralType(inst->getOperand(1)->getDataType())) + { + // Unless the operator is a shift, and if the integer operands differ in signedness, + // then convert the signed one to unsigned. + // We're assuming that the cases where this is bad have already been caught by + // common validation checks. + IntInfo opIntInfo[2] = { + getIntTypeInfo(inst->getOperand(0)->getDataType()), + getIntTypeInfo(inst->getOperand(1)->getDataType())}; + bool isShift = inst->getOp() == kIROp_Lsh || inst->getOp() == kIROp_Rsh; + bool signednessDiffers = opIntInfo[0].isSigned != opIntInfo[1].isSigned; + if (!isShift && signednessDiffers) + { + int signedOpIndex = (int)opIntInfo[1].isSigned; + opIntInfo[signedOpIndex].isSigned = false; + IRBuilder builder(inst); + builder.setInsertBefore(inst); + auto newOp = builder.emitCast( + builder.getType(getIntTypeOpFromInfo(opIntInfo[signedOpIndex])), + inst->getOperand(signedOpIndex)); + builder.replaceOperand(inst->getOperands() + signedOpIndex, newOp); + } + } +} + +} // namespace Slang |
