From 06a7a2ec9fffc5266af3d5f30d614fd2372afbe8 Mon Sep 17 00:00:00 2001 From: darthnoward Date: Tue, 9 May 2023 20:52:45 +0800 Subject: [PATCH] add script and pt for beta and Co-Occurring FD AMM --- benchmark/torchscripts/Co-Occurring FD.pt | Bin 0 -> 7548 bytes benchmark/torchscripts/Co-Occurring FD.py | 115 ++++++++++++++++++ .../torchscripts/beta-Co-Occurring FD.pt | Bin 0 -> 7903 bytes .../torchscripts/beta-Co-Occurring FD.py | 112 +++++++++++++++++ 4 files changed, 227 insertions(+) create mode 100644 benchmark/torchscripts/Co-Occurring FD.pt create mode 100644 benchmark/torchscripts/Co-Occurring FD.py create mode 100644 benchmark/torchscripts/beta-Co-Occurring FD.pt create mode 100644 benchmark/torchscripts/beta-Co-Occurring FD.py diff --git a/benchmark/torchscripts/Co-Occurring FD.pt b/benchmark/torchscripts/Co-Occurring FD.pt new file mode 100644 index 0000000000000000000000000000000000000000..8f75adeabeede78d9cc9db253b495192554a48b5 GIT binary patch literal 7548 zcmbVx1zeQP+Wt}!im-G^cS)B>cXz3D$1V-hjl_b0NQ0!5fb1e&f*_qsC{jyGEhVee zhwpp-=lI5V&iQ}y{AQouvop^%_w3wz-`C8gqmF?I0N~&N{!%CacmR1YkG74CkEf>- z$bnu-;l8bvw-v9ui>s2{Er8i?B@Erh!onNuY2#>Nfv4kYWnLg@0Vl zdOoi@OhH4~V)5#}0Sp2gvTItq-5V7@4h}t8jRvM+db*Eqgc|u2=~M&73xKMIY*Z{9hX{;j@vpoG38l*)DNNRulNQ73fS9o5I3do zew;aE2SP*#$B#kUFtFBbhiD1|ih;+w92*nQB?lh!4LTLZ$hjMB!K`2OJyVQ4y+}Rv z6cp~H9ZA?gD)D9Y6M7gdJyuXJb8I9ZY1BVXo?31syy&a+iBp+C){wm^PGF#2WJ!Ll zlw=^1Tblr+F_oZ+eD@f1PAr%(f}34t6DCDV)a>=z76G(2iqd~;JeK(=rewlDYIt1< zN9bN)X;fS)Eq^n{t;7mhRVm0z>oDv&OF@no{FF38-=*dOB@T zPROA&u{Aiz)FhW+1BjFCJ#j#h``c!Si;b5_{?Iz+pq%8MnF!9=DJZgeK4f zm^uU*UH>$XvdguwH2_qGI?KP3l;2v?Of&+0a=%sVt8HFV zQ3L%H0soc@pMn0MDvtF@oB@&7+F6uJ82j`e;uYZ^2Q)w!I}PZ3toC!kOr#+u=VVEf zsWZ(R?I)S|IPREP{XK(UFqW$mHov{;FLWu37lm-o@34pk6MLp>n_3cSpHv@n0%uJ6 z)k68n2I|ikPv7t3VpLNSgfCrWFg4qM93ql|qS|>zi_dlnG0=>{uIgC#pu4?|mM^Da zw#_CJ^;}z=aky*ectcINYV7y1oKo zo+dyS!~#y*Nx8irg+f*fg|4Pi;Y{8?ezYOZ5v)_$+|!RO@Iy9Zbe%6o8L6btXKo)0 z(IZ^IRby(Mn%ZF2z-EGEt>bZ$d(ZnNg)S?`iHoW@4)Pd9SyVR6Ni`+{>Y9!|=h*rY z7ca-u)tY~KfC1&^Y!mattOQ$Yk^AxRBoc~)!JMtR)XXGh6 zHsd?N008*z{JY4*1rYqtk%!mT&f3Sp;V@dJ$A+7(sI&DlG2iPGU4`qyZ7n|#`ayQ8p}CJJPfDt@NnXPJjZj34%0)) z0Ci+j!sH4Ga0|(z9Ce3^VOY4p#nOik_P)X3+}NVwtBugn=CgdSgRGnfk$sghF)=-C z%2w{&KF1SpQ>M4yP872xmGs?;O`@x(_QqJh^}v$5ryuRL!RT5l9W2+DMgN_BZl*mp zB(L0dgu~b0pHPNw*CrpHXAK z^0H*637U2fn{6{fIjhb1p%l^SBd)?W^`0&Mvd|tS!KVaoLquM9pETJkCjvSqcFQM} z)LQiThwf5q7yvJT4sPz|22#vG7? zl2f>}0XgXOjha{MB%w6WT;FOvpzoz#U@8iiU#WsqT%ck z(tu8(Ko3poS{Hd?x45@)LfQJ9`h&)|Kd1#4B|VZ4E^;AqXr5L-CXEr$@;C8)m?2TIB^w76^IfK4;%B7>yRRv(9#wr=Cb4PgrW&fKm)DM za8GFpKXsZ?ZDvu;dJ45%rT#!n3>(Zu=DJHPntQjO&j%UbZ~vf|VvRDWAUk1T>aI6J zi`|V+b1;@PeOp!bG^b(;r|JuWx{6#pP?h62yG4?B;8d+v_?`cekJ`sJsv5zT%UD2e_k|%vwhY#d!%j?kf~j;|lNTPmO-c;`(a(7WYoH8x!>8XQz6OeH zEs5}|-I%fi2RChsz>mBoqFa}`S!VlXr?m?pY%XOv5a5`Q%f{QnQ^EEPLD|5nTOKp> z0|BgAo$DLsV7rF~8=x~n$8Tgzkd?WHtH<}nb}Z)(Kh2-^^<(eE363l}E3`ZH*KIE> z3*Q^maw9#Zzc($~NEw$MZsDc2OzLFQXr?tv@M^wgV*fSL6)(7QYwF;2(6rmv;<+gw z&Z3=g;dpARF?nbsHMI*OeFV>bvDXvvRz(}(iwx*t83ENAJAvq)YYuYfib4Z|Efe$j zW=2N(`n*hQU+=;7p%0vi7l$ha=V1t6ptCGYe&Zm7J9u&{O-p~Ig{;{ZR(e|X+7$Z0 zC2+IyO`{oXe0m9<1uN#-W9?qk{i3;O_{cUG>p0t~@on*p)K>%wTM13#jl-(<#|3XI zWPZyg>h$3upukYl&)))4zilw22v5uY^4w2FxOH|e(l^hc4s0ahwXG-hL}Pt;)C%g^ z%`xKL(g;b)G8L8C?AGfa&?qzO8xez$r>qTEw@>YT_2%!+B{>n=v021t9+mi}kJ0Z7 z8@e19+|XvtV^*gL^!wtstEPOw_G57?M2=<2x`9tHhl=U?q0e=S%a1ckF2hhT-D@TG zF7vZ8Sv0)a{B9|m6;fcIP;2Pu@h6*uf=%(Fd5rz<4BBTcUZuw#RjV47pI9fZ!`>R4 z)|xDm(RzQr?1-0-o&ys!V`s>juP1Cq_c>xi;4r2-@_+;OG<$QuPQNeZ1c3z?KVBXt zA_Z1vzCNE&sk{ys-#z2~Kvv7;bIi2=&Fk1}Z$C0zDG}A?mOb$LwxDCo8+e0yU4)^x z^hkE_^xNd^HU>dIMXeE&tlLlZA6PUTnM!79(F&H&KQ5Toba)|(yk$k_M|X1OnS@EZ zEK9p;MZ2QA_%X*n(WNb`5s4MVV9`sCIA;msxN1$jzsZGCTO2e#ga_eUBpl2fo(5qo zeWbhsA#zam9ns)*QF=v`NX@aBRO7L4ewPEb_K{YQ?;&UrVYa)BFZ9N|Lusx{dN6t! z4ehPWUDWIE#+U?nCOh5o5wVysV#zK#5=m})2CjSc+HBoSm_Rga+X7ir%GBT(+-0TN^J=$41>H(Zl@VGy)# zYOJ6qcaJV*v}W=%8ps{DIuHG@z(mNF>I+k&MWrw-G?QosYjkwkfjW+IN2(PW07!Wn z69z0p=Zl`>UUt+Og!3CdP!%H!^LoaE&O1;Bt@MpeUTc==l8;If-H=0*e)+J#u%Uq1 z46P1Gj@!hYb0rVI+F^RiOIF7kbVZi)kk)C>{4nf9d(Vb} zeAsJ;j#ufRYv>?1dE#nVMJR%`a-iFtaogWeBv=!k-IF>)WCV+j62-ETU`OGhk|YzMi41$z3^8 zkA|>V+(?FJ(U`OWsEW3eO@lp__d~M&1P^ENP_HG1Lk&)!h>u?z@(>+9U#%e~WP3~; zA<}z_KB~Yxb~E23cUc@Wpy=`>;~@EWQ9Kl@6&NQ5m=$JG;c@+{ zGBlpuCs5yTDOEZ65&Vis*SSTk@(z@{vYgAnC%y6;(IRmKHlHUotERAs!717M+L-0m zXE~S)HB?K5m3~AWgb^7)Q5z_I5DO-l9t1=RzZaE9-d#09TQY(lly_(o_x02nqag90J~?R)Q{Au1&bN4Ik;Lu4S52!xO|k+c`q@3cVm{AaPT5u{AQn2zH`bs;{AIBp(EEFh38ty=N?q1tR zD6cDd8F4^bm*B2SBEq=elN_XBN^1;9R36~cH8OldkcJ^mNVJ&t`%imII2OtEM=yHn zzeRmLW3~zQg=p!g9u=|HM#Ey1cI6`o5X}sa zCvmt439K{?PnYT)0-=7l@2z5g=acRwnB|4*S{3vrJe5WEptoduu^sw$s3V#Db&;?1 zkwEVH>RkFcRdi$ud@!GfVqIb}9 zB@YhT)E_m`?5Q7#X$z}!ID!uK+@q4G$(A^G^--P4=H1#jj&1{a?+x*BLWqtPb@U~0 z#x;@p=&ALS^8%t_HcQN@2N*v&G$yyittLDGpo!w&acBVm`+wlj_kTM-xNq+Rvhj8T zgRESC+s-rpS5$gLuNxG}cc=ZBREms4e*DYxIgkFm$w!8!cM_=}tclF$SK`dBd?h~Z zq_`2*=k1#gjnimJwqMG19y}0hO^fejp;4iH7>|isp;LFm1)H70t-MKUZ1M~Wgs9?I zHd{n@{EPGI?WsZQrBSopYHz+?v8O*Yylk7~J- zmr9qUV;-~knq2B3!INj@k1tX=XC|XE80puKmcs!H$LH*lMk&L2y^fOsQ)Dnd>n!R3 z2_^hGDeXnod0BS?U8+)^RnXT!Id+F2uxM*w*+!n3t209My%t%iz7`(?IbqJSz1jTp zl^pudOqfMmz>LvjB<6MQV5*zEh!;QH=iwa%SF2_Kx5)N~T=SrzWWmKHm`3Bkm3sf! z;$BTx=>Yl9r5p5_Iub+!06yMqf&TY{7fArq|6%EV-4D_Ir$u~Y{Mqxf0)=lX`kQRb zPkLLxmn8AI6256Lngt0V%}3e1&J3DX4ic1)Qf3bIOeepNTh=rk>z}-&+;3jSYO=p& z9?r#U8N;@`^kg8#w!sNalhH+cDdgJY`Vy3yTW-MFxbH?$m9tknBMNRj0E4QRE@AD> z0hiF6wjyGR8YFH#mbA!Tk*`d#nEPjQn(rPZg5-Ic$T`os%s0> zu$C-yn>ovCD_7*WeWW^2Nfny4pskbM;arX%T7J5eV!@zHycSjUNd|#EfOkhoQrh{X zP3p_tEl`iuXZ%pYQ%S$(j;1*rQw8#%sQ}yk@j8anNtIJ;37^m|f|Q|(_Woxgy3S;Z zO&pr&0;)SMkC_VWgyN1kHA|^@IFFOwzv?4MgA1sPD+nIo)+x25xj*CfTB)&Ia}rQV z7ZDYG-FGcY{!)wL{BCgOSGMcc(nHD$qqo;FQFyx3#VUs~(}lr!T_1M_K?AoHJ>8xl zpuHM1Gdmwo89!mlpO`$;8)87cGi_AIkdgh;wNqS1;o@5oy%Q(FJ%(RFs=Uj;*G21j zlhO;3b`jTYY9I&id+WY{xQELFKZb>wsL=8K7$Rd`2@6d5qG)rpG~{1hQWe_3W>;n0 zy>D$_xe$P=8LikV(De7#KRNFnV}B_8LhBQ~7~NZ_%`;<1C`u(Q_fBxy%i1sOyha?S zG>fH1k6XadC1;;+)^7hU#nb*fuCXyp^4`5v032WIZnP)~wd!G|Gb@v7%uaV#8M#8l z9$qX_mc#M~+0hpfux-p~7uIoh(z_WVn*t$RpG0E_#i1Q!_uTptL$MC zg$i5|!(}RI&7_Q}Xr8e(1t1#w0FPQCt|ngz4&B?F1NjqMDddP$(FdF@N@ zaFr@2u{a8)l#i{cwYzs@gP(I}BIzK-xh%cVu#L@h?U6C2VDptRws<{^$wYMi2}Q@O zB+A4f4A|w}W;660F`Ea6F!ZZ@oafX+Qp)>@c`lmcuH(CF!L3X})k&-?d?b4U{lJR8 zsLXn|G8y5JBnq1_rfUt7vRHe%we2h^(zm6Y^qK3}B%^3q0|JFbhkojO%7;*kx)V(c z)FTOY!$3vELCPt=hPKMPAvq44Z__FD0lZ?aM0Bip;id9tcl(cRn|sMPWAfc1&eSv4 zG+*iD&!Pw5E=rN9&}gCR2fY~1Uw$nbjud>OhNPgS;bLs^8yP+{ww4v$2g`e+(S74q zr){3hKQ7UA9@EIy?QgUcnnYp=bf)#>T9lXFUZ&X{p*34o4_-(;OJ<17T~>sSm&=8`XUIhnJkoirFq(^MZmoFYk|X#F7~GJzK^LW~2M08OR5tgDV4S*5i9Fr;88a zx;`AVI#Z{**wHRT?Dl(Dq{VGQ&AmL~)}X|wk6-)?;tLORfSsj?_|jJ$_PXtRf;%bl z4SSj5BHYhTIPK*$ z6p0jMY#BO!P8LGFh|?Lli4?!!4IwD#Kd`xHd?w-sOXGTTVD~0p1^&@tkal zx-Wl@v3$XBg%f&bNizt`;sZ}xgmm{TlCvXe_Ep|HP>H~$g!r0Y;=yYOHdMZiU+5}( z>kEHKrbtLO?drKaqKjE5)L>+!(8FCoG(bh@Y=$;}c&nJ9FQh=^DyD)UrZ4ScT3YdV zlM>OG#c<{6=&H4@RG_hkrbt7ubA6NQv2nJ$W-vk60Khw<-&6nJt3USu*ngAyL0;Zg zAaAeVB_NESxjhV><2NP8?LMyYjgSQRS$`6Wzegqi8zslr&eQ9L_5C098zH=(b0WiG zr9!_M`DfKpNBepEr>6C*#&kne|5W}rh3j{ef9X-bo)CY&CQK$Y4_bYYg z`U4y$)_(!_FN*bN;9rTIz#o8B2>%oC-|_H2t`GG&GobY?|uIV*WeGr literal 0 HcmV?d00001 diff --git a/benchmark/torchscripts/Co-Occurring FD.py b/benchmark/torchscripts/Co-Occurring FD.py new file mode 100644 index 00000000..7272058f --- /dev/null +++ b/benchmark/torchscripts/Co-Occurring FD.py @@ -0,0 +1,115 @@ +import torch +import time +import os +import math + +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 + +def attenuate(beta, k, l): + return (math.e**(k*beta/(l-1))-1) / (math.e**beta -1) + +def paramerizedReduceRank(SV, delta, l, beta): + elements = [delta*attenuate(beta, i, l) for i in range(l)] + reduceRank = torch.tensor(elements, dtype=torch.float) + return torch.clamp(SV - reduceRank, min = 0) + +def medianReduceRank(SV, delta): + return torch.clamp(SV - delta, min = 0) + +@torch.jit.script +def FDAMM(A: torch.Tensor, B: torch.Tensor, l: int): + # if beta set zero, the median is used to reduce rank + B = B.t() + + 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] + + + # shrink the singular values with delta + SV_shrunk = medianReduceRank(SV, delta) + + + # 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, 200) + 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("Co-Occurring FD.pt") + +if __name__ == '__main__': + main() + \ No newline at end of file diff --git a/benchmark/torchscripts/beta-Co-Occurring FD.pt b/benchmark/torchscripts/beta-Co-Occurring FD.pt new file mode 100644 index 0000000000000000000000000000000000000000..80048d964d42dc7a2a06ae12d0f1e7d9de55e3a8 GIT binary patch literal 7903 zcmbVR1z42Z)*c#Zq&tR|7J;E#>6DTh8U}{$kP?s(0g)O)q+yUw2?eA>(xJOM1{8r` z&$<7-o_p^(|God7=X<{Ad-mS%S~F{XYrnHzEmbsh000{s@GpZ7fB~?y@dWe8yYOgQ zS$Vm;J3#Fil@$1_!Jc4VSBR66916fV1l7vi+|$L~%HG@@N6QIpWn=H+WNqWF?qcob zWP{@d;0z7oF#HMc=?->){(x{piT1K|L*>;{#ljW^fY-%-%>BHyet1Iq!;^n?g$*G1 z?>(_{v9{s+yLnz$KP5R{fboAdQVeFHv!X;nBL+l4c>wqu)g(5|W*I!sMenNJ+DLpL zqy2O52%RK#ZKERdMJCSY`GF2sma4Gy`<=_nOet5&R{Q!gu?W`v!Yf*3Z5xOG6NBWb~f%HYZck>pm6{yrEv{#D@M7qiDv(h?8wNp_BFi1?Y#J}|!SA$q*5rDYH) z&>GH3jmV|FU$xGEll8p$1LRp(ilce6MKQVMU&zRqa*TL0Q#RkK0MWqhPziQbm?B zzbE?f1+i@OW)Jy_G+yhmWkGO>xsB%d!cSoEX0ma<0TLm_hlD|fEGHHQeQQ94cH-gE ztibL$einjAR&FQfhz?%v1xn6R7@dG0028l3R+S!Wq%;c4&_z(7R$vE>^ee-agd@$@ zS+^b*1bJQ}`qW$Nl(R%RQoJ{c!4X@|&Zv2S*yaL>Y%9>&`a71oo(|H?{%dx?{hV z^uj0?H{sP3DC;Po+`YR!J(i`fKyBWl{X?wkQ}N_^JfRU3p*wu)0{FZa@~o)b zPw$k{`a0R9dfVwlyMJx7%Q=AGOUdoyN#SO9dDS4++dpE_y@NxH>@Ys%{^YtHS|gNG z^M2?HOn))*EGDdZ>qsa|zP1aoGiuYljoT}{XVr<$WXjB*uf_L};<1D^+xOfJfgvJHgmV=Up(RV){XIsYOqtwlQ)YxGBzp8Qk0QCQ`#_?L)SbEu+|5ktjfS^I?Y0&RE_3Kn?aMQ&BpY~l9{y5r zqwmOq!OGE`!Y#fq&Vw?EtNdIZ=Y(3QR|PbH?8Ua=;@sP z+ZeKmz$3*fxfy!BGo3z;nS4Gom0b+0`k)tQxem`r@0B~twCw4`tue;f6&n&CUR7VU zwx6`>&&DraeDGA8ow&c@jLFGE?JcNY-E1n9tfXAEQ%g5oD|F;J?Ut!Mn5st>)u#Z? zh_Mr>>9<7G6iF!uftq_~$A6%ux#|a~3wWSkm*k zus62+(vZ3#jt^?KEtXxdPO|29&Z*`YVcp4Z_h$T^`5Z-y6a;fP?K}#06rAl;ClK|Z zv~^68p;W3E(b)wI6q-6PpL5FdPUby{zU4NtC)UFx*GW*@vsfyiQk>eIlMLRBrtZn@ z`C{kRu`qCv>riHGV@*Ynf+qQZ?Pu4!b{jgM?)G&u4)67{O6m&hbY3J-O)DX-;1bWy#ute#=m6tlMAbNtf?ST>wzowv_zGA zdyQufZ1P7wEO6J~pSEPNu?6?mbNJUQuCkUnQ8rLX61jh5Q9MB((&a`&Py>|b4Z%sfSAt;&7y`gIhWvVGR6t6SxA-o_RYg7duM6Xs> zJWl~833&M1GC|3}A1ce|AJu1E3PptkbE`CsK+{CO{w?&Kbta4%w?-&vyo_B3v)M#5#)-2=;963 zE*yas4*k7f#J-e}%-JmYT*E6Nx{rxv*FNy2{AV&f z;%XHSq`ouKY3CVU)m4h`i`k1|%@u|tO3L%7^D(YAWDh$+WM%nwjL?Q3oPC^#NaSLV zza!aDY&iWGiFp#s-jE^K>TjuwbpNpB)-Gn+T7!mEniMmAI3EesodQopWSU}YUl*>} z5CP3xoiby{CSSB<8>Ew7w(Now^t<)T<%e4%G9|J?RCvXa2Wv@Gl21aMBVU=Zdd67~ z!8{(3?N*X|qH*f0K`K~E?O!pbp?KkJ8kvWPd1uZfS=Q+dKng-Wvw^Mos5$Gz7^3N> zy(`ql!XkBe6b5({x*aD83_fvF=u~#8BvsZLFg!8z%ba4Z4IQ{B8%vI;?R9=sHFHn# zX}9=W*!>R^;xClv^$Vn^fGfqr+k3^_)y-j#8)WPX9z{T=#JBHSWHu3XA1vkLhB1`* z51Y)Ah>g^)`AxJQ#6+6%ZdhyXja@^1w9BoG`M%Vdg+!V=s_%Y{%T~>gpOSrBYrlyL zA1*5@sN6%h?-qno^)|pX*#;Etb8wHiVkln2SBHr^4eqkAmF?umHjcdEaM?v7~!&}=isX}Rk01L?Q;e!5BGZ!WlBBzE7E&d#EZM#b~iry<=zlm%Or#UI+K!f-6_NN0MONZ%Pw*6>D=pIc(h7&H{}zw#1_p ztT{?;+zHp-#QHw+u}z3;;^DQBOt_Fjr;r@Pf$p)rq;oE0%ea}K?OhSg(xOsxA~Y{^ zo4h#%`c~ajJaAUlXUx=7#5r>D&b<7^Sdgc*B!Dy@Yuih&K-)i>9y8`MIdxPHyf(hR zbA(_p#%Q#KIMEb{M0jYVX!2jX20#sFQd@MZD*5vLoRU~U^kDA;pj%aR{l`W zUol4)JC#Td79T4`k|)Ud%|vFD_#-T+JJn2i@_yb)dGY%tR)I$e6jV-E8Cr9tIF4n# z4<`KVfp51cFR}Wf>2-~d!ibN}eo&Vg0}ZAL78x~>8W2SpazrFK?cq&%oR_c;5OGyP zlfxW-qkklTTdOpjRvh0SY!!X+&Ef(?%F-f_e6Op2NM7WVP(B}f*_L1fgh5yK-sNpC zT07fmje29S!l{mPbG)uV^Mgi3QQ93mS~u0lMn(nIe7GtQs+38AuApG{3H{*>qLQ3> zydB|T@^HC+DL#lvcGuhLZ=x#>(&D&Zx)2w{i^nyDfCdkwxZCz&-wXrs;!%w`+QCw} zFblpeL()VY6Vs@e(d-J-r};i3ij|&VBw>dsY$!GxxB9i>_tB9X)JB@uY?{lE>TeO0 zTQt?kS5cG==`}fwQy7N2A-W$CIctr#uO^>9%MZG0D8DhFJtr}`GJiBPE(vx>y9O&u za+GNs(aBAql>-tlWX7zmboD8S4+$#y&<@u3WD%EOqagI-ZNL@5l9WAp^df#G1K?nD zuP0g>r_un0&=nb3exruG2CWpK=+ms7G1@TzvvI23I(9{$khr-4w4P~DCQg_}Tk$G^ zW_ilc6VVX^AyK&q-yNYP@UVtLgydFIU2?S@p`}o-2I?o#7DOkd1+HB~zoSyNuD<@M z`&x_S9%fwUIay-qhPF|V9P(OtB}*2GZE?@;U4Sp41&v*M;-t^I(Y3z*qWpVO>jac~ zdBy(mOFeyvqv_&1RPT-^IN3>~*_7~vN;VB~4o)=%Kzy10^oXhh==ycQ@hMzr6~&0| zgt;Bt_f-w)%CxD(Cq#j<14>TH(WQjalAAMuGyjdm2vUyr9*8~^JDv=4h#rah!Ut#p1z;Z| zhyn;IqiQyH^1G$^%J`TgJDpV1#?5?-8S5hvcRFZ9PdP(!aBHqmG(UWm3uGLm$J$#h zH%(;r>2)mH3z+C+N4Wy2GPaT1x}Xwtyyp1mm`YkQJ;&~t&>x(@)JAIi4wb+inZ(>0 zXZCuMAtux;v2I#n-^(aKeqtf1u7OxYA@|%>tPAllx&^pLw8RgK_g98Lq>ZZ-=olY` z6ZZg}2wc&{UX~L~sSaJe@aN>@n~a6XI)gqQpzY(?wg2$wV;j0JIZ{vmmeF5WxRGyB zR!kvk+d)jnuQ#gc0l`xOYfu5q+4)gIr#VKCs3ey2#4~*jDyQmVT1C+q{(gjjHt2=c zbK#IttekB>--_6NJEFVTRCs2nwQ&ZPX%?BPX&HzL@^-ypxxd`(QVii!e6l zWld_>{WB+<V~TghEz%vGpUO9Qg^$e{f$NH3%qK2f?q=swI77tp ztAL6`+`YS=Nfc3kUtj$g4;$`)80aZ=k8Ud3sKp6~S^($JQ!tA7dO@u`9bBMbr{B67 zlz-aQ7}EIw4d=&eJtdPO=ae6vfBD6&XMgO;VUg#|ma7{OixYpoS1Z|_Fw3jf zO}mB(lo!_X#aiOxqAkfW?X0wqs3c?1FP9lqo$t69Urd5MNh_?f^m2r#W0p6YMRtAP zW>wqLK-UW*ra!2_cWyXRBp>q`YvT}VCbq6g48B2CDC^mXOH;R-0!>CV<7Fi=ywEa> zTqvTDk|Z>ZE;jj|#5FY*k;=@tezFt_m^;1VkkC&Y%<8fq^B*U#^|4H+@fTCVuaweU zP)5kQ5^7Ty@T@?00^~UC0$oH}0tz>>jGY`;MLuef7wBs6Gf@y_F4-C*UM^=c&H~Z% zwun=QPmj=VVf{(Y^1>bhlU|Z|3QpiA0JrdVH_R+>AmRSPVy#+3-;HX|$ijX_M?oLO z&zSXlji1~{0RVc@{wij&0Q&!L%zpisM*XMJGd6}B%(~BhqAFxzuL3JS-xoa=K=eLG z>nC+k_9mz;F?_vmxX_{D+=QSqU^VE_nZrb{6g9MY=igJ*>ul%07~3ZEPJFyg?w$YV z;j=s64tZrz7-`od5@q(bXAh+h+MmK$Hy!I{!r3+9Ls6_}78o!IO zt32DgloijZs&S|767ae-5QGV@!N z(?}A&7u6b3gyIQ736rytX{*{%-C2G z{cX32Aqaf8aZj-@XXUVhiTsuE`i&-;Uc?HDd!FO48TDcIcWo*9QeBs0>GMaXp4 ziFqW!hz3t!iMXFTYH4=)S_&5h<6Qn}^o)D?ye>JkaHlee&(izuX=*uEKrx*8OT-5X zj9b5CiNS3#-aWwGtBz#CriWJKDEKa{QBkx+IkWr`W;!BZ;?dyPO3s-l${e8Pw@$b4 z%?shg+T=07WPslb`m1k#ep?nD5+y#IZlwH~7oe_qCXK{a+Vik4l1srd4-*AzK2F9c z>NOz0NL@%A?1v{_IBYd?XHe#`PXlBca~=^ePB>U@3j1Sry0czB+XdW>F}alydo?Hl zl)i&o_$t4IMsO#Njd;WgK_fWJ@R5K3tuN@6!XS}duRgR0xbMfmgamv}mGY@|QOuxm z4bG(G9J6g?kmZN5p3~MArrFN<2lS)i9~1b>)T|b(Of(Zf!d(6gm7fg-nnTe&Kkh zmzrDygmjp>hZ@zM$i1pIhH9ayH8tvjCj1CdgI=C)4){tQt+4t_W9F#YrN+u(sjOLoUlO?NB{OWcpd2wOt4;qx4{ig_Dv}__#)RcdXc%bio>QzV*iT45v2z8 z5F0n2fIY})vnl#EG`)H-{61Xch?Ji8o?s)P{lR5fWo|0o1!GeP;Z;5Sg^A6G%~)xz z9q-WA>Jz~X`%7vgHG}zE$&?;98o5U8=<@jZc(vCA{W;o&PKV7Nd$n~E0$*rCG@0E4iC{lA$CXH*_}N{kfUI__E3oXj@v*EO{{a7l+`?pYIQLV z-ktuzy5#(jsuhVYe55Qvc#mS~DRKLq!-gZyK(cMZoz-wA>#aymDUF#?IrKc>F9 zko+$7p7;j9w|wtX{yTA7jPxhAidCDfcFbG+_hFewvmd7?r@72@;peqpXp0d@_KvT* z->GhB+h#Ax+v3H2(Vm7LE~y7nxX6k_<_E05_3Rl+RNn}>NSf2G&R3l;rhfqLj+{=J zZ}wa^SnYn{7P`&LoUif*w*Wl%JU9js*bJ-~Ijg|Mv 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, beta: float): + B = B.t() + + 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, 28.0) + 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("beta-Co-Occurring FD.pt") + +if __name__ == '__main__': + main() + \ No newline at end of file