From a3d960fc48e098a44b2d3613e27a18246733a04d Mon Sep 17 00:00:00 2001 From: tony <292224750@qq.com> Date: Fri, 19 May 2023 09:59:38 +0800 Subject: [PATCH 1/2] 1. add the energy meter subsystem 2. include haolan's crs update --- benchmark/src/Benchmark.cpp | 38 +++++ benchmark/torchscripts/BernoulliCRS.pt | Bin 0 -> 2723 bytes benchmark/torchscripts/BernoulliCRS.py | 66 ++++++++ benchmark/torchscripts/CRS.pt | Bin 0 -> 2614 bytes benchmark/torchscripts/CRSV2.pt | Bin 0 -> 3136 bytes benchmark/torchscripts/ColumnRowSampling.py | 70 +++++++++ .../torchscripts/ColumnRowSamplingVer2.py | 71 +++++++++ benchmark/torchscripts/E2E_Minist.py | 104 ++++++------- commit.sh | 2 +- commit_info | 4 +- include/AMMBench.h | 15 +- include/Utils/Meters/AbstractMeter.hpp | 125 +++++++++++++++ .../Meters/EspMeterUart/EspMeterUart.hpp | 70 +++++++++ .../Utils/Meters/IntelMeter/IntelMeter.hpp | 84 ++++++++++ include/Utils/Meters/MeterTable.h | 74 +++++++++ src/Utils/CMakeLists.txt | 3 +- src/Utils/Meters/AbstractMeter.cpp | 16 ++ src/Utils/Meters/CMakeLists.txt | 3 + src/Utils/Meters/EspMeterUart/CMakeLists.txt | 1 + .../Meters/EspMeterUart/EspMeterUart.cpp | 129 ++++++++++++++++ src/Utils/Meters/IntelMeter/CMakeLists.txt | 2 + src/Utils/Meters/IntelMeter/IntelMeter.cpp | 93 +++++++++++ src/Utils/Meters/MeterTable.cpp | 13 ++ test/CMakeLists.txt | 1 + test/SystemTest/CRSTest.cpp | 63 ++++++++ test/scripts/config_CRS.csv | 6 + test/scripts/config_CRSV2.csv | 6 + test/torchscripts/BernoulliCRS.pt | Bin 0 -> 2723 bytes test/torchscripts/BernoulliCRS.py | 66 ++++++++ test/torchscripts/BetaCoOccurringFD.pt | Bin 0 -> 7882 bytes test/torchscripts/BetaCoOccurringFD.py | 110 +++++++++++++ test/torchscripts/CRS.pt | Bin 0 -> 2614 bytes test/torchscripts/CRSV2.pt | Bin 0 -> 3136 bytes test/torchscripts/CoOccurringFD.pt | Bin 0 -> 7534 bytes test/torchscripts/CoOccurringFD.py | 118 ++++++++++++++ test/torchscripts/ColumnRowSampling.py | 70 +++++++++ test/torchscripts/ColumnRowSamplingVer2.py | 71 +++++++++ test/torchscripts/E2E_Minist.py | 146 ++++++++++++++++++ test/torchscripts/FDAMM.pt | Bin 4416 -> 4416 bytes test/torchscripts/FDAMM.py | 86 +++++++++++ test/torchscripts/RAWMM.pt | Bin 0 -> 1408 bytes 41 files changed, 1662 insertions(+), 64 deletions(-) create mode 100644 benchmark/torchscripts/BernoulliCRS.pt create mode 100644 benchmark/torchscripts/BernoulliCRS.py create mode 100644 benchmark/torchscripts/CRS.pt create mode 100644 benchmark/torchscripts/CRSV2.pt create mode 100644 benchmark/torchscripts/ColumnRowSampling.py create mode 100644 benchmark/torchscripts/ColumnRowSamplingVer2.py create mode 100644 include/Utils/Meters/AbstractMeter.hpp create mode 100644 include/Utils/Meters/EspMeterUart/EspMeterUart.hpp create mode 100644 include/Utils/Meters/IntelMeter/IntelMeter.hpp create mode 100644 include/Utils/Meters/MeterTable.h create mode 100644 src/Utils/Meters/AbstractMeter.cpp create mode 100644 src/Utils/Meters/CMakeLists.txt create mode 100644 src/Utils/Meters/EspMeterUart/CMakeLists.txt create mode 100644 src/Utils/Meters/EspMeterUart/EspMeterUart.cpp create mode 100644 src/Utils/Meters/IntelMeter/CMakeLists.txt create mode 100644 src/Utils/Meters/IntelMeter/IntelMeter.cpp create mode 100644 src/Utils/Meters/MeterTable.cpp create mode 100644 test/SystemTest/CRSTest.cpp create mode 100644 test/scripts/config_CRS.csv create mode 100644 test/scripts/config_CRSV2.csv create mode 100644 test/torchscripts/BernoulliCRS.pt create mode 100644 test/torchscripts/BernoulliCRS.py create mode 100644 test/torchscripts/BetaCoOccurringFD.pt create mode 100644 test/torchscripts/BetaCoOccurringFD.py create mode 100644 test/torchscripts/CRS.pt create mode 100644 test/torchscripts/CRSV2.pt create mode 100644 test/torchscripts/CoOccurringFD.pt create mode 100644 test/torchscripts/CoOccurringFD.py create mode 100644 test/torchscripts/ColumnRowSampling.py create mode 100644 test/torchscripts/ColumnRowSamplingVer2.py create mode 100644 test/torchscripts/E2E_Minist.py create mode 100644 test/torchscripts/FDAMM.py create mode 100644 test/torchscripts/RAWMM.pt diff --git a/benchmark/src/Benchmark.cpp b/benchmark/src/Benchmark.cpp index a47aabec..82093ee9 100644 --- a/benchmark/src/Benchmark.cpp +++ b/benchmark/src/Benchmark.cpp @@ -9,13 +9,36 @@ using namespace std; using namespace INTELLI; using namespace torch; +using namespace DIVERSE_METER; void runSingleThreadTest(std::string configName) { + MeterTable meterTable; + AbstractMeterPtr eMeter = nullptr; 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); + uint64_t usingMeter = cfg->tryU64("usingMeter", 0, true); + std::string meterTag = cfg->tryString("meterTag", "intelMsr", true); + if (usingMeter) { + eMeter = meterTable.findMeter(meterTag); + if (eMeter != nullptr) { + eMeter->setConfig(cfg); + double staticPower = cfg->tryDouble("staticPower", 0.0, false); + if (staticPower == 0.0) { + eMeter->testStaticPower(2); + } else { + INTELLI_INFO("use pre-defined static power"); + eMeter->setStaticPower(staticPower); + } + INTELLI_INFO("static power is " + to_string(eMeter->getStaticPower()) + " W"); + } else { + INTELLI_ERROR("No meter found: " + meterTag); + } + + } + UtilityFunctions::bind2Core((int) coreBind); torch::set_num_threads(1); std::string ptFile = cfg->tryString("ptFile", "torchscripts/FDAMM.pt", true); @@ -33,19 +56,34 @@ void runSingleThreadTest(std::string configName) { matLoaderPtr->setConfig(cfg); auto A = matLoaderPtr->getA(); auto B = matLoaderPtr->getB(); + //555 /*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..."); + if (eMeter != nullptr) { + eMeter->startMeter(); + } ThreadPerf pef((int) coreBind); pef.setPerfList(); pef.start(); auto C =module.forward({A, B, (long) sketchDimension}).toTensor(); pef.end(); + if (eMeter != nullptr) { + eMeter->stopMeter(); + } std::string ruName = "default"; auto resultCsv = pef.resultToConfigMap(); + if (eMeter != nullptr) { + eMeter->stopMeter(); + double energyConsumption = eMeter->getE(); + double staticEnergyConsumption = eMeter->getStaicEnergyConsumption(resultCsv->tryU64("perfElapsedTime", 0, false)); + double pureEnergy = energyConsumption - staticEnergyConsumption; + resultCsv->edit("energyAll", (double) energyConsumption); + resultCsv->edit("energyOnlyMe", (double) pureEnergy); + } resultCsv->toFile(ruName + ".csv"); INTELLI_INFO("Done. here is result"); std::cout << resultCsv->toString() << endl; diff --git a/benchmark/torchscripts/BernoulliCRS.pt b/benchmark/torchscripts/BernoulliCRS.pt new file mode 100644 index 0000000000000000000000000000000000000000..e635a5f300ae36571e3ca4ec6c2aa2024be946a0 GIT binary patch literal 2723 zcmbVO2{@EnA0NA6XdC-d6KVzxZVY8q24fh@Bx|mOkz~w_u?sOsVVK4zBwOa{MwVm= zS+bQRt~O)&B-uhymidzIH`7Yp=YGroea>^<_xYdmJHPY(&vSl19F~tC01y@i{1g%Z zF#yVwLL>ze2;N3EM>Ozm0d8=z4*`t=0*-!-C|(a&*8md5*c?c_ zQdyXseg7e1dCC7prDgq%N8*saN2kLW3>0?qc!_7UP%)M%229A@&iA|M zUU$nfMFu*y4~g7FmNfCs$R}M6K};H);D7vrSAQvvkki2^5(HbT_x7Kwcpjcg3?+#6 z#?CB4PQ`omBH-SbC!NK%;_0~ph+|Hzh5nQDI z%iUCvDtiWdqMw){4Qc)8Yw-A1lv0Y=+v1DmI58IL)y?O|kvJw_gaL)I8tqVXY5ZE0 zQ_I~W-K8TBD!QUyS8tw_KrKo?DRxfDce>t*ExFkEZcSxj_{@&u3ORPRb02LxlbbD_c>R=D2M(wRX+Xy30WxSDEJj=>Pbk;Z2rq&l}5sIUDb{t>)%y=?2TvEv91%gq`LE6SXP@c$R&6;x%@(SK2F{69s-W7Y3y9a0zP{ zTz>A}-pduQnfh09VbpfbS3#QAuac{#!|Z6m75aY`s5z5n{wM~Wf9=$C?8zf#)9Oe3 zw9Rhw2(56Ny$$8sIz!Z}3)={)_p1mp&ZQpxZ6lh5F4j&vzIVeAHRIqQ`Rle)Sr;y- z-LRwWS+iN*S-Q1hzE(>6Ev6tZs~i24;!M=3rh?Az;YmrKIuGt1o<3i3+)r+TWjUzR z((|kf{TnaCOsC~s$7=tt?1)Q9xKqhs7;kI|9iqk-nqu|z*&GQg>Y(4SDOgA){MvE) z1a9RrlNmdgu!YafXYyG4t_6RI)+({9j#~b@7H9OyKDEf-@Ov&(hmIC#*Rnx)pZah* z>}37|B&bMBMI2E>q=#6{n}kZ94m!A3hB{AYp~dNSfu>2CZrtna3+5 zy2D;{sQh9CqUHJwNRPigOnA7xGBXV!)eXC#3hUP(HLHI)bxTq=r+*GgwU)Ex!;P>BxMe z%;t&Q(hHf23F)`?oABblwP_%({we>OX$vGS$Fv`IGE3HwI?qouwk-EJh0@PIP`U`V zXx*o`N$OC%_Jd}1{BiM$8|DX1T^)X>)xizRFj*_$_gh+shNR*&JyAv=xGtt#DDm;} zh}Z_|0Atv0s$O`NKag=f>oEoRMs0CZ>>4$)KdK3WgiY-hV=J}r7{0l7gAtMWve8_1 zh0!>IVdL`8D9M8t8)LIo6*>xKw=OyWO>Otn$d0`SRghP06h_bGr7&ij%u4#G^F`*Q z9rEoxkWtIDMk?6?XVCeustrN5qhzTr(dDETt&&995*OAsaj)6Hga{l^001D0e7JJJjj#c?@c29 z|N5;h!ri!}@GeDtK8tGaaaLC?z;fz;g0`2_TPvt B1SbFh literal 0 HcmV?d00001 diff --git a/benchmark/torchscripts/BernoulliCRS.py b/benchmark/torchscripts/BernoulliCRS.py new file mode 100644 index 00000000..25b5859a --- /dev/null +++ b/benchmark/torchscripts/BernoulliCRS.py @@ -0,0 +1,66 @@ +import torch +import time +import os + +def get_first_element(tensor): + if tensor.numel() == 1: + return tensor.item() + else: + return tensor[0].item() + +def is_empty_tensor(tensor): + return tensor.numel() == 0 + +@torch.jit.script +def BernoulliCRS(A: torch.Tensor, B: torch.Tensor, k: int): + # Get the dimension of A + A = A.t() + n, m = A.shape + + assert n == B.shape[0] + assert k < n + + # probability distribution + sample = torch.rand(n) # default: uniform + sample = torch.div(sample, sample.sum() / k) # sum = k as per the paper + + # diagonal scaling matrix P (nxn) + P = torch.diag(1.0 / torch.sqrt(sample)) + + # random diagonal sampling matrix K (nxn) + sample = (torch.rand(n) < sample).float() + K = torch.diag(sample) + + a = torch.matmul(torch.matmul(A.t(), P), K) + b = torch.matmul(torch.matmul(a, K), P) + + return torch.matmul(b, B) + + +def main(): + + width = 1000 + A = torch.rand(2000, width) + B = torch.rand(width, 2000) + + t = time.time() + + aResult = BernoulliCRS(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())) + + script = BernoulliCRS.save("BernoulliCRS.pt") + +if __name__ == '__main__': + main() + \ No newline at end of file diff --git a/benchmark/torchscripts/CRS.pt b/benchmark/torchscripts/CRS.pt new file mode 100644 index 0000000000000000000000000000000000000000..b405bcb921eef0ffe38e3d37a147b179867e546e GIT binary patch literal 2614 zcmb_edpy(YAOE^+hR}o#PDbPwCS$H0MQA4GmUN*tYnY5}xg6K1pZhgbqYkN)G-Am; zrKA)!my9AuE_)Se4JD`Px0TaTIp?qQ`M#d-^L)S0=kk7@=kt2r9}HSV6aXL);0IX) zC;+x@?r=XW6>AVe46?V81@`_R@xnM?Un+@=^Y`_Yzyx7&cz;rmAD+C6{M`kwaOz;&tc^rr*P4^wB!-fwqGGTkl4 zv;wpCG2g{@TRzcOxxkH>_rbyPOs4BOD#I*xKE+QeQ)Zlcuy9T12;=>uN?r}2)Th^? zznF0G^1`V@b66$Mv7u*D7TP$=O(x3TMygR+!M6tBlO8vIM?HWip7gLRp&7S>6|)a) zhCOKNzSM_rcJjnyG&^*qPrMQn(Gc z+<})?mV4A+JD#tT>v?A`0bAHVGokC;G3mW=tF`KZuI~LGH`1QCbQhhRv`Yow`J?k) zgb|ZB852;Lk_d|t7UYK>G#y710sxlCzvTxEP+UDf4E*qijuL!7C$IeKnG)lH4pR1% zn;-b#y!0uK+j9ZymayS+P*tvkR-R^Kq@Cj$D;>wKwsUQ^R)|={Sb?fqb=h270!TG% zSXJ%#fc7XiAdfHe9cmvMSkfVcmW#Z(;WQoBF%t#a+EFQFbi)qOJ?$%FeuW z5k||6(g}u!UWalmyB_WL`r~cag44A(PkTnDW|T5wGKUPKU!0P1Gt$>H+E8dFrg4d@ z*{>!^_RAmdX{jya9`4`lHNEK&D9y+Fs*)}Vi)x^A`wt>sH?kEgU4j^mj{UdK4^M2( zudAMc#O&CeT)_=PDz=w3K5A(p?RNU~Q)xsCafTJ`R7%djh)>buRFUe{?jfzZd^8lN z5e^~1uN5x{s4k?;&n&k{F*FML`zyafte8t(xx9#SVOWQrl8SlAhP3*|82@a1jh$9s zuI+3s)_Rlk`l#)TXK6)XM%B##%Yw^ZHLAMOQ7QHI*Uzt)Dx%+{g%iQf)rC4`UXe0| z5X$SKNnMM48lFWN6`%W`e1pyDW{sqleQ%m}FwA;dIBO4VC{-iAj5R-C+?jfK>|kb< zxAi`A60}|aOr1$ZwceEpH8@E5eb}|ZTND|6@NM^uaq3yi`D~&tY>jQDodx}_t#tKE z=E33Wa+49rb7cA}8mS}r;$EolIJS1}bQ>omlU#Z!Hh|tX7)X#Z?14;)r}1dSlbBMj zYr&l-omUoqjXcAc@t`(gkVbID=Q(dY5ax~@&LKyPTAiK#G`yzb6#cA$(wJMRGP?!o zdDFdK65Wd`IL;#EScUKMfi-U(qstH>mmQeB4B3;riUMZYCq~r+wFWW!i7fi7P&N3x zJ7Fjgf=YmSh- zD$@d8QhCKrv}PnfM~oo-W**X|Qsg7Mr|%G>TiK8Ov$c^+ZzVjhMR;uK~3Y}$whM6Kam}0yIV<7Rq zRI6e7Hk+pzHd%jo6^kYb%L&OmkeYHRIcl2TYUXrH!>QG{6FBBwFG?|lDebpcj4R00 zp@Qd8LRsnxarswoS#omX^Wty5{ZJVlEL(%x&N(OAl_2fN5sKC>wqurPBj`Jf)gAAo z2kT@Q%P2ZHPz2u%Wb;Ee4ITcI~SOT7#sRi_MG-l z2e#D}|K_-G-R6z@G;~hmQLN4=?!+(K2huF_I9ex`QmR~-q$!TfabCjG&xE)X;wO~7 z^wfyidkvcDnbgo-mmw?#_M*l|pAs3UG-NTlhU!U{L^Y<9`c=-F1OD~Fk5z+su3R)ih)!XoyRUo;` fB^Te0j{^Yu#075v5KLpxAh2Kz;y(iYFS`E+dR5cY literal 0 HcmV?d00001 diff --git a/benchmark/torchscripts/CRSV2.pt b/benchmark/torchscripts/CRSV2.pt new file mode 100644 index 0000000000000000000000000000000000000000..12de4027388101bf603aa34ec610682416e6f181 GIT binary patch literal 3136 zcmbtWdpwlc8-HB~VaVOcCB)=z#{EttOfH$hL>gwqU|fdDRLUS)m&^=mBW%${R_iXI zGI9?gwMoLVj3`OQuHFZ!IF&OR0&@cyE5di&_C2$7%`^Q8?2L}853%G@00&&3+VOU%=A_5y5h7*Va zpp=GbF#ai7LiJe9 zS>mlEmGFUhKk*~fdRu%cSfq7KT7oLKO+6!C!ca{ZDGIt^WO#h?#EO#IHqHcuVAl15 z%IBoxcP!c!Yp*RbL%|&d>H-Ee?J!rvnl^Y6U8o}awWFcC6x}{+62=eLzZ3>VT+SiY z0m#|?jeDjSx%0YvcWd-K)lu4E*@W{Zhrk=gj8p4;EWII4NiWl5y0ewJ)=xbzP}ex! zYE%BuWM7P_(Ano*)7zN3B)p5?3g3dt3GBfKSR7*4qf-0+$E=GtBD{=VzMC^d7wx~{mi_htJUka;+miQ0*GJ^`QEPo)8>qnq(bU(Qnk9Bx0CsAk zQ!z)bDOhJb-TA{SE}Tgp!E6xQr^vdi(FB(3uuB)Rn*v{>L; zkB`2~UD$w136jbyrqRo_!!yekNm|XnbFqtZ@_d0Pm=gdhHxI@?7o`%g^Ph@R8;c8w zJ>vhziIM-0^2|gzhe7P6R|tdfu@bVu>zNW}V1X%VaM{Mrch2w%bcjQ?QCC$cC_G|5 ze7=^2$vc;7Z7XZ+#_y!Vr=%2RZPBvBSaCe%e(Cl}-pRCtZHg-Qts47+j8xrA?l-K5d(2RlA4Id{BLUF#Uq(G3H>l6pT65&?A*Wb)g8$a4sZv_%sIm znpO$P+ff~T6;-QhS5j>~HH?fo2~|~_p+e}W*s=hJfrfIfP@?0|Lgz&+{Gp9oirPc_ zVF!}CJ-J-a25rAL7nzfj;V{DBqN`KrdI!%sw#$#S>|IvhZI*;`FTFt@*kf<0DRaD_ zEm@++0|8;G)2wGkf0$-UbOm1xw>&1)tL&%Gi1`^&j7L>Dp)^8Pv>e;25LySTl^|Wz z*&GNX6%%mO{)4mz3Hb!25#&>N_VNWjdWNJ|AaokP?;ZT= z0LxbM_c(a`TY;JEzgxhOF zk02T(V_@L64S7wFN=#W74(XD9;YqM!QBrf#>3TNc$GLq5mBvD4%J(z%ltadWa-_bJY{3W~U0*GvxLI1&GZj*D-CAaK$~*z|Mnt|{IZiwalicx!!<+)X z&6{c+q|%vS-EOU;|31Hek~L|5lII2P7G>h@_`1-Cu};TgU|dUhMR%~b#cb6tC10BP z2G*SDI;0ksN9$qysyRirddmr7E$$V3jVrm=0+NHnJl?`s>%2v47j#Tkv-v9Wbd|yH zqpY~ED{#>4DZU))dQ0Hcets2j#c3K!?_S+^kYV$Ttiv*U7ga3NisJ-2BIR6vk?}X^ zr@rKbQf}soUU=5Dn`*@IDiHg)+Jx7z`Y=&0Y=1)Uz6ukG3bJ6pe!ZXltEAx#Dw695 z4)>koi55F-=2_i`2F#?PQ6xV_9gcbI9&lCe50A$eAsjc_Bc+rTh^uARAAnm$G}pxW zc;O96M`bv|=#k$t^7ogDtn#9@LSpMq6JcMlvejCHzwwjZj5L{7xG8@tDLh}Km+Xn~ zdhw~7QE_OMZNQpz<@%*sMOJiDqq5TLgu>IHA1jJzthibGC$R~%<3eqSSB<=<@R>K& zz;elmnZ*#AFUa`5V`8fWEf`3ZAt$y%c0#%2FjIJHv*aUfqPJL=={7h`En zLXm=x8528bO<)ndXtmC<@MYY4r5Q`HE~$H&u%*(h>P(gNbr73!7PY-IIC%j;Sm-0= z{JkO*11vsM4uKehA;b_r$}*GxAen!WKna}CKa^Ow4!_qX-vR8#CcVGUK 20: - sketchSize = cols / 10 + rows,cols=x.shape + if cols>20: + sketchSize=cols/10 else: - sketchSize = 10 + sketchSize=10 # Your implementation here - return CoOccurringFD.FDAMM(x, y, int(sketchSize)) - + return CoOccurringFD.FDAMM(x,y,int(sketchSize)) # Define a custom Linear layer that uses your custom matrix multiplication function class CustomLinear(nn.Module): @@ -30,11 +27,9 @@ def __init__(self, in_features, out_features, bias=True): else: self.register_parameter('bias', None) self.reset_parameters() - - def mySqrt(self, a: float): + def mySqrt(self,a:float): y = torch.sqrt(torch.tensor(a, dtype=torch.float32)) return y.item() - def reset_parameters(self): nn.init.kaiming_uniform_(self.weight, a=self.mySqrt(5.0)) if self.bias is not None: @@ -49,7 +44,6 @@ def forward(self, input): output += self.bias return output - # Define your neural network architecture class MyNet(nn.Module): def __init__(self): @@ -57,41 +51,36 @@ def __init__(self): self.fc1 = nn.Linear(784, 128) self.fc2 = nn.Linear(128, 128) self.fc3 = nn.Linear(128, 10) - def forward(self, x): x = x.view(-1, 784) x = nn.functional.relu(self.fc1(x)) x = self.fc2(x) x = nn.functional.relu(self.fc3(x)) return x - - -def testNN(net, test_loader): - # first, load parameters - pretrained_params = torch.load('pretrained_model.pt') - custom_params = net.state_dict() - - for name in custom_params: - if name in pretrained_params: - custom_params[name] = pretrained_params[name] - net.load_state_dict(custom_params) - correct = 0 - total = 0 - # then, run test - net2 = net - for data in test_loader: - images, labels = data - outputs = net2(images) - _, predicted = torch.max(outputs.data, 1) - total += labels.size(0) - correct += (predicted == labels).sum().item() +def testNN(net,test_loader): + #first, load parameters + pretrained_params = torch.load('pretrained_model.pt') + custom_params = net.state_dict() + + for name in custom_params: + if name in pretrained_params: + custom_params[name] = pretrained_params[name] + net.load_state_dict(custom_params) + correct = 0 + total = 0 + #then, run test + net2=net + for data in test_loader: + images, labels = data + outputs = net2(images) + _, predicted = torch.max(outputs.data, 1) + total += labels.size(0) + correct += (predicted == labels).sum().item() + print(f"Accuracy on test set: {correct / total}") print(f"Accuracy on test set: {correct / total}") - print(f"Accuracy on test set: {correct / total}") - return correct / total - - + return correct / total def main(): - device = 'cuda' + device='cuda' # Load the MNIST dataset train_dataset = datasets.MNIST(root='./data', train=True, transform=transforms.ToTensor(), download=True) test_dataset = datasets.MNIST(root='./data', train=False, transform=transforms.ToTensor()) @@ -107,48 +96,51 @@ def main(): if os.path.exists('pretrained_model.pt'): print('find pretrained model, run test') print('first run default version') - accuracy0 = testNN(net, test_loader) + accuracy0=testNN(net,test_loader) print('then run coocuuring 1 version') # Replace the Linear layers with your custom Linear layers and load the pre-trained weights net.fc1 = CustomLinear(784, 128) - accuracy1 = testNN(net, test_loader) + accuracy1=testNN(net,test_loader) print('next run coocuuring 2 version') - net.fc1 = nn.Linear(784, 128) + net.fc1=nn.Linear(784,128) net.fc2 = CustomLinear(128, 128) - accuracy2 = testNN(net, test_loader) + accuracy2=testNN(net,test_loader) print('finally run coocuuring 3 version') net.fc2 = nn.Linear(128, 128) net.fc3 = CustomLinear(128, 10) - accuracy3 = testNN(net, test_loader) - print('default accuracy=', accuracy0) - print('co-occuring 1 accuracy=', accuracy1) - print('co-occuring 2 accuracy=', accuracy2) - print('co-occuring 3 accuracy=', accuracy3) - - + accuracy3=testNN(net,test_loader) + print('default accuracy=',accuracy0) + print('co-occuring 1 accuracy=',accuracy1) + print('co-occuring 2 accuracy=',accuracy2) + print('co-occuring 3 accuracy=',accuracy3) + + else: print('build pretrain model first') criterion = nn.CrossEntropyLoss() optimizer = optim.SGD(net.parameters(), lr=0.1) - net = net.to(device) + net=net.to(device) for epoch in range(10): running_loss = 0.0 for i, data in enumerate(train_loader, 0): inputs, labels = data - inputs = inputs.to(device) - labels = labels.to(device) + inputs=inputs.to(device) + labels=labels.to(device) optimizer.zero_grad() outputs = net(inputs) loss = criterion(outputs, labels) loss.backward() optimizer.step() running_loss += loss.item() - print(f"Epoch {epoch + 1}: loss = {running_loss / len(train_loader)}") - net = net.to('cpu') + print(f"Epoch {epoch+1}: loss = {running_loss / len(train_loader)}") + net=net.to('cpu') # Save the pre-trained model torch.save(net.state_dict(), 'pretrained_model.pt') + + + # Evaluate if __name__ == '__main__': - main() + main() \ No newline at end of file diff --git a/commit.sh b/commit.sh index 19788286..5a6df3b4 100755 --- a/commit.sh +++ b/commit.sh @@ -1,4 +1,4 @@ -BRANCH=SYSTEST +BRANCH=METER_CRS git init git checkout -b $BRANCH git add . diff --git a/commit_info b/commit_info index ea5f45f9..391bc0c6 100644 --- a/commit_info +++ b/commit_info @@ -1,2 +1,2 @@ -1. simplify the cmake of test - +1. add the energy meter subsystem +2. include haolan's crs update diff --git a/include/AMMBench.h b/include/AMMBench.h index 4b39c821..6ab70a6e 100755 --- a/include/AMMBench.h +++ b/include/AMMBench.h @@ -20,14 +20,20 @@ * - aRow (U64) the rows of tensor A, required by @ref RandomMatrixLoader, @ref SparseMatrixLoader * - aCol (U64) the columns of tensor A, required by @ref RandomMatrixLoader, @ref SparseMatrixLoader * - bCol (U64) the columns of tensor B, required by @ref RandomMatrixLoader, @ref SparseMatrixLoader - * - "aDensity" The density factor of matrix A, Double, 1.0, required by @ref SparseMatrixLoader - * - "bDensity" The density factor of matrix B, Double, 1.0, required by @ref SparseMatrixLoader - * - "aReduce" Reduce some rows of A to be linearly dependent, U64, 0, required by @ref SparseMatrixLoader - * - "bReduce" Reduce some rows of A to be linearly dependent, U64, 0, required by @ref SparseMatrixLoader + * - aDensity The density factor of matrix A, Double, 1.0, required by @ref SparseMatrixLoader + * - bDensity The density factor of matrix B, Double, 1.0, required by @ref SparseMatrixLoader + * - aReduce Reduce some rows of A to be linearly dependent, U64, 0, required by @ref SparseMatrixLoader + * - "bReduce Reduce some rows of A to be linearly dependent, U64, 0, required by @ref SparseMatrixLoader * - sketchDimension (U64) the dimension of sketch matrix, default 50 * - coreBind (U64) the specific core tor run this benchmark, default 0 * - ptFile (String) the path for the *.pt to be loaded, default torchscripts/FDAMM.pt * - matrixLoaderTag (String) the nameTag of matrix loader, see @ref MatrixLoaderTable, default is random + * @note Additional tags for energy measurement (please validate usingMeter first) see also @ref INTELLI_UTIL_METER + * - usingMeter (U64) set to 1 if you want to use some energy meter, default diabled + * - meterTag (String) the tag of meter, see also @ref MeterTable, default is intelMsr + * - staticPower (Double) set this to >0 if you want to manually config the static power of the device + * - meterAddress (String) set this to the file system path of the meter, if it is different from the meter's default + * @warning For some platforms, the staticPower automatically measured by sleep is not accurate. Please do this mannulally. See also the template config.csv * @section subsec_extend_operator How to extend a new algorithm * - go to the benchmark/torchscripts @@ -71,6 +77,7 @@ * @{ */ #include +#include /** * @ingroup INTELLI_UTIL * @defgroup INTELLI_UTIL_OTHERC20 Other common class or package under C++20 standard diff --git a/include/Utils/Meters/AbstractMeter.hpp b/include/Utils/Meters/AbstractMeter.hpp new file mode 100644 index 00000000..3035caf7 --- /dev/null +++ b/include/Utils/Meters/AbstractMeter.hpp @@ -0,0 +1,125 @@ +/*! \file AbstractMeter.hpp*/ +#ifndef ADB_INCLUDE_UTILS_AbstractMeter_HPP_ +#define ADB_INCLUDE_UTILS_AbstractMeter_HPP_ +//#include +#include +#include +#include +#include +#define METER_ERROR(n) INTELLI_ERROR(n) + +#include +using namespace std; +namespace DIVERSE_METER { +/** + * @ingroup INTELLI_UTIL + * @{ +* @defgroup INTELLI_UTIL_METER Energy Meter packs +* @{ + * This package is used for energy meter +*/ +/** + * @ingroup INTELLI_UTIL_METER + * @class AbstractMeter Utils/Meters/AbstractMeter.hpp + * @brief The abstract class for all meters + * @note default behaviors: + * - create + * - call @ref setConfig() to config this meter + * - (optional) call @ref testStaticPower() to automatically test the static power of a device or @ref setStaticPower to manually set the static power, if you want to exclude it + * - call @ref startMeter() to start measurement + * - (run your program) + * - call @ref stopMeter() to stop measurement + * - call @ref getE(), @ref getPeak(), etc to get the measurement resluts + * + */ +class AbstractMeter { + protected: + /** + * @brief static power of a system in W + */ + double staticPower = 0; + INTELLI::ConfigMapPtr cfg = nullptr; + + private: + + public: + AbstractMeter(/* args */) { + + } + //if exist in another name + + ~AbstractMeter() { + + } + /** + * @brief to set the configmap + * @param cfg the config map + */ + virtual void setConfig(INTELLI::ConfigMapPtr _cfg) { + cfg = _cfg; + } + /** + * @brief to manually set the static power + * @param _sp + */ + void setStaticPower(double _sp) { + staticPower = _sp; + } + /** + * @brief to test the static power of a system by sleeping + * @param sleepingSecond The seconds for sleep + */ + void testStaticPower(uint64_t sleepingSecond); + /** + * @brief to start the meter into some measuring tasks + */ + virtual void startMeter() { + + } + /** + * @brief to stop the meter into some measuring tasks + */ + virtual void stopMeter() { + + } + //energy in J + /** + * @brief to get the energy in J, including static energy consumption of system + */ + virtual double getE() { + return 0.0; + } + + /** + * @brief to get the peak power in W, including static power of system + */ + virtual double getPeak() { + return 0.0; + } + + virtual bool isValid() { + return false; + } + /** + * @brief to return the tested static power + * return the @ref staticPower + */ + double getStaticPower(); + /** +* @brief to return the static energy consumption of a system under several us + * @param runningUs The time in us of a running + * return the @ref staticPower +*/ + double getStaicEnergyConsumption(uint64_t runningUs); + +}; +typedef std::shared_ptr AbstractMeterPtr; +/** + * @} + */ +/** + * @} + */ +} + +#endif \ No newline at end of file diff --git a/include/Utils/Meters/EspMeterUart/EspMeterUart.hpp b/include/Utils/Meters/EspMeterUart/EspMeterUart.hpp new file mode 100644 index 00000000..9b6d8612 --- /dev/null +++ b/include/Utils/Meters/EspMeterUart/EspMeterUart.hpp @@ -0,0 +1,70 @@ +/*! \file EspMeterUart.hpp*/ +#ifndef ADB_INCLUDE_UTILS_EspMeterUartUARY_HPP_ +#define ADB_INCLUDE_UTILS_EspMeterUartUART_HPP_ +//#include + +#include +//#include +using namespace std; +namespace DIVERSE_METER { + +/** + * @ingroup INTELLI_UTIL_METER + * @class EspMeterUart Utils/Meters/EspMeterUart.hpp + * @brief the entity of an esp32s2-based power meter, connected by uart 115200 + * @note default behaviors: + * - create + * - call @ref setConfig() to config this meter + * - (optional) call @ref testStaticPower() to test the static power of a device, if you want to exclude it + * - call @ref startMeter() to start measurement + * - (run your program) + * - call @ref stopMeter() to stop measurement + * - call @ref getE(), @ref getPeak(), etc to get the measurement resluts + * @note config parameters: + * - meterAddress, String, The file system path of meter, default "/dev/ttyUSB0"; + * @note tag is "espUart" + */ +class EspMeterUart : public AbstractMeter { + private: + int devFd = -1; + /** + * @brief The file system path of meter + */ + std::string meterAddress = "/dev/ttyUSB0"; + void openUartDev(); + // uint64_t accessEsp32(uint64_t cmd); + public: + EspMeterUart(/* args */); + ~EspMeterUart(); + /** + * @brief to set the configmap + * @param cfg the config map + */ + virtual void setConfig(INTELLI::ConfigMapPtr _cfg); + /** + * @brief to start the meter into some measuring tasks + */ + void startMeter(); + /** + * @brief to stop the meter into some measuring tasks + */ + void stopMeter(); + /** +* @brief to get the energy in J, including static energy consumption of system +*/ + double getE(); + //peak power in mW + /** + * @brief to get the peak power in W, including static power of system + */ + double getPeak(); + + bool isValid() { + return (devFd != -1); + } +}; +typedef std::shared_ptr EspMeterUartPtr; +#define newEspMeterUart() std::make_shared(); +} + +#endif \ No newline at end of file diff --git a/include/Utils/Meters/IntelMeter/IntelMeter.hpp b/include/Utils/Meters/IntelMeter/IntelMeter.hpp new file mode 100644 index 00000000..98aab374 --- /dev/null +++ b/include/Utils/Meters/IntelMeter/IntelMeter.hpp @@ -0,0 +1,84 @@ +#ifndef ADB_INCLUDE_UTILS_IntelMeter_HPP_ +#define ADB_INCLUDE_UTILS_IntelMeter_HPP_ +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +using namespace std; +namespace DIVERSE_METER { +typedef struct rapl_power_unit { + double PU; //power units + double ESU; //energy status units + double TU; //time units +} rapl_power_unit; +/*class:IntelMeter +description:the entity of intel msr-based power meter, providing all function including: +E,PeakPower +note: the meter and bus rate is about 1ms, you must run on intel x64 with modprobe msr and cpuid +date:20211202 +*/ + +/** + * @ingroup INTELLI_UTIL_METER + * @class IntelMeter Utils/Meters/IntelMeter.hpp + * @brief the entity of intel msr-based power meter, may be not support for some newer architectures + * - create + * - call @ref setConfig() to config this meter + * - (optional) call @ref testStaticPower() to test the static power of a device, if you want to exclude it + * - call @ref startMeter() to start measurement + * - (run your program) + * - call @ref stopMeter() to stop measurement + * - call @ref getE(), @ref getPeak(), etc to get the measurement resluts + * @warning: only works for some x64 machines + * @note: no peak power support, tag is "intelMsr" + */ +class IntelMeter : public AbstractMeter { + private: + int devFd; + uint64_t rdmsr(int cpu, uint32_t reg); + rapl_power_unit get_rapl_power_unit(); + double eSum = 0; + + uint32_t maxCpu = 0; + vector cpus; + vector st; + vector en; + vector count; + rapl_power_unit power_units; + public: + /** +* @brief to set the configmap +* @param cfg the config map +*/ + virtual void setConfig(INTELLI::ConfigMapPtr _cfg); + IntelMeter(/* args */); + ~IntelMeter(); + void startMeter(); + void stopMeter(); + //energy in J + double getE(); + //peak power in mW + // double getPeak(); + + bool isValid() { + return (devFd != -1); + } +}; +typedef std::shared_ptr IntelMeterPtr; +#define newIntelMeter() std::make_shared(); +} + +#endif \ No newline at end of file diff --git a/include/Utils/Meters/MeterTable.h b/include/Utils/Meters/MeterTable.h new file mode 100644 index 00000000..6237fa07 --- /dev/null +++ b/include/Utils/Meters/MeterTable.h @@ -0,0 +1,74 @@ +/*! \file MeterTable.hpp*/ +#ifndef INTELLISTREAM_UTILS_METERTABLE_H_ +#define INTELLISTREAM_UTILS_METERTABLE_H_ +#include +#include +namespace DIVERSE_METER { + +/** + * @ingroup INTELLI_UTIL_METER + * @class MeterTable Utils/Meter/MeterTable.h + * @brief The table class to index all meters + * @note Default behavior +* - create +* - (optional) call @ref registerNewMeter for new meter +* - find a loader by @ref findMeter using its tag + * @note default tags + * - espUart @ref EspMeterUart + * - intelMsr @ref IntelMeter + */ +class MeterTable { + protected: + std::map meterMap; + public: + /** + * @brief The constructing function + * @note If new MatrixLoader wants to be included by default, please revise the following in *.cpp + */ + MeterTable(); + + ~MeterTable() { + } + + /** + * @brief To register a new meter + * @param onew The new operator + * @param tag THe name tag + */ + void registerNewMeter(DIVERSE_METER::AbstractMeterPtr dnew, std::string tag) { + meterMap[tag] = dnew; + } + + /** + * @brief find a meter in the table according to its name + * @param name The nameTag of loader + * @return The Meter, nullptr if not found + */ + DIVERSE_METER::AbstractMeterPtr findMeter(std::string name) { + if (meterMap.count(name)) { + return meterMap[name]; + } + return nullptr; + } + /** + * @ingroup INTELLI_UTIL_METER + * @typedef MeterTablePtr + * @brief The class to describe a shared pointer to @ref MeterTable + + */ + typedef std::shared_ptr MeterTablePtr; +/** + * @ingroup INTELLI_UTIL_METER + * @def newMeterTable + * @brief (Macro) To creat a new @ref MeterTable under shared pointer. + */ +#define newMeterTable std::make_shared +}; +} +/** + * @} + */ + + + +#endif //INTELLISTREAM_INCLUDE_MATRIXLOADER_MeterTable_H_ diff --git a/src/Utils/CMakeLists.txt b/src/Utils/CMakeLists.txt index 597ed0ff..3dcb886a 100644 --- a/src/Utils/CMakeLists.txt +++ b/src/Utils/CMakeLists.txt @@ -1,4 +1,5 @@ add_sources( IntelliLog.cpp UtilityFunctions.cpp -) \ No newline at end of file +) +add_subdirectory(Meters) \ No newline at end of file diff --git a/src/Utils/Meters/AbstractMeter.cpp b/src/Utils/Meters/AbstractMeter.cpp new file mode 100644 index 00000000..7ac4feab --- /dev/null +++ b/src/Utils/Meters/AbstractMeter.cpp @@ -0,0 +1,16 @@ +#include +void DIVERSE_METER::AbstractMeter::testStaticPower(uint64_t sleepingSecond) { + startMeter(); + sleep(sleepingSecond); + stopMeter(); + staticPower = getE(); + staticPower = staticPower / sleepingSecond; +} +double DIVERSE_METER::AbstractMeter::getStaicEnergyConsumption(uint64_t runningUs) { + double t = runningUs; + t = t * staticPower / 1e6; + return t; +} +double DIVERSE_METER::AbstractMeter::getStaticPower() { + return staticPower; +} \ No newline at end of file diff --git a/src/Utils/Meters/CMakeLists.txt b/src/Utils/Meters/CMakeLists.txt new file mode 100644 index 00000000..2c1b191b --- /dev/null +++ b/src/Utils/Meters/CMakeLists.txt @@ -0,0 +1,3 @@ +add_sources(AbstractMeter.cpp MeterTable.cpp) +add_subdirectory(EspMeterUart) +add_subdirectory(IntelMeter) diff --git a/src/Utils/Meters/EspMeterUart/CMakeLists.txt b/src/Utils/Meters/EspMeterUart/CMakeLists.txt new file mode 100644 index 00000000..6d680aef --- /dev/null +++ b/src/Utils/Meters/EspMeterUart/CMakeLists.txt @@ -0,0 +1 @@ +add_sources(EspMeterUart.cpp) diff --git a/src/Utils/Meters/EspMeterUart/EspMeterUart.cpp b/src/Utils/Meters/EspMeterUart/EspMeterUart.cpp new file mode 100644 index 00000000..ba7c2b6f --- /dev/null +++ b/src/Utils/Meters/EspMeterUart/EspMeterUart.cpp @@ -0,0 +1,129 @@ +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include +#include + +using namespace DIVERSE_METER; +enum { + UART_VCMD_START = 1, + UART_VCMD_STOP, + UART_VCMD_I, + UART_VCMD_V, + UART_VCMD_P, + UART_VCMD_E, + UART_VCMD_PEAK, + +}; +void EspMeterUart::setConfig(INTELLI::ConfigMapPtr _cfg) { + AbstractMeter::setConfig(_cfg); + meterAddress = cfg->tryString("meterAddress", "/dev/ttyUSB0", true); + // openUartDev(); +} +void EspMeterUart::openUartDev() { + devFd = open(meterAddress.data(), O_RDWR | O_NOCTTY); + if (devFd == -1) { + + METER_ERROR("can not open device meter"); + } + //char *welcome="hello world"; + struct termios termios_p; + tcgetattr(devFd, &termios_p); + /** + * @brief set up uart, 115200, maximum compatability + */ + termios_p.c_iflag &= ~(IGNBRK | BRKINT | PARMRK | ISTRIP | INLCR | IGNCR | ICRNL | IXON); + termios_p.c_oflag &= ~OPOST; + termios_p.c_lflag &= ~(ECHO | ECHONL | ICANON | ISIG | IEXTEN); + termios_p.c_cflag &= ~(CSIZE | PARENB); + termios_p.c_cflag = B115200 | CS8 | CLOCAL | CREAD; + termios_p.c_cc[VTIME] = 0; + termios_p.c_cc[VMIN] = 1; + tcsetattr(devFd, TCSANOW, &termios_p); + tcflush(devFd, TCIOFLUSH); + fcntl(devFd, F_SETFL, O_NONBLOCK); + //close(devFd); +} +EspMeterUart::EspMeterUart(/* args */) { + +} +/* +EspMeterUart::EspMeterUart(string name) { + devFd = open(name.data(), O_RDWR); + if (devFd == -1) { + + METER_ERROR("can not open device meter"); + } +}*/ +EspMeterUart::~EspMeterUart() { + /*if (devFd != -1) { + close(devFd); + }*/ +} +void EspMeterUart::startMeter() { + openUartDev(); + uint8_t cmdSend = UART_VCMD_START; + //double ru=0; + write(devFd, &cmdSend, 1); + close(devFd); +} +void EspMeterUart::stopMeter() { + openUartDev(); + uint8_t cmdSend = UART_VCMD_STOP; + //double ru=0; + write(devFd, &cmdSend, 1); + close(devFd); +} +double EspMeterUart::getE() { + openUartDev(); + double ru = 0; + uint8_t cmdSend = UART_VCMD_E; + write(devFd, &cmdSend, 1); +//usleep(1000); +//double ru; + int ret = -1; + uint64_t tryCnt = 0; + while (ret < 0 && tryCnt < 1000) { + ret = read(devFd, &ru, sizeof(double)); + tryCnt++; + usleep(1000); + } + close(devFd); + return ru; +} + +double EspMeterUart::getPeak() { + double ru = 0; + uint8_t cmdSend = UART_VCMD_PEAK; + write(devFd, &cmdSend, 1); +//usleep(1000); +//double ru; + int ret = -1; + uint64_t tryCnt = 0; + while (ret < 0 && tryCnt < 1000) { + ret = read(devFd, &ru, sizeof(double)); + tryCnt++; + usleep(1000); + } + return ru; +} \ No newline at end of file diff --git a/src/Utils/Meters/IntelMeter/CMakeLists.txt b/src/Utils/Meters/IntelMeter/CMakeLists.txt new file mode 100644 index 00000000..64fe7c21 --- /dev/null +++ b/src/Utils/Meters/IntelMeter/CMakeLists.txt @@ -0,0 +1,2 @@ + +add_sources(IntelMeter.cpp) diff --git a/src/Utils/Meters/IntelMeter/IntelMeter.cpp b/src/Utils/Meters/IntelMeter/IntelMeter.cpp new file mode 100644 index 00000000..42a5a1c4 --- /dev/null +++ b/src/Utils/Meters/IntelMeter/IntelMeter.cpp @@ -0,0 +1,93 @@ +#include +#include +using namespace DIVERSE_METER; + +IntelMeter::IntelMeter(/* args */) { + //system("modprobe cpuid\r\n"); + //system("modprobe msr\r\n"); + + //en=vector(maxCpu); +} +void IntelMeter::setConfig(INTELLI::ConfigMapPtr _cfg) { + AbstractMeter::setConfig(_cfg); + maxCpu = std::thread::hardware_concurrency(); + power_units = get_rapl_power_unit(); + uint32_t i; + //printf("we have %d ,%dcores\r\n",maxCpu,i); + cpus = vector(maxCpu); + st = vector(maxCpu); + en = vector(maxCpu); + count = vector(maxCpu); + + for (i = 0; i < maxCpu; i++) { + cpus[i] = i; + } +} +IntelMeter::~IntelMeter() { + +} +uint64_t IntelMeter::rdmsr(int cpu, uint32_t reg) { + char buf[1024]; + sprintf(buf, "/dev/cpu/%d/msr", cpu); + int msr_file = open(buf, O_RDONLY); + if (msr_file < 0) { + perror("rdmsr: open"); + return msr_file; + } + uint64_t data; + if (pread(msr_file, &data, sizeof(data), reg) != sizeof(data)) { + fprintf(stderr, "read msr register 0x%x error.\n", reg); + perror("rdmsr: read msr"); + return -1; + } + close(msr_file); + return data; +} +rapl_power_unit IntelMeter::get_rapl_power_unit() { + rapl_power_unit ret; + uint64_t data = rdmsr(0, 0x606); + double t = (1 << (data & 0xf)); + t = 1.0 / t; + ret.PU = t; + t = (1 << ((data >> 8) & 0x1f)); + ret.ESU = 1.0 / t; + t = (1 << ((data >> 16) & 0xf)); + ret.TU = 1.0 / t; + return ret; +} +void IntelMeter::startMeter() { + double energy_units = power_units.ESU; + uint32_t cpu, i; + uint64_t data; + size_t n = st.size(); + for (i = 0; i < n; ++i) { + cpu = cpus[i]; + data = rdmsr(cpu, 0x611); + st[i] = (data & 0xffffffff) * energy_units; + } +} +void IntelMeter::stopMeter() { + double energy_units = power_units.ESU; + uint32_t cpu, i; + uint64_t data; + size_t n = st.size(); + eSum = 0; + for (i = 0; i < n; ++i) { + cpu = cpus[i]; + data = rdmsr(cpu, 0x611); + en[i] = (data & 0xffffffff) * energy_units; + count[i] = 0; + if (en[i] < st[i]) { + count[i] = (double) (1ll << 32) + en[i] - st[i]; + } else { + count[i] = en[i] - st[i]; + } + eSum += count[i]; + } + +} +double IntelMeter::getE() { + + return eSum / 1000.0; +} + diff --git a/src/Utils/Meters/MeterTable.cpp b/src/Utils/Meters/MeterTable.cpp new file mode 100644 index 00000000..0a51a4d9 --- /dev/null +++ b/src/Utils/Meters/MeterTable.cpp @@ -0,0 +1,13 @@ +#include +#include +#include +namespace DIVERSE_METER { +/** + * @note revise me if you need new loader + */ +DIVERSE_METER::MeterTable::MeterTable() { + meterMap["espUart"] = newEspMeterUart(); + meterMap["intelMsr"] = newIntelMeter(); +} + +} \ No newline at end of file diff --git a/test/CMakeLists.txt b/test/CMakeLists.txt index 96fa3cd5..f1fa57db 100644 --- a/test/CMakeLists.txt +++ b/test/CMakeLists.txt @@ -20,6 +20,7 @@ macro(add_catch_test appName SOURCE_FILES SOURCE_LIBS) endmacro() add_catch_test(cpp_test SystemTest/SimpleTest.cpp IntelliStream) add_catch_test(sketch_test SystemTest/SketchTest.cpp IntelliStream) +add_catch_test(crs_test SystemTest/CRSTest.cpp IntelliStream) diff --git a/test/SystemTest/CRSTest.cpp b/test/SystemTest/CRSTest.cpp new file mode 100644 index 00000000..f2afbf87 --- /dev/null +++ b/test/SystemTest/CRSTest.cpp @@ -0,0 +1,63 @@ +#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 COLUMN ROW SAMPLINGS", "[short]") +{ + int a = 0; + runSingleThreadTest("scripts/config_CRS.csv"); + // place your test here + REQUIRE(a == 0); +} +TEST_CASE("Test the COLUMN ROW SAMPLINGS, V2", "[short]") +{ + int a = 0; + runSingleThreadTest("scripts/config_CRSV2.csv"); + // place your test here + REQUIRE(a == 0); +} \ No newline at end of file diff --git a/test/scripts/config_CRS.csv b/test/scripts/config_CRS.csv new file mode 100644 index 00000000..b6afe1ef --- /dev/null +++ b/test/scripts/config_CRS.csv @@ -0,0 +1,6 @@ +key,value,type +aRow,100,U64 +aCol,1000,U64 +bCol,500,U64 +sketchDimension,25,U64 +ptFile,torchscripts/CRS.pt,String \ No newline at end of file diff --git a/test/scripts/config_CRSV2.csv b/test/scripts/config_CRSV2.csv new file mode 100644 index 00000000..a63e79b2 --- /dev/null +++ b/test/scripts/config_CRSV2.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/BernoulliCRS.pt b/test/torchscripts/BernoulliCRS.pt new file mode 100644 index 0000000000000000000000000000000000000000..e635a5f300ae36571e3ca4ec6c2aa2024be946a0 GIT binary patch literal 2723 zcmbVO2{@EnA0NA6XdC-d6KVzxZVY8q24fh@Bx|mOkz~w_u?sOsVVK4zBwOa{MwVm= zS+bQRt~O)&B-uhymidzIH`7Yp=YGroea>^<_xYdmJHPY(&vSl19F~tC01y@i{1g%Z zF#yVwLL>ze2;N3EM>Ozm0d8=z4*`t=0*-!-C|(a&*8md5*c?c_ zQdyXseg7e1dCC7prDgq%N8*saN2kLW3>0?qc!_7UP%)M%229A@&iA|M zUU$nfMFu*y4~g7FmNfCs$R}M6K};H);D7vrSAQvvkki2^5(HbT_x7Kwcpjcg3?+#6 z#?CB4PQ`omBH-SbC!NK%;_0~ph+|Hzh5nQDI z%iUCvDtiWdqMw){4Qc)8Yw-A1lv0Y=+v1DmI58IL)y?O|kvJw_gaL)I8tqVXY5ZE0 zQ_I~W-K8TBD!QUyS8tw_KrKo?DRxfDce>t*ExFkEZcSxj_{@&u3ORPRb02LxlbbD_c>R=D2M(wRX+Xy30WxSDEJj=>Pbk;Z2rq&l}5sIUDb{t>)%y=?2TvEv91%gq`LE6SXP@c$R&6;x%@(SK2F{69s-W7Y3y9a0zP{ zTz>A}-pduQnfh09VbpfbS3#QAuac{#!|Z6m75aY`s5z5n{wM~Wf9=$C?8zf#)9Oe3 zw9Rhw2(56Ny$$8sIz!Z}3)={)_p1mp&ZQpxZ6lh5F4j&vzIVeAHRIqQ`Rle)Sr;y- z-LRwWS+iN*S-Q1hzE(>6Ev6tZs~i24;!M=3rh?Az;YmrKIuGt1o<3i3+)r+TWjUzR z((|kf{TnaCOsC~s$7=tt?1)Q9xKqhs7;kI|9iqk-nqu|z*&GQg>Y(4SDOgA){MvE) z1a9RrlNmdgu!YafXYyG4t_6RI)+({9j#~b@7H9OyKDEf-@Ov&(hmIC#*Rnx)pZah* z>}37|B&bMBMI2E>q=#6{n}kZ94m!A3hB{AYp~dNSfu>2CZrtna3+5 zy2D;{sQh9CqUHJwNRPigOnA7xGBXV!)eXC#3hUP(HLHI)bxTq=r+*GgwU)Ex!;P>BxMe z%;t&Q(hHf23F)`?oABblwP_%({we>OX$vGS$Fv`IGE3HwI?qouwk-EJh0@PIP`U`V zXx*o`N$OC%_Jd}1{BiM$8|DX1T^)X>)xizRFj*_$_gh+shNR*&JyAv=xGtt#DDm;} zh}Z_|0Atv0s$O`NKag=f>oEoRMs0CZ>>4$)KdK3WgiY-hV=J}r7{0l7gAtMWve8_1 zh0!>IVdL`8D9M8t8)LIo6*>xKw=OyWO>Otn$d0`SRghP06h_bGr7&ij%u4#G^F`*Q z9rEoxkWtIDMk?6?XVCeustrN5qhzTr(dDETt&&995*OAsaj)6Hga{l^001D0e7JJJjj#c?@c29 z|N5;h!ri!}@GeDtK8tGaaaLC?z;fz;g0`2_TPvt B1SbFh literal 0 HcmV?d00001 diff --git a/test/torchscripts/BernoulliCRS.py b/test/torchscripts/BernoulliCRS.py new file mode 100644 index 00000000..25b5859a --- /dev/null +++ b/test/torchscripts/BernoulliCRS.py @@ -0,0 +1,66 @@ +import torch +import time +import os + +def get_first_element(tensor): + if tensor.numel() == 1: + return tensor.item() + else: + return tensor[0].item() + +def is_empty_tensor(tensor): + return tensor.numel() == 0 + +@torch.jit.script +def BernoulliCRS(A: torch.Tensor, B: torch.Tensor, k: int): + # Get the dimension of A + A = A.t() + n, m = A.shape + + assert n == B.shape[0] + assert k < n + + # probability distribution + sample = torch.rand(n) # default: uniform + sample = torch.div(sample, sample.sum() / k) # sum = k as per the paper + + # diagonal scaling matrix P (nxn) + P = torch.diag(1.0 / torch.sqrt(sample)) + + # random diagonal sampling matrix K (nxn) + sample = (torch.rand(n) < sample).float() + K = torch.diag(sample) + + a = torch.matmul(torch.matmul(A.t(), P), K) + b = torch.matmul(torch.matmul(a, K), P) + + return torch.matmul(b, B) + + +def main(): + + width = 1000 + A = torch.rand(2000, width) + B = torch.rand(width, 2000) + + t = time.time() + + aResult = BernoulliCRS(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())) + + script = BernoulliCRS.save("BernoulliCRS.pt") + +if __name__ == '__main__': + main() + \ No newline at end of file diff --git a/test/torchscripts/BetaCoOccurringFD.pt b/test/torchscripts/BetaCoOccurringFD.pt new file mode 100644 index 0000000000000000000000000000000000000000..abd515618154a4b739e58738319627c1772694af GIT binary patch literal 7882 zcmbVx1ymi`vi89tI6;D2f(F-v!@+}FaCbTA!JPzmch}&W-~e+jsW@q zA3DvS@Gee9R^UevN0?BEu_G+A(kmonUVzay-{Z5Nr_!S)xR09rt18q0uD{X5#NO15 z<^NxSWEncGxzP5sv0t0?y9eSNQrr>xKdCM_xNBFb{YeN=+ztPVj_C!nNpvE zmo*r*(ND%CFgn$R{K^!u@E z&A3;_x$ErMLX3UfG&C92L_*>Ujr&g*o@pV|kz^+Vf3e-(7rhCJCly1``Ms-~jOO)^fYQ!@8M9qY34x`Oz}nlWS9@sDn(Lt=78 zVLEimgqxxZ=Q$bX`y!({;;d!`UDhnsZn#lrfPD8WRk@(yyNNJ>%?EfYBr^VmqQf-YeBJi% z3n^YZvUrwGAiUv0;_WCDr_;|mjshtBng#9{d7bNBGU?-5E2g#zeKA7KG`aLZ0r{_! zFHgAn&NMl$Ve3#m^<%Hqx%KNJKRaRGefaTQ<3>Wib2qxismEz41QQ!mrC9~Toc}UB zk=GGBX+nn->ir!Xave7xNl+xj8ToM=4&dYRU9RX!wUqW2b$jM(4tYi#n3pv>8Md0C z_qjDJ*Zl@bhWFQ_$P_zAu)PRJPZ+5Lh{T;G@U(>t_2D*`Q1&M!pvfP!)>j4Pe4$ ztDPpiB ziD+!o+m!-Z38R$RfpZcUGje0LM@X!(El8Ur;6abqI{6+iu|U54HS%c zO{jX)8fL)>-U`1GOMDIn*R+JCIa4i%>s503+xrBb#fy(9Q-_mI;o3GE^LpA_RDf!G zzJP5fA|wjdV(&HC3dV_5vpyFR_zV@!bA8cn7-kli!RWpU4%gd%&*FI)=H!*tgt`dr zQj+VRX8MsX8txpNFce#HggcU`tajjO*UB%1F3=9t<#=+H|IV041O8L_C4R#2OMIUQ zM5a*u@a0mwAD%Akh-qx75+RzsE|GypsaGUtvxRmdL22FrsjM$mCY`3J|#-%Q=s4 z9fL}}Wn*CEHhtpL&#$uqa-RkhNUADZ)xM%lR__%T8PrD7u~krQ@u;#oI}d+u6J7K! z59j(B`!3!pK@(cV0|(?P__Y2 zW$MQA&A%(-uIt>~6NU!GOC^gdv#^y@riZ-_u|YaJHkav3;tqVz$wlG*@~h-l&u^pK z&d+*EjOoj18t7y)w07&L1G7ovG$)w}KXHNj$eLP^>N@J?zadyPzgORq5v{7YbhWFd zT{TAudmqf-pdefuTcmrB?-7}tBBxmNG)!8#4_#;`6|T_@ZcIzXVMl zR=ZZN#x_=t8sL(AJl9llc8{LQo7$9@?Q5f$7rL*2iY7AicNlWbQ<%AY+}L4B@=(j3 zY^yIdB%}w^b>p~4mF(iu0tzj~VV&E_={hn5%+4iV$hcnB#US3gshmXw+Cz8-1L9`A zhhxzeD)CyKMk4$9b0zN!k3OV&)oh6u()_dI{I_s$h(A0<$B7)O@@R^sqFm_BBKMks zeHz0ya94c^Ix=bcz*;K{$Ccf4ZHtPwBa+>~#xP4H0jcExYt^Z|xxy4ZWMwPW+C8Z- za+2W^(hd-X||C!L5Kd}uPU zisi;l9FlBkk1Cxf)z)bqpYkRu)>Oa&A8z;Ko3E5XP|eo9o-@onr(|*>R5_?RE$-aT zQXWb(+#~N#JStVmio~gX4Rs--Q>>04x~{q3T`v51qnGw@KWot3~X$SbDUSOA9qPVL^CL7hRwCySIQbd=?b z)G~D_3%#_i#*7TlNwysJlrL;GEnQRp(heHsedzf!?QTFTW5*^Bzo)@R=!Y96EN|zc zgIi-Kk~=Sk`yR$?sC9;^pM5=T2wN8lRPZfK9b-(cq7Fp$jboos^H#Id>fNwF2p<1B z8T*+FIfRElCRM`X$rLyGeaX^J`I-S|yN}SCzWHOWF}MCka5o!_<`HqTgo-{H+)(+! zE_ceyE>)&SWqGmf>DlL%pyIk}0cv-KW*e@GnfyJ29E;nbmkZq{6~?|A#AhMTUk`yd z@aO5$bk+l+fLI8-k=AhCzEH)GTGr|@QJB752ai&GeU7-$8{yMy6G$H81VpnwI zKRdx+DogDmY*(P1>q=DMczKYU)qpC5IT3UgYvY!xu|5$!alX`%5IylD6Op$Z)AXp& z%2T*RbgaoMcC#V9o*O>G9Ek3SkC)^+NgJ&sg6!<7a`ufN>CG;OKWmgy&q4}pnr!jB zYW}IbvIqCt3rBQ8*_cVNXXeF(zGvXYWI{s@%`UyYCmbvN4RaG$S2|qgtkXr{v#mxE z$>|mRMz@-MgGo2u%FToE{L|OUSMcj_*VrZJnFxV_T6#)KxPT^_F3DWlxAn`Rq%f}V zH39V;qLhjjW`_dCHh{VjD{2eW{BQXcg1gk7+>RHdy`t6Q>Y44#o*k4nouC7Zt~!U7 z^!jn}&wfK+{Jh1-xiam_5!*-%p|Dlk6|LF>U5?WgfZ@QV=z22c{iM5VtZJnBP5*u1 z@K82TE2|nm*kV{4`=lMfVFO(5S|FGBEIgUsB>|MBIEomX-25`Yc8=hhXr+LNst^I1l587BsL^l=c-QD+fz4$t%$IIwWcwa zEqUt`+m87ms(v+8)$i5piiEe|zzV|B*h1$PSSV%YAuOOHw$Wi86xkY7kG#<@dV$%M zm8pk}n*FrS73Cy|;s(osK#u7p`Z~6jj~d!pYK7oB^(VWp2NJywT#K5c^y>$Lz#{|G z^PZXpF!|tKzj1ZWz?aqTsr0I{xE3t~>M~+8NJIshcxIJ3zW55dOhOcwR(t5B2-nRM zLsSds;3v$XJxH&`y9G=bs7S4dcU5{j@`sfA9c+Ut^J>^8rwa}(AFYc!;Kt6(OWX|? zatm^bi#N~U+JfpIuIG_T#a>h-O-uG-V^U{w4enX@l2Tr%cHJr2dP8t)qBt5gG=oAb z4pRgRLMtg<;aP_W&y^8ucWhNOO19Ya;%A1PwmisVU_Y&0ze{o*?WwyiJ-A`&y%9So z_jerAKRprc!Wi>qDz_gf@Kn`U6?^$+KEXe$hI-q`o1xqf>XQRF2y|xrdJ3Sf@~2p+;f>;Aw+p26!r|GzR^e_rdvU3iOif~jcrUnZpMJ(B4k0U3tUPI5o^2` zpA|wVB+A-2p$_IB;uXUKX(4Fd4E~%wgrA*YCh6I$X%LC#HTM+Jj;T zaoOsjm-skTQ#Z}jfc$CmSSgvOYgphedX&`wEuyFXtUgvFf!W~; zWiN?YeUR-tGaaJ}P-c6~hsxwbI05tsy7C7}0j*(4KuBH<51>xFlqNHK-PRze(3M;I ztxuJx>`*=WhOQFhOkIcw|L_6}HhGy@42+7Vs+Nz=S=;1VdpOcJEJ68@Z)vC0oxzI* z$IdX(m&g`*pGN3z&3Q-YgTnVkGe0}m4CGG*Iv}qLSa^w>v4{b+k;H_(3*5R}_0{Oa zaMqc}j!r0gzR?MoVVUo~EHOEb#ZE#{*1}D)F>bi&^Ch#vRXGlkWI?c?XHxX#df&fx zkd_(Uv)|4Z96|VKIk*OfbwL<^tq=ZPBV<(N zBRGf^t^Eg{5I#LHM^7+j@@VPB(A*3A-CD8b$eaGVcG)%PV za;XwOKW}S7c;^ceX+ptp_{(+jS9U1&+Rzyz7u+h7EVUd?qVV;d7M^|gf~-1oV(@m+ zo4Jp&m3#MeNrLJu+R9Hb7312s1V-{;#moD@MyASJ%mQcLw4i0hlSeD*2d|V62nk~9 zg_i2x#xu-Ly-A^>*gjng04)8up%c)E8_DXmobvpP-{5APM(oKaiC!b5xFWM4>VT!vDv9j7A_kQsHRsZY#ZV2-E zlmi9;7=ZgH8wvp^|0^5*TJ4ekvo$MUFFGxXKXZ+TtxSjiqOyyfj2oWG=bF(#vc$(J zuanMfohNT(!S``Rdf#0%xs8uAmGS5ncZt16kf04PH`2xE(1$Q`41e3o2#l}5QAih` z&F;9n-@Lz>;RGs|sexDh3@h!Lgcs)9KKp*SGTdCbIUH>yCSStAOXzZtkc(rwbkrnh zgv0L(n7i1t#FC=4l4!w(@bEeJ~_oO<&4e-Dh7PQLJCdm zI$_?MtRcM$ShjI98Bcz-SSuBj_m#!XFL+6@cvvbyf;B86wb8?G?W+WntX8a|j0RV4 z{gzSSyK{Ar-6g2xtEg$DQnx7kU1ZkfSlmX(5+6gdY2DacqMi8t zdxeg-e2ykV--`6;46RWj_3m(^<a$S;2P|yaDd^kTGXSIrnPR}NA*r;(Kf`34C>l7 zz|ZNRiyw_X8D|5am=Ak_gtbRk%Rwyo#;6(Rwi^OKd0W-VNGfZdlShsORF$2YMt=qg z-YngTJ3vo8gc>1*1Z)X~gUmxsOa<#dcU&wTI<|)FKF#6(@?LBdowrju65fhw9)XawyInmxDmJ>5bPedG(0jl9Y;@%SL`=sQbQHn)=$KI&GjrJn0tUWX&VUxZX%^W zjnAh^YYtNr$Hs{UJ>eEnEOTHx**h`WQ5hrSxeQZWkDY_=yqKplfJjJ2h=FwgCpqHQ zKgd~$XEKCaNkBS;tLoN_7q7|rlM&8T>) z+V?O(^G7Vyt)t%NS@==JQH7QF{3<=ymhl)p6U=UUY1VYTiJK1S)}hY5=>U=&i+om& zLZ8?YDI~d1$yI@dvlNC77^QN_D0rYY{Xk%J&Rt-UBGT0AlxWeH@?iHMTRFpPl{!5# z2EzK1==d3Emr=rHE**R_^CCaVctR46&hE2oo2PAB-Oue|qfXtJ43pOmO#*t?3t9FB z2u@kj=wAu)etrk9xc0G@D#$CvJh~GH)zi_BI3D$~#(r^cK>b4Zd6bHm`BbD($BTm6 zV3yb`)0s+^Qi13UtnNl#)5@x<;<^TWoAMaovFtQb%<8$==~Y|9Mtul z(kgSDNl{I3X@cs)hAE~}-p47#X5ALtv3YMpe0rh{s&A&rvw6D8Ubk?2etP0V3Wn}s zne?nLSvh*yfjU&%5{Q7q4GJO^pT13ujscx$Y4K$#=RHS_tr~+qCfLbbJ*xnG>M*3%KWJuJLKWXbKSh-8N@tkJjkp{EgQn7N2FLduVpxbZg6+Xi~9L>5rsRlG@Jm4R0hMqhYy@^ zLrhMEtgwu^eo2|kXFCw1Q0Ty0C!?Hq+562K4#MZpjO}#whLEOROqM>brNf77J%of$ z>Zj{&B#W}~^ownQ+ZFp*dn=`vxE=dum2yr>SI8W0oKpxF+1yAqu?AIFu)b&L&_s}} zJ*kT9B|WEp@I4rW5qMj2f+pmj69BEh1nS1AVrt)yNZV@8(C-gTBilz+T$Qc3m!0n? z>U$v*U$SPuO>#aKj>!puqFy2kg~xoHi@49ti+w!TEjf|9kq39032fdq3FO#R%-;{5yvR0Qkew|Hg#bv7o$u zY!*Mh^iQCKf6IiqnmIW?rhoqx|Hb*_=O^JK(-6TvO8$8&y@L69`9DtZSB~y6WBtqc z|1*TY torch.Tensor: + return (torch.exp(k * beta / (l - 1)) - 1) / (torch.exp(beta) - 1) + + +@torch.jit.script +def FDAMM(A: torch.Tensor, B: torch.Tensor, l: int): + B = B.t() + beta = 1.0 + assert A.shape[1] == B.shape[1] + mx, n = A.shape + my, n = B.shape + # initialize sketch matrices + BX = torch.zeros((mx, l)) + BY = torch.zeros((my, l)) + + # the first l iterations + for i in range(l): + BX[:, i] = A[:, i] + BY[:, i] = B[:, i] + + zero_columns = torch.tensor([0]) + zero_columns = zero_columns[1:] + + # iteration l to n: insert if available, else shrink sketch matrices + for i in range(l, n): + # acruire the index of a zero valued column + if len(zero_columns) != 0: + idx = int(get_first_element(zero_columns)) + # assert idx == get_first_element(torch.nonzero(torch.sum(BX, dim = 0) == 0).squeeze()) + BX[:, idx] = A[:, i] + BY[:, idx] = B[:, i] + zero_columns = zero_columns[1:] + + # if no zero valued column, shrink accrodingly + else: + QX, RX = torch.linalg.qr(BX) + QY, RY = torch.linalg.qr(BY) + U, SV, V = torch.svd(torch.matmul(RX, RY.t())) + + # find the median of singular values + S_sorted = torch.sort(SV).values + delta = (S_sorted[len(S_sorted) // 2] if len(S_sorted) % 2 == 1 else + (S_sorted[len(S_sorted) // 2 - 1] + S_sorted[len(S_sorted) // 2]) / 2) + + # delta = S_sorted[-1] + + # parameterizedReduceRank + indices = torch.arange(l, dtype=torch.float32) + attenuated_values = attenuate(beta, indices, l) + parameterizedReduceRank = delta * attenuated_values + SV_shrunk = torch.clamp(SV - parameterizedReduceRank, min=0) + + # restore SV diagnal matrix + SV = torch.diag_embed(SV_shrunk) + SV_sqrt = torch.sqrt(SV) + + # update indices of zero valued columns + zero_indices = torch.nonzero(SV_shrunk == 0).squeeze() + zero_columns = torch.unique(torch.cat((zero_columns, zero_indices))) + + # update sketch matrices + BX = torch.matmul(torch.matmul(QX, U), SV_sqrt) + BY = torch.matmul(torch.matmul(QY, V), SV_sqrt) + + return torch.matmul(BX, BY.t()) + + +def main(): + width = 1000 + A = torch.rand(10000, width) + B = torch.rand(width, 5000) + + t = time.time() + + aResult = FDAMM(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())) + + FDAMM_script = FDAMM.save("BetaCoOccurringFD.pt") + + +if __name__ == '__main__': + main() diff --git a/test/torchscripts/CRS.pt b/test/torchscripts/CRS.pt new file mode 100644 index 0000000000000000000000000000000000000000..b405bcb921eef0ffe38e3d37a147b179867e546e GIT binary patch literal 2614 zcmb_edpy(YAOE^+hR}o#PDbPwCS$H0MQA4GmUN*tYnY5}xg6K1pZhgbqYkN)G-Am; zrKA)!my9AuE_)Se4JD`Px0TaTIp?qQ`M#d-^L)S0=kk7@=kt2r9}HSV6aXL);0IX) zC;+x@?r=XW6>AVe46?V81@`_R@xnM?Un+@=^Y`_Yzyx7&cz;rmAD+C6{M`kwaOz;&tc^rr*P4^wB!-fwqGGTkl4 zv;wpCG2g{@TRzcOxxkH>_rbyPOs4BOD#I*xKE+QeQ)Zlcuy9T12;=>uN?r}2)Th^? zznF0G^1`V@b66$Mv7u*D7TP$=O(x3TMygR+!M6tBlO8vIM?HWip7gLRp&7S>6|)a) zhCOKNzSM_rcJjnyG&^*qPrMQn(Gc z+<})?mV4A+JD#tT>v?A`0bAHVGokC;G3mW=tF`KZuI~LGH`1QCbQhhRv`Yow`J?k) zgb|ZB852;Lk_d|t7UYK>G#y710sxlCzvTxEP+UDf4E*qijuL!7C$IeKnG)lH4pR1% zn;-b#y!0uK+j9ZymayS+P*tvkR-R^Kq@Cj$D;>wKwsUQ^R)|={Sb?fqb=h270!TG% zSXJ%#fc7XiAdfHe9cmvMSkfVcmW#Z(;WQoBF%t#a+EFQFbi)qOJ?$%FeuW z5k||6(g}u!UWalmyB_WL`r~cag44A(PkTnDW|T5wGKUPKU!0P1Gt$>H+E8dFrg4d@ z*{>!^_RAmdX{jya9`4`lHNEK&D9y+Fs*)}Vi)x^A`wt>sH?kEgU4j^mj{UdK4^M2( zudAMc#O&CeT)_=PDz=w3K5A(p?RNU~Q)xsCafTJ`R7%djh)>buRFUe{?jfzZd^8lN z5e^~1uN5x{s4k?;&n&k{F*FML`zyafte8t(xx9#SVOWQrl8SlAhP3*|82@a1jh$9s zuI+3s)_Rlk`l#)TXK6)XM%B##%Yw^ZHLAMOQ7QHI*Uzt)Dx%+{g%iQf)rC4`UXe0| z5X$SKNnMM48lFWN6`%W`e1pyDW{sqleQ%m}FwA;dIBO4VC{-iAj5R-C+?jfK>|kb< zxAi`A60}|aOr1$ZwceEpH8@E5eb}|ZTND|6@NM^uaq3yi`D~&tY>jQDodx}_t#tKE z=E33Wa+49rb7cA}8mS}r;$EolIJS1}bQ>omlU#Z!Hh|tX7)X#Z?14;)r}1dSlbBMj zYr&l-omUoqjXcAc@t`(gkVbID=Q(dY5ax~@&LKyPTAiK#G`yzb6#cA$(wJMRGP?!o zdDFdK65Wd`IL;#EScUKMfi-U(qstH>mmQeB4B3;riUMZYCq~r+wFWW!i7fi7P&N3x zJ7Fjgf=YmSh- zD$@d8QhCKrv}PnfM~oo-W**X|Qsg7Mr|%G>TiK8Ov$c^+ZzVjhMR;uK~3Y}$whM6Kam}0yIV<7Rq zRI6e7Hk+pzHd%jo6^kYb%L&OmkeYHRIcl2TYUXrH!>QG{6FBBwFG?|lDebpcj4R00 zp@Qd8LRsnxarswoS#omX^Wty5{ZJVlEL(%x&N(OAl_2fN5sKC>wqurPBj`Jf)gAAo z2kT@Q%P2ZHPz2u%Wb;Ee4ITcI~SOT7#sRi_MG-l z2e#D}|K_-G-R6z@G;~hmQLN4=?!+(K2huF_I9ex`QmR~-q$!TfabCjG&xE)X;wO~7 z^wfyidkvcDnbgo-mmw?#_M*l|pAs3UG-NTlhU!U{L^Y<9`c=-F1OD~Fk5z+su3R)ih)!XoyRUo;` fB^Te0j{^Yu#075v5KLpxAh2Kz;y(iYFS`E+dR5cY literal 0 HcmV?d00001 diff --git a/test/torchscripts/CRSV2.pt b/test/torchscripts/CRSV2.pt new file mode 100644 index 0000000000000000000000000000000000000000..12de4027388101bf603aa34ec610682416e6f181 GIT binary patch literal 3136 zcmbtWdpwlc8-HB~VaVOcCB)=z#{EttOfH$hL>gwqU|fdDRLUS)m&^=mBW%${R_iXI zGI9?gwMoLVj3`OQuHFZ!IF&OR0&@cyE5di&_C2$7%`^Q8?2L}853%G@00&&3+VOU%=A_5y5h7*Va zpp=GbF#ai7LiJe9 zS>mlEmGFUhKk*~fdRu%cSfq7KT7oLKO+6!C!ca{ZDGIt^WO#h?#EO#IHqHcuVAl15 z%IBoxcP!c!Yp*RbL%|&d>H-Ee?J!rvnl^Y6U8o}awWFcC6x}{+62=eLzZ3>VT+SiY z0m#|?jeDjSx%0YvcWd-K)lu4E*@W{Zhrk=gj8p4;EWII4NiWl5y0ewJ)=xbzP}ex! zYE%BuWM7P_(Ano*)7zN3B)p5?3g3dt3GBfKSR7*4qf-0+$E=GtBD{=VzMC^d7wx~{mi_htJUka;+miQ0*GJ^`QEPo)8>qnq(bU(Qnk9Bx0CsAk zQ!z)bDOhJb-TA{SE}Tgp!E6xQr^vdi(FB(3uuB)Rn*v{>L; zkB`2~UD$w136jbyrqRo_!!yekNm|XnbFqtZ@_d0Pm=gdhHxI@?7o`%g^Ph@R8;c8w zJ>vhziIM-0^2|gzhe7P6R|tdfu@bVu>zNW}V1X%VaM{Mrch2w%bcjQ?QCC$cC_G|5 ze7=^2$vc;7Z7XZ+#_y!Vr=%2RZPBvBSaCe%e(Cl}-pRCtZHg-Qts47+j8xrA?l-K5d(2RlA4Id{BLUF#Uq(G3H>l6pT65&?A*Wb)g8$a4sZv_%sIm znpO$P+ff~T6;-QhS5j>~HH?fo2~|~_p+e}W*s=hJfrfIfP@?0|Lgz&+{Gp9oirPc_ zVF!}CJ-J-a25rAL7nzfj;V{DBqN`KrdI!%sw#$#S>|IvhZI*;`FTFt@*kf<0DRaD_ zEm@++0|8;G)2wGkf0$-UbOm1xw>&1)tL&%Gi1`^&j7L>Dp)^8Pv>e;25LySTl^|Wz z*&GNX6%%mO{)4mz3Hb!25#&>N_VNWjdWNJ|AaokP?;ZT= z0LxbM_c(a`TY;JEzgxhOF zk02T(V_@L64S7wFN=#W74(XD9;YqM!QBrf#>3TNc$GLq5mBvD4%J(z%ltadWa-_bJY{3W~U0*GvxLI1&GZj*D-CAaK$~*z|Mnt|{IZiwalicx!!<+)X z&6{c+q|%vS-EOU;|31Hek~L|5lII2P7G>h@_`1-Cu};TgU|dUhMR%~b#cb6tC10BP z2G*SDI;0ksN9$qysyRirddmr7E$$V3jVrm=0+NHnJl?`s>%2v47j#Tkv-v9Wbd|yH zqpY~ED{#>4DZU))dQ0Hcets2j#c3K!?_S+^kYV$Ttiv*U7ga3NisJ-2BIR6vk?}X^ zr@rKbQf}soUU=5Dn`*@IDiHg)+Jx7z`Y=&0Y=1)Uz6ukG3bJ6pe!ZXltEAx#Dw695 z4)>koi55F-=2_i`2F#?PQ6xV_9gcbI9&lCe50A$eAsjc_Bc+rTh^uARAAnm$G}pxW zc;O96M`bv|=#k$t^7ogDtn#9@LSpMq6JcMlvejCHzwwjZj5L{7xG8@tDLh}Km+Xn~ zdhw~7QE_OMZNQpz<@%*sMOJiDqq5TLgu>IHA1jJzthibGC$R~%<3eqSSB<=<@R>K& zz;elmnZ*#AFUa`5V`8fWEf`3ZAt$y%c0#%2FjIJHv*aUfqPJL=={7h`En zLXm=x8528bO<)ndXtmC<@MYY4r5Q`HE~$H&u%*(h>P(gNbr73!7PY-IIC%j;Sm-0= z{JkO*11vsM4uKehA;b_r$}*GxAen!WKna}CKa^Ow4!_qX-vR8#CcVGUKTr5hv^5JWnb?#^9ekq)K1r3FDcWl8CB0qIx?rAxXS{PCRc z`%eDnk7uvBW@l&Sp6A_p-kE#mR#iYo0RYg^0e>-K02V+RqGE3D=Hg-tww9IQvM`03 zayr>N$V#CD^!{Z5Nan`IP>74Ujj=JNs)MPyr47Wv!qPR>EEvXukvIqW5}!^p8s$m&7}-5owiRtoqQ^QbYo7?LD_za8%ZjjC6wiv2w}3w z&GF-egy;P+v#nE0Hy7Uy&YS4(zpdxE3|&X{1$)S^8Y@jZN9-8aG)RTTNpl{hv$5$* zjqKT+*}PXV1fkL~x-z9s<)M1nk>eyJXcgJ5_&2gUD|prdgeH5J(&DD&2S_RLJ{sq6 zY1Ty7e70;@x!)ZWF#_>FUJAxfL~(H%S@YL&%T+DI7s~WJ6i_v9{+>ABLlEX z4Mh3Abug(WX`jJ1>)L3XSf2rRzwNsaDJQLs8ndKNVb20jt`kmN_ypRh1|!$@Wx3Pf zc+Q&h1~Li-HuWTf^%_U9;|uk8*Plw@sR}h!yMbc8uF~%0k#%Ky=Rz~7e>WWeZ z))|1W2>2ogG13dn10<*(G`bdAEaRJL1!;WH9!}K{c|YnAG_Wd*&i~LmKPWtbil-6z zUR3c@d5P#ZW&vnZCVZ?(JYSH4hQbI<#kz=g&k4;UJXYD!jsAZ62B}DM zK)6`Yv zqfL$BF}p;&feaNNpH+6MQMR&obKhlR55yATmhSTJ$afriGW>_ z=Y_gOA;2V+md5#_+HUXz%v|gvui9w3Hce_b)7?x6J&~Wz6>-d10-|xPq@ulkJcq;nSo zxr_|=)%~-8i#17T> zk6wDme>-vhiXM7-*}QzU%ru_PK4D;t?YADHW_LYAM=p7VxPQb?vuqD39e&ZKqyk~` zZp4jMJ{lo>7zY>QzbPIe$Sq~v%c2uve7*<0Dv3opH?1L zfWZH>@^D&Mnz>mU|FSId{NL9g-xPZXN^K%|i;6dV=xVXqHA@agQR(gomnUbAF^4;6 ze$8ZX`wBIHmkBnR?@e6oZx0nXN0vz8Ydn71QHQqun?weNfn{e z>P;N7I6jpMa%Wv17^v8bK2l9MfZVQUdha!^zin)tt*o@Qu_2WZz(h7L){lvfem@iy zP4(E5UqF!CV<~JBZ&D3h<#UPYLs_k`s>tb_Vav>xRxM&L$N#xRgGDJ>VDfkm*&ycT;EH(f<&;Th%qRXw@2~cwt%6Sk)yI+})as_6W5Cx-?J-`wu?emsgZ>274iQUL(z93ebRa_r@jHQ*vv#l6I zH1=|OfRt!(?D%A;26_x7QGQe+lgF5EftV*ey>2}6%$auzv zVjb!9W1M8nO2StH(w!DxX=3DR2@{5Kw>&yQlyZ=|M z8hT#{hw(jrNt;z^_Ovl^qHOX~7LH*ZFBGBX3FJW_;G#TEr6j^;R$NPT9|H2mL$X}z z7%W$!U%!G7swj?-`*yl@es8GV@HqPrTmklq%%|(wFH+}_K|2xtoQ_Ia%_sXd>w`>B z3HhgovPrT!Ta@$9eG_3M$&Nm}da&#HMeszwgUoI%MiSBgalT46>??+WTL5@mV|WOFuc zb*0a$e;!XG)M|@XsE7?D6?YYUQR#+TsrV=~GBudpF$0NSs zWrI-mb#-j6R_hp`X!z#AJ~pm{1BNsd>F%5urdJ|Y%aa8k$qMRGlbXX%u#aFL_6-nc z<1S|!Fi|^#x&9K^bj=PF-|ndFpTQo@hlf zrF((0x^Q9yBvtxh)Ta}9a5*+}J^G&2>3-YC?Qku@cB+uysP;EzZ% z;MI6OO!{!W)JkGDUSLSOFme66FRkHQdC`t<87~+;8=!27#cw_s6o!UAfVm>Tg*)Se zsEJx(1-e|Wl`K36=QvTZ(^Byvr2rmC-FFn-rIk~JFT4T@^_Q8a$byuu=$8>&t*e`_ zv{mj(x+wiH@~z$6BlfZjLfUq(pd-w2+99JHAu+}`RfJuxiC2oBcA&ncPoLbQ=*) z9$Jp|%4Ol2k!(uuX5|=ZRN&7EX10jt3MFj%z6xnsE0|jHEdcvX6ss^xegUhH+Q`Im z6fG-C7^Xhs+ir<5%p@AfJ=#H(nb*0A^;7U3xo=FArKR;zq-!*LKzGEfRAC-aUc0Hh z=2t@KTR}%TZ7Hq{PSj^d0Ng`7+#}h!ZhC{|!RBDDrO1m1t%UPj%byQ^5cx0~_h=t! zwT4u#Zk>Di)W3qPWSmtB36^s8MyTE6)urpKmW z)5AIUWPEdRD+u+P-T-}<0PkQmz(2xY*QqVwgj0;pZ}O z=orEbS6IFMIoo_PPSj6z+CzOCz1}R0S09q;Uc4SwG6~yT^FBB7qo0R+U-|$!mwGFzFXdqA{7?Djq>>xCQ1(^G>YQsT0go~-TL$r114J;W$#WRLf3F~yg zSiG6t)JW9p=chhKlYX)`m+evpbPLrEm+qje*wB?m&GwIgr5lNJ$|YW43y>A5*M>TGER4x?k^7vo?GGJw*wq9i#KUA6`)fzLUM` z3tiA8XHjB<90=8P8W+@Ei$ukEj7IYS&tqee`EblwwKfHV>uS;W*nn$+dXxKy*5ujB zlOGzjs=d!5nw~z1axw(W;0a+ih%k7*1=YF+MIwzvG+?95`IomWWv zc(UNG_r;$+mdw34PX_p)IY$OO1}!&fo`5?QfF7vMj|A^`Y9sG6EeQgR8GRlB_UOPQ z`y8N~y3H$gQ+D*G*RSVk7KDQMo}SYz;Q)V}RE3fDC{Gm?aRj-4BqgYN4obm0q`Ag% zW)`GSxxq7Q3b4tmLE{0<7JGfdfWPJrCM9VXZ$|wf-QG~5T;wWRSVFBzKkOwOo!YLV zlcz3lR)|^kI+I_O!^$4(d zz@7P2ENmLD8$-cTwTBG4$?bA zt-0gaiPp6*yhIB$#Xbn<_Ic?1q+L%*tseiDAFGlEXURHof0Da7+(W0ms6^TB1;I^O z0mCYzlF@xz?C7ab6zWb_>4q|VEn2D$anPiH-`kzRT@Gj*^_uiGI(3C%Nmr+6Fv-#D zg)St~4A;W_m#uR9c^+y`w5;do;OeqWhAAZ{tV*A1m-HOs9m)=XB<;C#nU`*37R|GE z)UWGO$~lZ;vfMp|sA$0Klw`Q&Du%A7dd%rF&rIv|d_w;E23FyIIlD**5yJL*Cr@ra z5lK(qrO+94^E$iJ2Y(_72MJ1WT@w&g#-P83R#l!Mv-IO9w}u&3AdO%G0Mn%Z#I3mj z5B`x`bNzCwz-8qIHiz0mz@`qr9P@GhM6n0eJHUb5SS?3H62z?1Bj4htoZ&lT`XF7b zDDr5gD2A(B5e5hD_iimj7_ZH)TGp-WCx9^)--=WpKNfCEjA&z|d`>1FfpWe`t>B0O z(L0+og%VbnXK7~hlSeGBe-zsC$jhp>q5!Yv2Yu;yQNDT0k}MA5(o@64Rf%s|5gT}i zBvT4si%M0pMuB3YjY;kr2l<3@b9YPGCcf1*Oo@tR61B0(?EXViD7d zAIR#q8S@$^u5mX@qwo@y#jcW2nUkM>>V&ICp3kua-t?AYvG#!oHF+1TW$8KCEemxj z6X$CvbJLRGWh_|fO~)-}(9F`Ku+vk%C+@fXJsf}CLookqs*vp;J+b^$P|5E?MGQk}P$E-&ZJ~AcS;L;6h>|jhEx8F- zsJyoP)S5#7Y$mo&P>9NBYEG}|Wb~mL4+fTYH{o;$b*P)#V-qz>4@b$Ed(IXvxN`!8 zmruJh5n6aK@)uT4d#v}R5FZu3IkkN`HRr+{R}MYX$MuN?WrUT#L~au5MUn+bT1%)ymj8PJ@BQ9o(;jMRp@jFaos-Iv zpS~7M_k;up^yk0Vut-x#L6sg(FfC~a4n`Hf&2IeQ((SQO&0Oqv*ftxIs$rLW1z{_~ zdnO*jkK1j4um(h3S8QrEtG-7duz6H$s_Ts8-co;HBF%+MC+XpI)OR%*5(6+d@A**+ ztm?q>*bP)BOs|Z3w}qEfbkR;#^@FQ00=Y0Hj_DdxM%|KTfefdwnD<;`W!X{4=hW}v ziqJ^Zo1%;7g9oH?ZXTEt)O=Zf;sOo*e!CS5sid`K&Igv99a4&)9xN}77{z+3@2M|2 zwL8qM;yn%#mVL=P;cDg{a7Fbnq<;WY`w_9)w;9u>uSX5__hskuUoBTVIfTU+AOkwU z>>>fG@bIkZ4ScN6$qOVE$j@1zH{#z_d9Sc;gnQG}6AHlXMSPTU;;|Nxc850H$HG_) zXKIwZ(wQ_{X<0Zm>qfD&;xgHnUq@Ss?7;08mSzmk?K{Be9Zm_Fd?V>*32vo>7`Xi1 zk5-;qwHCxEk`;5`z^DR=Uag2DoJK1`mGg(F+UOY!LCef8YlR=AH;Hh!QfqH+`Ksc#&6!!fmVPA20DEQpw- zR)vF;&Q)!hiz#9Zxj*Uw&*(Gsbr$STt-u+-z^^Q?lrleLdFTQCjZu2BLv(1)n1^ba zYy@py9=3jqCr8L7LP8 z(N;mHZx=n{%aW#~;3L@?3p)>(JPB-)+U$TZls&+G1{`Kz-Sx&e{|-gn&5krJMh_(1 z7$UN--u$lD$vPvESC8Wi?%4LMmRv>kI-PCO3}2$v@c{=}w17|e0CVZ?b$0)90|u?Z z=Vyq`O*=M)`>P(mrsoFHrK!NNm%t$P;v5w)MK9PJP=XCv=0*%N#syI z1rm|7JI0mfD=bD0qFKe~(sCv$TiHJBC#r?d=iIJ8d%ukmqTSnw)P=++&7OF~sBWZ1 zp`e?e^Qc=St~~d?Og2mBkJJN@n+p4wnonJabY=WFV^1S*f1Sl8j!0&d^q^yydtXoA zXD!?k@f4y8Ta(1IPwqv>qiZGa-aAwxR9pSH+Md^9p0P zAy4X}izc@@2c&h*;^%7^c52o0YVNY!A8)c{L^Gi|jLtlYRp86{N(aZJuR^-i?wgkK zHzEl+4=`XOIZfMaOnZcbs~$^>_%xaC^F*os(?hHZ`JB{~wqnwq{60o78KJ8;yUZ@) zN)b!?Yk0RIfwlwoi6qtJEHc=Cm)D$mtpfQzoo7GAlC5p>HYM zH9W|N7o^WcJS1a{+lfR@s#b@ix$>Gg!M<VgIUe0=ch_W$<^5Cs7BFWW!Z6>18Gy8bEy0RegeNUXmoIgVTy+IK<{;Ai?tDE=Oi z=r2mnOG_8mJ9hW~)^GSRe?AW}Iukk4-N-+assiw5_8(2__Zk!N&x$|T-z!|dV*I8? z{r(sBO*)&%dLVgHV4{~7)F z5%W{6`^}K=V(6d#fxkol^HcmbTfdJZir)sg`-VjR*KvLqjlVkl8_@4Vg6I!GntuWM zR}iWy+`GsAuclPU9RKNb_oDn_|IS5!i7x

<;Ls-SzvW3hUvo20: + sketchSize=cols/10 + else: + sketchSize=10 + # Your implementation here + return CoOccurringFD.FDAMM(x,y,int(sketchSize)) + +# Define a custom Linear layer that uses your custom matrix multiplication function +class CustomLinear(nn.Module): + def __init__(self, in_features, out_features, bias=True): + super().__init__() + self.in_features = in_features + self.out_features = out_features + self.weight = nn.Parameter(torch.Tensor(out_features, in_features)) + if bias: + self.bias = nn.Parameter(torch.Tensor(out_features)) + else: + self.register_parameter('bias', None) + self.reset_parameters() + def mySqrt(self,a:float): + y = torch.sqrt(torch.tensor(a, dtype=torch.float32)) + return y.item() + def reset_parameters(self): + nn.init.kaiming_uniform_(self.weight, a=self.mySqrt(5.0)) + if self.bias is not None: + fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight) + bound = 1 / self.mySqrt(fan_in) + nn.init.uniform_(self.bias, -bound, bound) + + def forward(self, input): + # Use your custom matrix multiplication function instead of torch.matmul + output = my_matmul(input, self.weight.t()) + if self.bias is not None: + output += self.bias + return output + +# Define your neural network architecture +class MyNet(nn.Module): + def __init__(self): + super(MyNet, self).__init__() + self.fc1 = nn.Linear(784, 128) + self.fc2 = nn.Linear(128, 128) + self.fc3 = nn.Linear(128, 10) + def forward(self, x): + x = x.view(-1, 784) + x = nn.functional.relu(self.fc1(x)) + x = self.fc2(x) + x = nn.functional.relu(self.fc3(x)) + return x +def testNN(net,test_loader): + #first, load parameters + pretrained_params = torch.load('pretrained_model.pt') + custom_params = net.state_dict() + + for name in custom_params: + if name in pretrained_params: + custom_params[name] = pretrained_params[name] + net.load_state_dict(custom_params) + correct = 0 + total = 0 + #then, run test + net2=net + for data in test_loader: + images, labels = data + outputs = net2(images) + _, predicted = torch.max(outputs.data, 1) + total += labels.size(0) + correct += (predicted == labels).sum().item() + print(f"Accuracy on test set: {correct / total}") + print(f"Accuracy on test set: {correct / total}") + return correct / total +def main(): + device='cuda' + # Load the MNIST dataset + train_dataset = datasets.MNIST(root='./data', train=True, transform=transforms.ToTensor(), download=True) + test_dataset = datasets.MNIST(root='./data', train=False, transform=transforms.ToTensor()) + + # Set up the data loaders + batch_size = 64 + train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=batch_size, shuffle=True) + test_loader = torch.utils.data.DataLoader(test_dataset, batch_size=batch_size, shuffle=False) + + # Train the neural network using the default Linear layers + net = MyNet() + + if os.path.exists('pretrained_model.pt'): + print('find pretrained model, run test') + print('first run default version') + accuracy0=testNN(net,test_loader) + print('then run coocuuring 1 version') + # Replace the Linear layers with your custom Linear layers and load the pre-trained weights + net.fc1 = CustomLinear(784, 128) + accuracy1=testNN(net,test_loader) + print('next run coocuuring 2 version') + net.fc1=nn.Linear(784,128) + net.fc2 = CustomLinear(128, 128) + accuracy2=testNN(net,test_loader) + print('finally run coocuuring 3 version') + net.fc2 = nn.Linear(128, 128) + net.fc3 = CustomLinear(128, 10) + accuracy3=testNN(net,test_loader) + print('default accuracy=',accuracy0) + print('co-occuring 1 accuracy=',accuracy1) + print('co-occuring 2 accuracy=',accuracy2) + print('co-occuring 3 accuracy=',accuracy3) + + + else: + print('build pretrain model first') + criterion = nn.CrossEntropyLoss() + optimizer = optim.SGD(net.parameters(), lr=0.1) + net=net.to(device) + for epoch in range(10): + running_loss = 0.0 + for i, data in enumerate(train_loader, 0): + inputs, labels = data + inputs=inputs.to(device) + labels=labels.to(device) + optimizer.zero_grad() + outputs = net(inputs) + loss = criterion(outputs, labels) + loss.backward() + optimizer.step() + running_loss += loss.item() + print(f"Epoch {epoch+1}: loss = {running_loss / len(train_loader)}") + net=net.to('cpu') + # Save the pre-trained model + torch.save(net.state_dict(), 'pretrained_model.pt') + + + + + +# Evaluate +if __name__ == '__main__': + main() \ No newline at end of file diff --git a/test/torchscripts/FDAMM.pt b/test/torchscripts/FDAMM.pt index c320062c3126a8ec26944963562e5372abea6ed2..5f6d9fce0c1d1e7a13adf70c62742d3a8770b7eb 100644 GIT binary patch delta 3495 zcmV;Y4OsHPBETZB0s((ZZ`&{ozUNoyCXI}sVKDa)7OCM$oDJ`(crB&H7z2;D5v zPs-~`w8xcuWJ^?V>7fr_I`LAA%V!^PAM()PAM4_k;u}w^b|5cbM!Hq2TWtsND-~6Z z?-dYx=(O>P_Qlo)pkBotkUnhn0SZP1q=`hr_Y`xl%@LuxDM?zf+T_xd7+igHA9Sd2 zufd?^6<2|Ad!m0w+ynS1@Opb?L7o!U&`jx?lGntPRUaElprlw<2sCXt3X zMc!uxXraoWFAFduz%l0n}m=87cW>3IV88_UW=B zRRsG85-BYE`CImI)1625Z24RaK52D$A(+g42;3V`R#Sgd=mCahL}uW!8IRb!JkMru zj?Va$&dVXd)qCiSr|G;J(Ruw3bjBk(gPzZ8>#i535>1K2)e5pubXxXd@FtG0F%X8PDvR3=Wqb605ZDHaPDrOX=W^Y)AoR8(g7e(K~I*^dVNhmXp#bGPR6+C}odvX+dFpL!txsIo-V$15zN4FavZA@WLS(j;<$g*L04bfVG#B1m~yS@ux zBQ%o%Yh2>njcsANTF?@{D)5nTF5&_%!devW{X;v4*8MBI>j5ukRVA|G=`A?*;bJlw z_nuB%by8^cw9A5uwRmSN5(R&Y)GwADNEWIp*Drs3cO=u_NtVH`H8DmdMs6(kK2#vU za?o!#q<$0)|7B*6CFKE|5>&%z9J`5?V_*R%dqT?{&G3NdhByjoZA0BR7S?KO3hesO zIj9l1^ZXvka*WS7!c=Ug9bx~hwxNCtMCxxeIi8MPp#BE4Ueg}6rZv(UL1iPjo0EPl z^ixx%2GX0UC)VkPmhI!HdnD=i5^vyM^?S*lmhL0al$3Q&o|zsJnN3tkTJ{z*yE+oh zM@G)RZXa*R`pC7Y;5HkI1kvmdP)i303dJ;sK?48)iVBg490MT$Fq1$KA_PW4EC7+w zER(zYO@}Dr`d;RcQ&y@5iKBY zA)XtW|L8sa>$OePKcBWr#$swz258W>y8-g|o3Wv4=>`3E=hH^nn#Qutp#LguJ@qdemvlMr)mj4zhs&5rVSdnD zqisMuY86YFZ2S?!ikCAbW5J3j=31@G9qh!&5P&+=?-)4)#@I0BTtAK+#gW%(JJ2uE zo-UX5lEsGN_FSjCX4t?08?KUKuD6#ql7C-}u>u=dHdM?F&iiwPSbm9BT@-VpUAIU( zc7|aY=^Q4+%I9?}lV$e8Wc0FPHfqh%Xjnr}XMUITQn9S6a+q8c3$V$xNZw(w=k}(uH|lEyygJv=J%cZ03x@ER!MFg>&$3K0SIHM*7=N6d z({TY7a_J1dzGzrkmd00>rI%P<&t%gDqnuaFt@f%$BIQaxg}&wLyzqrOti>3%JihD1 znKf$*wy7L0CTXnLE>F`hOyJVL687vX?*YzE{lcTFiDu$DCeJ@0xeTO!5&$1`qaoH8K{q0MTysZQ~-)#w!vIM)d&29;jo)QdC8-KBWtiW!2 zhr5eeDBiV+%UL`p&a-qTQ!38liIH1U%m=iocf69a1}u%8^|lEQI0>s7bEfp;`l2a|*Au+d^A`>Y*3@t2RbsR)>Vk#P6 zHto`>?iZYO zFg%s_vipsXKdd{-6f@!WE*8Y?30acc3gH}bPGXALshwGNQLU^|W`Fp1E?u^d6>PAM zVkT3OUm>=D4NziprDC3ScXsv3tv5GU!p?Ty;gOl{I9BX9F@g)zrFFWMo=V|FDmanl zQfG4M`~ub*^RlN9V|kiQT30dq?46Ow z7V~<-Di)R!3#H;Cdd4az@GG9Mzvc0TCG&vR>VDquJWnj=PJz4epw=V@#U^y#$t#qe zQ|;Z_g?hhxz%0+oF)}ROzN8(}>ZK|h6>Se|Tko`;?wqNb4}VzZpvvnqwaZ9q!N{D? z>B_l^k(*87;g`dYi((#8XH?x9=XJbJm`Cj&`FAZV1@7a19^gS9;$dFT8+e2(d<9?0 zSMfUOw3_>bIRs5UX!Qwv=6}f6LcEdnD*yH;+t`}Exd_u z<;{E>Z{gc{D}QJBu8r^D?R+QSh0nYB16<`%z6YQ8@)+;naePYfPTs}4c@IACP!HD$AepYmS{9uN@on)P^(^?kv`XAmO--jG=D2gNVABB0LiT`_3$6SZAp&<~FyuD}YpY7B9ufP;QAjee>g zqP_zX5BcFV;wk~B{GQ}>s=Z%gpC1DXzFNT35-U!iH3D9e_<9ij5)^!`fS0NIV~JTm zln~cpwts%_;{2V8Z_-G&sQQ+RGVpoh4u|>eupm%XzqT&yOqzQgR)t|d418>I$qUqF! z5T6Rbrx9BPw1h|*2dR40#hZvsz?h3y5ZiF+LnI-Ms!8fQO$WM6)z3(L(ZL-8zU1N$ z5P#bRJWJJY({aB;)$h{C*Fv84`Zd*lBk_{~yUIHS+?4oc0P3);y95No-X*9H6EPxj zT@V=JZUK9UZ3h(-5_^Jh81VrCrzH*qVGL0fFeC9y5XKRs0JTBlBSCuW65O6(AI_|er z{T=oFy~Ha)Y(&931^mawtB75=`Ra+P)RS4Qpz5l661mmIU!dA8psSuVLN`?pP~Vdh z*My)Au?N41F0M!1C%~fW1v;yXR9%u74#953UICw_)A$@!zu@}rLfkLl>-D6P-+!R$ zH>vNpBzA{j1hG%Rt1k8;9uV+Hs@|Y;evPU(-N?rg4+`)%AVwYR7qGE`q|iv!<_6;1 z(m)EEa8-$i@Ca!j(>Oxaqttgu;`I>x8S$`yw8Wbs_#5JYfC^PF(n&5-^^(N8Fs#Kr zIEdE|JB2XxARZC$6Ny7%0K}tM9Dj*t!tgNSkbr+u^?wqrFg%7hjGxg6;^Sf5ZGs=e z-5nw7fe1N7$`S8LaDj@8M0_t5A0^_aBA#>Wi_-EdVfY?uBY5ITd_IhOR`BEab)f3Y zUe8+*=MH;-ALRY~5I?+p+YAcx7~Wd*f}arZuSoV7US0PKeiE-&iVX?DM_C1IP&`Vb zgkA|_0ycS-%@UTQwM!V{rv*QS_fW4gC?O?*3z%>qE6@_M66W~4;DCpoS2^!ha$Y6x zjaKxI^FF0k0?Ql`@&5x*O9u#m@tFa`2><}6761T{8XFS;JOD;SK}}6BV{dMAbYX6E zb1raeY(_#j0Fluzk>ClE58eq1#WaUO0{{SuljjdS8-Ved0mBIZ0H+oJ03iSX00000 V000000000{lRgkm0h5!B5VPiLsDc0h delta 3463 zcmV;24S4dvBETZB0s()?Zrd;rz56S8qyRObRA;e)7KjhM6=IWF8IpQB7D4}JRDis!XBfANvWbuWHnvMBL` z2I7ycHagP2+FAnYMbrY-hjo2`f>8!(A{mPv!-9BCk_eSmPPYXwO)^!E!PiF*q=kjX zl3pzep(Eqw#14PB1Mm^?I(uzFpJQIJMC+yf-zK9{}Ew%ET?Os75s?hF_&nJIsC58E;zlela~12(VDvq{d; z86DGk)d#qK51r9Co!0|8Z~lYMXh0|F_`I=tJv+(IF`ofB=@1!0iR~2(DGi{QXn)IKVjpR=4#-Xi205l zWW5;13W$GPMq^fyWp(ePyOobtrm*Lv%`}fyUU9sHXpKPOCCr^&-zBgSn9YEDT;aQw zZDFQbvK&4u@Re{HaS4sEE(*8*FwTM1e~ow3Z`rV3_PqGrf z&F=9y-(ypP9+(=(9%AiCEZ}%YS-xQj?(y6Z2dS*vu$~(WYq2&3wsYtl^a|X2eurc^ zMrRyhD^}xZ$LmS45x>YXbSt>3qqly`<_)gsD=scotNZUkx^X1&a zQgTB2hb#W0xAm{pHdFsx$|@R*$x#`gPHS`n7ubdeby_p}MOsp&l3ujfaNJ(& zbmt5k9At?KDdu{6Ya_YE7|VaN!DU0m+~AHcIpc}t@r`!fA}!b%hGnF(m=G&n(5-Za z*&CDAONzNk+b)d~8hSeGyQ~)rB~_KvWTRMs&924NO%_Y8#TLykEflj}+bAFA3Pon% ztgM*I&+BSlX1P@xkrJ+E&dAR)8G>C*Wy|_3OBb@`Tt0@u**P6IU?G2-O5^Y)!^*G} z4q297WH~*ZN#%`FPBFLHyBdj<%DE)^mMZJQ0X3n;7`8l)b>PajYYVoioG$51mvp8p zC94_d3wddvRs%}!c*!j5x@-d0!cby*3M=P$%eC5hSQ(w38ro+LF+p$!IUh)u1X1}JS9jhmtePj#NFM@7jA#q#ia~h6c<=3oh}w; z@yf_9Ddrw+@@;RWECNemYrSR0gHFb(#+)_38(V*@+SOSjpUTe1%_5!&Of6b5!>t;d zx|P`!n(fuP9J8}AHa;0+lhtj$IIB8Oczfi^S#?a#Hx;+ks4U$+O?4bwswX3-&l;&b zE8zW$l_+OZMW=sSctm1uH$}#icpF++z{)&`>iA?dzHHj2jXI{Yde%yzox1EZEo-z4 zQP$qhx@MSL1?H9_GsP9cS}We+lYYd;j>ZMdUhoA z?5%iC-PwN=+j^3C_GwMh(|w~8WBHayIIXLg{r1Ud&lGZcyH&_9wJ#J4=k>HzYRA`j zyZthcFD#k+wVm$!{pJhBay}~XG#=2l%1N;aoiX_hrDs+9w05FCplz|!SzeW6WLUcW zm3C08m8xu1v^}J4yV-WSW2Rz0XqiJQuSr+GMv{LEM*2clSI(7<>}(Qmzbw986!Wk; zqw3Zeui-VqJYxUIzh7jezJj`o(9glE@uiz{BDqbU)hheJ^cKQT9 z^OrN#5U*#v%9|eL-%ju~d@Wzc*Ygd0Bj3a~^DQ`CJ>SZ=@$I~U@8FGmCuca;#G82w zZ{>fx@Od}i!&M&Td+~W6kMTAh$ESAQ!8>^u@5bjI-i!bHc%QSxejX8Kf*rZd;pkXNWlF>{4f=TH|7IU z(|j<6T3En?UhN~)?_r5$9|jk^R=`&ze#n0Zk090w_&yEv1FBx3zE@p*3^5|$ZxU<$ zp!fw>1o-^k9fM{+QClSj{ctbh3IT1d8be$uV8Bn7ageHqsPB-(gMK)LxJm%`dy=1{ z+I-TQXKdATyjdXvV zs{e6O23{{_Wd=hbkfOiGF^NdsRY`~N2WPlXyZm*_Mzk4KJ^}|mQHwq|8 zeBKYgMcjm~MFTxR)d#8XLoU98xLLrH694Ll*Ace}cp*TV;Ga~zF7Z!(mOzR03czECI|Vd`NEr`P zb=bvgh)lpq7q1{T;ns&p8KIbS#3z4(fHygh3wYjDWs9{7cqvRe?sruEJ@tKA;*}saqTn3@ z{_EmXh@E)&YKgj{maJ+eRae)N$jrr`quM2)tCloEH&qW%-%*KcL(q)ajc-I3Hz4*1 zC{y(kUDYC0FH5Wq!5+k30iUDG_&im===$zP>=W>{T2jfcQ}r9v_nUtbyF+jcv0uQ8 zF7_er7w`&IU!`k)jjGq&$j1>62nf_6?saegzcA`Z3Y)3gP)B?l>quccTvg&hyh7^8 zGLBGnnEECpz7&E#A|4WOPU5u?{2g&nz(uMq(nT&&^?ee3VOWP}a0ov?>=eS#gLqiL za}pC_IE{D&izD$&7^Z&_hXwqLsx=Y$eh$M25fk_tm3V&`Pn+OJ@pM!5K!jW(R>XS~ zT%zI<5$~nq`-%AJi09t=lC=D47`}_z2;O)SpA6%f75o^!9jN-O*YkSB`G(!k5AXqg zkRMw9+6)TwD1No(1V1j|%}C}be!BJveiwdTDK@kVJ}O|N;!#mHOX!tw0>4AOO1*^T zXsr^4`6Z~+qzWCU7*Az_Zs3l4bcd6leJ$$6E$H`;=Co)0M15?JPl zi2ol@O9u#s<1sqE2><}7laLM{0YQ_n4i=HYKavmL3C&uP4>0E?674?G)&<1sqE p2><}7761St00000000000000007;WR5KaS*4giy|4iuB45S_Tqibenc diff --git a/test/torchscripts/FDAMM.py b/test/torchscripts/FDAMM.py new file mode 100644 index 00000000..6dc46068 --- /dev/null +++ b/test/torchscripts/FDAMM.py @@ -0,0 +1,86 @@ +import torch +import time + + +def get_first_element(tensor): + if tensor.numel() == 1: + return tensor.item() + else: + return tensor[0].item() + + +@torch.jit.script +def FDAMM(A: torch.Tensor, B: torch.Tensor, l: int): + # assert A.shape[1] == B.shape[1] + mx, n = A.shape + bn, my = B.shape + # initialize sketch matrices + BX = torch.zeros((mx, l)) + BY = torch.zeros((my, l)) + + for i in range(n): + # find zero valued column, to be replaced with a better mechanism + sum_cols = torch.sum(BX, dim=0) + zero_valued_columns_X = torch.nonzero(sum_cols == 0).squeeze() # sum each column to find the zero valued ones + + # if a zero valued column exists, insert a column + if len(zero_valued_columns_X.shape) != 0: + idx = int(get_first_element(zero_valued_columns_X)) + BX[:, idx] = A[:, i] + + sum_cols = torch.sum(BY, dim=0) + zero_valued_columns_Y = torch.nonzero(sum_cols == 0).squeeze() + if len(zero_valued_columns_Y.shape) != 0: + idx = int(get_first_element(zero_valued_columns_Y)) + BY[:, idx] = B[i, :] + + # if no zero valued column, shrink accrodingly + if len(zero_valued_columns_X.shape) == 0 and len(zero_valued_columns_Y.shape) == 0: + QX, RX = torch.linalg.qr(BX) + QY, RY = torch.linalg.qr(BY) + U, SV, V = torch.svd(torch.matmul(RX, RY.t())) + + # find the median of singular values + S_sorted = torch.sort(SV).values + delta = (S_sorted[len(S_sorted) // 2] if len(S_sorted) % 2 == 1 else + (S_sorted[len(S_sorted) // 2 - 1] + S_sorted[len(S_sorted) // 2]) / 2) + + # shrink the singular values with delta + # (this is based on co-occuring paper from 2017, diffrent from beta-Co-FD) + SV_shrunk = torch.clamp(SV - delta, min=0) + SV = torch.diag_embed(SV_shrunk) + SV_sqrt = torch.sqrt(SV) + + BX = torch.matmul(torch.matmul(QX, U), SV_sqrt) + BY = torch.matmul(torch.matmul(QY, V), SV_sqrt) + + return torch.matmul(BX, BY.t()) + + +@torch.jit.script +def RAWMM(A: torch.Tensor, B: torch.Tensor, l: int): + return torch.matmul(A, B) + + +def main(): + A = torch.rand(10000, 1000) + B = torch.rand(1000, 5000) + # A= A.to('cuda') + # B = B.to('cuda') + t = time.time() + Aresult = FDAMM(A, B, 500) + print("approximate: " + str(time.time() - t) + "s") + print(Aresult) + + t = time.time() + Eresult = torch.matmul(A, B) # exact result + print("\nExact: " + str(time.time() - t) + "s") + print(Eresult) + FDAMM_script = FDAMM.save("FDAMM.pt") + RAWMM_script = RAWMM.save("RAWMM.pt") + + +# print("\nerror: " + str(torch.norm(Aresult - Eresult, p='fro').item())) + +if __name__ == '__main__': + main() diff --git a/test/torchscripts/RAWMM.pt b/test/torchscripts/RAWMM.pt new file mode 100644 index 0000000000000000000000000000000000000000..6f9bd35e981213cd8c4b1585b5a32e75c2f05466 GIT binary patch literal 1408 zcmWIWW@cev;NW1u0DKH03_*_JzP|b?i6x181=%@nP67;3XrO^9IX=E5zbH8)KAtNe zCowrSBR?l4wa7O=r8Fm%tB^snu~s7jWPC|cVrE`uUV0&8M`=38quT862T5NzPA6)rZ@sS5WEZWQ*Y>G#BFMrk*{>)odWZa$$G8p*!nq zc9y6M)eGD2d2wFMS+zkeb@GM#>yN}OS)M9)=G+lAw>vj{OxU-nTu(3#-aFqdwBcI5 za@tZCp{>Ux`_nHPT9+?SN&DDc@0iHbXuk2ugQl9E8&*QrTdxL9$h@}e_wAhROc%88 zoII{jEz>tqRBXP~96kH6*47}g{!hE2eWQ%8*LE*lxepXBp+Ar9p9Bn(2YACpkwG3; zxag&%CY7eggCk8EUkHWv?e#lsAkunw%bz3K+|16$LYz7yj307wT|KJM>S4;{x>Sl%BQ@-4i}(AK$85Pu|5Zty=O+qlQm)nL|TF?H}dGHPaR|3V%r8N?}Qx5b>6| zU7cH)JN~JDj_KBm=Ad9+plDz96BxuMjF9AqJtGJ+u)%{luec;JucR1~8<>Cr+rXp= zPRY6YoZ-NX!~nvel#JadUU*6_OD!tS%+I4Z{Tg$DJaTbfdMFdnW)KeWW&~02JdK>k z_&^dU05yk0^dfR0ayk@8(bom!LG)teNpwSyLs=2UkVU|7$6*K}0i*i|Ii!VAjN@cR z7>5*5=msDM7B7kcA#ekr(G%d!#-;;RBFC%?SIi1!!Dt_#M?l~b&;Sqs>SG7da!>^z S=>Tt5Hjo%A5Q5Z0)B*sHU3{wm literal 0 HcmV?d00001 From 5173e6256fae0dfbb680685987f34cb6d35741a9 Mon Sep 17 00:00:00 2001 From: tony <292224750@qq.com> Date: Fri, 19 May 2023 10:00:10 +0800 Subject: [PATCH 2/2] 1. add the energy meter subsystem 2. include haolan's crs update --- .github/workflows/cmake.yml | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/.github/workflows/cmake.yml b/.github/workflows/cmake.yml index 509cd206..3255493e 100644 --- a/.github/workflows/cmake.yml +++ b/.github/workflows/cmake.yml @@ -45,4 +45,5 @@ jobs: # See https://cmake.org/cmake/help/latest/manual/ctest.1.html for more detail run: | ./cpp_test "--success" - ./sketch_test "--success" \ No newline at end of file + ./sketch_test "--success" + ./crs_test "--success" \ No newline at end of file