commit 4a265d34937ad339697218e4fd919b17b678ab6d
parent 5d227d1abca84be81af16d07d3b9bc7be49f5b19
Author: Andrew Laack <andrew@laack.co>
Date: Mon, 21 Sep 2026 09:47:30 -0500
Added return double for edge traversed cost from primonestep, added inadmissable heuristic test and assertion it never performs better at mst induction than prim's algorithm.
Diffstat:
5 files changed, 108 insertions(+), 7 deletions(-)
diff --git a/include/prim.hpp b/include/prim.hpp
@@ -11,6 +11,6 @@ void explore(
std::size_t cIdx,
std::priority_queue<Edge, std::vector<Edge>, std::greater<Edge>> &toVisit,
Edge ¤t, Graph &g, std::vector<double>& minVertWeight);
-void oneStepPrim(
+double oneStepPrim(
std::priority_queue<Edge, std::vector<Edge>, std::greater<Edge>> &toVisit,
std::unordered_set<std::size_t> &visitedIndices, Graph &g, std::vector<double>& minVertWeight);
diff --git a/src/main.cpp b/src/main.cpp
@@ -114,7 +114,8 @@ int main(int argc, char** argv) {
EndDrawing();
usleep((int)(sleepTime * 1000000));
- oneStepPrim(toVisit, visitedIndices, g, minEdgeToVertex);
+ // onestep prim returns the cost of the edge we chose.
+ double _ = oneStepPrim(toVisit, visitedIndices, g, minEdgeToVertex);
}
UnloadRenderTexture(blankGraph);
diff --git a/src/prim.cpp b/src/prim.cpp
@@ -30,18 +30,18 @@ void explore(
delete edges;
}
-void oneStepPrim(
+double oneStepPrim(
std::priority_queue<Edge, std::vector<Edge>, std::greater<Edge>>& toVisit,
std::unordered_set<std::size_t>& visitedIndices, Graph& g,
std::vector<double>& minVertWeight) {
bool found = false;
if (toVisit.size() == 0) {
- return;
+ return 0;
}
while (found == false) {
if (toVisit.size() == 0) {
- return;
+ return 0;
}
found = true;
@@ -51,13 +51,16 @@ void oneStepPrim(
if (visitedIndices.find(current.v2Index) == visitedIndices.end()) {
auto cIdx = current.v2Index;
explore(cIdx, toVisit, current, g, visitedIndices, minVertWeight);
+ return current.length2;
} else if (visitedIndices.find(current.v1Index) ==
visitedIndices.end()) {
auto cIdx = current.v1Index;
explore(cIdx, toVisit, current, g, visitedIndices, minVertWeight);
+ return current.length2;
} else {
found = false;
}
}
+ return 0;
}
diff --git a/tests/algo_test.cpp b/tests/algo_test.cpp
@@ -1,5 +1,7 @@
#include <catch2/catch_test_macros.hpp>
#include <cstddef>
+#include <iostream>
+#include <random>
#include "../include/prim.hpp"
@@ -207,3 +209,98 @@ TEST_CASE("Small Prim Test", "[Small full validation]") {
REQUIRE(visitedIndices.size() == vertCount);
}
}
+
+// traverse the graph randomly, ensuring we visit all vertices then return edge
+// weight sum. This ensures our admissible heuristic always beats or is on par
+// with an inadmissible one (this one).
+
+double randomTraversalCost(Graph g) {
+ std::uint32_t seed = std::random_device{}();
+ std::mt19937 rng{seed};
+
+ std::size_t count = g.getVertexCount();
+ std::uniform_int_distribution<std::size_t> pick1(0, count - 1);
+
+ std::size_t traversed = 0;
+ double cost = 0;
+
+ g.traverseVertexIdx(pick1(rng)); // start
+ traversed += 1;
+
+ while (traversed < g.getVertexCount()) {
+ std::size_t idx = pick1(rng);
+ if (g.getVertex(idx).visited) {
+ std::vector<Edge> edges = g.getEdgesOfVertexIdx(idx);
+ std::uniform_int_distribution<std::size_t> pick2(0,
+ edges.size() - 1);
+ std::size_t selection = pick2(rng);
+ if (!g.getVertex(edges[selection].v2Index).visited) {
+ cost += edges[selection].length2;
+ g.setEdgeTraversed(edges[selection]);
+ g.traverseVertexIdx(edges[selection].v2Index);
+ traversed += 1;
+ }
+ }
+ }
+ return cost;
+}
+
+TEST_CASE("Prim Correctness Test", "[Correctness test for MST]") {
+ for (int i = 0; i < 100; ++i) {
+ std::size_t edgeCount = 10;
+ std::size_t vertCount = 5;
+
+ float xMax = 5120;
+ float yMax = 1440;
+
+ Graph g = Graph(edgeCount, vertCount, xMax, yMax);
+ do {
+ g = Graph(edgeCount, vertCount, xMax, yMax);
+ } while (!isConnected(g)); // this passes by value
+
+ double rndCost = randomTraversalCost(g); // this passes by value
+
+ std::unordered_set<std::size_t> visitedIndices{};
+ std::priority_queue<Edge, std::vector<Edge>, std::greater<Edge>>
+ toVisit{};
+
+ for (auto edge : g.getEdgesOfVertexIdx(0)) toVisit.push(edge);
+ g.traverseVertexIdx(0);
+ visitedIndices.insert(0);
+
+ std::unordered_set<std::size_t> vBefore = visitedIndices;
+
+ bool havePrior = false;
+ Edge prior = toVisit.top();
+ std::unordered_set<std::size_t> visibleAtPrior;
+ std::vector<double> minEdgeToVertex(vertCount, -1);
+
+ double primCost = 0;
+
+ while (toVisit.size() != 0) {
+ auto current = toVisit.top();
+ if (havePrior) {
+ bool wasPresent = visibleAtPrior.count(current.v1Index) > 0 ||
+ visibleAtPrior.count(current.v2Index) > 0;
+ // anytime we use the same source node two steps in a row, the
+ // second weight must be smaller.
+ if (wasPresent) {
+ REQUIRE(current.length2 >= prior.length2);
+ }
+ }
+
+ prior = current;
+ visibleAtPrior = visitedIndices;
+ havePrior = true;
+
+ primCost +=
+ oneStepPrim(toVisit, visitedIndices, g, minEdgeToVertex);
+ bool valid = vBefore.size() + 1 == visitedIndices.size() ||
+ vBefore.size() == vertCount;
+ REQUIRE(valid);
+ vBefore = visitedIndices;
+ }
+ REQUIRE(visitedIndices.size() == vertCount);
+ REQUIRE(primCost <= rndCost);
+ }
+}
diff --git a/tests/crash_test.cpp b/tests/crash_test.cpp
@@ -52,8 +52,8 @@ void run(size_t mv, size_t me, uint32_t iterations) {
}
}
-void runs() { run(100, 500, 10000); }
-void runl() { run(100000, 1000000, 100); }
+void runs() { run(100, 500, 1000); }
+void runl() { run(100000, 1000000, 10); }
TEST_CASE("Small prim algorithm not guaranteed connected", "[Small prim]") {
std::vector<std::thread*> threads{};