From 407242bd7ebe040cdcf0e997255128f046706393 Mon Sep 17 00:00:00 2001 From: lucasdelimanogueira Date: Tue, 30 Apr 2024 20:09:08 -0300 Subject: [PATCH] autograd add function --- build/cpu.o | Bin 3248 -> 3752 bytes build/cuda.cu.o | Bin 43896 -> 51224 bytes build/libtensor.so | Bin 63840 -> 72880 bytes build/tensor.o | Bin 18496 -> 21200 bytes norch/__pycache__/tensor.cpython-38.pyc | Bin 6686 -> 8028 bytes .../__pycache__/functions.cpython-38.pyc | Bin 0 -> 594 bytes norch/autograd/functions.py | 6 ++ norch/csrc/cpu.cpp | 14 ++++ norch/csrc/cpu.h | 2 + norch/csrc/cuda.cu | 44 ++++++++++ norch/csrc/cuda.h | 7 ++ norch/csrc/tensor.cpp | 76 +++++++++++++++++- norch/csrc/tensor.h | 2 + norch/tensor.py | 75 ++++++++++++++++- test.py | 10 ++- 15 files changed, 229 insertions(+), 7 deletions(-) create mode 100644 norch/autograd/__pycache__/functions.cpython-38.pyc create mode 100644 norch/autograd/functions.py diff --git a/build/cpu.o b/build/cpu.o index 4220158a4b3ffa19bbdbab34a25da7299c6faa1c..124881cefc47431222c470c46e4998834a4db702 100644 GIT binary patch delta 786 zcmdlWxk7e=hPVbd0~|PjSq=<*47?ld-!j*K=9l{w>e2bsqqFvfM`!7SmSC~=O)A+CTx}xx>=GS z!_xTQAAlQl31|=lLxcTh9Tp8nMw!W>tnyYeAgdTb1OtN)h;Zc-Xk&8bWn*&YVdr3A zPyh-kK;=D<fE6y*9PcA48Fbe^* z1JW3PN(`$~i}LZQftZ;CwG(88AQ1og4*}{>JDGrDpm2tXD=p&^B?V4=^(Y z7=tDoa*H#5nC!`|&X_g1kXxJ+lB5_I941fXmS;?z{E<^!^aj*vH(*W!(wvj;a%+nQ out9`=Kt0U^yFqe-hIj!d0~|PjSq==m3_Kg{-!gAL!lJ>*s4)30tNdhRZkfp$Y-yAA*o7uP zU^AP1j!kH?1N)859vs4qLSSt<%n)-}fRw=Gvz(hJZ{RYT{EBNH(;Mc=UEF4Z4jd2# z5)g`kYw|^IeMX(hj6C*0(ve4B&;Tj}qjV-8MgY~SOL_tM>_E%} zQwF6YCOh(~OC=71k`~mpn_aLnq%@+UTv-l oR*29B#>uC6;yqc?xKPU^jvs zKm$pBaKp@~h`gOunRYxJ9|Nc}G9x&^Ad3%FTu^Yug+06Ao^^)(|EhbF+tuka?*8Xg z)xF>Uy#IaFt=s*>d9|a}(pso?96#L0Z3|DaWGLaOmR?GDux*i|?sLKQe(I$RffZ<4=C!$cwNd7Pg1b9S0#+M-%|Hkb;nDsh*rZ9 zhabk-Z>w&Lj()0kqKA8VtR$|D?)}C^Hzlb4L+X zxRJ3EteL5hkZCVvIt2ogrUfQW9$7bGRAB5##%izEy`eA{JGQ{!t4vX9u=wCKXwQ5g zr_Mmxc;h)fZak}D1DOTyeVGEAeHE!CD;t|umSB95F=s&U+$2*wGT^xaM-;X;*JZ-; zY!PkMO_?FIWQNdCBy<;i<68lPb9#nPWZkx!otWxBHDHf3KdQ0maJ3-aWY3ctSk~}x z!{Vr}`+v30lT2PC(xN+!Y_JE1L>_$ZI|qk7{_q!h>*9krBCxwWydS$By%fS-T>vjN z>3BvIUWNnC!g@yG8+AN03h!$nd3P|1u})`XMd8&}Qb>uyf2rfyQTVJll247okLh?0 z$0h$e<4GYcO5wJSd!z7OJxIP+6n;a;bEELK1d`Vxaq@pjXZRu+Slg0F3h7b!8#<27 zHc>g&-7cvsv8#xt;M47rx{TwE*fl}7OX?dmOf(a_!0qblf>5%7uRxz}my~L5s2x-4 zc8QdtEXO~>%(-1s<~V*DGvs#Vl2N07LKq3rQOtzf<%2-qx$r_!x|+w})4p}EtH`7J z8H_KQ2FLpb+KXLKT2v3q@VN`0U*faABo|7HYoNBIl%A@~0$&v`Q9~A3w>kwHOT0m9 zjU6xRT?tII%<7iPJ*9HL1V&lGXmojzA(HPbqkLx>cnKiMMlo7ECMBT;U_0l;|RO(U^U`T0DvBBG=0mTm4WdS8w zzq?R)vf_r~@&!Ta=p%R$QE$^cE9B#W>g2scs#9bMd|1$fn${%MDK%H6I;D0xj$0yv zsc%kn)m4@iRw~}Ng7Iq~raC%gNRXnS#Raq8OeCqpLXsPR3Z z&#WrjXq*>B-zLICv#N?q`D%JZY9_6JG%SfkZ=c6$roiFMl*FridAtGFW(`QZz{bl@)7|LP@)oI& zYx*AQe#W8#DCJ8}tvguF@56G0K8=rmMFWZYnD~;b zv_kKIxA$@6Q6KUlg(ZWQ@l|~#F^ORL9L@6t*1D6v3@_9dcq-6Hciic4slFh1GY0pZ z&v3ATR5^7w2##QGy0w0np|nQQ?{*NhV^X>^a4tXEKzD;+682$t`nCRO{cZ=r4>1g! zn%)Q^$ru`LC!{TuZ z3*^WcbOOa^5icw7(A=EF*HK0ngD+R7;vmNnX7I@_uHVj(QSH3X<8(M5=hL6_P1~6Le z`y>OnxY=xi70mfuWxVqq=1<;JnS2j(9;!^wDp%5yJ1V`EpZVjHk1%I_WrnwW1k3W4 zPxL?P9Wc?KK&3cu#iCTz-@mAmrB}ShGO7|&;?JrntIn!-^`AW_yQi&^ktRVTUnmm{}|Dw)FeuFvBsovywnDZmG)jRMb^?q;ZN!5`$@T6L6 zTWhn~SopbBd1`zvwl7P8|9B+NU&E(JIzII^aa6KvqdLM7Z*eO29>zBRM+HhA$x*XZ z`1;Wtb+8KYO*!g&3jD3fuWW#cOJqc?p`Ao-04E5qii zxC~6TL&ln|YDo?pTJ!pp^ZJq)&+5697$h0!Q^CRYpa~y`I9Yj#Z+-NbQ08HdpU;LT z)^6<4riVQ;y$%&9*YLI=r#{>t#4>*Aypa= z_JQvAGrGrRCgD0jdnx7L^?DpmL8IP8^@dceQ(XVWmv$O=MCBY0(MJGMsqtV_IPTDM z#+Zn71;_jFeZPoh2iy#4o_9F+W4B>OaBqNboi?@?emYPLn~e`@Hi-MG=J-;45>4W< z&E@z8`tCv!Z0M~;8NU;Lt_*E#+T|CSIfO~GVSRJZBF6-mes zc2Ol8V?_7*O^m=EyPybZq>62ULIL>51&n9n6NjFbGo%6mX) zPT_cfXIkJpIiAeFqaNdSKIZtlJiEs^o{8;3`=i;1JdHwOs=;9Ec>+a1Z5-G5ToCx5 z5!V|V_jPCKZ=EHN7=UrPyp#%_^-0!-(yFR(;rfl1?8>r0V1DBx3+FWkR;*k#HxP9e zsMJEn#$3}8WV%uteqrMq7W>TdvaZ8p-&H6={aPPK#=2esLs~TWRf`?! zQHth+kh&vDRl&0@nxjJE-VeU58eC|xD+A$5OTP&Ft-NPxYF0q;R!!*-qqg>our8KB zk|~Ckt(t>gyu$Y)xP;bEY_&TEOXxE&Y?}r@Y_%)Z5Vx(LLyq=+SdP5u+w7hiN#k3- znvj=auy&i~pNdM$q@?pQG$5MZN-^Bs8rh~VBzr?) zajO;)i%M>gj98vUt{lreiCYY(S|elG1HLv5658yJ(P%-${~@eHbYL6BGO8^yA9tV) z%MZ?{FpuwLR8!}Yev7~X?n?HUxe+dcQ91rfvdlFksA*sftLjqQ=K zc;O6c9zo3-Npm=^#Y9Fd7uzFa*&}g_!2y~_9v@dE?maLBG{-o!Bw}*GanPU%C@)V! zWGutrCh{&oWGwS>?Iki|`73&p`v%YWI=RJAyj}C$gLUC;fyA8vv$t!GArdzTm$qX+ zY|mFJ;ce6$D)Dwpnl*48xw6+(TrY_zreet65!oAcuzUxO_Z@b}B;<=&LU0DrO*^RN zJ9b32{4glb;CS8<+46Iemto*|MhnU@l;QG8q{o0y;b&Y7{9y@~a|9AzE7|FgaGtfG zzDB(yF{&kn-*TMhHAUVVFQTK8yu2*@QNrcae<LOO>s<+#GjxjM<_!IdBrj*^l7v%+p!A;- zf?rEujTDTP?;gzANtAFoJ9!+ZU`U5oS$`AGw*;)lz=OD3HWUQoUamk!?#D}HXMu!I zmGEmE?}M293c0~?VV}OQiOTubmF!QWa}nWS4QH4Qg=65SBtyf{n6Pt6!sQI5@!b|V zt`qWBoC8n5@(wMynX^T_`6xp4fuuGK>%vVM-=tBvBPD!74E*ODC--u=s|~!5-fMEa z2PK6EBnL%Mwo@ysz>O!72)D?nvQ-=>H*&Z~NlUp!_pvU%k2HspFX3_|(ILiTjIzM<=ca-wu=G8B(}U!Qo$M z!FJ9S;nJ8TIwxt#iOb;=gTft$muP%l416ia&B0V5<_ZonD|pCzv?*CYq>Tk!_RMGL0#7L@ZZ zLvsVxaGU~_?S#9->HT_2G9(Y)OA;=(K>Y4aj%Q1mTg$(|$ng{j-@tLvfLX{R`UB$n z`_AqM^h$I1RU$^JJS z=Uou&bZpBuEQ(B66{FrW2A7Fxly^+sE-gu3H9U_el=q@4Oi? z9IP?fnZ+BhJ_f#paGHPAf{}QME@+S8;I}42;>WVLV&Ly`oa_&l?0+EH4;s_JNYNF^ z5WR3jK;KFV^%8zllBaa@mlU(m&lE^v8Vfun8Xv9;W^x7@X~0W3E|BoK65b@?^Puke zinQ}Q6x|%&DEGtR+2i9X}KL^xgITUh3niQ!H$l(cT2LJ39+0?H6vu;yC7iT28Bwv* z5>7>NaC^gqi{r%qmTTopuDJg z#DU~yiA!2XpbWX<&{z7ui|*@PXh7}>Zuy3!OO`i4+lyLIoLqWnI3W>-EMK4s(m)9( zo;VVoZo3ovUO7DJ0=!haVLO*)9`99E~27a5D4*SwT7>;bcghAzwA&;s=Wa ze#%3#bcw-VsN*=k1fwJ;{K<=36r~Uv_YPGC!QQJvHe$=)j4)fZ6_sZfKAeO2oD z9QwTtT31_P(Z2k|3w(ktgPr>Zs26hJ?7r$m@wd0#;MqUG_M2Q?pCu?TbAOe3%m?lJ zt8J%!26}n{e6oLlvKib5s#VvUFy+8dth8gL@J%?YSKJ4y)f(g+9GbA%P|DZ;+V<;% zccZ9?<~zOSJ~T9Oiy>4K2Qv><*`~i~R14$a7l*3SYPT8HdU|^!(Vndct{m!XdzBC2 WvLwiOsmc~QMr>TG6#Rcu&i?`&rbA!= delta 7751 zcmZ`;4OCQh7N7SRaAusr0g*SO2n->hs9?%Bh?s6_iuj$G0WB4JL#;F z+pNuuL`6Z1rZ%lCzb)H3dN$k6+3mxVp6$y{di1P$Ecf2`?jQa$ceF%a`XN0e=zeQ0ON9geA2!!)}mt)o|&~Km>{v10} z%{3#I<@+{vo}%tkb+HvHiNStX7TDrk=5^oSzU?~~H&=o3xI9RWAEJ7Ghtl{ob^8@) zjIV&t9{#9}aU( zVR*g^7AB2_>*@7SG13fiDGq2zPJrnteSBw<&zaFfDkLP=!=mB$hcBjSJHg{ z9lj>QI}0s^nFi<<<n~5pn?-I>Ou;L#zdshGe-HIvCU#qVOrs1iW%N2^upA_`=PnWqF4~i zgHMZ6U`$ay{bIiR3+{n!W9@LQ=)o<;_DV8wP&aNeQI#o#G@b{#rX03O*HkA9PY3-_ z+cZko)Ggmh*Rj6M#s9QH;*4vvd>lNSYcwQ2MMc{R>N1NYd2)fzmHS32fD6-N|f&6IzKH*->C>=pp9(vLX~absR)tasYBbm|H#sJD#96(wx$iHNKK~I zP=dMCLuH+&w8Im%sY(*G;n)R-YEz+Ng*z?D5H#@?Me~NO-LHvZ@$(f_u1N5R4U43( z;WGoJUWE%Qw4qITti@#@X>2fCp!8z!)T4BMurw|fd!Y0bsH}HY7X%AqF9iwTM4#k6 zqGzcadQFA=%L|18xK){i zD4-Nd4 zS4s*>5RF1@6`unVY;d)8qgv&Hr#AmR={R;4p|f~X(S4xwkBr-0uwlzO#RgGOuDGEV zG$l*qh|=cn}Zv--JJKR-`rcL$~tjqbv^Q_t$5nM_BBJL)6m zo@Z_%|NeluNA2wtF{=@vY}MRJxHFQ?xBD8#ve@PX;wp3K9KIHd;iTPDsitd4Y%qg) zYm9P*{CZp1!gXG4F815Wcp2ZeGTHLSj5~N@M4Ue{ev_YFEDVg zX)os{6c(|;(#_WKA-P{VoQ3P@n!g=M^wu4|7LtP-^ZpHY?F#|6wqP4lBvB8 zJj`%&I|v8k5!dS%{-78Qu3_0CdQ3uE#ZEhfg!ed0{KPX;#S%9e7mou%GYi9^urm0j zFYwijf5m0`6We)@ao}vQE znVFem#yp&r={vG*U6}8am-dJGoNd}f)qqT-z8SvmP zyCa?ZmGsi#@Gi}rE@>>0xOs4Sm*$|^6!!iK7j|iow%e{0Le}mK$7qRn8lrY}FBZlAh_{Y#Lar0nAhvt|d#qcS3&|-B5#n8~~6UZwBHbs*e} zp4*$@mggh^mogD)7QvUODaSM(&Y`AjpMPE2C7pCA+UFn3C5f8{HOQTWwnZ#XT%$zf zuK>31(=evD_W9SR1UB!}z|!eIM^$hI(cDh|SnzE|C$g6g3p+J8J$(spk4fAq(1Kig zfD&A<#4Uuwof*n#xZJ5ZiX@p=p}Q0Nxy#>c3@&6M!|Svz{|d~JxOwnom*$v?vchW< z)OKONchPK|?DEgN8`_ZP+V7wFXQ1rIe%^2Q$}8k5d!gk>anM^Rk(u{Rf%9Hva)!WN zjFXLNc!)T2B)rtXz51SUh#A8Wl$ZVy#)V^gh9FXTn7vC{Rt*-uc0u09uOG}Ma0v;luQz7tgBs)C`P!C*_ z7;*+_e1j#&GbIO|jFV$|V|s&ekpS5ux+d9?J7z52sNrIFjM4xXx;5`6yd(-N@Jn@;_D$1g*ycg;e8?FB3zncq8$bv9Pala3?)LPPxI!-80vxHmP)uB?(-5p z2USH_g9aFODBP(Mp2WBamv}@~2JQ{2?tu`7E``8vGai^ge|`uhTuxxBgwHk2>djy= zGeo4b^*BwY=OmmqW6>Wkg}`@$=S|JqDY==4TB1&$2542Wm-8X;|1cid1{Xuv?_+$+ z@#>%JDBNat5NR-i@BxgI19@2tW}NCb4-b*RxBxqEvi*DgaNNHHqm&usK;9PSG9KvQ z2|q(7(?xdB6axQ?WPgGPl=iPpAsoCBV93r7(1$|cM#_B98eRUrcEV zhO>jK0S*NItK>kw)hhg@EwDm^7^hB+!9z??^N`cA&PK=FRT#iHcs5MLs0 z1mFoaUEky)AtWI$UMM{hPNEbS(KrdGp?LAHF!1~yL*l)2k)%L|7LygE^%73QGPrVB zbBnjqZi!1;tHATV<`(bvzerpfX26X1waFWBMC9W49^%j;Uhs1RaPdZ4BH`rgMd(Ir z;x)fX;*vG-zWc0%lh!^wh&m*kF44;#@r1cJze4tGcZ~RW`3?vb`_%M_i zfQv3D3&5*fzEwvyC`t;1A1hXJVdAmT>OmL&Cy;GB$QxNlG`w>xTeZ9ClO0sln8AKL z*(&}gs({)?4a$xus!li59Up6*#D5^oghR)3tS=5U;yVK~YWjoyM6$J-zc}|Z6i+1D z8W~A@09KvIu{0POE}eh&M50=;1FoJZ7P!UGJDI39?tr?J0r*+OTT#9jPu4w8Jv9WM zZbSmjr-le~iH4g~9_vcOZ?ze2p31g$>@Y6zuh{(4*#jC^8JDl{yA3U>Veri9bgOHv Y!Fd6W;@>7;oOap1*=Zm#3jMe8f8z0t0ssI2 diff --git a/build/libtensor.so b/build/libtensor.so index 03b64d409a0ff8a167e269b8d088937026f0d9bb..66d9ac138b805f36d4b19a7117f35604a14454d1 100755 GIT binary patch delta 22245 zcmcJ130PIt-uK>{0|Fu(=0TB35k&beu-5wh$Ms*sUVERj zSM?2#&7ayzBZBQo3Qw6LQ&L0JRmo|-N{p<>teVo+xwxOG>A7MTbrqS7ius~$%?suW zT+hvNULX}*mGnRB=9{@vMd|Fx@dQN|)KDkEQYly}1xx>A52#1tE(kUIEV&+XDG!8d7HzJB9RdTq;ZzNYth--un*T%|ca@%VV-u zzc=4X&8Tb?Fq|FqNV0V~7k`f3@@Uq~<`;i()r{$_S2Vs99@S|_oAt*x{HOihr|wO3 zja$<#^p6|YuLQi!Qf)~gr?1|+)gZ;I<3v|{%A)oC5ASK;a> zyc|tsX1`@)SG+>hgB%;i|G&XC$8NN-0S$UNt1#Yz6o2J!<3&6u+>|DG2)r@H6X_fu z35t>@`P=4kSb0EX%4v)w(i0-%I!N-99O3Kw2*w4HKU45Km1ycDB%N>*WfIslfWu0a z$dtPc1fDcb0OO>A3)0coh6_AG;tq)?%Za>5>S-?W_oQgsNiqh}01#an3YpEz$ z;>TzNA;n6(r8F>xMi5fG#CJ+ZXN?i~e2Gt$dYX0-cq`NuKYW@>Mc3I;?{MdcHiG;x zP)eI2PVva6Nr0r|bEMRr9s&<1bx40m{z};{oj)H74C(pXBZ5zFc1y`Wk0uEPL~q|5 zNqBXT5YPi#Bm>xTjKF1l6(;%LmcS%IRu;iAGJHeMK0Qviv8oYn8z~s)L7sQW3}i(3 zCP@&E$(Grr;UkX=TyMc98Hqre>PT|nDE*{E%`HPOLv~O(9|4M@cTgKdDBcOfNmGO( z-N0L@l0zTMHR-k>E6+;9%W?&vN2o@Iy7^-Q*XNW?;v=P_a-b<6$c9}T%iFIQlj^86 z{7!43=sl4sZDr_!MhHNJ4(lPc;64r5FT21&izd?BaGcU-GD05?7IkzzCIr%YDX>Pm zI6=^rWzz62If(R0_?g6WWs7x(rbq)H$PkF@iJ}~mk@L?MiuJ*J+w)-&>ez{bpoi`& zRH+Mg%Ftbw4lR^@IVn#7x`8xr!M}aHz{i!ZYPi6o>5S>)XN{UV`N{F4W{;mfYsOQ` zs7KqkGeFkuN@TCMfg9KQ*JCSo>O(3d79fZM-2*l!(lQ_4bqBKyN&?Ov+=;uxJ;~zYUaphcd2q1>xu~ANfWt_r_K?uw*~Jk>k$^*^^D*zb&DL~ z;QJy`Q6OKMH(pD3$?_Y>mVZt4p+bLxBnFzfF z#Q#^BaJjZ{{EP`lXBw&6gzHy)BDzd?V}n~!YMdrQjEQj9gvXlj8renELE$Dm0)vzi zrnxNfsuF9$<>i4Bk|gdV0l8LiLaIqXu6`WvY{Id68Y$C+XBsGNDQg{~{?}8E3CESx zspHBB6QPBHVk2e3Tbb}lCR|>5IW^ye=a~2lOn46ihy6kmfuf$s zl$$b6IAbCt8fp}!+Jw_aOiwNojup#DH6}d6Ko#Y#3HKMc?0-3`DH3f=0^TNE?j1PQ zZo>6dlL!$e9CuwtiZ$Uu2C67YsK2Ac?}CIPt(<Z)(D;IF73c4LW&G z#K~t&0@)^kY7-t;KC^K{^~G4fmz&G5@zqwGJ9NX8! zIjJM?8$Yi5!|g7;8PN2%Nol;k%BZLF`eLI_7r7fh&NS+@&fNHMv{9!qdE>_+dL5=d z#?CJwrkA3Se$?MH)tv;qWdf+a)>41jQeSDQFSpbeTIzHDy6)s2{cjtbW^rJ=QFn@v zj5GjI&$ZMCTI$)BdWNN*W~sNa)SK&dr|gqx9gy`Q2+VjJ7a@TI!7~bx%wE&$E^>V5wiV z)W1W0w0U9p+QM*BuRB&|TtwNi((^x*b-6n`)KTKOgOfx9OE?smJzx!vx)o^#H! zsus|nI>|Tq6gx|6KcHREd?cPaio~sE;j^gaQj-?ClDs$Svp$|{7tn|>p zDo062==wYaLEUi#U;XCY&q$KUq@9bbwq!z%X?5QNxn55u>+ZD(B zT4u-nPC~_>J4#@tILDjkC7Iq>*15rDC&{4wEnEv+rQ{el_zL@b%g3C<2yBjJ7{M4v zGr;RYAx)5z;#zO{>1fVke7O5=B zW;0vGcS$g4Ki_q`OR^6xb*WvlZ!HXjoUcJU4%*V}Tk7R^0pG@ia1CJxTRrD92}O!l zQ#LBOt@;9cKDiT~r0z{laQ>&K?)sWL-1QpA%CriA#QrqgKzv%tV5NezPm&%JB!3al zO4paLe7PHHB>$#lN*7EmO!DI;sj-yDn4u1-bS3|W%O-oiog+T|7_66Ih^93Y%p$a> zVL|lDKvvy4N?pQ!ZynYB0-g#9=1?p6L>`;voS5yZA0bs^ow7=K1Q zjvOKf6KZgG;`$1l!iS60%#z$15m&D5c{ZwDSlC3cAT$FhG=6$$hO&}&{S%rVf+u`P z9u$s>Zrlk~LTTDfcCmG&HsvP!qjhBW%7Z%lpRf&o1^Xm}{XB0GTEiEEO4sw(YeRqO z1|O8MH~64@i#RQVvdcjp_(yTu=*mW`Md15$ot}Q2Xm3Nr2qN;pH)R1G<|a-3lKW+h z$#G`y8bDq!FJzlLWc%#z%3XIk*}V=4k@J7%QL~Q7JwH=Fs4H1c$FPv^Yq+nb9$0XI zJhiWV0L~N{`k!b2>=>p!Ttfpg+zz^HCMR4tEI)nM9t{nb-IB@V3R^{>tTxmW#I8FSF~A zz5OF!=U#yAdTpD>#y-%eMZZ0~Hyp*gw$V-V4w@jiTy~{U4TqRe_c%NHK$NsL$t*UA58kVZFM>seRalu8Cn4Kk~LzV~wiTFEPV* z^VrSl0P}fhtBBnM;w= zoC_g7-MN)%e+QL9K9Hqm#Ha&VPDWU_R|HEHR=+bPxd(B5s3w+CaN=@7N%o%FrGNVm z+^A`!4c*0!&T0pvPe4GpIF9)EDm$=@8y8owg{(RwEc|?RRsiDyer4D3h5G$QiYb^V^jcFFOX z50skNjt%IZptfQ8-6OT*-}7PjAY0cxEaYYKqt38P-bu#n1MuL))d|fL`Z8DzM8|$h zqfXt<;xa?jf3ft;2_r`A5Q7H8Y^rIP&80YvgNH(V%C~$sR~z8D0K`7^4ECvZ#13n2 zrd{8sYCE_S2^hI9e?v@nm@r=#&Bw8U1&1BoN#<4Rg@&21G z!h($Oy-R#*KX{2x?Om66o9l!)cRLU9(|9oM8V%n>h$j+fiDwcg2KFUlM-0mr&&Z=@ zL-DxQl>s4Pv?uY^y{uPX;LF^8*vDwAdtLvinthp);?FKXO2*=cY95Q`)jSrP$Wdc< ziQQngZ9EpE@Pyjc4_ZYma(I(G1359UJrSdc$YYUE&0{eZD8(WYyhbcS&vTc}u~>4R z$08Yxxlb&Xk+*eX5%npLMOQqIbR|Ma7;Q~_bz`xnnqm?4A?(+S#aj5WpVjmXFK^Rh zxQ90X99y1U>~uMb=Tui_|ADcKy-&6az0xK7kA;aP*?*vSR)#E>gldOcbu+;3nTpNi z&455s)+l6x(qtJGcc!7h6BN{XmpgoB@aI)vwO7ExltqL~!zS;>IYt%xNbbT*c$5wo?!pCrJM$)V z>Mq^+Iyh6=UC+2@unJg3Y~D#<>F#-1m&{D$(8=DIpxs1t*;3<9ppXV_6klLn38 zwJYf~wpOiKY~NN|*-36@#Yt|)d4{d*8CkQwz!)^2bMZR}~@m_^6IvmfFXD;uCkScyx}tsDcduri54 zP=1PAIY}04P8e3k9p{bdf1Ed_=P7Px$1yOy4S9>1XQ4=#QR8(pAA(nyaZm`@Np9w2 zGFf$;4H)p4*7F!Q(drmC5r2~XG@w~V>!V;;0a=TQaZn;mT!__8l!8~7NTl}9Ji$#c zGWY&5v}A*Js)~gTO4JS?XWa(19-31HN-~5jy1gMSbT5t3b$0}>(0zbB3Iz)$dfhRd zca*o|_z~WYZAW-JK0L|ca?R=&-=hh=JW4zqs49KK68 z2<8oT_S&#EsRLwnh8^(IJ@(6fv}%&JR`A8zf9VI=;D;hu#qcOC=n(s9$df*6-{LmX zH?k*&W@_gSva+ELUtbPZx*7$r8$;Wq9$zm^p2D|rb#1Id%dIxn9)OKD+7k!ZqK6jy zeEBBV|I!BDhW20bHgx7-rE6>>-Ud}xTjSfoy83@OP+$MeeMTE%_Hi2n-rzQ_uIDzE z?&mf>T}L1u2sGq2mdNTteDzt^M(6$YZS<-%Y~)w6bHkVTj4R_dvft!3j#qLUNB=2o z*nGK-S7r5(kJUy|WqlhjeZfYI=%jtQhpig1&u7c)T>sE@Y|JB>+Lvyl^<7m$3d>lQ=K|C_E@p05qR%hUGOI`n4 zeO_OGdIcXx4R&$;FRtPGXT8StU##Hz-Kz<#bOm{E{rhC~EBu6^uKrgl>g#`F2iMOk z*pSDY*fahC8F}e#i;XB7J+g5`dq03HLz5{hNehwc_v~PsACHTO6T~P%EdM)r;7W-g z+6B=^5I0HUP(gIl7+e{8m$H(yRgxGfh}Q)12SKz;V)ehd&94RVlpqF3;=cs(gdiRg z#0{?qho%eSZb95Gh%ZZG8Yen83F4cA_@E^Iu^mBMx+k`Df9+0e=@o33GL|3`uB)Gc zT^VRLfXzAUog9E{ok3fZd3iOjsWvw)28~~7auXj-?#2h)mX*UhQ4*ZCav`$uJnHJ zh$)S%aK%(V=AJf-)l9ar)l>c1yVIk~?@c;utLGDp*TepA+H`)FS3P<6tVk;sh9BwW zkA7G z|HE+Wkt=n0(KUb)D7d_ zLoxsML@@k=4<*8=>xFGT?7-7aoTffQpHLc#Jjib9 z8;Orc9|YM=eE`Lj9!02w?6sjT#O{XDEQB=3ZVIW?^CChQWS61qL7b9s5+Mq*n?fY< zDug1)9!s1&6~YlpJJI4GdmMY_nKA6sf(R{MVK<-2Wm^j(wM2!DEf~(uJ~N6%%(XK| zK^`kYz7_dh%rQ4pvwN`L=1$iNJy_Q5FjhFPiIbXxYmQ<|@K*Gu z*@8^vKvOxzTcNPP9nP{%_>D67DZILVJ=DZ6%TC1x0fPVm>L*=5?=uoG%HTK3 z;HM7vl@95hPy9w1{6-o4G@|e%jyE7&X@L?|AfOL5a>yt{fQD=UTQ_f%>c!OgNj6_^ zMPt7#?8!10*nH@+M2JmCv@(SqUho9#wZN`=v)lzv_R_)!nE%2=)sJ1B--S(In8LQ?Tr`V!ifOo*;$I_X4x;!XL;-XpOvIxMd>FW_S}T@)~RyY-D(OAEtE8@ zDe0r-6s9S=J|V=%C700}roP?D=ZqX2n_1S|G3qE5_jb5XksOgZZ1JWTHJ+_|J6t;x z$+m4uRKKQ*Pm&zslbHHW489JJ+Y}yJDzQRtl@>2WX|gISSRKYjgRHMWc_wl+k)a;s zd^yX^|*wQ69m5 zMfp0FLr>61Ls}*WJ(YQcwz9Z)!$U_T5P(;gOk^IR7uh+(A;8K#hFv7?bdZtss}*r$1g91fXiAAJ z?)`|QUJ@}mg6ZH}Xkw(|HO`zZQ4=%ffnTA7_3SWP_kOsuuy=iKOn`dai!eS(p_gSh zRoF+!7s^>n?Z6nX7o$Y#9xgBV<^NWS(oYsm?qYn_bKnJb@%;!^eci_X{2)Oc!s0#< zwT#UPY6CV})&j2wv(n9Wb*GNV!Dbkxp>H>8qOhD@6yjE9niE#$M_AluZaK-Kg?1EH z<~T&no5gJmA;+x(w9~K(9QyABV(%j`7jZQ#ShdBmm7ANWxtR1}tYUMlmLE;U_yO`q z*Lv;^y_Agp=?}IYvV4^pb zJYzac4bc)iwoX$bI((!=rTYYX9n*s69Mcjyj(j3I(07#*ptkZGO!r$=S`+_{P9@-| z)*#?PrG@`7&Hq6q;JEgP|28Gygcjrfxe{7$9JKxFKtFf zjT$t3)WE^LatCCO>Wew0HumUBuLNg;HX_sQPD0tcm)o6#d?WHn$ko1XHx2H_$jgvl z?eBKFD^U1mfZJV-ymX-3O&5tDa@_7%&-I>VUxo&qJ^5Ks_0rJAfpa6NxQEqny z@`px49(l@G$lLH{?s&)}A2|{7$j>6rL!LVc^2m=OFGJpGvJ(nW7&`?DkY7Y@!?@}D zq}$yL`OC zY4iLma=M=wIp6J$z+hX6yfg9!3(-R44dW9A-{}#AMy!>kV9UE+>TXs^b*J+ zAN~#sBTz_O20`R`$jgxLSq?$uw~)J#k5~agbbJx=ROIEwXesiB&$-?C$n%gFAtI}h zm-@0LyAwT=y1Ct(*?YSa+HcN+QVa||(KhgLd^-0K{ix%#qRhr;Zf`baZEhby-ERImrNL09d#HQ_`e;ms`zyfs0{zL67<0&L8*tocz%Ae{Fp2ov zk%oSM;Q;I_A()j-F}v=|06J};Low?Xplle>W*H$K2KvnZhCUngLw`Z%cWy6%q2)NY zdqN!pBMk!|fj$Ft{w}rAf&|{eSje0JeLj|hMP~Wzr5;*7SHZ{0UGSU)&m=RCHtHe6 zNCeiF=m~E3fI194IXxCwXVCkC?l9Ax!T>=ut{(@_vy%wNYwTwN7;%g}OnZP!o zxE;lyiR{4Mma%Oh(6}+=OVK~E5N?Lg8(73*QMbKdW6sLZ4rJ_gEM)Xfux%a|FQfGV zM(d}7=f%lVZakM``IWYY&rgw94q=gr&tYAQCp}%wdLoT+LT5)5O=UYOn>s(ilE+!S z>a)heTob7I2+JXN>pI$WAGb2Wvk^-r=Q&rOhvM}h7L(WV-EMvhq3amp$L8<#8BC9I zUI8r;e1?88CU~qDNLl`5LYK`I{m~2}dtm=n$kR*2wuv6>$$gPtdmwXVCVOsQ)Ns0v zsaUz`-D6vd$vy3ETe~)v>me8Yu##tBWluG`F|pRnFz6WwIyZBNP1zUaH4GyEoK05z zUjy|cTn{*PC#Sv!s@Ky*4fqz+#<)1ptK+s68Y|f!>Gdth)1P6R_D4P38eSFQnn5d* zZ7_KCR(j;Bx{u^yYg{YhaLwpvc2UdXK2juRLr;IuxtVZbW;SGApGRh#G$yyslaaUp zTe=xbB(_837A{496$yGT-L^pL`PX0Tv0>%T#3fAX=_HnX8B%IbGce?UAsJ0JobiY^ zE(yIrf6U3|9f6f!v`!7R$ z-&dbkKi8D~nkQaV(zG%9E5iCK!TKw}Xai6DU1XH>I=$~qNw3rU&E`72FKw=S`mmo4 zH4U0QMZoedTPgUHMICOc72yb$!)@cM`~(&K&(lVPlM=4GJmuiF|8Pv8G)keNl< z9c~(5Ds`fZcq){4_tFu?uCd=>sw#{{9cik?hOq~Z>}5HcUBme=*y5H)RF$b?@)M-6 z%u8inCiBfQuaLP*zVan)DqDJv=}ms=FGTAvKod@(ay`R-I=WZg$hI5{@2fXg7pB|; zw@b0~c*v4Gu#Kv^FYb*VzqmQr%L6CMlt!{v$9t<)Z07NBXFd&lOY*>*|Mnych^hyk z!{~iurt&A@`DR3v>4~@AMl+S3Y;V25kyd6L27(p+XdN?wo>ZS<*H45xLzVoW%oI_k z2hIBH>}D!GP}W~tuZu@0`oTVR@km8~+q^DL3op8+w;4s*1NWrX-X#_fth{$Ey)N)* zB}r3R!;{%gYlP|XvWCcKS|T4OaecCp9!j0?Q5CL_*Y!z8KC1U=_%z_8UryrM?+WH} z0g52Ibh528)$$ttb}6@C`o*t7z)o|qXhk2r`V_heT)9ur-{btXanYZ<1;cTha^20v zuu8pfD$AKEebC!99R+Hu!{Qrl38_tQ$?xWL-#(Z8ITn7Ge6mn6OltaG@=vnxCwK@w`4)Vf#0xC=7Q%5%kWy$N+>ruB7Q8!_Uvjk6f-jQz zY74$g;$;@x6I%|_vk`b(YLq^n(mX`xG8<2Vgaj9F9a$F83bS8a6ORNq)!1LNVG^#q+1;B# zqym1%*^B)DXY_yc5y@#i&Z=7C*3+ZBu)8A7{@KE5{RE*1iCa%78qC&x9p3jzNwuC) zRU~oi@j_)1x1L^BLE0kEN$T}(!Uukzk76OUXy6Z$TK9Y_MQPhmcy2w^sk6kz!DO}P zl1(6-#-Zmtlme-g_Ppe^o>jI*;*;`(M|1;1=_cVM$!g9=s40BL_ND*MC$M>?I9gBn zS|D-j;aRUq+n!jf#d=KIOm^`?xVn`7w|hwnTaOz1Oybrf z*<2(Q>4%NW{T89dOuZN$`5>s&;4$(-6)!{d7^kj~RO?AuFB26nJCKI+lE->N-EDTX zIy^cES8CF4J>)B$aNdoB^iuSE$!k67tXSu0|FPt-9-vhraqE#}`vX|_tc zefbUK-#L{XaReDu4xj z7i|574Nj~IC|~nk3$;>;+0O5G`PMy%3+HLE55^_1fBsO?&~o+`>*k-}f1lqv z4GGLBU*YPd`qZAk>PfLv*pS?wGj>vt5wFt zpG>%NP9ohRR~Em#JIXuInj5`u(GxCwVjNM%y6$ zzDv(>TiI_yQG}I?@r@K;{l*LTj@xdeN9uQ$dTz%CjgpNWoe<7qZ`y{dEozUOi`30**Ugl`sszEV=VeU2l@h4mt?79g%es{kq~A&C`IsJT$*s1&`mNgR zDQw%VPJ#Ljn#@~yGVPbPLHhlfo)=AFJ$^~?5qD>}OJnnX>C{Bl^jkCiv{UXOeF0ki zi#_PfMd7fXuSOVew++;9*Yw<*W&Mh;>2KS8^&7T=t8CuwbYJ~Gt>`M-hN^yhR(h5F zMpb!xQ2H~zfDdCicjD~&zNkjnb+N^FQhfEBG1qmr?M}L{eh+r{I{WQTdSH^=b0*yo zGDR%wZYM422Ag-cTWEmMN>O%yL05n>D`40_l; z1_eb0jiMhSq67h{LO_ttC)hXzB+`ru?Dzf5%x*Rt`2G2tyzZWv_fzND*_pY`PRSR} ziVKdysK%+S6kgH=Q7S5xqPA)7xv`=>H@0Y6iZ8LRXc<)Th_(viq-qv@p=9xIlP-Bk zv<+(2s{_`e?*xC*iJCN<8&#Bmpv@FRQG*;Q7E8rqsaSl1^-%XFZW3m;+vo<06J>+S z+Guws-puP}Dcu&R8zGT~v($$LQZq-T_tMn%$ON!x>Vz$aBvZk3Rt zbV98gKHc%TMWQ_=?FHH!pLBfs;WGfA+wd8L&tTS1YwS-ASDj6UpxJWBml@sv969s! zPrVlpy!LG0#jo%C;yra^eOBd2X3d?=qg?fqb}x+mF*2sx*X>t)|L)PwRgYco8-Dcy zd)S#AG&zve!MEb}ov<0j%ABd{t*q$ZNEYPsy8h=-lu+`eE+tHD#5v1cUiA%Lt3s`c zJ^D|iYndU*vg@R}W*IG4rBP$Wb)sBTU1aZ>3t>#H8|CB;7-y+F6rTae2wtApeAq8!hl2f+&YHfj@SK0Pd9Xq(|~h=Km=vpDyJa z^%sDFUqln~J|$?NV3bP*uMDVZh5(+HxK9SO`fdTFN&H{w;7{qGoeZe4#8Yn?qm>epTY9yVS(LlGfG_6S&d6)>8hkOq0=#SrV`Ea(SN;1%S#ynWs^*8-oR`Fsbmm z%QEATafC@HeX>6rpBdYpHYfOg2tO@U$bUs2wd*Zxb{V~@GU<+YzAQ~jmHb7jYjfTl_Nduhr;*`z8dymAcR)U&2Em7$ywM42yB{ffMq zH?7@K|0&1XD$o1s+CM^%scb;_NtPJ#rthh`V#P>+VrI)H>M2THWjQb^D*o=b@W21~ zA;ecldEdf+%)(z-^a|X%l+P^u5f*;Zv5Qz2913B=zUEV2kOBmRO`Y0+mWz!Nn-w!P zSWz@(hg2{YzaXxzC>5xq-8d|Sjv~45_&=6Dk_Qrl$#78aCnlWHazG@ zfd_1O;begq|4&e)Lun?e{CkxJ=L-`bMPQX!7@aK&r4}5yH%plX?`NWlvfYBGTkvuV z-r9ou4qF(wIA*D^;9X4=`wR;%H$I$OX~E+y@>Lc*!Gf2G%=s|H6eZHah^i0`y5z(& zNf{k5O%5G&IE(rP6}SMTm2&EVx`n zg_aoAJfj|g@b4fC9%8{gJ_{q%!iciqjV*YB1vi#kVkcYhUKaUO3oh3n&h0L7A6?_- z7KL<+LVFADwcueEJkx@QTkz2qysZV#vfvSN{Biaa3q$U#IG%06qb!EzTJUHKo@2pd zEchY|-pa%uz-?e*G%*>9vdn@vwcy1T+}L0cdzA%`v&fg6Qr4b{3s=_`PFv(yl2AW> zqoQ~hl+D!ia&PXIpv_!VPx}$@1-->j8`ri?807Gw+G2yEG>XTZcCk_1KWI0#wZm=g z2DY}#*1q;-O?#LBA8p_N4~OS$0sKs+=+}2)WYp*oh zK9TY_O+d6?wzUgv?FF{>d|P{#tv%h=e!yt^f3RNBC_aa7{>_p? zegc24%&Y_sh)s-u&KIzQDqDG6Owx7z zTUU$qtR%U)JBV!LZ40DcWhaxHH9OV~QZ29R`bJu^1LLmgx;HPQ%9}gCN}1Ij0WDd_ z61#=QTFCb=wB-!Mr!OOeKS7*t4yZwWO{DGJ`nHPY=SR#~n@*std5K3=|?9;qV7Gr_ESj7Xn9;enIa6 zR3Aij7|Ksi&ifW%e}I$(C9wBzgVWB&`CM|m`RS=%&g*V!y-K`--knWg5rAzOm&jy( z1{>Bk=@y5{8wy^2#_rsJdP~No`61Bz9gD7T-x<8zj7#d7)qsBh9C(JUZu>&87gfsE zeXMW06m>6~(XN|1ioM&enQuKdWkUDwzj)}C-h#AZfF%B3m_vSA%Ved3wMw#H7p#UN zpW6cU&_2@#Ez-}GLbHWXD~o=HWQ~>j7&CN1wgtxiqw9riMkjAl#9DaI#}G|(3Skjn zBs9cTDPU#oW7KchvGy^2H-2geU&c=q7!tx|Rfcd^)Zk;g$osayl`Ck=!PBFPrS?zI z=KjTo^^K0WMhVEv2w)Tjlv(4>&^YY)DUGP8;dqD_xR(OpBkCQVDwmTpLs}3s)r@ZT zU#zn8+}4vm;eqAPD({y^IoSd!GraLIAo_Cf3M=X|z`dy>=LQa4VgKur6m{_u1wHbx zH~)?@9(>*&%F!E4O^wjD{>3_`g@^T_j^rWVx$`T-pg0z#6Bg>Z{fMM6$Taa=1YbX2?7=O`739e<{4+PbxbR=B9jFC5eBgqAxAk#H+ z^6^I$U$b&NnFt(%qzECAhcL2IUt9qceI9Pks;Vi31tvl1P3&N#?uGt|q3Z0WqV9$&-H{JKy<<2QweL?!Y3`{vPH_i_rX)d}O8qt~D}yR+_1Q?*k+C=R`m-@MOL{}Y~D1pXkKd}8Rr zZMt6O%#i;QbCcrzmbl#m9TB4NFy|bf7&6cD^dEueTIv7!ET0%sZb@jk=S+2a zLeKH^K$GUpgA|i3J;Cqs^t8c`nm`)$iKpjD9=+@MyL>QQ!vl&yklEoonSRSKcM3G9 zzXQ&)_UTDl*6(asdX#qC89ur?u#b8~gsq}TYK*RvD|!EP{0tAc(RGO`4#czZCa1}K>3<(6s(`Il&-s>Cd|FTKIs{w?qj=q#`M{< zLTF%#LMP0@wfA>JI0H3IE~|`H$Msuv_6D)$y%N=h?6zJlBNUR$%c#I4USX^`T~_ei zjcgA0S?H_C-MeHV-&L10J;}uEOiX@PKfw<4Ivch+z>Uc%f|joz-zHP79C+LpXlMp^ z7s(mH6{O7pmr`E9JbmKn4(=1t=UX9Cj;Zi){>a@}Fe*sIhak?(&Cl3Uy$oDE&7E47 zp-#f-h$J|@4+#;n+NInUpyMEIYGeu(_7`jz~&*4tt2l7YIW z9jRb!aGtvgKCmrwU4Naezdc=Bae`gFz156kCn<*`13Z8{4}L`a)`U`qO$crLkh)h- zKtf)q6C|BSJ8p_Pl~daM(jX8Ecr3TaJ0thR%VAG0smEE-kj7dJ{v$B9G5cajsCNHJ z_T7-Sed`D%1i>%X9J5JgDW((=lkCE_z~{%SWBU69>vl&%VmAsiA6q?hGyYYq;W!rJ zg{)-gXzk=NR(W@%w)GhMb7V_x{z-OwX0zT^NAZ)w0hzpvIYxGmfL@o3L_rrI4truoeGHWUIUPOx*so76oI$v~Uq7~FSj-yG#G-apD+OgOa|H1-J4Tce9RMUN~a9dHyhWGyO1k)8ZI+Q*j7FA41>O zo4GI~+`JxXxH$k>;bt)63M@Fv-T2An{=;w+#-uGmBi79DW;g?wIm~`#>_iXxu zUia}AMKq7sV;diApSn^ud)2f1c<1}tKKAZ|kM>&OK9)K5nc$WA+w3&PzJ{BNOVAO7A~+kT63vk$}eavyPA-}Taq+{emt?jt`>7>x)I1o|L@DIWHi1rvxuX z@_KOIA;BvbJXP{k;`zQ7ymf+C@{EXO`!;kucS}Oyj_R8@_be8--e039fiJ!VJJ0nG z^{T*?c%52UuJgkKgz!`x!h`&L*J0W3((Mo_==~;smj&Ly#shZw#E{F5OdX&-w3WqA z>l6QAN~3sO85HvYKTNJcxE-(j_VrelJ+1Q7>>LMM^mu=^BCP4!?lVgrp;2SF2LAi0 zs_Cds8N{@?qvE^*RR}`4-E}*v-SKbsbo}dCWq;1B} zWhrN1Ikc7OxGfxt^}lCAU>C=&{{j;DOc%+A+N>zb7LH@%z>CHI0~i<%#roeiA@VuL zv1zfD`!IeT3f2l+S;%yiyl^1kj2sqLRiPe z<5@|8hkZOdh^<;Ynk6i*t{0-7^HSE5mr|Lh(CMPc>$6*m#;f(&t3}Da4naJNSh^Hv zGf^jZ&QMD|#8QVmT_(4bH!c~;n{~*Wb;#2&z_T}Q-xv}kPxTCZ6eko+1{CPOu{1a- zAgWU#1+xx$v+fR3G3|o+*0Axan}1`uUOH55$i7-S zIg9l9&IFWFO#YMQHJ~j5jN(O+zJM;j@s7jn|gFyvfbFzZ1H4 z@X9;mv}axHwas1K-{1;U`MNGUy18qcoCczqY*D7p2MJg)hR$#otGpD3<$smDZ$0du zE!~1kWu`OPt6Sn(PM)(h^s=R{NrL3apCS?5wn}(WzsBwBi7wVs$d(E@7KI2>D8Gq9 zrBukWC`6LNe?-4(dB~JEmsKGO3W<+Va5v~fiJ`!bY>AIa5cg~?vz`bRy0vBE&YQ5a zjlWSS<#dq}_{j$5m~U3{>E#0Tf}tSyd#R!jUJN)klt znDc_r%jqSl#*XfY)82N%YGk{@ZfiIjFhRAwmNY`KwWS!u<_} zX<*rfRg8=q3e#w_3m?3hFpWFA@S8}4X_+0_74I5^@llOL?~Zp}wj%xq!tiSIl-(XT z78tw_s3>J#+z}BhYfgfi%KBYw%rZ4+(xhZ?Pz)L*nlkBVq(iACtC$nSiuVKsC&;+s z*e82hX^}nF9^KQ*<10(+;ePCanLXX}-EJ%(;Dk zQSeMknTgMD_p=_q#Q2t?`qTZoz8{xN$!O?#UhChWj$?;bziyC1qMzaOGA{iUP?N8n zBR1zB0Q}HoU0+j!zry5S0$v3^f8X1*f2Tv*_bFi*a~&Sib-c7&?Z8p08Ub_xUjFb+ z4JdbkUpM>a{F&hQ!C2(4sha^L$_`?KyA%T6S-Q^ORyGBSjR=*`z~6_ldt(ENasd2e z^K|_rd6l(BmRQw80KT>M!0*0DM~2hR(l<0r2kv|3mN}tl{8Z z^AgMi|4;DwNdU8h&3J&r6QHyd0@rXiK5aFSk?YX1xebM6N+C1Khw*8Z(Z&k_U^G^a zL*{iK4XZj=QSGvfjX|{u2Hz|nn}=#IRMT_VYrnKk_znwnY%J6Zk)8x7H^-+OQ$PaR z&UFqp^w;ptUSPX14bZ#%&bOg#M&F1>PkImIQ-&#m>lJW4mj0`=-Udt)9Ggp6bjwmq z6Z9^>^M0$HX!A1B6!YH&Z1JygEq}%g!bMYVEK0M_j|+J}m!f(c)zdL`J`?De(=iQc zzfiGI*Xady=ND{D00{xNQy!Zah-tJ1=s?Ux^jf>~8IA2a9bIPt=Aqmd*zZuTkJK?t zQCxWtSAG%PU-C)0Azkc$FmZ8i2;u2}IYK1f9u+%>7%8uO0GIgjtZgnR;_K ztD*8UX2%=3p)huXe;xcUV4^zI67!mz1>cX^bTrr|UmMzi2TBI$$;n;_+(iQYJ4vKWFl!l2mret0!RQ5Yu1%)y9`8aJxW7g*Uw;E2iVUw{o zsGNbmh0qlG0 zD#kCNn({$r)HU&7Wf^|Rg|itKJE_@h`Nc@J65rvzhKg~(u%>(?#Yl8b99KWfiikQr zAu&2><Yffs=iC&sC3*C%6LTgoR#7@hLx| zjbe=84XBgeZ!#|Qm6t>KCjZ6&HRLaHZtZJ(oy#K=GG61SePaafRQ}Yt8m8NgSNi!p zG6thhGf*L4QMTA(xe5jAJ)aTrrYTaY`T}DH!1ozz{i3wK;xQI+%Q)P~fvI zs=rG4?l$>mSRTkux(%Nsajy;kgm51eGHr~%q{3($-WRhaImoi%PfL7?4gUx@W$x~A zq62)rLG=fT(^A0N{oUM6Kh#v*+FN$AF@Lx8Ik2)Scv{D&0HMmiBY>-Pqsa^!tY#7~ zlejTF`vNEbid++nXLYx6y!I3}gyZfpijrWM7IZjwP;+V^&RC9Dk8?w7GRJGr6!loq z^+?xzaQT7Tv{z9}LV)%&hO0;dc+YfUk?(O)U1w+xIwB?P=Qk@QZa*&@g%yu1+D}FH zpo6_l=^-<5?WaB$0#BmHM(R5x=cE5pkevNQA^j{P+CyCon-8RA!O7iv-h+QOZ zKWX@gz;CT}jCGz=upi!BFLC>s%VTulc+fS;wV&;c#lWGk?1x-?N!)(2a~%8A8L7_Z z-*3@($F)+@ev0*o#O+5bJ?xY#a!@-wt|GtoGqHmtZa+skLE`p9&d-pw=r<)dq>mUp z#zEmp6ycy%l50OtdtT!9)2M$-+__w+eN4*MJROav55(;hjwzA6Sq|=U zuL%;jA8l^kkUTh)?wrOb!=;@4DEg$AA5|mQP0`ejsrB0>ag*Cca+*KBuf}oB;D`y506q)LXEAoK`X}j zmqe)@f;S2A(oO!I(2kN#V-=+*j=ZqtTO$3b(W)oZ*h(8T9jS~)DvhnQL3gs{vyn=u zxQ*TPqVN1Wqt&i%V;8;bJHI;yZDaRb{GER=+9AfF0nN#rdaInWdx#J(mau}9O6nmZcor6( zLYTu|tTc!SLX3Cu;K70pb*Y;b5q79!1sx(V>Nma__sffIVD`;?e!uzs`M%k=HmALL zsLeI%nZBp7Scv<*DSp>@#FN2{*30LD5nktO(ynr|qL1J4_3>(N+kfmyoiY1dsMknV zGsdoK^g|;^l%>kd+j>`NaZ#D>-{ED#K3~|*pLOr!%PEh&TKmLP(^AA(1=Cn%NQ9FC z>DfWtJhp{SGmG(&e!J_c^Sax#o^2nv^;8b>Oa(sdf z4#ZoxA=erz`Eu0Q89$t7SE7LfnPGMmD8}d+r zyJ&9yamH!jrNBplR{);{j{ftK$H=z~{B4I@OuQL`rk|!5D-fuh_Nt)gz)j#U3nbW( z+#UZj#3^T1KEXFIV1g4>qAa#u=sjo4PDw75+f7pyItsi7_%!e%id&?14h$vAKzs=} zO1dVwThevO-5q~`IH>3s;?zH5q66&s7aVAVf=qd#g+CS=1s(+d^T6?L4duDTu(O%G zf&OiIN$&$ke_L|560an8EAblQI7$fpX}kjiDzQ-@!OsE-HX#m5T`VpF4Y$($mNckJ z)#bVX2k>&u02fPj>2kmmz?XpID|;ciB`QKm*QLQN=~u~Bq3ZAY0}^1zzk#EoW|$Jv z5P91fJG9+F0~oNFXcA5ZNq zIoYzcXZ>}2D0Ni78K$rJ9{S)ns^xcD&W3p?-Ju&1K0uGwh~txP)Y~KceY(TfA8~@L r@clgr-@+MZ*Um5YBudPBnZV!=dJ^6yXScyWk>}g9&Tfh?rfUBJ(FFLH delta 216 zcmcbxl<~j>#t9ls2V6F4X-REvk@I1l9ON%C*&%>svyRFdmdzQOy39-r44d6`S4wYQ z;LE`{`G(&OUJ)q<1eh!sC_MRszXaC|b%@{!>B*n`&6q-@CJU;F0gcjtF&lu)2&u`d z0_>Qs%1(X-q!{HOY63NwIOHa01=?|H$UzLLn7lF2oYTP>!sD3yG0>cAfjoqFf@QK% ZkU576gr_^XFvy(K!v(@CnQRy!4*=MrG?oAW diff --git a/norch/__pycache__/tensor.cpython-38.pyc b/norch/__pycache__/tensor.cpython-38.pyc index ce4243d14b95522f7734ac1604ebfece37fd3ef4..335bdd63192cfadf55266284e06f8596395504b5 100644 GIT binary patch literal 8028 zcmc&(U5^`A8J;sUw#Q?yYTvjSb-Ydw@K;kzbG2#lY!Y@D(5_6_SzOS&@mi6wKP z6*>D(!^KQq7BExTw;MSr-nHu9ePFOQ?+d`XNX6wA$A>=!3|_+%zXK3jz;88dDMDxL zhE9|_5L>pivr%4;oVgKq<*_H1wRyCKEqgoP@MM8+kQT1r2)l7lrLJry&9o3Fsx8Af z&BZ&-eweyl+1^dvxIGNhybSNOTj6HCn0i6bZSICakorNe+mnM1`{f|GJ!p2am3%9? z+Xtq?+dus7=8f;XlNMURR;SlY(p)QOcazj>1+AT?3i=H%*lLFz83*?)bsE1^g9l&R z*y-(t8=X#fUvjN~SCL0K(yapE z*%iUB$L|BZb~2&=iUSiugIMlZlkPwWD=dsXJySe9FU5gpjt^IE4%!{L(P`hz2p9J@Q>WD*ywM{4CLhu%oAr7QsvdT>6nU<1 zRFAW#-(sqmYL(bw0@~en5(J!&zY1X4o~YQXVpUW{+0F}Jy*094w>~((=K!6f(YbrMkr>+}mui1|4HD`kkdl2! zsjx7zN02vRb$#rDZ0cDtvU$Bahburvza6)`ankIzU_(3YI01h@0+D5!Z}$6W)p>5r zr=>1f5XwOBDN>38ii{?CqYw{na-re3yWLRfl}5?*wH)0O>FG@J0vN`DiraU>&Rt_f z$tHCb;2U^ijx5aqA9k>eZv|ZP#qha|{u!QYuXK8?W+#62+DVA(u6i1{)MbL_2xjsD zhiPf~Q0wS@9ZyUSSfadC7Q@Q?RdbA*a9Nv~2uBiYEE1^rwn$)yqg;@m%w*#N+i?y` z?2!eBM|(K44QCyFL1*}f~N>bfYfdE;CiOYde6~hN|&KzgHDp24El*mOWBm>^nzX$r+0lH zNDbTU6{u9ju4r{C(~a4~YDUZH_WuDjNiER!2sQn2NmJ9GQqhW*KutUoYF199W@C96 zs@_A%d2K>dwZYz$r0RJN7YOFh(7pdqiY;dpo9#bhYH0Vg0QS*ItOx=)^Jugg6iLp2 zC=N!j3-B*&v6sC_Y&%IV%14F3lUCL%vK`ag_%u@o4y*|MUCH1=uRMZJF}5QULplp> zBaBs_7sy5h0>t?2A>yd1mZ@r^*bP4n^oho)X*Y916;8|?(IR>2BY4#<(w2_Rf;3j_ z8qE6&jl4b8>TKg2QEAKh?wY9ILr4E`nBEAa%-lZ0hOwK_g(Fzqwv$5SMa94e`Ao0G zc1&3Wnc(zokdH3d-c)CSvA$$H)nn+Z3k2E+&-8fe5w$Wmib*dpK6XKi>$b6L^t6jq zu#?_hOvhya3tp>YuW14difsqL|0Pa5LGUC%>fLO%ZhhEP##N=xu-$JI zjQNY}IZdjD$IzpVz`T)TFE-XtC^10T z0QCor1w!Dp1j*RG_!yEqkzlltVtB2o9Am5{7&2Yvn!#X#WLVaTY%hVYh-7F%A=$&P zYz}sBhN{-vsHRF*bu1GxdMm>2dnRBaWeyotaN1D4syo(xA3Ibpuk|6jT2z2UY zG3~N1fVJ>o_HZRrmP1?T_)BUtTwoj1S{#n{q_a)BTdU9+>5@_i*euDE!+D9m9%(N5 z@;MT1v^Ltp?{a1ZH1#%HL_7IVL3tHrB>o!!l`ldzf6jtRZAz$3VpP*vN8|@pdKzlv z;u7Z-#EHmJEX0IRL9R;otc8o5+chX!?f~M*^HB8MEmx_L!^|He;FjJhet_9+@Q*YG z;D5*99|_`@MdBMVCZ)j;XA!*YJJL5NnJ3Xetef|>)X%4P9KwGav*VmB9iY~rmA<^G zl|HrZ{~Mp7eWj-LKNbE1QM*que<0o;+82hBo**_sV-NfNP;L)Kyg~b&s2yf0q zXk4`C^vJZZ4fUz`;PN1tMXE=yprsgk##%l5{;>4Hm71~n-@G_nzK9=wIFI_~15t;R z&%z8QJ(<`milJgCW(58R?CAuk5Cr-)evg^VFO*N0@e8VW|EW{VHqD9ooEhV-%%mwb zCIS%ZZC>&f29)xVh?v2zCXg&`zI28I z2b{GqURew?nK>qs5s^=WmE;g;>+~yNIB4n6&a*k&rF$PSQBWN@#_(M=OwP>RPV-KvtzC?AQ%wgVwXL42HKlAPuz2-llYB76Su}Dp z`{rZ&38B(LfJx3bO5c+5+Y{$UUwF@CJZY{6!>iN_yzE7S1@=HBcvJ8vti=pPkoQ28 z&%^dHR?2fkhsWM5Zpsdw%(f8_k+lC8wqu48=`0xPiJKns8}wV8miBlxIm1E+Ad?)pw~rC6 zLNT4q;cJg10bwU(q4w&B?Knil(n;F=PP^5lL=2yQuNx2g{hmU>qt@$(a4oXN`qQ&< zU2FB_uh%a9Z0}Ny^X$~u)H7g9>abj#)5T+xmFt4Tg6b}WI68=_;dEm;Mm+;{X1vV_ zqDJe{Sj%3WZ+>WRTD*Zb)KOHLhv`88mb!fDNg3o_95&WK+l8Yhf%gXgz^<6p5qi0a zAy|W9I0sM1`1h!pa1WaZjbu)6Y5f_|Xu)MLJgLWKegskQ`~nt1q7=+PRMI{)KlPz3OljQu4&1d(?}#)5knd)EDvfe} z^FBpjSpPv2q1ZQ#HRukCUz1YK?9|Cj z?SM!u+~~ZS#TwY?tW;wss4GT!*zn1$8GD=`)4~@`R+hL^2C)ud>R*lyc;C_?VgH}32jw!yn8Mh|b)@n6o9szhuo^r$6()b7Qp3r*Qc3(cFc z?&MKvmu-_(H{9#5zx7t?M!j}d)p+WQ1nUG}BKR`FR|xngi7tdF7RJ@91aA}kh~Ngn zdjvGw>H`A)VW9$oCc#aD7QxJr>mFIHJtIX;yBYiyBc8DROWv|q#kb=5{u$5l%HDav zGnlQR<#*Lv)+dqhX{DCH&{$~>k{&BM*R}>-9nQs#;=73&w32}e^`A$^5g+p&srQQc SNBC7x--^j41ijGO75_i*@4ZI= delta 2685 zcmai0TWl0n7@jki-RX39_TKJpOD|lzg47}gA_%om@PZZOB3U&jYtOW7-R^GA%wjN0 zmZXg_K9JxczDe2`V@P~a;)~DXHNlve7$*|ZHy(U4#)K#T|4dh9J? zX^|AWS%`%fm12aMtP5zAS*#nV%_1xcG{$Te0~%*>mH?VyNtObdWNDTGnqoby7igMg zSq^B1~%Wh7b>S8Ou{veu^V7!7ac(cM?cRH4TY&s0jD=TvzFvgO!0U#mFP zqQy#HiB5ax>JE>ISMMx9rnZAaZ$w&fl0T<3%aknPGyE; zkpn$xCYTzOB~i_!v|KHFl;WNH2!NuR#8%TJP2wb?hKMDO897lmejq(!JTyYGq8u6^ zhIl9RR7@H}O+Es!L;MtaZWLt&oe!dP6T)VMw(l)?D|LE-8xt>t2gRRV7dQ97T_GA2 zBZh0T2r^epkf#v(#Ilv{??-kW!T`d(@}OJsqxD#}BT1~ynNr=E5V`KltqkaZBU#EE z>gN$;6WZQ_^5TO&fMSF^u40AG1X;DnZ_E=a!oD+9Kyc^pS`9RyHoL8H;Q^VJOiRnbWO~h~Uty^TY z1}{}-9a(i>b3N`GUZY-dd;|D$yw#` z2$5bC7gD3I*}>FoMOPiIZ&%wYF-C*!)X{cdpRU!he7IV<*Mf!GnHD&3$~dIt;0h{mE`F4i5tDW@%wiK-zxs@-L*~b|4FdU4$)QT$sQURL&Lbu_|pg-tDyu{XO_}$ z6yfCdj-ZcV$paSGYc8eVE7!iwr3e`nH}g-na{l?;C+pLt{l-t1GF3- z^hb7Tbp`7oFiihY#Mk{FwnibmLgXQP#%;_$EQ(c|Urxqao;9%a7nTH%0)-@DY;9HC zT-UWheaku_Uvk@FNkN_BX$-25cLCw}sl~fRbYQTBy(;Pl4jwt;>vOeoRi5R12p8bH zhZhj=oyWH$j3Z1SJcIBo0*)g;hVT-?afFiyrx0F7XctU!^3{f2v0DJ{MgVxBk^$8k zG!v$6T4uz|n;P&5a|kG2X)_@=9vq%>zD|QuvFo_U8`EBcJMtSY*qLiM^Ub~UHP)y& P`v^M30>g)&nza4}ev|q) diff --git a/norch/autograd/__pycache__/functions.cpython-38.pyc b/norch/autograd/__pycache__/functions.cpython-38.pyc new file mode 100644 index 0000000000000000000000000000000000000000..35e8a00d98dfad81f7ab04203f38f7c8e170334a GIT binary patch literal 594 zcmZuty-ve05VoC^hL$oQMqVIGqwWZmKe2VGSh83lcG`%IToiJM@*j@(659 z+&NS#5+{A1@6LC=&v!o>4G7r8Oo|UUKQY+10D>vNo&enuK?HpVys?xBHYahAL;jC3 zh^7F01|&&J1xcBpg1wV85P=Bc3`Gxm58Uulh^s7re9nx>an#Px$A_)~_84eGH*CQ+ zgoA`R+0t1|UB8jaYGdQTg;rLDqc4uGx5V_J;rljO>RKjMrIH(+x3yG_wuzCsky#`8 zYTam)FOx#+M0qc38@(`@NQzeFO{tX~uhuT&yi{ewIm)p}gub5c=^pP2a({^hVCWsize; i++) { + result_data[i] = 1.0; + } +} + +void zeros_like_tensor_cpu(Tensor* tensor, float* result_data) { + + for (int i = 0; i < tensor->size; i++) { + result_data[i] = 0.0; + } +} diff --git a/norch/csrc/cpu.h b/norch/csrc/cpu.h index 449a3f0..a8c5ac0 100644 --- a/norch/csrc/cpu.h +++ b/norch/csrc/cpu.h @@ -10,5 +10,7 @@ void elementwise_mul_tensor_cpu(Tensor* tensor1, Tensor* tensor2, float* result_ void 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); +void ones_like_tensor_cpu(Tensor* tensor, float* result_data); +void zeros_like_tensor_cpu(Tensor* tensor, float* result_data); #endif /* CPU_H */ diff --git a/norch/csrc/cuda.cu b/norch/csrc/cuda.cu index ad4a79b..52a8ed5 100644 --- a/norch/csrc/cuda.cu +++ b/norch/csrc/cuda.cu @@ -223,4 +223,48 @@ __host__ void pow_tensor_cuda(Tensor* tensor, float power, float* result_data) { cudaDeviceSynchronize(); } +__global__ void ones_like_tensor_cuda_kernel(float* data, float* result_data, int size) { + + int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i < size) { + result_data[i] = 1.0; + } +} + +__host__ void ones_like_tensor_cuda(Tensor* tensor, float* result_data) { + + int number_of_blocks = (tensor->size + THREADS_PER_BLOCK - 1) / THREADS_PER_BLOCK; + ones_like_tensor_cuda_kernel<<>>(tensor->data, result_data, tensor->size); + + cudaError_t error = cudaGetLastError(); + if (error != cudaSuccess) { + printf("CUDA error: %s\n", cudaGetErrorString(error)); + exit(-1); + } + + cudaDeviceSynchronize(); +} + +__global__ void zeros_like_tensor_cuda_kernel(float* data, float* result_data, int size) { + + int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i < size) { + result_data[i] = 0.0; + } +} + +__host__ void zeros_like_tensor_cuda(Tensor* tensor, float* result_data) { + + int number_of_blocks = (tensor->size + THREADS_PER_BLOCK - 1) / THREADS_PER_BLOCK; + zeros_like_tensor_cuda_kernel<<>>(tensor->data, result_data, tensor->size); + + cudaError_t error = cudaGetLastError(); + if (error != cudaSuccess) { + printf("CUDA error: %s\n", cudaGetErrorString(error)); + exit(-1); + } + + cudaDeviceSynchronize(); +} + diff --git a/norch/csrc/cuda.h b/norch/csrc/cuda.h index b18693c..799c056 100644 --- a/norch/csrc/cuda.h +++ b/norch/csrc/cuda.h @@ -25,4 +25,11 @@ __global__ void pow_tensor_cuda_kernel(float* data, float power, float* result_data, int size); __host__ void pow_tensor_cuda(Tensor* tensor, float power, float* result_data); + __global__ void ones_like_tensor_cuda_kernel(float* data, float* result_data, int size); + __host__ void ones_like_tensor_cuda(Tensor* tensor, float* result_data); + + __global__ void zeros_like_tensor_cuda_kernel(float* data, float* result_data, int size); + __host__ void zeros_like_tensor_cuda(Tensor* tensor, float* result_data); + + #endif /* CUDA_KERNEL_H_ */ diff --git a/norch/csrc/tensor.cpp b/norch/csrc/tensor.cpp index 3571bba..782d365 100644 --- a/norch/csrc/tensor.cpp +++ b/norch/csrc/tensor.cpp @@ -461,4 +461,78 @@ extern "C" { stride *= new_shape[i]; } } -} \ No newline at end of file +} + + Tensor* ones_like_tensor(Tensor* tensor) { + char* device = (char*)malloc(strlen(tensor->device) + 1); + if (device != NULL) { + strcpy(device, tensor->device); + } else { + fprintf(stderr, "Memory allocation failed\n"); + exit(-1); + } + int ndim = tensor->ndim; + int* shape = (int*)malloc(ndim * sizeof(int)); + if (shape == NULL) { + fprintf(stderr, "Memory allocation failed\n"); + exit(1); + } + + for (int i = 0; i < ndim; i++) { + shape[i] = tensor->shape[i]; + } + + if (strcmp(tensor->device, "cuda") == 0) { + + float* result_data; + cudaMalloc((void **)&result_data, tensor->size * sizeof(float)); + ones_like_tensor_cuda(tensor, result_data); + return create_tensor(result_data, shape, ndim, device); + } + else { + float* result_data = (float*)malloc(tensor->size * sizeof(float)); + if (result_data == NULL) { + fprintf(stderr, "Memory allocation failed\n"); + exit(1); + } + ones_like_tensor_cpu(tensor, result_data); + return create_tensor(result_data, shape, ndim, device); + } + } + + Tensor* zeros_like_tensor(Tensor* tensor) { + char* device = (char*)malloc(strlen(tensor->device) + 1); + if (device != NULL) { + strcpy(device, tensor->device); + } else { + fprintf(stderr, "Memory allocation failed\n"); + exit(-1); + } + int ndim = tensor->ndim; + int* shape = (int*)malloc(ndim * sizeof(int)); + if (shape == NULL) { + fprintf(stderr, "Memory allocation failed\n"); + exit(1); + } + + for (int i = 0; i < ndim; i++) { + shape[i] = tensor->shape[i]; + } + + if (strcmp(tensor->device, "cuda") == 0) { + + float* result_data; + cudaMalloc((void **)&result_data, tensor->size * sizeof(float)); + zeros_like_tensor_cuda(tensor, result_data); + return create_tensor(result_data, shape, ndim, device); + } + else { + float* result_data = (float*)malloc(tensor->size * sizeof(float)); + if (result_data == NULL) { + fprintf(stderr, "Memory allocation failed\n"); + exit(1); + } + zeros_like_tensor_cpu(tensor, result_data); + return create_tensor(result_data, shape, ndim, device); + } + } \ No newline at end of file diff --git a/norch/csrc/tensor.h b/norch/csrc/tensor.h index 7a8798b..cc03666 100644 --- a/norch/csrc/tensor.h +++ b/norch/csrc/tensor.h @@ -24,6 +24,8 @@ extern "C" { Tensor* matmul_tensor(Tensor* tensor1, Tensor* tensor2); Tensor* pow_tensor(Tensor* tensor, float power); void to_device(Tensor* tensor, char* device); + Tensor* ones_like_tensor(Tensor* tensor); + Tensor* zeros_like_tensor(Tensor* tensor); } #endif /* TENSOR_H */ diff --git a/norch/tensor.py b/norch/tensor.py index 022c343..265cb6b 100644 --- a/norch/tensor.py +++ b/norch/tensor.py @@ -1,5 +1,6 @@ import ctypes import os +from .autograd.functions import * class CTensor(ctypes.Structure): _fields_ = [ @@ -15,7 +16,7 @@ class Tensor: os.path.abspath(os.curdir) _C = ctypes.CDLL(os.path.join(os.path.abspath(os.curdir), "build/libtensor.so")) - def __init__(self, data=None, device="cpu"): + def __init__(self, data=None, device="cpu", requires_grad=False): if data != None: data, shape = self.flatten(data) @@ -29,6 +30,10 @@ class Tensor: self.ndim = len(shape) self.device = device + self.requires_grad = requires_grad + self.grad = None + self.grad_fn = None + Tensor._C.create_tensor.argtypes = [ctypes.POINTER(ctypes.c_float), ctypes.POINTER(ctypes.c_int), ctypes.c_int, ctypes.c_char_p] Tensor._C.create_tensor.restype = ctypes.POINTER(CTensor) @@ -45,6 +50,10 @@ class Tensor: self.shape = None, self.ndim = None, self.device = device + self.requires_grad = None + self.grad = None + self.grad_fn = None + def flatten(self, nested_list): def flatten_recursively(nested_list): @@ -63,6 +72,38 @@ class Tensor: flat_data, shape = flatten_recursively(nested_list) return flat_data, shape + def ones_like(self): + + Tensor._C.ones_like_tensor.argtypes = [ctypes.POINTER(CTensor)] + Tensor._C.ones_like_tensor.restype = ctypes.POINTER(CTensor) + Tensor._C.ones_like_tensor(self.tensor) + + result_tensor_ptr = Tensor._C.ones_like_tensor(self.tensor) + + result_data = Tensor() + result_data.tensor = result_tensor_ptr + result_data.shape = self.shape.copy() + result_data.ndim = self.ndim + result_data.device = self.device + + return result_data + + def zeros_like(self): + + Tensor._C.zeros_like_tensor.argtypes = [ctypes.POINTER(CTensor)] + Tensor._C.zeros_like_tensor.restype = ctypes.POINTER(CTensor) + Tensor._C.zeros_like_tensor(self.tensor) + + result_tensor_ptr = Tensor._C.ones_like_tensor(self.tensor) + + result_data = Tensor() + result_data.tensor = result_tensor_ptr + result_data.shape = self.shape.copy() + result_data.ndim = self.ndim + result_data.device = self.device + + return result_data + def reshape(self, new_shape): new_shape_ctype = (ctypes.c_int * len(new_shape))(*new_shape) @@ -86,6 +127,30 @@ class Tensor: Tensor._C.to_device(self.tensor, self.device_ctype) return self + + def backward(self, gradient=None): + if not self.requires_grad: + return + + if gradient is None: + gradient = self.ones_like() + + if self.grad is None: + self.grad = gradient + else: + self.grad += gradient + + if self.grad_fn is not None: + grads = self.grad_fn.backward(gradient) + if len(grads) == 1: + self.grad = grads[0] + else: + for tensor, grad in zip(self.grad_fn.tensors, grads): + tensor.backward(grad) + + + def zero_grad(self): + self.grad = None def __getitem__(self, indices): if len(indices) != self.ndim: @@ -122,7 +187,7 @@ class Tensor: index = [0] * self.ndim result = "tensor([" result += print_recursively(self, 0, index) - result += f"""], device="{self.device}")""" + result += f"""], device="{self.device}", requires_grad={self.requires_grad})""" return result def __repr__(self): @@ -142,6 +207,10 @@ class Tensor: result_data.shape = self.shape.copy() result_data.ndim = self.ndim result_data.device = self.device + + result_data.requires_grad = self.requires_grad or other.requires_grad + if result_data.requires_grad: + result_data.grad_fn = AddBackward(self, other) return result_data @@ -252,4 +321,6 @@ class Tensor: result_data.ndim = 1 result_data.device = self.device + + return result_data \ No newline at end of file diff --git a/test.py b/test.py index b4c088c..0dc6a80 100644 --- a/test.py +++ b/test.py @@ -16,8 +16,8 @@ def matrix_sum(matrix1, matrix2): if __name__ == "__main__": import norch - a = norch.Tensor([[1, 2, 3], [1, 2, 3], [1, 2, 3]])#.to("cuda") - b = norch.Tensor([[1, 2, 3], [1, 2, 3], [1, 2, 3]]) + a = norch.Tensor([[1, 2, 3], [1, 2, 3], [1, 2, 3]], requires_grad=True)#.to("cuda") + b = norch.Tensor([[1, 2, 3], [1, 2, 3], [1, 2, 3]], requires_grad=True) import time import random import numpy as np @@ -29,8 +29,10 @@ if __name__ == "__main__": #d = b-c - b = a.sum() - print(b) + c = (a + b) + c.backward() + print(a.grad) + print(b.grad) #print(a ** 2) """#print(a)