[llvm] [SelectionDAG] Improve wide integer squaring (PR #226403)
via llvm-commits
llvm-commits at lists.llvm.org
Fri Sep 25 02:01:52 PDT 2026
llvmorg-github-actions[bot] wrote:
<!--LLVM PR SUMMARY COMMENT-->
@llvm/pr-subscribers-llvm-selectiondag
Author: MITSUNARI Shigeo (herumi)
<details>
<summary>Changes</summary>
For a wide integer `mul x, x`, the number of limb multiplications is
n(n+1)/2: n for x[i]*x[i] (i=0, ..., n-1) and n(n-1)/2 for x[i]*x[j]
(i<j). The existing expansion already reaches this count. However, the
way the products are added has room for improvement. The existing
expansion adds each cross product x[i]*x[j] twice and materializes the
carry between partial sums with setcc/zext before adding it back, so the
add and carry-handling instructions outnumber the multiplications
several times over, and register pressure goes up as well.
This patch adds TargetLowering::expandWideSquare. It groups the products
x[i]*x[j] into rows by diagonal (the products with equal j-i) and
accumulates them from the narrowest width up as
`acc = (acc << U) + row`. The products on one diagonal do not
overlap, so a row is built by concatenation alone, adding a row fits in
a single adc chain, and the carries stay within the row. Finally, acc is
doubled (shl 1) and the concatenation of x[i]*x[i] is added. This
removes the materialized carries (setb/movzbl) and minimizes the number
of adc chains.
The limb type is getTypeToExpandTo() of the result type, i.e. the legal
integer type the result is ultimately expanded to (i64 on x86-64,
AArch64 and RISCV64, i32 on i386 and RISCV32). The expansion is used
only when this type has UMUL_LOHI or MULHU, and falls back to the
existing expansion otherwise. The operand is usually zero-extended (to
get all the bits of the square of an i384, it is zero-extended to i768
and then multiplied), so the limbs that computeKnownBits proves to be
zero are left out and the value is squared at its real width. When the
result is truncated (e.g. i512 = mul i512 %x, %x), the row widths are
clamped to the number of result limbs and the square is computed modulo
2^Bits (a product that reaches the top limb is computed with a plain mul
for its low half only).
The result is as follows.
The tables show the number of non-multiply instructions for an input of
n 64-bit limbs (an i(64n) with all limbs in use is zero-extended to
i(128n) and squared to get the full product; n=6 means i384 x i384 ->
i768), i.e. the whole function minus the multiply instructions. The
number of multiplications is the same n(n+1)/2 for org and opti: mulx on
x86-64, and the two instructions mul + umulh on AArch64. The value in
parentheses is the throughput time ratio to org (ns/op, 4 independent
streams; measured on a Xeon w9-3495X for x86-64 and an Apple M4 Pro for
AArch64).
x86-64 (+bmi2)
| n | 6 | 8 | 16 | 32 |
|---|---|---|---|---|
| org | 215 | 390 | 1654 | 6790 |
| opti | 138 (0.74x) | 246 (0.77x) | 1014 (0.77x) | 4037 (0.53x) |
Breakdown for n=6 (both use 21 mulx): org has adc 58 / add 38 /
setb 13 / movzbl 13, opti has adc 30 / add 6 / shld 9 / setb 0.
AArch64 (-mcpu=apple-m1)
| n | 6 | 8 | 16 | 32 |
|---|---|---|---|---|
| org | 145 | 293 | 1434 | 6344 |
| opti | 61 (0.62x) | 98 (0.48x) | 501 (0.49x) | 3117 (0.60x) |
Breakdown for n=6 (both use 21 mul + 21 umulh): org has adds 46 /
adcs 25 / cinc 36, opti has adds 5 / adcs 25 / cinc 4 / extr 9.
Tests: added llc output for `mul iK %x, %x` (including zext / sext
variants) on X86 / AArch64 / RISCV. The only existing test that changes
is the i192 square in X86/dagcombine-cse.ll, which gets shorter.
Assisted-by: Claude Code
---
Patch is 88.01 KiB, truncated to 20.00 KiB below, full version: https://github.com/llvm/llvm-project/pull/226403.diff
8 Files Affected:
- (modified) llvm/include/llvm/CodeGen/TargetLowering.h (+6)
- (modified) llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp (+7)
- (modified) llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp (+135)
- (added) llvm/test/CodeGen/AArch64/wide-int-square.ll (+288)
- (added) llvm/test/CodeGen/RISCV/wide-int-square.ll (+723)
- (modified) llvm/test/CodeGen/X86/dagcombine-cse.ll (+44-49)
- (added) llvm/test/CodeGen/X86/wide-int-square-i512.ll (+302)
- (added) llvm/test/CodeGen/X86/wide-int-square.ll (+809)
``````````diff
diff --git a/llvm/include/llvm/CodeGen/TargetLowering.h b/llvm/include/llvm/CodeGen/TargetLowering.h
index c08d0e53ec34b..9b56f98887791 100644
--- a/llvm/include/llvm/CodeGen/TargetLowering.h
+++ b/llvm/include/llvm/CodeGen/TargetLowering.h
@@ -6026,6 +6026,12 @@ class LLVM_ABI TargetLowering : public TargetLoweringBase {
SDValue HiLHS = SDValue(),
SDValue HiRHS = SDValue()) const;
+ /// Expand MUL X, X of an illegal scalar integer type. Only the low bits of
+ /// the result type are computed.
+ /// \param N Node to expand
+ /// \returns The expansion if successful, SDValue() otherwise
+ SDValue expandWideSquare(SDNode *N, SelectionDAG &DAG) const;
+
/// Calculate full product of LHS and RHS either via a libcall or through
/// brute force expansion of the multiplication. The expansion works by
/// splitting the 2 inputs into 4 pieces that we can multiply and add together
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp b/llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp
index a90bb06fb424d..2c9f0c90e553b 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeIntegerTypes.cpp
@@ -4556,6 +4556,13 @@ void DAGTypeLegalizer::ExpandIntRes_MUL(SDNode *N,
RTLIB::Libcall LC = RTLIB::getMUL(VT);
RTLIB::LibcallImpl LCImpl = DAG.getLibcalls().getLibcallImpl(LC);
if (LCImpl == RTLIB::Unsupported) {
+ // A square needs only half of the limb products.
+ if (N->getOperand(0) == N->getOperand(1)) {
+ if (SDValue Sq = TLI.expandWideSquare(N, DAG)) {
+ SplitInteger(Sq, Lo, Hi);
+ return;
+ }
+ }
// Perform a wide multiplication where the wide type is the original VT and
// the 4 parts are the split arguments.
TLI.forceExpandMultiply(DAG, dl, /*Signed=*/false, Lo, Hi, LL, RL, LH, RH);
diff --git a/llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp b/llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp
index 7c97dfb2c806a..e3bbfafa74739 100644
--- a/llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/TargetLowering.cpp
@@ -12634,6 +12634,141 @@ void TargetLowering::forceExpandMultiply(SelectionDAG &DAG, const SDLoc &dl,
}
}
+SDValue TargetLowering::expandWideSquare(SDNode *N, SelectionDAG &DAG) const {
+ assert(N->getOpcode() == ISD::MUL && N->getOperand(0) == N->getOperand(1) &&
+ "Expected a squaring MUL");
+ EVT VT = N->getValueType(0);
+ if (!VT.isScalarInteger())
+ return SDValue();
+ LLVMContext &Ctx = *DAG.getContext();
+ SDLoc dl(N);
+
+ // The limbs have the legal type that VT is ultimately expanded to.
+ EVT LimbVT = getTypeToExpandTo(Ctx, VT);
+ unsigned U = LimbVT.getSizeInBits();
+ unsigned Bits = VT.getSizeInBits();
+ if (Bits < 4 * U)
+ return SDValue();
+ unsigned R = Bits / U; // number of result limbs
+
+ bool HasUMUL_LOHI = isOperationLegalOrCustom(ISD::UMUL_LOHI, LimbVT);
+ bool HasMULHU = isOperationLegalOrCustom(ISD::MULHU, LimbVT);
+ if (!HasUMUL_LOHI && !HasMULHU)
+ return SDValue();
+
+ // X is usually zero-extended (the way to request the full product), so
+ // square it at its real width: only the limbs that may be nonzero take part.
+ SDValue X = N->getOperand(0);
+ unsigned NumLimbs =
+ divideCeil(Bits - DAG.computeKnownBits(X).countMinLeadingZeros(), U);
+ if (NumLimbs < 2)
+ return SDValue();
+ NumLimbs = std::min(NumLimbs, R);
+
+ // Split X into limbs in least-significant-first order.
+ SmallVector<SDValue, 16> Limb(NumLimbs);
+ for (unsigned I = 0; I != NumLimbs; ++I) {
+ SDValue V = DAG.getNode(ISD::SRL, dl, VT, X,
+ DAG.getShiftAmountConstant(I * U, VT, dl));
+ Limb[I] = DAG.getNode(ISD::TRUNCATE, dl, LimbVT, V);
+ }
+
+ EVT Prod2VT = EVT::getIntegerVT(Ctx, 2 * U);
+ auto LimbsVT = [&](unsigned K) { return EVT::getIntegerVT(Ctx, K * U); };
+
+ // Return A * B as BUILD_PAIR(low, high), or just low.
+ auto MakeProd = [&](SDValue A, SDValue B, bool LowOnly) -> SDValue {
+ if (LowOnly)
+ return DAG.getNode(ISD::MUL, dl, LimbVT, A, B);
+ SDValue PLo, PHi;
+ if (HasUMUL_LOHI) {
+ PLo =
+ DAG.getNode(ISD::UMUL_LOHI, dl, DAG.getVTList(LimbVT, LimbVT), A, B);
+ PHi = PLo.getValue(1);
+ } else {
+ PLo = DAG.getNode(ISD::MUL, dl, LimbVT, A, B);
+ PHi = DAG.getNode(ISD::MULHU, dl, LimbVT, A, B);
+ }
+ return DAG.getNode(ISD::BUILD_PAIR, dl, Prod2VT, PLo, PHi);
+ };
+
+ // Zero-extend V and place it at limb offset Off.
+ auto Place = [&](SDValue V, unsigned Off, EVT WideVT) -> SDValue {
+ V = DAG.getNode(ISD::ZERO_EXTEND, dl, WideVT, V);
+ return DAG.getNode(ISD::SHL, dl, WideVT, V,
+ DAG.getShiftAmountConstant(Off * U, WideVT, dl));
+ };
+
+ // Concatenate parts that do not overlap.
+ auto Pack = [&](ArrayRef<std::pair<SDValue, unsigned>> Parts, EVT WideVT) {
+ SDNodeFlags Flags;
+ Flags.setDisjoint(true);
+ SDValue Acc;
+ for (const auto &[V, Off] : Parts) {
+ SDValue P = Place(V, Off, WideVT);
+ Acc = Acc ? DAG.getNode(ISD::OR, dl, WideVT, Acc, P, Flags) : P;
+ }
+ return Acc;
+ };
+
+ // Write the 2U-bit product X[i] * X[j] as i.j and the concatenation of
+ // non-overlapping values (least significant on the right) as ||.
+ // For NumLimbs = 4:
+ //
+ // X * X = S + ((2 * Acc) << U)
+ // S := 3.3 || 2.2 || 1.1 || 0.0
+ // Acc := 3.0 // D = NumLimbs - 1 = 3
+ // Acc := (Acc << U) + (3.1 || 2.0) // D = 2
+ // Acc := (Acc << U) + (3.2 || 2.1 || 1.0) // D = 1
+ //
+ // Each parenthesized row is the diagonal j - i = D below: its products
+ // sit at limbs 2i, so they tile without overlap. Acc is accumulated row by
+ // row as in the loop, shifting the sum so far by one limb before adding the
+ // next row. The sum never outgrows the row, because the shifted sum is
+ // one limb narrower than the row and the high limb of a product is at most
+ // 2^U - 2. Row widths are clamped to the R result limbs, which computes the
+ // square modulo 2^Bits.
+ SDValue Acc;
+ for (unsigned D = NumLimbs - 1; D > 0; --D) {
+ unsigned RowLimbs = std::min(2 * (NumLimbs - D), R - D);
+ SmallVector<std::pair<SDValue, unsigned>, 8> Parts;
+ for (unsigned I = 0; I + D < NumLimbs; ++I) {
+ unsigned Pos = 2 * I;
+ if (Pos >= RowLimbs)
+ break;
+ Parts.push_back(
+ {MakeProd(Limb[I], Limb[I + D], /*LowOnly=*/Pos + 1 >= RowLimbs),
+ Pos});
+ }
+ EVT RowVT = LimbsVT(RowLimbs);
+ SDValue Row = Pack(Parts, RowVT);
+ if (!Acc)
+ Acc = Row;
+ else
+ Acc = DAG.getNode(ISD::ADD, dl, RowVT, Place(Acc, 1, RowVT), Row);
+ }
+
+ // Compute (Acc * 2) << U using a shift.
+ unsigned W = std::min(2 * NumLimbs, R); // result limbs that can be nonzero
+ EVT DblVT = LimbsVT(W - 1);
+ SDValue Dbl = DAG.getNode(ISD::SHL, dl, DblVT, Place(Acc, 0, DblVT),
+ DAG.getShiftAmountConstant(1, DblVT, dl));
+ EVT ZVT = LimbsVT(W);
+ SDValue Cross = Place(Dbl, 1, ZVT);
+
+ // Compute S in the formula above. The squares are created last so that
+ // they are scheduled close to their use.
+ SmallVector<std::pair<SDValue, unsigned>, 8> DiagParts;
+ for (unsigned I = 0; 2 * I < W; ++I)
+ DiagParts.push_back(
+ {MakeProd(Limb[I], Limb[I], /*LowOnly=*/2 * I + 1 >= W), 2 * I});
+ SDValue Z = DAG.getNode(ISD::ADD, dl, ZVT, Cross, Pack(DiagParts, ZVT));
+
+ if (ZVT != VT)
+ Z = DAG.getNode(ISD::ZERO_EXTEND, dl, VT, Z);
+ return Z;
+}
+
void TargetLowering::forceExpandWideMUL(SelectionDAG &DAG, const SDLoc &dl,
bool Signed, const SDValue LHS,
const SDValue RHS, SDValue &Lo,
diff --git a/llvm/test/CodeGen/AArch64/wide-int-square.ll b/llvm/test/CodeGen/AArch64/wide-int-square.ll
new file mode 100644
index 0000000000000..adfd9efc2d73a
--- /dev/null
+++ b/llvm/test/CodeGen/AArch64/wide-int-square.ll
@@ -0,0 +1,288 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 5
+; RUN: llc < %s -mtriple=aarch64 | FileCheck %s
+
+; truncated square, 4 x 4 limbs
+define void @sqr_i256(ptr %out, ptr %in) {
+; CHECK-LABEL: sqr_i256:
+; CHECK: // %bb.0:
+; CHECK-NEXT: ldp x10, x8, [x1, #8]
+; CHECK-NEXT: ldr x9, [x1]
+; CHECK-NEXT: ldr x12, [x1, #24]
+; CHECK-NEXT: umulh x16, x9, x9
+; CHECK-NEXT: umulh x11, x9, x8
+; CHECK-NEXT: mul x13, x10, x8
+; CHECK-NEXT: madd x11, x9, x12, x11
+; CHECK-NEXT: umulh x12, x9, x10
+; CHECK-NEXT: mul x8, x9, x8
+; CHECK-NEXT: mul x14, x9, x10
+; CHECK-NEXT: mul x15, x10, x10
+; CHECK-NEXT: adds x8, x8, x12
+; CHECK-NEXT: mul x9, x9, x9
+; CHECK-NEXT: adc x11, x11, x13
+; CHECK-NEXT: adds x12, x16, x14, lsl #1
+; CHECK-NEXT: extr x13, x8, x14, #63
+; CHECK-NEXT: extr x8, x11, x8, #63
+; CHECK-NEXT: umulh x10, x10, x10
+; CHECK-NEXT: stp x9, x12, [x0]
+; CHECK-NEXT: adcs x9, x13, x15
+; CHECK-NEXT: adc x8, x8, x10
+; CHECK-NEXT: stp x9, x8, [x0, #16]
+; CHECK-NEXT: ret
+ %x = load i256, ptr %in
+ %r = mul i256 %x, %x
+ store i256 %r, ptr %out
+ ret void
+}
+
+; full square, 4 limbs -> 8 limbs
+define void @sqr_i512_zext_i256(ptr %out, ptr %in) {
+; CHECK-LABEL: sqr_i512_zext_i256:
+; CHECK: // %bb.0:
+; CHECK-NEXT: str x19, [sp, #-16]! // 8-byte Folded Spill
+; CHECK-NEXT: .cfi_def_cfa_offset 16
+; CHECK-NEXT: .cfi_offset w19, -16
+; CHECK-NEXT: ldp x13, x8, [x1, #16]
+; CHECK-NEXT: ldp x11, x9, [x1]
+; CHECK-NEXT: mul x18, x13, x8
+; CHECK-NEXT: umulh x15, x11, x13
+; CHECK-NEXT: mul x16, x11, x8
+; CHECK-NEXT: mul x12, x9, x8
+; CHECK-NEXT: umulh x14, x11, x8
+; CHECK-NEXT: adds x15, x16, x15
+; CHECK-NEXT: umulh x10, x9, x8
+; CHECK-NEXT: umulh x3, x11, x9
+; CHECK-NEXT: adcs x12, x14, x12
+; CHECK-NEXT: mul x4, x11, x13
+; CHECK-NEXT: cinc x10, x10, hs
+; CHECK-NEXT: mul x2, x9, x13
+; CHECK-NEXT: umulh x1, x9, x13
+; CHECK-NEXT: adds x16, x4, x3
+; CHECK-NEXT: umulh x17, x13, x8
+; CHECK-NEXT: adcs x15, x15, x2
+; CHECK-NEXT: mul x5, x11, x9
+; CHECK-NEXT: adcs x12, x12, x1
+; CHECK-NEXT: umulh x14, x11, x11
+; CHECK-NEXT: adcs x10, x10, x18
+; CHECK-NEXT: extr x1, x12, x15, #63
+; CHECK-NEXT: cinc x17, x17, hs
+; CHECK-NEXT: extr x15, x15, x16, #63
+; CHECK-NEXT: extr x12, x10, x12, #63
+; CHECK-NEXT: mul x19, x9, x9
+; CHECK-NEXT: extr x10, x17, x10, #63
+; CHECK-NEXT: extr x18, x16, x5, #63
+; CHECK-NEXT: mul x11, x11, x11
+; CHECK-NEXT: adds x14, x14, x5, lsl #1
+; CHECK-NEXT: umulh x9, x9, x9
+; CHECK-NEXT: mul x7, x13, x13
+; CHECK-NEXT: stp x11, x14, [x0]
+; CHECK-NEXT: adcs x11, x18, x19
+; CHECK-NEXT: umulh x13, x13, x13
+; CHECK-NEXT: adcs x9, x15, x9
+; CHECK-NEXT: mul x6, x8, x8
+; CHECK-NEXT: stp x11, x9, [x0, #16]
+; CHECK-NEXT: lsr x9, x17, #63
+; CHECK-NEXT: adcs x11, x1, x7
+; CHECK-NEXT: umulh x8, x8, x8
+; CHECK-NEXT: adcs x12, x12, x13
+; CHECK-NEXT: stp x11, x12, [x0, #32]
+; CHECK-NEXT: adcs x10, x10, x6
+; CHECK-NEXT: adc x8, x9, x8
+; CHECK-NEXT: stp x10, x8, [x0, #48]
+; CHECK-NEXT: ldr x19, [sp], #16 // 8-byte Folded Reload
+; CHECK-NEXT: ret
+ %x = load i256, ptr %in
+ %z = zext i256 %x to i512
+ %r = mul i512 %z, %z
+ store i512 %r, ptr %out
+ ret void
+}
+
+; full square through a promoted type (i768 -> i1024), 6 limbs -> 12 limbs
+define void @sqr_i768_zext_i384(ptr %out, ptr %in) {
+; CHECK-LABEL: sqr_i768_zext_i384:
+; CHECK: // %bb.0:
+; CHECK-NEXT: sub sp, sp, #128
+; CHECK-NEXT: stp x29, x30, [sp, #32] // 16-byte Folded Spill
+; CHECK-NEXT: stp x28, x27, [sp, #48] // 16-byte Folded Spill
+; CHECK-NEXT: stp x26, x25, [sp, #64] // 16-byte Folded Spill
+; CHECK-NEXT: stp x24, x23, [sp, #80] // 16-byte Folded Spill
+; CHECK-NEXT: stp x22, x21, [sp, #96] // 16-byte Folded Spill
+; CHECK-NEXT: stp x20, x19, [sp, #112] // 16-byte Folded Spill
+; CHECK-NEXT: .cfi_def_cfa_offset 128
+; CHECK-NEXT: .cfi_offset w19, -8
+; CHECK-NEXT: .cfi_offset w20, -16
+; CHECK-NEXT: .cfi_offset w21, -24
+; CHECK-NEXT: .cfi_offset w22, -32
+; CHECK-NEXT: .cfi_offset w23, -40
+; CHECK-NEXT: .cfi_offset w24, -48
+; CHECK-NEXT: .cfi_offset w25, -56
+; CHECK-NEXT: .cfi_offset w26, -64
+; CHECK-NEXT: .cfi_offset w27, -72
+; CHECK-NEXT: .cfi_offset w28, -80
+; CHECK-NEXT: .cfi_offset w30, -88
+; CHECK-NEXT: .cfi_offset w29, -96
+; CHECK-NEXT: ldp x9, x8, [x1, #32]
+; CHECK-NEXT: ldp x13, x10, [x1]
+; CHECK-NEXT: ldp x11, x12, [x1, #16]
+; CHECK-NEXT: umulh x15, x9, x8
+; CHECK-NEXT: umulh x7, x13, x9
+; CHECK-NEXT: mul x19, x13, x8
+; CHECK-NEXT: mul x5, x10, x8
+; CHECK-NEXT: umulh x6, x13, x8
+; CHECK-NEXT: adds x7, x19, x7
+; CHECK-NEXT: umulh x17, x10, x8
+; CHECK-NEXT: umulh x23, x13, x12
+; CHECK-NEXT: adcs x5, x6, x5
+; CHECK-NEXT: mul x24, x13, x9
+; CHECK-NEXT: cinc x17, x17, hs
+; CHECK-NEXT: mul x14, x9, x8
+; CHECK-NEXT: mul x22, x10, x9
+; CHECK-NEXT: adds x19, x24, x23
+; CHECK-NEXT: umulh x21, x10, x9
+; CHECK-NEXT: stp x14, x15, [sp, #16] // 16-byte Folded Spill
+; CHECK-NEXT: mul x20, x11, x8
+; CHECK-NEXT: adcs x7, x7, x22
+; CHECK-NEXT: umulh x3, x11, x8
+; CHECK-NEXT: adcs x5, x5, x21
+; CHECK-NEXT: umulh x29, x13, x11
+; CHECK-NEXT: adcs x20, x17, x20
+; CHECK-NEXT: mul x30, x13, x12
+; CHECK-NEXT: cinc x21, x3, hs
+; CHECK-NEXT: mul x28, x10, x12
+; CHECK-NEXT: umulh x15, x12, x9
+; CHECK-NEXT: adds x23, x30, x29
+; CHECK-NEXT: ldp x29, x30, [sp, #32] // 16-byte Folded Reload
+; CHECK-NEXT: mul x14, x12, x9
+; CHECK-NEXT: adcs x19, x19, x28
+; CHECK-NEXT: umulh x27, x10, x12
+; CHECK-NEXT: mul x26, x11, x9
+; CHECK-NEXT: stp x14, x15, [sp] // 16-byte Folded Spill
+; CHECK-NEXT: umulh x25, x11, x9
+; CHECK-NEXT: adcs x7, x7, x27
+; CHECK-NEXT: ldp x28, x27, [sp, #48] // 16-byte Folded Reload
+; CHECK-NEXT: mul x4, x12, x8
+; CHECK-NEXT: adcs x5, x5, x26
+; CHECK-NEXT: umulh x2, x12, x8
+; CHECK-NEXT: adcs x20, x20, x25
+; CHECK-NEXT: ldp x26, x25, [sp, #64] // 16-byte Folded Reload
+; CHECK-NEXT: umulh x14, x13, x10
+; CHECK-NEXT: adcs x4, x21, x4
+; CHECK-NEXT: mul x6, x13, x11
+; CHECK-NEXT: cinc x2, x2, hs
+; CHECK-NEXT: mul x15, x10, x11
+; CHECK-NEXT: umulh x16, x10, x11
+; CHECK-NEXT: adds x14, x6, x14
+; CHECK-NEXT: mul x18, x11, x12
+; CHECK-NEXT: adcs x15, x23, x15
+; CHECK-NEXT: umulh x1, x11, x12
+; CHECK-NEXT: adcs x16, x19, x16
+; CHECK-NEXT: mul x22, x13, x10
+; CHECK-NEXT: adcs x18, x7, x18
+; CHECK-NEXT: umulh x6, x13, x13
+; CHECK-NEXT: adcs x1, x5, x1
+; CHECK-NEXT: ldp x5, x19, [sp] // 16-byte Folded Reload
+; CHECK-NEXT: mul x7, x10, x10
+; CHECK-NEXT: adcs x5, x20, x5
+; CHECK-NEXT: umulh x10, x10, x10
+; CHECK-NEXT: adcs x4, x4, x19
+; CHECK-NEXT: ldr x19, [sp, #16] // 8-byte Reload
+; CHECK-NEXT: mul x21, x11, x11
+; CHECK-NEXT: extr x20, x4, x5, #63
+; CHECK-NEXT: adcs x2, x2, x19
+; CHECK-NEXT: ldr x19, [sp, #24] // 8-byte Reload
+; CHECK-NEXT: mul x13, x13, x13
+; CHECK-NEXT: cinc x19, x19, hs
+; CHECK-NEXT: adds x6, x6, x22, lsl #1
+; CHECK-NEXT: extr x22, x14, x22, #63
+; CHECK-NEXT: umulh x11, x11, x11
+; CHECK-NEXT: extr x14, x15, x14, #63
+; CHECK-NEXT: extr x15, x16, x15, #63
+; CHECK-NEXT: adcs x7, x22, x7
+; CHECK-NEXT: extr x16, x18, x16, #63
+; CHECK-NEXT: mul x24, x12, x12
+; CHECK-NEXT: adcs x10, x14, x10
+; CHECK-NEXT: stp x13, x6, [x0]
+; CHECK-NEXT: extr x13, x1, x18, #63
+; CHECK-NEXT: adcs x14, x15, x21
+; CHECK-NEXT: umulh x12, x12, x12
+; CHECK-NEXT: stp x7, x10, [x0, #16]
+; CHECK-NEXT: extr x10, x5, x1, #63
+; CHECK-NEXT: adcs x11, x16, x11
+; CHECK-NEXT: ldp x22, x21, [sp, #96] // 16-byte Folded Reload
+; CHECK-NEXT: mul x3, x9, x9
+; CHECK-NEXT: stp x14, x11, [x0, #32]
+; CHECK-NEXT: extr x11, x2, x4, #63
+; CHECK-NEXT: adcs x13, x13, x24
+; CHECK-NEXT: ldp x24, x23, [sp, #80] // 16-byte Folded Reload
+; CHECK-NEXT: umulh x9, x9, x9
+; CHECK-NEXT: adcs x10, x10, x12
+; CHECK-NEXT: extr x12, x19, x2, #63
+; CHECK-NEXT: mul x17, x8, x8
+; CHECK-NEXT: stp x13, x10, [x0, #48]
+; CHECK-NEXT: lsr x10, x19, #63
+; CHECK-NEXT: adcs x13, x20, x3
+; CHECK-NEXT: ldp x20, x19, [sp, #112] // 16-byte Folded Reload
+; CHECK-NEXT: umulh x8, x8, x8
+; CHECK-NEXT: adcs x9, x11, x9
+; CHECK-NEXT: stp x13, x9, [x0, #64]
+; CHECK-NEXT: adcs x11, x12, x17
+; CHECK-NEXT: adc x8, x10, x8
+; CHECK-NEXT: stp x11, x8, [x0, #80]
+; CHECK-NEXT: add sp, sp, #128
+; CHECK-NEXT: ret
+ %x = load i384, ptr %in
+ %z = zext i384 %x to i768
+ %r = mul i768 %z, %z
+ store i768 %r, ptr %out
+ ret void
+}
+
+; sign-extended operand: no known-zero limbs, truncated 4 x 4
+define void @sqr_i256_sext_i128(ptr %out, ptr %in) {
+; CHECK-LABEL: sqr_i256_sext_i128:
+; CHECK: // %bb.0:
+; CHECK-NEXT: ldp x9, x8, [x1]
+; CHECK-NEXT: asr x10, x8, #63
+; CHECK-NEXT: umulh x13, x9, x8
+; CHECK-NEXT: umulh x11, x9, x10
+; CHECK-NEXT: mul x12, x9, x10
+; CHECK-NEXT: mul x10, x8, x10
+; CHECK-NEXT: mul x14, x9, x8
+; CHECK-NEXT: add x11, x12, x11
+; CHECK-NEXT: adds x12, x12, x13
+; CHECK-NEXT: umulh x16, x9, x9
+; CHECK-NEXT: adc x10, x11, x10
+; CHECK-NEXT: mul x15, x8, x8
+; CHECK-NEXT: extr x10, x10, x12, #63
+; CHECK-NEXT: extr x13, x12, x14, #63
+; CHECK-NEXT: mul x9, x9, x9
+; CHECK-NEXT: adds x11, x16, x14, lsl #1
+; CHECK-NEXT: umulh x8, x8, x8
+; CHECK-NEXT: stp x9, x11, [x0]
+; CHECK-NEXT: adcs x9, x13, x15
+; CHECK-NEXT: adc x8, x10, x8
+; CHECK-NEXT: stp x9, x8, [x0, #16]
+; CHECK-NEXT: ret
+ %x = load i128, ptr %in
+ %z = sext i128 %x to i256
+ %r = mul i256 %z, %z
+ store i256 %r, ptr %out
+ ret void
+}
+
+; a single limb: one product, not expanded here
+define void @sqr_i256_zext_i64(ptr %out, ptr %in) {
+; CHECK-LABEL: sqr_i256_zext_i64:
+; CHECK: // %bb.0:
+; CHECK-NEXT: ldr x8, [x1]
+; CHECK-NEXT: stp xzr, xzr, [x0, #16]
+; CHECK-NEXT: umulh x9, x8, x8
+; CHECK-NEXT: mul x8, x8, x8
+; CHECK-NEXT: stp x8, x9, [x0]
+; CHECK-NEXT: ret
+ %x = load i64, ptr %in
+ %z = zext i64 %x to i256
+ %r = mul i256 %z, %z
+ store i256 %r, ptr %out
+ ret void
+}
diff --git a/llvm/test/CodeGen/RISCV/wide-int-square.ll b/llvm/test/CodeGen/RISCV/wide-int-square.ll
new file mode 100644
index 0000000000000..c66d4249da362
--- /dev/null
+++ b/llvm/test/CodeGen/RISCV/wide-int-square.ll
@@ -0,0 +1,723 @@
+; NOTE: Assertions have been autogenerated by utils/update_llc_test_checks.py UTC_ARGS: --version 5
+; RUN: llc < %s -mtriple=riscv64 -mattr=+m | FileCheck %s
+
+; truncated square, 4 x 4 limbs
+define void @sqr_i256(ptr %out, ptr %in) {
+; CHECK-LABEL: sqr_i256:
+; CHECK: # %bb.0:
+; CHECK-NEXT: ld a2, 0(a1)
+; CHECK-NEXT: ld a3, 8(a1)
+; CHECK-NEXT: ld a4, 16(a1)
+; CHECK-NEXT: ld a1, 24(a1)
+; CHECK-NEXT: mulhu a5, a2, a3
+; CHECK-NEXT: mul a6, a2, a4
+; CHECK-NEXT: mulhu a7, a2, a4
+; CHECK-NEXT: mul a1, a2, a1
+; CHECK-NEXT: mul t0, a2, a3
+; CHECK-NEXT: mul a4, a3, a4
+; CHECK-NEXT: mul t1, a3, a3
+; CHECK-NEXT: mulhu t2, a2, a2
+; CHECK-NEXT: mulhu a3, a3, a3
+; CHECK-NEXT: mul a2, a2, a2
+; CHECK-NEXT: add a5, a6, a5
+; CHECK-NEXT: add a1, a1, a7
+; CHECK-NEXT: srli a7, t0, 63
+; CHECK-NEXT: slli t0, t0, 1
+; CHECK-NEXT: add a1, a1, a4
+; CHECK-NEXT: slli a4, a5, 1
+; CHECK-NEXT: sltu a6, a5, a6
+; CHECK-NEXT: srli a5, a5, 63
+; CHECK-NEXT: add t2, t0, t2
+; CHECK-NEXT: or a4, a4, a7
+; CHECK-NEXT: add a1, a1, a6
+; CHECK-NEXT: sltu a6, t2, t0
+; CHECK-NEXT: add t1, a4, t1
+; CHECK-NEXT: slli a1, a1, 1
+; CHECK-NEXT: sltu a4, t1, a4
+; CHECK-NEXT: or a1, a1, a5
+; CHECK-NEXT: add a6, t1, a6
+; CHECK-NEXT: add a1, a1, a3
+; CHECK-NEXT: sltu a3, a6, t1
+; CHECK-NEXT: add a1, a1, a4
+; CHECK-NEXT: add a1, a1, a3
+; CHECK-NEXT: sd a2, 0(a0)
+; CHECK-NEXT: sd t2, 8(a0)
+; CHECK-NEXT: sd a6, 16(a0)...
[truncated]
``````````
</details>
https://github.com/llvm/llvm-project/pull/226403
More information about the llvm-commits
mailing list