From 18c50dbeb88b30455094307c6cb7c3c474808a97 Mon Sep 17 00:00:00 2001 From: lucasdelimanogueira Date: Fri, 3 May 2024 13:34:43 -0300 Subject: [PATCH] Small fix shape of reshape --- build/libtensor.so | Bin 90336 -> 90336 bytes build/tensor.o | Bin 27232 -> 27496 bytes norch/csrc/tensor.cpp | 21 ++++++++++++++++----- test.py | 40 ++++++++++------------------------------ 4 files changed, 26 insertions(+), 35 deletions(-) diff --git a/build/libtensor.so b/build/libtensor.so index 6912a056ea71261f54a1ff237e8db51a48743418..e6e11d3e7979d8148970629574aacdc8aea4f7d0 100755 GIT binary patch delta 7303 zcmaJ`2~<>9wypc52$LuZ6o@hki312IM58FEC52+QCPHFF1vO4l)E01visFbRC|>*D z^zTf5{dP+!l3-|npWCByrQvr(;_J^z`)OL@VoG_cG*_ssM3uFCx-{7B zUpvXH>6y|Vq||@LC0Dsi@41-#*Kd;E8A|KFRMs&X;@U4}>pbbrkYK3nAags;k%m>u zZc!awqibAE{uwi*xv%2vsw%ug55v16PkKX@H^^0THqvigleLH<+S_fE7MDT);1-3s z?Q}EwS7pnLV3m;z7s-~jhQ_)_ zk&86T-Q?eQjSPxY$ysVFKhqNTez679q}e{|cNInFwbK5UDs+U=UC1t$>3Ww%rS_)Slw-n|cDoooq!(w#`e*tNhLi3}TDvQ;8++t0PwiMRc zKEI~xmbF!cf1Y_Ap9{?`mbEKe3^{{k*tlq5i>2@ik68DWu3OFf3eA_T=F671EoWOa zZjO|RHYC~%xf!jI7HbFNpT4DDulo2^n+*or4>xrEjE*1MB`{d5#(Bay=#Si4#hrYB zwmopxnd^CfR&)Je^W~NOfzT=PUeArF&%Tjz-n$B zY~Hsr9M~qy#^Ye?g6ve%Q4)_Rqq4f1txvAVoy^{RAYrf=xY0V;7POp@DP4?|G z2E%Gys>J^@!yTNs(0ms&ylcqi1!=eV0K@u>L;>kYk~1123P`*A8{P)GbE>cDb5yX+ zVjcTvhNZuUSnV96?cJ+1WT`jRlL9G;XMo!z-Fp8Le&b*HN{#=8!Ik^2Hr3aACNHZ^-PFm(PGWD@r1`9Z_R0Qh9vWv+f3x>bzWE4jRZ)$Vh(a5KV23gd?GmB#Ho})tFnB3 z*Mqa&K*p8b$@!S_SWqxq~OWed0#2H(Q zLu3(>jzh;_ORf`Dn)zKS&t%-CH|uH1maa5=eit&Bu9@F=*bT`5Sy9^WMHOaT;=UJE zJjJg!RPhKhk9ZXcptV==Sv|coKe)$F++^iR{G{hq6d_4gF`f3x&d}m(Y1R@GcI5ln z&ui-&Xnl5ALgyMB(kE+mz0myC+Ld1!a$c~Szk)zLuSl?*8(?S2L9E;?%b{g!s;5tv zbbsz4B5n6;aR)}xWeWnG9)OHSFEP>51%3hd&f^zI?qy4nxut^NJ}3Bpyn*KP^n>gO zl1M!k29r)Sa^XNN{X8#EOAwxE{|gBmu)fx`%h*i(cUYCB@FGur>m2{3)Qn$Bx3hF; zUxS82^)}QXtN&Z9z4{$Ulhyw>@A;Qn>awW2j|nnAcM5#WiQRcu&T!Jm?ANs3=R`SP z=XjyovRW~&2WJs-5KIn&r;sHD%VX?Mv5?eV~{IrL9FcRci^<`JulYeQ!ZFrO_wa5q}{F(5{*@2-~*|5 zPI$uqAjLKxsSXOu!66lL``HyDs_Y7VQIc)rInl;muC^NtXd{x==Pc1yoS{j%5!$V4 zIw!a9kSmq&dm+_9ArKr=!9U8b@CjmN4~9HN(^)Z~L5Q>sID<`_PBLib5ea6z_rbp(RV7*XS7` zF#IIke?_K)z*SI4fu`Pefenb20$Wg#?Ym0Rf-7&w=`>pCLuV|Tqy2b-&woKW2yox{ z!7z9U&m+O@o;+Ltw5*#V4^|!t6cY1bF;Cu zhJIj8)GQ~dE#E?X>5nVoqrW~YGeW$9YtMIX^{}4{Pf&tm=eLjJ2=}80S9jBzD(Fue z-*kx<`Nk0X&y`8q!3sKbl|_5Eif&pJM`qH>RsOy+(4-^HxN^`0&;^BEh!+jWFL0f5 zMiiJ|%?l((`dxmC_Gu-JT^$#_qf!=lqqAdyC(7G1?>YLc!0IEkd;S{N?@kNqW*h|Q zJ^h&IeFAM+6Ccn+xaw^EIyq9i)0l#;+QTwC8R<`#6|8lAcv5ghpAlSpjtH(9blcka zfJou0v*mYmyosOOX=|S7OkR!Zdo7K zub_gv@OOy2BUff=`<7GbU8L?hM3)y7xK72+!b-0`DJuN?Azq<&p`4E27#ID=a!Gpz ze{Z#2>OUW9PrLDeytHkv#A#+pK6XrC?d(T_`|*e!fhYCD~(kK42F+Y9z+ZS!7H z{gcCjHR`CSKIbD*{aj(hDs~Xn_g9x=@yByp{$3xo=lA*0UcGZX+}j%HnYT_eqqwrUHX1GDD#^BLbJAp_;ry^ zcj?@robN~{kibNX>hDgyu!*tdq`1 z%6VKmYoxPMIyWk3zH}au&i&Gvr<|jN)AVQQ+##Jkl(QrKX`BCulbDF@#%>I{(D;x~ zr^uM_gkSM5mpW-x!Ojw?uJ6zX4#2(r*QrGK&yG9`ySZ&dD7o)*2G)3260Gtlu0cKCqbtl}$ zr!sWCC-BB(6bk%x3WR{3nb?oOL0QlNj-RIMmw|5rAGqKVQVBkI^xuB zZvm~@$iq=t1MHiD54{#55jYW;13bM5iNGI#r+_mTBN6x~pp%Ee@FWL`z{p&T6!^&! zbP5bzije~IfQ^{sq4LP?H%kDx&BN~G zk8z!nS#bzbR)-KLZT5J!Aq1Q=US=P0oc0R47y^oI6WEXZx;Tw>2?Z9Uvv`hqud>%e z;a~F_Th6a@C-PbXwFMJdB}ZJ};JA1qdm2h2w3^p>ecf1M7;z$3nK=xL=9|H?!w?gZ z!HUBWvvjhIX+(@xG?`t2k$labgkcm9SdX4aNzP=WdLkt≶c2?_*Zf6BFt=jUB+H zwq!cH0yimPzaX0YmGuZm?sb+Bj$HqlT&^2i9S-_lGuei4WGtV>4upfSXcoK3uN!9b zAYbMYL0rg}%qIfu>2I?52*eb>$Z?i#zNl+L2>yVB(WVm_L8*J`iGAcG^mZbfP?`iyesQiwvOB$8E?;3V< zFo_~RGp`}UL~gTQLr5~Yz!nW5?cZSB!2QL=9N(6>;lSbcW9&0d@+-TIY=6f`(_2+k zDT_%U$=w{E_xs@V;Ur@=vF!;YD%SBeNuKgMzR+OGUXai3Cy-Iv^i1~BP!csQE<;Yh z{=_^QhZN^;_Xt zImb6C`O5D2Xr3wiQ?_UrnW|;2VV@7fx?E!l?HmsCu(|B>;iNJ)qDZp#mtb(keZ!^wH%o*4 pMM<)`@&pZ3OyzlpYd#xq2cT2Ut zrCMLQ$%Ulr#D^UCXt2`le+nxyj&BcGvdj6-@q5J!)f)3pGY;KTUeuFh8sY34|2vV< zi=1@2E;Dd4}=1OmX=Jj=voFnLWE+dt93h6VhiOMI_>07R$nA=CL z2A{Ajnc=5pG+gt<g_2mB z=4SBOk}rdzG}1I}Ea|k&Eo{((G-b;rLVAAZh+DyyW(>d;;J~x(2`!bFDu68SC(*y2?9y`jU*IsWoDDx5h zl1Xp6hmxVRb4Nq^EbYEzXbQKqbybun4fZwAQ)|cH+*)#wCi(bvY5!lXluXUOErY(@ zF?7_K!q#+av<8pOkamN%5ZyKTdv8m-eIOUK1<{KV4}(w6TWs}?>~4t3s^3G?rg@tijXIr;e5eZg1p!T)Ohw(qDga;yQ}1)rlMHCQBc`(11j8LTLl$aXs5VvQ&H3{6u~PkZ>Do+^oVTH3ZHzHhb6U$ zP}xfY@78LBX`IcPRAI8_R`JB<>-1czSGOR}Ycbx(Ebi-ac|G(N?L5V6LdAO=jd&r*Jwman8)=iVA1Zvf@?G@ z#ZM`@N=s9Gydyck{hs8eh;eWGKVC`q`UF(+5OvCVUlp?Yg~I8%l{@RWQ;ca3?L1?m z%hxdSQa+|{&It4x(IiT5`b+7(nyG(+mr{F8%)o;)wwl3;n_Ri?r!RoK0hLMeBfR?7 zGw9kxPsDo}#v-Oo&xR|a=hK&wrS&}Ris*SsljylO=X2=U@}21UU_6wsqb}KVjc7ri zBAlJqa_4^T6g|J&L}w;Xa6JMmZ@RDvDt8Zqu=^gG`(d6lux5o+4OGFf!-*bowK;}y{ArdBaKU&=~@CSZ+ccE-Idm7qW7&f>ugq(aW}_myomi;-J)<)6R)||L=tcCx&pI@LwYk3 z~Y-P_w6A#Dco-J!hnU!0WD-_YF5FO-rP8a}63@%@@s&FMz!4_%rw zP*HpvH_nSwiq6m%=JgMM@k_yO6E#~rO-1_AgL#kGT-pXZcm@&LY!knv@6Gdj?)7?} z#GAz7U^$KBLqXE#*n2u#Q=O>F$oFG=`+i(lPw&kOBLARWvVuos%WO!BS++dA#C8`I zlthaPeuc!Zd{SpYX|Eu?p5|rw_0H!-@dNZ}1$~etD|krDvr?40({$W&16E&nc7MhD zOFA_>F!rNb9H3*tW-&Gw=wlN86O%0Z)Aoho1}$4gJuO+%tHTCF z+6wBh)4I|7*+bf|1sSbgY@iAAyDO_}>5TaXC9{@BWcMTI==<~i$TnIze~5DX6fZ_O zdy0F~qd~Z~D%abZbP>y>3Vl?WEN6MA?x*eXa`rKauLth8yaV z3QqXXUf$!Qh?NS*dGp;*i8f+*OHnYOjSxCDcd^p3iq2ftPsy#JPjZ7_egi4C?MQV{ z*oPFUuq)KAFdeZ{A&qxlfe>uj)|=yYm1xHGglOhLrI^>$lXOQ;-`K8|$SMYdgVq+X zNv(H7>{=rcE49Y)8SJbU^ZJtO)}GK-Y!BM9WB~bzK3Lpcxp-VC4XdUf%=eAmQh~Td zpm0zs0GCu+8EjVyJ8oBsf}pLmiYqxCDcAnm8IpP z*YA#tX;yuK}vmy&@4?UEM}DWEQfByf0~Va}yz6d1CyG&lft#q1DvK97P_}@n#>dM`+2B<~y2a4kg>^Zu4rFf1MBno;ZLCw~vUnzo!dUM~Ck{ z&Ruo3uuhKDTaUD*-f-AXO-@nQf&!Nf6@s}MdkM^gKNHN4sQ;Siu*zeS+1$~Q`QG8S z%uRINnju7`C)U(EUoRJ=XKCr$cxBijdSR_e8C_0$uZs%nDqQ$=;O0pB;9y(Q>vYvR zlM;D=x)rQ;xsUS=_r5b0DoR~^fS2kzLm2C9zqrt^@dlxb={vzJ*_nt*IyOZc81z3)DSnY>}7JbC)Q;_#n^M zpZH9aum8I!|33#L>6&(;{8QTHxOR>j<-6Ky?AXc6cdCP{w4P2_VQ8nf9Hves0}YK{ z$S?BxhJ!WAQ)6wS=8^zIt@Kq&UkA{C zrLULfdq^KZZ}}9W=}39W=kh-w-UrgJk&Ly;9}_nqTIVFct%_Wq*q4TJ#(EFpBiCfAA~(d1>`p z#4X}4&C6;7j?)ZnT$O)ANKx#>A8`(Eq|Id$+|GRr3y#bBno=6QG2X?9Mk9W?v+|Qe z`5W^)($c`vgWD!8GvA6=)ty*D|4C5QGyM5OA{c-XCRKIT>vVg71AqfZfe!d7a5m6w zw5o0Z4jzMiV9QwK11nxaJ}_pys=8y*ZwCeg{|OumjGmyXbAiKOR@Fi$omsaTU-rXr z7vAup~)UbAZdIsOl!*mSk0}03MsFsx83l(^R!RcJH$(s@fMg zZUzbk&UzI>z`RtfM&Qnw&;eGZsp>`GRp3KsGkyi6t7G zELHskm;ci9xw;E6&N)ac|bGp7VybDE`Aj5x5mN7Wn*PbP6m376N|& zHeix|M?-p%{z`B0G}Y+gxlT#w=+y3K;@-ho*9(0|;_*@>ThNF2l4n^NhX8hfgFidR zA&)~({Y8XbkzXI;YZQy{CBE^?Uqli7C-zg~!2c5Cv{-LcRq@#Csejv3|0MW5Ss7@Q zf0)=o=#qZy5-v#=d&I4qM=@_d(nmQtibZqWHHJ<0BLT|IG3-@oa zVIteb@!ZSotUoB!SJ*%K^&gX1mjK{{Bo@u_`^jut0Q~o-ux0%ERx+=pk8&fKRdU4j z4URXH*`Me~iJZpk>%rmzNqaJi83S=E)=p}|>Ie4PIV`##V*Z%J zCilbLnSFF=Kl0cK4-eT}A!L#Aue;E`>65Ek!Z0TSUrsOYYy9SfQ9-Fmi zu>G(=_-7avj)JGGRm+*{5b`RScrhw=3mj3tQ=bxK6Ie6S}rSdwNeH2A}*_r1_ zevg(iY3m|E_cE5aqU)r4SjN(0NHK|J4`N7|kK+U0CWP?^yvgkOp~Sa~{nbK8939;J z_;ogCC`lZ&Ekn{a^EHjnPR;N5>Y+iiZ^>sjhmuf|&N>ex29nGM4kIJUP_|$gY5O+d z2F@EU=J@8p6~_p-2eEUUWEQ)FY#+zRyd_$cKZ}SZBYQYLJNLqw!b!^4vyHJNbdcjK zheev-@#O&1>2`c?K{dPMOP~hL?)c8Zt&`+& zd_l$knCJ2%)(KrUC7#4b%w8so=&lv9TMKf0UvojTPiOV<*lRn=*pqk?s@%DBNg+Gh|391{L`whw diff --git a/build/tensor.o b/build/tensor.o index 6657ec4f6ab10cbdf09b24d52131444bf996f433..a4e4a4102036d1b18de9d7a98344aacbfddfbec9 100644 GIT binary patch delta 2193 zcmZ{kZD>)6Lv%6{Mv4)u!aj8E=7)e|Sf+^Gn0rpnz3kcDdf;$- zfB*A8=YO8(rVDd&<(B+eNU5Z9>%VGp$IZ;>jG1|VIx{c@j{{p@=rXg$jG5iD0iO$I z=9ZZ$R3sc(V=80J!J#KUY%3T=L5f|*-00@4Q8Z7Q6HguI5M9%zT+R*@;a=_9NKwA~ zeZeTns1%H{aFx4^snI~zDF4e<;VvmVP-?vJUB)PVYLw->H+Yt8Px}s|1gCwQ;PP{< z<~7bhtMfxFxCT_O2jT&jJOi)$_rP4xo1gZ-s^m+}Oo4KfEB{wOQ}Sx4wN{-_@N~nY zo=|o*)UGFqv5I9xazaoXz3uo2arK14X2-u|k0&{q}z4%V>8)=TE`FOD zoZqez?~N9BCedMhh$TV}*BN1wHW1g-$n5<6gg;MsitufOAA#HulP95hi0NNjYn;h1 z!Cq9ak{xKpW_eD1fvDTb{tx;7W9}A2IQShIPCM;R!ZpGpgf|dALb#Xk6NJ$(k#sLrq8eNgo- zui$mT#w3gE3y3>9I})|H3uW4oqlD9rOv3pDuY4lOrQg+qs@EM9w{vRtcFF5#v&z?3&9++>4S$lIyg7fLs5KpbGdz1gOz;P1BQ`?lYAy_=H2Fj_B z{im7^$>(54x?7%wkJ8=Bqjs?(9a5wYSWbr(Uk7v^3@i8EfsgUA{U!JYAL*^Id@x*h NC?$4(53Zz|{syJ$u-yOv delta 2105 zcmZ`(T}TvB7~PpnG}PRzrX{3R+qT5kTuiZYZL5tbuwS-^GV^D()d-UGp=7h|kG3Ju zz95Mxe2Id{WM6vlL5s+j{wQP-iLHlNs2-$x6Y9IOGkBQT3wJKxIrrS}e)rrvsc9+s zNg7-)Cp*W=AFjzepvEFoYOHS}796K%#U4*ey&Or|dfloPm{PR^n|MsAv1v8-I%VVb zI*^d`VOB)n$-d4-0|mLYT3~i-U_9dDQ5{xqF25>h)=sF>d(EZ=6H7;5!~%)Sfm!MM zN3piHy`nCVAWw;x{^V46l8{5yf=NC7LRt}OXzO(Zh5$jzVa~YikELb;@(isiURKQ zf77kV$pJT+`(o~SBH*@)CYuyTsZ$SZSv(YSH&tyk#OyhPHMQf5(z7oJ?Z9<|wQBEJ zRQJ;&+XI-tf-y7KOTdc(p9Z`Z@SkbCT;KE(AhJq0;3hVd&)rP0lKjnzW23>E#h3@- zrRZ_Ru?@6%5c_GQS&=%)-mDDt!yE8iTcJ8+GlW4B954wLdjjnO7Be=RWnEC<>x?GoVNBegeJ) z_LNn|0XJb0TDO{_Eo_m&TKBq(#+cm$TCDRheQ8lzJAvKGtDELW0LP5C07o-l0*+?B z6FaP$0~W2zTPv#i|LQquq*Za$8mxJa=GEIqADG<@TD*FVx;B(S+a-4B+U|H^A`$<*ZxSIl5LxW9{sY!CG(KKC*`t$6?UotvgPBW%@Z>{`VA`dXt{f#3VdsBk~$QJ4_LhC2dRqLc7wInc%FK=iUH7KjWHSt zDc*a)`oW_(;FxjB9iQoFmb_8pcb1gW`xbSIK5V=pJuonVkGiP)~?L!z4G TQ@-4AXL`9_@>dp7Q#kJ*$w8zx diff --git a/norch/csrc/tensor.cpp b/norch/csrc/tensor.cpp index 532798d..f85759e 100644 --- a/norch/csrc/tensor.cpp +++ b/norch/csrc/tensor.cpp @@ -565,14 +565,25 @@ extern "C" { exit(-1); } + int ndim = new_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] = new_shape[i]; + } + // Calculate the total number of elements in the new shape - int new_size = 1; + int size = 1; for (int i = 0; i < new_ndim; i++) { - new_size *= new_shape[i]; + size *= shape[i]; } // Check if the total number of elements matches the current tensor's size - if (new_size != tensor->size) { + if (size != tensor->size) { fprintf(stderr, "Cannot reshape tensor. Total number of elements in new shape does not match the current size of the tensor.\n"); exit(1); } @@ -582,7 +593,7 @@ extern "C" { float* result_data; cudaMalloc((void **)&result_data, tensor->size * sizeof(float)); assign_tensor_cuda(tensor, result_data); - return create_tensor(result_data, new_shape, new_ndim, device); + return create_tensor(result_data, shape, ndim, device); } else { float* result_data = (float*)malloc(tensor->size * sizeof(float)); @@ -591,7 +602,7 @@ extern "C" { exit(1); } assign_tensor_cpu(tensor, result_data); - return create_tensor(result_data, new_shape, new_ndim, device); + return create_tensor(result_data, shape, ndim, device); } } diff --git a/test.py b/test.py index 093d5aa..e86e5c2 100644 --- a/test.py +++ b/test.py @@ -28,39 +28,19 @@ if __name__ == "__main__": [[7.890, 8.901], [9.012, 1.234], [2.345, 3.456]] ], requires_grad=True) - b = norch.Tensor([[ - [1.234, 2.123, 1.5], - [5.678, 6.789, 1.293], - [3.635, 4.456, 1.0202], - [7.890, 8.901, 1.91], - ],[ - [1.234, 2.123, 1.5], - [5.678, 6.789, 1.293], - [3.635, 4.456, 1.0202], - [7.890, 8.901, 1.91], - ],[ - [1.234, 2.123, 1.5], - [5.678, 6.789, 1.293], - [3.635, 4.456, 1.0202], - [7.890, 8.901, 1.91], - ],[ - [1.234, 2.123, 1.5], - [5.678, 6.789, 1.293], - [3.635, 4.456, 1.0202], - [7.890, 8.901, 1.91], - ],[ - [1.234, 2.123, 1.5], - [5.678, 6.789, 1.293], - [3.635, 4.456, 1.0202], - [7.890, 8.901, 5.91], - ]]) - b = norch.Tensor([ - [1.234, 2.123, 1.5]]) + [1.234, 2.123, 1.5], + [5.678, 6.789, 1.293], + [3.635, 4.456, 1.0202], + [7.890, 8.901, 1.91], + ]) + + #b = norch.Tensor([ + # [1.234, 2.123, 1.5]]) #print(a.shape) - b = a.reshape([2,3,5]) - print(b.shape) + c = a.reshape([2,3,5]) + print(b @ c) #c = result.sum() #c.backward() #print(a.grad)