From 9d5be65fe714a5b7cfe179e52c1551fc7b3385c3 Mon Sep 17 00:00:00 2001 From: darthnoward Date: Wed, 17 May 2023 13:45:10 +0800 Subject: [PATCH 1/3] add pytorch script for Column row samplings --- benchmark/torchscripts/BernoulliCRS.pt | Bin 0 -> 2659 bytes benchmark/torchscripts/BernoulliCRS.py | 66 ++++++++++++++ benchmark/torchscripts/CRS.pt | Bin 0 -> 2550 bytes benchmark/torchscripts/CRSV2.pt | Bin 0 -> 3072 bytes benchmark/torchscripts/ColumnRowSampling.py | 82 ++++++++++++++++++ .../torchscripts/ColumnRowSamplingVer2.py | 72 +++++++++++++++ 6 files changed, 220 insertions(+) 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 diff --git a/benchmark/torchscripts/BernoulliCRS.pt b/benchmark/torchscripts/BernoulliCRS.pt new file mode 100644 index 0000000000000000000000000000000000000000..94629ff510887b3df26e54d9fad31d2b197461f4 GIT binary patch literal 2659 zcmbVO2~<-_77ZX0g6vUtWei&+1R}_)KtMr^k+2D3Lk2@yf}eyykVV-8G-*OmKp-f@ zg+-buM4^#gP>}&;QMR_(TE)m>H%u$LPT~sMr)SQ*f9k*br>btSpNgNAp7lp8y=}BthPvEKZbw$I~bwgaAAqf&=^r!~hE5 zPYl6Q{KEhugbISkN9%bB;b|d$Br=J7oGS7|*fFXo5{Hohn}htGTMEX64JTj%6)^dh z%2vVnIwk~)Khf~>w#eW}2RqX*ls-}WpFG#KA+Xg*Ze$tis3Yl~iO|?&pp^lK+P&P8 z;4&S$eLLo%6B$$3$m7t6h$H<)}B%YcUv@_q2Ud^ySbNrb9n|cZ1k{p7+`E zZrCkEqT}F<)t}v2&o|6zWxGs(;p*jmBJc`&Z*Sd|1dGQ+Rba$5=3l&Lj?VI%Kpjxu z9&=1|ExY&yG;3DJ77O3ZP5a*cu^aVPDkQJI5M(5 zNu?UT=%aO4bUd^TT`j(gedcwFD4lO`!9K8cmrJhb_hvIUkH0W_{O)i#uWXmhsk0{H zPx4lC&&vWAN4!|6@(S~mvinq#<0DG<8yo1i%GxJf@zP3jJt)@POljtB2cN!xa}UH3 z3f2XIsDK7ytJpB`k8ApL{>5-g@t$Lq1({-=PNQ#DwP(lGz=O#*USd;hu)Myir9&9L z+D^AcbTeZ#T0|J%w~f67ZxIk^m)PIMR}UojU*e1OCmstsj(>moHm!>^)BOSE*S%1@ z6*kj6$U46(2Ft$rT1Sl{R>UcV0|BJ60l4;1nOIO&#EnD!ccz9b8}#B2Xd+=_An@y6 zDaSiD%^H%_t9PU&C6RINGV1(1srt$b%nN6Cxbg}{S4MuHAE{qlz!5L!9h|5d$RSxO>Bq_&REpI|AX`)M9LDZ`r!Jc+BzX4y*K0MwP>_44+5Xoz9GNQv(f%yoSAi zS?)lD+Ekme-ERGWC?rZ+C9iKG{#(zx!xpc2e&W;$p%I3V{9%hD5H~-DkYT#XPlwz5_ zIyuMi&DC`QvdkQ}ggXPB9o&v%(yM^FL9I@^hrTAdZ+lI+vj2dOD;_4!l&!E<8JpP8 zL}Vn$Upw$L@nH5NRU3tZ66B+K1H%lIls^hmYaH(U)ZUL({ZcyQ%wrSl&h&jwYQ7D? zxu7#trw{5Sh-%5QKS3oBg z6|Wp*!U;tVPZg(2(^y^6JIc~ZHp31j=~~4~IWr5dv%}hG7!Uo}p*gdJC>WlTRZh>* zY*P4Hr?6hsf$*bq^04gniq7-_qiJ+{)_l!Fw{o*cjShcyuBy#252c^II(U94$QTy% zFxdby>T<#c85J?3fK-&~)p^3JT!;aR!=E7NrpDZTix1{vPKS&e_speMcUJYk4l(r} z1v*@O8PVpdz;7hSF7K*YY_ru4X6xG*VljtOIwv1ac?J&tq?TTkonP$+AGewq`-wAa zv06YD@5qQx8|zBiwm47lP~gNtf`UQ6jBC{L1cnSv(jg zzT&b+c3V3&Kdy<2c(xE17fW4HIgRA_bf@v`ar=f`a2^SRjKyl^<~Sb< z*gk#lZYjLl2TQVgJoA7}hnq_uA4Bjnl1RDn>YWK^2FsQBK1r6PR$6@e%|+Xb+hDs9 z{KzV{zv2FV557*zR=~Y;34BVt!C=TBX~c0*m!D61!@c?D8l44Q(rt-BuPVPy4z}%H zMy*JWb2xTblbKADd6M^{=FBp#MHjLvF)oj8%C0zC1dBMGydhd^y!1ecoQ%}c4SI=3 z8BAUN&v}z&bM7jpCC5EVWn6FTP1B*+?8f-ugWd7fP^;c^Ud`}!?$yWkWc5?$$3$b$ z%J1IiRi)}SlR1HIAAx+xtADNFs-TTu=T&kj&5ukA{UGC_f=n7G3V*M}L57l^0zi#La!Op^l!-xpo=bB_~jj&lbu+G-i1)nOE z4-j99iM2Acb=`7?;AJhbSZj}7Zyi>kfcD>GeUY)Re{${dtFA+Bk@yC6O%T7v zS$mMsbvRH#UA2zg$I#(08#f|8&fF?y@Oh2+$KEGT5y-gE&VH=V1nUIMKp+vp6@2dn SFC0dEqp&Y7c!cvm%>56t0Psox literal 0 HcmV?d00001 diff --git a/benchmark/torchscripts/BernoulliCRS.py b/benchmark/torchscripts/BernoulliCRS.py new file mode 100644 index 00000000..e4b17656 --- /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 + n, m = A.shape + + assert m == B.shape[1] + assert k < m + + # 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(width, 2000) + 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.t(), 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..9b68fa65bfc45490e6e545d77d94f3ce9168c1f0 GIT binary patch literal 2550 zcmWIWW@cev;NW1u06Yw049-Ep`YDMeiFyUuIc`ou3{e=MfhjpYz9hdWIU_!vD+oo4g>E>jQ#Yq@$#3!&e(AWR40ngsg+B$wKLzWpn z2;yAgcdJ)lF)vg!b?xazOFcb}X8ibC;d?F5IYwZ5jm_=z?{7{OI(Klhkm$OeRTu4M zR4ty+s(tbtgT?z6QQfL7eZ@Ok7rnXobXjJfLVoL&p#6?)-pL}0re+r;erx8Om06{> z*eE||>!X@ai_^VW87kSp|$p^P`KaSx<+J?4GUhi= z#5dJQIc@6H*k|mO%&z%aMRQrIPzIybEWiIQPYd{Go%6NQzJ0QL0-yM8E^GbXwC6jE zGs@4mZdKtu=YLD{{Iw;o>Tg0gc$=?Wty+itP8qc-nQ1!*o|~It-TL)_ zci|l)zdbenHr5f3UZiQaCcDy3?b+rNv;E>%~VsM;;Q7C!OAua*M4TYG;t z-~E|iwtnSv-xC{On1{U&SG37Z{^I4w*7xL%G$=8oJX`UrosofIDKoyrAkUzKJu&E| zq$ZW7$AeR<2GPl)HY7SjIaK1D%HGsl%fA19`)1oZ;kRnt`#f&wsg;&T|C2jv@~ScD zvB^xeWMku*-n+sCL|U5j;{<9husaIWTy&Rr37B6KAQl@IG3C{-M~nCUd{LivDQ(%M z)=M+r*Zlwf_x{i4bN0xkU)~t;=9sbTEd8nRH*|8RZ1+BTc6t6Z@1KXRm0B+|SdzT$ zh}*}HOVTZ;Ry|9-dwPecdw$g>+4O$pz^mU*Rux8{eC`ntbjjy>QJCk$H-{YiUM!n_ z?qpfh*+oV>ee$y8KE9G#`HD^K4cpLP5RND z7n_^4>F(i2Ns-&C|DII|y?#}ue@04fuSnCPx!Suo-8~+p_^TmCwJk^eW{tq&ml4&k z=H0unSGng=Ro}krpA33)PU|RpxiogyHuFXS>D!%XI^*F*VlXm*@pUEjJEY^j)EYDiq)LzmQE9Bbu z>`~7VIln(f5uL@se$oj`Jk3Nzs&$X6G`ekD9w8UL`TWC+g_}#SX01ruxh{47r8~^$ zv(i3KnUwqAE9LLp7iqe?-|=tGxN8-8(xm@Q=C%v_^mlMO2YU9gge&{%#S4cs=O5#d zJ}BjMOYln1iX$nTm*2V+nfRgWmZh2C)mgfWD?{`+|NdE7_c-P+)9jvC&5_g0TMjO= z&q+*}Ge60WscKFn~MU}1{ zjy<=#{tMLKzsAZabVQLc_5!c%1@77(_MApv4m0pLakV8#xP7qXNFF)S!tZv3n z<266I8Taa4d|oK`cTV3i=T$o`isU`N1V3T3*eml)E3l9IN4Lv=2a6LgZ!^`YUz{(n z%;N6?-h5@5M_)c`{B{lb?QCQ6P=4{9w*@WMf$#EW>=u7oC;A~`(SBu_XJ58!RO>vL zF7%h}A;#47+@6#s_n4XyMhdo&;l;6xFj*Jq!?29iefL)8kjV}HEFIs zXE@M61`x(=HaD~;ElVvb&dkrFLj7sN1@f5~hnf@<&}I-0@MZ*2@HPc<9nS-jKmo`t z42Vud8vwaN5Jl0|1>}KrVzda*jX(}^c@!h2109Lg2t+-C?i1uN7DO?uf*ocUQk0zuFw6Hp(}AO?mhoFI~c0VEyZ P&B_L1vjQPVJwz=4ZivQe literal 0 HcmV?d00001 diff --git a/benchmark/torchscripts/CRSV2.pt b/benchmark/torchscripts/CRSV2.pt new file mode 100644 index 0000000000000000000000000000000000000000..89180b79b9a9c1541e08893df64f4564da65baf8 GIT binary patch literal 3072 zcmbtW2UJtp7JWoY=%JZVR4_C_Ko~^8p*JB0B3&E^fy4xi2}!5|PACG>YX(F>1w;Xn zB2DQa5{gI}WoQB-APNYg=pT_!hWU-}^_9at231~DAiiq*V5y(U=j*28>{fIao z3UDASRMD9Q?@Psak-SKF3R{ey2L+^xLU4h#0nBsV?KRdwZC4_?UCH;QKmnsaEAb>_ zacZBxr|J`AZmjoRN#9iX-`!hwOw!mPNaMURr%nppJMiMmVMCsgVjY%b7STiQj zbg3d>bTrAksCYo_=8yUfbIx2i{(EjgNy|PVK8eNjmvO$j<$7XoD;C`xt2?jSC?H#H z({eZ;nA}12m%f$T{OMBho7ar}iG=9-;+($1WsCJ$N3{YTbdPPPmelVp{97SFz#UJ3WB->LT~Pve<O;n9Cw66OD;BNH4D zL`gH)yRspz-?UcdHWvzzV}6UHmP7KI9bRco!bq-yj=~`e?BjRxZXUNA1#4wN_7r!5 zT%@5WWhxwEEH)zpkvJ$ z^dnI$b90PIG1;oIb6Pf~$kjb@hKr7)yxmbI3#s1lqec5iG{A$gx=LkNPBIDkDnA+~ zT?1KN`7N^w#v~N&&tt@IK2-Nc9Yh>HyEn~{MtYP`fnzTCxY!jwd8Day%v@n9AdS)M z;bM6vsLw6BBP8oZC4aqKeq1ia|I3XB&K1gcYiu=n*9f8sW{}8YWh))B&DBc8 zG4ms?KCTA~QQalU=!l{06DoH+cJp@GuRDk(gzJb3@Z2e$uzh>Bc#&dl(x*24S{=%) z+?&vH{q5a7tH$GlGER<Ze>0w)%a;IWPJz*y}WT>sbtNWD0*L6 zgn=1dlaRQnSWlooq@1{(sRmuQ_xsf$8g$2UD2@Zk3_khfn$IVzAi4ZDwsH|@Mnhrv zfb`OgIU?8JDC!C}BSOH3r3!?~5abd&>B#}hPMOgM&y;^rfn-G8 z0|~XSgG@Vaw`5h%=oDr;t=n5C3se9t+yUpSm`IU95I50ZZKa;s|dWO`zGEy zHzwn8Li?yU3mNuq(*A}zAr+PSh8Q-LQhgId*Wh3!*6r2+Kbwa zR5yyBf-3ys(ndX>TwR}fzxiQ*50MFCnh!BCIIgn$Mb7(<%TK!BP&TZcoSyr3xT42t zvruoY#C+cciiKSt;WZdL*4F4vUvR2DGxnNtPP0yB%sg-Xn)1X` zt$(FESyeN~qHF@Ee-$7}>U&ArbZMvG+_bUExYO(6mNAG^$-59oo}a#$9zV2JSvnf0 z9cj{=N!9;2UnHeaPB)Y4^b3SeCTvWd8>h=PKv&umv^&g@&hDLQ1s2j=?3|t-5oM=W zOU?|$pYl;AwJ!x4EYq8F>UfD`lq1Wnjdg+>?R9$ur;I`XIE;Tf;$g4Mk^6^89-W$6 zt0jDKdXhBmiGpPyp`Qbl68Rc_*EPyIXV2Xv5DdFW3$&5F&Grg6>zj;Z zuf}UPsHFw~?^@?LPWDa*@BaXEHyCFVFc4xs0k~l55wDhX(f8G8%lMR!oa54P@(1lP z?H-RUHWts|t)Eygasa>*-wx0Gz48(P#6I&3iROzT`Ozr0<3{Tu(77kn(|AHC%{_x%Dz^Z zJGp%Md;t|=gX58`3mkF_g`>d{I_ob zzC0}cuKZ`**>-$eo}GC+kBjw|%e&*VmQe@} OPL{{9{jlP9V*d@j0;!t- literal 0 HcmV?d00001 diff --git a/benchmark/torchscripts/ColumnRowSampling.py b/benchmark/torchscripts/ColumnRowSampling.py new file mode 100644 index 00000000..7196447e --- /dev/null +++ b/benchmark/torchscripts/ColumnRowSampling.py @@ -0,0 +1,82 @@ +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 CRS(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 < m + + Ai = torch.linalg.norm(A, dim=0) + Aj = torch.linalg.norm(A, dim=1) + Bi = torch.linalg.norm(B, dim=0) + Bj = torch.linalg.norm(B, dim=1) + + print(Ai.shape) + print(Aj.shape) + print(Bi.shape) + print(Bj.shape) + + dot_product1 = torch.dot(Aj, Bj) + print(dot_product1) + + + # probability distribution + probs = torch.ones(n) / n # default: uniform + + # sample k indices from range 0 to n for given probability distribution + indices = torch.multinomial(probs, k, replacement=False) + + # Sample k columns from A + A_sampled = A[indices, :] + A_sampled = torch.div((A_sampled / k).t(), probs[::2]) + + # Sample k rows from B + B_sampled = B[indices, :] + + # Compute the matrix product + result = A_sampled.matmul(B_sampled) + return result + + +def main(): + + width = 1000 + A = torch.rand(5000, width) + B = torch.rand(width, 5000) + + + t = time.time() + + aResult = CRS(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 = CRS.save("CRS.pt") + +if __name__ == '__main__': + main() + \ No newline at end of file diff --git a/benchmark/torchscripts/ColumnRowSamplingVer2.py b/benchmark/torchscripts/ColumnRowSamplingVer2.py new file mode 100644 index 00000000..c6c3aedd --- /dev/null +++ b/benchmark/torchscripts/ColumnRowSamplingVer2.py @@ -0,0 +1,72 @@ +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 CRS(A: torch.Tensor, B: torch.Tensor, k: int): + # Get the dimension of A + n, m = A.shape + + assert m == B.shape[1] + assert k < m + + # probability distribution + # dist = torch.distributions.Uniform(0, 1) # default: uniform + + + # sample k indices from range 0 to m for given probability distribution + # sample = dist.sample((n,)) + sample = torch.rand(n) + sample = torch.div(sample, sample.sum()) + D = torch.diag(1.0 / torch.sqrt(k * sample)) + + + column_indices = torch.multinomial(sample, k, replacement=False) + S = torch.zeros(k, n) + for row, col in enumerate(column_indices): + S[row, col] = 1 + + a = torch.matmul(torch.matmul(A.t(), D), S.t()) + b = torch.matmul(torch.matmul(a, S), D) + + + return torch.matmul(b, B) + + +def main(): + + + width = 5000 + A = torch.rand(width, 5000) + B = torch.rand(width, 5000) + + t = time.time() + + aResult = CRS(A, B, 500) + print("approximate: " + str(time.time() - t) + "s") + + print(aResult) + + # exact result + t = time.time() + eResult = torch.matmul(A.t(), B) + print("\nExact: " + str(time.time() - t) + "s") + + print(eResult) + + print("\nerror: " + str(torch.norm(aResult - eResult, p='fro').item())) + + script = CRS.save("CRSV2.pt") + +if __name__ == '__main__': + main() + \ No newline at end of file From cbf7aef18ee5a2d3be1b17507e0f7b7ac9b66c1a Mon Sep 17 00:00:00 2001 From: darthnoward Date: Wed, 17 May 2023 15:34:31 +0800 Subject: [PATCH 2/3] fix input matrix shape and some comments --- benchmark/torchscripts/BernoulliCRS.pt | Bin 2659 -> 2723 bytes benchmark/torchscripts/BernoulliCRS.py | 10 ++++---- benchmark/torchscripts/CRS.pt | Bin 2550 -> 2614 bytes benchmark/torchscripts/CRSV2.pt | Bin 3072 -> 3072 bytes benchmark/torchscripts/ColumnRowSampling.py | 24 +++++------------- .../torchscripts/ColumnRowSamplingVer2.py | 23 ++++++++--------- 6 files changed, 22 insertions(+), 35 deletions(-) diff --git a/benchmark/torchscripts/BernoulliCRS.pt b/benchmark/torchscripts/BernoulliCRS.pt index 94629ff510887b3df26e54d9fad31d2b197461f4..e635a5f300ae36571e3ca4ec6c2aa2024be946a0 100644 GIT binary patch delta 1867 zcmV-R2ekO(6r&Ze0s())YQr!Lz4t3bPd0cVz3dR!poN~wpp4y032vgarjDJ;&NlY- zbKrn4xc+csA!fW(3sY0z37EY@WuFss(b z&^9h40${%wCf)VP2nj}7vxmatF_*+S~G;6#2%-SRPFv$$&T0YdoH1G$=uR+ z0v3Go4Nyx52#?Up6=DGZ080Y^08mQ<1QY-W2nYZG005JB0uKQQlTHaKe?~$C09smF zT8&lDZ`?!_UOP#fIL*)IchYv!q^&pE?6wrDDrzb=v^4F@Ku8ubQVEMU>ts#5_KxjE z2||Js0|yRpgbKKD$%TIaap2k`q7p(#z=ac1xmCb{H}-a8r=&`h*#7*z@0&Nz?@iuU z*+?ptdhz8|nlBVzyIoi+f36i)Z@+Y_SjaCf-{dKsA)Vq1>&M&kf$MElvTs>I*5B}~ zTA*>2T|=U5ZjwdQbKJUZTQ97x>G^pQ#THD*cfAZ*JosXJo>)$hp4WH_8$*kZDRRpU zNKi3J*{WiY4fM3e+gPuB<77xReCbz=f10WBV@#E=HGVwY zRfpmGrWX*0EG&@4DCq9Rb&a>P?ijsIt`SG$9c&^>qvp9Kqh#4uu#2hqfoGNK0cN9d zja`x?xZ`SpDBXHFzR+NPwI}f>3UU-cd|SGqbpl)>(Q$WiSAjy zUe(hiN4C=hjSyk*e>O3EQZv1v3Y|GgqRRi@VOT>@<^kgo6Sbna$9aSwt@_aEdh_0=!q1~-RX(4 zU~wDf*-hIu!q3;sMx_FhNj1WV#{1Y@_Nu;_f1P@@wp%Pn{a91Kbf|uu4JGx_H!dzog9&!JVNg9hgHz#1#T^P7KfSL! zeC;z&D3>%&*mOe~PVGVQGLGe4EZaKjC&QyDruVIPOqust){ewAIl#`=bv_kNLYR!~ z1T+l$F+)m+BsR@@;>IO1Lkx|dYuXs)@=NAwb`3kiJc^4Ro_H~r$pvYX zZrYgqC^j<(Z5BeC+2*GFdg&-WnefbR`b6X7m|J=&9qn(f=_rr6BRFQ+Ov5=i4;Mg(G|a#(WMB@mu}VKjUstG&x-LPlNSpT*`VqozL2s*- ze>e-0briaTum_9OV10zWg4&YZRp`42`vi58teezw3GXZP4}|@Ke%OG&Av`JQCk^-` z!T~|=sj-rMndH7oD)*AU-$?jSp?@PB6!gA?ohlv1rWz9TR|!W{I*o8x&_5+Sr&1l^ zh@g*>?7vB^(t_}+N^c+>#m^z(O_jcaC~!V>R0sf*F$fit z0SYiO(g*+m00000P)h~}00000K?(o>000000RR91P)h{{000001poyAZvg-R(g*+m F007Q*Q!)Sm delta 1849 zcmYk6X*k=77RF<%gxb^Er>@vSM1tCj&`MD|r3`MXrszK+q+@9%L~vz9RIJewwbocu zQ9GrjF14@Iu~n<7ZH7Coy`9d?huil&&pGdh_c_n|=~sGY723ehNA3*kLrE?kL9rST zPbNb7w60nv9BT97EX8dzN>$Zy#6H}yxQkQU@}sZW9m%#M(l%CWcEWTntYCF;o&Ur9 zS-Ra!phc&^OO9_=%^0i|k%XPwGXG;N`<>vfT8`Tq7_L+^%?EFwPE53yQB0?T6@W!| z(tkN`u`W`lK|@e-z_^KHd>gCcJv4hqX}ehU!H*46u4yPcZ*#)B%zRdKx;E@TvS-zR zG9N;h-)3weA6r~vJh@6a0OQwQrs1Wo=WhBv`K{00V6j5ohqhsT=iT!8e=y!^y7eA4 z_2p_br~173-7G_)*99l}4)Tcs^H0Vesd!i6CD-5a>J;&wq|jlF8uLA)f~Mej~c) zEBvn)p@HU3OcR`C;V8+8)$~zA5id~Au7U$0q1Cf+_4#T695?2P^UU*&g~m>;LkM;Yg;^h;mE*t0@o9P2 z5e1zW0ptAH7{!eNIU5U|kT@h-Ouk@xKQXtt#ZOMm?e$2?gzwX9w}Bj*JGr4e@E4Dz zXDO41EcuvEKZM#V5IHi97~F3nZ9b#O)JOeXvhusd*7jm|WazPEAuWeA4fwRbn8jJO zB;ndq5BmU>+P$qk#IwJC)2c1eayMuF-2-7lbO~~U zrTp^RWjZ30B30+Gnsg;+Qo%~Ps1iB3r>mEV7QKlUX)}m+S+%{+X#OA;d4I~#aya9X zy<$Kokc%4-Z$>qKk8hR#cQ~EVkXI^cN>3|2WTLwWL)_W*04J+FSkyI2$qZ|rNQs?D z2h$mxL7Z$<=Qc}x$Jfr1??;tP^D zqasr6P=O@FupDHHL6w6i-H8TL__~d#e$%(5km9?T&CE;B&_*&jI{%xwomL%u`4KRO zJc=m`nJ%wf?++hX>t}fw3&b`%H=rLU5ZwrJ5S2M?dtT(Gcf05O>?rS3%$6Wb!UL+=%quR-!-Ph!2;RY zHPk$nZt15Q#U1>1K|i^qEJ)Nh3}?)uHZd95do3d#HO8^ZLpNFZ3RVjov`)^++=F?X z0Sq^iqARlGcH0UW7en5XM#_jzXuRe$?vsFu=+_9UkwNXHgIBwxdyy*!r%2E#5={coT{p_2uq3^_XkZ`#!%HAUUTFua9haht2)0m{Fcn*z5scF<)E$ znZ09rQWP#Ul$n^mJeqv&U@yo^nw=nmivazyqC&LRBqp9(AV+}3ZS81r)>6l@W))*h zS`YfzeW7vxJcpNT$xG#vqryC9lJwi zv0ZM2sBZSZsDd~9JB6B0ZM|Yq;U;3YmJz(nt$&1)IOr`o1ZbyNoQP_vhWIg$IE z5@9`dgg%C>u-P_FhFeU-y-?o+m8~P!UQLk`;T0)NcTR&}IV|R0ieS9jN|i&Uq`2te z6KbVbHB3qBk3GX9liCx<-aOAZ`4z42nvCYjIbDenSH=>Xq2?30K0WZk+Q(D2;R;H3 zA1w2eFmhl1CZufhfZ?tn2y{&Z1j}KsI9w)UJyP!ig-Po7=?_u z4cOW6gP4}!_mKZ{6$!Cb`u2ww=v#&HGcgbqW*xr>^C?7|ub24-f)#**K}-Vp3=<=Q l($50_6L)nyCH$A;`g`*pcG3Z44P3w<5@SBU?U{S&67Q}h4; diff --git a/benchmark/torchscripts/BernoulliCRS.py b/benchmark/torchscripts/BernoulliCRS.py index e4b17656..25b5859a 100644 --- a/benchmark/torchscripts/BernoulliCRS.py +++ b/benchmark/torchscripts/BernoulliCRS.py @@ -14,10 +14,11 @@ def is_empty_tensor(tensor): @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 m == B.shape[1] - assert k < m + assert n == B.shape[0] + assert k < n # probability distribution sample = torch.rand(n) # default: uniform @@ -38,9 +39,8 @@ def BernoulliCRS(A: torch.Tensor, B: torch.Tensor, k: int): def main(): - width = 1000 - A = torch.rand(width, 2000) + A = torch.rand(2000, width) B = torch.rand(width, 2000) t = time.time() @@ -52,7 +52,7 @@ def main(): # exact result t = time.time() - eResult = torch.matmul(A.t(), B) + eResult = torch.matmul(A, B) print("\nExact: " + str(time.time() - t) + "s") print(eResult) diff --git a/benchmark/torchscripts/CRS.pt b/benchmark/torchscripts/CRS.pt index 9b68fa65bfc45490e6e545d77d94f3ce9168c1f0..4e482a12d2e489f1ebada521bea7a1ff3758b9a4 100644 GIT binary patch delta 1778 zcmV<+PL33+^5`WJNU78tFClJGWVy+PT6OGTyHN4< zY-2ZRi?|>zi9NeJGdr%7a1Nf8sBmRejrPIR8-**R49bhNFkmb`*l1)VMolzS*XqhKTF)foo3(!$I{SlSJ zJ!m>N>8xxgFG4z8Y%A=%G@6~~&Y|5M7uJ$BI6sk7o~={R|I6N4vvk9rwfu$k< zx=eP8!H2G8g=2qjyHQ?h(@LRO&yvDAbn-Dzr)Y+fOW`usE~zH8HePsS_tN>|sdp$; z5#FRMQz)WX8HRH1MB6E>2VSb%5eDW3he7e4VhJxY{ck048-IONOK2oj6z%HHxAOo0 zlQ)yhOl8b;kGK$;WUvz_A8wvP(&~c{+e&vOU8bw5Usd&t5fm4~!=;Kl6?7vAgBuqv zbP!#+6BQLjgmEL{!dWP@5cR&QPP!_Rfu_68JNKM>-mUlgr6e0oBoeESZ_!-7@YF_r zy|A6%+IVcOkbloDZ(ipKjUk=j3hT$8M|{UEmebFge%dR#X4O}Dl0Ae($=D|=TiaSL zOTv=8VSA3dK$cH09A=4W`}0|qr`R|u+h@s|;S;}XkdkR(k?GjP*(W)bH?a;`kRwYZ zm-e-Jl{d3qxneI6D_Hc(deu;Q3rosnmA3|)Y1ci^aDRPblcgoH92V_d-Bo!T>yFTe zWRch^Z)cNXjH>JG>3e3y^p7x#=ey=!&Brtzdjie5Buvi{PlUhxbZV>eUKK;=DbA`G^D z9IWV6YL;COhOeMEtjHC@=p{MzoqFR1lEsTW-F&na$?w8X=!>njFPUaAs^|Qs( z7Gd6!c|+40%d%R&<%MA9-cB|fcrO?t3uPnosNo&@PBoLGlxzq0CKK$$v=1^X!Nj+mo5w5Z>JcAfQ}s)Q zyMM4Won%RPf9A=RYx4f!^%-YYgW>t?0d8wZ| z(XXGWpJu~xeK^a)x-^&xhAj=OvlEcuDn=dhDnHjyp56QW|CH-0pJj7(WpLY1i#PH9 z;EC{eQ9l>_lSTBN`Me=>Utn!WEXtc^AAelq7kz$7gG90MmshBo#lwoB?$yjnsepDB zv{kvL?P`X<1qnz9J|FzZ&u3*xfC5QKK@&7X3$#KTez$`P9nc9~&<#D%3w_X!PXjQB z+z<@I2qdK4DD((^8A@;5a%gejPY86a>D&ciIsg$5T;#(y86 zV+x%?eNw<55sSO!%kTocE+8e~9zMS;pPK}2QV_f%L9?KP61)np31|^?D9Q{g5zUCw zFr%AFqYHgH$~+^%8}N>RRzVLWcvFHlL2FUw{iwmmQRcP;U%(Fn+6Db41YZkK1^qdy zxf5l6iE4gDK(8zGNt{QAp#8~+YJVb$+*=BL54lc3uSB_GoD&j$sL)Rkb_x2ngtry? zIl^v1Z^hXUB)qH8e-ZWw`jLb`DfD-Qy@GyPhrc206ZE?{`+Z#dXEO3S{)%(=qIv(0 zbN@uSWGcpWNcfdPe?-_XXgU=seOPjrlJr45I|G8QreafAi?i!d<%@Cd-6JS>OLB!I zb&wkr^vek6^97Oq4^T@72w`5NE{X&I0Fnp*0FW9O5db#;LsC;OV{dMAbYX6Eb1rae zY(_#ek-#*Qzz5t3I_YT;l>q<%(UYYJIvZhLq%MjC005E*0018V000000000000000 zrIY#yOaYpcM+rUx6bO@t2|Y5B2mk;80000`O9lr30002g2><{9000010000`O9ci1 U0000500jU-0RRA!2mk;80O3kA?*IS* delta 1729 zcmV;y20r<=6!sIa0s((bPs1<_z4I$9PN`I;G2l`Oq)yt6FClJK6=kic?}QaCdbEynxP_wZ(dDLT zjv#M^)Kb4e21{cUAvHG882c$~%^a#GqcU%7wc>p8qi0G50C|4~uc)1|;X6P6t}~F@ zCvb5W;VcwQNkdOO@hnezOsCMM1MK7%vV0msvWY0ZGfHCzDf1*LrezXn0Xmva|3u|* z51P(RIxE}Bi;xZ%+X_1`jb`V$b7*(Zg|#FN&QIDY&(@^sf7v^0mTvG_%a3TsXmqU5 zWs-u9I&>{79CLr$jq*~PRtm*>mK4?@laKRsifU+bDO|?bCB=l+#tV<^UOHbq^$LY5 z!cNLEg(5a9Ls8D1Xgh`Vz)Nx5GETgFPdqzV;^<~5-KmS2S10nKdLnxy+xv5N(tCO| zB|f9pSCR1ggm6CEz4?RP`EPcwspn0|vFJBe?^hu_ZDSVnN=^ojp~3fl}i%nI>a1Gnq@eRuB|~gKu6Ce9#9$5PZ=`A4E`}d{7o$R#=qk`r@O$ z`RIRU(q^X8h0>YteE;|T-}&d9ydG!Mu~_WJGew#&m0#E_ZIri5#m%R0mP`46!q#;= zrZA*qw!}v9@uX)tm3ro7-OIQYM{jzv9cPatQPuXzdU0FH=SWae(oENKR!QOP!ad+MIv z(7hv!;(Ct0*YYrpY|AVoN<@K+eS-~IrsgVUnq;xhGQk0;t%jE)EmPmO9Q5sH2_eE% z18g(s;HpNmpB^#*b`?ff;d(EP*FO|S#ppTA92;+n7fn9J+_NcMuRUFir`qc+WiZ% zJr?vYgxm|!{^M*hFj%*LjAqNz{F7v;IIdSjr)sN$hjfBn5=D;Y;*7)|PW&=hiE~p< z8(^eKHWUau z?hmtQ-8xyxG>`0Qw$Q#WxC=v3Nsa_pBTX(}6ITOo#ta+58baEC5!8&R#ALIl+R|BJ z@QT_wHWAgDEu(oJ{Q{ee>VsL9H-y2WKWt%OoS$IHzhz;El5CTX^8DTx z{-@lK?Tc)st@JPQS@9NL4LlLvF6x*4AE}Vu)nC;_?iJRD#8q*L%-S{kA~K&;aF zH7hsu%0WYu_gZ>?qgqD02HMK@!^*Crc}0jpjN6L;5P#p+1pyMoApu>`4L#5cefU2K zGW5d$3_=QqU>HX5ISOOQjKc&>LQI%VL5kaHs6vf{#NlqeV|+g_PQf(Hz%0zcJS@N> zECGQFa1oXRgB6t7kI)&3E}=fo;qQ>eKlL=c2Jdl5aQGL0-`^DPU7U7F2;LH)o6~Ut z-hp>H^l&;6W+tVOW=iUq(JiUdg}xYOUJ~Ga_?Sa4r!@gS6rhjOR+#xLZ17c>`C5SQ z;U0%1r*{JIGY6T|DvN+C()k}ra1ke4R0eH;`H+<`$bgyLp(I$AEVsuaD+Qi?rxa77v+8x z@K~b1BOK=RwA5E_SIB$d%HF~s;bGU;n zkLb3osx!!2A+^*yNMUP?3P_0+RK~stTM~z=N`uTBTRd_;d6$2a1p&Y+h2N;1vHf;^ zzTKxFwWsju>j0;{Ce6Lq6-Jg61A zB$>`ngU9Za@|l117&|nktp|km5`JV`EM#L%A=IHoPu361x>bWNVX&uz`2Hl5*qu;$ zU`;0jt1zmqHJ7@S1-elZcwTB#OL1l$C7Y4;hSoTeb&VVj$*0`}ewoeY^W`j=f4*JZ zESHPhWO+EFLQT$d(mX1|Y=Si}VWuGTZHHuRel0cOQ#yZMRgfQN2^ALp!jVBJHh9rN zz-O2I%{$8cIJI$@$&^9q|1@|v))BXa&uEC?b4YX&iu*1$8r}p$p|$bCBl#ws&wj+= z5i0ASOP({xLgD!+S&TJ=gWxM~TF-+g3ufV)_dIdDBjLO2{Kv+3&x^r5W#dcS43%yF zD;EqlQ5YJz7JX~=Iep7uO|Sj}P)i30DY)o_3jzQDtCJ4{9|0zk3nBzYLLvZ>(IS71 zR^4wDRTSTumX94u`D_dQyrl@UWxLx_pjum`r9jE)mC!AEBbIe`ciJ7gJG1k#KrtF) zOfEhMQDb7FFD5?u=#vkI7!zaSA3$P^G4U_ZCkeiI?#y;NQ*6_u-I?G0o!{@Cd*+WfE(NOeuCV+am&|$RwG{cxpq zy5kX(Oiq&NAn4}sqGIo6tr2>KOb}DCldL;Pqvluztzc9PZxvH-B8t5y(70N)R?$O^qnRbuOeK@TOkdLK6>pr>1F2%S zFbOq^y_XHF?WZkUrdAs8&+JY9_^AiC_Z-?+(S?&g|e5xqwgyN!+b&9I)O=3%00 z4ETMQCNub-6StT{J4*>&VN0Ec@xD&6JA&dlx>tA1xTJvdfe}a8C5*`4z^-(&WzDK+ zX1=UvXLNVPvufEXN=bjtpSEnSQUB~k|D5yI-L;k1bmw%YwrU?`J&ke15>!=^=Cp&;KgfvppG9nW5=nAz98->9wYP(rmT#Kczt>zbmz9-brZmREP?Qwn3 zN%>h}(8qck2Gz|oApQ_X9TpV3e_grx-p8I(&MNi*JKj+G?}>kn;yE1hn^?9*)F1a- zDWZ3c`?|>c1-2WB3Gq5Gm(SZL{Gs%d5v_oROBGA=(xk9SVkg<&^)9jkxqr3Y&CjqublD8QN$w>Dy$&L%H?Bc6_q_w?ya@ z%#W!rtA$i>zTtn+QM~rH;W)zjpLcZhqJduk_?i{Qe*I4~b;D!CKG@uLGz?#RTw8{p zVJG8{Wy|QZVRWPVTSlJ?Q&}tWmf>Sz_*wmYQuNlRaimnvuMx}%= zU$25wquAr>qN;oIkbnfYU-mEJ@5izrKmr+>pc!_+PS}41yYVv#3haRv*bA+&5Al9D z00-d^9LD=LXvfc;&;bcyc?4RyJpm z>62*OuQBR3M08f7*AXh5{wv@c5-lOz!)c2gV_W65YMy|Tph5};;S{7{2r`g`VHk;2 z^E`c1q7~G&aH>h~)>9shUi)(KPct@tI2zxlq$5~Cl_htGJVK1IVz(+Fu z0%0Gg58~`c&5?{h7u>fpJ}3?Z2PVRK_7oG({0C4=2M955^0D^?008j}005908xsHs z07Ft!Rx&SRZ*FsRVQzGDE^upXMnVJtle`Jy2`RYfgbM-y0IQRy3OpMzZ}PGC1pol? f3;+Nj00000000000000002`D33QhrslT8b=;W#Ou delta 2196 zcmV;F2y6F%7=ReC0s()?Zrd;rz2_?iIgx;Bc!}%OKnmDMZY^5Wy$J$AOCyIcMJgnf zr1|=esl`i*o^lYxnR)Nc!eOqZ_wZB893QPN(0#W>qmi{gzPthuqLxlp7|?kk;11F( zdbdqkp1FdP$|$pgC2XzLCs?cB@gHVIS-l@YZ^UcDh(@fcBM=fo#Q5O#uCbDmBz-YNjpG>CXa?Xh>- z+K+0ds?jg0JU+bUSO;YdyIy&Z?r?l?j;es`i*!1$Jr(t_>Rwj0#-UFqO~F!=od(B4 zv{28i^(FmWnyP>E@RlNuXotC~oz120NA%=EcdQ!S@x}@4c}KoK$s~6tP(i54s6*ve zHqQPzV>4v6Yf%(vh+mP;cn0XDJP`Z&FMS%l2hDoy`{0@$B>M{ARJ3-;Ni3 zj~Wd-&&jjMq%{O*Uub=R=C(sN48Kx_i3uOBis;8xLXCfUxKK=hiX|=H5ybQozj+6l z4UiD`mQ4~U!k-%NhB~6j#FSG6pDEEvDB6eEAb3N9(ij_LK=w`fkp76vBX!!oj4Vqa zrNXm;G9PN7M(|bCO%SA+kQjZl7AKE)B*Wo4|6v4rT?yRte(2P2SJ?KyY?;6&5+l}O zXq-8RZxaz%^Q*s5O9u$aWXM_q0ssJ&lNbXZ0WXpZA_PW4D*%zvD}RkuOKclO81}By zHrY1s=KbiTP1|+q+DX$UG{k9>z7kCbq{OAuM`YIaCSGE%y$`n}0#tEeXeER|6mj5! z5JDB=ArKrmfrJp^fW#R%m2%~T5GVL&*Iw_YB#M+|=bP{UzW<+p9=RjZ?ocRn@##e} zn@n9@NY1C0l8XyZU4Kd?XXh?oFhVi~I%Eh`fRAK|bO63tbp`2JtspCC3MEzDDw+{= zd|2WJEIRJZ*Bb&>ezy7wrgY^>U%Ko=S54ZcjhIrG5|X z4C)JtZEDY_=D9&V9jq7>wyvOoHf?f9N`_QZZk>DYW6F8S*ndZdD@rS;l+?}Q%bI21 zfVM4zzLB;C^_I4w^17R7JrL6ttVCa(F`AvYvv$$UR&Ix4Q9S;=UlEj&Pvq5?{w z*_g8!$!Mn{pMRE&{mwCJP%KL|Z61?Ab8gU$v1Q5Va9UHP*U&U7NJb~^caxUPVp>US zdCgvj#C)0QvJJTs&f;kip3bd--Qh)Vrs_Jo0L(MPY@G%$3vL21V zruDiV@j~3vAy-z`f$pOt|3hcBTe3S9t(eEGGF%=I^wSs@ocF~ltMXWE6vZnOmSNiR z8XBo6xoCjSaZ#q2FBkMwb=6ge0<=lfiv>-|ds)|_aQ;epC8NRtuoD-RyrqKCAnoEt zOP+Q^HGhki4AFcs2WXVDymw~tb`(KAXfaaZ@~Wm{t8%aae&UdGh#MM~Zd4V^;5{Z8 zLv+kJDzz#V8w!;JU^ZsUktjM3ze{`wH5{Z7u1jyLQy|{aNyZ`C$KxxPEfAH4y8;oz z^a4c0ufUnL;yI<5SM*d)jh|PoHM>}f&k}-`oPV1aU*gS!yBR;{taGt=yE^lXYMzLd z){Vo?b~_IUF9RNdm3&dLqbR+_V{(Kx`yYceY^x#}M?D)Qo85_xyxHa$w~08ta$a$D zCoUt-?f1D1CAiHo&qg(NViVnLlW=T~dyBHl>7Dq*oY%_paraJFSLJlXU0>XDi8{w`<$*ae6ekE!#$)^r9=>-!}S` zS4y?X+lEhg;nfaq8$M|pPsliwsr^q%m9)%SUX{{iEuV#pZ|C8Ol#J8zvaH&RID|vY zm~uY+=cgPeuz*Dz#&viX-i`O*dfb2|ynh$(!;Lt^MNL>>#*?H+AO{2nWBg05X7Q_Q z(TrPgD{jN>ct7rdzn!=XY};`+?!mpd5BI~k9S`6^d;lXD4dH{%%po_$Fm7VTG&b?8 z42uju%uVBUJ`OYdCyZa><2r`_h4CwxF}#bBkWeFtAk-8@p}!fww)-hX8~10Psn zCnM`&-|5{5=W5L7e#b`@N1nWkk;XcJC6T-Uu$z&C9KI-$y8wF_iF0^QB(DMNW#n{_ zO$4=PIJ_^CcL4Um<^@^I;d>(a7+^o#0S-SA$=3h}7%_wF?K)q@FLUl^k^BbSAR}-2 zINLbOV*de9O9u#NPS;NA1qJ{B?F;|_kQyix02Kg3Qd3qkFJo_Rb97;DbaO6nYiveB z5dd0RT3T9KT3T9KT3T9@yb0n7$z;e{0s;U4m6N9mJR4_D*G}pM008X_001EX00000 W00000000005|jQ4P62w8O$)P{J`HIA diff --git a/benchmark/torchscripts/ColumnRowSampling.py b/benchmark/torchscripts/ColumnRowSampling.py index 7196447e..3d68d97c 100644 --- a/benchmark/torchscripts/ColumnRowSampling.py +++ b/benchmark/torchscripts/ColumnRowSampling.py @@ -1,6 +1,7 @@ import torch import time import os +import math def get_first_element(tensor): if tensor.numel() == 1: @@ -18,21 +19,7 @@ def CRS(A: torch.Tensor, B: torch.Tensor, k: int): n, m = A.shape assert n == B.shape[0] - assert k < m - - Ai = torch.linalg.norm(A, dim=0) - Aj = torch.linalg.norm(A, dim=1) - Bi = torch.linalg.norm(B, dim=0) - Bj = torch.linalg.norm(B, dim=1) - - print(Ai.shape) - print(Aj.shape) - print(Bi.shape) - print(Bj.shape) - - dot_product1 = torch.dot(Aj, Bj) - print(dot_product1) - + assert k < n # probability distribution probs = torch.ones(n) / n # default: uniform @@ -42,7 +29,8 @@ def CRS(A: torch.Tensor, B: torch.Tensor, k: int): # Sample k columns from A A_sampled = A[indices, :] - A_sampled = torch.div((A_sampled / k).t(), probs[::2]) + ratio = math.ceil(n / k) + A_sampled = torch.div((A_sampled / k).t(), probs[::ratio]) # Sample k rows from B B_sampled = B[indices, :] @@ -54,7 +42,7 @@ def CRS(A: torch.Tensor, B: torch.Tensor, k: int): def main(): - width = 1000 + width = 2000 A = torch.rand(5000, width) B = torch.rand(width, 5000) @@ -75,7 +63,7 @@ def main(): print("\nerror: " + str(torch.norm(aResult - eResult, p='fro').item())) - # FDAMM_script = CRS.save("CRS.pt") + FDAMM_script = CRS.save("CRS.pt") if __name__ == '__main__': main() diff --git a/benchmark/torchscripts/ColumnRowSamplingVer2.py b/benchmark/torchscripts/ColumnRowSamplingVer2.py index c6c3aedd..f4b00c41 100644 --- a/benchmark/torchscripts/ColumnRowSamplingVer2.py +++ b/benchmark/torchscripts/ColumnRowSamplingVer2.py @@ -14,22 +14,21 @@ def is_empty_tensor(tensor): @torch.jit.script def CRS(A: torch.Tensor, B: torch.Tensor, k: int): # Get the dimension of A + A = A.t() n, m = A.shape - assert m == B.shape[1] - assert k < m + assert n == B.shape[0] + assert k < n # probability distribution - # dist = torch.distributions.Uniform(0, 1) # default: uniform + # dist = torch.distributions.Uniform(0, 1) + sample = torch.rand(n) # default: uniform - - # sample k indices from range 0 to m for given probability distribution - # sample = dist.sample((n,)) - sample = torch.rand(n) + # diagonal scaling matrix D (nxn) sample = torch.div(sample, sample.sum()) D = torch.diag(1.0 / torch.sqrt(k * sample)) - + # sampling matrix S (kxn) column_indices = torch.multinomial(sample, k, replacement=False) S = torch.zeros(k, n) for row, col in enumerate(column_indices): @@ -45,9 +44,9 @@ def CRS(A: torch.Tensor, B: torch.Tensor, k: int): def main(): - width = 5000 - A = torch.rand(width, 5000) - B = torch.rand(width, 5000) + width = 2000 + A = torch.rand(1000, width) + B = torch.rand(width, 1000) t = time.time() @@ -58,7 +57,7 @@ def main(): # exact result t = time.time() - eResult = torch.matmul(A.t(), B) + eResult = torch.matmul(A, B) print("\nExact: " + str(time.time() - t) + "s") print(eResult) From 5a42073cad6faa6cebec511e732ef376e69d2008 Mon Sep 17 00:00:00 2001 From: haolan_he Date: Thu, 18 May 2023 18:32:17 +0800 Subject: [PATCH 3/3] change to sampling with replacement --- benchmark/torchscripts/CRS.pt | Bin 2614 -> 2614 bytes benchmark/torchscripts/CRSV2.pt | Bin 3072 -> 3136 bytes benchmark/torchscripts/ColumnRowSampling.py | 2 +- .../torchscripts/ColumnRowSamplingVer2.py | 2 +- 4 files changed, 2 insertions(+), 2 deletions(-) diff --git a/benchmark/torchscripts/CRS.pt b/benchmark/torchscripts/CRS.pt index 4e482a12d2e489f1ebada521bea7a1ff3758b9a4..b405bcb921eef0ffe38e3d37a147b179867e546e 100644 GIT binary patch delta 1731 zcmV;!20Zz;6t)zw0s((ZPxCMkyz?tooNOsyRbH1!A5toD<0Zr`iYzzTfK|s1wgZS? z&l0;yTf_x%$zyhRW_GsL(tEhqvc_kt8+13e88ou^;`tGP5WRG=!+_2U0so<_Msz*2 z?H&}IR7RO6$YE`*CZxd@TI=4VYnVsf<}?XWO z6FB_4OmH5DmgHe1UIbR~5i=Qd*#vv_fvi}@l4=u*FO1sQPRok(a#`_63(!+K<0op5 zThR7?(L2>G9;EVk^V>=}^8(A_Sdyj6w*e~W(y{1XZFN2d*JE591#s(RXT`M2T zyV0lAWxPyPWl$!`Dh?IG%dXehOuSUbcP5BeBw{CtWvbVxG@Cc0V{LdcaJqByxnq*p z={!4FBkf9}w-nO$ISL>6T<+y6PRZt3B0WDKUCq&4{30ou)lbo!ljcQ=2E*XY_INVT z{pNjXn~8%TeF9KR2M8Lp(EnGvk&ZA+f0R#DWDNcHJN$ zgxE&xfJGNbED@-SijPz_AU0)(1gf}i#))TwQpvLCymQaF=iN8Y?l#Bzi zzP+R6vLr0Y8@A`TOJwc*!cmr(w!fHFd5VprvVD%+Fnr>d4N@{KEHWLNI0q!B@+Q_H z3vy(YR}%%{C7l*HU>0OGy!?>SUW?2T!-E6$2ZacF8On9yxFwi@3UdXpm*JMJLpf zL-V#_KP65VdqIEtotm5mvar8tJC>+-Ddx^M4wW1ddoUFf^cWVYtv70HVS9uSc z2y^Y6<|;asnq}8>;R*DH6}du~yd;OdTkpF>vN+Pd@DN=cqakbQemPw<%!+2?V2-d^ zKU+C(5ymZ7P?EUV=qF9jR-cC*>Qa}O;B!c&dpbe(@Y^^j{S9}GjTMLMj-Lk_X& z&|uxMsx{vT-ZqTtns3Q&rCJeRGQmzv`!KT}Onlq9bGnkQ9`i9a zRlifX3rl~~NtT4yXOTR2O#qoYY3@F(9^OKlTDmy^XG-(BWfqvP+V)*tm<(g zFZEMr`i%?q(`-1d4`*4}kOnisu%&@@aRL(D#Hd4F<+Bat#l0{6Pr0G;IW}Kc2ABQ3 zcnhx&z7hT&>gR($vWVU@-!^3K%d8EF6?xI@!)t$h!ROCtkSI3(@d{P5cvLae{hC=R z70|ANwkp@OJr_j%k>k{;8lr!Vp z`x2fg7+pZOpr1=E=N&VDK3V}(9J*o!-q@F#`-iLg)5Z|m?6g#Cj49A|%t zYyU|`4(s1I_avG(nTl~usR-8==lUf4UZD>W4hZ^uv~*Q+3rV_+Z`L55dnz^s8fPkR zMwNHt+~-m5L~^&1)JASd&<7FD=Sw2}A5cpN2tzNA(vbuJ0IvuD0FW9O5dbg%LsC;O zV{dMAbYX6Eb1raeY(_#Xk-#mI2;2!Q5sYJ+0RRBkljR6H8$&OT(vbuJ0IvuD03QGV Z000000000000020lR61Z0ilzJ39lY@GrRx* delta 1714 zcmV;j22J_46t)zw0s()^PQx$|yz><+PL33+^5`WJNU78tFClJGWVy+PT6OGTyHN4< zY-2ZRi?|>zi9NeJGdr%7a1Nf8sBmRejrPIR8-**R49bhNFkmb`*l1)VMolzS*XqhKTF)foo3(!$I{SlSJ zJ!m>N>8xxgFG4z8Y%A=%G@6~~&Y|5M7uJ$BI6sk7o~={R|I6N4vvk9rwfu$k< zx=eP8!H2G8g=2qjyHQ?h(@LRO&yvDAbn-Dzr)Y+fOW`usE~zH8HePsS_tN>|sdp$; z5#FRMQz)WX8HRH1MB6E>2VSb%5eDW3he7e4VhJxY{ckdv?F|NoOW zlgvzI%yf^q5SnDL6DJ>TotEyjB^^6e|7sA7(iaQl_BM5^V7cO)V zUAYq#6-9(`BjUnYD6gha{MCo5aqS}se% zlDuJij=MmXPcIy1iD~=uS(T^QI4av`$(rF4zig0_XcROC*=} zwRx2{vtGGkFAytO^vZhGP+I=!X2tZ6FpB59=3dRmG%8nFT9(KH5Br82a%{uX?0J&GKC=V|pw}vX zmeg!>-*GLKcd(QcVX98H5q9u&t6DLzv1ymgqT!K!*RhDJ+Xn`@ini#4T5@3CFzoxp z>0()b>33{$9LU1{rtMg!UeSW6IP3-D8r6zkG;p&1(yCtZ3>;%OQ_w)=J!~QjwtXC| z=u~Q!T@Qw@pf{|@6~gExIrW`-;{}q%nf8Tm(bX{fD?W6%l8eu(Qv0rm-5^7XF%5u1x zc>>RumnrK`MYjuOBlD=?9r{i+lcSVu2lplu?8LMWGAqHvx15{DE9vSHA7fMXONG0C zur!@yNqB$e$(3vJ{^0c)XMI>hNIim{mX(-n;zXN2Eesz~JIRLPTDxXdj|zFIpE}X6 zpQ)c_!*P8$%fh-emhkl-pt9r7wa*HE6_`~3fu>nfjRb9H5K+fR!(@&4e6 z@OM!^7yOe&^q%>=A#-0~ZAdK2n`R$>T;msgeo2EwvGJEzsG7yYilOe+%u1<%b``W$ zxu)%EhQ9>~NC-Y3{K(H|Wl4YnNk~BxG(!utLK}X!g9;tc30=?)J;3a<%h5p*cZ3@Z`Mh|(~l zn@Xb#eLBiKBf%T+j(}D{46kxni6X5`L)APY`wq`nH6(75X{C zZb5Iw*$*VVtI&TD_6Yisgg+_tcZ9uyep-jWA?y?MyEyxOT>ED-@;d&CbN8Zo|BiG2 zM7d-t#&t;el|p|+*e_^06)AmKa+i|yK|DJHg07}wQ&@|$>rv&4aqit9D0fS8g(P*5 z8x-`*2JxQd2KuZ*FsRVQzGDE^upXMnW@@ zz%-Hw+zC4AX%LkG007aGOV0000000000005JbkQ@srjC~b0 zCk|B+2bniEd*XcdAtwt0fK3dKsGYIhc6+`(#vrw)@abz0$DMyD3N^IGGfy_O#WaJG z4X~5%NH$Ryl0}c=GYK>@p=>ElqLg;F09EKtyF}&i1e(gtDl5y`FCiU1yk%HR=qA_aHPHXev(NXCR83+^I6ZzjH*k`rOsuBZjc0?m)ew4OsyF_ z+8R6IZs5M7HI5`rJNu*Zc{7J!7K`O_y@-~dZ&x?#_3AcSpH8Y!lk?0pZIxj{!J5}F zSrEE+Kr&oGDK+79I$T(gj}wIoGk@X8Ae0=uXfNQ4Oa6cE17$hRZadOs&LH={E4&-( zXt#tfXo%oTNHhtx`zAIR-ULITwei9u`6iuDezYSbRMI_~G-Z&4!qY*r8fpj!!8cyj zo(E4B%)+g4jKMu+!+YHImF@p27Yuf-u;p6xwbkeJJ%cU1`U6l) z2MFSqoe>2O0ssK9lNbXZ0V0wMA_PW49srTi9)FEiOKclO81`<`G~T3nHSg!7>0_Nb zcGC1gLtN@Mr76<^8oMwh4Q6d`>}~9|_hFirhd@FaI8X^nF9;zaaX>x69XKH1hQz57 z;s)Z1fD0fb#6P>vW;Z2Kq{O@5eE;{&KmYtQ_G=>T4~0UHKea@X>CCg2(u^W3(rIOJ~1<**wK02{6pHiskN(osxrKITvRDV`r z%5}7WWPM$Sg0#`e|CuGPy(qgwQRF^jR;?^YK+sNOT(IDaRSNPHspDx?)y zXR3D1)Ps^N91n=-rxze1eifS1iWikiNzpS!HNK!)Yj&j?PZEMs&S8tE>VM^rpLf=| zRJpk^@{DR8i&bwK19YT*8lePbSwVBiJ$Ml`b%q~2D9Ir0gcyYCDJd}?33&{8+VobT z`wH4&+7r}5>8rOgX|5mf^o!f-M`>?R?>16qksFNB(Yir-`wWnChW!p{$r#^MZa?>d zhm?zwF+mU3mCiM>ReTxF_5d5f}Y+1>cWoi7Q`$2~u~T9C65cYV{IBfs@_;eR+o#~*ey)tmG3nYf=|;YY=4@O%fjX?KLn^B0g*+4TmdLC@(+j42_z4& zg^^YQg7BtD9xcnZgG98cpJUp3E=>jEi(u9Xp)+q^1}Hv#TrWY)7e z1~#)`a~y0=fPc+NusId5c?@DsG1ykybW*+*23W(k-QHOF>)u!-fi?{ z{2}K)6CTMIz>PEVsgJXbB#ZqAP)i30@Ge-40tNs84h;YRkQy5k06hRhQd3qkFJo_R zb97;DbaO6nYiveBIRKH-FOlF1lfVhy3F4QX5f1_Y0I`#&3OpO|E?A5L1^@sK4FCWk z00000000000000003Vb73Qhr%lT8ai0u&09j|)E$k_rF-00000P)h~}00000-jm-8 L8wQdJ00000@)ikh delta 2201 zcmV;K2xj-d7=ReC8v%drO%MoLS~-L%QX#1%&DWQrWciVzryK;4voo{H<+4!1IrynW zfsaO&Xuq4XRY>|DO|JkrZ-f;!dbEynxPvT@=(eq@Gss&ZwbVOEVQY*ENQo6x#=Zz! z5{IfvgUlOSJaRsHmy-nnz$%5`sGYI>c749xry#Yb@agLSr@enD2{lCHi6<)>G0mW6 zBkbfGl2y`&WHq4pL;{VRP}Y=YNltrPfF^V&E>Sr=fTnSi#>#r~OGt+guLU+<8cnWO z&Y|s=57vg#;QAt+@?_6If2_L`b)#@Rs1>>-na)px$L^H!ne-StG^VWwg!U4CWLqp` zV@)B{p+!&D56geLRf8^Ju&0Cg{v?ywoltpTO(z4ZFsiLJm%5Y%x=|8%UTRZIab_JQ zo00W~);N-NjT{cir`-g8na$?&_7XHGKK`1tO(Lumxm;B8;%KU#gwQ-oqltJnLGZI-ES0C6yKSaj~z<+YzzIor3kZS zyW3KrT3e*0K*{Nq&@Fl+mUVV_+8w$(v-7b)F&bk`EXgvyXk*b+O4mkXZd6w8_0hL>@Rj#2XzTV|t3l=LNXaXzO`jguf` zK{s8?Nt5Y~3oGNqFul~cVmGl)R3-<=72P9VStli*1 zYD%$pu>E4iOp~g==$5sbuGl-7ES42}mw$hlq~^N1;}MfgPLk;$=;rXEV((_H5qgD8 z5L2;}tUE}f=2!)-U{nlm6;p9N$0*c2%to;lHWY-2W!yM8Q8dQQ+%SzL%c-ivX)=-` zioGY$xLUPV(L;@+nI+XsC6mHTU()IoZ=BQvsbaS<2{nqnmkq4#r!8BiRv~UttA7~g zGO22~gkUBTt1YV|nPHM`EV#Dg zsVii#p%kM1tSiXJUy9_BBkGEu#{o7iB8pa}UN!TESu%>c8}3oXV;ZJaHMB~Y*Gif= zdaa@rb$saFS%P3bqJ$esDHh$JC>_rwAkhlOA8jOMI7sHn&5ueO++{ARo&4% zU5zIoy6mvuxXEPh=9*#=y+;(gjg9!tu%4ynVWMgb_~q=55*5l7f1jL6==u5`0y&8lc-zN}|wba%zGYS}4DNq^3twrs9Z z|LjHoob%S*wUyU&=X9pFY9D1ijd8>hR8@_PlVI{vWazw@JS<5kJAg5SG*Z$sA`0rv<`;y%C)Cews_$j(aedH9`B`Dm$9fwE)y*>?{t!nU78JXG zUAg(*$DUKpD)s<7-cb7QiGPjaIUMqvShhvfANN}+qIZq^y2$$lwi}5F@j5V<&)X;b zq4bjxt$>D06-)Edq_9b1C)vKJaRG-N0Kl>d&6 zxb#m7n{;TS+uN}j+GsQB+hoE+x%I+!e6s$xMCcRDkEt)Kg;a39;eXIky!N)?IKuj$ zcXafkfnNalnia-={ZBJ>!(+rg*xYtB3}1U(TZW%uC*zN0%jmOVbffxPMxP5)Su66E z;bUR=S_QWZf6=pFQXx@X|1FTO8O4=~t`zDjU8Z|1aFLhzvgJ2*Woz{l_jhn<|ZMVWRfqJQa-*3Iagw7v_y6=mKM z;4}D&!!Ayj1^7~c-JI5=%r{Yk@1x9P0sf&9o+ioZlW5$pG3qx&bXKC*5h|SiE8rUv zEg{^)X^R|VTjjNCo`92}LJ9`q6r^DYGLVH~7>QK#JbhE571XtGstTKTCHes2UQS0t zo6~4BiZ*A^=6@{OoI{(jn9VB~Gsf}x(%HDlRkSu$m3lMYx~SieT?a z^lgL(I5p#}71!DV{v=`5I2`2kYXP50bO*M@Ax^)OBei@S=YEKCzX+~Jre~2m%<1n@ z?hnE3mGKZ9+Bkh0<(|bkXo@t}TvLq82{|dFP=rfM>72aVIQXt;_OGwk&HhV+_y45C=LV% zCc=646cNw-2T)4~2r+N+vG)Z40PzfykO>|E2$Qb~76JqSll%$d2`RYfgbM-y0IQSd z3OpMzZ}PGC1pol?3;+Nj00000000000000002`A&3r+!tlZ^{N0@MkU&kH{hQVIY7 b00000P)h~}00000o|7L88wOGe00000qFf(6 diff --git a/benchmark/torchscripts/ColumnRowSampling.py b/benchmark/torchscripts/ColumnRowSampling.py index 3d68d97c..9e7d57c3 100644 --- a/benchmark/torchscripts/ColumnRowSampling.py +++ b/benchmark/torchscripts/ColumnRowSampling.py @@ -25,7 +25,7 @@ def CRS(A: torch.Tensor, B: torch.Tensor, k: int): probs = torch.ones(n) / n # default: uniform # sample k indices from range 0 to n for given probability distribution - indices = torch.multinomial(probs, k, replacement=False) + indices = torch.multinomial(probs, k, replacement=True) # Sample k columns from A A_sampled = A[indices, :] diff --git a/benchmark/torchscripts/ColumnRowSamplingVer2.py b/benchmark/torchscripts/ColumnRowSamplingVer2.py index f4b00c41..ef67b4e2 100644 --- a/benchmark/torchscripts/ColumnRowSamplingVer2.py +++ b/benchmark/torchscripts/ColumnRowSamplingVer2.py @@ -29,7 +29,7 @@ def CRS(A: torch.Tensor, B: torch.Tensor, k: int): D = torch.diag(1.0 / torch.sqrt(k * sample)) # sampling matrix S (kxn) - column_indices = torch.multinomial(sample, k, replacement=False) + column_indices = torch.multinomial(sample, k, replacement=True) S = torch.zeros(k, n) for row, col in enumerate(column_indices): S[row, col] = 1