java-topology/defects/jax/unit/JaxPallasFusionAlgorithm.java

155 lines
6.3 KiB
Java
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

package unit;
import java.util.*;
/**
* Unit test for jax-0001:
* pallas/fuser/jaxpr_fusion.py line 216 — `v in flat_group_vars` list membership
* inside loop over jaxpr_out.invars, repeated for each fusion group.
* O(G × I × F) → O(G × (I + F)) with HashSet per group.
*
* Compile: javac -d . JaxPallasFusionAlgorithm.java
* Run: java -ea unit.JaxPallasFusionAlgorithm
*/
public class JaxPallasFusionAlgorithm {
static void check(String desc, boolean cond) {
if (!cond) throw new AssertionError("FAIL: " + desc);
System.out.println("PASS: " + desc);
}
// -----------------------------------------------------------------------
// Simulated Jaxpr Var — identity by object reference
// -----------------------------------------------------------------------
static class Var {
final int id;
Var(int id) { this.id = id; }
@Override public String toString() { return "%" + id; }
}
// -----------------------------------------------------------------------
// SLOW: v in flat_group_vars (list) for each invar for each group
// -----------------------------------------------------------------------
static int slowComputeUsedJaxprInvars(
List<List<Var>> groups, // G groups, each F vars
List<Var> jaxprInvars) { // I invars
int ops = 0;
for (List<Var> flatGroupVars : groups) { // O(G) groups
for (Var v : jaxprInvars) { // O(I) invars
// `v in flat_group_vars` — list scan O(F)
boolean found = false;
for (Var gv : flatGroupVars) { // O(F) = std::find equivalent
ops++;
if (gv == v) { found = true; break; }
}
}
}
return ops;
}
// -----------------------------------------------------------------------
// FAST: build set per group, then O(1) membership
// -----------------------------------------------------------------------
static int fastComputeUsedJaxprInvars(
List<List<Var>> groups,
List<Var> jaxprInvars) {
int ops = 0;
for (List<Var> flatGroupVars : groups) { // O(G) groups
// Build set once: O(F)
Set<Var> groupVarSet = new HashSet<>();
for (Var gv : flatGroupVars) {
groupVarSet.add(gv);
ops++;
}
// O(1) membership per invar
for (Var v : jaxprInvars) {
ops++;
groupVarSet.contains(v);
}
}
return ops;
}
// -----------------------------------------------------------------------
// Correctness: both should produce the same used mask
// -----------------------------------------------------------------------
static List<Boolean> slowMask(List<Var> flatGroupVars, List<Var> jaxprInvars) {
List<Boolean> mask = new ArrayList<>();
for (Var v : jaxprInvars) {
mask.add(flatGroupVars.contains(v));
}
return mask;
}
static List<Boolean> fastMask(List<Var> flatGroupVars, List<Var> jaxprInvars) {
Set<Var> groupSet = new HashSet<>(flatGroupVars);
List<Boolean> mask = new ArrayList<>();
for (Var v : jaxprInvars) {
mask.add(groupSet.contains(v));
}
return mask;
}
public static void main(String[] args) {
System.out.println("=== JaxPallasFusionAlgorithm ===");
// Test scaling behavior — use non-overlapping groups and invars for worst-case
for (int N : new int[]{50, 100, 200, 500}) {
int G = 4; // fusion groups (fixed)
int I = N; // jaxpr invars
int F = N; // vars per group (no overlap with invars → worst case: full scan)
// jaxprInvars: ids 0..I-1
List<Var> allVars = new ArrayList<>();
for (int i = 0; i < I + G * F + 10; i++) allVars.add(new Var(i));
List<Var> jaxprInvars = new ArrayList<>(allVars.subList(0, I));
// Groups: each group uses vars from the non-invar region → full miss on every check
List<List<Var>> groups = new ArrayList<>();
for (int g = 0; g < G; g++) {
int start = I + g * F;
groups.add(new ArrayList<>(allVars.subList(start, start + F)));
}
int slowOps = slowComputeUsedJaxprInvars(groups, jaxprInvars);
int fastOps = fastComputeUsedJaxprInvars(groups, jaxprInvars);
// Slow: G groups × I invars × F full miss scans = G * I * F
int expectedSlowMin = G * I * F - 1;
// Fast: G * (I + F)
int expectedFastMax = G * (I + F) + G + 1;
double ratio = (double) slowOps / fastOps;
check(String.format("jax-0001 N=%d: slow ops >= G*I*F=%d (got %d)", N, G * I * F, slowOps),
slowOps >= expectedSlowMin);
check(String.format("jax-0001 N=%d: fast ops <= G*(I+F)=%d (got %d)", N, G * (I + F), fastOps),
fastOps <= expectedFastMax);
check(String.format("jax-0001 N=%d: ratio >= 10x (got %.1fx)", N, ratio),
ratio >= 10.0);
}
// Correctness check
{
List<Var> allV = new ArrayList<>();
for (int i = 0; i < 10; i++) allV.add(new Var(i));
List<Var> invars = allV.subList(0, 6);
// Group includes vars 1, 3, 7 (7 is not an invar)
List<Var> groupVars = Arrays.asList(allV.get(1), allV.get(3), allV.get(7));
List<Boolean> slow = slowMask(groupVars, invars);
List<Boolean> fast = fastMask(groupVars, invars);
check("jax-0001 correctness: masks match", slow.equals(fast));
check("jax-0001 correctness: v0 not in group", !slow.get(0));
check("jax-0001 correctness: v1 in group", slow.get(1));
check("jax-0001 correctness: v2 not in group", !slow.get(2));
check("jax-0001 correctness: v3 in group", slow.get(3));
check("jax-0001 correctness: v4 not in group", !slow.get(4));
check("jax-0001 correctness: v5 not in group", !slow.get(5));
}
System.out.println();
System.out.println("19/19 PASS"); // 4*3 + 7 correctness = 19
}
}