package unit; import java.util.*; /** * Unit test for pytorch-0003: * python_function.cpp tracer — std::find on trace_outputs (list) inside node+output loop. * O(N × K × T) → O(T + N × K) with HashSet pre-built before the loop. * * Compile: javac -d . PyTorchTracerAlgorithm.java * Run: java -ea unit.PyTorchTracerAlgorithm */ public class PyTorchTracerAlgorithm { static void check(String desc, boolean cond) { if (!cond) throw new AssertionError("FAIL: " + desc); System.out.println("PASS: " + desc); } static class Value { final int id; Value(int id) { this.id = id; } } static class Node { final List outputs; Node(List outputs) { this.outputs = outputs; } } // ----------------------------------------------------------------------- // SLOW: std::find on traceOutputs list per output per node // ----------------------------------------------------------------------- static int slowBuildSubgraph(List blockNodes, List traceOutputs) { int ops = 0; for (Node node : blockNodes) { // O(N) nodes for (Value output : node.outputs) { // O(K) outputs per node // std::find on traceOutputs: O(T) per output for (Value tv : traceOutputs) { ops++; if (tv == output) break; } } } return ops; } // ----------------------------------------------------------------------- // FAST: HashSet built once before loop // ----------------------------------------------------------------------- static int fastBuildSubgraph(List blockNodes, List traceOutputs) { int ops = 0; // Build set once: O(T) Set traceSet = new HashSet<>(); for (Value tv : traceOutputs) { traceSet.add(tv); ops++; } for (Node node : blockNodes) { // O(N) nodes for (Value output : node.outputs) { // O(K) outputs per node ops++; // O(1) set lookup traceSet.contains(output); } } return ops; } // ----------------------------------------------------------------------- // Correctness: both should register the same outputs // ----------------------------------------------------------------------- static List slowRegistered(List blockNodes, List traceOutputs) { List registered = new ArrayList<>(); for (Node node : blockNodes) { for (Value output : node.outputs) { if (traceOutputs.contains(output)) registered.add(output); } } return registered; } static List fastRegistered(List blockNodes, List traceOutputs) { Set traceSet = new HashSet<>(traceOutputs); List registered = new ArrayList<>(); for (Node node : blockNodes) { for (Value output : node.outputs) { if (traceSet.contains(output)) registered.add(output); } } return registered; } public static void main(String[] args) { System.out.println("=== PyTorchTracerAlgorithm ==="); // Worst case: none of block node outputs are in traceOutputs → full scan every time for (int N : new int[]{50, 100, 200, 500}) { int nodesCount = N; int K = 4; // outputs per node int T = N; // trace outputs (different from node outputs) // node outputs: values 0..N*K-1 List allNodeVals = new ArrayList<>(); for (int i = 0; i < nodesCount * K; i++) allNodeVals.add(new Value(i)); List blockNodes = new ArrayList<>(); for (int n = 0; n < nodesCount; n++) { blockNodes.add(new Node(allNodeVals.subList(n * K, n * K + K))); } // traceOutputs: fresh values NOT in any node output → full miss on every find List traceOutputs = new ArrayList<>(); for (int i = 0; i < T; i++) traceOutputs.add(new Value(nodesCount * K + i)); int slowOps = slowBuildSubgraph(blockNodes, traceOutputs); int fastOps = fastBuildSubgraph(blockNodes, traceOutputs); // Slow: N*K outputs × T full miss = N*K*T int expectedSlowMin = nodesCount * K * T - 1; // Fast: T + N*K lookups int expectedFastMax = T + nodesCount * K + 1; double ratio = (double) slowOps / fastOps; check(String.format("pytorch-0003 N=%d: slow ops >= N*K*T=%d (got %d)", N, nodesCount * K * T, slowOps), slowOps >= expectedSlowMin); check(String.format("pytorch-0003 N=%d: fast ops <= T+N*K=%d (got %d)", N, T + nodesCount * K, fastOps), fastOps <= expectedFastMax); check(String.format("pytorch-0003 N=%d: ratio >= 10x (got %.1fx)", N, ratio), ratio >= 10.0); } // Correctness: some outputs are trace outputs, both paths must agree { List pool = new ArrayList<>(); for (int i = 0; i < 20; i++) pool.add(new Value(i)); // Nodes: 5 nodes, 2 outputs each → outputs v0..v9 List nodes = new ArrayList<>(); for (int n = 0; n < 5; n++) { nodes.add(new Node(Arrays.asList(pool.get(n * 2), pool.get(n * 2 + 1)))); } // traceOutputs: v1, v3, v5, v7 (odd ones) List traceOutputs = Arrays.asList(pool.get(1), pool.get(3), pool.get(5), pool.get(7)); List slow = slowRegistered(nodes, traceOutputs); List fast = fastRegistered(nodes, traceOutputs); check("pytorch-0003 correctness: same count", slow.size() == fast.size()); check("pytorch-0003 correctness: 4 outputs registered", slow.size() == 4); for (int i = 0; i < slow.size(); i++) { check("pytorch-0003 correctness: output " + i + " matches", slow.get(i) == fast.get(i)); } } System.out.println(); System.out.println("18/18 PASS"); // 4*3 + 6 correctness = 18 } }