Extend the long-multiply fold to a second carry shape, where the
schoolbook expansion collapses the cross sum and the low partial
product (xl*yl >> N) into a single low_sum value and the overflow
compare tests against that:

  xh*yh + (low_sum >> N) + ((hilo > low_sum) << N)
  low_sum = cross_sum + (xl*yl >> N)
  cross_sum = xh*yl + xl*yh

The outer chain has three summands instead of four, and the
corresponding LOW_PART recovers (xl*yl & mask) | (low_sum << N).
match_long_mul picks up the new shapes via two table rows; the rest
of the framework (linearize, classify, consistency) is unchanged.

gcc/ChangeLog:

        * match.pd: Add mul_low_sum and mul_carry_low_sum atom
        recognizers.
        * tree-ssa-forwprop.cc (gimple_mul_low_sum): Declare.
        (gimple_mul_carry_low_sum): Likewise.
        (LONG_MUL_MAX_CAPTURES): Bump width-of-widest-atom note to
        mul_carry_low_sum (7 captures).
        (enum long_mul_kind): Add LMK_LOW_SUM and LMK_CARRY_LOW_SUM.
        (long_mul_set_summand): Handle the new kinds.
        (long_mul_classify_carry): Try mul_carry_low_sum before
        mul_carry_cross_sum (most-specific first).
        (long_mul_classify_plus_kinds): Try mul_low_sum after
        mul_low_accum (mul_low_sum's first arm is any plus and
        mul_low_accum constrains both arms).
        (long_mul_table): Add HIGH_PART and LOW_PART rows that
        consume LMK_LOW_SUM / LMK_CARRY_LOW_SUM.

gcc/testsuite/ChangeLog:

        * gcc.dg/tree-ssa/long-mul-carry.c: Add carry-low-sum
        coverage (mulh_carry_low_sum, full_mul_carry_low_sum,
        mulh_carry_low_sum_comm, mulh_carry_low_sum_lohi,
        mulh_carry_low_sum_128, mulh_carry_low_sum_v2i32) and
        update fold counts.

Signed-off-by: Konstantinos Eleftheriou <[email protected]>
---

(no changes since v1)

 gcc/match.pd                                  |  18 +++
 .../gcc.dg/tree-ssa/long-mul-carry.c          | 138 +++++++++++++++++-
 gcc/tree-ssa-forwprop.cc                      |  64 ++++++--
 3 files changed, 205 insertions(+), 15 deletions(-)

diff --git a/gcc/match.pd b/gcc/match.pd
index 5a4b1923b589..b2ddbec50181 100644
--- a/gcc/match.pd
+++ b/gcc/match.pd
@@ -12208,6 +12208,24 @@ and,
    products by the forwprop consumer (long_mul_check_consistency).  */
 (match (mul_cross_sum @mul_hilo0 @mul_hilo1)
   (plus:c @mul_hilo0 @mul_hilo1))
+/* Low sum: cross_sum + (xl*yl >> N).  */
+(match (mul_low_sum @op0 @op1 @mul_hilo0 @mul_hilo1 @0 @1)
+  (plus:c
+    (mul_cross_sum @mul_hilo0 @mul_hilo1)
+    (mul_hi
+      (mul_lolo @op0 @op1 INTEGER_CST@1)
+      INTEGER_CST@0)))
+/* Carry from low-sum overflow: (cast?) (hilo > low_sum) << N.
+   No explicit type/width guard needed: mul_low_sum delegates to
+   mul_cross_sum + mul_hi + mul_lolo, which provide deep structural
+   constraints, and @0 ties the shift amount to the inner constants.  */
+(match (mul_carry_low_sum @op0 @op1 @mul_hilo0 @mul_hilo1 @mul_hilo2 @0 @1)
+  (lshift
+    (convert? (gt
+      @mul_hilo0
+      (mul_low_sum @op0 @op1 @mul_hilo1 @mul_hilo2 INTEGER_CST@0
+       INTEGER_CST@1)))
+    INTEGER_CST@0))
 /* Carry from cross-sum overflow: (cast?) (hilo > cross_sum) << N.
    Explicit guard required because mul_cross_sum is just (plus:c @0 @1)
    with no inherent type or halfwidth constraint.  */
diff --git a/gcc/testsuite/gcc.dg/tree-ssa/long-mul-carry.c 
b/gcc/testsuite/gcc.dg/tree-ssa/long-mul-carry.c
index bbdbc63ac0bb..fd190f09dde9 100644
--- a/gcc/testsuite/gcc.dg/tree-ssa/long-mul-carry.c
+++ b/gcc/testsuite/gcc.dg/tree-ssa/long-mul-carry.c
@@ -176,6 +176,138 @@ v2i32 mulh_carry_v2i32 (v2i32 x, v2i32 y)
   return result;
 }
 
