EnzymeAD / EnzymeAD/Enzyme-JAX

Reverse removal

Open
#1,237 0 comments 0 reactions 0 assignees View on GitHub
Dominant language
MLIR
Stars
131
Forks
53
Avg merge
1d 10h
Merged PRs (30d)
193

Description

given reverse(f(x,y)), if x and y are both either reversuble (e.g. splatted constants, broadcast), or themselves generated by reverses, we should change to f(reverse(x), reverse(y))

```

%1532 = "enzymexla.wrap"(%1135) <{dimension = 2 : i64, lhs = 1 : i64, rhs = 1 : i64}> : (tensor<3x1522x3056xf64>) -> tensor<3x1522x3058xf64> loc(#loc3976)
%1533 = stablehlo.reverse %1532, dims = [0] : tensor<3x1522x3058xf64> loc(#loc)
%1534 = "enzymexla.wrap"(%1142) <{dimension = 2 : i64, lhs = 1 : i64, rhs = 1 : i64}> : (tensor<3x1522x3056xf64>) -> tensor<3x1522x3058xf64> loc(#loc3976)
%1535 = stablehlo.reverse %1534, dims = [0] : tensor<3x1522x3058xf64> loc(#loc)
%1536 = stablehlo.reverse %14, dims = [0] : tensor<3xf64> loc(#loc)
%1537 = stablehlo.multiply %1533, %cst_236 : tensor<3x1522x3058xf64> loc(#loc2470)
%1538 = stablehlo.add %1535, %cst_235 : tensor<3x1522x3058xf64> loc(#loc2471)
%1539 = stablehlo.multiply %1538, %cst_234 : tensor<3x1522x3058xf64> loc(#loc2472)
%1540 = stablehlo.sqrt %1539 : tensor<3x1522x3058xf64> loc(#loc2473)
%1541 = stablehlo.negate %1536 : tensor<3xf64> loc(#loc2474)
%1542 = stablehlo.multiply %1541, %cst_233 : tensor<3xf64> loc(#loc2475)
%1543 = stablehlo.multiply %1533, %cst_232 : tensor<3x1522x3058xf64> loc(#loc2995)
%1544 = stablehlo.multiply %1540, %cst_231 : tensor<3x1522x3058xf64> loc(#loc2707)
%1545 = stablehlo.subtract %1543, %1544 : tensor<3x1522x3058xf64> loc(#loc2996)
%1546 = stablehlo.add %1545, %cst_230 : tensor<3x1522x3058xf64> loc(#loc2996)
%1547 = stablehlo.broadcast_in_dim %1542, dims = [0] : (tensor<3xf64>) -> tensor<3x1522x3058xf64> loc(#loc2475)
%1548 = stablehlo.multiply %1546, %1547 : tensor<3x1522x3058xf64> loc(#loc2477)
%1549 = stablehlo.multiply %1533, %cst_229 : tensor<3x1522x3058xf64> loc(#loc2997)
%1550 = stablehlo.multiply %1540, %cst_228 : tensor<3x1522x3058xf64> loc(#loc2709)
%1551 = stablehlo.subtract %1549, %1550 : tensor<3x1522x3058xf64> loc(#loc2998)
%1552 = stablehlo.add %1551, %cst_227 : tensor<3x1522x3058xf64> loc(#loc2998)
%1553 = stablehlo.multiply %1537, %1552 : tensor<3x1522x3058xf64> loc(#loc2709)
%1554 = stablehlo.multiply %1540, %cst_226 : tensor<3x1522x3058xf64> loc(#loc2709)
%1555 = stablehlo.add %1554, %cst_225 : tensor<3x1522x3058xf64> loc(#loc2711)
%1556 = stablehlo.multiply %1540, %1555 : tensor<3x1522x3058xf64> loc(#loc2709)
%1557 = stablehlo.add %1556, %1553 : tensor<3x1522x3058xf64> loc(#loc2998)
%1558 = stablehlo.add %1557, %cst_224 : tensor<3x1522x3058xf64> loc(#loc2998)
%1559 = stablehlo.add %1548, %1558 : tensor<3x1522x3058xf64> loc(#loc2479)
%1560 = stablehlo.multiply %1547, %1559 : tensor<3x1522x3058xf64> loc(#loc2477)
%1561 = stablehlo.multiply %1533, %cst_223 : tensor<3x1522x3058xf64> loc(#loc2999)
%1562 = stablehlo.multiply %1540, %cst_222 : tensor<3x1522x3058xf64> loc(#loc2712)
%1563 = stablehlo.subtract %1561, %1562 : tensor<3x1522x3058xf64> loc(#loc3000)
%1564 = stablehlo.add %1563, %cst_221 : tensor<3x1522x3058xf64> loc(#loc3000)
%1565 = stablehlo.multiply %1537, %1564 : tensor<3x1522x3058xf64> loc(#loc2712)
%1566 = stablehlo.multiply %1540, %cst_220 : tensor<3x1522x3058xf64> loc(#loc2712)
%1567 = stablehlo.subtract %cst_219, %1566 : tensor<3x1522x3058xf64> loc(#loc2714)
%1568 = stablehlo.multiply %1540, %1567 : tensor<3x1522x3058xf64> loc(#loc2712)
%1569 = stablehlo.add %1568, %1565 : tensor<3x1522x3058xf64> loc(#loc3000)
%1570 = stablehlo.add %1569, %cst_218 : tensor<3x1522x3058xf64> loc(#loc3000)
%1571 = stablehlo.multiply %1537, %1570 : tensor<3x1522x3058xf64> loc(#loc2712)
%1572 = stablehlo.multiply %1540, %cst_217 : tensor<3x1522x3058xf64> loc(#loc2712)
%1573 = stablehlo.subtract %cst_216, %1572 : tensor<3x1522x3058xf64> loc(#loc2714)
%1574 = stablehlo.multiply %1540, %1573 : tensor<3x1522x3058xf64> loc(#loc2712)
%1575 = stablehlo.add %1574, %cst_215 : tensor<3x1522x3058xf64> loc(#loc2714)
%1576 = stablehlo.multiply %1540, %1575 : tensor<3x1522x3058xf64> loc(#loc2712)
%1577 = stablehlo.add %1576, %1571 : tensor<3x1522x3058xf64> loc(#loc3000)
%1578 = stablehlo.add %1577, %cst_214 : tensor<3x1522x3058xf64> loc(#loc3000)
%1579 = stablehlo.multiply %1537, %1578 : tensor<3x1522x3058xf64> loc(#loc2712)
%1580 = stablehlo.multiply %1540, %cst_213 : tensor<3x1522x3058xf64> loc(#loc2712)
%1581 = stablehlo.add %1580, %cst_212 : tensor<3x1522x3058xf64> loc(#loc2714)
%1582 = stablehlo.multiply %1540, %1581 : tensor<3x1522x3058xf64> loc(#loc2712)
%1583 = stablehlo.add %1582, %cst_211 : tensor<3x1522x3058xf64> loc(#loc2714)
%1584 = stablehlo.multiply %1540, %1583 : tensor<3x1522x3058xf64> loc(#loc2712)
%1585 = stablehlo.add %1584, %cst_210 : tensor<3x1522x3058xf64> loc(#loc2714)
%1586 = stablehlo.multiply %1540, %1585 : tensor<3x1522x3058xf64> loc(#loc2712)
%1587 = stablehlo.add %1586, %1579 : tensor<3x1522x3058xf64> loc(#loc3000)
%1588 = stablehlo.add %1587, %cst_209 : tensor<3x1522x3058xf64> loc(#loc3000)
%1589 = stablehlo.add %1560, %1588 : tensor<3x1522x3058xf64> loc(#loc2479)
%1590 = stablehlo.multiply %1547, %1589 : tensor<3x1522x3058xf64> loc(#loc2477)
%1591 = stablehlo.multiply %1533, %cst_208 : tensor<3x1522x3058xf64> loc(#loc3001)
%1592 = stablehlo.multiply %1540, %cst_207 : tensor<3x1522x3058xf64> loc(#loc2715)
%1593 = stablehlo.subtract %1592, %1591 : tensor<3x1522x3058xf64> loc(#loc3002)
%1594 = stablehlo.add %1593, %cst_206 : tensor<3x1522x3058xf64> loc(#loc3002)
%1595 = stablehlo.multiply %1537, %1594 : tensor<3x1522x3058xf64> loc(#loc2715)
%1596 = stablehlo.multiply %1540, %cst_205 : tensor<3x1522x3058xf64> loc(#loc2715)
%1597 = stablehlo.subtract %cst_204, %1596 : tensor<3x1522x3058xf64> loc(#loc2717)
%1598 = stablehlo.multiply %1540, %1597 : tensor<3x1522x3058xf64> loc(#loc2715)
%1599 = stablehlo.add %1598, %1595 : tensor<3x1522x3058xf64> loc(#loc3002)
%1600 = stablehlo.add %1599, %cst_203 : tensor<3x1522x3058xf64> loc(#loc3002)
%1601 = stablehlo.multiply %1537, %1600 : tensor<3x1522x3058xf64> loc(#loc2715)
%1602 = stablehlo.multiply %1540, %cst_202 : tensor<3x1522x3058xf64> loc(#loc2715)
%1603 = stablehlo.subtract %cst_201, %1602 : tensor<3x1522x3058xf64> loc(#loc2717)
%1604 = stablehlo.multiply %1540, %1603 : tensor<3x1522x3058xf64> loc(#loc2715)
%1605 = stablehlo.add %1604, %cst_200 : tensor<3x1522x3058xf64> loc(#loc2717)
%1606 = stablehlo.multiply %1540, %1605 : tensor<3x1522x3058xf64> loc(#loc2715)
%1607 = stablehlo.add %1606, %1601 : tensor<3x1522x3058xf64> loc(#loc3002)
%1608 = stablehlo.add %1607, %cst_199 : tensor<3x1522x3058xf64> loc(#loc3002)
%1609 = stablehlo.multiply %1537, %1608 : tensor<3x1522x3058xf64> loc(#loc2715)
%1610 = stablehlo.multiply %1540, %cst_198 : tensor<3x1522x3058xf64> loc(#loc2715)
%1611 = stablehlo.subtract %cst_197, %1610 : tensor<3x1522x3058xf64> loc(#loc2717)
%1612 = stablehlo.multiply %1540, %1611 : tensor<3x1522x3058xf64> loc(#loc2715)
%1613 = stablehlo.add %1612, %cst_196 : tensor<3x1522x3058xf64> loc(#loc2717)
%1614 = stablehlo.multiply %1540, %1613 : tensor<3x1522x3058xf64> loc(#loc2715)
%1615 = stablehlo.add %1614, %cst_195 : tensor<3x1522x3058xf64> loc(#loc2717)
%1616 = stablehlo.multiply %1540, %1615 : tensor<3x1522x3058xf64> loc(#loc2715)
%1617 = stablehlo.add %1616, %1609 : tensor<3x1522x3058xf64> loc(#loc3002)
%1618 = stablehlo.add %1617, %cst_194 : tensor<3x1522x3058xf64> loc(#loc3002)
%1619 = stablehlo.multiply %1537, %1618 : tensor<3x1522x3058xf64> loc(#loc2715)
%1620 = stablehlo.multiply %1540, %cst_193 : tensor<3x1522x3058xf64> loc(#loc2715)
%1621 = stablehlo.subtract %cst_192, %1620 : tensor<3x1522x3058xf64> loc(#loc2717)
%1622 = stablehlo.multiply %1540, %1621 : tensor<3x1522x3058xf64> loc(#loc2715)
%1623 = stablehlo.add %1622, %cst_191 : tensor<3x1522x3058xf64> loc(#loc2717)
%1624 = stablehlo.multiply %1540, %1623 : tensor<3x1522x3058xf64> loc(#loc2715)
%1625 = stablehlo.add %1624, %cst_190 : tensor<3x1522x3058xf64> loc(#loc2717)
%1626 = stablehlo.multiply %1540, %1625 : tensor<3x1522x3058xf64> loc(#loc2715)
%1627 = stablehlo.add %1626, %cst_189 : tensor<3x1522x3058xf64> loc(#loc2717)
%1628 = stablehlo.multiply %1540, %1627 : tensor<3x1522x3058xf64> loc(#loc2715)
%1629 = stablehlo.add %1628, %1619 : tensor<3x1522x3058xf64> loc(#loc3002)
%1630 = stablehlo.add %1629, %cst_188 : tensor<3x1522x3058xf64> loc(#loc3002)
%1631 = stablehlo.multiply %1537, %1630 : tensor<3x1522x3058xf64> loc(#loc2715)
%1632 = stablehlo.multiply %1540, %cst_187 : tensor<3x1522x3058xf64> loc(#loc2715)
%1633 = stablehlo.subtract %cst_186, %1632 : tensor<3x1522x3058xf64> loc(#loc2717)
%1634 = stablehlo.multiply %1540, %1633 : tensor<3x1522x3058xf64> loc(#loc2715)
%1635 = stablehlo.add %1634, %cst_185 : tensor<3x1522x3058xf64> loc(#loc2717)
%1636 = stablehlo.multiply %1540, %1635 : tensor<3x1522x3058xf64> loc(#loc2715)
%1637 = stablehlo.add %1636, %cst_184 : tensor<3x1522x3058xf64> loc(#loc2717)
%1638 = stablehlo.multiply %1540, %1637 : tensor<3x1522x3058xf64> loc(#loc2715)
%1639 = stablehlo.add %1638, %cst_183 : tensor<3x1522x3058xf64> loc(#loc2717)
%1640 = stablehlo.multiply %1540, %1639 : tensor<3x1522x3058xf64> loc(#loc2715)
%1641 = stablehlo.add %1640, %cst_182 : tensor<3x1522x3058xf64> loc(#loc2717)
%1642 = stablehlo.multiply %1540, %1641 : tensor<3x1522x3058xf64> loc(#loc2715)
%1643 = stablehlo.add %1642, %1631 : tensor<3x1522x3058xf64> loc(#loc3002)
%1644 = stablehlo.add %1643, %cst_181 : tensor<3x1522x3058xf64> loc(#loc3002)
%1645 = stablehlo.add %1590, %1644 : tensor<3x1522x3058xf64> loc(#loc2479)
%1646 = stablehlo.subtract %1645, %cst_180 : tensor<3x1522x3058xf64> loc(#loc2297)

....

%1647 = stablehlo.multiply %1646, %cst_178 : tensor<3x1522x3058xf64> loc(#loc2153)

%1648 = stablehlo.subtract %1531, %1647 : tensor<3x1522x3058xf64> loc(#loc1883)

%1649 = stablehlo.multiply %1648, %cst_177 : tensor<3x1522x3058xf64> loc(#loc2021)

// reversed
%1650 = stablehlo.reverse %13, dims = [0] : tensor<3xf64> loc(#loc)
%1651 = stablehlo.broadcast_in_dim %1650, dims = [0] : (tensor<3xf64>) -> tensor<3x1522x3058xf64> loc(#loc)

%1652 = stablehlo.multiply %1651, %1649 : tensor<3x1522x3058xf64> loc(#loc1647)

// reversible
%1653 = stablehlo.broadcast_in_dim %1414, dims = [1, 2] : (tensor<1522x3058xf64>) -> tensor<3x1522x3058xf64> loc(#loc3003)

%cst_368 = stablehlo.constant dense<0.000000e+00> : tensor loc(#loc)

%1654 = "stablehlo.reduce_window"(%1652, %cst_368) <{base_dilations = array, padding = dense<[[2, 0], [0, 0], [0, 0]]> : tensor<3x2xi64>, window_dilations = array, window_dimensions = array, window_strides = array}> ({
^bb0(%arg39: tensor loc(callsite(#loc1013 at #loc1500)), %arg40: tensor loc(callsite(#loc1013 at #loc1500))):
%5673 = stablehlo.add %arg39, %arg40 : tensor loc(#loc1649)
stablehlo.return %5673 : tensor loc(#loc1649)
}) : (tensor<3x1522x3058xf64>, tensor) -> tensor<3x1522x3058xf64> loc(#loc1649)

%1655 = stablehlo.subtract %1653, %1654 : tensor<3x1522x3058xf64> loc(#loc1649)
%1656 = stablehlo.reverse %1655, dims = [0] : tensor<3x1522x3058xf64> loc(#loc2943)
```

Contributor guide

No contributing guide indexed for this repository

Assessment

This issue has not been assessed yet.

Get new issues in your inbox

A short digest of beginner-friendly GitHub issues.