Extend the long-multiply fold to a third carry shape, where the cross
sum and the low-half accumulation each carry through a separate
unsigned overflow compare:
xh*yh + (cross_sum >> N) + carry_low + ((hilo > cross_sum) << N)
cross_sum = xh*yl + xl*yh
carry_low = (xl*yl + (cross_sum << N)) < (cross_sum << N)
The new HIGH_PART row consumes a LMK_CARRY_LOW summand alongside the
existing LMK_CARRY_CROSS_SUM, and an extra check validates that the
two LMK_CARRY_LOW operands are a (cross_shifted, mul_lolo) pair
consistent with the multiset's canonical (op0, op1).
gcc/ChangeLog:
* match.pd: Add mul_carry_low atom recognizer.
* tree-ssa-forwprop.cc (gimple_mul_carry_low): Declare.
(enum long_mul_kind): Add LMK_CARRY_LOW.
(struct long_mul_summand): Add carry_a / carry_b for the
LMK_CARRY_LOW summand's two operands.
(long_mul_set_summand): Handle LMK_CARRY_LOW.
(long_mul_classify_carry): Try mul_carry_low after the
more-constrained mul_carry_low_sum / mul_carry_cross_sum.
(long_mul_find_summand): New helper, find the first summand
with a given kind.
(long_mul_check_two_carries): New extra check, validates the
LMK_CARRY_LOW summand's operands against the canonical
(op0, op1).
(long_mul_table): Add the two-carry HIGH_PART row, gated on
long_mul_check_two_carries.
gcc/testsuite/ChangeLog:
* gcc.dg/tree-ssa/long-mul-two-carry.c: New test.
Signed-off-by: Konstantinos Eleftheriou <[email protected]>
---
(no changes since v1)
gcc/match.pd | 7 ++
.../gcc.dg/tree-ssa/long-mul-two-carry.c | 112 ++++++++++++++++++
gcc/tree-ssa-forwprop.cc | 81 ++++++++++++-
3 files changed, 198 insertions(+), 2 deletions(-)
create mode 100644 gcc/testsuite/gcc.dg/tree-ssa/long-mul-two-carry.c
diff --git a/gcc/match.pd b/gcc/match.pd
index b2ddbec50181..85a60848b2e6 100644
--- a/gcc/match.pd
+++ b/gcc/match.pd
@@ -12242,6 +12242,13 @@ and,
&& TYPE_PRECISION (op_type) % 2 == 0
&& tree_fits_uhwi_p (@0)
&& tree_to_uhwi (@0) == TYPE_PRECISION (op_type) / 2))))
+/* Carry from addition overflow: (cast?) (a > a + b).
+ :c on gt also matches the LT form: (cast?) (a + b < a). */
+(match (mul_carry_low @0 @1)
+ (convert?
+ (gt:c @0 (plus:c @1 @0)))
+ (with { tree op_type = TREE_TYPE (@0); }
+ (if (INTEGRAL_TYPE_P (op_type) && TYPE_UNSIGNED (op_type)))))
/* Low accumulate: (xl*yl >> N) + (cross_sum & mask). */
(match (mul_low_accum @op0 @op1 @mul_hilo0 @mul_hilo1 @0 @1)
(plus:c
diff --git a/gcc/testsuite/gcc.dg/tree-ssa/long-mul-two-carry.c
b/gcc/testsuite/gcc.dg/tree-ssa/long-mul-two-carry.c
new file mode 100644
index 000000000000..288363429dd9
--- /dev/null
+++ b/gcc/testsuite/gcc.dg/tree-ssa/long-mul-two-carry.c
@@ -0,0 +1,112 @@
+/* { dg-do compile } */
+/* { dg-options "-O3 -fdump-tree-forwprop-details" } */
+
+typedef __UINT32_TYPE__ uint32_t;
+typedef __UINT64_TYPE__ uint64_t;
+
+/* High part using two separate carries (cross carry + low carry). */
+uint64_t mulh_two_carry (uint64_t x, uint64_t y)
+{
+ uint64_t x_hi = x >> 32;
+ uint64_t x_lo = x & 0xFFFFFFFFUL;
+ uint64_t y_hi = y >> 32;
+ uint64_t y_lo = y & 0xFFFFFFFFUL;
+
+ uint64_t lolo = x_lo * y_lo;
+ uint64_t hilo = x_hi * y_lo;
+ uint64_t lohi = x_lo * y_hi;
+ uint64_t hihi = x_hi * y_hi;
+
+ uint64_t cross_sum = hilo + lohi;
+ uint64_t cross_carry = (uint64_t)(cross_sum < hilo) << 32;
+
+ uint64_t cross_shifted = cross_sum << 32;
+ uint64_t low_result = lolo + cross_shifted;
+ uint64_t low_carry = (uint64_t)(low_result < cross_shifted);
+
+ uint64_t high = hihi + (cross_sum >> 32) + cross_carry + low_carry;
+
+ return high;
+}
+
+/* Commuted operand order. */
+uint64_t mulh_two_carry_comm (uint64_t x, uint64_t y)
+{
+ uint64_t x_hi = x >> 32;
+ uint64_t x_lo = x & 0xFFFFFFFFUL;
+ uint64_t y_hi = y >> 32;
+ uint64_t y_lo = y & 0xFFFFFFFFUL;
+
+ uint64_t lolo = x_lo * y_lo;
+ uint64_t lohi = x_lo * y_hi;
+ uint64_t hilo = x_hi * y_lo;
+ uint64_t hihi = x_hi * y_hi;
+
+ uint64_t cross_sum = lohi + hilo;
+ uint64_t cross_carry = (uint64_t)(cross_sum < lohi) << 32;
+
+ uint64_t cross_shifted = cross_sum << 32;
+ uint64_t low_result = cross_shifted + lolo;
+ uint64_t low_carry = (uint64_t)(low_result < lolo);
+
+ uint64_t high = hihi + (cross_sum >> 32) + cross_carry + low_carry;
+
+ return high;
+}
+
+/* 32-bit variant. */
+uint32_t mulh_two_carry_32 (uint32_t x, uint32_t y)
+{
+ uint32_t x_hi = x >> 16;
+ uint32_t x_lo = x & 0xFFFF;
+ uint32_t y_hi = y >> 16;
+ uint32_t y_lo = y & 0xFFFF;
+
+ uint32_t lolo = x_lo * y_lo;
+ uint32_t hilo = x_hi * y_lo;
+ uint32_t lohi = x_lo * y_hi;
+ uint32_t hihi = x_hi * y_hi;
+
+ uint32_t cross_sum = hilo + lohi;
+ uint32_t cross_carry = (uint32_t)(cross_sum < hilo) << 16;
+
+ uint32_t cross_shifted = cross_sum << 16;
+ uint32_t low_result = lolo + cross_shifted;
+ uint32_t low_carry = (uint32_t)(low_result < cross_shifted);
+
+ uint32_t high = hihi + (cross_sum >> 16) + cross_carry + low_carry;
+
+ return high;
+}
+
+/* Full multiply: both high and low parts. */
+uint64_t full_mul_two_carry (uint64_t x, uint64_t y, uint64_t *lo)
+{
+ uint64_t x_hi = x >> 32;
+ uint64_t x_lo = x & 0xFFFFFFFFUL;
+ uint64_t y_hi = y >> 32;
+ uint64_t y_lo = y & 0xFFFFFFFFUL;
+
+ uint64_t lolo = x_lo * y_lo;
+ uint64_t hilo = x_hi * y_lo;
+ uint64_t lohi = x_lo * y_hi;
+ uint64_t hihi = x_hi * y_hi;
+
+ uint64_t cross_sum = hilo + lohi;
+ uint64_t cross_carry = (uint64_t)(cross_sum < hilo) << 32;
+
+ uint64_t cross_shifted = cross_sum << 32;
+ uint64_t low_result = lolo + cross_shifted;
+ uint64_t low_carry = (uint64_t)(low_result < cross_shifted);
+
+ uint64_t high = hihi + (cross_sum >> 32) + cross_carry + low_carry;
+
+ *lo = low_result;
+ return high;
+}
+
+/* Folds land in forwprop3 because the LT_EXPR low-carry compare is
+ only canonicalized into the gt:c-plus-overflow shape that
+ mul_carry_low matches by an earlier pass between forwprop1 and
+ forwprop3. */
+/* { dg-final { scan-tree-dump-times "Long multiplication high part folded." 4
"forwprop3" } } */
\ No newline at end of file
diff --git a/gcc/tree-ssa-forwprop.cc b/gcc/tree-ssa-forwprop.cc
index 288c9588b17a..ea989c809181 100644
--- a/gcc/tree-ssa-forwprop.cc
+++ b/gcc/tree-ssa-forwprop.cc
@@ -3604,6 +3604,7 @@ extern bool gimple_mul_low_sum (tree, tree *, tree
(*)(tree));
extern bool gimple_mul_low_accum (tree, tree *, tree (*)(tree));
extern bool gimple_mul_carry_cross_sum (tree, tree *, tree (*)(tree));
extern bool gimple_mul_carry_low_sum (tree, tree *, tree (*)(tree));
+extern bool gimple_mul_carry_low (tree, tree *, tree (*)(tree));
/* Replace STMT with a high-part multiply of OP1 and OP2, emitted as
(N)(((2N) op1 * (2N) op2) >> N).
@@ -3673,6 +3674,7 @@ enum long_mul_kind {
LMK_CROSS_SUM,
LMK_LOW_ACCUM,
LMK_LOW_SUM,
+ LMK_CARRY_LOW,
LMK_CARRY_CROSS_SUM,
LMK_CARRY_LOW_SUM,
};
@@ -3693,6 +3695,7 @@ struct long_mul_summand {
long_mul_extract extract;
tree op0, op1;
tree hilo0, hilo1, hilo2;
+ tree carry_a, carry_b;
unsigned HOST_WIDE_INT shift;
unsigned HOST_WIDE_INT mask;
};
@@ -3801,6 +3804,10 @@ long_mul_set_summand (long_mul_summand *info,
long_mul_kind kind,
n_hilos = 3;
shift_idx = 5;
break;
+ case LMK_CARRY_LOW:
+ info->carry_a = res_ops[0];
+ info->carry_b = res_ops[1];
+ return;
}
if (n_ops >= 1)
info->op0 = res_ops[0];
@@ -3826,8 +3833,9 @@ long_mul_classify_carry (tree leaf, long_mul_summand
*info)
tree res_ops[LONG_MUL_MAX_CAPTURES];
/* mul_carry_low_sum's inner is constrained to mul_low_sum (cross_sum
+ mul_hi(mul_lolo)); mul_carry_cross_sum's inner is just
- mul_cross_sum (any plus). Most specific first, so the less-
- constrained pattern doesn't shadow the more-constrained one. */
+ mul_cross_sum (any plus); mul_carry_low matches gt:c (@0, plus(@0,
+ @1)) without a baked-in shift. Most specific first, so the
+ less-constrained pattern doesn't shadow the more-constrained one. */
if (gimple_mul_carry_low_sum (leaf, res_ops, NULL))
{
long_mul_set_summand (info, LMK_CARRY_LOW_SUM, res_ops);
@@ -3838,6 +3846,11 @@ long_mul_classify_carry (tree leaf, long_mul_summand
*info)
long_mul_set_summand (info, LMK_CARRY_CROSS_SUM, res_ops);
return true;
}
+ if (gimple_mul_carry_low (leaf, res_ops, NULL))
+ {
+ long_mul_set_summand (info, LMK_CARRY_LOW, res_ops);
+ return true;
+ }
return false;
}
@@ -4071,6 +4084,18 @@ long_mul_canonical_ops (const vec<long_mul_summand>
&summands,
return false;
}
+/* Return the first summand in SUMMANDS whose kind matches KIND, or NULL. */
+
+static const long_mul_summand *
+long_mul_find_summand (const vec<long_mul_summand> &summands,
+ long_mul_kind kind)
+{
+ for (const long_mul_summand &s : summands)
+ if (s.kind == kind)
+ return &s;
+ return NULL;
+}
+
/* Run the cross-summand validation invariants and return the canonical
(op0, op1). Returns false unless all summands that carry operands use
the same (op0, op1) pair (in either order), every LMX_HI/LMX_SHL_N shift
@@ -4166,6 +4191,50 @@ long_mul_signature_matches (const vec<long_mul_summand>
&summands,
return true;
}
+/* Extra check for the two-carries high-part row: the LMK_CARRY_LOW summand's
+ two operands (carry_a, carry_b) must be a (cross_shifted, mul_lolo) pair
+ consistent with the multiset's canonical (op0, op1). */
+
+static bool
+long_mul_check_two_carries (const vec<long_mul_summand> &summands,
+ gimple *)
+{
+ tree op0, op1;
+ if (!long_mul_canonical_ops (summands, &op0, &op1))
+ return false;
+ unsigned int halfwidth = TYPE_PRECISION (TREE_TYPE (op0)) / 2;
+
+ const long_mul_summand *cl = long_mul_find_summand (summands, LMK_CARRY_LOW);
+ if (!cl)
+ return false;
+
+ /* The two carry_low operands must be (cross_shifted, mul_lolo) in either
+ order. cross_shifted = LSHIFT_EXPR (mul_cross_sum, halfwidth). */
+ tree cs = cl->carry_a, lolo = cl->carry_b;
+ tree inner;
+ unsigned HOST_WIDE_INT shift;
+ if (!long_mul_is_lshift_def (cs, &inner, &shift))
+ {
+ std::swap (cs, lolo);
+ if (!long_mul_is_lshift_def (cs, &inner, &shift))
+ return false;
+ }
+ if (shift != halfwidth)
+ return false;
+
+ tree scratch[LONG_MUL_MAX_CAPTURES];
+ if (!gimple_mul_cross_sum (inner, scratch, NULL))
+ return false;
+ for (int i = 0; i < 2; i++)
+ if (!long_mul_is_cross_half (scratch[i], op0, op1))
+ return false;
+ if (!gimple_mul_lolo (lolo, scratch, NULL)
+ || !long_mul_same_ops (scratch[0], scratch[1], op0, op1))
+ return false;
+
+ return true;
+}
+
/* Long-multiply variant table. Each row enumerates the multiset of
(kind, extract) summands that compose one long-multiply form. Rows
are sorted by long_mul_summand_compare, matching the input summands'
@@ -4195,6 +4264,14 @@ static const long_mul_row long_mul_table[] = {
{ LMK_LOW_ACCUM, LMX_HI },
{ LMK_CARRY_CROSS_SUM, LMX_NONE } },
NULL },
+ /* xh*yh + (cross_sum >> N) + carry_low + ((hilo > cross_sum) << N),
+ carry_low = (xl*yl + (cross_sum << N)) < (cross_sum << N). */
+ { long_mul_row::HIGH_PART, PLUS_EXPR, 4,
+ { { LMK_MUL_HIHI, LMX_NONE },
+ { LMK_CROSS_SUM, LMX_HI },
+ { LMK_CARRY_LOW, LMX_NONE },
+ { LMK_CARRY_CROSS_SUM, LMX_NONE } },
+ long_mul_check_two_carries },
/* LOW-PART folds. Recover the lower 2N bits from xl*yl plus a
shifted cross-half term. */
/* (xl*yl & mask) | (low_accum << N),
--
2.55.0