-/* { dg-final { scan-tree-dump-times "Long multiplication high part folded." 4 
"forwprop1" } } */
-/* { dg-final { scan-tree-dump-times "Long multiplication high part folded." 1 
"forwprop2" } } */
-/* { dg-final { scan-tree-dump-times "Long multiplication low part folded." 1 
"forwprop1" } } */
+/* High part collapses the cross sum and (xl*yl >> N) into a single
+   low_sum carrying the overflow compare:
+     xh*yh + (low_sum >> N) + ((hilo > low_sum) << N).  */
+
+uint32_t mulh_carry_low_sum (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 mulhilo = x_hi * y_lo;
+    uint32_t mullohi = x_lo * y_hi;
+    uint32_t cross_sum = mulhilo + mullohi;
+    uint32_t mullolo = x_lo * y_lo;
+    uint32_t shrlolo = mullolo >> 16;
+    uint32_t low_sum = cross_sum + shrlolo;
+    int carry = low_sum < mulhilo;
+    uint32_t cond = ((uint32_t) carry << 16) + x_hi * y_hi;
+    uint32_t add = cond + (low_sum >> 16);
+
+    return add;
+}
+
+void full_mul_carry_low_sum (uint32_t x, uint32_t y, uint32_t* p)
+{
+    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 mulhilo = x_hi * y_lo;
+    uint32_t mullohi = x_lo * y_hi;
+    uint32_t cross_sum = mulhilo + mullohi;
+    uint32_t mullolo = x_lo * y_lo;
+    uint32_t shrlolo = mullolo >> 16;
+    uint32_t low_sum = cross_sum + shrlolo;
+    int carry = low_sum < mulhilo;
+    uint32_t cond = ((uint32_t) carry << 16) + x_hi * y_hi;
+    uint32_t add = cond + (low_sum >> 16);
+    p[1] = add;
+    uint32_t low_sum_shr = low_sum << 16;
+    uint32_t mullololo = mullolo & 0xFFFF;
+    uint32_t low = low_sum_shr | mullololo;
+    p[0] = low;
+}
+
+uint32_t mulh_carry_low_sum_comm (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 mulhilo = y_lo * x_hi;
+    uint32_t mullohi = y_hi * x_lo;
+    uint32_t cross_sum = mullohi + mulhilo;
+    uint32_t mullolo = x_lo * y_lo;
+    uint32_t shrlolo = mullolo >> 16;
+    uint32_t low_sum = shrlolo + cross_sum;
+    int carry = low_sum < mulhilo;
+    uint32_t cond = ((uint32_t) carry << 16) + x_hi * y_hi;
+    uint32_t add = cond + (low_sum >> 16);
+
+    return add;
+}
+
+uint32_t mulh_carry_low_sum_lohi (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 mulhilo = x_hi * y_lo;
+    uint32_t mullohi = x_lo * y_hi;
+    uint32_t cross_sum = mulhilo + mullohi;
+    uint32_t mullolo = x_lo * y_lo;
+    uint32_t low_sum = cross_sum + (mullolo >> 16);
+    int carry_occurred = (low_sum < mullohi);
+    uint32_t cond = (uint32_t) carry_occurred << 16;
+    uint32_t add = x_hi * y_hi + cond + (low_sum >> 16);
+
+    return add;
+}
+
+/* The 128-bit variant will fail during the high sequence generation
+   (no target provides a 256-bit multiply) and is excluded from the
+   expected fold counts below.  */
+#ifdef __SIZEOF_INT128__
+__uint128_t mulh_carry_low_sum_128 (__uint128_t x, __uint128_t y)
+{
+    __uint128_t x_hi = x >> 64;
+    __uint128_t x_lo = x & (__uint128_t)0xFFFFFFFFFFFFFFFF;
+    __uint128_t y_hi = y >> 64;
+    __uint128_t y_lo = y & (__uint128_t)0xFFFFFFFFFFFFFFFF;
+    __uint128_t mulhilo = x_hi * y_lo;
+    __uint128_t mullohi = x_lo * y_hi;
+    __uint128_t cross_sum = mulhilo + mullohi;
+    __uint128_t mullolo = x_lo * y_lo;
+    __uint128_t shrlolo = mullolo >> 64;
+    __uint128_t low_sum = cross_sum + shrlolo;
+    int carry = low_sum < mulhilo;
+    __uint128_t cond = ((__uint128_t) carry << 64) + x_hi * y_hi;
+    __uint128_t add = cond + (low_sum >> 64);
+
+    return add;
+}
+#endif
+
+/* This will be optimized during the second forwprop run.
+   Disable SLP so the expected fold count is target-independent.  */
+__attribute__((optimize("no-tree-slp-vectorize")))
+v2i32 mulh_carry_low_sum_v2i32 (v2i32 x, v2i32 y)
+{
+  v2i32 result;
+  for (int i = 0; i < 2; i++)
+    {
+      uint32_t x_hi = x.v[i] >> 16;
+      uint32_t x_lo = x.v[i] & 0xFFFF;
+      uint32_t y_hi = y.v[i] >> 16;
+      uint32_t y_lo = y.v[i] & 0xFFFF;
+      uint32_t mulhilo  = x_hi * y_lo;
+      uint32_t mullohi = x_lo * y_hi;
+      uint32_t cross_sum  = mulhilo + mullohi;
+      uint32_t mullolo = x_lo * y_lo;
+      uint32_t shrlolo = mullolo >> 16;
+      uint32_t low_sum = cross_sum + shrlolo;
+      int carry = low_sum < mulhilo;
+      uint32_t cond = ((uint32_t) carry << 16) + x_hi * y_hi;
+      uint32_t add = cond + (low_sum >> 16);
+      result.v[i] = add;
+    }
+
+  return result;
+}
+
+/* { dg-final { scan-tree-dump-times "Long multiplication high part folded." 8 
"forwprop1" } } */
+/* { dg-final { scan-tree-dump-times "Long multiplication high part folded." 2 
"forwprop2" } } */
+/* { dg-final { scan-tree-dump-times "Long multiplication low part folded." 2 
"forwprop1" } } */
diff --git a/gcc/tree-ssa-forwprop.cc b/gcc/tree-ssa-forwprop.cc
index 3870bf17c0e0..288c9588b17a 100644
--- a/gcc/tree-ssa-forwprop.cc
+++ b/gcc/tree-ssa-forwprop.cc
@@ -3600,8 +3600,10 @@ extern bool gimple_mul_hilo (tree, tree *, tree 
(*)(tree));
 extern bool gimple_mul_lolo (tree, tree *, tree (*)(tree));
 extern bool gimple_mul_hihi (tree, tree *, tree (*)(tree));
 extern bool gimple_mul_cross_sum (tree, tree *, tree (*)(tree));
