From 06e5b85a9ba69ce53c1aeaa9fd5bcdacdec9c284 Mon Sep 17 00:00:00 2001 From: tqlong Date: Mon, 1 Nov 2010 03:08:16 +0000 Subject: [PATCH] auction algorithm: start allnn dual tree --- .../anmf/allnn_auction_max_weight_matching.h | 196 ++++++ .../anmf/allnn_kdtree_distance_matrix.h | 664 ++++++++++++++++++ 2 files changed, 860 insertions(+) create mode 100644 fastlib/trunk/contrib/tqlong/anmf/allnn_auction_max_weight_matching.h create mode 100644 fastlib/trunk/contrib/tqlong/anmf/allnn_kdtree_distance_matrix.h diff --git a/fastlib/trunk/contrib/tqlong/anmf/allnn_auction_max_weight_matching.h b/fastlib/trunk/contrib/tqlong/anmf/allnn_auction_max_weight_matching.h new file mode 100644 index 0000000000..c9c27ed21d --- /dev/null +++ b/fastlib/trunk/contrib/tqlong/anmf/allnn_auction_max_weight_matching.h @@ -0,0 +1,196 @@ +#ifndef AUCTION_MAX_WEIGHT_MATCHING_H +#define AUCTION_MAX_WEIGHT_MATCHING_H + +#include "anmf.h" +#include +#include +#include +#include + +BEGIN_ANMF_NAMESPACE; + +template + class AuctionMaxWeightMatching +{ +public: + /** typename W should implement n_rows(), n_cols(), get(i, j), setPrice(j, p) */ + typedef W weight_matrix_type; + + /** Constructor */ + AuctionMaxWeightMatching(weight_matrix_type& weight, bool doMatch = false); + + /** return item that matches person leftIndex */ + int leftMatch(int leftIndex) const { return doneMatching_ ? leftMatch_[leftIndex] : -1; } + + /** return person that matches item rightIndex */ + int rightMatch(int rightIndex) const { return doneMatching_ ? rightMatch_[rightIndex] : -1; } + + /** check if matched */ + bool matched(int index) const { return index != -1; } + void unmatch(int& index) const { index = -1; } + + /** check if matching is done */ + bool doneMatching() const { return doneMatching_; } + + /** the auction algorithm for max weight matching */ + void doMatch(); +protected: + weight_matrix_type& weight_; + double epsilon_; + int n_rows_, n_cols_; + std::vector leftMatch_, rightMatch_; + bool doneMatching_; + + std::vector price_, bid_, best_surplus_, second_surplus_; + std::vector winner_, best_item_; + + /** The forward auction algorithm (i.e. use only price of items as dual) */ + void forwardAuction(); + + /** Place a bid */ + void placeBid(int bidder, int item, double price); + + /** Clear bids for next iteration */ + void clearBids(); + + /** naively get the best and second best items in term of surplus = benefit - price + * use this function if typename W does not implement it + */ + void getBestAndSecondBest(int bidder, int& best_item, double& best_surplus, double& second_surplus); + + /** add a match (person, item) to the assigment */ + void setMatch(int new_bidder, int item); + + /** check the epsilon-Complementary condition + * for every match (person, item) the surplus is at least max of surpluses minus epsilon + */ + bool checkEpsilonComplementary() const; +private: +}; + +template + AuctionMaxWeightMatching::AuctionMaxWeightMatching(weight_matrix_type &weight, bool match) + : weight_(weight), + n_rows_(weight_.n_rows()), n_cols_(weight.n_cols()), + leftMatch_(n_rows_, -1), rightMatch_(n_cols_, -1), + doneMatching_(false), + price_(n_cols_, 0), bid_(n_cols_, -std::numeric_limits::infinity()), + best_surplus_(n_rows_, 0), second_surplus_(n_rows_, 0), + winner_(n_cols_, -1), best_item_(n_rows_, -1) +{ + epsilon_ = 1.0 / n_rows_; + if (match) + doMatch(); +} + +template + void AuctionMaxWeightMatching::doMatch() +{ + forwardAuction(); +} + +template + void AuctionMaxWeightMatching::forwardAuction() +{ + long int pruned = 0, total = 0, iter = 0; + while (!doneMatching_) + { + iter++; +// std::cout << (checkEpsilonComplementary() ? "e-CS satisfied" : "e-CS not satisfied") << std::endl; + + doneMatching_ = weight_.allMatched(); + if (doneMatching_) break; + clearBids(); + pruned += weight_.getAllBestAndSecondBest(best_item_, best_surplus_, second_surplus_); + + for (int i = 0; i < n_rows_; i++) if (!matched(leftMatch_[i])) + { + int j = best_item_[i]; +// std::cout << "i = " << i << " j = " << j << "\n"; + placeBid(i, j, price_[j]+best_surplus_[i]-second_surplus_[i]+epsilon_); + total += n_cols_; + } +// std::cout << "Done 1 \n"; + // for all items, assign them to best bidder + for (int j = 0; j < n_cols_; j++) if (winner_[j] != -1) + { +// std::cout << "winner = " << winner_[j] << " j = " << j << "\n"; + price_[j] = bid_[j]; + weight_.setPrice(j, PointStatistics(price_[j], winner_[j])); + setMatch(winner_[j], j); + } + if (iter % (weight_.n_rows()/10) == 0) std::cout << iter << " calculations = " << pruned << "/" << total << "\n"; + } + if (iter % (weight_.n_rows()/10) == 0) std::cout << iter << " calculations = " << pruned << "/" << total << "\n"; +} + +template + void AuctionMaxWeightMatching::setMatch(int new_bidder, int item) +{ + int old_bidder = rightMatch_[item]; + leftMatch_[new_bidder] = item; + rightMatch_[item] = new_bidder; + if (matched(old_bidder)) + unmatch(leftMatch_[old_bidder]); +} + +template + void AuctionMaxWeightMatching::clearBids() +{ + std::fill(bid_.begin(), bid_.end(), -std::numeric_limits::infinity()); + std::fill(winner_.begin(), winner_.end(), -1); +} + +template + void AuctionMaxWeightMatching::placeBid(int bidder, int item, double price) +{ + if (bid_[item] < price) + { + bid_[item] = price; + winner_[item] = bidder; + } +} + +template + void AuctionMaxWeightMatching::getBestAndSecondBest(int bidder, + int &best_item, + double &best_surplus, + double &second_surplus) +{ + best_surplus = second_surplus = -std::numeric_limits::infinity(); + for (int item = 0; item < n_cols_; item++) + { + double surplus = weight_.get(bidder, item) - price_[item]; + if (surplus > best_surplus) + { + best_item = item; + second_surplus = best_surplus; + best_surplus = surplus; + } + else if (surplus > second_surplus) + { + second_surplus = surplus; + } + } +} + +template + bool AuctionMaxWeightMatching::checkEpsilonComplementary() const +{ + for (int bidder = 0; bidder < n_rows_; bidder++) if (matched(leftMatch_[bidder])) + { + int item = leftMatch_[bidder]; + double surplus = weight_.get(bidder, item) - price_[item]; + for (int j = 0; j < n_cols_; j++) if (surplus < weight_.get(bidder, j) - price_[j] - epsilon_ - 1e-10) + { +// std::cout << "check bidder " << bidder << " item " << item << " " +// << surplus << " " << weight_.get(bidder, j) - price_[j] - epsilon_ << std::endl; + return false; + } + } + return true; +} + +END_ANMF_NAMESPACE; + +#endif // AUCTION_MAX_WEIGHT_MATCHING_H diff --git a/fastlib/trunk/contrib/tqlong/anmf/allnn_kdtree_distance_matrix.h b/fastlib/trunk/contrib/tqlong/anmf/allnn_kdtree_distance_matrix.h new file mode 100644 index 0000000000..8bdc7d989f --- /dev/null +++ b/fastlib/trunk/contrib/tqlong/anmf/allnn_kdtree_distance_matrix.h @@ -0,0 +1,664 @@ +#ifndef KDTREE_DISTANCE_MATRIX_H +#define KDTREE_DISTANCE_MATRIX_H + +#include "anmf.h" +#include +#include +#include +#include + +BEGIN_ANMF_NAMESPACE; + +class PointStatistics +{ + friend class NodeStatistics; + double price_; + int matchTo_; +public: + PointStatistics(double price = 0.0, int matchTo = -1) : price_(price), matchTo_(matchTo) {} +// PointStatistics& operator= (double price) { price_ = price; return *this; } + double price() const { return price_; } + int matchTo() const { return matchTo_; } +}; + +class NodeStatistics +{ + double minPrice_, maxPrice_; + bool allMatched_; + Vector minBox_, maxBox_; +public: + double minPrice() const { return minPrice_; } + double maxPrice() const { return maxPrice_; } + const Vector& minBox() const { return minBox_; } + const Vector& maxBox() const { return maxBox_; } + bool allMatched() const { return allMatched_; } + + double distance(const Vector& x) const + { + double s = 0; + for (int i = 0; i < minBox_.length(); i++) + { + if (x[i] < minBox_[i]) s += math::Sqr(x[i]-minBox_[i]); + else if (x[i] > maxBox_[i]) s += math::Sqr(x[i]-maxBox_[i]); + } + return sqrt(s); + } + + /** Create bounding box for prices and coordinates of leaf node */ + void InitAtLeaf(const Matrix& points, const std::vector& pointStatistics, + const std::vector& indexMap, int dfsIndex, int n_points) + { + minPrice_ = std::numeric_limits::infinity(); + maxPrice_ = -std::numeric_limits::infinity(); + minBox_.Init(points.n_rows()); + maxBox_.Init(points.n_rows()); + minBox_.SetAll(std::numeric_limits::infinity()); + maxBox_.SetAll(-std::numeric_limits::infinity()); + allMatched_ = true; + for (int i = 0; i < n_points; i++) + { + int index = dfsIndex+i; + int oldIndex = indexMap[index]; + + double price = pointStatistics[oldIndex].price_; + if (pointStatistics[oldIndex].matchTo_ == -1) allMatched_ = false; + if (price < minPrice_) minPrice_ = price; + if (price > maxPrice_) maxPrice_ = price; + + for (int dim = 0; dim < points.n_rows(); dim++) + { + double val = points.get(dim, oldIndex); + if (val < minBox_[dim]) minBox_[dim] = val; + if (val > maxBox_[dim]) maxBox_[dim] = val; + } + } + } + + void InitFromChildren(const NodeStatistics& leftStats, const NodeStatistics& rightStats) + { + minBox_.Init(leftStats.minBox_.length()); + maxBox_.Init(leftStats.maxBox_.length()); + minPrice_ = leftStats.minPrice_ < rightStats.minPrice_ ? leftStats.minPrice_ : rightStats.minPrice_; + maxPrice_ = leftStats.maxPrice_ > rightStats.maxPrice_ ? leftStats.maxPrice_ : rightStats.maxPrice_; + allMatched_ = leftStats.allMatched_ && rightStats.allMatched_; + for (int dim = 0; dim < minBox_.length(); dim++) + { + minBox_[dim] = leftStats.minBox_[dim] < rightStats.minBox_[dim] ? leftStats.minBox_[dim] : rightStats.minBox_[dim]; + maxBox_[dim] = leftStats.maxBox_[dim] > rightStats.maxBox_[dim] ? leftStats.maxBox_[dim] : rightStats.maxBox_[dim]; + } + } + + void ResetAtLeaf(const Matrix& points, const std::vector& pointStatistics, + const std::vector& indexMap, int dfsIndex, int n_points, bool resetBoundingBox = false) + { + minPrice_ = std::numeric_limits::infinity(); + maxPrice_ = -std::numeric_limits::infinity(); + if (resetBoundingBox) + { + minBox_.SetAll(std::numeric_limits::infinity()); + maxBox_.SetAll(-std::numeric_limits::infinity()); + } + + allMatched_ = true; + for (int i = 0; i < n_points; i++) + { + int index = dfsIndex+i; + int oldIndex = indexMap[index]; + + double price = pointStatistics[oldIndex].price_; + if (pointStatistics[oldIndex].matchTo_ == -1) allMatched_ = false; + if (price < minPrice_) minPrice_ = price; + if (price > maxPrice_) maxPrice_ = price; + + if (resetBoundingBox) + { + for (int dim = 0; dim < points.n_rows(); dim++) + { + double val = points.get(dim, oldIndex); + if (val < minBox_[dim]) minBox_[dim] = val; + if (val > maxBox_[dim]) maxBox_[dim] = val; + } + } + } + } + + void ResetFromChildren(const NodeStatistics& leftStats, const NodeStatistics& rightStats, bool resetBoundingBox = false) + { + minPrice_ = leftStats.minPrice_ < rightStats.minPrice_ ? leftStats.minPrice_ : rightStats.minPrice_; + maxPrice_ = leftStats.maxPrice_ > rightStats.maxPrice_ ? leftStats.maxPrice_ : rightStats.maxPrice_; + allMatched_ = leftStats.allMatched_ && rightStats.allMatched_; + if (resetBoundingBox) + for (int dim = 0; dim < minBox_.length(); dim++) + { + minBox_[dim] = leftStats.minBox_[dim] < rightStats.minBox_[dim] ? leftStats.minBox_[dim] : rightStats.minBox_[dim]; + maxBox_[dim] = leftStats.maxBox_[dim] > rightStats.maxBox_[dim] ? leftStats.maxBox_[dim] : rightStats.maxBox_[dim]; + } + } + + std::string toString() const + { + std::stringstream s; + s << "price = (" << minPrice_ << "," << maxPrice_ << ")" + << " box = " << anmf::toString(minBox_) << " --> " << anmf::toString(maxBox_); + return s.str(); + } +}; + +class KDNode +{ +protected: + /** Global properties of a tree */ + const Matrix& points_; + std::vector* pointStatistics_; + std::vector* oldFromNewIndex_; + std::vector *nodeFromOldIndex_; + + /** Node properties */ + int n_points_, dfsIndex_; + NodeStatistics nodeStatistics_; + std::vector children_; + KDNode* parent_; +public: + /** Constructor for root node */ + template + KDNode(const Matrix& points, const std::vector& stats) + : points_(points) + { + // create global properties for the tree + int n = points_.n_cols(); + oldFromNewIndex_ = new std::vector(n); + pointStatistics_ = new std::vector(n); + nodeFromOldIndex_ = new std::vector(n); + for (int i = 0; i < n; i++) + { + oldFromNewIndex_->at(i) = i; + pointStatistics_->at(i) = stats[i]; + nodeFromOldIndex_->at(i) = NULL; + } + n_points_ = n; + dfsIndex_ = 0; + parent_ = NULL; + splitMidPoint(0); + visitToSetStatistics(); + } + + /** set point statistics and traverse up the tree */ + void setPointStatistics(int index, const PointStatistics& stats) + { + pointStats(index) = stats; + leaf(index)->resetStatistics(); + } + + /** convert this subtree to string (to print) */ + std::string toString(int depth = 0) const + { + std::stringstream s; + if (depth == 0) + { + for (int i = 0; i < n_points_; i++) + { + int oldIndex = oldFromNewIndex_->at(i); + s << i << " --> "; + Vector p_i; + points_.MakeColumnVector(oldIndex, &p_i); + s << anmf::toString(p_i) << "\n"; + } + } + for (int i = 0; i < depth; i++) s << " "; + s << "-- Node (" << dfsIndex_; + for (int i = 1; i < n_points_; i++) + s << "," << dfsIndex_ + i; + s << ")\n"; + for (int i = 0; i < depth+1; i++) s << " "; + s << " " << nodeStatistics_.toString() << "\n"; + for (unsigned int i = 0; i < children_.size(); i++) + s << children_[i]->toString(depth+1); + return s.str(); + } + + /** choose a random point to set the bounds */ + void randomBound(const Vector& x, int& minIndex, double& minSoFar) + { + minIndex = dfsIndex_ + math::RandInt(0, n_points_); +// minIndex = dfsIndex_ + math::RandInt(0, n_points_); + int oldIndex = oldFromNewIndex(minIndex); + Vector p_i; + points_.MakeColumnVector(oldIndex, &p_i); + double d = sqrt(la::DistanceSqEuclidean(p_i, x)); + double p = pointStats(minIndex).price(); + minSoFar = d+p; + } + + void randomBound(const Vector& x, int& minIndex, double& minSoFar, int& sndIndex, double& sndMinSoFar) + { + minIndex = dfsIndex_ + math::RandInt(0, n_points_); +// minIndex = dfsIndex_ + math::RandInt(0, n_points_); + int oldIndex = oldFromNewIndex(minIndex); + Vector p_i; + points_.MakeColumnVector(oldIndex, &p_i); + double d = sqrt(la::DistanceSqEuclidean(p_i, x)); + double p = pointStats(minIndex).price(); + minSoFar = d+p; + + sndIndex = dfsIndex_ + math::RandInt(0, n_points_); +// minIndex = dfsIndex_ + math::RandInt(0, n_points_); + oldIndex = oldFromNewIndex(sndIndex); + Vector p_j; + points_.MakeColumnVector(oldIndex, &p_j); + d = sqrt(la::DistanceSqEuclidean(p_j, x)); + p = pointStats(sndIndex).price(); + sndMinSoFar = d+p; + + if (minSoFar > sndMinSoFar) + { + int tmp = minIndex; minIndex = sndIndex; sndIndex = tmp; + double dtmp = minSoFar; minSoFar = sndMinSoFar; sndMinSoFar = dtmp; + } + } + /** get Nearest Neighbor index (newIndex) in term of distance + price */ + void nearestNeighbor(const Vector& x, int& minIndex, double& minSoFar) + { + // check bound + double d = nodeStatistics_.distance(x); + double minPrice = nodeStatistics_.minPrice(); +// std::cout << "lb = " << d+minPrice << " minIndex = " << minIndex << " minSoFar = " << minSoFar << "\n"; + if (d+minPrice >= minSoFar) + { +// std::cout << "PRUNE " << nodeStatistics_.toString() << " n_points = " << n_points_ << "\n"; + return; + } + else + { +// std::cout << "Search " << nodeStatistics_.toString() << " n_points = " << n_points_ << "\n"; + } + + if (isLeaf()) // at leaf, do naive search + { + for (int i = dfsIndex_; i < dfsIndex_+n_points_; i++) + { + int oldIndex = oldFromNewIndex(i); + Vector p_i; + points_.MakeColumnVector(oldIndex, &p_i); + double d = sqrt(la::DistanceSqEuclidean(p_i, x)); + double p = pointStats(i).price(); + if (d+p < minSoFar) + { + minSoFar = d+p; + minIndex = i; + } + } + } + else + { + children_[0]->nearestNeighbor(x, minIndex, minSoFar); + children_[1]->nearestNeighbor(x, minIndex, minSoFar); + } + } + + /** return the number of pruned calculations */ + long int nearestNeighbor(const Vector& x, int& minIndex, double& minSoFar, int& sndIndex, double& sndMinSoFar) + { + // check bound + double d = nodeStatistics_.distance(x); + double minPrice = nodeStatistics_.minPrice(); +// std::cout << "lb = " << d+minPrice << " minIndex = " << minIndex << " minSoFar = " << minSoFar << "\n"; + if (d+minPrice >= sndMinSoFar) + { +// std::cout << "PRUNE " << nodeStatistics_.toString() << " n_points = " << n_points_ << "\n"; + return n_points_; + } + else + { +// std::cout << "Search " << nodeStatistics_.toString() << " n_points = " << n_points_ << "\n"; + } + + if (isLeaf()) // at leaf, do naive search + { + for (int i = dfsIndex_; i < dfsIndex_+n_points_; i++) + { + int oldIndex = oldFromNewIndex(i); + Vector p_i; + points_.MakeColumnVector(oldIndex, &p_i); + double d = sqrt(la::DistanceSqEuclidean(p_i, x)); + double p = pointStats(i).price(); + if (d+p < minSoFar) + { + sndIndex = minIndex; + sndMinSoFar = minSoFar; + minIndex = i; + minSoFar = d+p; + } + else if (d+p < sndMinSoFar) + { + sndIndex = i; + sndMinSoFar = d+p; + } + } + return 0; + } + else + { + long int left_pruned = children_[0]->nearestNeighbor(x, minIndex, minSoFar, sndIndex, sndMinSoFar); + long int right_pruned = children_[1]->nearestNeighbor(x, minIndex, minSoFar, sndIndex, sndMinSoFar); + return left_pruned + right_pruned; + } + } + + /** Basic getters and setters */ + int n_points() const { return n_points_; } + int dfsIndex() const { return dfsIndex_; } +// NodeStatistics& stats() { return nodeStatistics_; } + const NodeStatistics& stats() const { return nodeStatistics_; } + KDNode* parent() const { return parent_; } + int n_children() const { return (int) children_.size(); } + KDNode* child(int index) const { return children_[index]; } + const PointStatistics& pointStats(int index) const { return pointStatistics_->at(oldFromNewIndex(index)); } + PointStatistics& pointStats(int index) { return pointStatistics_->at(oldFromNewIndex(index)); } + int oldFromNewIndex(int index) const { return oldFromNewIndex_->at(index); } + KDNode* leaf(int index) const { return nodeFromOldIndex_->at(oldFromNewIndex(index)); } + bool isLeaf() const { return children_.empty(); } + bool allMatched() const { return nodeStatistics_.allMatched(); } +protected: + + /** Constructor for a child node */ + KDNode(KDNode* parent) + : points_(parent->points_), + pointStatistics_(parent->pointStatistics_), + oldFromNewIndex_(parent->oldFromNewIndex_), + nodeFromOldIndex_(parent->nodeFromOldIndex_), + parent_(parent) + { + DEBUG_ASSERT(parent); + parent->children_.push_back(this); + } + + /** Split the points in a node by the median point at certain dimension */ + void splitMidPoint(int dim) + { +// std::cout << dfsIndex_ << " --> " << dfsIndex_+n_points_-1 << "\n"; + if (n_points_ < 10) return; // the node is too small to split + + int n_left = 0, n_right = 0; + double mid = findMidPoint(dim, n_left, n_right); // find mid value at dim dimension + if (n_left == 0 || n_right == 0) // no mid value found + { +// std::cout << "cannot split\n"; + return; + } + + KDNode *left = new KDNode(this), *right = new KDNode(this); // create two left and right nodes + left->n_points_ = n_left; + left->dfsIndex_ = this->dfsIndex_; + right->n_points_ = n_right; + right->dfsIndex_ = this->dfsIndex_ + left->n_points_; + + // now move the points to theirs right location + n_left = n_right = 0; + std::vector *oldMap = new std::vector(oldFromNewIndex_->begin()+dfsIndex_, oldFromNewIndex_->begin()+dfsIndex_+n_points_); + for (int i = 0; i < n_points_; i++) // for each point assign to left or right node by its value at dim + { + int oldIndex = oldMap->at(i); + double val = points_.get(dim, oldIndex); + if (val < mid) + { + oldFromNewIndex_->at(left->dfsIndex_+n_left) = oldIndex; + n_left++; + } + else + { + oldFromNewIndex_->at(right->dfsIndex_+n_right) = oldIndex; + n_right++; + } + } + delete oldMap; + + dim = (dim+1) % points_.n_rows(); + left->splitMidPoint(dim); + right->splitMidPoint(dim); + } + + /** find the median point at a certain dimension + * try other dimensions if cannot split + */ + double findMidPoint(int& dim, int& n_left, int& n_right) + { + std::vector vals(n_points_); + for (int k = 0; k < points_.n_rows(); k++) + { + n_left = n_right = 0; + for (int i = 0; i < n_points_; i++) + vals[i] = points_.get(dim, oldFromNewIndex_->at(i+dfsIndex_)); + + double mid = selectMedian(vals, n_points_); // using median of medians algorithm here + + // check if the split is ok + for (int i = 0; i < n_points_; i++) + { + int newIndex = i+dfsIndex_; + int oldIndex = oldFromNewIndex_->at(newIndex); + double val = points_.get(dim, oldIndex); +// std::cout << "val = " << val << "\n"; + if (val < mid) n_left++; + else n_right++; + } +// std::cout << "dim = " << dim << " mid = " << mid +// << " n_left = " << n_left << " n_right = " << n_right << "\n"; + if (n_left >= 1 && n_right >= 1) return mid; + else + dim = (dim+1) % points_.n_rows(); + } + n_left = 0; + n_right = 0; + return std::numeric_limits::quiet_NaN(); + } + + /** The median of medians algorithm */ + double selectMedian(std::vector& x, int n) + { + if (n < 5) return selectMedianSmall(x, 0, n); + int k = 0; + for (int i = 0; i < n; i+=5) + x[k++] = selectMedianSmall(x, i, n-i > 5 ? 5 : n-i); + return selectMedian(x, k); + } + + /** Select median of a short array (n <= 5) */ + double selectMedianSmall(const std::vector& x, int s, int n) + { + DEBUG_ASSERT(n != 0); + double mid; + if (n == 1) mid = x[s]; + else if (n == 2) mid = (x[s]+x[s+1])/2; + else if (n == 3) + { + if (x[s] < x[s+1]) + { + if (x[s+1] < x[s+2]) mid = x[s+1]; + else if (x[s] < x[s+2]) mid = x[s+2]; + else mid = x[s]; + } + else + { + if (x[s+1] > x[s+2]) mid = x[s+1]; + else if (x[s] > x[s+2]) mid = x[s+2]; + else mid = x[s]; + } + } + else if (n == 4) + { + double min1, min2; + double max1, max2; + min1 = min2 = std::numeric_limits::infinity(); + max1 = max2 = -min1; + for (int i = 0; i < n; i++) + { + if (x[s+i] < min1) { min2 = min1; min1 = x[s+i]; } + else if (x[s+i] < min2) min2 = x[s+i]; + if (x[s+i] > max1) { max2 = max1; max1 = x[s+i]; } + else if (x[s+i] > max2) max2 = x[s+i]; + } + mid = (min2+max2)/2; + } + else + { + double min1, min2; + double max1, max2; + min1 = min2 = std::numeric_limits::infinity(); + max1 = max2 = -min1; + for (int i = 0; i < n; i++) + { + if (x[s+i] < min1) { min2 = min1; min1 = x[s+i]; } + else if (x[s+i] < min2) min2 = x[s+i]; + if (x[s+i] > max1) { max2 = max1; max1 = x[s+i]; } + else if (x[s+i] > max2) max2 = x[s+i]; + } + mid = (min2+max2)/2; + for (int i = 0; i < n; i++) + if (x[s+i] < max2 && x[s+i] > min2) mid = x[s+i]; + } +// std::cout << "x = "; +// for (int i = 0; i < n; i++) +// std::cout << x[s+i] << " "; +// std::cout << "\nmid = " << mid << " n = " << n << "\n"; + return mid; + } + + /** Traverse tree to set node statistics */ + void visitToSetStatistics() + { + if (children_.size() > 0) + { + BOOST_FOREACH(KDNode* child, children_) + { + child->visitToSetStatistics(); + } + nodeStatistics_.InitFromChildren(children_[0]->nodeStatistics_, children_[1]->nodeStatistics_); + } + else + { + nodeStatistics_.InitAtLeaf(points_, *pointStatistics_, *oldFromNewIndex_, dfsIndex_, n_points_); + for (int i = dfsIndex_; i < dfsIndex_+n_points_; i++) + nodeFromOldIndex_->at(oldFromNewIndex(i)) = this; + } + } + + /** Traverse up the tree reset node statistics after a change at the leaves */ + void resetStatistics() + { + if (children_.empty()) + { + nodeStatistics_.ResetAtLeaf(points_, *pointStatistics_, *oldFromNewIndex_, dfsIndex_, n_points_); + } + else + { + nodeStatistics_.ResetFromChildren(children_[0]->nodeStatistics_, children_[1]->nodeStatistics_); + } + if (parent_) + parent_->resetStatistics(); + } + + double findMidPoint1(int& dim, int& n_left, int& n_right) + { + for (int k = 0; k < points_.n_rows(); k++) + { + n_left = n_right = 0; + double min = std::numeric_limits::infinity(); + double max = -std::numeric_limits::infinity(); + for (int i = 0; i < n_points_; i++) + { + int newIndex = i+dfsIndex_; + int oldIndex = oldFromNewIndex_->at(newIndex); + double val = points_.get(dim, oldIndex); + if (val > max) max = val; + if (val < min) min = val; + } + double mid = (min+max)/2; + + for (int i = 0; i < n_points_; i++) + { + int newIndex = i+dfsIndex_; + int oldIndex = oldFromNewIndex_->at(newIndex); + double val = points_.get(dim, oldIndex); + std::cout << "val = " << val << "\n"; + if (val < mid) n_left++; + else n_right++; + } + + std::cout << "dim = " << dim << " (min, mid, max) = " << min << ", " << mid << ", " << max + << " n_left = " << n_left << " n_right = " << n_right << "\n"; + if (n_left >= 1 && n_right >= 1) return mid; + else + dim = (dim+1) % points_.n_rows(); + } + n_left = 0; + n_right = 0; + return std::numeric_limits::quiet_NaN(); + } +}; + +class KDTreeDistanceMatrix +{ + const Matrix &reference_, &query_; + std::vector price_; + KDNode *referenceRoot_, *queryRoot_; +public: + KDTreeDistanceMatrix(const Matrix &reference, const Matrix &query) + : reference_(reference), query_(query), price_(query.n_cols(), 0) + { + DEBUG_ASSERT(reference.n_rows() == query.n_rows()); + referenceRoot_ = new KDNode(reference_, price_); + queryRoot_ = new KDNode(query_, price_); + } + int n_rows() const { return reference_.n_cols(); } + int n_cols() const { return query_.n_cols(); } + double get(int i, int j) const + { + Vector r_i, q_j; + reference_.MakeColumnVector(referenceRoot_->oldFromNewIndex(i), &r_i); + query_.MakeColumnVector(queryRoot_->oldFromNewIndex(j), &q_j); + return -sqrt(la::DistanceSqEuclidean(r_i, q_j)); + } + void setPrice(int j, const PointStatistics& price) + { + PointStatistics pStats = queryRoot_->pointStats(j); + DEBUG_ASSERT(pStats.matchTo() != price.matchTo()); + queryRoot_->setPointStatistics(j, price); + if (pStats.matchTo() != -1) + referenceRoot_->setPointStatistics(pStats.matchTo(), PointStatistics(0, -1)); + referenceRoot_->setPointStatistics(price.matchTo(), PointStatistics(0, j)); + } + long int getBestAndSecondBest(int bidder, int &best_item, double &best_surplus, double &second_surplus) + { + Vector r_i; + reference_.MakeColumnVector(referenceRoot_->oldFromNewIndex(bidder), &r_i); + + int minIndex, sndIndex; + double min, sndMin; + queryRoot_->randomBound(r_i, minIndex, min, sndIndex, sndMin); + int pruned = queryRoot_->nearestNeighbor(r_i, minIndex, min, sndIndex, sndMin); + best_item = minIndex; + best_surplus = -min; + second_surplus = -sndMin; +// std::cout << "bidder = " << bidder << " best = " << best_item +// << " best surplus = " << best_surplus << " second surplus = " << second_surplus << "\n"; + return pruned; + } + bool allMatched() const + { + return referenceRoot_->allMatched(); + } + + long int getAllBestAndSecondBest(std::vector& best_item, std::vector& best_surplus, std::vector& second_surplus) + { + int pruned = 0; + for (int bidder = 0; bidder < n_rows(); bidder++) if (referenceRoot_->pointStats(bidder).matchTo() == -1) + { + pruned += getBestAndSecondBest(bidder, best_item[bidder], best_surplus[bidder], second_surplus[bidder]); + } + return pruned; + } +}; + +END_ANMF_NAMESPACE; + +#endif // KDTREE_DISTANCE_MATRIX_H