From b7cf03828ecde2331e996119f06bafaa665e3ae0 Mon Sep 17 00:00:00 2001 From: lucasdelimanogueira Date: Thu, 2 May 2024 20:21:19 -0300 Subject: [PATCH] batched and broadcasted matmul --- build/cpu.o | Bin 4776 -> 5384 bytes build/libtensor.so | Bin 86064 -> 90288 bytes build/tensor.o | Bin 24984 -> 27232 bytes norch/__pycache__/tensor.cpython-38.pyc | Bin 9138 -> 9351 bytes .../__pycache__/functions.cpython-38.pyc | Bin 3473 -> 3473 bytes norch/autograd/functions.py | 2 - norch/csrc/cpu.cpp | 22 +++++- norch/csrc/cpu.h | 1 + norch/csrc/tensor.cpp | 74 +++++++++++++++++- norch/tensor.py | 20 ++++- test.py | 29 ++++++- 11 files changed, 136 insertions(+), 12 deletions(-) diff --git a/build/cpu.o b/build/cpu.o index d681b7686f31c09521697ef680103c12a50ebd06..f1d237e13074bdd24178712377b4b07ce85b0899 100644 GIT binary patch delta 1013 zcmZ9JUr19?9LLYOySaOEdhS)JBxyzpB@sxF4Qfbv;06^zL>fZ=SeEY!VSU zo)najtjNoVU#>dQJWuX%qXqaHXt_|}$v>EOmH(PAb3?3VMX9C=hekw+CoOLbHq&jK zF(=loOzmzWPcybCLASbV!b#b1Q{2r^jJv)lSs!gtpe4F!PMevWV$!8D*GZWVWbi(- zTIT)uz*Q2WM2kb!YSkN8u@-!l|9j;v99N)=bP0PUTOP>O7qRDFY$mXoK@ylz{JmMN zCUCvtXZMk##ZY$K-eGL&aJ$aMs=R*AJ|zhou@k RuL`Ho(+$rXbNfHxp?}FD|Ih#c delta 542 zcmeCsTA?~YgGqybqt;oL&6nBcGBGnSFig(p{K={SWZaqD%cac7FnJ@FJfp(oi$Kz0 z@<$+sf#dAgF^-3~ZAbdBquJKxDllBSeiZD?~B_ zs>Ti~J{2nN0~Jq%ia&&kdqBloq2iis5P`W+@fA?M?+nUvTkzY;WKO5mNWBxJ{w#6CHYTWoTLkOqvo1ISEapZo|YnZZ6; zN!X5Q3RI1-2Gb1o$%};TICrpvy}__x@=IZJ&I!B_p2B2J5p&K6P6*Foawd>BgB!v- TF?pqkIoAe02(N?z=mt3et5{Yv diff --git a/build/libtensor.so b/build/libtensor.so index 411895840f183a23fb29024bf2ed54e9d24c29e5..c430d1986fbb36dca99ac3279f2f5b53a273deee 100755 GIT binary patch delta 19288 zcmai62UrzX+MXG>2nbT7D}s$EiekZ%NW^mWf(7h72ns>h7R8oCFAEv9MM&np{gsM`r>27haUa$NyKLmKFRnD!)GKuWAGV^&jfs4!Y36UdQ8M8om9BA zaX;qmGxC?q^Y*QFEp}ZST@pcxU1!PfkfxSY`D5~iC0M?O^I&;A*#_hm@pDU+tqd7U zrKHF$R_@C5wzI#t+22FzkT!C#e3UWs7(*u33z07}U=sVAO}&uPdN%u2bdJB&Ksqh@ zFP3JAWF5x&W4n0ny7pNcqIh-FX^oP~s_{VQ?-02^3=-Z`1 z>0IzSCsKlNbkc-h#pt5}wu>@GK#v{|@S*sbzc5xR=R_Lf!SN|?a$vGZXsz(`w|Nw| zOHzms_=*+}N#O=r;OXHd4CYMb{CG~J144g%76%p!yqBoZ!5j|63H)c#uuYRV(3ZwU ze5s)jl+WUVNRiM-BEZCKjyDr{h-lf)r5t}v;M0VkK}j5s78P0}{Jb=(Hr^IpW2dF_ z%;$nmoJjTIkm3to#|{D5e2`!_kDuLl)iY2dHpIlfuwFBj{n%FQSY^)9D24LMU?ZCxr2>j089~h4tH#qE(yU@(V^e<;s)gcukXq06ffqW9(bv! z_j?n#zCNccBGI%wj_dvNjR*I$O5h8bfAsoA=+Tb2C6=!{+)g?tTDEZs2g*c3Zwp6H zm+yGu9pTbr_RO@&lhWp7XG~5_Ps^T^1v)(~d-~jIso5Db=FFTerH<{?HT~II>N`E; zh#ejLO6rl@A%{ppsH>P4vLs6<(Q;qJj>IP=YGHjzswdq8Mh_}4Hpsu<`WJM!oAj$e ze$XIazG^MTp+%D6MYk_9$eUe!;3G*cQjy*1K#|T11!@p5Wr{bw!SR9vn$+B|DX4=Y z1qp}x{RE1}jw*30HA~#1==h9xxO{;LFB1CpR=s(hbpxGGFyXPEaz{23zCiG!O}JwN z*U#z6{i7UM5NeJU9Jiy0lGLIGmkJFARs+7yfQ!Y0NsA0P`owu`FyN_96c-Hxo-EdT zCN41u=&h4}lp63T9RxjL!1Ym2#pMPZo$Nd+4R}u-rQ@&4AZVluKph4=!hlyB@U8~@ zkpYi1;IX1>Xh#Ga@Q?(Vw;xj&E3TzHru)_{xa2E)4<@UaFz z2{pJKlc46X)d(bU?PkIh1CDjbd5kvT{hTQ7Xbm_WU;2?{!08<^zuy6wWf17M-Bi55 zfVZN6IC2blTO9;lVZg=R4U-lcaB&06@O5^BAkpBc$bi#s(2orUyuA*B78`JJg=5kZ z1FlbTDkv4WojS+L;yg|m48+>RWaS3Dms4AkDh>E>172moM;ma50moWK$3H)+4T6a} z2>Qr?Cm3+4SIreYrwsnIO2yG1Xy#ON2CGo?L;N1l>u*Hz+(+~2Ls;C zfTtVq1d7x4(=MQEynfgW27WpSnqt8H4ftpS-q3(gG~nXSnMtz@xZP$jm}L;OG2jaf z_z(l0W58P)@D&D}KD_Ejp#cxzAn*Tm20@^%4O(QthZ^t=2E3yIFE-#o2E4?8>kk!F zT-wWCGl)VA1}6*#p$5F%fQK3IN&_Bkz^e@SU<2;ZaoYaI20^th02Xe*9~tmI23!(r zI~_w!47j@ipJc#&4LH4%i_v8v27&%GKxL5zyy=c@&Gve&j`fQ4ktN%*(riU5x8)vj zKg^tHan}Lg*INGIBkyP%8{%ru>IK35TNS@+THynqCQl%1iDx;V2wYSIA8z6UO}w{> zx0v{+p0)k`V&LsG;Liqt@;6QV<^P8_>;JESrg(PNVzUEtz_X?VPMi3nCVsz(|I);7 zcJg*!^7T%D^Bo7n%6ECVslk+eIH|=z!oSnD`MUevpZ`n)rAVA7|oY z1aH>|Sxb{bxQP!m@!lrhV&b2Am?C2by?q6K^r`Pu)#pK<8}*@iKI61(r@vfZ4h~$_}*UTh_YqCDTKz@`=cy z3+7=#VNgL83*t_>%VbdZFqi3WGMU=Fahv^srMYcJ*mC1|J7k1maH+NY8dcJ#fi;Ed}csTb|(+<+Fj2(acp=q`}>myt4qDRu)7$mdO zB9jib0A)u8dExn$E24Os2dr#ol=`d|0TDg807l{4v&ZTfX&C zenP9fV*vXAETJuO{QHqEX={?46%t#1LadF+x;d@(P+38ICnvB0z;WwC>N4L-hV*C_ z=i-#jge>3sO|C=kZhg4O4|cb4edFPWL6&QMDBt-4@VQB3Lyxx=+a%I0epT>Sxbxyl zp}(;lt5NTwbW#-`s%%InzsEOMa??rVi=m1&oxJekB(DZCx+I_WifcOg{KWxYyCog? zC5^~Em$use!boM>pUmRjm&ll27nKE1DHrlFF4B%Ke$}0?en^wI3}8Yf zdjfg2ca-N9=-Bd(JJviTS4M`B)I@hVft>E$Us?AD^H?^X#PnI`_0#Vx2Sop10;%XT zz^l*+d^?`B>bpcyerGY;;L*#`1(A4z8sBFz&pu(!{v1oH`!4s7ds4eWapOszb)53I zvE)Z@{oeN~e5A*}GM2RJKftTI6PQ1S*!wS0HvPuppTU!p z<7%}Y|L8pye=!p(pN%5@ZOi?y;EO|JLEjuhs%+zwKBGu{Vzj)TWF?MN?)}2T2ahD@ z6MvBJlOZnz1veT2-%;?$+vatb`94Zz%Dv&M2c#+yBS_kSX!#BD)_^g}@dwOB^)T}1 zfUo_3xX*(2%6rUv@rM!Q;=lo3Z*rJ+>oC%B(A&!0|FCuh;+EABc}I`sc8kS2$b^bx z$PTx`Ub6f*vaM&3ayW&wNp7rk9YPY5qvQ|BtmL7}?SHe_4O7U)_W0Uz zOWpHh4V-IV3{xS@??&EAY1HA@d%V=Scil%?^U8YVJwkFFkPe0kp9O1vL^8QYd=Jw5 zUxyp{E+tHUf!t0Bvlj~6g1B#!4F=s|z?HY&r0uYl@Tq*<(f$T?RVvdVxP!{_td&@5 zDy8fmTrGuWTWQNXM>Q&-ecbwu3!B^D*w)#~suYSZ4~o&0-f;B%iL?41Sq)+04JKm373sYaA)9%ad0 znw0M*+(A2b!GI_AK1mrCq#V6QDjIEas=w=0uU``i3;S_LeicHC(cnX6|8q3N zTzmFtI13z>RSk@zA?-GsRx6WG+;?ePvD|8o*Qjc~r85)yO`yUoXCVe(!*FP>B)Agy zp{W55np>?=G4B?2+3q8J_2T&4*``9OS2HRa!Y6IoNx;rgvG58T6$*?@qav9{F^r1K z4mMQXXcRwZl;5tBA;SWdgez=VWu zNbK0qpb=M%t>P2s+E1)ho_M-<1i~{f)$i0tzGWv9`d#i1AuS2-_iwKdp8+Yp4PZhW zdjz^@>@ZT?yRkCs5-AuEqRhKWuJsOVa}@#a*ioPp`3Y@nW3{uVO`9_>8(NwD9F1jDdFAx!AEj|z`d zAzN}cU!ViW^Jkzmm%m@o??M8TD_H|74Gm~~i4LLYqX_+s$~2&HYE|ZaQZ^{4#Zi3! z<|u~^kGP9AsdizSSF+JG5IS`V5&SP!h;!#*B?>Y2U#t)dDp?`IsXtR8HeO?e_-8Lf zzKpbZAr7+^Szc#Czx7mDOoglvYb#kH&H<%`c&pM`h;HXt!p1_JJ;&zK7KDCAA%3H| z)+)rZDwuHOygFYnCzcE*^n05M|LPRl&ykbk zCpO#)lvXDLx>#Q@G|rG~!vj+S&p1m|A|~|x?6j}f=!+3T|2X*-b;I3Nobl+ z1d?L6Vy^6xt#j`AlF!yddDe%dWm+_eNei|Aaok!(mm;(S4Pe)GYn4xd6$9qD^@fn> z=B~KvTW|OjG+|PQ(GJ}O=iRLipMu|jXZrlR^ns6`_3m+NwaE!ntM)0#18$3A<@3eT z$hW>`Csp#T4PD^yE1#Svy)zON_fuqP#&fY-zh?d9VA|yH99Zp7m1hqDxYCt9}H{QW*8Zm)kLX2L8fJOP@WzmA7uqAH;$3gX;Dg# zQ=}>@LTPl2$Wyv1>Bnd?N+Oe`e+-A1kwN(*&S1usq9LWWyl-h!mmj6QT$yJ*b1MxU z?l^y}POdkQKF?LdH~zfpMKI&JuB9#f&k1sIO5?zW2t&UXqmJv*HYgCTetERgVw9S&%Spd&2A z)59#ph~upMDTkqv4_i}+b#TH%?C7Y67!F+?A_*=X>yNSWAE)6C9U)26#wouXVj(Ua zVj(&oBaUe;;4d1}P@f=jJ)YYCFLl3dW zEGuP=nO4dgbL22<%-b+_JcbifEH`AwV>QKpCV8ww(B-j8XbBn~X7!4LrX#jY4A>TCEKGris zOUbF(jfT$M2WcS;P5wWC2kyTF|7zpa8UtPKKbbcEbDY=dnX~(y?eX5n+VgNPYtO)g ztUWL6h2@)YWQs@Nk;mH-t;b7-E{|tLf{vmCtUZTl?EU*la&~a5r6(k*aKd3aUq|5E zM|`dJDiFs%H8{(+%)mLF#)J2gx3j}t+ryB2yl}C{(IXVM96L@HE=W-N?k9T}*gQKj z@VH}i0P$WJ6Z`vNZrc?92&7*`-0{*Eg~v$VJRWa;MM~#~ke>6xl#0D%+`PYd3}p6} zxMSqIc?rstF_C0V9e8(u>d&s}%zv1y}8MU>L zBP@6ME1J9VQz@}5j8lqBNsomgLA&XMzWIF-I_nUZa{2Y<%G4c1S#-i9Yafd-P7=7-rkveNQWwX?ez=#1 zT3X-id;HG2zHQsdk;VO$kG4DWf9)&gd(D29|Ml%Gf5#r~^j8m-|2DxM^f1Rzx7Ur4 zTkOn#p;P}P-_Gl$^xa0sif0(ZFdp5>OSv(zKkViao4T7rJT9&q;^r2F zh*6GhWn-}yv$f3rijBqMt!ylA*+oX?#wqi6k%HWipp30`oebH`oTP6>EkcRQ(w9A| zzGQh;l(1UV+rkFlpf8EGG%hxX$&Nd2xSA8(x4EwG4V&P*5wR}|RNmagk__F+eD~hX zlAQGgE5e&QD0tkl)P)v7N!~$DzZnwL@{76-{XS>OcKm`wFK?`P?IO12twYj+VO|h8 zmXiFq_qKCqeLg2c@|uREaA6`Bt`t`0m$|tCHZmQLiqCr7V|b2{=kI+Lb#I)uW{i8F8q5j zPv|u+Ea$?LT)0{Ydo!VZHy0LjVRs?)*nk4(9*HbGS+lg}-o=e>{CB8?k&H&-3`td%oqCs9WoeP$`HXavLjm0q*Y!6D&JQn@#lGUKr^@s#b?IH2iKdQvAA-{I)u3(1Mne zJ>hpa`-SD4$mw(EWJhMD&C7_KlQunrYer6#kBmXjmXW=sa5 zPNW@8h}%3n-6nyrTcwFXxXR(YFig`*KvNoNS~+OzNHhkt52(9~Bn39tv{sHbb%U97@O=|&qXdDs+y)YheK<`gLr9s0}5yy^` zxQUu}1~eJ;A?RGt04%Wtv@@u8x~7c>?G5@m=%)6 zjcsnw?Og`D2>*5|0N)naF&Wq|xSa~eWj%q;khhQeuMVW)kY{4@{FeYyHsl|))ikz! z!|A`NXduS%8Yq0$PSY-6wmDnqQlV!k?SwoHvoNccJk=>b1Np`{O=EivoB`f<${#}B z4zu=|2_U`{fMufF3!29E1UVf%sHp(t0hsOVH|tJ$17{ZvhrB*kf&V6_OO=Avfo<_| z+6OUvI~@;epde$}cpc}?3Mg}&eYg|y3I9WW2J$`sL;eu*z`mN6Sfg==Q2uS3>?O~<_i_G5l6u=m zZXF!cRFWDsg1v|;QuKm>yec`dtS9VnfCaY_y zk}XC3$`|1avGDgBqF%wdl08)YlgeffQ=9pLW)4^T`9Y^KR;IBd)TuaEGDfN^DK-(< zm{C0K9GnjurGAT5#;u6s_o(iuQECH!xtW^-*v8ScSz&65Kaw~wS{+FRvA7sK8KW-t zhqC8b^+Sp+05*OcSFWPUuZ&ZFq=Hfi4v(jTCQ9XawOvDC?pT}Jyrhn#nDr%f9>r3C zT}f3xq?r3e^#H}z0UMX5{s^r8+%!$2n?+pTQq)v`In=ER!fokly8sxJrK^JjV9*Lz zi>H%lW?^b!0J89&tgfYkSrELGp&p>huVkoKC{_Y&b0$>|N}{u@68wH#Vm`oZs2Rj- zoYVLb&^Jryn(0 z^w4>_6U4~VogYS??j$kt=hc(pa)kfnLJsRcD9wsjp8%KR)vzXVyJiUuxEQ_4jw%dH zdI)Dyv-;}PCUS&t0#Ab$fgZ8yUxE5XhI1CZ$&M=ZWD~i)5*wyIX@Z)jhpJ%_@=r>< zOnNDSyF(&ncjftn&RkdD^C3*7ZGF=FQTO_u4}FgZ#c6QI26bpN`B9X13wFIE zs2!SP$7i~l)?5zkTu+KM{!~^dyGZU*)iZGqDKcDL-CXV$nj-S6?cZDSRsRG3c76MB zZLyD}*SR+CiyI>&!p~hKKW8(IQWxw|Di#}jYvT6>+G;~|p10h!|>zKawp<@9`30903umD-LPN#+ixau=-O-NcW5 zYwL$gSW1rBvrLKySgFEWuRTn>CN_9K>PvA0Zwt68-e#xFA-0Dfq zzzI!ug()1p;h}B(ZYpyu>BIPt&@VLUJB0o(3v27UVUo!;1I;esYan)fg1uzFB5o`3Af-n zP5n54XV6CJVh);3-5B#O;7TDkZ}R_B;O0&K z-ngn!f9Bn-ec9Ut(lUkEysvwaz{TcQ=ib*10yn=ipj_bQ-O~42m=56>8#D{^rs)KM zWBWKe;zR+{Rll}!uw^k7)3T7N&8}>JyLl7v4&lJOzx{x~%{%?C2;BTi0#Dok(FFOc z1Zv(S&|cu!WXvxbnIgmCLSWv!`LWgT?2m+xOv}ouE5Rjz4$xBCDKRn3q2@w%x`Si#JcO0keXjs;OEUVG;g_Y zB=Cb`ROs_1Q7wzcZQ}&BqP-m4;T55^X%a6vd$h&r8iAYliJxQEYH73_V)55kB1>ES z%-)Nhp9HlzyN{AL5Yz|HUL*e7uD!T{&{0)7*?`Gq5ma8XBGf64rMjNxk24rp#B zBZ>35nq+?6z&e4OUny}=;O3Wf+-2@UeQ~?n46QY@Nr51K3 zj@u-}<~J4W61e$Q5(j~^@Pob;8s_(j{37tD%X#zkcN}!WRg+f5{O*J4z?sJm%Y}yd zl_q}`xcTiB`*)Xhl3Ubk^#8n*pz-dT&&#uw-P7acM{-lcm+;Si#lh~$J>|FDNs3}o ztNY7E@?CY6O>Xa}KaS}Xx8Q*N6EHC8q0?sd$*YbgYHXtH>z6K`MRj^Y9g--w_tYQQ z@btJcQSRcQKd9lMGEwg8eAbK<&z!NUUy^K>W7Snja(fT`sVzl4nFM8DF}-yfsrn5- zR{C>;PCJfLhYpb2dFl^ulMB_A17Kdna(6zvMZ2m0q-y%3Y3!$}=RmoOr~a%~9I7UP z_tWpEbegTM8i*3JMluqEyd8a#s&=Ba{`RCK=cTG3rVKn-im+qO5rK z%4x594wk#>kLOY>Zu)VwQilvigjV8ytyM>LQT~>yTNT^GGh7mAFqxQYjK=N`MowYZ|6}E#@T|f4X;MpFAwLOP9-nS%OI}wG?(*c`rS&7>(prH z2GQjD)B|WlRa(KE!mWjgiNEaQEmSjvYDNUMtPP>G*Gw6TnXV!Gi_ZRwJ)i|!LIgx5r=Xh}j2gZwp1_?hEvnXztqyQl}BPzZ; zo*U4?qHVFj50B;iGESryn5#5!{ly%hlZCdQMFICj0WF39Is%WJ%Jr-iM|^372yko$ z7la9euibe7AF715Sm4uyqli8nZw{Wefx=JTAkORfeKBN{S8$xpY}#tTpPg2a+Jyr{ z;efV$2q$M?FQ4ez1_^i)7KYGr`je?ufPV zOHtV_FvDmp;Gw9% z#>LWNF$6go9MG@(%V25XPhvT~ixa7bxbFAP=(bK6Biugqn5H)lrf$mqXSc3C9U zc070VjtCGS91SJ6JpJvpUgfYplRgwPz%DMriNb1=NGn}TNqtqgEbu3yJM?rLizwb> zxPu*{j910D?TKf&T`CoV60z+3pB3Cd4-hZP_?q4e(DZBPgy`2VGC8h~LzeLKtH2iu z|JR^K`@7c!uK$3MhLS|YjH$Mi3yuqeMBylH`M%b*8n~!$CCa|)+C;hKzAN6j)yTNO zlk6IhC0SbOftdUHiSZB)Ley_%$d?inLKE(15;ErqrYl8)pPIKWGXnQKYYgL;c>4vRNw>m+ma_ zyrARfv$%YL3C|Gv*4FeIMAHqdtvTO8P#w>Gn>&g#;j;vvWWsHETt6d<`$KnNf>*ZH zI!^1y06JT)E&$TQiAqwQ0e{7SZ!zE)6KC6Q!1bpND$FY5emCeco9PeHUOHB>9xJ58&v;n6NetPR*z_F4zTbuzO=InonZ3Y2W z3};I+;0aC?PnQOqPB6Wt8}Jkz1f6Wa^(S2xf0jYe)L^i{fVVW@83ufa0bgyv>7%dS zat*jm2SM|q#XXn~0X_7##b6MrgP_|DxVYIdX}$s1FIg%mFyL4~oUPD_qkgfbFky*P zz#sH=i&6tFp0Jpx%z%prBZik7a4eXa5*|>Y_ zD5%X~;A6m(40s&_o@&6wa|n~B8}MER{mBNrMdkV@NwW-sB!j^M1K!AhXBcq(v6_lk z8*pEPey#!c({Vcfc?Q8iT>!epfVVQ>+YNYK1DsVAO585hHwYe83Z!lZyuyHo7;s6f?X-X~2Hf3%rx|cB z1D?omv04PQ-q$Jov?oaq&wDG{)2d+_t|*+RXr;EyBKMO_Q;WF`_&2T59ZlNTGCIK3 zj`NtmB&l`-+^B!A9r6r(go$rt;_I7uKNDZm#9K`K^JkUu>@?ti3;^YS{(tah{r{hU zrg(N%u{q#9Qv%)-{@A6$vfb&~T{5vN8Z4;kk;#ZjXr6zv9 z&fCS9PuBs#PcZQ#O?--p?`Pr@Onj_~j}g4xIq6IajZA!f6Ypo@YnphAiGTiQ<)mYC zfb&>7M8!Xw6z-Y$?@j#ICjNqnKV{-SH}MA;Z?C-Uc9|5mn)r82{M#lz$HcEN@k>qo ze9GGm*Zp*p!UPjP(!{5j_HQJU?5u<=|7kB39SRYd6TC>*2}Tb@&Mk?pIKQnA;T+ncgwTCFiL`)^rSR@iky+ zev1@aW~>M7C0o|Ld7ifH*k!;S-#ybb$9PTCPFr@t?sCU)mt4?4IyMjg3lFW1Z=P$~ z7ad^b`1@~7vt=bz*fN(^Nb@6*&By;CkB9gvU!{<=PHU8EoQqB&FFMUpmONpyt%FI{ zD^u))2g8AdI>?C~PtX7S3+LGh71?o3vrYpX2e1H9B!^a^OIn&F_k_fj9T#n5vQ($l z87j+(9pnTG02C)Yp)Ru%$n4HxG2u>`9kT3%^O+90W5ScA-mvp<%Em#Kneaq@uoLjp zL8PSfI%V=8k{Y|g|0>=pc#h5fU^zZQXBQ^F)g@3VNhbBWgexB;lQ?h_l1W;Zah{Qn z(bw5@L>eZOvt1HBzkQ?we#s=X>vE;{Bicy;bMTns82KkU8snfjx-;*VU?x-&lF0q8 zt9+LH%51aem&Rw&xObDt^0+kR)&O!dt~q{L_J|*%to?=c`GNjqVf-~E?`P)xIGzL? zJD=*#-+Dq*w@hV1Wn(|Gsas=@7tpa~6+1RQAxB1p$w}nrZoQO(hs@ipzNBCGJkM$m z_1r)2M;>=i^vuVv9=y6=`Vw1@8QEdCQbWIAp>(&L}{g~flD36yh&Rc=QFRsfZu=Vd_mq{x<;?$`IFdLZSQ#2mrVMB-?;^Pak`)gwY7c%#Kz0re96ksD9>6-_Qq8wz+oC}WY}rYm zmXaHK!ftT}7XzIxito`}sXo12iygi0g3m0J>G`>Pq(O>Rxqp|$r_`;xfof(Ylw)Nn zm*$UnK-ZSo8?47C;x`J%z#sMc+cTdQFDCTvNQDilkmc-gj})ealdj3>zFTRmtON%< zI`~>O=zHp~=@Pt}I3_uB>__$V9CKnGK}yqGd6&!{5E9t)HkGE7+OqqhdQ_5ikxDP! zC0>L4nrufd4;Ex66xs4@$IBHe`Oh5|_7TE9;$c6yPQ!Yg$6K=F5Da*x?~vJp>MB3n zBB4os{!xH)5u#$et&L}+rs(~77RR&OBqPn2%ou8|wf`2Jihk>GhkTx9RfgOqH`4;^ z{zdb1R>}Kq+cnm2rLQ315md)#>iet`%OWQ9{(uVCQ6cNMdAG@!5lLQGfYKsfh3<;N z`Gd&u5y491aypHb;5$U_?x(D{Lw-&5^Swwj%1SVlx#TL1>bQ%TL#~t8Y$&WE9)C7z0ZvR^>bvW97QdK6kdLdSV#N$o}>9! z>6L|LG|!}ea{*Y_QoR+XX?udt%$wT1W5D3yCppn6q{ zK5?BDeGN&yRP;wz$?nlnwIovSew|Sozomu$@5( zfcllQb#cCf8ni-L59VbjT&Y|Mm(iT7^hnDuS&yXP7ZFDs40w+;<~>U<20aKVBkzy# z^FNA=U+R%nW&19SIpRvjPadV5ze0*8f2Qn>+Gn3KPF6Z!Aa`GlQ#M^9ji)w=ZgrU@ z?||PjTb36dqrZ)y)#1Vt7uTLk&=9@a<}z72)!JrYDK(3N-i!8t~MWKSqPQ zVaUpq&_r4@OGw5~SG=&+FyqrS0I50VUL;Fr2G@BTn)Jn%_EZ*4^v(0QS45FZ(^@Dc zXIb<;XIb<=F0ej$>kO14;L8-U6P)spU*qn|L%N;SLq3Foqt^x22UBSH@#majwx3~P z)}LWvzCF*v+%ADqP1u{l1ff_S=0FoY%n#7zm7Rxyqse&|rZ*j|9%sqJ*=fqE5*A`c z2@7%X99c7`VQgG66du6V6ruu7cnE^O0r5(TpvyyiOzXFtV05ol8=<2Vsdb9 zV`c1Fa$|1GTwCEjB!6jMs8n{x->hM2J>yTd% z>#Qb9&oitH>nRv5h7VJy)o{i`g*Vhg4Tdfc)rl7Wb_pxv5Dok3X-vvU(qK+g%ohvoP<`keg09!`wc}!kjtD z!gv?6Fy$wpln8rMm{d68VSW$Q!!$%nyc!SM2NO=SFso?D z20V570&#q)!CAJ&9q05~E-fU{i|Q+TPm%$PGOABHOi@ev=cL}UIOX&SGIE*Cfq@g`X8FZXHUp%$?FPP1Uap?$| zy(CVFI8Ju2wkc@^{yMzaa8qv8}*ufUpxli5z6{wWW&-oss|rpej^UE z+%vzRxhq)*$+l%NO7cN+V_87m4(KtnlZKx=6O22`9Bn4Q&6 z?D7D=?~c6e^b{EbNA~66%EJQk;qsE|9-lIQ5B8J!nQ=<9!{kt=O?iDUk>8Aoj@wIR z#f~>?fL0ly8=t-G`s^Vx@y%XJXn`}=`uWUv}2je^O^hB2iR;%+DR%_$0&h2NmN#V{399oW*z0LPnp@>{U|e#Y|WZf zefcNMbn0$aX4XD7VSfFXyuT(Uy5M7;;9Xa9N?rE7oKo|>@Li9DX8S4bds$jXb}--9 zce1p8+smqM%0P^Z3#wKgZznU?2GrfQ_ho-%56gA`UUGhIu=3F-*j5Wsctn zO0r{Ox6=T{j#GQc8|y*?lDROE3)>3eS}u&|!cJVcZyUG72vW;vE^Nky$wC;!g$=kc zfD3yI;jP`wvKALsB@fWXbOQz`5Pj*8kCwAaZSZdefbrycv`kG13ukWE; z{FwOU#@4Cbfc+T&ZYH2#T0Ol#~7J}}G*0eIvlb{u#Z@1Aj`h{0^TTSZ#+95{MQbE^u(6j}hqdRF@9_YbXO)CJc z84r8V_T6C*>fHea=kOT*y=ab<8^ub^_0G*Tq2cUOA z11tzM1P(w?fu@2!9twNVO~YUh8Zcbb=*M&!p!BQ2^Pqo&`le}Gu$v_91?}vHzmwD) zp=s#=nt*Nr^%;rEL2rZ7KY%NvHLW4&FX`xT&{|_PZ5n9gINV!7lg4Y>e$aO(XxbIf z!ih)}v}_Urxl7Wc$(j}iTJKfF0c|r?)7}K_2l}BMCv$O91iAtAE@+MEn&ycwBH}<> zf_^wd(*}cH0i6fhf0m|g2F(OL3i^CD>_A)3(KJ^NN!m3RcA!6k_66NOA9kR>fzl7T zroRq5&~2b&uH&S}LKtGUz5?11^u0we1ib`04b-*-hMN53Yvi)%mqzXTLsHaEDIWHnpGVcEQhzA)fD|pXPjubf3xqe1Da`?7Y0~w z^pDc>aGZ%y-wl(4>>F`@GD6eXC!h(nyjCjRYr46dQ9P@;OH_A1j`^6S^sgC`E?0i0 zVkW-=II)$cVU;(?KXl5U04{@^eN?O4yYzPv{_T=4ehjD`qiLn|M**iov8*T17WfzK z{?UOn81knb|5-i{@(5fu?DJr!|4>mNX6i;LwC}2ES8*LxCU9BLP}&dq<8GS9K4Nt` z80C~-fxJV{e+nQ;PayAR!5B0C-9jZm2M4+?A!Kz{Q7BEJH8gMpehu}XlE zm2&{{36Qf-!<`9+(;l$1N#u*JpN*A?eQxSh*uuKV3a4!$e>OKXma**2yoc9@`UbV{2{QbhP zoYjNBT+9^eu_2DDL)4&9IZXLxsM-b8<6SJoR>)m@sp~@FF(XuVRW=P*S5l?_NmD-# zMffDF$Qwtf7jdqf7^(h3Wrd?upD@tO_ipVYz8JRb5B1 zAyd`e6iWrRW19L6uo@?)X&U|Z*mbR<4hWY6-L^oeO;MMK<#E(`F;TQ9TnOH}EN1!{ImM6aoYrp3Q-`;fPs?R$XdBtTm7cI}PN|FJ zCOHC(2tRj0V)^R8HnN|KyOgfJ0efXju)3*@+*6KIZ-EVD8=y>ngTzJhlk)uZ`A{iv zk?K0Pn^mH;Tl^*coA9c5fTZs^tBMCo`ls9+AIyo16ogM7|6AaCl6g3(g$q{aXz>vm z_n`Z!FQvCu7sa6ez0~6|=>NEmMg>U<1*@R{Gt|cr_emFC^w>Qimj6px(}kISQRz+> zGJV=TYH~X{(mq-E!>YjAQT$88oe^oEPS(D~zcN1B(cjYayM1x^GWB5=J|^!)$=mh@h+3IDuT@cc7@WmI3v zGX?%g1ZXH0EiwP3AB6rElYThvxwPlDoAB`h&o|+FDNaAeEHDZF5C(-NyepPN8lc34 zFB5pF2|q6IG81mWotOG42R@2+k3OOL2!X?7utwmIO!#5obZDYh@lIy9CY;_D`1ve# zQ%5;c+32U91s!W%V8XD_%hJo2ZIr;9n&Kq_r}6T|h~RfT){+@sTPj+k4(}vKwH8}X zgjAZu+*Mu=IzE@-Y&B8S;^anZc_-WxyymDsb&@0F7l_`86Iz zucd|Bs>A{D4O?mEfKhR%i=15n>TN66S#Rd=~n{hyFe@V#Z_?JE}1uu z)yESG&D^|CtB1hN8^`PdH}Cj+Tj1usm0JYPH~fWH?)EGff}9NAWt}*YWV-#gcQ_tJ zXj2;96wFNfih`fV54+rC36nRmxI1a97KTmwHB&?GvH<4I5E zL~1JVs=MhVDNf+#&B^Hum()pJW$(ZxLdm?}b`2YuRYF?1*ObS+C~)(J$ufbPw-dXm zkGta12&#$q0(xn*io_YM#k{|Gpuo)=010y-Z{y?=C-)h#zYsR&4Y$E~@TD=#+fq9V zyf}jwqwj89C~)&`&;0^7?+P#DxQ}Gs+w85z$IJe5I3vB=;$BSC5SvAvdqqON%w6CGh{NE{ypmq$;^)HJybn8A9n%dDW2?l4IJ7UXo&7w5QyUY@Gr86V zJpR&@&D+s;s2ix0Q$lRsUVEQAk=5f=E7F@gxg#FD^zF!hm1pm6XiLIx*<{aJ9Mg_{HJ`Ve1M$ouFgu8x5#%@?_qM3xBkAR(`a=d zU?2T0N2fPENYI}lYW^_U%Ugfb(P@!-873b2fxI@ne0dC)+f`2&kBjMQ@8NQL=c}6j z$`+t*7%tmotLmL5M^%q&%Z=mI-f2*Fz7>hPmzTN$TF(DhL;wFZJ{_ojnI=bi=x<4Q z81NVYbN%H=&@~n{c?8t-H@w`p)iopJb{_gmkymx~EO>AISw*MgRqv7LBJnPyzZ2or zY#?9{{Y@y(PhB%oZeLwIflRKio;9!w>Z=~3Dy13q)jpIJyC^dnsB6ILZ@2m`eXDwT z6e3u~vye4R^%#wOqXi!wruLzHy5Q5p)HS2!E(+*b(CYbOGk#OmI~}pg+VDHmI&~mr QC$hI%`~;+K7%4aSUpqFR^#A|> diff --git a/build/tensor.o b/build/tensor.o index 3dbd4caff2340dc14b78f7f8f694863588229b57..6657ec4f6ab10cbdf09b24d52131444bf996f433 100644 GIT binary patch delta 4699 zcmZ`+4NO&K7(VB6s|)@F{_x!sg$EQbApTq}h%^UF4K`5`6&3Yz;X=SdRBo>9=fW(P zv6}Ykv}IazrEO-?U77uCk-Dr^&dFMtZLof3W-GSbYGc*=o%0=dIP8Ag;hyt6@B6&( z`@QcuoOAH3-q)+|$@TO#cIUs5p1ZxQYw@YFu9DQUouBx+=D?MAQJ3aicHK+Are}{S z+j+ycX7TQx(%pYk|CmYMvYqEr<2mA{uHQo$Mkf6_`5`)Ol+#P;=CrQVA!n`l)98Ys zD<7Y|FC%uPQpdy@$=11y9{)J~nC>a@QiK29VZr!mqdUc;Y4N6MX3~uE&Gl`qTf))i zx`!Lg_NMTbhIX^DZJSxQtt}j>54T4fB4%AUTHn+Vp-^CY#iKcq6iw4sHEeEg+tzNj zKGGgFo5I`KH{9A_TbhyPR;HWVHrqYAx@T}@&OPLj}Q*f3rdb#Kcny$TPShBV*pJvejvr9ROWNYrS*?hK^t6H-9cFSg; zvUz5T%?+vyk6Je8ESsN{O=p&EBb3ff%f_Jo*}_=hwPPKeI?`sDWc7B-=76$!JIA(} zqI8a1HebrW@sc^8dxFD_3aLhg?9C2EttM%q_4A6LogvDe7l>!hyQm-i^KQR>nFf{$ z<662MMFzQrk2|>vxo%Qch0v3zoBw7~WrZ+`h0Wyi&v-rt|FhQv)XD4$#d7ji4;^;B zR4`hAO#|QU)X(feU?%{3%3{wYi2FU2E)zx{_}&h_*Ja$`KmuQ%pG|o=RgbkVwQ&?& zE`ma_vX)dnoLUY6J`V7cq%Rl7`2?C@NKK;`d{Og93UkEYfSnA=W7NdTOyqbvnPgMX za=~|Y_X?puNCT|DNm;rps;8_K!e|8@jQ%uLvck)j+9~Q>ABRnqaRgs zGDYF;zUNVCr6})6pti@Qc2&yqdL@C{F{!41N|!3dxF1lH-m1*VnxyW=zY=)=E4_{J z8FmbrDbC;>6kN%7Ca^ezt1PxjvF=@oQJ8(7249@P7icfD`xE${uzb4{_3LhcnW(slXwMi z%|BUdLi%bETn((ey`4W_RE0Z4fu5%&$p7sw(060nR~D%QvVuZ1&1IH#rzD_b_EsaIGk>Gc9t|1MM+| zJ2if!9%xS~+^zPkREr&Oc^$UJwIkq#4Lb`sHtZX~i&(>HL z>xE%vb1+`cZacR}4sB(2wqj+@I8lX`Z$$#%mC`r1A%X8nmzUG2_rM-Ibp>$j)bD`b z$N4#(@=lfYX!KA`nvr3jk9thcq`fu5Sfp4Fo6A=$r}lE{sS&{{@Wt920iVxaPTVIc zTq9yH0*jqG3pjS_BH(y`4FHb48#|3d@{%1vd(*~aqi&}n<=`%I0SMufehJ`Mc7?*- z!yOr+9~<7qW)Q*v?SNx|^CN^b&6$g@z#ivf0Pvx?i1DkhSG}Kmd5iAn!)Z_k;CKtq z1su0@k;2^-*C^b*qqU54{iC0EEf4@JeiHD7kmx?ZG0`^x$Jx3vjK@a2LH#mJj2IT* zbeX1|C=S`d2APcW75=IVcfPzoapBI_ zcApD(zOXG})|J}wP+!$_WaD?9uLHDWQ(?-5V#axyGw6d&h5C3J2nEO=@~gA8A%9A0 zu5B{YOFKe^p8E^vozOy0a{-m+r<1QX*R!R7=GHDu*;8QK`8~9wwlMix#J=pHGqpv@ g$7k7BbLbzw`c}TD@%n>`!$m1gjrLVkSC@lmVm@4bgS!X8ZG0jOrEVe&H46jWxLJPwcX|n*ql+# z`8r$KxQ=t$=6qpuZfVYMIsfI{Nr50|CGu=FB&vFn@)RdW$9dZ39FE3<^KKfi;L*P; zE;8ndezL)x8JW&dpIwi#ziiB(hUD6z*8X3axNakkoPrxfOzF$q;QnT@u!zIMy+$64oh4mX*R^ zj!!ADjkC*kmN}bwBM$SX&0MmX2Pn3UTPHJB4Cm~P)1ruZng+%BJoy^9^|@x69*32q z>k8F1@WAidDrZT)1!k5!`qh$n5#T-=5HW*b#o5?G6T;f2ttO-Cx04ox^;z&my(r$P@KUJiA;3L=w*U^kyO(-sv5{K` zwAtDGVe&Qc2!S(7H4bjvlUL~n!MIV;Gn3DG`fSc z9$IMP)?VPCRQ==?ob!?yxS~0x(>3C|CdZd?e7}>v)}q79$aagWcJPO@^>@=LQVHPq zi5#Zi8o+UWLe$aBtrxV_xug$KteLZ;bV+192@z0{cL6VljFW(4#;X))=8@%VqJ>)rG7UWeWn1pS(?gqYRyZqiin&j{HG|Q%%Tg zy`wp9QG?t`ng8%O?Fn(~97MoR$pxws)~j)RzmmR@AL96Kxkq(vGMXj#xW=8@JqI{y zHvo902x+t%QaE+*^0N1+uAQ@wX>b>}W;NgBHJ#2;Y!`pzdx(ae{sH(#$da6|x{%d+ z0ME;hUVbWrMa;H)KT|xJRUT%mJ>tO4_IlERn=SRc12@}fSMiH4_4Iahifiv)bq$L< zigd-2{klfG0vZjcQG0i$=SEOfR9ir=beFrs`LxiLK{MTcbDI+Od&1iFwFIgTm%DEk zP+z#x?ei;c*zflH>9265r`E56{mH0$cX`tKR`nRBf!!5Jo8*#_gJSZjL_LfVd!{03 Mq)k2bQmi}YA4O@+r~m)} diff --git a/norch/__pycache__/tensor.cpython-38.pyc b/norch/__pycache__/tensor.cpython-38.pyc index f6d3f536f149a35beabf14531895fb9aaed26447..15f473bdac5e4b577e38551a4775c2a66be43b79 100644 GIT binary patch delta 414 zcmdnw-tNg4%FD~e00eX7Ow!COH}d65OT1xZV5nipVn|`kW+>{ZVFa_7ih3q@N|*Rd zV9c|CN-|AgERuprgV-gE3z$ln7c#==bfyxPEY>WZcy^cwl+OX87w}BpDXq%LGWou= zyyYsenatVD6Bvt)z-EHkMOG!eDXh(mO^iT!t`rtY2A~{Q3X>!QNM%t(3GZY>8Cgj_ zG#mKhAvW+$PL~1M&?_Tp!9Rhq-~v!Le+_dy5N84TMK4MOYJes-GXcc}YnX!>G}--L zG6O@S2p9wmFPT8B$v0*6C-cf0FbYk!la=RHOe)GxOi4~GE=f&^pByKvH#tUDVsfu+ zkiH90H3OpvvlOEMlMoXdBO9X}5c4oepows?NHB6RiEaKPTgk{6y}3-@nvth9KZRE?9=JR1>V5nitVn|`kW-4l!JXyNLByUL#15kiD zo4I%jSPIB0np47)!qUvx#0Zq*N@0}*Y2r#@l4JnN?I_`yd{IVLk~f~cgmD2=3G+h6 z5|%91EZ%qy5WRqRGP|rQBkN>cSxIBQ35*2>K)rl5%<({+1>_fbltyv=ishape[1] * tensor2->shape[2]; int result_data_offset = tensor1->shape[0] * tensor2->shape[2]; @@ -65,6 +65,26 @@ void batched_matmul_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_d } } +void batched_matmul_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data) { + + int tensor1_offset = tensor1->shape[1] * tensor1->shape[2]; + int tensor2_offset = tensor2->shape[1] * tensor2->shape[2]; + int result_data_offset = tensor1->shape[1] * tensor2->shape[2]; + + for (int batch = 0; batch < tensor2->shape[0]; batch++) { + + for (int i = 0; i < tensor1->shape[1]; i++) { + for (int j = 0; j < tensor2->shape[2]; j++) { + float sum = 0.0; + for (int k = 0; k < tensor1->shape[2]; k++) { + sum += tensor1->data[(batch * tensor1_offset) + i * tensor1->shape[2] + k] * tensor2->data[batch*tensor2_offset + (k * tensor2->shape[2] + j)]; + } + result_data[(batch * result_data_offset) + (i * tensor2->shape[2] + j)] = sum; + } + } + } +} + void pow_tensor_cpu(Tensor* tensor, float power, float* result_data) { for (int i = 0; i < tensor->size; i++) { diff --git a/norch/csrc/cpu.h b/norch/csrc/cpu.h index b6275b3..ae3eae5 100644 --- a/norch/csrc/cpu.h +++ b/norch/csrc/cpu.h @@ -8,6 +8,7 @@ void sum_tensor_cpu(Tensor* tensor1, float* result_data); void sub_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data); void elementwise_mul_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data); void matmul_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data); +void broadcasted_batched_matmul_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data); void batched_matmul_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data); void pow_tensor_cpu(Tensor* tensor, float power, float* result_data); void scalar_mul_tensor_cpu(Tensor* tensor, float scalar, float* result_data); diff --git a/norch/csrc/tensor.cpp b/norch/csrc/tensor.cpp index ed4eeb7..532798d 100644 --- a/norch/csrc/tensor.cpp +++ b/norch/csrc/tensor.cpp @@ -388,11 +388,11 @@ extern "C" { } } - Tensor* batched_matmul_tensor(Tensor* tensor1, Tensor* tensor2) { + Tensor* broadcasted_batched_matmul_tensor(Tensor* tensor1, Tensor* tensor2) { //MxN @ BATCHxNxP = BATCHxMxP // Check if tensors have compatible shapes for matrix multiplication if (tensor1->shape[1] != tensor2->shape[1]) { - fprintf(stderr, "Incompatible shapes for matrix multiplication %dx%d and %dx%d\n", tensor1->shape[0], tensor1->shape[1], tensor2->shape[0], tensor2->shape[1]); + fprintf(stderr, "Incompatible shapes for broadcasted batched matrix multiplication %dx%d and %dx%dx%d\n", tensor1->shape[0], tensor1->shape[1], tensor2->shape[0], tensor2->shape[1], tensor2->shape[2]); exit(1); } @@ -435,7 +435,75 @@ extern "C" { float* result_data; cudaMalloc((void **)&result_data, size * sizeof(float)); - matmul_tensor_cuda(tensor1, tensor2, result_data); + ////broadcasted_batched_matmul_tensor_cuda(tensor1, tensor2, result_data); + return create_tensor(result_data, shape, ndim, device); + } + else { + float* result_data = (float*)malloc(size * sizeof(float)); + if (result_data == NULL) { + fprintf(stderr, "Memory allocation failed\n"); + exit(1); + } + broadcasted_batched_matmul_tensor_cpu(tensor1, tensor2, result_data); + return create_tensor(result_data, shape, ndim, device); + } + + } + + Tensor* batched_matmul_tensor(Tensor* tensor1, Tensor* tensor2) { + //BATCHxMxN @ BATCHxNxP = BATCHxMxP + // Check if tensors have compatible shapes for matrix multiplication + + if (tensor1->shape[0] != tensor2->shape[0]) { + fprintf(stderr, "Tensors must have same batch dimension for batch matmul %d and %d\n", tensor1->shape[0], tensor2->shape[0]); + exit(1); + } + + if (tensor1->shape[2] != tensor2->shape[1]) { + fprintf(stderr, "Incompatible shapes for matrix multiplication %dx%d and %dx%d\n", tensor1->shape[0], tensor1->shape[1], tensor2->shape[0], tensor2->shape[1]); + exit(1); + } + + if (strcmp(tensor1->device, tensor2->device) != 0) { + fprintf(stderr, "Tensors must be on the same device: %s and %s\n", tensor1->device, tensor2->device); + exit(1); + } + + char* device = (char*)malloc(strlen(tensor1->device) + 1); + if (device != NULL) { + strcpy(device, tensor1->device); + } else { + fprintf(stderr, "Memory allocation failed\n"); + exit(-1); + } + + int ndim = 3; + int* shape = (int*)malloc(ndim * sizeof(int)); + if (shape == NULL) { + fprintf(stderr, "Memory allocation failed\n"); + exit(1); + } + + shape[0] = tensor2->shape[0];; + shape[1] = tensor1->shape[1]; + shape[2] = tensor2->shape[2]; + + int size = 1; + for (int i = 0; i < ndim; i++) { + size *= shape[i]; + } + + float* result_data = (float*)malloc(size * sizeof(float)); + if (result_data == NULL) { + fprintf(stderr, "Memory allocation failed\n"); + exit(1); + } + + if (strcmp(tensor1->device, "cuda") == 0) { + + float* result_data; + cudaMalloc((void **)&result_data, size * sizeof(float)); + //batched_matmul_tensor_cuda(tensor1, tensor2, result_data); return create_tensor(result_data, shape, ndim, device); } else { diff --git a/norch/tensor.py b/norch/tensor.py index 32a5bd4..fa69bd6 100644 --- a/norch/tensor.py +++ b/norch/tensor.py @@ -291,8 +291,22 @@ class Tensor: return self def __matmul__(self, other): - if other.ndim == 3: - #batched 3D matmul + if self.ndim < 3 and other.ndim == 3: + #broadcasted 2D x 3D matmul + + Tensor._C.broadcasted_batched_matmul_tensor.argtypes = [ctypes.POINTER(CTensor), ctypes.POINTER(CTensor)] + Tensor._C.broadcasted_batched_matmul_tensor.restype = ctypes.POINTER(CTensor) + + result_tensor_ptr = Tensor._C.broadcasted_batched_matmul_tensor(self.tensor, other.tensor) + + result_data = Tensor() + result_data.tensor = result_tensor_ptr + result_data.shape = [other.shape[0], self.shape[0], other.shape[2]] + result_data.ndim = 3 + result_data.device = self.device + + elif self.ndim == 3 and other.ndim == 3: + #broadcasted 3D x 3D matmul Tensor._C.batched_matmul_tensor.argtypes = [ctypes.POINTER(CTensor), ctypes.POINTER(CTensor)] Tensor._C.batched_matmul_tensor.restype = ctypes.POINTER(CTensor) @@ -301,7 +315,7 @@ class Tensor: result_data = Tensor() result_data.tensor = result_tensor_ptr - result_data.shape = [other.shape[0], self.shape[0], other.shape[2]] + result_data.shape = [other.shape[0], self.shape[1], other.shape[2]] result_data.ndim = 3 result_data.device = self.device diff --git a/test.py b/test.py index a65593d..27f3ddf 100644 --- a/test.py +++ b/test.py @@ -26,17 +26,40 @@ if __name__ == "__main__": [[4.567, 5.678], [6.789, 7.890], [8.901, 9.012]], [[1.234, 2.345], [3.456, 4.567], [5.678, 6.789]], [[7.890, 8.901], [9.012, 1.234], [2.345, 3.456]] - ]) + ], requires_grad=True) - b = norch.Tensor([ + b = norch.Tensor([[ [1.234, 2.123, 1.5], [5.678, 6.789, 1.293], [3.635, 4.456, 1.0202], [7.890, 8.901, 1.91], - ]) + ],[ + [1.234, 2.123, 1.5], + [5.678, 6.789, 1.293], + [3.635, 4.456, 1.0202], + [7.890, 8.901, 1.91], + ],[ + [1.234, 2.123, 1.5], + [5.678, 6.789, 1.293], + [3.635, 4.456, 1.0202], + [7.890, 8.901, 1.91], + ],[ + [1.234, 2.123, 1.5], + [5.678, 6.789, 1.293], + [3.635, 4.456, 1.0202], + [7.890, 8.901, 1.91], + ],[ + [1.234, 2.123, 1.5], + [5.678, 6.789, 1.293], + [3.635, 4.456, 1.0202], + [7.890, 8.901, 5.91], + ]]) result = b @ a print(result) + #c = result.sum() + #c.backward() + #print(a.grad) #a = norch.Tensor([[1, 2, 3], [1, 2, 3], [1, 2, 3]])#.to("cuda") #b = Tensor([[1, 2, 3], [1, 2, 3], [1, 2, 3]])#.to("cuda")