package unit; import java.util.ArrayList; import java.util.HashMap; import java.util.Iterator; import java.util.List; import java.util.Map; /** * MESA-0001: O(n²) ACO register allocation — update_renames parallel-copy linear scan. * * Models aco_register_allocation.cpp update_renames(): * Slow: std::find_if over parallelcopies for each entry (O(n) inner search). * Fast: unordered_map tempId→index built once, O(1) lookup per entry. * * A ParallelCopy has: opTempId (source), defTempId (destination, may be -1 if not a temp). */ public class MesaParallelCopyAlgorithm { static class ParallelCopy { int opTempId; // source temporary ID int defTempId; // destination temp ID (-1 = not a temp) ParallelCopy(int opTempId, int defTempId) { this.opTempId = opTempId; this.defTempId = defTempId; } boolean defIsTemp() { return defTempId >= 0; } @Override public String toString() { return "Copy(op=" + opTempId + ",def=" + (defIsTemp() ? defTempId : "X") + ")"; } } // --------------------------------------------------------------- // SLOW: std::find_if scan (original aco_register_allocation.cpp) // --------------------------------------------------------------- static long updateRenamesSlow(List copies) { long ops = 0; int idx = 0; while (idx < copies.size()) { ParallelCopy it = copies.get(idx); if (it.defIsTemp()) { idx++; continue; } // std::find_if — O(n) scan for matching def int otherIdx = -1; for (int j = 0; j < copies.size(); j++) { // inner O(n) scan ops++; ParallelCopy c = copies.get(j); if (c.defIsTemp() && it.opTempId == c.defTempId) { otherIdx = j; break; } } if (otherIdx >= 0) { // simulate: update op, then erase the other entry copies.remove(otherIdx); if (otherIdx < idx) idx--; // don't advance idx — re-check current position } else { idx++; } } return ops; } // --------------------------------------------------------------- // FAST: map-based O(1) lookup (patched version) // --------------------------------------------------------------- static long updateRenamesFast(List copies) { long ops = 0; // Build tempId → index map once Map defTempIdx = new HashMap<>(); for (int i = 0; i < copies.size(); i++) { if (copies.get(i).defIsTemp()) defTempIdx.put(copies.get(i).defTempId, i); } int idx = 0; while (idx < copies.size()) { ParallelCopy it = copies.get(idx); if (it.defIsTemp()) { idx++; continue; } ops++; // one O(1) map lookup Integer otherIdx = defTempIdx.get(it.opTempId); if (otherIdx != null && otherIdx < copies.size() && copies.get(otherIdx).defTempId == it.opTempId) { // Remove other entry, update map defTempIdx.remove(copies.get(otherIdx).defTempId); copies.remove((int) otherIdx); if (otherIdx < idx) idx--; // Rebuild shifted entries in map (entries after otherIdx shifted by -1) for (int i = otherIdx; i < copies.size(); i++) { if (copies.get(i).defIsTemp()) defTempIdx.put(copies.get(i).defTempId, i); } } else { idx++; } } return ops; } // --------------------------------------------------------------- // Test helpers // --------------------------------------------------------------- /** Build N parallel copies: ops 0..N-1, defs N..2N-1 (all temp). */ static List buildCopies(int n) { List list = new ArrayList<>(); // Half: non-temp defs that reference previous defs (create find_if work) for (int i = 0; i < n / 2; i++) { list.add(new ParallelCopy(n + i, -1)); // non-temp def, op = some temp } // Half: temp defs (sources for the find_if lookups) for (int i = 0; i < n / 2; i++) { // defTempId = n+i so the non-temp copies above can find them list.add(new ParallelCopy(i, n + i)); } return list; } static void testN(int n) { List slowCopies = buildCopies(n); List fastCopies = buildCopies(n); long slowOps = updateRenamesSlow(slowCopies); long fastOps = updateRenamesFast(fastCopies); int halfN = n / 2; // Slow: each non-temp entry scans half the list on average → O(n²/4) long slowMin = (long) halfN * halfN / 4; boolean slowBad = n < 8 || slowOps >= slowMin; boolean fastGood = fastOps <= (long) n + 2; boolean speedup = slowOps >= fastOps; assert slowBad : "slow not O(n²): ops=" + slowOps + " min=" + slowMin; assert fastGood : "fast not O(n): ops=" + fastOps + " n=" + n; assert speedup : "fast not faster: slow=" + slowOps + " fast=" + fastOps; System.out.printf(" %-30s N=%-4d slow=%6d fast=%4d speedup=%.0fx%n", "parallelCopyRename N=" + n, n, slowOps, fastOps, (double) slowOps / Math.max(fastOps, 1)); } public static void main(String[] args) { int passed = 0; int total = 0; int[] sizes = {8, 16, 32, 64, 128, 256}; for (int n : sizes) { total++; testN(n); passed++; } // Verify semantics: both paths should produce identical final lists total++; { int n = 20; List slowList = buildCopies(n); List fastList = buildCopies(n); updateRenamesSlow(slowList); updateRenamesFast(fastList); assert slowList.size() == fastList.size() : "final list sizes differ: slow=" + slowList.size() + " fast=" + fastList.size(); System.out.printf(" %-30s final sizes match: %d%n", "semantics check", slowList.size()); passed++; } System.out.printf("%n%d/%d PASS%n", passed, total); } }