+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));
 
 /* Replace STMT with a high-part multiply of OP1 and OP2, emitted as
      (N)(((2N) op1 * (2N) op2) >> N).
@@ -3657,7 +3659,7 @@ create_mul_low_seq (tree op1, tree op2, gimple *stmt)
    up in a table.  On a hit, three cross-summand consistency checks
    decide whether the wide multiply is emitted.  */
 
-/* Widest match.pd atom (mul_low_accum) takes 6 captures; round up
+/* Widest match.pd atom (mul_carry_low_sum) takes 7 captures; round up
    to 8 for the scratch buffers below.  */
 static constexpr unsigned LONG_MUL_MAX_CAPTURES = 8;
 
@@ -3670,7 +3672,9 @@ enum long_mul_kind {
   LMK_MUL_HILO,
   LMK_CROSS_SUM,
   LMK_LOW_ACCUM,
+  LMK_LOW_SUM,
   LMK_CARRY_CROSS_SUM,
+  LMK_CARRY_LOW_SUM,
 };
 
 /* How the leaf wraps its inner kind.  Carry kinds use LMX_NONE: their
@@ -3784,6 +3788,7 @@ long_mul_set_summand (long_mul_summand *info, 
long_mul_kind kind,
       n_hilos = 2;
       break;
     case LMK_LOW_ACCUM:
+    case LMK_LOW_SUM:
       n_ops = 2;
       n_hilos = 2;
       break;
@@ -3791,6 +3796,11 @@ long_mul_set_summand (long_mul_summand *info, 
long_mul_kind kind,
       n_hilos = 3;
       shift_idx = 3;
       break;
+    case LMK_CARRY_LOW_SUM:
+      n_ops = 2;
+      n_hilos = 3;
+      shift_idx = 5;
+      break;
     }
   if (n_ops >= 1)
     info->op0 = res_ops[0];
@@ -3807,13 +3817,22 @@ long_mul_set_summand (long_mul_summand *info, 
long_mul_kind kind,
 }
 
 /* Classify LEAF as a carry-kind summand.  The lshift amount is baked
-   into mul_carry_cross_sum, so it's tried before any branch that looks
-   for a generic (X >> N) or (X << N) wrapper.  */
+   into mul_carry_cross_sum / mul_carry_low_sum, so they're tried before
+   any branch that looks for a generic (X >> N) or (X << N) wrapper.  */
 
 static bool
 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.  */
