Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
68 changes: 65 additions & 3 deletions cpp/tests/testsearchnonn.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
#include <algorithm>
#include <iterator>
#include <iomanip>
#include <unordered_map>

#include "../core/fileutils.h"
#include "../dataio/sgf.h"
Expand Down Expand Up @@ -2638,14 +2639,25 @@ x.x.x
.........
)%%");

//Check that exactly the capped nodes that reached the cap have a snapshot, and that their child weight distribution
//matches it. Relies on this test net giving every visit the same weight. Multithreading and virtual losses can
//make the distribution overshoot, so the tolerance is looser then.
//Check that capped nonterminal nodes that reached the cap have a snapshot, and that their child weight distribution
//matches it. Terminal leaves have no children to freeze. Relies on this test net giving every visit the same weight.
//Multithreading and virtual losses can make the distribution overshoot, so the tolerance is looser then.
//Returns the number of snapshot nodes well past the cap.
auto checkTree = [&](Search* search, Player cappedPla, int64_t visitCap, int numThreads) {
testAssert(search->getVisitCap(P_BLACK) == (cappedPla == P_BLACK ? visitCap : 0));
testAssert(search->getVisitCap(P_WHITE) == (cappedPla == P_WHITE ? visitCap : 0));
testAssert(!search->searchParams.useGraphSearch);
vector<SearchNode*> nodes = search->enumerateTreePostOrder();
std::unordered_map<const SearchNode*, std::pair<const SearchNode*, Loc>> parents;
for(const SearchNode* parent: nodes) {
ConstSearchNodeChildrenReference children = parent->getChildren();
for(int i = 0; i < children.getCapacity(); i++) {
const SearchNode* child = children[i].getIfAllocated();
if(child == NULL)
break;
testAssert(parents.emplace(child, std::make_pair(parent, children[i].getMoveLoc())).second);
}
}
int numNodesWellPastCap = 0;
for(SearchNode* node: nodes) {
int64_t visits = node->stats.visits.load(std::memory_order_acquire);
Expand All @@ -2654,6 +2666,32 @@ x.x.x
testAssert(snapshot == NULL);
continue;
}
if(node->getNNOutput() == NULL) {
//Terminal playouts add their outcome directly, without NN evaluation or
//child-weight recomputation. Replay the path to require an actual terminal
//position, so a missing snapshot on a nonterminal node cannot be excused.
vector<Loc> path;
for(const SearchNode* current = node; current != search->rootNode;) {
auto parent = parents.find(current);
testAssert(parent != parents.end());
path.push_back(parent->second.second);
current = parent->second.first;
}
Board replayBoard = search->rootBoard;
BoardHistory replayHistory = search->rootHistory;
Player replayPla = search->rootPla;
for(auto move = path.rbegin(); move != path.rend(); ++move) {
replayHistory.makeBoardMoveAssumeLegal(replayBoard, *move, replayPla, NULL);
replayPla = getOpp(replayPla);
}
testAssert(replayHistory.isGameFinished);
testAssert(replayPla == node->nextPla);
testAssert(!node->forceNonTerminal);
testAssert(node->state.load(std::memory_order_acquire) == SearchNode::STATE_UNEVALUATED);
testAssert(node->getChildren().getCapacity() == 0);
testAssert(snapshot == NULL);
continue;
}
testAssert(snapshot != NULL);
if(visits >= 3 * visitCap)
numNodesWellPastCap++;
Expand Down Expand Up @@ -2697,6 +2735,30 @@ x.x.x
return numNodesWellPastCap;
};

{
SearchParams params;
params.maxVisits = 80;
params.visitCapContempt = 20;
params.visitCapContemptPla = P_BLACK;
Search search(params, nnEval, &logger, "visitCapTerminalSeed");
BoardHistory hist(board, P_BLACK, rules, 0, BoardHistoryModes(false,false));
search.setPosition(P_BLACK, board, hist);
//Only pass is allowed for both plies, so the capped terminal child must
//accumulate well past the cap regardless of NN preferences or scheduling.
vector<int> avoid(Board::MAX_ARR_SIZE, 2);
avoid[Board::PASS_LOC] = 0;
search.setAvoidMoveUntilByLoc(avoid, avoid);
search.runWholeSearch(P_BLACK);
const SearchNode* afterPass = search.getChildForMove(search.rootNode, Board::PASS_LOC);
testAssert(afterPass != NULL);
const SearchNode* terminal = search.getChildForMove(afterPass, Board::PASS_LOC);
testAssert(terminal != NULL);
testAssert(terminal->stats.visits.load(std::memory_order_acquire) >= 3 * params.visitCapContempt);
testAssert(terminal->getNNOutput() == NULL);
testAssert(terminal->visitCapSnapshot.load(std::memory_order_acquire) == NULL);
checkTree(&search, P_BLACK, params.visitCapContempt, params.numThreads);
}

auto runTest = [&](int numThreads) {
cout << "Threads: " << numThreads << endl;
SearchParams params;
Expand Down