Skip to content

Commit 6644ce3

Browse files
committed
[SYSTEMDS-3333] Fix new mmchain-opt rewrite code style and test
* Fixes the remaining invalid method names of the new and existing mmchain-opt rewrites * Fixes the handling of stdout streams with default output buffering (which is on in the github actions)
1 parent b748091 commit 6644ce3

4 files changed

Lines changed: 40 additions & 55 deletions

File tree

src/main/java/org/apache/sysds/hops/rewrite/RewriteMatrixMultChainOptimization.java

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -50,7 +50,7 @@ public ArrayList<Hop> rewriteHopDAGs(ArrayList<Hop> roots, ProgramRewriteStatus
5050

5151
// Find the optimal order for the chain whose result is the current HOP
5252
for( Hop h : roots )
53-
rule_OptimizeMMChains(h, state);
53+
ruleOptimizeMMChains(h, state);
5454

5555
return roots;
5656
}
@@ -62,7 +62,7 @@ public Hop rewriteHopDAG(Hop root, ProgramRewriteStatus state)
6262
return null;
6363

6464
// Find the optimal order for the chain whose result is the current HOP
65-
rule_OptimizeMMChains(root, state);
65+
ruleOptimizeMMChains(root, state);
6666

6767
return root;
6868
}
@@ -73,7 +73,7 @@ public Hop rewriteHopDAG(Hop root, ProgramRewriteStatus state)
7373
*
7474
* @param hop high-level operator
7575
*/
76-
private void rule_OptimizeMMChains(Hop hop, ProgramRewriteStatus state)
76+
private void ruleOptimizeMMChains(Hop hop, ProgramRewriteStatus state)
7777
{
7878
if( hop.isVisited() )
7979
return;
@@ -87,7 +87,7 @@ private void rule_OptimizeMMChains(Hop hop, ProgramRewriteStatus state)
8787
}
8888

8989
for( Hop hi : hop.getInput() )
90-
rule_OptimizeMMChains(hi, state);
90+
ruleOptimizeMMChains(hi, state);
9191

9292
hop.setVisited();
9393
}

src/main/java/org/apache/sysds/hops/rewrite/RewriteMatrixMultChainOptimizationTranspose.java

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -51,7 +51,7 @@ public ArrayList<Hop> rewriteHopDAGs(ArrayList<Hop> roots, ProgramRewriteStatus
5151

5252
// Find the optimal order for the chain whose result is the current HOP
5353
for( Hop h : roots )
54-
rule_OptimizeMMChains(h, state);
54+
ruleOptimizeMMChains(h, state);
5555

5656
return roots;
5757
}
@@ -63,7 +63,7 @@ public Hop rewriteHopDAG(Hop root, ProgramRewriteStatus state)
6363
return null;
6464

6565
// Find the optimal order for the chain whose result is the current HOP
66-
rule_OptimizeMMChains(root, state);
66+
ruleOptimizeMMChains(root, state);
6767

6868
return root;
6969
}
@@ -74,7 +74,7 @@ public Hop rewriteHopDAG(Hop root, ProgramRewriteStatus state)
7474
*
7575
* @param hop high-level operator
7676
*/
77-
private void rule_OptimizeMMChains(Hop hop, ProgramRewriteStatus state)
77+
private void ruleOptimizeMMChains(Hop hop, ProgramRewriteStatus state)
7878
{
7979
if( !hop.isVisited() ) {
8080

@@ -85,7 +85,7 @@ private void rule_OptimizeMMChains(Hop hop, ProgramRewriteStatus state)
8585
}
8686

8787
for (Hop hi : hop.getInput())
88-
rule_OptimizeMMChains(hi, state);
88+
ruleOptimizeMMChains(hi, state);
8989

9090
hop.setVisited();
9191
}

src/main/java/org/apache/sysds/hops/rewrite/RewriteMatrixMultChainWithTransOptimization.java

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -45,7 +45,7 @@ public ArrayList<Hop> rewriteHopDAGs(ArrayList<Hop> roots, ProgramRewriteStatus
4545

4646
// Find the optimal order for the chain whose result is the current HOP
4747
for( Hop h : roots )
48-
rule_OptimizeMMChains(h, state);
48+
ruleOptimizeMMChains(h, state);
4949

5050
return roots;
5151
}
@@ -57,7 +57,7 @@ public Hop rewriteHopDAG(Hop root, ProgramRewriteStatus state)
5757
return null;
5858

5959
// Find the optimal order for the chain whose result is the current HOP
60-
rule_OptimizeMMChains(root, state);
60+
ruleOptimizeMMChains(root, state);
6161