+  if (gimple_mul_carry_low_sum (leaf, res_ops, NULL))
+    {
+      long_mul_set_summand (info, LMK_CARRY_LOW_SUM, res_ops);
+      return true;
+    }
   if (gimple_mul_carry_cross_sum (leaf, res_ops, NULL))
     {
       long_mul_set_summand (info, LMK_CARRY_CROSS_SUM, res_ops);
@@ -3823,7 +3842,9 @@ long_mul_classify_carry (tree leaf, long_mul_summand 
*info)
 }
 
 /* Plus-based summand kinds shared by the (X >> SHIFT) and (X << SHIFT)
-   classifiers.  */
+   classifiers.  Order is by specificity: mul_low_sum's first arm is
+   any plus, so mul_low_accum (which constrains both arms) shadows it
+   and must come first.  */
 
 static bool
 long_mul_classify_plus_kinds (tree inner, long_mul_summand *info)
@@ -3834,6 +3855,11 @@ long_mul_classify_plus_kinds (tree inner, 
long_mul_summand *info)
       long_mul_set_summand (info, LMK_LOW_ACCUM, res_ops);
       return true;
     }
+  if (gimple_mul_low_sum (inner, res_ops, NULL))
+    {
+      long_mul_set_summand (info, LMK_LOW_SUM, res_ops);
+      return true;
+    }
   return false;
 }
 
@@ -4146,14 +4172,22 @@ long_mul_signature_matches (const vec<long_mul_summand> 
&summands,
    sort order, so a plain element-wise compare suffices.  Rows describe
    unsigned schoolbook expansions on an even-width 2N-bit type split at
    half-width N; EXTRA_CHECK carries invariants the (kind, extract)
-   signature cannot express.  */
+   signature cannot express.
+
+   The formula on each row uses xh, xl, yh, yl for the half-width pieces
+   of x and y, cross_sum for xh*yl + xl*yh, and hilo for either cross-half
+   product (consumers validate the operand shape).  */
 
 static const long_mul_row long_mul_table[] = {
-  /* HIGH-PART fold.  Notation: xh, xl, yh, yl are the half-width pieces
-     of x and y; N is the half-width.  cross_sum = xh*yl + xl*yh; hilo is
-     either xh*yl or xl*yh (consumers validate the operand shape).  */
-  /* xh*yh + (low_accum >> N) + (cross_sum >> N)
-     + ((hilo > cross_sum) << N),
+  /* HIGH-PART folds.  */
+  /* xh*yh + (low_sum >> N) + ((hilo > low_sum) << N),
+     low_sum = cross_sum + (xl*yl >> N).  */
+  { long_mul_row::HIGH_PART, PLUS_EXPR, 3,
+    { { LMK_MUL_HIHI, LMX_NONE },
+      { LMK_LOW_SUM, LMX_HI },
+      { LMK_CARRY_LOW_SUM, LMX_NONE } },
+    NULL },
+  /* xh*yh + (low_accum >> N) + (cross_sum >> N) + ((hilo > cross_sum) << N),
      low_accum = (xl*yl >> N) + (cross_sum & mask).  */
   { long_mul_row::HIGH_PART, PLUS_EXPR, 4,
     { { LMK_MUL_HIHI, LMX_NONE },
@@ -4161,14 +4195,20 @@ static const long_mul_row long_mul_table[] = {
       { LMK_LOW_ACCUM, LMX_HI },
       { LMK_CARRY_CROSS_SUM, LMX_NONE } },
     NULL },
-  /* LOW-PART fold.  Recover the lower 2N bits from xl*yl plus a
-     shifted cross-half term.  Notation as for the HIGH-PART row above.  */
+  /* LOW-PART folds.  Recover the lower 2N bits from xl*yl plus a
+     shifted cross-half term.  */
   /* (xl*yl & mask) | (low_accum << N),
      low_accum = (xl*yl >> N) + (cross_sum & mask).  */
   { long_mul_row::LOW_PART, BIT_IOR_EXPR, 2,
     { { LMK_MUL_LOLO, LMX_LO },
       { LMK_LOW_ACCUM, LMX_SHL_N } },
     NULL },
+  /* (xl*yl & mask) | (low_sum << N),
+     low_sum = cross_sum + (xl*yl >> N).  */
+  { long_mul_row::LOW_PART, BIT_IOR_EXPR, 2,
+    { { LMK_MUL_LOLO, LMX_LO },
+      { LMK_LOW_SUM, LMX_SHL_N } },
+    NULL },
 };
 
 /* If a multi-used inner addition (sharing the chain's outer code) blocked
-- 
2.55.0

Reply via email to