Skip to content
Merged
Show file tree
Hide file tree
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
1 change: 1 addition & 0 deletions include/AMMBench.h
Original file line number Diff line number Diff line change
Expand Up @@ -105,6 +105,7 @@
#include <CPPAlgos/EWSCPPAlgo.h>
#include <CPPAlgos/CoOccurringFDCPPAlgo.h>
#include <CPPAlgos/BetaCoOFDCPPAlgo.h>
#include <CPPAlgos/TugOfWarCPPAlgo.h>
/**
* @}
*
Expand Down
67 changes: 67 additions & 0 deletions include/CPPAlgos/TugOfWarCPPAlgo.h
Original file line number Diff line number Diff line change
@@ -0,0 +1,67 @@
//
// Created by luv on 5/30/23.
//

#ifndef INTELLISTREAM_TUGOFWARCPPALGO_H
#define INTELLISTREAM_TUGOFWARCPPALGO_H

#include <CPPAlgos/AbstractCPPAlgo.h>

namespace AMMBench {
/**
* @ingroup AMMBENCH_CppAlgos The algorithms writtrn in c++
* @{
*/
/**
* @class CRSCPPlgo CPPAlgos/TugOfWarCPPAlgo.h
* @brief The cloumn row sampling (CRS) class of c++ algos
*
*/
class TugOfWarCPPAlgo : public AMMBench::AbstractCPPAlgo {
double delta = 0.2;

public:
TugOfWarCPPAlgo() {

}

TugOfWarCPPAlgo(double delta): delta(delta) {

}

~TugOfWarCPPAlgo() {

}

/**
* @brief the virtual function provided for outside callers, rewrite in children classes
* @param A the A matrix
* @param B the B matrix
* @param sketchSize the size of sketc or sampling
* @return the output c matrix
*/
virtual torch::Tensor amm(torch::Tensor A, torch::Tensor B, int sketchSize);

private:
torch::Tensor generateTugOfWarMatrix(int64_t m, int64_t n);

};

/**
* @ingroup AMMBENCH_CppAlgos
* @typedef AbstractMatrixCppAlgoPtr
* @brief The class to describe a shared pointer to @ref TugOfWarCppAlgo

*/
typedef std::shared_ptr<class AMMBench::TugOfWarCPPAlgo> TugOfWarCPPAlgoPtr;
/**
* @ingroup AMMBENCH_CppAlgos
* @def newTugOfWarCppAlgo
* @brief (Macro) To creat a new @ref TugOfWarCppAlgounder shared pointer.
*/
#define newTugOfWarCPPAlgo std::make_shared<AMMBench::TugOfWarCPPAlgo>
}
/**
* @}
*/
#endif //INTELLISTREAM_TUGOFWARCPPALGO_H
1 change: 1 addition & 0 deletions src/CPPAlgos/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -8,5 +8,6 @@ add_sources(
CoOccurringFDCPPAlgo.cpp
BetaCoOFDCPPAlgo.cpp
CountSketchCPPAlgo.cpp
TugOfWarCPPAlgo.cpp
)

2 changes: 2 additions & 0 deletions src/CPPAlgos/CPPAlgoTable.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
#include <CPPAlgos/EWSCPPAlgo.h>
#include <CPPAlgos/CoOccurringFDCPPAlgo.h>
#include <CPPAlgos/BetaCoOFDCPPAlgo.h>
#include <CPPAlgos/TugOfWarCPPAlgo.h>

namespace AMMBench {
AMMBench::CPPAlgoTable::CPPAlgoTable() {
Expand All @@ -23,6 +24,7 @@ AMMBench::CPPAlgoTable::CPPAlgoTable() {
algoMap["ews"] = newEWSCPPAlgo();
algoMap["CoOFD"] = newCoOccurringFDCPPAlgo();
algoMap["bcoofd"] = newBetaCoOFDCPPAlgo();
algoMap["tug-of-war"] = newTugOfWarCPPAlgo();
}

} // AMMBench
48 changes: 48 additions & 0 deletions src/CPPAlgos/TugOfWarCPPAlgo.cpp
Original file line number Diff line number Diff line change
@@ -0,0 +1,48 @@
//
// Created by luv on 5/30/23.
//

#include <CPPAlgos/TugOfWarCPPAlgo.h>

namespace AMMBench {
torch::Tensor TugOfWarCPPAlgo::generateTugOfWarMatrix(int64_t m, int64_t n) {
double e = 1.0 / std::sqrt(m);
torch::Tensor M = torch::randint(2, {m, n});
return e * (2 * M - 1);
}

torch::Tensor TugOfWarCPPAlgo::amm(torch::Tensor A, torch::Tensor B, int l) {
int n = A.size(1);
int p = B.size(1);

int i_iters = static_cast<int>(-std::log(delta));
int j_iters = static_cast<int>(2 * (-std::log(delta) + std::log(-std::log(delta))));

torch::Tensor z = torch::empty({i_iters});
std::vector<torch::Tensor> AS;
std::vector<torch::Tensor> SB;

for (int i = 0; i < i_iters; ++i) {
torch::Tensor S = generateTugOfWarMatrix(l, n);
SB.push_back(S.matmul(B));
AS.push_back(A.matmul(S.t()));

torch::Tensor y = torch::empty({j_iters});

for (int j = 0; j < j_iters; ++j) {
torch::Tensor Q = generateTugOfWarMatrix(16, p);
torch::Tensor X = A.matmul(B.matmul(Q.t()));
torch::Tensor X_hat = AS[i].matmul(SB[i].matmul(Q.t()));
y[j] = torch::norm(X - X_hat).pow(2);
}


z[i] = at::median(y);
}

torch::Tensor z_argmin = z.argmin();
int i_star = z_argmin.item<int>();

return AS[i_star].matmul(SB[i_star]);
}
} // AMMBench
11 changes: 11 additions & 0 deletions test/SystemTest/TugOfWarTest.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -54,3 +54,14 @@ TEST_CASE("Test the Tug of War", "[short]")
// place your test here
REQUIRE(a == 0);
}
TEST_CASE("Test Tug of War in cpp", "[short]")
{
torch::manual_seed(114514);
AMMBench::TugOfWarCPPAlgo tw;
auto A = torch::rand({400, 400});
auto B = torch::rand({400, 400});
auto realC = torch::matmul(A, B);
auto ammC = tw.amm(A, B, 20);
double froError = INTELLI::UtilityFunctions::relativeFrobeniusNorm(realC, ammC);
REQUIRE(froError < 0.5);
}
2 changes: 1 addition & 1 deletion test/scripts/config_tugOfWar.csv
Original file line number Diff line number Diff line change
Expand Up @@ -3,4 +3,4 @@ aRow,100,U64
aCol,1000,U64
bCol,500,U64
sketchDimension,25,U64
ptFile,torchscripts/CRSV2.pt,String
ptFile,torchscripts/TugOfWar.pt,String