6262
return root;
6363
}
@@ -69,7 +69,7 @@ public Hop rewriteHopDAG(Hop root, ProgramRewriteStatus state)
6969
* @param hop The current high-level operator node.
7070
* @param state The rewrite status.
7171
*/
72-
private void rule_OptimizeMMChains(Hop hop, ProgramRewriteStatus state) {
72+
private void ruleOptimizeMMChains(Hop hop, ProgramRewriteStatus state) {
7373
if (hop.isVisited()) return;
7474

7575
boolean isMatrixMult = HopRewriteUtils.isMatrixMultiply(hop) && !((AggBinaryOp) hop).hasLeftPMInput();
@@ -91,7 +91,7 @@ private void rule_OptimizeMMChains(Hop hop, ProgramRewriteStatus state) {
9191
// .toArray(new Hop[0]) this prevents ConcurrentModificationException because the optimizer
9292
// may replace or modify parts of the HOP DAG during recursion
9393
for( Hop i : currentHop.getInput().toArray(new Hop[0]) ) {
94-
rule_OptimizeMMChains(i, state);
94+
ruleOptimizeMMChains(i, state);
9595
}
9696
}
9797

src/test/java/org/apache/sysds/test/functions/rewrite/RewriteMatrixChainDPTest.java

Lines changed: 28 additions & 43 deletions
Original file line numberDiff line numberDiff line change
@@ -28,9 +28,6 @@
2828
import org.apache.sysds.test.TestConfiguration;
2929
import org.apache.sysds.test.TestUtils;
3030

31-
import java.io.ByteArrayOutputStream;
32-
import java.io.PrintStream;
33-
3431
public class RewriteMatrixChainDPTest extends AutomatedTestBase {
3532

3633
private static final String TEST_DIR = "functions/rewrite/mmchain/";
@@ -52,76 +49,76 @@ public void setUp() {
5249
}
5350

5451
@Test
55-
public void testMatrixChainDP_Test1() { runTestMatrixChainDP(TEST_CASES[0]); }
52+
public void testMatrixChainDPTest1() { runTestMatrixChainDP(TEST_CASES[0]); }
5653

5754
@Test
58-
public void testMatrixChainDP_Test2() { runTestMatrixChainDP(TEST_CASES[1]); }
55+
public void testMatrixChainDPTest2() { runTestMatrixChainDP(TEST_CASES[1]); }
5956

6057
@Test
61-
public void testMatrixChainDP_Test3() { runTestMatrixChainDP(TEST_CASES[2]); }
58+
public void testMatrixChainDPTest3() { runTestMatrixChainDP(TEST_CASES[2]); }
6259

6360
@Test
64-
public void testMatrixChainDP_Test4() { runTestMatrixChainDP(TEST_CASES[3]); }
61+
public void testMatrixChainDPTest4() { runTestMatrixChainDP(TEST_CASES[3]); }
6562

6663
@Test
67-
public void testMatrixChainDP_Test5() { runTestMatrixChainDP(TEST_CASES[4]); }
64+
public void testMatrixChainDPTest5() { runTestMatrixChainDP(TEST_CASES[4]); }
6865

6966
@Test
70-
public void testMatrixChainDP_Test6() { runTestMatrixChainDP(TEST_CASES[5]); }
67+
public void testMatrixChainDPTest6() { runTestMatrixChainDP(TEST_CASES[5]); }
7168

7269
@Test
73-
public void testMatrixChainDP_Test7() { runTestMatrixChainDP(TEST_CASES[6]); }
70+
public void testMatrixChainDPTest7() { runTestMatrixChainDP(TEST_CASES[6]); }
7471

7572
@Test
76-
public void testMatrixChainDP_Test8() { runTestMatrixChainDP(TEST_CASES[7]); }
73+
public void testMatrixChainDPTest8() { runTestMatrixChainDP(TEST_CASES[7]); }
7774

7875
@Test
79-
public void testMatrixChainDP_Test9() { runTestMatrixChainDP(TEST_CASES[8]); }
76+
public void testMatrixChainDPTest9() { runTestMatrixChainDP(TEST_CASES[8]); }
8077

8178
@Test
82-
public void testMatrixChainDP_Test10() { runTestMatrixChainDP(TEST_CASES[9]); }
79+
public void testMatrixChainDPTest10() { runTestMatrixChainDP(TEST_CASES[9]); }
8380

8481
@Test
85-
public void testMatrixChainDP_Test11() { runTestMatrixChainDP(TEST_CASES[10]); }
82+
public void testMatrixChainDPTest11() { runTestMatrixChainDP(TEST_CASES[10]); }
8683

8784
@Test
88-
public void testMatrixChainDP_Test12() { runTestMatrixChainDP(TEST_CASES[11]); }
85+
public void testMatrixChainDPTest12() { runTestMatrixChainDP(TEST_CASES[11]); }
8986

9087
@Test
91-
public void testMatrixChainDP_Test13() { runTestMatrixChainDP(TEST_CASES[12]); }
88+
public void testMatrixChainDPTest13() { runTestMatrixChainDP(TEST_CASES[12]); }
9289

9390
@Test
94-
public void testMatrixChainDP_Test14() { runTestMatrixChainDP(TEST_CASES[13]); }
91+
public void testMatrixChainDPTest14() { runTestMatrixChainDP(TEST_CASES[13]); }
9592

9693
@Test
97-
public void testMatrixChainDP_Test15() { runTestMatrixChainDP(TEST_CASES[14]); }
94+
public void testMatrixChainDPTest15() { runTestMatrixChainDP(TEST_CASES[14]); }
9895

9996
@Test
100-
public void testMatrixChainDP_Test16() { runTestMatrixChainDP(TEST_CASES[15]); }
97+
public void testMatrixChainDPTest16() { runTestMatrixChainDP(TEST_CASES[15]); }
10198

10299
@Test
103-
public void testMatrixChainDP_Test17() { runTestMatrixChainDP(TEST_CASES[16]); }
100+
public void testMatrixChainDPTest17() { runTestMatrixChainDP(TEST_CASES[16]); }
104101

105102
@Test
106-
public void testMatrixChainDP_Test18() { runTestMatrixChainDP(TEST_CASES[17]); }
103+
public void testMatrixChainDPTest18() { runTestMatrixChainDP(TEST_CASES[17]); }
107104

108105
@Test
109-
public void testMatrixChainDP_Test19() { runTestMatrixChainDP(TEST_CASES[18]); }
106+
public void testMatrixChainDPTest19() { runTestMatrixChainDP(TEST_CASES[18]); }
110107

111108
@Test
112-
public void testMatrixChainDP_Test20() { runTestMatrixChainDP(TEST_CASES[19]); }
109+
public void testMatrixChainDPTest20() { runTestMatrixChainDP(TEST_CASES[19]); }
113110

114111
@Test
115-
public void testMatrixChainDP_Test21() { runTestMatrixChainDP(TEST_CASES[20]); }
112+
public void testMatrixChainDPTest21() { runTestMatrixChainDP(TEST_CASES[20]); }
116113

117114
@Test
118-
public void testMatrixChainDP_Test22() { runTestMatrixChainDP(TEST_CASES[21]); }
115+
public void testMatrixChainDPTest22() { runTestMatrixChainDP(TEST_CASES[21]); }
119116

120117
@Test
121-
public void testMatrixChainDP_Test23() { runTestMatrixChainDP(TEST_CASES[22]); }
118+
public void testMatrixChainDPTest23() { runTestMatrixChainDP(TEST_CASES[22]); }
122119

123120
@Test
124-
public void testMatrixChainDP_Test24() {runTestMatrixChainDP(TEST_CASES[23]);}
121+
public void testMatrixChainDPTest24() {runTestMatrixChainDP(TEST_CASES[23]);}
125122

126123

127124
private void runTestMatrixChainDP(String testName) {
@@ -144,23 +141,11 @@ private void runTestMatrixChainDP(String testName) {
144141

145142
programArgs = new String[]{ "-explain", "hops", "-stats", "-args", output("R") };
146143

147-
// print HOP DAG
148-
PrintStream originalOut = System.out;
149-
ByteArrayOutputStream bos = new ByteArrayOutputStream();
150-
System.setOut(new PrintStream(bos));
151-
152-
try {
153-
// Execute the DML script
154-
runTest(true, false, null, -1);
155-
} finally {
156-
System.setOut(originalOut);
157-
}
158-
159-
String output = bos.toString();
160-
161-
System.out.println("Output for " + testName + ":\n" + output);
144+
// Execute the DML script
145+
setOutputBuffering(true);
146+
String output = runTest(true, false, null, -1).toString();
162147

163-
/* the following uses the intermediate matrices dimensions to check, wether
148+
/* the following uses the intermediate matrices dimensions to check, whether
164149
* the rewrite rule has found the optimal plan, which is commented in each script
165150
*/
166151
switch(testName) {

0 commit comments

Comments
 (0)