From 0c5d7b7cefc33d610f5b9551e3b970710f952303 Mon Sep 17 00:00:00 2001 From: Lutetium-Vanadium Date: Mon, 29 May 2023 19:41:33 +0800 Subject: [PATCH 1/2] Add TugOfWar implementation in PyTorch --- benchmark/torchscripts/TugOfWar.pt | Bin 0 -> 8523 bytes benchmark/torchscripts/TugOfWar.py | 69 +++++++++++++++++++++++++++++ test/CMakeLists.txt | 1 + test/SystemTest/TugOfWar.cpp | 56 +++++++++++++++++++++++ test/scripts/config_tugOfWar.csv | 6 +++ test/torchscripts/TugOfWar.pt | Bin 0 -> 8523 bytes test/torchscripts/TugOfWar.py | 69 +++++++++++++++++++++++++++++ 7 files changed, 201 insertions(+) create mode 100644 benchmark/torchscripts/TugOfWar.pt create mode 100644 benchmark/torchscripts/TugOfWar.py create mode 100644 test/SystemTest/TugOfWar.cpp create mode 100644 test/scripts/config_tugOfWar.csv create mode 100644 test/torchscripts/TugOfWar.pt create mode 100644 test/torchscripts/TugOfWar.py diff --git a/benchmark/torchscripts/TugOfWar.pt b/benchmark/torchscripts/TugOfWar.pt new file mode 100644 index 0000000000000000000000000000000000000000..6d252c08c709ed88e42d2d2ffb6d192b3440be47 GIT binary patch literal 8523 zcmbt(1yEeew)P+afeb)7y7u1FtG`}T`&-qk*Qc(Gf(ig&U;zH55da7P+8(xQHhLCtZYv9S z3oaK&sN7Q=fZ^{g2pM8#=I#uK*qNDOszWUx)^^TND{Hu_vy}(b8q*cP8XN%9{{`<3 zx3GuV!)#rVB0Vf!k-5~B(J_Pp7CWLp$9{G7pPmr?^yFV%p#Vt!!xM6VrsswFfNxYi<>RTfE3Q61 z+Pn&1k-wic;te?IYk^pY?OKXc#WCx!7e^=LWi{TTEPxf$%;|03&Gx&p40;)~s1Bs7 zyH7z1?jhqe3PcZH?Pt?zjI$!^L%y!h&WPE-!yk8H9!3_#?dC<)W#|o!Xzs<yWIzT0Bp&=X zo&o{b|AnVqR@Rmtwq}0>5yqd4jnh$vQqE9Zs2nh7ytThlxl&>d073?8Q#7wIRJnxC zJnm`e8moQPogJ7jfH~9E##~cD8V4Ci{%iNy_xEEx_er1CK}4eWd&J4ZyH%#crlZz%dR!84`ZU!d z?%evP6tbdLirU=~mP2Rwl34x8hOFt`$CvNNH<&Gy*=Gl?1&r2EQyDH>}(?^k4Y=i^jdh*^~r@3)D_RE)Lku8V! zZg~PdX`IogM|rKUf9Nz)>9K#Q$W#&*VV`}LTiQEA!S|V@J_grYr3)nJUig-9fY`BF zOORH`&8DYjUtyhN5y&Oat5L?~7%v{(_ci~M_~&Rk37S^<%4a2agA&h{9l3QRss-I$ z_XUO<`^(ho7zZa?WTT)hB2X*RAGiqa%B-a;haOKTHr-@&jjFer+)J)YB+yl7 zYYt!*_HeC_eKm697&w1vZ?z>b`<}9v_$JVZO~*sU(KXu-o(#iVb&GLOnIUH1SlSEg zwt*&R3q4x4iIy{ZqTX*~{I(ry)QdcOM`$2`vS%?C5s)oZACGwWp~TnCm0$95>Q--sPIUi#~Idt{IYp4ImbC4rrKek*J00kdjFGut>j^&M?9T{t^#IQz~NFHXMZE*T?6*FfaMmQx(Wa5Cd= zTGu{bb))gveqs#iNCVbSUH4D8;K;zle5`K#V9s$IiE%7UXqeT49%22YRW4W# z7=HB4r`G8>%(0dl%^dvzTyFAx0mtu~A2)V4zT{zVcY)2O056j#y!&sRX2^CWLmWS zcXVUh@sYu1hRbZS9}*GO%G(kzFc~dG63v@AT-vDu8Io*k8Y*M4(Tf-vywZEjm?lGO z5|Y5`zNXt;F7xaul^D&ua78QHEgrTM`xm#dFL$IY)S{Q2aCI{5Fq-*K#3K?$dZmOr z2id7iUk!@w=D4 z;f6}f9Y!wS3%Bib)61EjtbZUlglwB61?2XmS)VX>-^TiPZD7^8x#RM{=R&y|#3RM{ zIrL?_k(Ve*EYXE;q?mK*#d`;V5oROX=r4L=#YZ<1_D$-Y7mYgsSoM%X^`xbf^6#p6A$zY=WY5! zcPA9`SZnF{30RvRWkxq(LG?-XH@Saph!HNM^z~=}KqlUQ+Ys3S0{>}4>1Y1-i5h1EU{|Lou4g8=`G2?NLOO7!I~ex%8yrrDzwP-Iu^Ry!6V{R8#uwahDpnN=jP@0 z(m$ApBZz-g#SU84J-=atpMnZ~0!^Qf0J&9Hh4p3+m_{!4xrZdZ(8wm1OczpW97!UiBUpQiO$Mx;cItHrM(j;rOXezZHrA*RzFT z@v?s?pSPSCe8>|9#`)^AB5+S$A`H}hU1b4GpCz6?tAEoWlZQU|4EvJO#~%&&8Wy2b zXpJ$!aMKe_*>~nSPuR{m{U>Za@T@wdUNh z@$1X7(Xw*lzPSc(-yB|sFL@%Pjboc));@l^`-WC{*~`$@SKZ@?6mN&wXElxDr0#-8 z_bbCBi-7);6F-jsM;C(6;6v{qvCf!K=$KI0u-RR|O2QU1HOlJ@ zS~Q_~97W1&nnNzF>rf-;ve(jSxMl5JlvrM@?q#x3N*XDrmukK;m!7AE9rp#nC+$+a z%BAmqww%ObjDmuf74zUK7`(+#=(IJP_&Mxpd+cYzFZMW%tpb!%UJ6{<%p^;tmyIS3 zOyV7h-}v?pl9e2)w(c@*?p|uNYJJ`YHbm02d;3?Rcph7b37snXe`-4H6Y>A7^PQ_2 zR1%(Xk=Aq&>~o021LIRl=#d+^-QTyduB<3yM~(kNX$QT8h<_bDDMfO3d?2c( z(wPU5jkr@-C7q0vaRixUMm^U_aBZ}MVR9Z%A)fz=YL*i9ZABCSAR6nxMYSM+=>HPc zzvsR9{}SJoA_h>Zfz)ZmQfo01&g5kOMvxKJa1VV%p@I)osQ_0~vlNt3P*_Lq`&LA& z$G)`qXoFzknb1YSV@i=H*V!wm^L@yzF@0zYx#qPo58!mtAwVfYx*r&XbW({ZB0&Jf zM3g~49SMFDewjs=WlgoqeZ8ZZaxcAOMzMvBs;a7meJpny-shi&H+OCjePQgUZa?}m zPUpCMOjyzOxSY{jQ4$Eb_nWr#wBCN5?QS@HhW9eSmaZ~w^vhWr<#*^I@n(ME=BWs;!^OIt$;J>5>l*y^J6hG{!yzlkDg{#Kf$!AHi?6x?_#E% zW%@%HHa2-_@Agj%2u_k=^GxYJT4o#X$I3ra`13iXR~n-Xy$gi65wIfJx-zcw9GS{J z@~?i;iC~>ATgHtWcZYVlgz3c!J}OKQ$?}{ywvZcVEI{Cg)ljd%4|EP4O&|O5bxnc7 zYqT97m8h;eEViMZ9ztTGc$<-{J`D@Drn3(g8N7AUFSiT}6Hp<0u?#BGNXiaZ+}Z83 zLoAK$KB|~w&)Y4~dQczAmg0%=;w<~S;*w%uNbs(G^j+2sp$#tIRokPG9_`m29sRcr z-$FZ$8*%b)_(rkJ%u9ut93GFip+K3tNYBXhcO95eJyJ*E8Q10*bbJTT+9;j@uKXkI zd2`913j!j1(m8?TZ8>n(XxZGMSc}Nbo

&8LOTu}VUjLBArfZSM3aFM~sIgctGM;_R%zKQ;@k;>maK1K|7Csi-GB8|KW zsd+IHW3#Z$#M(3#nn{ey+}Hp!e6>q$)oiFAx|#SCthg@f`A|=9&b)%Fi>{BqoIGj* z0XG5E5)w(D`p$20wr$Nk%%tbJfs87|x_At_;HM$S1#3Oq@7eUiKUtv@KQJVqiI2In z)nCU1m2RbV4#$rsk>4_KQ`XGBHU(kSA%kb=g;p6JUa2qCa@bPqF)@~5Hyi>R>Y?a)LIzD=&)ba@N5sY zOgS%5Elt!4VClSRd>9hN+U2RkM$n+W^C`Qmd;bmRtJHZ((_IL>E5T&smkzG=+ zqeXoly672Hyi^@QX-R5`dfWmXbWd`qc~up{#j~zsr1DrLn77J@WkLIs1Dk`bTZ5(& z7!p^sE#k>6{Z)W6<>LKIV@*v_gmN>}TIGXPB~b*Wjb1VmV+_2?j$B&(=%ID!IJ03C zI`i=40owR1d6-#+<@}@0&VGjB%52#>f)smxDX$&?h5c4>6DNKO{yf)ZK7B>M1d+mU zz(hOCbOL+$d3>FG<`Pzo@Ibq)?oeY|Z%5(gv@SbCQa`oS%2;!iD4aLnER+kTqUfxi z&rTIc7f(m$v(TVg+X6Uw1!jfkY`iC7em!n4d@9=Dn)$d&z}3jAtMV~n2v)_zFEEG> zTEo(Ogjz`!8Wq{mrl?|0xb89+eEAB<-$>M!6DV**G__`6ipYYJ?M#LTuTWE1M^h_~ z=ZLt;Juc@PTYoQbBp|jD^0>2crvPUyC|m!_kE>O|I}Sg&0R7zYwoeOESU)fmDak_M3f%qnD=%JV<8n(A|>DHuqY3b;6?>5NXY&qWl0h&tvw9w7gGa*3? zi37T5w{VGB$t5rvk8LxgVX}XZijlJ?2Isa9n)d1pYDkMB%Tb&bzCnDa#=*T$JyYhx ze=t9bOthncQnoEl>0s8aDa4d%id<>!Y` zBecacUSkdcsnnvnO(55n$ijq^;)~yuL1n~@gE3Jj=lntb2%#HQAFeR!`8G>x1Bswf zE`zo(z9OHwrHf6_bHjv$mheg6t+tU=iWjXu)5K49EYd=^3iErpE4Jkr8=Rm|-ZC{~ zw-p1fHG2f(Wk|4VAHSTDYv{~u9^Ye9)nKk^uplQSUP$sgN8?a-K8ayhOMQC$fXFRH zB3pbs=uoiJjUMI11(SsIg=Cp?ty2#+qshF%m(5N*xn)Z|BRL>Z_6Lb9);VQONgqql zvU7>5S6;+YG=HWwy}SIzVU3>Q+mFiaTp63-f!??vr3h z$JdYbfwiG!FY+n3<>Ykpy=Qg#T-13tVVG%7;t0BuV{C%%*OYx{YNqH9>){fpTJf~G%5i!nE-4cuqGa1v6mH^K+fy^!#_YrH=sZ_7a3xw3 zEqKtJ{i?m$CwUP^bT-4Q%Z>dWV*zoQWT82o`JFPuVLdF9x|>-x^PMjSJg&XgQz#o{A600f91D<~Dp#4$}saTNFu#IN@2#^w=G3DyH&G zW3+n!8OCHA>fTJ&Alx2gO{d01JaAWnCXU zB1TMbST^s+_H}>AUb1X+HMYwQ5HEiw4c~iCSG6ef0q`y=GLS^NIRb|4%~o)qV6GaH zLU3CUeVeRwZH;=LA905jq(yM+0na=%!K|`Pdn)pdo%c+HxH-MIy<+o3aXB;EXsw9j z>Pd>slk&|wmZ)o*%|qG(^iDxF`LhU&4%f(z01(5N{8@7=QPz!g-%f>b@B(TT3ThzJIc+f^GVjoRAcay(DifEVe>sKe7siT(*V=HK=|DA#b}$Z8 zqplplCa9N}Q*Qvtk#Z~1%#u}>v?UF>Qaqepjs&+yus+$20?Q9&lXs|Hb5~*70tFwW zzxCXnZ3Sf~?5zewlPyX<@H5YSdd(A(dr*uO!3d`&6bZ@wQhbja>I@dBtfSDWHAwZO z>f*kfIzql+0vWMgJh~<%xiag;p?V7mG78!(!2P;V9VuMMof3N=qo zA+!FWT!{Zag7YV^;yKd1kTEd9?ozgJ9wf5GAVS#kN# zIR7Z0{4?V3)f)D{Aj)9=N5p?ucl~GJ-x(tK7hrPS{}b>(`~HuQ@;hBf{<6d!^TFAHO~LS^EY6(0^XP?Ehx_KmKy$i~s-t literal 0 HcmV?d00001 diff --git a/benchmark/torchscripts/TugOfWar.py b/benchmark/torchscripts/TugOfWar.py new file mode 100644 index 00000000..2c1c8928 --- /dev/null +++ b/benchmark/torchscripts/TugOfWar.py @@ -0,0 +1,69 @@ +import torch +import time +import os +import math + +def tug_of_war_mat(m: int, n: int) -> torch.Tensor: + e = 1/math.sqrt(m) + M = torch.randint(2, (m, n)) + return e*(2*M - 1) + + +@torch.jit.script +def TugOfWar(A: torch.Tensor, B: torch.Tensor, l: int): + m, n = A.shape + n, p = B.shape + + delta = 0.2 + + i_iters = int(-math.log(delta)) + j_iters = int(2*(-math.log(delta) + math.log(-math.log(delta)))) + + z = torch.empty((i_iters,)) + AS = [] + SB = [] + + for i in range(i_iters): + S = tug_of_war_mat(l, n) + SB.append(S.matmul(B)) + AS.append(A.matmul(S.T)) + + y = torch.empty((j_iters,)) + + for j in range(j_iters): + Q = tug_of_war_mat(16, p) + X = A.matmul(B.matmul(Q.T)) + X_hat = AS[i].matmul(SB[i].matmul(Q.T)) + y[j] = torch.norm(X - X_hat)**2 + z[i] = torch.median(y) + + i_star = torch.argmin(z) + return torch.matmul(AS[i_star], SB[i_star]) + + +def main(): + width = 1000 + A = torch.rand(10000, width) + B = torch.rand(width, 5000) + + t = time.time() + + aResult = TugOfWar(A, B, 500) + print("approximate: " + str(time.time() - t) + "s") + + print(aResult) + + # exact result + t = time.time() + eResult = torch.matmul(A, B) + print("\nExact: " + str(time.time() - t) + "s") + + print(eResult) + + print("\nerror: " + str(torch.norm(aResult - eResult, p='fro').item())) + + TugOfWar_script = TugOfWar.save("TugOfWar.pt") + + +if __name__ == '__main__': + main() diff --git a/test/CMakeLists.txt b/test/CMakeLists.txt index d5433de2..8577f40c 100644 --- a/test/CMakeLists.txt +++ b/test/CMakeLists.txt @@ -23,4 +23,5 @@ add_catch_test(sketch_test SystemTest/SketchTest.cpp IntelliStream) add_catch_test(crs_test SystemTest/CRSTest.cpp IntelliStream) add_catch_test(weighted_cr_test SystemTest/WeightedCRTest.cpp IntelliStream) add_catch_test(block_partition_test SystemTest/BlockPartitionTest.cpp IntelliStream) +add_catch_test(tug_of_war_test SystemTest/TugOfWarTest.cpp IntelliStream) diff --git a/test/SystemTest/TugOfWar.cpp b/test/SystemTest/TugOfWar.cpp new file mode 100644 index 00000000..6e623525 --- /dev/null +++ b/test/SystemTest/TugOfWar.cpp @@ -0,0 +1,56 @@ +#include + +#define CATCH_CONFIG_MAIN +#include "catch.hpp" +#include +using namespace std; +using namespace INTELLI; +using namespace torch; +void runSingleThreadTest(std::string configName) { + ConfigMapPtr cfg = newConfigMap(); + cfg->fromFile(configName); + AMMBench::MatrixLoaderTable mLoaderTable; + uint64_t sketchDimension; + sketchDimension = cfg->tryU64("sketchDimension", 50, true); + uint64_t coreBind = cfg->tryU64("coreBind", 0, true); + UtilityFunctions::bind2Core((int) coreBind); + torch::set_num_threads(1); + std::string ptFile = cfg->tryString("ptFile", "torchscripts/FDAMM.pt", true); + + //uint64_t customResultName = cfg->tryU64("customResultName", 0, true); + INTELLI_INFO("Place me at core" + to_string(coreBind)); + INTELLI_INFO( + "with sketch" + to_string(sketchDimension)); + torch::jit::script::Module module; + INTELLI_INFO("Try pt file " + ptFile); + module = torch::jit::load(ptFile); + std::string matrixLoaderTag = cfg->tryString("matrixLoaderTag", "random", true); + auto matLoaderPtr = mLoaderTable.findMatrixLoader(matrixLoaderTag); + assert(matLoaderPtr); + matLoaderPtr->setConfig(cfg); + auto A = matLoaderPtr->getA(); + auto B = matLoaderPtr->getB(); + /*torch::manual_seed(114514); +//555 +auto A = torch::rand({(long) aRow, (long) aCol}); +auto B = torch::rand({(long) aCol, (long) bCol});*/ + INTELLI_INFO("Generation done, conducting..."); + ThreadPerf pef((int) coreBind); + pef.setPerfList(); + pef.start(); + auto C =module.forward({A, B, (long) sketchDimension}).toTensor(); + pef.end(); + std::string ruName = "default"; + + auto resultCsv = pef.resultToConfigMap(); + resultCsv->toFile(ruName + ".csv"); + INTELLI_INFO("Done. here is result"); + std::cout << resultCsv->toString() << endl; +} +TEST_CASE("Test the Tug of War", "[short]") +{ + int a = 0; + runSingleThreadTest("scripts/config_tugOfWar.csv"); + // place your test here + REQUIRE(a == 0); +} diff --git a/test/scripts/config_tugOfWar.csv b/test/scripts/config_tugOfWar.csv new file mode 100644 index 00000000..a63e79b2 --- /dev/null +++ b/test/scripts/config_tugOfWar.csv @@ -0,0 +1,6 @@ +key,value,type +aRow,100,U64 +aCol,1000,U64 +bCol,500,U64 +sketchDimension,25,U64 +ptFile,torchscripts/CRSV2.pt,String \ No newline at end of file diff --git a/test/torchscripts/TugOfWar.pt b/test/torchscripts/TugOfWar.pt new file mode 100644 index 0000000000000000000000000000000000000000..6d252c08c709ed88e42d2d2ffb6d192b3440be47 GIT binary patch literal 8523 zcmbt(1yEeew)P+afeb)7y7u1FtG`}T`&-qk*Qc(Gf(ig&U;zH55da7P+8(xQHhLCtZYv9S z3oaK&sN7Q=fZ^{g2pM8#=I#uK*qNDOszWUx)^^TND{Hu_vy}(b8q*cP8XN%9{{`<3 zx3GuV!)#rVB0Vf!k-5~B(J_Pp7CWLp$9{G7pPmr?^yFV%p#Vt!!xM6VrsswFfNxYi<>RTfE3Q61 z+Pn&1k-wic;te?IYk^pY?OKXc#WCx!7e^=LWi{TTEPxf$%;|03&Gx&p40;)~s1Bs7 zyH7z1?jhqe3PcZH?Pt?zjI$!^L%y!h&WPE-!yk8H9!3_#?dC<)W#|o!Xzs<yWIzT0Bp&=X zo&o{b|AnVqR@Rmtwq}0>5yqd4jnh$vQqE9Zs2nh7ytThlxl&>d073?8Q#7wIRJnxC zJnm`e8moQPogJ7jfH~9E##~cD8V4Ci{%iNy_xEEx_er1CK}4eWd&J4ZyH%#crlZz%dR!84`ZU!d z?%evP6tbdLirU=~mP2Rwl34x8hOFt`$CvNNH<&Gy*=Gl?1&r2EQyDH>}(?^k4Y=i^jdh*^~r@3)D_RE)Lku8V! zZg~PdX`IogM|rKUf9Nz)>9K#Q$W#&*VV`}LTiQEA!S|V@J_grYr3)nJUig-9fY`BF zOORH`&8DYjUtyhN5y&Oat5L?~7%v{(_ci~M_~&Rk37S^<%4a2agA&h{9l3QRss-I$ z_XUO<`^(ho7zZa?WTT)hB2X*RAGiqa%B-a;haOKTHr-@&jjFer+)J)YB+yl7 zYYt!*_HeC_eKm697&w1vZ?z>b`<}9v_$JVZO~*sU(KXu-o(#iVb&GLOnIUH1SlSEg zwt*&R3q4x4iIy{ZqTX*~{I(ry)QdcOM`$2`vS%?C5s)oZACGwWp~TnCm0$95>Q--sPIUi#~Idt{IYp4ImbC4rrKek*J00kdjFGut>j^&M?9T{t^#IQz~NFHXMZE*T?6*FfaMmQx(Wa5Cd= zTGu{bb))gveqs#iNCVbSUH4D8;K;zle5`K#V9s$IiE%7UXqeT49%22YRW4W# z7=HB4r`G8>%(0dl%^dvzTyFAx0mtu~A2)V4zT{zVcY)2O056j#y!&sRX2^CWLmWS zcXVUh@sYu1hRbZS9}*GO%G(kzFc~dG63v@AT-vDu8Io*k8Y*M4(Tf-vywZEjm?lGO z5|Y5`zNXt;F7xaul^D&ua78QHEgrTM`xm#dFL$IY)S{Q2aCI{5Fq-*K#3K?$dZmOr z2id7iUk!@w=D4 z;f6}f9Y!wS3%Bib)61EjtbZUlglwB61?2XmS)VX>-^TiPZD7^8x#RM{=R&y|#3RM{ zIrL?_k(Ve*EYXE;q?mK*#d`;V5oROX=r4L=#YZ<1_D$-Y7mYgsSoM%X^`xbf^6#p6A$zY=WY5! zcPA9`SZnF{30RvRWkxq(LG?-XH@Saph!HNM^z~=}KqlUQ+Ys3S0{>}4>1Y1-i5h1EU{|Lou4g8=`G2?NLOO7!I~ex%8yrrDzwP-Iu^Ry!6V{R8#uwahDpnN=jP@0 z(m$ApBZz-g#SU84J-=atpMnZ~0!^Qf0J&9Hh4p3+m_{!4xrZdZ(8wm1OczpW97!UiBUpQiO$Mx;cItHrM(j;rOXezZHrA*RzFT z@v?s?pSPSCe8>|9#`)^AB5+S$A`H}hU1b4GpCz6?tAEoWlZQU|4EvJO#~%&&8Wy2b zXpJ$!aMKe_*>~nSPuR{m{U>Za@T@wdUNh z@$1X7(Xw*lzPSc(-yB|sFL@%Pjboc));@l^`-WC{*~`$@SKZ@?6mN&wXElxDr0#-8 z_bbCBi-7);6F-jsM;C(6;6v{qvCf!K=$KI0u-RR|O2QU1HOlJ@ zS~Q_~97W1&nnNzF>rf-;ve(jSxMl5JlvrM@?q#x3N*XDrmukK;m!7AE9rp#nC+$+a z%BAmqww%ObjDmuf74zUK7`(+#=(IJP_&Mxpd+cYzFZMW%tpb!%UJ6{<%p^;tmyIS3 zOyV7h-}v?pl9e2)w(c@*?p|uNYJJ`YHbm02d;3?Rcph7b37snXe`-4H6Y>A7^PQ_2 zR1%(Xk=Aq&>~o021LIRl=#d+^-QTyduB<3yM~(kNX$QT8h<_bDDMfO3d?2c( z(wPU5jkr@-C7q0vaRixUMm^U_aBZ}MVR9Z%A)fz=YL*i9ZABCSAR6nxMYSM+=>HPc zzvsR9{}SJoA_h>Zfz)ZmQfo01&g5kOMvxKJa1VV%p@I)osQ_0~vlNt3P*_Lq`&LA& z$G)`qXoFzknb1YSV@i=H*V!wm^L@yzF@0zYx#qPo58!mtAwVfYx*r&XbW({ZB0&Jf zM3g~49SMFDewjs=WlgoqeZ8ZZaxcAOMzMvBs;a7meJpny-shi&H+OCjePQgUZa?}m zPUpCMOjyzOxSY{jQ4$Eb_nWr#wBCN5?QS@HhW9eSmaZ~w^vhWr<#*^I@n(ME=BWs;!^OIt$;J>5>l*y^J6hG{!yzlkDg{#Kf$!AHi?6x?_#E% zW%@%HHa2-_@Agj%2u_k=^GxYJT4o#X$I3ra`13iXR~n-Xy$gi65wIfJx-zcw9GS{J z@~?i;iC~>ATgHtWcZYVlgz3c!J}OKQ$?}{ywvZcVEI{Cg)ljd%4|EP4O&|O5bxnc7 zYqT97m8h;eEViMZ9ztTGc$<-{J`D@Drn3(g8N7AUFSiT}6Hp<0u?#BGNXiaZ+}Z83 zLoAK$KB|~w&)Y4~dQczAmg0%=;w<~S;*w%uNbs(G^j+2sp$#tIRokPG9_`m29sRcr z-$FZ$8*%b)_(rkJ%u9ut93GFip+K3tNYBXhcO95eJyJ*E8Q10*bbJTT+9;j@uKXkI zd2`913j!j1(m8?TZ8>n(XxZGMSc}Nbo

&8LOTu}VUjLBArfZSM3aFM~sIgctGM;_R%zKQ;@k;>maK1K|7Csi-GB8|KW zsd+IHW3#Z$#M(3#nn{ey+}Hp!e6>q$)oiFAx|#SCthg@f`A|=9&b)%Fi>{BqoIGj* z0XG5E5)w(D`p$20wr$Nk%%tbJfs87|x_At_;HM$S1#3Oq@7eUiKUtv@KQJVqiI2In z)nCU1m2RbV4#$rsk>4_KQ`XGBHU(kSA%kb=g;p6JUa2qCa@bPqF)@~5Hyi>R>Y?a)LIzD=&)ba@N5sY zOgS%5Elt!4VClSRd>9hN+U2RkM$n+W^C`Qmd;bmRtJHZ((_IL>E5T&smkzG=+ zqeXoly672Hyi^@QX-R5`dfWmXbWd`qc~up{#j~zsr1DrLn77J@WkLIs1Dk`bTZ5(& z7!p^sE#k>6{Z)W6<>LKIV@*v_gmN>}TIGXPB~b*Wjb1VmV+_2?j$B&(=%ID!IJ03C zI`i=40owR1d6-#+<@}@0&VGjB%52#>f)smxDX$&?h5c4>6DNKO{yf)ZK7B>M1d+mU zz(hOCbOL+$d3>FG<`Pzo@Ibq)?oeY|Z%5(gv@SbCQa`oS%2;!iD4aLnER+kTqUfxi z&rTIc7f(m$v(TVg+X6Uw1!jfkY`iC7em!n4d@9=Dn)$d&z}3jAtMV~n2v)_zFEEG> zTEo(Ogjz`!8Wq{mrl?|0xb89+eEAB<-$>M!6DV**G__`6ipYYJ?M#LTuTWE1M^h_~ z=ZLt;Juc@PTYoQbBp|jD^0>2crvPUyC|m!_kE>O|I}Sg&0R7zYwoeOESU)fmDak_M3f%qnD=%JV<8n(A|>DHuqY3b;6?>5NXY&qWl0h&tvw9w7gGa*3? zi37T5w{VGB$t5rvk8LxgVX}XZijlJ?2Isa9n)d1pYDkMB%Tb&bzCnDa#=*T$JyYhx ze=t9bOthncQnoEl>0s8aDa4d%id<>!Y` zBecacUSkdcsnnvnO(55n$ijq^;)~yuL1n~@gE3Jj=lntb2%#HQAFeR!`8G>x1Bswf zE`zo(z9OHwrHf6_bHjv$mheg6t+tU=iWjXu)5K49EYd=^3iErpE4Jkr8=Rm|-ZC{~ zw-p1fHG2f(Wk|4VAHSTDYv{~u9^Ye9)nKk^uplQSUP$sgN8?a-K8ayhOMQC$fXFRH zB3pbs=uoiJjUMI11(SsIg=Cp?ty2#+qshF%m(5N*xn)Z|BRL>Z_6Lb9);VQONgqql zvU7>5S6;+YG=HWwy}SIzVU3>Q+mFiaTp63-f!??vr3h z$JdYbfwiG!FY+n3<>Ykpy=Qg#T-13tVVG%7;t0BuV{C%%*OYx{YNqH9>){fpTJf~G%5i!nE-4cuqGa1v6mH^K+fy^!#_YrH=sZ_7a3xw3 zEqKtJ{i?m$CwUP^bT-4Q%Z>dWV*zoQWT82o`JFPuVLdF9x|>-x^PMjSJg&XgQz#o{A600f91D<~Dp#4$}saTNFu#IN@2#^w=G3DyH&G zW3+n!8OCHA>fTJ&Alx2gO{d01JaAWnCXU zB1TMbST^s+_H}>AUb1X+HMYwQ5HEiw4c~iCSG6ef0q`y=GLS^NIRb|4%~o)qV6GaH zLU3CUeVeRwZH;=LA905jq(yM+0na=%!K|`Pdn)pdo%c+HxH-MIy<+o3aXB;EXsw9j z>Pd>slk&|wmZ)o*%|qG(^iDxF`LhU&4%f(z01(5N{8@7=QPz!g-%f>b@B(TT3ThzJIc+f^GVjoRAcay(DifEVe>sKe7siT(*V=HK=|DA#b}$Z8 zqplplCa9N}Q*Qvtk#Z~1%#u}>v?UF>Qaqepjs&+yus+$20?Q9&lXs|Hb5~*70tFwW zzxCXnZ3Sf~?5zewlPyX<@H5YSdd(A(dr*uO!3d`&6bZ@wQhbja>I@dBtfSDWHAwZO z>f*kfIzql+0vWMgJh~<%xiag;p?V7mG78!(!2P;V9VuMMof3N=qo zA+!FWT!{Zag7YV^;yKd1kTEd9?ozgJ9wf5GAVS#kN# zIR7Z0{4?V3)f)D{Aj)9=N5p?ucl~GJ-x(tK7hrPS{}b>(`~HuQ@;hBf{<6d!^TFAHO~LS^EY6(0^XP?Ehx_KmKy$i~s-t literal 0 HcmV?d00001 diff --git a/test/torchscripts/TugOfWar.py b/test/torchscripts/TugOfWar.py new file mode 100644 index 00000000..2c1c8928 --- /dev/null +++ b/test/torchscripts/TugOfWar.py @@ -0,0 +1,69 @@ +import torch +import time +import os +import math + +def tug_of_war_mat(m: int, n: int) -> torch.Tensor: + e = 1/math.sqrt(m) + M = torch.randint(2, (m, n)) + return e*(2*M - 1) + + +@torch.jit.script +def TugOfWar(A: torch.Tensor, B: torch.Tensor, l: int): + m, n = A.shape + n, p = B.shape + + delta = 0.2 + + i_iters = int(-math.log(delta)) + j_iters = int(2*(-math.log(delta) + math.log(-math.log(delta)))) + + z = torch.empty((i_iters,)) + AS = [] + SB = [] + + for i in range(i_iters): + S = tug_of_war_mat(l, n) + SB.append(S.matmul(B)) + AS.append(A.matmul(S.T)) + + y = torch.empty((j_iters,)) + + for j in range(j_iters): + Q = tug_of_war_mat(16, p) + X = A.matmul(B.matmul(Q.T)) + X_hat = AS[i].matmul(SB[i].matmul(Q.T)) + y[j] = torch.norm(X - X_hat)**2 + z[i] = torch.median(y) + + i_star = torch.argmin(z) + return torch.matmul(AS[i_star], SB[i_star]) + + +def main(): + width = 1000 + A = torch.rand(10000, width) + B = torch.rand(width, 5000) + + t = time.time() + + aResult = TugOfWar(A, B, 500) + print("approximate: " + str(time.time() - t) + "s") + + print(aResult) + + # exact result + t = time.time() + eResult = torch.matmul(A, B) + print("\nExact: " + str(time.time() - t) + "s") + + print(eResult) + + print("\nerror: " + str(torch.norm(aResult - eResult, p='fro').item())) + + TugOfWar_script = TugOfWar.save("TugOfWar.pt") + + +if __name__ == '__main__': + main() From 27d1c4719d9b5bf78d2b86ba5a52400d5cd7916f Mon Sep 17 00:00:00 2001 From: Lutetium-Vanadium Date: Tue, 30 May 2023 13:19:36 +0800 Subject: [PATCH 2/2] Fix failing test --- test/SystemTest/{TugOfWar.cpp => TugOfWarTest.cpp} | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename test/SystemTest/{TugOfWar.cpp => TugOfWarTest.cpp} (100%) diff --git a/test/SystemTest/TugOfWar.cpp b/test/SystemTest/TugOfWarTest.cpp similarity index 100% rename from test/SystemTest/TugOfWar.cpp rename to test/SystemTest/TugOfWarTest.cpp