summaryrefslogtreecommitdiff
path: root/source/slang/slang-ir-legalize-binary-operator.cpp
diff options
context:
space:
mode:
authorAnders Leino <aleino@nvidia.com>2025-01-10 21:05:05 +0200
committerGitHub <noreply@github.com>2025-01-10 11:05:05 -0800
commit803e0c9f9a9dc4b01e29ebbf3b37a5bba782ac83 (patch)
tree4996c9f415c64692e8381ae8c9ab1ab914ee86ea /source/slang/slang-ir-legalize-binary-operator.cpp
parent6437f2d37b08972db5e4515bd124639c2903dda1 (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.cpp121
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