From 36175bfc0f247391ac34af244c97503e140baf39 Mon Sep 17 00:00:00 2001 From: Lutetium-Vanadium Date: Wed, 17 May 2023 14:02:26 +0800 Subject: [PATCH] Added count-sketch --- benchmark/torchscripts/CountSketch.pt | Bin 0 -> 2782 bytes benchmark/torchscripts/CountSketch.py | 55 ++++++++++++++++++++++++++ 2 files changed, 55 insertions(+) create mode 100644 benchmark/torchscripts/CountSketch.pt create mode 100644 benchmark/torchscripts/CountSketch.py diff --git a/benchmark/torchscripts/CountSketch.pt b/benchmark/torchscripts/CountSketch.pt new file mode 100644 index 0000000000000000000000000000000000000000..71459484bf4c18c130d408d41d0a9c4ce5404b5e GIT binary patch literal 2782 zcmbtW2{_c<8XxlF~h8n{tV~IvHW0=V>%wmg>C5*B(zLfctY!#-A zC0ck0Wl3G#lzj=2y`<1`XTG|9sOP))dG7mv&i|bMIp_VI_nh-v^YWsRDL8_cmnf2i!QlyHk~f}$ zAbZnEc+nt$W=y1p8xNjF!4Lz80lqOhMvnh%I z^L3Ws$lmzfKUReVhQX{&)|m9ew*RG($wZ{~=t(8at!hWm@@e4gIlEH9lP|?Sn=U4# z;S&XZKaKD&DqVUJ!m6%sQE`(H3>ygXVfhL#fp5!xzV;ElosI&&L&?$A0#`Ff9fA76 z;|Z7Jq?-*AEb<;#CC$F!Ovkf^45RgXT-AX;m&iK*W?@Yu(<-GG6@+_hc@- zFfID2fvsgcr>Ei}>MMb&nhAeev|E{FGAw;C{m}3Z#DcH(pdfnJ*?f-79}`v{l|5o} z4dey_!=yH?Uzl#tV_!c_U)#^sYW+FvG z8Uv|OCo@!MJXBQg-bmpyS@zJoF1jFRIB(Q-+&iC|?3HUKZZ)lV%AaMB)W;zSC&*8l zK{thCx_sFC$D0i9xM|hMeT=sq3C&V6zfc=h_fqhar{+wxbf|7=Z-k}!b?{B&nUdQc zRv!`C=?!C5PU96*jf~Ur7JTi%c4$mm zHT=CxE}CU!VRQO1PB?BL`j%ZJ$}z?}sAEFwOxKLkVx7CQi%l`K`9oP+OX=9oXCk*_ z+jqE28qAg@6}rs{lMz<8HD)C&iWsAl(Fh2kgApCc( zsSc3+d)I_`W3Z+ zWBZP@MUaUkKh;blzC&1vAvPwVbwejhZ-~ZxP$Zq%8_KnuJEV(?NA|pD8{IC-`?Vv z=`dre-o!Di@r})i?!`VXOH}zBGI7*3wO#3%aMQhuA;rz*H7NAXLd1$td!9PsxtXN% zBaNy~`yKsvXOI+BS#Ra5Q=OXQ+AH<*mximetNq>ELkRb=1NxAOh*<~SgdwPiD|RNQ z>h7D`GyRlGq&>@|FYd@S1+Z(`-Aa*9=(e#-V>!Lg`mLQ1MTx%edzxE1>$f>1Th_R5 z&$XSpnQE07X5gU1vFl7)=}YeKwQ^x)6-nv#!y2A*7OYLeK4@DM>r%(Ww86&4VFn;nf(Z96G72b;Fk_+P_# z?DaN(Ua{Hl-Pz^GZ>!469s902RF{)7tZSSNb>F58%aR8vDaYVtwiQm6RY%w$feRhW zY<9=1u7&>X?QfQa?dT8fL}zn_Yx&`5j29>S6$xE@bk8Q+8#>9abt1WDY zlN!Os$}Ana!$BDy>fIZWqZW}avs!%#MGyHSfsoi8^5W~q58>$oT22=eh;cBY*ghDp zB2%m4Y}GRfg^^qbE&Wx&6gxdN{dRL= zFo-d;ww>InSR_r$a2Qq`G0HA{m=!4*$w1?reV%F>X*WbWlQNVc76q3L6?~*5duGKl zgdnwfmWB$LKx2Z3WQMq06jc0$6h3eVxZ=>2Ddqpzvg_-Dr6s%=8Udg^y-is9NBw3s#q$P?O~^dz~>RW%fhmy3Pv zjk)pqK=?t4e_V0=WX+<#y1pqd!u7PD#>uUHO0o%VQpDu8V$i)aoD20Zw}~4`iw7Z= z%pn{zQXIm_^pcs(3_RZWgsb@8FzXz!*ZASsZBRriEt>Bvkp9AE#q`3U-Mm06R2UD5 zVh+hM_lu($K(Wwc`D9yu!~S`(AmjF-T#Z%_&?g2Z4r;6=zQD8-^ZWuCo;*%t){8mi zPN}7%$2kYsH-l5^SFvM~HY-kh^hOYMQogS~2HX~WDXsemhBu8N@q2XxsQ#I1UqpBz z(6$+xoyBDbgPXpEz+dwOd4TeU5EwwEVFGB>Z)8F77eQ}LVvgq5{F4dz@80CfT?)V} zJSOn#wWK#NfnYp^N+bvTANk)=Cv(5nBK*W`HbSkv67_XZY6AZS>aR?+affTqUVa^@2A7Pjw{OkeZiHET zh|=p|m>XceUmFrGEUfc=;tm1bA1XrMTmPv~@O|q5as7}dTs;w99U;nl0|2~XBwUD_ Pr~v>W?!%k^ZSKDT5!ED) literal 0 HcmV?d00001 diff --git a/benchmark/torchscripts/CountSketch.py b/benchmark/torchscripts/CountSketch.py new file mode 100644 index 00000000..84685928 --- /dev/null +++ b/benchmark/torchscripts/CountSketch.py @@ -0,0 +1,55 @@ +import torch +import time +import os +import math + + +@torch.jit.script +def CountSketch(A: torch.Tensor, B: torch.Tensor, s: int): + m1, n = A.shape + n, m2 = B.shape + # initialize sketch matrices + Ca = torch.zeros((m1, s)) + Cb = torch.zeros((s, m2)) + + L = torch.randint(s, (n, )) + G = torch.randint(2, (n, )) + + # modify the random column with random sign + for i in range(0, n): + if G[i] == 1: + Ca[:, L[i]] += A[:, i] + Cb[L[i], :] += B[i, :] + else: + Ca[:, L[i]] -= A[:, i] + Cb[L[i], :] -= B[i, :] + + return torch.matmul(Ca, Cb) + + +def main(): + width = 1000 + A = torch.rand(10000, width) + B = torch.rand(width, 5000) + + t = time.time() + + aResult = CountSketch(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())) + + CountSketch_script = CountSketch.save("CountSketch.pt") + + +if __name__ == '__main__': + main()