From 5dd3f3d2026e2ea8a84c85c8f6a242ec9d77a987 Mon Sep 17 00:00:00 2001 From: lucasdelimanogueira Date: Fri, 17 May 2024 19:08:09 -0300 Subject: [PATCH] implementing keepdim sum --- build/cpu.o | Bin 11248 -> 10976 bytes build/tensor.o | Bin 43056 -> 43312 bytes examples/train.ipynb | 33 ++++++++++------- norch/__pycache__/tensor.cpython-38.pyc | Bin 15409 -> 15530 bytes .../__pycache__/functions.cpython-38.pyc | Bin 7881 -> 8163 bytes norch/autograd/functions.py | 12 +++++-- norch/csrc/cpu.cpp | 25 ++++--------- norch/csrc/cpu.h | 2 +- norch/csrc/tensor.cpp | 34 +++++++++++++----- norch/csrc/tensor.h | 2 +- norch/libtensor.so | Bin 172912 -> 172864 bytes .../modules/__pycache__/linear.cpython-38.pyc | Bin 1036 -> 1040 bytes norch/tensor.py | 22 ++++++++---- test.py | 24 +++++++++++++ tests/test_operations.py | 18 ++++++++++ 15 files changed, 119 insertions(+), 53 deletions(-) create mode 100644 test.py diff --git a/build/cpu.o b/build/cpu.o index 0b75e023751f2236eacd5535a5a69885d1d11141..adc39e5238ee388ad81ef0ef6924b00d82d01d9e 100644 GIT binary patch delta 1324 zcmZWnQAkr!7(VCD%&B$n)NN>F=``mwu!3|*BISV#CFp}B5+y~3kfxy_#M;Bf=ynsz z15rfaLr@Qf+CT;wY!9&&Qd;y979)I^C}=szKwSUXIWxTd@N@t3{onun=X~4R*;w0s zx71lN&c8fl7Kow zRs^JzHMRg&YH>*?3C)*4LxToMsE^bUf59-L0bvr}Vq_U)?eXtF`aHi<3!D@6BgJ2{|`KXZ%y!`eZ{b+XuZsm_!sO)rG%DWIf?mTJ%Qs9d7h$Nqspd@>h+tR5$ znHF#0d7(GMngRizrv3Y!{x?9$Lm1u60|?>_|Dc$_BLY_io&oBNFXwaMg1|ij?-lq{ zfwu@er+@>gKMRyMIPO8KqH#_PJQM?`?HuSAcz}kBorOWsoiTO41b!ArrJR=T?u)4J zdUx+YS0GYSf-_+F1H~@qE7v5(pXpSY)BjE3imX=(xd}fCi_7Ar%(GpPE=U&F1a9U} zoA_={N6*C77~cb#*0(!SgiOR5)q$7gPq2U6b+#0OeORVlziPtMuHQ4^Y1hY1c-r+@ z6D~}zLMx_5+I0ndX2IGlXYBM>#dmy|&bw=^b)~F@P#?8b)=E7NI#uMLot5r%IacYm z_VL%E8afZ0F$c}|G)S|8>TyeR4tf`mH3xkUpBH5`yQ;x@u!sxS(blS3YkM(2eWefJ N6hFjI1GKgBz+Y#@V0Zuk delta 1535 zcmZva@k>)t6vyv2>&3KfkJkD|wG zAl+Yz=sy^^8pJ9HSHujYR_q6fRig|h4jL#_60URi?z1^MaNoW6bI$pmd+&K^U#@pl zm4?(u?h}eTsu?3;!xz&1q5R0v0r|)HoLmi)aYzs}&4j4g-R6OWzm2_Ewf@MQtQ?q} zlaIs-Q92~(>;wjdkpuE-GILfU1jWdeN-J6QKX+Tnk6rH_mLULJp_(eg- zh!<`YhqRF8XUF=>&hCM<98~afahT>|CQNriQV!1JIt~Q^L`DFz;!=7>G9ah>bL7g1 z3Ats3@!Scs!i4PopTCXw*6F>t&*2OBS6k7sv@~0ZOMRSTF766XD38S`Wd5ty&h-L6 z7Y)ejgE=~T6VJ1(Y8Vx;&lOfZeA2W$2J=SDgbbZns(}&wa{tLJ{4_*|X*xsl`MeHk zI+Iw;=ij`Vlj|3Jn(`jmeI;}}>&DA~9c9ccTliSSs?PYbRIzN*aPVZr+Z-!8bk&H@hw zKO}fa@Io)qch%x0x9nh(E8=YmQ~H|GW>?^yV#%z)sR_p)d}eOh zBEF`GFADBPR#n7nt8s(IU%ecMty3pB*O!ZUv&9DpcW_NR;Y(4%<|5uJp4%?aj6QB4 z-hu0E!(r<@DdO8B51Sv*z*2{~1{$CEuseLNgxejyQ^M^Ihf28J;mHy%8oZ5UOB{BG z73?WD?&jmTlYQO#8JiVr>2PG4xQmbkqiTm#PUclrdpWAA4#j1~w6Pe*>~^sQ zwM#lJLR(e7yw=(!-FC5OI7PRxh1M=dw9ZQMvL1JbqoTpuePTCo_p8O)4YMA#;V<+^ Bv|#`M diff --git a/build/tensor.o b/build/tensor.o index b8560d82b24a628ce8cc4ff6e0ead48af48ba712..0071d08d1c19ea6783ac31468df999592c89e296 100644 GIT binary patch literal 43312 zcmeI53wTu3wa3o{l!!=zS|1g4v{ex?CgEj7F@S@kfYM58MQs?8iA3{iW};|~cpF8> zF+Q)gUbS+$wc3_es;IG|2E+!HYFepMMVl(N&|t+%Z>iU-cddQan)#nS=MX^c?IZj9 z!p!{kKL5S;+H0N1J~M{}k?B+OLLo;>$eHM5KPh#brT1pvmz%fc&QNDB=lz01KB$PM z-j23yE{~>K*F{q;-KmN-(bSC2r;ev~8@X*X71mnpsT! zb~?#p-0jPYNshL7In$72yUIs#k6n??Y(no2nfNgk(`!1ic|v4OD{U0nmYz9^K1RPC zZ69(qOA+gihE}JYf(N>}J~eYxTV!*n_xW^W4YgX)lb10gV{g6~njgKob0@vk71_e$ zAn~@1p;TmRTkBS5{>iGBuGku_`ek&*n|aaD%5-E)DlZ+`%C_zcwsxnTbj6l}2Ueyc zTiWWkhI$|8&7XZ5X1Q2rxjLF2c*F=ECe3b1C-u%_?uu;X7usxHx_m@gWNT}gZ$CVV z>!RuM{PYa6>)Ggv-;!O%RVzc(cSQQ;CdJYAu|`uS+I0+X&Bp4W+Q$?|+asGF%cG`0 z9mpm~^WKBquCQ^a`~8AKirrz6&Gd%mj7C`<9qo=Mxtp0;&!`J&)Y~Yww@}M1x%M2* z0{*5BHC1ZMME~uN3f_JQZhw=x`x;g=jM~xNP2B-I+X|MWgAcRCBW(mK8a_#EtY&#`ZW#4qR$xdhhXO z>}ncSDzY{mS!-ueWFu|7(P=qvC$kt#?|ZM*73pSYQF2d+tSxANhTe;Z3ey$c8jqpw zbfN2^^1}2C!r2(7896Ug<=G-XP)J>yZt0$sf(Q(lH+HaPjoVTCklWaJ+6``uhMr41 zsl2vGcWBdbv_F%+ebQ|vSEb8G!S14+@-G_>W+$DVLE%7imU=-`k3LeC≥3tc|8y zHY)DDJgeu~2~ujcwQkw@qXg{*{5sEe%tej=#Ycjr2OrQhFF*Kv@v5F$v>D{e^rBXbv#x~uGJ;PhEy}?_pr+}CcS)VLx zYh4#gj&5sRA4*#GA4BrnIVErfuGo!Udk3T=>nXI_{t!B*Vm0TTH@G>pSBEwYqtNb( ztZ@(Vkh!f#@ZlK`fRqvD0Jz2+0M~G0^uFYJ#P1+9+dYiy0koa=T%L^RNw=<{xZFqq zN8X8aOiIxaM0R<`%_dgtGX z;*%IAK=A^8xb=k)SS6@SZ*vZ}aQ+(Vv2eD{Mjw3sH4n-FA zbeR+2u1F_`%%%hA-z!^}It5py=p$Q95uS70>FJ8BqH3)#i>zuLiz+;C{$P zYENc2yeib&nT{;;ge28=yY@)WsYvIhtFs+Zs`=Ssa-4$RW76u@$tLZWCBZb4H+)jq zKBSQ1xqW7c&)zXB$b^;Ykef`s4~iM)43p0Jy4|!|%N^12ctWI;(~6JFnik|xNvm#? zR^4rpwcAarY2LK-JJ8TRDjO|i5$%j-Y~&OBbPMI68*-Z#8k0D~KZRq=DFu(vuB~;M zdo=5&jwqiZ9g)-uPgZBG*r`z1%q{Q)_)?Wmtg%T za^|#yuRV6@M6Unn(|`2oUq+xD`gE5OOj_|0&Hdi1wjDK(M8!`X>-NfRlq3o9ilp-al|6Pp~Uh(hWr_kB3L!QMI?EZa<(|>Qs z-O=~4UjDb+E6(y+;h77`Op?^>l|qA`z?3t^@MkqTtIVvU%S~FUHg_7RdZ+G6 ziJb3A>N>hnx6VmUbFDVLd+!HE?sxBLEuu4!P_aQyGpb+?YTddUG?t_^*%y358QQhvu;m$!n(AM&LX+%=&Q^+`YOJ;)%#L9 zvMT86)w;3P_jhgF3C}svn(p?YhP2aNvuupfVYr*Jn;byF?78#Hn9kJc7Fc#szIVQ# zy`SW+Z*IrgPQ7dJyMEWVgTo|~*-JB@IJ+;Vfh);3Ijx@tomQzA~zH=bCA6*kG8EU^n5jg`Z0OOvOZZ))w^~t+Lha0)}QE_pOf`;!)X8|+W%LY zXPjL~?|M@Fuj}%31?^>44&*O8Xd4P!X8Gyz{8ZlGN&2}HnVEXM&KbM?^rLM2TjUvq z6W?`%uXY6G*^b;c_WyQ;nJshlq~2)9%=yK!=6>=|nvA=$nc~S3Hv-x^^Rv1pll91K zv@=I$x2NJv|fHfzhk6dpzB|Lu$CYCmo%+!qj@5w)0|{wj$L}F|Ce3T%)Odky`9t0f9*RY)9Em> zdi(3%^ot3z*pgWZppS&Ii_UWAsds)>%X?Y=2`*F+$7h33ivt?9&q8+qUziTi2E3@ti#t2P1tUa(g#nOfdRMq&RxigKM3(Kjsq`S(%%VIk}LJ9;DTU*pi#9~S*#;qw#1{;=o| z3!ku<$tSwG=N3C~hReT@^gF|KPowBAc(?mY3h&9(XZXeX%#WhwmnU>1U_zw3^#s>W z^SvMQTdg&Czs-Ee-fznuO4fvWpOiB~I$;+LpK}fRy$Jr^WJiB5LZ_eIbZcs(etT@w z?pZ?8D?%zw3#$Pk!g}}^3Qt#P*knStywzTf) zjUk(J&)YNe;9l~%!`!K>FB6pf+pdb8{AHKi7sk+B2P1XolR|T~d*$N|&Cl>&aK?mfS4k++jpL;l$jkii@Mc?^2HJz!G#mDMKiGB9D7slL z+xZ7P{zA*|eyk_FNQpv#Z`aV-diGxIrlaYev){cs*+W>N*5~bG`Z|_3tg+8D{W4p> z?n=)VkYCHuovjJ_(Ys@NATIxB2;6wt1$) z(>x>M_F+0#i;W>WvNldz$z_i>vfX`1!Bw}A%iT7)KOvMSWDLWeH_fFDx+3fN+X1qj z`Q_bu`sFM=+kt-n>iLGiHuGqQb>eth=9l9btoC{)f zS)`|+{W1F771_X;@6SXw%2P3<^Rme1);aQwi_CK)sMqdvG0pHDfB7i2f?K+?wi&xF zcwj{+vcY%Le6#6aQxzM{(=qs!}DSn$HU2a@o*wm9}hRQ)X$AKhZ}3c)wT81 zj6@RQBdWu(hH8GPX>1P1s;g^Re6Uw%ZahpKW`wHa7uQzBPY558aJNki-fq8%d9kKA zbQDX5Ya6QL^K&~T`_JOWq}xpPHYUSWjrC15+_`meW8d87##nV#ERn2jxPUtS#N$o0 z%uP1Os@x&>rPtRuJYt`09^!RznxW*em((WWG?jJ9+NQc%ci!D$d*{bH&wbgsUGdWz zz~x3MVLV?SOE%Zew``#qoSz9GdPlYeZJ(LBcy)MgELk;=UUx{-mGv2~iD@I1jb}GV zSf0hnaoMmFzfFlXa9lOVT`wg~x)g6D>_mjm^nK*lm*tw={*5jp5PR zAtZ`uyHlr3IUzivVs1-AvL#$rG^VKd*pe3WspN{%;-ca)MWsi(+ypKC!@n#qq(A;@ z-l=C)IP{pAPGgSRziNHnJ6vt*&y#xd`J^_lw1rT%O-M(YQhF+lb2lKP`O^@n8Y zmrMOLPyOUf{rW!iKkP$)IFFl3kbNgg{YZXxPsaYk^3?wG`l!FWkNUg%=zqD?zv>zP zFS6s8`gxxE%1r%-Qor0&|4^oWIFE};(D+X5qy9Xp-{0e(p&9>NF7@r6`l~YaclELT z@;>U<_fh|0AN9l8FH~;hcA5FTK4agBQa{?`zmXaLohS9rvY#{CcbfcRUBALj2RiS4 z>{mCy_X%g5%hHjx(EyYE7Nm~@%-+C{56Xq6JteeH8Q|>ULs@_K96p4|-eAAkrx5!* z{qg?EF2w}=`1k-mF@Q$`_~`-sO9A|=#Myqe+wAj*#RA&T3*dLb_{Vpo$3I7tLVHd zp=I{%r1PDC_Ph`D(|<#BPS}e!rlNk>6wrP!<+z^?@B4UMH%q0CYf?b_uLtl&qSLXD z*PcfM+HVNpg*2~hPr?4Kj@dWU+rx?b*&hwy^`cY8*T+;;XF)*wdjt4$qEmH{S7&`d z`#d`N@U!zM;ykYQgIyh+ud;ylu>k&-=xhkPI_B6&Z?6q#za)UaC^}Dm(bX}>POFE{gETw_RDGfTt?7|m*0421@HyL*$;2>^Q5Tgyxbhn{?P!wUUY^Od3F96(0*Ua zTR%I?h_gN8id`LZET^}p1hk(Uz!!+lj8a#}9P{bz4FT;R2;eV@&a1qHNJZ`G31~l% z4rzYlI*K@t>%DPaow9)TUkl*h6rKDsuTFbF`yU4I)uMCg1h39Z0qx%l;6v$9ZRX27 zmsHMONb*qPetwt|z-t5ejRE{&;@t1^C%g8Wb0T{Cyzq}EyWE@$(c9MobR4-zs%mO- zs+!}mWW3TPPUWf7&zwAIdgYl@r=As=RXJ*5XkQq^>+b3r^=Nmcdcb3JL(>Iki!a8+exA{nc?sIqF_MU}jG;v^d@-K7kt za&}1xEumBr;SFYvn`N#GT~5GMG__=AyL3r&U&-8%cxe+KXWvDxRfAyhgq66~|Uh zE!fCS+e&WW_QT^Y&7r=Q`?T^}4h`J5wzxFFtktk8Ey>wBH@YYlk9U`f#5Eo}RjOpX z)h27D+XMqtmbWTBn~>e3$M@OF&%V-ezJk7XmW~Z*YuhWHX==R0XVS9N(}{*#&bDYZ zgd5AaZOx=f5N4?n$4H{M(OAY--f*jHydbMuGB*3sW2@|?ur})gs+9FVP&z)ANYq}? z;1>n1*m$dfN3XFVo~W#=y(oJWqp{d9*TT}`+%_KF%i_(AIl`qm)JM6*r;$fGQLBTc z>7m_y=3}l{$r!Bls8o@8nMp1goBimgg0&tMT7-93NwZp%y%x+>EFGO~Q*ugX>O9gE zRM4t)hgyveXqD|PHlOBO#$6wyX|HN*NY-A^(%6!4YMSG5=P>s)$gIrpn#hNIhnLD| z;UnplFIkoIIt3TH95kyo%Fi^^UEdu_@0Fheya4z@;Cvpa?Qa2&_B;-J7_@&`xNYY< z!fiYM1UhJ^uD#M_lY@3nrt5Pq%+a2haJ%36!tH)rK?nQ21~~S60Si!3o&DU8PNQ(P zhcDgLZ{GpVYv{^f2L3tVe-X|vbP49D{>RWB_4k(RM{O^c*UsmK+xCnE&PxMo&-ehH zlc7EK`(>-W+BrL*{e{2}fPNnpZqL^WOaF_s4VM={ht~wv{&m1{zTO9p_77#pP_g?x zK)Cv|&TZx#A)H^Z-{XK|zb6C7Jemd^uyfLH&1uqkb(1Hx=9dONFcbce%}+g~C}5{eKg1)PLF1(fE7=+G9NA z(YgQ^*5@2jza0pCIPhZNgMpt49OvsC;r4i&gzI><-xkn8fAabW7q%0}`z_$SG^RQ~ z0*-dB1djWkJX$Z{!usf+gN57n^ZE!EwO{jP4Cvr|O$3hiCxN5i77Dlfy<52UtK(e^ zI_S40z_H({w7$T_)<0jk>ThtHIdg@xe^9>;IO=b&}|=Hq?P9^>J8;FynZ0zU%u{|bCC@Isz!Dr_gt*WtqL@s2nI+L;E9`S?6=^v~~v+xBk}uJ&ub{4eO>d=28@ro#R~`%eOn zemhgR-EWO>?N|M95$K@bT7YA}1$!HFa^mnv;i|vEZRQ*!ob5;bvA|LPVM|Bj@Cj&- zetsP|#^L_}A3=TVc=zLtsjxjbzoUSoKPL*e{aGPg$E*6YK?lcM4IIb&P2fk9O{#M@ zaJ2J5;24J=3AgRcmoLhxol9i@@Hye^2aLBvfuo%#0OzjM&I;jnzm>wZUybuR(7}E$ z29EtM1kNd<{oVo`=W7{o4qEO9QBt5=yzKBL%1Uz|6oZ#wNvx)2zkTf#dtUkIOgLV;QSQ< z_2(tP2LrzWIL_B1;kG{?6|Q!w{&LVkfBp+Mdopm0 z+iAdYT=l>)Zm$K7?e7w9`{5zswjVk`2mP=TIQrq)0R5f-{kK5}_1^`K`jdFkjf&bS z?sZNRZpT$UaEz;KfESUT`u{t?alRe~j{fWtZu|2M;cBPqzXdut-gkiGcn8V%m3Umo zft`l{M>|IY$GEBmj&aoj9LIG#aEz-5fn)oh3Ag?5x^UYMZ-Nf`VKZ>_!}|gHgL!8} z#nvwn&VEDv1A(LdMV5}nRRY@M{_1w%7*{_5UP68Aysrk1^Yt6x=+F0r+y2bw&ZyYq z-AlMV-u;2&cuRnnlAhWb1&(%}4;nJ{Ih3299=K103V(CE>Q6uM1Z@^}O^Cpu?%7 ze%lHh?JPKi8c|{Yp#Mh-xBDFGsvTRYAq(7}GE0mpvN1&;Z00dSnJ6mZO!?*hm6 z9l~wuj>_yK?nW6 z1UTyd$<{SGlAIAb8PfgUH(}52LUIiTIt3|jy-s^?ycvb&q(82k- z6F83dLEy(bSv#Kvj&{BQ9P{xjhtbAVY&&a&tDV*ShDt&>&pXabD{!>)+rZKP4+*#X zeL}eQtNHjG=wQEVfn&cLfMY&x0gm&v_u*!ECLS;!%YbA18NzM*tA*#09N)^*`K<*V zoZkd+wEr^Uw*Jk+^DO;ufez{~0*?Cayy#5DK91ZY-0t@g;dZ}|gAVrl3~=mst)-)R zx*pnNoO}oz^K_ptX6-+L-m0HR0v`I+s>~FS3AGUZ>YqD52W{)FA3mi=MBKo|KAsG_xq4=?N{@(6LhfO z=YV6suL8$B{T*V8_G|pi1Rb2;*}&2MSb+Zg0R2|b zLH%ohqyCVi%y4t^^bq0pxW)>%$8`eeV85pT$9~VUbTm)Th4vUH-vEwzdJFKuG!*so zy})t49u;nn>jmNVxLyJs9M|i>aa>!0p9J>o$&12NIIhsnFmTMr`NC~GzbQOpC%>U` zBj{jW+yNZz+!CPg@S-pk)z{ z0FL>79dL|?CBQM?R{_WN>xA3(zbjnr*Erk?Iv7760Z03HV__<`{&3;8{$av-UQquS z;Hdu-;5e?Qh4TyFGFAIu6>j(Y8_>aiHvz|fw^=%x?~eR_i|xla2?NJ`9}oOwu=5Py zgMr6@<9HLo*(@CIb;9lO-UvE4-aCNfc$We%2Yb4Jqn+!4<9=Z8QM54?_7BGMA;585 z<-jqX&jpU{8-?3`SSZ}~LmG6@54QnFKin0d|4@K_2k4;wO5mu!@3F=`nf;Z<)o|f< zT$KaIxS9=oGD&Lue;qi^*R{ZhK>Pm?Zu|2m!fk&(0y;R}PT)A+SAkD~e%}L*cJ5x} z`iK31adjeajH@Vc9M^@wF|Mu#j_vOdZu?=WaN7?LgU(PG*HggJ4?h!b>%T7C)_)Up zhJyZP;HW?GIM;8gqj5D=xE)s)0>`+z3U~ziy%RXj*Mq>(pR0u1<9$uI+Np8d13Ea~ zKL9TPdj=M}A?VZt$8p^V9OLQ-z_I<4!fikNO1SNZS3w8; zumL#wp*KK3uf+9_tv^UO`wjJn0Z08ROGo4CB504t!5e{NTrCD31wTIq9OtVWIQsKV z;kG~jD%>9LN1%h_&F7l|RM-zV-h+UjMs*sWrNGh7Q-Nb#EdX8!I%(iIt{(v(2JKe? z$M)-l+kSXgxb26npo4z+2srv-x6xz-61f<;gZBH7w9fBvzS%&9?Z>!14)`>x)Amz=V?3M>9OtW1xb4q{!fk)1 zK?lcs8*m)&{lKS#olgNrJJ$loxZQ0mZA^vzfN?tvIF4&HaE#kCfMff)!fiiXD%|$N z6`+HDxE472;l=>{?*-`p5Oh%gVc@7gXq<6R&iPw`aP}L<)o9=tS5twXLEC6toev!6 zD+wI^xlp+6&%1@&{#*2;exb(}81L z#erk{ONHBhxJkI}huc60{csO(^uzZ8^q&aOUkN&>zXmw!AHp{+sHmO$xgVV)h1+p; zI&h4u^MQYfG&HV~z;V9Pz|o&~3%C8bOt{*q`j3GQj`wNcINq0me;Mq22RPc9SLQl| z{eW>b8aVE+P6m$SiUG&Cx*RyRze%|5hx>%vepmuJ=!ZvuqaU6K&|e#%|7*}e{kMRl z{@4>-zp0%XS0@R#<0=Lm<0=V!7R*-~IL_B%;ONg~!fk&(FI?@^xatNS9Pg{ZalG#U zuK+s-PH_Fe<3&3U1defaI&h4uuK>q!%?FNgbt`aef1hyM4^IlW{qQvCpdWq?9R091 zK)*LYe+%fK{s+KOfBK28Kh;i+t69SBxS9|AAefgl@Uww029ERf7;yCG^TKU^zA0So z)VTUR=-_zY1CHYz$cy7tcwAqlKh-%1INDhZ9OEhm9OJ4UIF9Q^;22jw0FLdS6mI+B zSHf*Syb3z#hYi5d54{2Uc@xqjP?gWl;^<&_( zsc*IC8Q?fyuK`DYz9ZcB=RjV}rDFSY58?KB_W_RMEeC!s>8YI+z|RAIA@K8oCxqK} zUMF1bTqf63H-Zk%%N@Yc&ZWT7Z_fb7IR6kh#(91@8BK-#gMJ=a&eY`?=aF)QbK0LH z+#YX}a2>Duvjuc;ybFNicozof-yNX87<5p732@Y}p5&OKc50m031|OcoL>(dqj4}bQTz3Ij9~`A_4$hB7T4z`&at>Y|L{tS>+=j3SzMn_ z*lBTn-e8Z#_4$GNC||DTKYboxSoX)t$I0=v+~WHDy*U=w=kGa1iA(E~vAOh`Y<6Gg z>AU85U92b>pPzJ!s$Yfp6^K?dE0y!I|s-sf6ausxK*g*4C0#~KE>{T6t!W^UERKfwq;G8e=PaZ|2vkj z+OKW-vX8efRv@o`j=FEyJ}hd#zD}k}eUI%9>Z~m7qc00aSgn9n=pg-+j=ijey-m;SqQdAD$ji?NvKR!O$+jF&L jA`!<~S7Bb99^JtGnB!|c*YVr__ffVE*nin-ifxhyyG+<6;1DANSfg4m7k{Vhq$_E z+wqJ|>4zD(`=x2WbOnhPM%~TsXlj<(PFJi=%^FR#J(}ut3%9n1GTW&&W{1?Q@lNnh|%0k>fPdYs-4^18QIDov|CTQZ1mK~)|RQh^>8Ql zMAK#Y>2k8`(dglogfjnziAl@>qVtUro>OyMuI~3bnsN|Mhd&aPU zzo~ngEVX5t|Nf(b_bdMn-j!QSDY&&>mfBcW`hYtVvEH(w>g;aB(v-V=&{{ z73^WtZqM**kEUq|(Qc-3*>SU^JvoI%qv@z*bFuA~6*<3w8|kGH+v_AbaH*K7z1KTp zyQo#E$hvf7ojr;o8)@f_PV<}{9K~Szz(=LdNH;r+2KTX%b;H|!P21w3!gNKq#$%{E zUFdqKtT0_pI2+^CBXcrYo;~sdh19j_=I#?y5P<>x#tycuaVu&oxQUIY!{Ekf=&`hu z%4>~uhc+Ee$1|zhXSmJctaRBpI9#;R_{)ZaIY_6=DIBQJQcvp8qpvhd%X`xm>!RuA zjf(rO%<6d#f;6;RTej@@S%QuNewAk{=0Ph%@J6zP$9kzbcKMCNgq?+%` z9VIe4e2?z#0Xlh@>sE7gr{hBAdaSkOz7K{kxsF`0gGca(`IK+YU(uADLU!=-i7e{v zG}pcoEt{5*B-`nRbk4b^S55?-k#$s1OHPe+w@f$ZhWnDGZYw=U;|v zh|$-}*JUC1c-;51Y`rPESTy0{7lSk1aswTiX;5c@#?KO=IYzy}dTWil^Q>a_=1YtWFb*zIP7)Uh_gNeC@x37b$%m zvX9Rc%-NjP{}fDlPrAZYNIf}!c>S7;AKmMH`$n_x>a65PG7D3_C8++LyiO_OD50f< z7M}^u_U(%lD&{wX*goK_%0F)g(L@S2X9wo%Or*Fj$jsoH>mZsN=SZRU^><5HPWOKm z?L2+KH)+pNu(Z3k&izD+NPa|-5>)@rMhee_4<1w8X$a5E?%*-S8C)}Q5A*}fBmeV` zc!O&uGEjLMj*jC00c$36=>U8Ot(h#%y{5Z^Pi@iG_Mv7yY^gcXIkW_JF*V$LGkH_= zTY5ZgW%tR~Z@jlIf}NOpcY7AWNZp+LS&gnLGn34+cuCo22A;BaXthSOuNYS-Nq@5p$m`;#;fcnkoQ>c$wo#Mr$*7iK}@B^*3 z(8Cy>b(^egd#LZ{oskZA*`|x_q;bN!v~0s7y6i%;_Ofi#VU}$=c+sct>2#zc=4WTjd9<~o(DTRw)nmpT%leFZ%HFv(T34>WQGc9ke$J?;^`s#tesdS~+?(2u z_j!ekvkPgfH^u+DE>Blb^=pRmLjYQaqK8?2x-37H_b)R1+>Xq2z24_c?0@+Elg1f^ zlfBlJ8Qk7+wmr+K{!bm4-ZE!Ts*MiJJia*A+;9GaCgZN`Nb!skHv(GQ^Rv2UBk zXlKsMZcXDek;7+m-+H(6%uITNW7KgxJ>vSfkbLYnM$3l&TMp9gIY&~%ts(R;H+Ba& zM-EPD+RxtTgJpx$57zX<;FPBIl$Iw_+Ra5)=G>*H`v2M~&D^``u6`~<|GQU6_)(~? z{^z~vmjq_AB{LI1UkPO=o%K(L&EwH7`$9+?g94 zd{sBV`q{B$2^BM!U%g;aI&x?K&zXMAQxVx){B|6Zhe2ee*jHi3_pd0>Q`#j=# z=MFppL4h+hveteWwU%z%`gs`DbtQQ`=ZN(op3BQwr{iFxXApTo=c5Zo?*oB58I0z- zCpQWQk5OjjHy|wj^N&%3VZqatG?jusvFPtg_`l=SH5e9yVd3-By1}p*3=5yIn8g>m zxz`rkafQpj9Q3=wbuXi66};R1;g9!X>SKJcKJ#N?x<=)`ruDvK=~p7h$mJtH_hWv; zv=-Of%$MwXTlQ44Hq`f!ToKX*yJ+}aYs_$eADLaA@-9T0dm-;tpxz60)6YWaM>+0N z$EH6Ui*tW!xl^lCtu5U{lGFI#tIO%%P%@wX%}egrN1bMFlF*{Q^ZX#KE9RE8*7Qor z=G^P{Oh34{d~P#$=|pRpImF!2Uv0)`c-Nz1?AcWt5^DShKk44b6{nktpMm z&PY4?j6##vZoAAPFgJ%n(=5=MUyTp2THBEg=j9m4_xZ2zW=op_T7%g&8}_;$%%$lT zsNdg0QNZ&8TQ=ak;r5n-hWeZ5nc*H>=b7n-CHp(5O=D=y*>6c+ju2L;4Y+(vSAllD zycFx!nWkSX>*KEUn|oSub9Gg{ zJ7wK%er}d6p2_eO&xp8nn99{+W61W*jniIo+vAOF_ZTw#ny-<|-7=$d?sDbL5qn|H zo959Dosk}XIzaX_?>tyfzbK`59?VWV>(!y|Pdut9rSmnfFa>p_eGo z8x8Vkn%QZM^oBM~rV%bXW&%wt@qVXcs#zK)BP-saw`}D1h3IVqPIX=j2*vvk0oO9y!iH4>`xUM;o3@?mb8V@HI#>0tNT|8Xh zTsJS?6mFOwuCA%0VkD9XA5k5S)mQV!{D!7*th%~}#fN!y=EcKQVMeGrerZir{Fv|& z33uPbuzu@JEQ~eAp`utaTvJ~iUzA%Z*?%?{Cfz2ow;>s>YN%_Z=FY2)8~f%pHN>i` zVu@r;{Q|1=L$^24JTKW4t8$w>klsM;@Q8i5eu&q`sfUtBT~?EbQ&-j|YZ_~7+0LKM!PQzv7>*2bUYAgzjSWr7 zMA$8p2sbx|lMUes*(M~4XunfWI_a43=!$vG^~vV&)S}X&@kbRmn{UNem5eVMUs_Z$ z#^okx=^y^tbs_!nqj{&EUg6N|SK68J123HRaX*~UyQ^hRK(;RXL23qE&kN}yjMBfrzpg)4!O(|&q<0OBqr~X4S_0Jg~ z|H=XKZyuojRg(X_$3IVH{Ig#27kct5Gx^&jf0ZZy-c0@oZWpDX_8m7s{v65Q&*Psx zGyb_!^4mQ5*JSc<9$^1f1LUtCAb;Bc`6Jjbly2g3nfSUsW8ZO-Kf&X_u^Io(k^D#5 z&zb$(&G=zmzr;m{I)D4XFK>wNo4Xp7CGMwF9l*aXI# z+||)`9UoA>E`Z-GI^BD_I%?0o0p)uGcmee*+cRokSH~QiY4afBe)gXnz-vV3%>BJO zR|J&*K>&Y5bgtlmOiAtO2`K-M0Dc%QGd1t zl)pEC_lVBnM|pMr5>UR7#;u>76N$4uryuR=m~%O8mIai*IDlUvI`Lvx$DH$N^D6=6 ze;U9a7oFE9c6<9XYw&# zo!H#)qf@o%uj5=DbL~Q#2NL)5!wCUA9>CK9{2t<5@5_85rev;(=(AJ!zS9kM z%(W10J{O?l$W2mJW1~~m6ptn2l`e5APn~(j^b==Ro^i@4XGhMdJmRpVl3K|EPWS#{=n9ck0-2+f^vR%K-(8LPUa zvTETal{|UkBpWK-sSKxbc5yLHp;Qv#9cE2B$J`gXoPfz_Y|e0+CsA|fH#XJOC+E8r zxki%O#B2$V;@RUXXU%8XQY?^FcJnz8ZkV;I%Go88{j+lGm{n6l)s##MDsNTE#S_hS z*_w(cX1{vW#eB@oYR#I@GV{DL`qd+obwsv&C3S(i*v;p>07v*|<<_B&2r6$?$;IX7 zpC?}`X$~RVnv%)>S-G_u-v<@8s&o=nvD#RZRaHq2_0%%o^g~IpPpPat7m@3HTEo(YWf-p;sI7Sl1jm9#T@`hV&!-A}8@x<&`kF9cq!rZJ2 zC{yPCKxs-Wk*HZv?-vEG*kr4KN3Wqio~W#?xg>iAqqf*E*TRzVxn(@Mm&cnLa)e8A zsE=~7Pa%(VqDC7_-9v}^%-3A8;!@1@s8o@8nNBXAnEmRfg1H_Qnt*p_NV8frdM%i% zSTZ47rugJc*LkEVsGwBIc9ohCP%2wnY!UUhw7V`w-CosDpR8HX+|ZnG<~POT<^t#t z_cF-L%I9jx~{ z;8^bh7NDd$a(nGG2xogfNq_3MTY&Q%y7Fg%e+u~9!uf-4!5r290Lr8O-gKYLNy{JT zesvBIZrd{!_J@ zGw7f{uLh3o{U-2{V9#B^(atr%@%WQR^97vP59ptf!fpFU30M1dyp)0t_SZDvXnzto z`t4fbcD>&duJx+lZUr6m+nvC%-cx9Pfs?I&k#N=D;1+Y{31|PHel2j+-)!k<9Bzg3 z=;wW8KHHAN65u>#sqHNXj{RE&9R1lW+-~pn!nM7s|5ebz_Wmz$Z101>!(it+;ArP2 z;24J^xez7x10Fxe0LOOC0FKAci-nIR9lTy@7H+rqa?ru{UIQH4djoK^XQ}Y9R=p1h zx9eR6I#}-`z_H$~z_H%OPf&%FRA1w5k#P1W#@o%n@wj(4@Iz@IjfY2ovkA(73mom) zAbc#9MgPAqTVF4xFwU0&NBe&gpuajm{}IqZ{chl>e;OB}WZOSmxZ2O_ken6> zxBdTl;HdwArK98VVJMIB@G@{5kM99Lg6vR#?!}c+VtaQHZnyVL;o4r+ zITv);1m$yq^QQ94fPV(~4ZzXPWx#Pfz6>1w^AF*+{rMc+l+=D5FP{=__t(L|SyJt< z0FHiJAl$C^D&bnM>a>Cm`mGH();o#^8zozRnsC+M;1+XE63*>H{c_-_|AeKZarg|B zM?b#_9OLklvM$Z_j;23tZxL|p-;;r(KhG6zx3^BXwpaC=KnL4<#vHbhOZTt6@2MN`F)!9$D9nXgVNBc($xAjj5&_5Y; zFn&%4j{02z`aPC@0XKxwE1-k^e+@Y5Pvjt@q<+xxc)W1-1G`ZD69bOpaS`w%q5S^< z9|rtKz_Gtp2)F(D8{ulF+OrOH(4Q{?$M(Jnd>q(QAPwjCqMe5T=Uuhl^MGUARsqL$ zEdh>kyA(KrK`n&Pu9VPY~>R%vS?Udb} z7?j7jS^^y7>ifWps1Ei2&w*oqJpmm3`I2zkpKl3QJ5~RE(82cZBF`V$?L8Ry(O~B! z;ArOz;22k51degl1{~XUH*k!r$ADw`7lqq?cvHCThxb4S{V`iBa) z^+$mY>K6mYxN5a@G_KN69^>k6;22kr0xza^XutOY$NqX7IQnyUKFCtC+Zz^cxAzF) z>^E%h(ZI32p9Nk*`>36Bfuo(xz%j0V3LN9=0pQrK-vP(C+6o-Y=d*Ju*?u@!xb26- zh1-580*-!|7@$8hK>tk8K|fyz9QE(AbTqE+hw>O#zXQ&j+AnVaF9kk~Cv7Qld$GR` z0gnEhAl&xnX~J!Po&h@8-V1{A^?HmOh{Xbo}UGGfcT5l~^M(G03!FsEJW4#IBI9{#>j{Wrm;5c4Z0mt&) z!fpHiEL`nZ|7->w?BBP5qy2vm(BJ1E*B`e2r-j?&@i5@1|2^Th{dZaVI$yC8bkP4D zz)^pgJP)gO>Ui8+INOQwFbO!0$8zA4X&?1xHSl4;uKPSyW0=wN@{ z3mn_~81N}z&mVxJoo@lh@mL?GohjLNUKKVtbDdw!22)Ch%}367ya_nkxg0q9|8e1V zz0U|&|LA!9Bj{i}YzB_?ZUc_vG5-+9B#sB{ud%>!Je~<0%g2S=_FpMn?brVOBIscM zrh%jVUklLxae)5+fez}g0FL@UJk;M04+yv0^`vmST~C7!*82i*tar1eqvLcdl*c&P zhbPA=aev`BJrejal&5~44tyB!3gFm}3xwP4y-IibHUpwzmy9w)ZyR$AZo; zfTNw=z;T>jcsT7$$+ojmxY~IeZ%}Fx&VIx3k_L`;{s=hwf3KKo{{bBPE6kI~l!nqDj?*c?vHVQow*6JY)qd^Yg`k7|+XNi#UmT!+V}Slm zpo99~2afuO^JFw7ZZEcLl5o3S<-+ZDodr5r?;PM*?qa2%&U0X~e{ zpnm=(aBSCagxl@YoH0^&bU}?fSiNyWZD@+wFQAbg@2=7S12M1gbi>3RgRo-wrz1-g|(fpSysUf&NRthXLOT9FGHI zdGeYP`xE1N5^!wS`M@!r6Tq?jb;7xN^usN}Z9gmn9rVLpz|jx)1?WE>px+HTsQ)~0 z)E_sFy^|ex8ds&l*>4zE=L5&MY63o;_R)S%0muGY3cLWy|4g{;&qsyZ{(KU2u)RIN zvAwSYKMCsH{V4iEiQ9{I9t0fY>Kx!0S95_Eke2%4OTaO%z6Tu3-y_`iLzi&d4^My& z`r$d?=!X{r^xq87e-Cs}e`pbVE%$hJu5h(ed$+DV}Csc z9R0aoxb4q3gsYvZ{|@Nv3H9b3&8^LC?+DI+0FHJ>fn!`P299yH1UR!JgM+xiC!=W&PnV}PUn6_$?1RV$Ro zxVjxU#??yTQL0<}>nY&aUmJm=Ki?B>`*W{i_Cof!s`~p1x7!;Aj_sWQ{IjH|I;R0g zJ1c=>Tzw08A?Pdvj_vw2@V%h?df-@oi*VZyyYV6iCA;4Xgxh}D7dZN1q;OlmG(dkU z=wMu(1RV9B4$%LLrO)p(;Pe{kp#HnSQ9plzyPx_&0LSuQ6>j_C4&k;R?gAb3!vnz44-Wbai5FKWsh#_}U!7^f?YNo;9OJ4P_?e`kadka#?5}0O(VzDVxBa|N$vM#z_Gtp0!M$Y z6>j_UW#P6zH-Zkf_jTad-d&Dy+so~$06RwjM>~swV_eMxj&W5F9NTpxaEz;;0>|^8zXu)k!z;ki51Rw@haBtr$JWmm&VEDvy@8{CwWXuSt6C`kN$~%Tzz+nz z4ETA#R|3cWdI~uD^JU?-Ki?B>xAz0k!S?Qc92r50{ebNq0sMT*)A*kN9PNw($GBPy z9OG&UaBSD@z%j1Y0LSvr3b+06s&Ly6uY(TyVHJ?rvt}0zX&*%|GaR!yb1km#|NMf*^?ja8Ew1m|Y`3_+KeN~3`aVp0Rhvm_kG|hBY;k>GWtqkG z{gZPouJ4olg2naykV`GD?|W>wxW2!!*W&s=f!eR{T?{h-uJ2JSv$(z&ajwPn zJ%?YgxW2b=sm1j@gzXmB_X_q}T;CI@uk7V){?qsVh2?mx{B$|*mRVfivp3h``kp8$BYhy*p_@bm!R2@skoT7P&gj3Yi;C`X^yz@Keog#XndAx|; z8Ge*`J)V4DP~TkC*woM%Z%QtveOM{JuyTGA{YL`q_79!Hq`+OitEE&X-(Qk(r`Zx7 z`qexG#bdQBL)?vSJ~>kAmHj_K+q`c+LrVL*AI@j!lexH6sN)=ba%R)BVb?#7a#?d{ z_b;Y>S(E!8Q&{SM)e=_wwJhHT@&4l#$mbqB$xeyw!=(1>=XA={Z7erc_D9Csc3Hl6 z-;Mrk`^#kiFy&#g?f)EQvZlTNT-kro0PW`}^y~lSKX7aPvIukk?@oVOzxKah-MnWH zpG+z3&pT1#a|)lI?EN|0GLeYmY~h<^O6zq8_iK*-_*%zr|0%b*J0xX4{jmN2|Dj%^ AyZ`_I diff --git a/examples/train.ipynb b/examples/train.ipynb index 3a37c54..9c92203 100644 --- a/examples/train.ipynb +++ b/examples/train.ipynb @@ -26,23 +26,28 @@ }, { "cell_type": "code", - "execution_count": 2, + "execution_count": 7, "metadata": {}, "outputs": [ { "name": "stdout", "output_type": "stream", "text": [ - "Epoch [1/10], Loss: 1.7035\n", - "Epoch [2/10], Loss: 0.7193\n", - "Epoch [3/10], Loss: 0.3068\n", - "Epoch [4/10], Loss: 0.1742\n", - "Epoch [5/10], Loss: 0.1342\n", - "Epoch [6/10], Loss: 0.1232\n", - "Epoch [7/10], Loss: 0.1220\n", - "Epoch [8/10], Loss: 0.1241\n", - "Epoch [9/10], Loss: 0.1270\n", - "Epoch [10/10], Loss: 0.1297\n" + "Invalid axis" + ] + }, + { + "ename": "ValueError", + "evalue": "Matrix multiplication requires 2D tensors", + "output_type": "error", + "traceback": [ + "\u001b[0;31m---------------------------------------------------------------------------\u001b[0m", + "\u001b[0;31mValueError\u001b[0m Traceback (most recent call last)", + "Cell \u001b[0;32mIn[7], line 56\u001b[0m\n\u001b[1;32m 53\u001b[0m loss \u001b[38;5;241m=\u001b[39m criterion(outputs, target)\n\u001b[1;32m 55\u001b[0m optimizer\u001b[38;5;241m.\u001b[39mzero_grad()\n\u001b[0;32m---> 56\u001b[0m \u001b[43mloss\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mbackward\u001b[49m\u001b[43m(\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m 57\u001b[0m optimizer\u001b[38;5;241m.\u001b[39mstep()\n\u001b[1;32m 59\u001b[0m \u001b[38;5;28mprint\u001b[39m(\u001b[38;5;124mf\u001b[39m\u001b[38;5;124m'\u001b[39m\u001b[38;5;124mEpoch [\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mepoch\u001b[38;5;250m \u001b[39m\u001b[38;5;241m+\u001b[39m\u001b[38;5;250m \u001b[39m\u001b[38;5;241m1\u001b[39m\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m/\u001b[39m\u001b[38;5;132;01m{\u001b[39;00mepochs\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m], Loss: \u001b[39m\u001b[38;5;132;01m{\u001b[39;00mloss[\u001b[38;5;241m0\u001b[39m]\u001b[38;5;132;01m:\u001b[39;00m\u001b[38;5;124m.4f\u001b[39m\u001b[38;5;132;01m}\u001b[39;00m\u001b[38;5;124m'\u001b[39m)\n", + "File \u001b[0;32m~/Documentos/recreate_pytorch/PyNorch/norch/tensor.py:167\u001b[0m, in \u001b[0;36mTensor.backward\u001b[0;34m(self, gradient)\u001b[0m\n\u001b[1;32m 165\u001b[0m \u001b[38;5;66;03m# Propagate gradients to inputs if not a leaf tensor\u001b[39;00m\n\u001b[1;32m 166\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m tensor\u001b[38;5;241m.\u001b[39mgrad_fn \u001b[38;5;129;01mis\u001b[39;00m \u001b[38;5;129;01mnot\u001b[39;00m \u001b[38;5;28;01mNone\u001b[39;00m:\n\u001b[0;32m--> 167\u001b[0m grads \u001b[38;5;241m=\u001b[39m \u001b[43mtensor\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mgrad_fn\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mbackward\u001b[49m\u001b[43m(\u001b[49m\u001b[43mgrad\u001b[49m\u001b[43m)\u001b[49m\n\u001b[1;32m 168\u001b[0m \u001b[38;5;28;01mfor\u001b[39;00m tensor, grad \u001b[38;5;129;01min\u001b[39;00m \u001b[38;5;28mzip\u001b[39m(tensor\u001b[38;5;241m.\u001b[39mgrad_fn\u001b[38;5;241m.\u001b[39minput, grads):\n\u001b[1;32m 169\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28misinstance\u001b[39m(tensor, Tensor) \u001b[38;5;129;01mand\u001b[39;00m tensor \u001b[38;5;129;01mnot\u001b[39;00m \u001b[38;5;129;01min\u001b[39;00m visited:\n", + "File \u001b[0;32m~/Documentos/recreate_pytorch/PyNorch/norch/autograd/functions.py:90\u001b[0m, in \u001b[0;36mMatmulBackward.backward\u001b[0;34m(self, gradient)\u001b[0m\n\u001b[1;32m 88\u001b[0m \u001b[38;5;28;01mreturn\u001b[39;00m [aux_sum, x\u001b[38;5;241m.\u001b[39mtranspose(\u001b[38;5;241m-\u001b[39m\u001b[38;5;241m1\u001b[39m,\u001b[38;5;241m-\u001b[39m\u001b[38;5;241m2\u001b[39m) \u001b[38;5;241m@\u001b[39m gradient]\n\u001b[1;32m 89\u001b[0m \u001b[38;5;28;01melse\u001b[39;00m:\n\u001b[0;32m---> 90\u001b[0m \u001b[38;5;28;01mreturn\u001b[39;00m [\u001b[43mgradient\u001b[49m\u001b[43m \u001b[49m\u001b[38;5;241;43m@\u001b[39;49m\u001b[43m \u001b[49m\u001b[43my\u001b[49m\u001b[38;5;241;43m.\u001b[39;49m\u001b[43mtranspose\u001b[49m\u001b[43m(\u001b[49m\u001b[38;5;241;43m-\u001b[39;49m\u001b[38;5;241;43m1\u001b[39;49m\u001b[43m,\u001b[49m\u001b[38;5;241;43m-\u001b[39;49m\u001b[38;5;241;43m2\u001b[39;49m\u001b[43m)\u001b[49m, x\u001b[38;5;241m.\u001b[39mtranspose(\u001b[38;5;241m-\u001b[39m\u001b[38;5;241m1\u001b[39m,\u001b[38;5;241m-\u001b[39m\u001b[38;5;241m2\u001b[39m) \u001b[38;5;241m@\u001b[39m gradient]\n", + "File \u001b[0;32m~/Documentos/recreate_pytorch/PyNorch/norch/tensor.py:487\u001b[0m, in \u001b[0;36mTensor.__matmul__\u001b[0;34m(self, other)\u001b[0m\n\u001b[1;32m 484\u001b[0m \u001b[38;5;28;01melse\u001b[39;00m:\n\u001b[1;32m 485\u001b[0m \u001b[38;5;66;03m#2D matmul\u001b[39;00m\n\u001b[1;32m 486\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mndim \u001b[38;5;241m!=\u001b[39m \u001b[38;5;241m2\u001b[39m \u001b[38;5;129;01mor\u001b[39;00m other\u001b[38;5;241m.\u001b[39mndim \u001b[38;5;241m!=\u001b[39m \u001b[38;5;241m2\u001b[39m:\n\u001b[0;32m--> 487\u001b[0m \u001b[38;5;28;01mraise\u001b[39;00m \u001b[38;5;167;01mValueError\u001b[39;00m(\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mMatrix multiplication requires 2D tensors\u001b[39m\u001b[38;5;124m\"\u001b[39m)\n\u001b[1;32m 489\u001b[0m \u001b[38;5;28;01mif\u001b[39;00m \u001b[38;5;28mself\u001b[39m\u001b[38;5;241m.\u001b[39mshape[\u001b[38;5;241m1\u001b[39m] \u001b[38;5;241m!=\u001b[39m other\u001b[38;5;241m.\u001b[39mshape[\u001b[38;5;241m0\u001b[39m]:\n\u001b[1;32m 490\u001b[0m \u001b[38;5;28;01mraise\u001b[39;00m \u001b[38;5;167;01mValueError\u001b[39;00m(\u001b[38;5;124m\"\u001b[39m\u001b[38;5;124mIncompatible shapes for matrix multiplication\u001b[39m\u001b[38;5;124m\"\u001b[39m)\n", + "\u001b[0;31mValueError\u001b[0m: Matrix multiplication requires 2D tensors" ] } ], @@ -87,11 +92,13 @@ "for x in x_values:\n", " y_true.append(math.pow(math.sin(x), 2))\n", "\n", + "batch_size = 5\n", + "\n", "\n", "for epoch in range(epochs):\n", " for x, target in zip(x_values, y_true):\n", - " x = norch.Tensor([[x]]).T\n", - " target = norch.Tensor([[target]]).T\n", + " x = norch.Tensor([[x] for _ in range(batch_size)]).T\n", + " target = norch.Tensor([[target] for _ in range(batch_size)]).T\n", "\n", " x = x.to(device)\n", " target = target.to(device)\n", diff --git a/norch/__pycache__/tensor.cpython-38.pyc b/norch/__pycache__/tensor.cpython-38.pyc index 72d3bebe7c6e0b2310c3fc5635d02abacb3002da..92de5802f744c677d973b17b28e9a0a9761ffd8d 100644 GIT binary patch delta 1207 zcmaKr-)|IE6vywmJ3G7SbbqkhW?LKD4c!uUyF%CP?-rBtV@=}^G@>F>7na%9PD{5P z2WWM${0xc`Nlx@Z(J0{sVhGI;HN@xx@Gls%G5X|_ycrUG()*o-NHB4e`Ruvpe9w=$ zckbHq)s^T=kw}wJzgxGTn%-Jli~ej--TU2W3s9Z4SRak^Kk@yL8VuDnbM{c7&VQv2(@w6neWO-hd;ma7D9wU|Zfq$Hy}BI_xrV0Psig(Sjv9 zj1#BZHgR|7M!nI5cXeK(S^uSiT}S?}^#WQOHB?Fi^M2Bw>5?@5owt<>gX?*^dk>ZP zhwj~m1!WO<0&*^7&PNWbWmeuF-1q*HdBJ!Zq9cI6WBhE-^Y=G^84D))G{Uqdy!p>ntp;pk9yl}wtG zw(ClJ-W8VNYGuQeGX%Rus8#dQPySk2rbKjMpVD4c@n%!FbVjxd#8UzumrR)wrhr#@ zRL;VNhn4@;tU}~OT+M1eMJTT%QS~47o5kSYA*K??g07|sE9M4AgcRl2(p{yCE1i0* z_yw}l&N^3(<4QAdKA`*wS9?)pwdZXAsHEDMEY(AE)#{@9ma|Daiv5eJ1DY%v;l297 za-~+A=iSAPdKi%qA1~(C>HfLbwFAquZ`x;+UmfSB!e(U%dW_;SA$8Jam^2DgJP@X^ z)F>x6P_y#s)T9*3M{zIM-%IqFR~UGP=sFh$dc7-y*NE=$p5fnUnLjLa@zh9_e&t&u z`#ZOz{{-OQ1|R;F*r9xX@Wg0GNcwM8;>zfdy4atL4z3PCKMafm+W_e8Nx;9`Q?TZM zL%<@i1mF|5j{+;eTfi}36*vx@1Wo~`fwRE7zNvl7Z1#`Eo#f9JRl=KAyG9k(TMAU>V#_T$%!3I#d@$uh4((BT zhJUlKq$X9NcM`CHr}@LX9pgEKje^ef2htcrKEYqK4L4PEXMS<8!hf}`&{00qe%&~Y zi&J5~<0DERMO+1jfF@vq-_7iyX}*@px6VMwC~WIqAL3egr1O1gokILP@FH)O?5O7$ zai;w~?%BKFXtd=Md#{iaZ6nt8>i;b^kHst({W8E@RP<-&VQKw4VVd*i9-i%amL~YC zo|9G$&V#@K*m-U-buEyAHl-AY7Nl2k4<@nimBw;OUl=xeCqGEc;EP`K&`NWC3gkd zq{wQId(nfyrnA(_-J@On^}$Xqm@y-+G;cRk$k?B~#KML&iz{eIoE%Q5rUHtu_8vVfK(rM!e);f$AJT|VvoBVmH)1YV<{8MQ} zeV2B5czqbI5ugk_1i|1V5k+O3{sonii{6TP4-TB8q|`6HAhQ^i4!56RNao z7gebb{D5GyP%OAmsvAWYDuPQFf@C3xD-l6lDya93ZLkCP&g8k3?G}M>SJLgu z6Id?IxJOpiCCFKLLe{_LXXbhIYz={c4jPP^SVz>usd4#whNhs|(+E?fOe~0!GSn!Z zsYOwSEmZ}DU}d$e$-44f`K{-V?jq5{7-4V9w#eEnHN7fubexGW~UGai9?Qsj^dJP z)N$vlD4aQThMr~Wvtgtxhqygb_9Dca0hLnrLOXxX3Kk#Log?PM+V^mVn{gAHNkzzz5gFy z;rp8WAJOCRNaSKS;Ia4`>kej-_iza}+vHBh55qI@F+L41L?w|4_95@-5^i6U`zvu6 zUWtKZ84loU!th3XNM?fl$bPvbqhGY8GBGIvm(SvIRxXF&vO4b^T4c delta 754 zcmY*V&ubG=5Z>AMHrb@jTDnQL$*w6>EDbFM4}yxpHf^KD5R_7+^#`ObElpbsFBPS1 zP-%NmdQb=ah1jHmPz6z9!Gnd~#EXAGPJ$;7J?KF^Sa4po!R|7%voqg(GxL2t@hN2t zheAPxp6$hx<98!(jSY^!_4AQ6o)bXnHK_6Y^~vHl`GYxa$Zx3g9lfrLl|Xy-vv!-Q zB32y_9A@BDpEhm*#H)9Vd;Fl!Rb93Sqs-;5wy2D97lxFizt%8xbtav^JasKE+#Kdv z9LIufrm%@V$%LOiLYL0Qvgw34g$n2}AEq|EWhZWElwf3iWW{N6o>EX;h$#gqsT8ZO zI;DyubIgMRFK7jytS!JUs_d^QVWq^bbkFIYx&S5aGLOr?FxG-W6L>-fuA zhWEI%KV=S5<=q~J!TQ8VVgxpDHc^1nxMf?gi9Zr4GfU_Xdpfe1ZcWAI^+FyX{{ik%uqprm diff --git a/norch/autograd/functions.py b/norch/autograd/functions.py index 30e452f..fe31be9 100644 --- a/norch/autograd/functions.py +++ b/norch/autograd/functions.py @@ -26,7 +26,7 @@ class AddBroadcastedBackward: # Sum along axes where the target shape dimension is 1 for i in range(len(shape)): if shape[i] == 1: - gradient = gradient.sum(axis=i) + gradient = gradient.sum(axis=i, keepdim=True) return gradient @@ -126,9 +126,10 @@ class LogBackward: return [grad_input] class SumBackward: - def __init__(self, x, axis=None): + def __init__(self, x, axis=None, keepdim=False): self.input = [x] self.axis = axis + self.keepdim = keepdim def backward(self, gradient): input_shape = self.input[0].shape @@ -136,13 +137,18 @@ class SumBackward: # If axis is None, sum reduces the tensor to a scalar. grad_output = float(gradient.tensor.contents.data[0]) * self.input[0].ones_like() else: + if not self.keepdim: + # Remove dimensions of size 1 from the gradient tensor. + input_shape = [s for i, s in enumerate(input_shape) if i != self.axis] + # Broadcast the gradient to the input shape along the specified axis. grad_output_shape = list(input_shape) - grad_output_shape[self.axis] = 1 + grad_output_shape.insert(self.axis, 1) grad_output = gradient.reshape(grad_output_shape) grad_output = grad_output + self.input[0].zeros_like() return [grad_output] + class ReshapeBackward: def __init__(self, x): self.input = [x] diff --git a/norch/csrc/cpu.cpp b/norch/csrc/cpu.cpp index 3b4e303..ffeadbc 100644 --- a/norch/csrc/cpu.cpp +++ b/norch/csrc/cpu.cpp @@ -205,7 +205,7 @@ void log_tensor_cpu(Tensor* tensor, float* result_data) { } } -void sum_tensor_cpu(Tensor* tensor, float* result_data, int axis) { +void sum_tensor_cpu(Tensor* tensor, float* result_data, int size, int* result_shape, int axis, bool keepdim) { if (axis == -1) { // Sum over all elements float sum = 0.0; @@ -219,35 +219,22 @@ void sum_tensor_cpu(Tensor* tensor, float* result_data, int axis) { return; } - int* result_shape = (int*)malloc((tensor->ndim - 1) * sizeof(int)); - if (result_shape == NULL) { - fprintf(stderr, "Memory allocation failed\n"); - exit(1); - } - int result_size = 1; int axis_stride = tensor->strides[axis]; - int idx = 0; - for (int i = 0; i < tensor->ndim; i++) { - if (i != axis) { - result_shape[idx++] = tensor->shape[i]; - result_size *= tensor->shape[i]; - } - } - - memset(result_data, 0, result_size * sizeof(float)); - for (int i = 0; i < tensor->shape[axis]; i++) { - for (int j = 0; j < result_size; j++) { + for (int j = 0; j < size; j++) { int index = 0; int remainder = j; for (int k = tensor->ndim - 2; k >= 0; k--) { - index += (remainder % result_shape[k]) * tensor->strides[k < axis ? k : k + 1]; + index += (remainder % result_shape[k]) * tensor->strides[k < axis ? k : k + 1]; remainder /= result_shape[k]; } result_data[j] += tensor->data[index + i * axis_stride]; } } + for (int j = 0; j < size; j++) { + printf("%f", result_data[j]); + } } } diff --git a/norch/csrc/cpu.h b/norch/csrc/cpu.h index be875ed..5fe358c 100644 --- a/norch/csrc/cpu.h +++ b/norch/csrc/cpu.h @@ -5,7 +5,7 @@ void add_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data); void add_broadcasted_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data, int* broadcasted_shape, int broadcasted_size); -void sum_tensor_cpu(Tensor* tensor, float* result_data, int axis); +void sum_tensor_cpu(Tensor* tensor, float* result_data, int size, int* shape, int axis, bool keepdim); void sub_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data); void sub_broadcasted_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data, int* broadcasted_shape, int broadcasted_size); void elementwise_mul_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_data); diff --git a/norch/csrc/tensor.cpp b/norch/csrc/tensor.cpp index 824737d..c945689 100644 --- a/norch/csrc/tensor.cpp +++ b/norch/csrc/tensor.cpp @@ -172,7 +172,7 @@ extern "C" { } } - Tensor* sum_tensor(Tensor* tensor, int axis) { + Tensor* sum_tensor(Tensor* tensor, int axis, bool keepdim) { char* device = (char*)malloc(strlen(tensor->device) + 1); if (device != NULL) { @@ -192,29 +192,45 @@ extern "C" { shape[0] = 1; ndim = 1; } else { - shape = (int*) malloc((tensor->ndim - 1) * sizeof(int)); - for (int i = 0, j = 0; i < tensor->ndim; ++i) { - if (i != axis) { - shape[j++] = tensor->shape[i]; + + if (keepdim) { + shape = (int*) malloc((tensor->ndim) * sizeof(int)); + for (int i = 0; i < tensor->ndim; i++) { + shape[i] = tensor->shape[i]; } + shape[axis] = 1; + ndim = tensor->ndim; + + } else { + shape = (int*) malloc((tensor->ndim - 1) * sizeof(int)); + for (int i = 0, j = 0; i < tensor->ndim; ++i) { + if (i != axis) { + shape[j++] = tensor->shape[i]; + } + } + ndim = tensor->ndim - 1; } - ndim = tensor->ndim - 1; + } + + int size = 1; + for (int i = 0; i < ndim; i++) { + size *= shape[i]; } if (strcmp(tensor->device, "cuda") == 0) { float* result_data; - cudaMalloc((void**)&result_data, tensor->size * sizeof(float)); + cudaMalloc((void**)&result_data, size * sizeof(float)); sum_tensor_cuda(tensor, result_data); return create_tensor(result_data, shape, ndim, device); } else { - float* result_data = (float*)malloc(1 * sizeof(float)); + float* result_data = (float*)malloc(size * sizeof(float)); if (result_data == NULL) { fprintf(stderr, "Memory allocation failed\n"); exit(1); } - sum_tensor_cpu(tensor, result_data, axis); + sum_tensor_cpu(tensor, result_data, size, shape, axis, keepdim); return create_tensor(result_data, shape, ndim, device); } } diff --git a/norch/csrc/tensor.h b/norch/csrc/tensor.h index 5ff62d0..a32e74c 100644 --- a/norch/csrc/tensor.h +++ b/norch/csrc/tensor.h @@ -16,7 +16,7 @@ extern "C" { Tensor* create_tensor(float* data, int* shape, int ndim, char* device); float get_item(Tensor* tensor, int* indices); Tensor* add_tensor(Tensor* tensor1, Tensor* tensor2); - Tensor* sum_tensor(Tensor* tensor, int axis); + Tensor* sum_tensor(Tensor* tensor, int axis, bool keepdims); Tensor* sub_tensor(Tensor* tensor1, Tensor* tensor2); Tensor* elementwise_mul_tensor(Tensor* tensor1, Tensor* tensor2); Tensor* scalar_mul_tensor(Tensor* tensor, float scalar); diff --git a/norch/libtensor.so b/norch/libtensor.so index a125f3cf42d7df3c5d394600b7016cafea233f0f..09af449b259d5e11ee27f21d0f339c1a2e53fcef 100755 GIT binary patch delta 24561 zcmZ`>2V4|K+rM2!!GILS^0cp*)`E<8u$I5nc4GL@~yw)xp{h<*>=veBW1FWl*tV64eMayN3;N@ z?5vO`9w|-4C?#uXs7J#8#2_JS^17lp*5p&uG~J{|mt_4)1jf_?b1**o z5u$$qtcy>W9ESro#HW!QH~_+_&1YjqREtOnvB6TiAV-ak>67RXUOj} zC7dPUY`}E+oj`{CK3~EGfLZt~#E0Rt2p`vDF$SOGv&2+dpqDbmt)u(+SHHK8>K6P3 z>+05{qsjO-6E#dyV0+_m$5%?o#1R_G>)6qo4|7x2+R4OUpI)qR7taV}> zEBPa(p33({D+HpRT)NK69?Q;?dpzlX+@R zQwK*KI6_BB=}?KaV*I|u|1J&tuNQclO?2#*8LIwBjP(M2Q&u<>eNNyf%K*xeqsd5x zLQ?|AQAv2_Ed=&UPje)GqU?ajX9a&30ZeW(vTBzFLhn4Qbg08Uf!7@xCH40NKh18_ z6lvfBzEf>{Hwb~#0+=qk3;g+S1mcdwH4f5(WN>jCi2rKOQgewIzNSAlYr0yju|T zuC9rb5)r(U1YU2(c@qEKj{+Yc4M$3RdQMS2!)5Jn*#%y&{bOnPi{%3EJ4`^+CfQtj zSd6zhxj=S6zNBuEm3~KdfSl|m2lX)A9e7GG>MfZsSA|{s1-^rHw}YGuf6K-2H;GT9 zB^G)4~ ze=LB>L35`a1pIahjJ_~aW=PR5vZ)tKfo@V!(|H21Sk48jG<=K&`I?jSHwvnLmg!7p zpe;ue_^DFcG?|ZSa`|r~@z-R5cF98Pb+{nIu9_vrG)$0wU~flS$ejb#Js zdG(QDRG%vp=`*f}8b|>9V`v}{*xfom96)$#Gjfk z@Qq}FQe^`y*(dON9e$8&N!hCcuP^7NWWLAJl^YKI{FO@q&@$h3k09tR87w_LbV%U! z)cHxngN6yb9&v4UohuV-|$?8s6F1wKIrR$rFv@G-f}>z+2pq^%Bz3SS5T{S4Xm z|G}rq8aG82>Cj8FUOGBVP7%F_>dPKFcSZ1HP6^3ca?x?E}a?qq@2Q=Ka#j!PBz z3^+o^a>>7U6(5?eNdx8N>e5cGmwKkJml1Y{CCd{JQo)Z@S^KMp#aNGEnvW28R3PyB zG8HZ}adE4_>&+xlh6+C zrfA^kAiC;8Kf7r-d@XH?b<$OnucU{8e|FaAaIKW7i^PqqI23JK($v?EHWX=>Kg=t)V=uk`yH`A~p3RS1VfRUou=rMj2bPzDmfbS${Pq4&? z=mM*xK=fxt73t$1tW#aHrJsSCDaWH(YTahaLbkuIS*gX&*A2D2Egd1yBT?u?N=%Wg zO1)54xt=HMQO{GE%?8#BwRCggo#VnghkakqYzc9pE^(puXV>e6Dm|E2eX|9rbRN50 zc)nm!pnl|}nM`>u)I>HC)G992OD@#vY(;&uSEvis&a*`@k5!~3~lV`UC#R#{x| z0i#9Km^7w%R=ZUpEM~%y8lF;le+JJ#JvOOmU$OaiNB@DxslD zI%^*q%9@9HvN@rimP{AkB`&-J*!SR#VAn&$VM7ixEw5pMPx=t7NUouDNz+|yFu1>3_#W+{6dZdS@Lw}zpXUM@Tn zU3gMi`-Wx<)(qz{$A#LNB{ejAjds#-v(LsfG_zd|J=wK}p31lECVa4Ch@= zWi^cb3rfAlo?Zbi*mA6XoX#FMHd|~i z)FCd^-mFSQD4W+9^ESfMGRB2#jtkdNmK0&O1h`O_xKJyy?;~1Sdb&_|xlji&FL=1G zv8R&9Vj|F$k)D=d7v4)Qyj9sq@Gg(^RP4+byx&IFRD#&aNV8Xhi-bE1YGP*YO+39m zcEJX)PQbp?#MAPr3pQBJV#@3^)~AWt($t0ejtjNsw(yrW_pQ-0N>%>NQS0kq)#P36 zLzC%u@2Y=cyh0zd%$E1q;8tO5LMtD90_pp<(y^9O&0;g6EVWcqY=)a3fu^s=`p4Sb zW(DE<%CRtrDIcdd?3r|;;(FCqdGQI7~9aX zqk6U)p;XV!%rEX!HI|Q*RJJ_siL#N|IxSEu|G+ADwy33kU`;yLz>}pI&vx)L{F8UC*SdI+ z_I7-Ir8VmsAFRez19x^M5Zw%z)wH7uN^H_PevZ9xA3K}^rt zQy2L@Kr1+sJ%goopN}c%)1$Up>lv$3d@A=%}NK4x8@SZsv zh9A~P*exrNn)=D>kdrT9Q{kfUb@|ZyjI3JIb7In8 zlLduZ!o{c!YJu+coO)3G`UxBJQJWee<@9MVTZF`K)9s(afH0N!goTf=s8yb_b|Wk; ze)iMZ=X}L8AlNH8*~dKL(dNb4qQl6VFC|1y@tIF(imRVJX3sydREQ(6vff~0n179ubTp= zA5JB_ExO*DvZYg$MankYhaw5iUIiDgLP(eo677LP6;b}+A&Z&jua-cn2bNPiJY;>w z*I~c)Ew7e+#1bc1Y8~>{vpxy>;Eu?;C$tH@pFBiZU$UBP*vDnd_r~B&VrePp-G}!6Q_m74_H?BX=VC;DN^-rIfCtwp}--3psb?_*|`}Td9u=jCh z=PSyFE#BC-Q*ZOWoq3P#j;X>WoYh!~T&61N?k==&7!3+-iy z$XB+wTZpkdi%M!<=37u`6&iGlZ5~p+@|qizn)X4Z@@jLN$BDi_5wQR%5njAqz-v4WSc<)Q z0k84IYrMv%Zt@z>zR7Ex`HL~I`6YObb8+8lzlCsw%HUskji+7XH7La}WSWzOO0ln`s+ZwYiqbdZC z#A68i9A|P)Q8J1zDBZ8{!fd|I3-k9CRw;3Oau<%Fx`<1eP*(G^UJVEREfSsf$i17qM$L$5jmVw7xt27h9rh9W|G2<5 ze`ZnJUZg9a7A7g35wQc$A)F$uAtKmUULbeq5=2~gwm3)1#I5DF3oLg+V1w021@0#3 z$uTXO&bymQBx)$zxW9PWe!$xs`%z>>l+BwPaf{_d-&7(lCn7J~t@CWi#2W10*nZ_R zDCtE?XZw@%H8_mhEBh}_r)%-zNT)L*vO%TNO#>RUFoIS4G(bIYhF9tyOdBfofv3w(WsS8CJ+ep$8qEGpHlM)+%eI!tiQ^JJ(HQRbgz z7skhj?0!{rP9`81&Z_)#R_E?~mKX7xvsA?QdJ7@UHNQD7>|Zr0;5YQ+AolRbs_bbm7(6}Z{d?MZZZ1uZ$TUvL=VzXwr}>%ZJ<{ZCRIy^c z`+}$Kvl*~CAITDFJ9nCFx_NP;Z%ZO35Rte3H>dfT+JWnQihuuUey0BOI6qTwI>FD> z?q?WXjro$EB2DZ4JK6J%g|NLpG9Xm0KgG|~9Z&KT!kl~_)7g{!YV7*)H}P~jMOR~$ zTbv~K{_|?AJjV@& z*I%cPlC}T3{)#`!ufJ{}s#?SukTaJl%c*4pYlx9Oo7Q3@Ka=BCUM%_4sfFHE4+QO z`w@|6qQ`Mw-_5{MeMcQi?7HaM-LH zjo07jDIbQY&phT)mzf6BwUA!981O>6t>r}D>q#KKMk0h??;T}BQs=jLf*-}ubs_GL z>9M)~EHIJs_r8;P4qS0r3Ti+VrRruLThXh!f0dtjh`3tL>xs#-XYtM@DA|=mcl-Egm*}( zLnKn9L$-+yshme0Qoa`HED~mo2RzK)lc9bBk|5OY7x|+w^%i;e5D)V!v=p9ZJvq#0 z!oNR4pNwoCW;*UJcU&)>fv)3UWX6%nV%qpb85T@lG53@7go!dtu4I<2Klq#FsqhE>WTZ@R% zMC4)4`jLluW4s7+`HwuzO9x1#NSHIT`GkLRh?nzK4!73)gVfp#BMK;)y_eioNex3LAfA9k0C(bvPq z$ic(+4akFpq+6dzwke}RaBZsIT(}@Y&gXT~pMKl1IPeGDM-NbLmGVrPA6d2v#UWFh zwU1fn2Zkr_rD*vpj&wZ7#fsv*73fMD4VRH|c*uctr{g`ZeJv6vY(L;i+-h*5Z*L+_ za1!_LW%uXrD6=V1q`7ZbMlJpJ9kA)xPvphJt? zr<|lvoV`F{qy}}hnhmSS zn!B)vcYJreH(CEl&#_q_Df;{sIsb1-UIJ5^)vnw(Iq-7TdPljRUr6_ouPGzVZPwRF zNm3)g|FVlp_R2<4?({)XkVupL>}t~TT6y@>QI_>>4Hmu5qAc0gYu#g2ImLFbuf_(i z@1#CniRKEyub=L$k5QkkV6`{Y4!gIC8k5h?OR?!bc-OzUCoK~^`F#DF`6W2ylJ(~* zHf=+o`QA!Q3&c*-v!p&f)m7~I>2fS@gP*dDUENUcy#`!0GH2tf{6Xgs5AM08pL=ji zuSpZl*ZnKrSPEy>?}CD=!4RlWL61?6Amj|lh(7GU98R#reLa+xtl+xtQG&a~~y#^Dp@s~7J%&*2R+hT4iiZY9 zyC24E<&8l0WLw7yA5$Li^3`TtE4v8#+OT@x2dFVW^0|Mt8!E9Ff?`p*x?Yu#FgHE@P*5gsB~tEZk|vPoSP}s;$;v z!D4nsw)fX4C5L_NBbHJu+wmHn$Fd!^n<4=smd#!CSO$V$oDud?1|o@tue-VScb6?Z zTDkO8VC=7!VkND|9&HX(LYUWTgJ2z>u?%e4mkZQxGZmm*o zIT-0DDbxU&pE@fr**N6&Xpm+Ven6V%vM;KcYryFYz zA4MjRx9r1Ro&hZRUHLMQPu6Ai>f}z zqN=KQ7qPmVDk|&Q$Q?n-H*CR{xBlXaHZehcGZea~$SF&qH zEuRO5R}eLH!C~kS7MiuuEj-K6EwqG7_MC-ep&W%aKa=&%`AA)z!7a?m;1*IAvVU@F zMr|Mq=Mjj(!W!rh7A&oG3#Y&}DJwk5l^GLl`+t`#+5CPaQj zPzFN_ph_4j^?`0^FB}u`_)_`i*(g&t$YxL0C#SbMVivD-mswP4r7inAr&g33>0bp8 z4f;nygV28`M%TX-yh8s|%JQ(e)JprvRFgS;4Xu`2N9{P18;P9B+o{1^T0Nsqrh#b= z+%s5-M+SwJ6VbYr)ET;!i*U={Ii0MmrG}|C+i9iA3~r_B3?51a7Iru=%56G$M#D#g zktjGTjBIbA8yO0RL?{R0mHpj0WMm#iReF}wNVPO>q*NLYrEoS6B_kC~J>a9kN-a1m ztSo;|x6%n72`k@Fnyb$yD`P3Dr!!!sF59=?o&Eh&f_gNSdN?r$zTn=^f+Y>tq9%S* zKtEZ_rt`s_6v$PB-v+~vpjsI2AEg@(fRn=TD0phW50+vrcR0-%=(Q9c=%ExI=-!zW zX!E(rkn~3c279-lQP}I*T(?&iz6pB+sM5bMM_!fE;S}@6=}s$$lev}elDUm}a__Kfo)jw4+kzrcpXA6yTy%ro87QrVXKvH0EQ_s1i{+I%Ir z8E*7wm3io?Ul6fW!h@LGLk4;yjDQepFy;)*Iy(!*7Du~TRM@|d}tE4#6ZqE**UVz!g9>TkC=rSVSCnacwbLSo3Ld3uQ+;UB}7b%5;w_xI;#_4_`tC zJjQU>Td`MHBNZRk;F_hvlhMVM^ctn7QrW>WuN`&&a3NQc#1gLK|H~(`oL^#eGY|0Qlk?Qn;Y~sb~9^>b6ZLcz+jb`NS*6M8=n|C`RY@ z*L-9UH36Q7 zvjcaBdt95tRh-Wt6(0=ZDq6Gszea}naZ0Y8{@ZQQ^(=T$ajlE`lUC(C_iM>)uC(<$ zQo3mXSL(((*(1YF%@V;reaonH(SYJgO`mXAa|WQ}EzDsb?9q|SyL~&C(pZXZzTaBe z$Nst>so>Q2z*6DiC&iU??*}E7*pvtUN<7PYFv4Tv4DOvJoxDp-An(fDO%s0o>dqQG zt{?V;9By(qy4f|MxP|P#+|AVqJfr_i=jt`$?pFBZWPG5>TOD-{LBkd_*eHlkJ~YQI50WPwNC^ zQ{!j&?tk|1nJqt2_BZ#X6aZ9ow-ej52s+;Ev+H67VCu!euT9TlR7qr2W zmdI%%1Z|+8oth?eMiMRHBSGsbXjV!4w--4R&_)nH5X1qJcvuj_1+k7GHj_k~AO;Ge zMG!Y636q@#(OVGR1#yxjS~$^MN)R38cy_WSv4Ca$R^9xPkQx3@1g*BDeZy&w1nsV% z{V`Qo9K~r@1?`-mMM=2`q6Hiiv_pc{PRjk$gC}6OAZ``JG)X)ph?@m*y&zVU#JPgF zNf5sk#NAVbJDmjaOF`5GFo*Nfg~~ z?uH&rcZ;S5$Z#w1H;wQB%;w*KF#K8AMMplaqN>Uk`J)F|GNOJtj|;Ld*whzsZ|&{^ z3@%xY9?E1FUPM$L)WEwMeglESyBZ!2;%SP_|0Qx5{>S~Q+p4|XQf9kiF6mKg5NAng zF*~rNvNN!x+Abrjp<@o|5w? zcf$rZIZTprEwLH1mBZ09^hd#}ggoG} zSjYq3Yy)|~@$DfGxIPZ@fS-1TJQk-XfWhu~^*^*56aaA-a0p<(UQhsd4e(pQ_dkLH zz>R=*K)>El;9)Y&2Mh+>(+Bc^{Xd30V2cFE1MUOd2>9+Ng^ql{?*Lx_J|0l$C||~8 zYC5pc(Ks1jQt_oLU;*G*z@$Nijzxe8g9{zo0rvx*18gv)(D52@C14=V0=ohG0nQm( z=tu$VF|5#$4LEUlq2mr9ODuGFdg5vNm_kPw;L1;72XN2$LdR6VD-#MG$*b_?pNWNz z!+=#MBSOIFDG(aC%{~BOhNn0PzfP1fY+%$@J^o zLdQFRm27BEK(7o$2-pvB72x6dC<)*TKs(^5OvvFP;48o|z{Lw82WSVJ2>8h&$N_4A z*?@kZLk_Ur=ZMJN7q1;O2m&fg5Gml)r4R(%0k{M(Y#9Ur2LcuVzP$o+e%S5_Fc|QI zm5>8m1UMSddo|<$TLbO_JiG>SfQ5i2EDWElg&g3*WPG9joBhYv5CnXCU7;fraPBt{ z1l$9739!+62x8eC1V{__u?>&|{0ndhVBd|91Dp@I8SuC7AqQCf2lNc!pIZwZ0sbaa zr)|hA;5xu)^oU~D23plnrJuIgs@4e4GohvEg@QaFY3hxSkBay(`ShbQzKq33(e7K- zYUTUcA9ej5VYgwOJ>IYt>O1xjfv$VU3 zRCR^PasJvd=ukHXXb%W(tE81_1)UGHkXF$0LTgQ+P#X%My4TX?f~d9#);18+pW0DE z$AoGRXn0-ohz0XKErh^xtu+C=HWWaOZ>G(Sg=o7dZ3Dp%T4=`zz8j-GAQ=9ER;D%J zI;&QP;O)Iw_7-~M z+NthJ1FdH}wYD-rO9e^ktF0!`Ny{NdQ(x^?J2j+2)Z3Vbm0&FsH3~$nYT}cr1+`ae zDDyOHd&nniBWcu1%OdcB_CG>KYS(GhP*XYpglV-1G|*ZTsHcr25Ugc&fWdCs76O`f zhLDT3mjsw*=?E}eiy|;ZOCXS_&Fl!y$=cU68l@d2@QL=AKsU`d4zxHeoIp3N6M?SU zm^k>dPFqT&z1nUX&D3tvXqTpR!svw-(h1>~OxD_TQs0K)7z~x+S{4x_wB0lsqFtuZ z9!=?t(P=HHvszoNK2vMc8EI~=4aHDdpk)zpy|#r$|7e#nQY*~S{^|^^L$shSASP)s zT_9P(rVZ-?$!IMdL*=2i1w^%ShW2w8@Fr@15^uU@=?daCEv73*e`pCbYN(~t=zVQH zjiza5Xf$7YPNQR*B_5+ITGM!p>S_rzdQVHG(PV8sjpk}cX>>$;PNNH&Z#Q`3w^Vz# z8yd3uQmq%Da#@?y4et1UseM7iIXicEQ$H@1pzqV5@5`WX!=Nt~v`GXV`j|GMa2?a8 z6Ru<0#KLtPqD|@pr@CrO`=}oyVSmzStY+zp#c86}y02Qz*S&!d!AYDSm$XSl4bawu zs$AQ7w68j$B(mDNKZcpw(EjRE%Mh$}#-jx0%r$My0Bm9xpsgOD`j>DwU7{~$#a-Jy zKy9U@Xv#pDQsCrXUPeXj-ND4(1XegVj*Apo7+Hu-a0Yqpb$ka$gEnl(H8A6;l<=zBAxPt1Cm0 zswG;&5HPQj|M43aDEz?F3vS_+mIbmlafs^UVVAf;TFwwC*sMJ#$&_P40u!Ad%~c_} zP74|et|3~dq2StO(7#S{t=CdP)`kp)ek*>Z* z+#D@vBpUXS7Bdp6e%E@AR6{MsZHk8D3X_C8Xsf|xG49c`N|G;YXNYXvrpX=G6eRtt zri@a}6=@GVI&c%jk7P;qTd6f1r8Y<~ZaqW~H%Q#L;Z7*7N(sMrmC$od>gy%(!)R{c zz!PnLOqcklw2XugCzO1BrnsRgNy~yBN3S}ZgrT5OIkAep4KT5u?)@@()tEd zJ0&@Dgf=}HQ&)78~3%ECds?C7?4{VcbPmU z$;K`E&P%*;=cB(Q{+zaYwCbN+)@0n}uqQ5~$y?(-U5OHJ+)!+h#78X=S($>lPsbsN zH*UvuU*h?WdARFx?$t`y?PS=vxzjBvH?r&|@y0FaKIeGuoAVEg zQ%gpCD+P>OoE?;Smu(ce;2(sieM^%Aa0fuH88_N`N8%&%geBU>fR1hwZ`{c%N#gb0 z0*489(;A64Zi1J~@g}YGr>bwlWyxgRVe~JFH*REF5hrMJT3aKmBuQf}B);eteb~^A zZ@!vJe_L;mJjVSA^CiC2L7`0F8|@#7H*T8e#rKF!4w9ucZvWU&;v4Q3>hzh`L*k8F zw$UC_6jIU6=y3NY@y6Zswh6pnvAtZ+3%{OAIpdC@3T}hQj&al0_7ZR0S96}k8@DRm zB=No*M0}@Z333I#g2}i8=wl*K%UK*!)VQHwHSNq;tRBs9nL{2LH^NM$@+EhcRO4Q_ z(1wA?5q~(rNeJz2K$++2QNA3AI)tt~@ zl8qbP&XRcJZmoF|zvHoRB24;Xmw4k2pMPmNFk>le7Oon1)U?v6(X5OVU^WNI%O$gM zJJLN8e~byEQ>E=I5^vlG&o@9AHSX|NS>lab4)zrIz9!mRkPdyryG2r0P7nrRQ zZ`=>`g2exQM)2#~eum~nEM3f>|&Nkzwt`VH9^yZ-bG+h0Xg9UQV4U6^}j0 zm~r3I5Q&%jH#@g)ZYc4_jaPd}ym7bWalrF%6J|&bEYWwaX&g@5(yulz8LzoL42@xGnU*gr_Qa_*NB;wAtBd ziW=^%<(E=QYVP%wMmtZ=Rb$>NF(PrW_FPlVJ6nFP)-J&ktGKara{%94ap#Do>RJ`| zSH8>DP+$GtP={l+PK3S|EL_xKq?QG=M?f8c4$zJ+SDScTlEX`y#|pKH^WIdyN1dp3 zTA?Q6vD%pxkWFbRWK%ScFTt$eSL!fM>-41>>Zfn8ufyu!XbZno8+demDWtpZJo=@Y ztKv=}Yc-Tc%iE-A?dWPa6eovqn#URpQ{*s3YrRHofqT+)0FV6s;@)Yrc7~Xur-)&= Q=KGb}L>Z&C{z|R=e-l?R3jhEB delta 25007 zcmZ`>31AG@|DPE`5(x=nv(6;qK9UsC78TVc)R7?WBJP{GuZV~%A|xKhU9@_jMN!9( zl2Sn(ao@yMOB_+O^vM1{-}m0E&7}Wn*LUCdzUIxFne01s)HC;}XPUonR6B(q5dtW= z=`|X9Y-=b++tP=Gc!cc|gRu1VgFHqT5reKfeAP^~!m>8$pI)ea>d5(7olAc;wC~k> zK@FD?Lo7a3*=wzB#bQc8SRr2}V1}kCN)cP(WKUI%RD6P!nTi%sg!L&D(5xz$Yv5xX zHH|NTwebm)x!!|-J|J|poNg-;wlWAGV=&jftvn266L{!LLPD@s{qDh3G>F#|ABexJ$K6t;NN zU>-j6=?lv!T#9?2hQVS&;IjmurTDlWluR9;6%x6Usc!8HPyF*`E7sMmtNUrMjZ?1v zw(iZ2oo?e)^*+mRx3Gw!k!oGGzNkh0kpKH7|F;4EH=LaXemn6kb`Mo6u`u@#^&=MR z-dkNsNcA&fW4Z4BYA_er#s7WE{~gAnJuK=m{=E`Q^axRJ@$UoqzsI?J4f@V}JR|j} zk(zr##9iU{htmCSo{CaZDXx5eRv?TD&PHz#K5K`-mn$sbhN1#L^eZur6hQev>Un)e zAVMX+l$+pxOU0sU3@#)n)d9&s?n*J9CV+BQ8Ypp0ApBJUg8_*@d9xS?NPGhyfnR@1 z;Ll6^QZIq;|5)IuO%xu{Ynb_L6%s|lHd_7sDB8c1j&4imrdM7 zhH+8G7jZ~121#v^GK`8cPi>^0nI6K5<)lCu`R-aw;A3xae1c*Sw#z(?8zcm3!Vx;2 zO9Rio5o0cZ3XbI`i2k8QToa1rUyeuIEmI;rP3@I>G zR{H39A<$6(#a%}5=88ZV9Tq7as`gOe4Ts{S{`am5{8Xv`5{4xcP@u=(3WDq?<-EIa z==-k){=NW81Bve~`Hf8MGD=3LD~_qMWU8;w6CxG#mxgype2!ht8j~eq3m|pObe3Py zQCWiKa|PZaOY@D)#JFsMj}|})lM%;X76>EamqmpB#j=z=jajRdG~kyc1PlYeNr#SX z6Zj|6(Q2|8D~%HPYck@(G82uy5cqr2(E(CVPK>};ll%h- zjb@xD@w0vw{KiT9p2YXf$ggLl?1}#&)pQs=@dyr(-E#@61YxKE$`7*X*hPUb&W>jo zlEB_g0-q)Mzmq-TD|xw4zd~{kc8qIrBP73J<1U{jiOlh<{Z7!S@gx@884lW;L zxQ#rI4g7f7L|q;Vextb@GUCe51>T4_KpJlUKY^E-PEcM*!^dRAMzbhV(d&Ccz&OoX zNJWiT3;aL9Y1lY0k+^{CQL`QuC61<#$<2{jnez0t;l94wp?j zP>>Zb8CmyafiOl zRv#94$Sq6 zZ)E`YWsP6P3cQg~i!50+_M(&}p>lgcrHP%6+Hx&Y>xmc}c{(huRG1?0MqUrg@aD}B z_*8jLEtmW+KNt8LQhymaxl|GJWrAYlHCxV)&E+M^2s>UDdgl?L$S8E8k0{(P*#d9O zQgvk5+kOyuqw_K(|CGC~{FsF+Nl#-Z2*Q5Z!oy^NjtvudBY$0G#AntEe7pe4ZCRg1 zJG%LWYCBRb-bIUy7%*b&fN?ukRJ3_|Dw( zsJy$<95^~u{O9IEKfiStJT0nx>ZEH*X-N+S|J?Lt@T`c^S>nc*9fFoDti(F`Rpm9R zYuF3A{xTeAVsF(7$GW-uHr{crq;4V zA8IE?kz?~yjd33nKD(&skHIGVayg#FhSX~0&7=(i5$sMa3p-XTSgppM*0QJ#nOp4; z>-#Poqeltb$ci$LwXGeZc46_gE!Jcgo-`Mpn(VvU7PUY71$yKah;wdp;VH$u>sYLE z_T|(=F4Uf^38fnmQKTY3-^LVvTj-iIN$h{FG(?*0FeZ zcA*BaJ9RDW{kn~<#zl^Z5u-)qICYfCEVXWkI-4D+Yq1V-k?iBbGmSm1Yf;~0ZXqG+ zeAXZ~oAoIlNlGo8hUSkwV*DR{zNv~6?YX~1qmPYm-8wW!NjmCz7tZ5PgL7tSSY zM5x6|UK&TP3w0b@6KY|rL%qC}xOeePXQv>u*bo5paiR8O|3apAn3s2e3$`Vz6K3(o z+2A}v+4e9C8yx0ko$JIQZ!R)jSt?`(x=^Fo9mpJmOkEdj9~YUG>}gmCYxABLtNxys zwTBDW3>U7MtnGUiD^4xvk>*0}$l^gY)efOS)ZF$ejdz%Y3gxv)82C#=* zu+5owJqz=y=jEN{f-S+C)U&9gSWLYTE3S#oBiDteI$HzD)sR%iyvuw^b_&>&!1}mg z!`Q#T{!`D31=R0kwK(-5n=V}KS!#Xw67J<4<$_(v4%D|;6J2EbxKNYW)A|;50ds2* z;w{%HoO1??YhYnr8+cjAxJahCNDgK34J@jaEp5s^EYG3v*sK&L9 zm;zj=eOR4_7Hb<9*-#hiSoRsHa=pZPB3!6db|g0J;r>I!tRh~62Y~wfqx829M;cq# z;vv&p2C$-$q0AELgHOPYqLDRgnc_Au5Z`AX8t382E>UR)d39Jq@Si2d)CMOA)X0m>$SV|?oEw>8 zZyKo{WeM%v)h29SyNWdqfUbRWwzA*bkI}00ks8aU zMLk1v4U1l+{qr6BC)%q0_8qfysH|qNMjgtk>sjXx{fnI|58Wx`a|HXPLz`k#o#QZ8 zzGGGO1Z&>0r}nfQ8St;VMN#are>ofn;hn?&^`8z05^{(WwYQtugN}231IlrQDdV%6 zA@bA)3XAX56FqvZQ!TC9OP1BkkCl(HXx-kh=uv*GbxcVuv@8#32pbpESiQlv$BZgA z+-i6g#AKQ`BTF$(E(&{P50jvhAeQ6N-u^ZKQ_Hw8y?+|IQ3* zo%)nKkGaA*M4pC^XXEoaUatG$EhXaS zEw2rHjJcs^57>Q_gdRvihda-%n8)U9)GN-t8IBuJ``1qApFU#ipz8TU_ET4@*wpH}?^5zVc4?b9lh!Ox?!C)J_-{NH!4G&Y-aaUpi;EsS7sXe?S}O_w z7dMA}5pkxteZz^qcUFK{-Wgrr2kh*eD9@F^Qgo9Zu)t|S<-fc`B5kjv#2n$3Nj^QW z#K@nCq3nGMxYl_@nv+1Zl!)5z8t&wM?v%^vx5=Gcs&#KUtX6zz4MQSsI?XDC~y{@7fU=j1+alkAnCRpev< zai+LU=S1J#MEo6T5*2^{9(xfVZOUynYP!FcmCY;u?R~axx?kxH zcezEEfM*op72k{JjrLavQ>gs#E_*J7X}+lee1Yit-{BSSb^Bd(ckl3u!z6Y8MvB`M zU7{P0ZXa$b?XMA&h^`SYpIawR^i3t=F(UH0xA+}i@fN^RbisFc#iMRgLjJ4bNAFU= zH4_l&e^h)@wmAWvS>UiTwHD#EsQmkte#vEk4aWQS(nTOI#gl4+VmI z<4r!zJW3a3$-K#@nU`1jG_&9;Z@MbC`7~3JYL#!A`A6YR*8sxy(asouyuqiLrq_9f zzk#U&8NPR&-5WLGlX438v03t4vDQ6|NSsNVj?@SrFI?wITPIhO--zV`Us2Y$!Pbnn zHvP$=%2oKdG#3@soGRjl3Mrat_rCfr@9Eds>(K#P!Zl`zt5PPCII%Po3r;cF2RqpM zseYEv79q~%$bzWCt!w0#+n+2tu8eOe(S{L?7h(G~wlOZJ^eZysJa_22eC}UTRlOLV zhTB^>-CIC>DK_C=YJ;Mn(n<2*D!VhLqULvrsbl@LcGuX934YaD!oi$LDYhf;=D$yb z8JFoGNSNt?drtc@2nerkTp=?txi}?|OqfgNt3ki3EPbq>BB>KM7^^O`MUw)wqNT-_*C;GM`;w&QaN}agGLMK;dTgUe)J)ZdTRjcD)yjlz9LnPDbqc_(l zeN1gI4^*nwz)P&pgo@g-bG%Hq;hw2XHz<~TC(ypPT>ncvaJ1KT>K{jZqDd|SHZDoF_snllP#PSs69Bx3wVGSO$A)>CokYR z+#lL0&7y!Exk|U;oanoeh}q6?s-8yyXO+%6$F+&OPky8L{U1EJ59feA$jLsK_zd?_ z8{7kxikESYPBs4*r-_Pty_B}Dsc+2^CX}gG88Mnl8X&B?oOyY#xz&bbb&NA9dngG7 zOZwt8FX<-4hmz*M5qOZrFU+^eBbmN_XZaN(I*VVJtCCJ$*Jm6kmxc2>d^;%Wuzc zpIq{B`6cC}mJjsXKZTI6nofNGd0}pNnqM|%Ah-f4`QK^265B=2$&`F>hOfk`o#rdC ze<)d|lwA3px0Ek#Q|xV^RitDwFSOfcPV~J$6U0hHj7sX3w_)6@|DRaQzy&!wOza>ItAQCTB4m!nGVogr+=sF|(0@0rj8${hc!gGD7H_4 zknm~0C_yeaEiyOxDDRmc5K4i}JXu zBxn}tI-W^kx!uhnqHj4Ob|xY}D^_Rnp1CkpWa5WR-ZQrjlL*g*QR>a2XG$OEJ>yL} z^F^NXm`5HAMfTpvi%`GuC|5+08;VLE;gPR`{{=Rg7p|^7IClQ<#XG~iM~%Kf>@o1JnoT)dE8$C zOT9nqFt7MFNSZ1>B7@Hoza8O$-~QPg_yLwNuT1bCczgDz@#?v);Y8m{<3aQ#A`iSv z1`m7@uoO7Lqv>(Q0q(GC!HBc8 zg#kG!5PAAREbdCPRtw7YP9qm&&PsYO8YH&m(Y6Mr}Hir{^4-)Y&4^wUznNY)Ed1NZ*V-xagLegzm-W{<>#uL!`ux`8qc}{IGhV(md*!O_l z)nWOVJOfmhBGKP9;5qp}`+1daEetIF8V(s%ntXtHFAC5G?Pnp2tbUD2qiN0S&t9aK z@sGv}UHbxP#{V#VZaW&x!au0!m`ceioNjvD555BtmOf zXV!n`T2Obsl*4m&@gm`W@ZwPobb7wWw!5fA&X%AovD;+5(aBM;UKBMES#K+D$V_;8Z*vH?44w@1--k6!heOoY!|B3tk*#{fM%dY^j_RZEjgcu%b(x6Qx>M|VO!V}?a6MIzPyxnc{kg-#G-wFo#*c| zMGypiB7(uYcRX8iNGm>JC%Fv^5Ffge?bpw%uUXpi(pvToF1`CQm(Csr=~_P;(s4T= zJw;=^SG6dfuM(6>XjL-*yA1}*jjwE%sNq0|)*WcmcLZH!r!S~!7L@;sU?8(jeM#0A3|g>CFVkjtv_Rhm}I-5 z<}4QL#CNacw*ADNK!Gd?t;j?sWPs+|EjIPq95&@#ov*v%>i(jB}BJ_+o+^%h?<+ ze5DCpEN&^H5?eD=o1Jqc7>R&9@{pAh$`$6pg`^R=YR>Ew+x>i2w}A=m;V$gqJ!N)@ zt+rRm88{);V@@03Kyd#h$$L#nvfWo7DVT;llWf`QqpUoK18cdjDziafS=ImTNc`%l zrunT$KUZW6zm3-Fufx|Gcr9LRQ!}mFTGnn;wa~ij>BNYkN#aUmvJd_SV1GB7t!Ney z=kflv;ETx%Nln!2o$^gZ=Q)*^C>Yx=E=p=t9foH9tD!|n@v z_66%$-2U=Ko`KWejs3c1Oz2k>J-;Y_Da!&)=>mBLGhUW5t|-72@osnU&wL`%px=H2 zO2jl4w06h7oy+k9DOP)Hu-cJDZ!M=4UBO0f4ZsWh=&hwZjAiL@ws~uydYPTt8iv2~ za%}b2hA(GG|QedyJPL65iSF?!2mDJ7b`1bd;naijc+K^>LN=Ssw*@woL!HJJc&LLPnB%fni zmeSqg70d;XVkz~RFk-6J5;!5k$%!_?*#u1YYvu)vR zH!gyA)I|#x7wDpHup{E#hQC!3C)@ZmBi`$jrOw16yQupT9{U!c>=wq;TP}^KS4-n2 z%P!^Vy|@qzebNi0w+>P%%+31LFxLw1i7Ib}fA*T(M3DvHwExNK?Cq)jwUAr5vyfYO z%b1!G+~oWM&_9BwrameLmxYDCtqlv+;Ek}53YYBejGyyEDY%QQ{l1>sl~iuwU@BR_ z>G*VCl_qz{KzaCPGVlm0gn@`wh5@%!!@w|jWdALl8;GQEcCfhr_0)b`zzwWfz(d%T zP9ZdQ%m@81h?*>9LWi(WG19PbFU7Dx>9QYTqxV-RT9u+nUr1$8m$K*k{k3T+Jc1=G z^gwWvGsz%dLl`FgUsFHuX)&*bq5m*^5hw= zA@w1)piD{Snde20RiK)V&2l`pPRi9D{qX7IR5 z&f#%oBtW_+95&gk0ndfa8I28_wi$-a&2Z9QbvDISWFE!%l7$`XqWzt~E!<7u7XHF= z0!2DL9rRD&uE~Ng{1z4lH!>{Lfj1&5Mx`x2i)U^$MR+~YH5e6-ACCKLC#G{tcayj! zl`PeUlO{`Vp-Whb`p~fC2VP-mDrNnTnLL)R6xU%Ed&0k3cl<0S?X$ylZI8jfnh)k49e-&TaKK>}*zqwla|gUW@en<_p1iqB{HO zN*y(h6+3Ayvj`^hFDcyYIn!CzwbHEJ$%%rR%uL(_gZTClT8Icot%h z^wc;q)BYrwCEDwR{`$GlzAM`w8 zqArGq+Exk!mP%#=rASM@1)`gM)Oe$saP!f{B#-CIx$EoJaLpqpaLqjzUk)!VpOt9iMKtJyn_tLe>lT?tpKvnN-qWxgL*&>7D$ ztlZUx+MRLi@x_@Q2bOXbEmm?Bg~xIgHQCT>;VRzEU9i%$#=kvMxv_*VIs5qEr z1hprNdJrC3g_APvmJ(*Ir3M$&I%)u~BPIehf@?06%rzHY$TiO&$mg2T^Mt4SylJjc zKVan^TFWF1EU4~!KSP~XegKQUKg{FcT&{1@0@8;+h^=dHrPgMX?BRGwx7BVflhePT zmX>{;TKor~cV~D^n8USPp3l?SwjZz3Jr?*lJao%!;ady*urU9$`f zkNvdOeYlgK&*JJMlX=iT#`2)OI5N}zn>!EsB72@=Eweecpw6Pbo$kJgC4KIdleoV9 z?9h{DT2LQeoix_wX?RGtnZnz0Ma|w?`V`cR%g58Y+Lyh!lV=jS($G2F$%VbR*2XOR zS$H`&PRg{;bu(*CVqwoqX}5bA9#)FzhBNH=y9;E}6Q9WBc=@Qs?Ju_Yc}@R^bV8;2 zWfMqg5IsY<9NUw<$PM&AEr^+dcqKs)s|(^TLHu42=St%99$fNEL0lt<3nlSKL1cos zKoDO~7i!}Lai$Lv_70xN61wbw1<*bjcESm1 z?Tnyh+PN_tzqa+)l!&#F*pOa z9QbiEP0MdzzSva9D@PJRe(I;ssP9X{$3gKfaR+~91BzjHCBQy_u|x74 zGXTSe<~a@lMi0w#+y@*!D$n8Nr6^tF@*JUnlg7XfAe)fqm1oEx9#hyt`vM~?t{&VYk}wG;CkdjLO~ndi6x*fS~5fxQHj zakKIq!GNF7&U1VMxO)yV19)a`o?`*v!+Cj*ZGdWWUVj2j*L`eWo0onmura=zNi)nzNfFl<}4lo075@4exkOLeI zcnC0ODdYgFEK5Kweet7I9fE+@0HXk7RzMJN5#Vw_k5v!^{0J}`(6JhFSTD6&gB}1Z z|2gCUM*xlnw6B94V9At!GI07Kn`#a;5NWtzk?j$-+*rb&;5|+@b^=cnmdqL z!0CVy=n-{~KES4hsH63zHm!2l6$LFtkGJIkNu?(~0UGc)`HZ7HzKp{sP=938Drx0B zbk7!AMfHweqlH#U-L1DGP+K2Dpr4)$5Hz%8o`Y|JVk_=5M{Dj~q)>OwtEjs|sye~s zcct|c(4pP8=#L5Jl+`_3LT8X(qa|0?ia@wN1VD=j*5`w$O{}JGBBn6?IH4!k)gRN) zPxpufb0xh7ff{-%0s;CE04?=HeSRcFXEoM05uEyweuChb2>mg^As_3WtpIO*qSqw& zqNU!N;L29|P=bH7(&rP5ZmVx1_+4B51fc4!KWYW_{iAiy)_+bZq zD8XeN^!WrgbksKisF)9<#?YLu~6Rg`k&P^dK2C;&bL<1<7LY^zmP2k5r8 zkngIGpizLHPN2NLi;yMsn=~q>tL*?h^eP03>a7U4=_3fJdU`t;4AHj}7_FZrJLCfe<~KKpj0U3ce)h zD`>Pt-%F!j`W+gj>uNMcr}P@p2>0f6y>+x!6@qaXs-NiTL@ciFrBO5eDvg+~cEIQx zJ+OmTP5XP6-ns+twILX)gY|SG&eXTl=%RiVBklQI{jUzt+Ds4Z2x4cwSw~1dpRW(? z2uZ7+jG=l&-wvYodVzkfBX}eAx5V35w{`+?m)@)sMrZU`8Wq!%X;e=CmPTFlvoso{ zzo5|?X2}K5^HRbQS%k*kHgp!!;3CeCTzxW8L-lV#Rp0DB z-b6>FQ9;vTP_wxhB#sJq#p0lxWI0nm}ns8gVq!V~N%WG^>w^b{6YYa@Jo@1dnfk zi9AHlAkjl48fv0F!E{Wb@9TF+G{Qt4s|OB)XuvRtMww`D(AOl|Rc{6!YabJN3_K_D zU_Fk=9iNLhE@Mqj#|eihPj@0T}%{xmym&{Z2xRFSidF!X{SSXXO=K6Y&TbcI*-6YB8 zZRE~K{J3o4TC@O)v9I0dx;hF8s;GyJf>HDCmg!h!lhFu07G&+%Fn!7>t%Vw^XArk} z*T4wOEX2J{e?esPW>;e+d5Io48WFhcn#g@xE6L`KMNde)dAG)k62DVVggQU-hF?Kg zrm=U1XYFB$H}3>{UEoV9 z=B;2={mwYd7^QHhNcPP;L)H`I_L6Mg_B4qqmoQLL%{w12lX&w6x;G`>ywBxBi8pUq z>Q_cMXWk_;fa4YY>+zbewZ3FBZ#mmfKRaHt#7>lC^B%@4CEmPe?+uAxBrhzXBLB)i z5^vtmwK%^dT7xXYRr7AC-8flI6krtx3EL#|k!?b`u_^6oiRb%zf;3IqekSqeJr1j2 z?M1~iZ`Bwo@#eiFM+kf`#k@!4H&US5abeM4o@z%W-q`Ea*y{DM#GAKzZEl4_bOmVJ ziN3)#*PTL1_jrjn?>M?z;>}y59yc^rc#3cTwU6aL28UWfZ#D@neO;cR=B<5OaWd{g ziEF@OETqXn^LDtKB;LFM@G*%$A*anaS(^tE|LqomFt%LvE2k&{l!~x0d?QQd{cq<> zym^1{@Aa_B2;-C7*3|eG8zp(o`_%T3c=Hb4{UzSC z)2=dK;>{arf2{{jK}vs=RnML%OjA!Q95^rP1V9mloA_ZMntr1)9T~qz2B?*Ufl+}W9Ms= zi?PubEAuU1areumTJ=J^o2<~j(A0JMomEFSFDFaUFE%2SN$pt=gHwb-MT?* ms?O890eD>D5B#uSrM`XxnEFf=>-lE-9Sqg=y7ddK+W!HRAh>k^ diff --git a/norch/nn/modules/__pycache__/linear.cpython-38.pyc b/norch/nn/modules/__pycache__/linear.cpython-38.pyc index d4a4c2cdca0a08a5e80a29ad666355d6ae35f6aa..15ee46288c31ddb6189927e914a6e86122f654aa 100644 GIT binary patch delta 45 zcmeC-n83ju%FD~e00j3JyKm&~XJqu4JeyICJ)qJrzbH9l^ASc4CdR1AXPNB*2@DLa delta 41 vcmbQh(Zj(V%FD~e00ieVEjM!aGcvkNp3SJnoR*)z`2-^e6XUJP7n$t=&vgq{ diff --git a/norch/tensor.py b/norch/tensor.py index 18192c3..a4a84f6 100644 --- a/norch/tensor.py +++ b/norch/tensor.py @@ -630,20 +630,28 @@ class Tensor: return result_data - def sum(self, axis=-1): - Tensor._C.sum_tensor.argtypes = [ctypes.POINTER(CTensor), ctypes.c_int] + def sum(self, axis=-1, keepdim=False): + Tensor._C.sum_tensor.argtypes = [ctypes.POINTER(CTensor), ctypes.c_int, ctypes.c_bool] Tensor._C.sum_tensor.restype = ctypes.POINTER(CTensor) - result_tensor_ptr = Tensor._C.sum_tensor(self.tensor, axis) + result_tensor_ptr = Tensor._C.sum_tensor(self.tensor, axis, keepdim) result_data = Tensor() result_data.tensor = result_tensor_ptr if axis == -1: - result_data.shape = [1] - result_data.ndim = 1 + if keepdim: + result_data.ndim = self.ndim + result_data.shape = [1] * self.ndim + + else: + result_data.shape = [1] + result_data.ndim = 1 else: - result_data.shape = self.shape[:axis] + self.shape[axis+1:] + if keepdim: + result_data.shape = self.shape[:axis] + [1] + self.shape[axis+1:] + else: + result_data.shape = self.shape[:axis] + self.shape[axis+1:] result_data.ndim = len(result_data.shape) result_data.device = self.device @@ -653,7 +661,7 @@ class Tensor: result_data.requires_grad = self.requires_grad if result_data.requires_grad: - result_data.grad_fn = SumBackward(self, axis) + result_data.grad_fn = SumBackward(self, axis, keepdim=keepdim) return result_data diff --git a/test.py b/test.py new file mode 100644 index 0000000..89801e0 --- /dev/null +++ b/test.py @@ -0,0 +1,24 @@ +import norch + +W1 = norch.Tensor([[1], [2], [3], [4], [5], [6], [7], [8], [9], [10]], requires_grad=True) # A has shape 10x1 +X = norch.Tensor([[1, 2, 3, 4, 5]]) # B has shape 1x5 +B1 = norch.Tensor([[1, 2, 3, 4, 5], [2, 1, 2, 3, 4], [3, 1,2,3, 4], [4, 1, 1, 1, 1,], [5,1,1,1,1], [6,1,1,1,1], [7,1,1,1,1], [8,1,1,1,1], [9,1,1,1,1], [10,1,1,1,1]], requires_grad=True) + + +W2 = norch.Tensor([[1], [2], [3], [4], [5], [6], [7], [8], [9], [10]], requires_grad=True).T +B2 = norch.Tensor([[1, 2, 3, 4, 5]], requires_grad=True) + + +Z1 = W1 @ X + B1 + +# Perform matrix multiplication + +Z2 = W2 @ Z1 + B2 +print(Z2.shape) + +l = Z2.sum() +l.backward() +#print(l) + +print("Resulting matrix shape:", W1.grad) + diff --git a/tests/test_operations.py b/tests/test_operations.py index d6552a6..a09d351 100644 --- a/tests/test_operations.py +++ b/tests/test_operations.py @@ -249,6 +249,24 @@ class TestTensorOperations(unittest.TestCase): self.assertTrue(utils.compare_torch(torch_result, torch_expected)) + + def test_sum_axis_keepdim(self): + """ + Test summation of a tensor along a specific axis with keepdim=True + """ + norch_tensor = norch.Tensor([[[1, 2], [3, -4]], [[5, 6], [7, 8]]]).to(self.device) + norch_result = norch_tensor.sum(axis=1, keepdim=True) + torch_result = utils.to_torch(norch_result).to(self.device) + + torch_tensor = torch.tensor([[[1, 2], [3, -4]], [[5, 6], [7, 8]]]).to(self.device) + torch_expected = torch.sum(torch_tensor, dim=1, keepdim=True) + + print(torch_result, '\n', torch_expected) + + + self.assertTrue(utils.compare_torch(torch_result, torch_expected)) + + def test_transpose_T(self): """ Test transposition of a tensor: tensor.T