From 087422a8b002146010deef10731ea572cfa470d6 Mon Sep 17 00:00:00 2001 From: lucasdelimanogueira Date: Sun, 28 Apr 2024 18:43:07 -0300 Subject: [PATCH] Sub, matmul, mul, reshape --- build/libtensor.so | Bin 16600 -> 20824 bytes csrc/tensor.cpp | 161 +++++++++++++++++++++++- csrc/tensor.h | 2 + csrc/tensor.o | Bin 9192 -> 14488 bytes norch/__pycache__/tensor.cpython-38.pyc | Bin 3092 -> 5223 bytes norch/tensor.py | 119 +++++++++++++++--- 6 files changed, 260 insertions(+), 22 deletions(-) diff --git a/build/libtensor.so b/build/libtensor.so index d7ec51345feccacd17e5d7e8494a9c0765d9f534..7d6f26406a4a1a948a1cbb3bb44b490d0f4b98dc 100755 GIT binary patch literal 20824 zcmeHPe{fvIeczL8kd3hf7^;9V7Y8R0@KrXLAX8z<_A}>Tfv~OM@T1rB-O0L{)17)h zu0(Dd{v4o1=YN#DNLg#46fD%fJSk- zI0gUD6X#0aX+_ec()N=TilL}cb+R#?in{KGulc6jfx-__U~u)}kN5|* zv$OTU*rD-F^H9LWdPE?wsUr~&q|(7;IuH&8pe-5?#-dvzB9x2-(~&?r5>F+PqBE0DiH=|_mI#Sd8oHCBZF4f3j)+KC zG%Yqo(g9RDL@*p?k+#lcG@fn~smw-3MPiYTNIbncnu-KEGBHMX1k<)!GLmW!b{fmE z=*EyQmGE6Au3EEt`HH||-(ueq)&2_19O@D~^k6m=|0zDI{~AS}x0$n}Gf3(^MHVK1 z>iU50?;6gl8h@9@x3nk(=Q*wqUaRml7j;)k#IOyoQ@X_GZFo3jt}z>arojp^VZ)zd z!`Etm==MbUfG@8NuZONAU0}oOp(XKkHazD>B3y36Q``7zwBhxwQIh;NJjaM-%CH=P zasa<`QA9>o;X`wVe@{p<10&jfe zZd0cP)A-1(T!(c_T3;!oebjGu)mvTlW>?)U-{7Jzan%=c-QU}=2xWio^luWVpGcpH zP(6s<01DNN!f^F@+u*I(qapY`l7Goj;<@lMe|F5j>%(RKT@#gl&x`)oK1Gz`Uxf!w! z`H{^alA>0~p>rGk+3rDqHZzcI-4Fj}%!87>Q1S~eROjDE$)B|*{5{93XnOWp0mk&~xbJoV6cZyM`tWolHj?LGBoXEB2tDoGRweme3qn<<0BT%o;g{jbgAw0be zZ&UetDpPwKYqKv>B!S#X4bC?k%3K%a5+z?`oexC6HAZ4NYm8#JM!Q;*YaKJL!Xv5a zX=7vLSfQZWSZ|GWp9cF67}_hK^;lz`Lr(zxSLEyCz@WU#ZZw^tndwT50--=glvF`eublpQuWI-DyoNc(dXk5*leQdX0Quq(p! z2bZcD)VpLoq%dezZzmwFCeLlG2P487RRc>8?;D7ewZEw1 zu~Dgp(tKD}X7fj7m!4w~zXGD)I-p~@FV{G?-a62|j{L{G&$S+KhD3L30`^q5-cJgt zv5j+O*B+2vyPW8xGrt-19_xVT&_jrukKAlk5(t=b?l!}FHBovSzGo0if#|mewR0p< zZw*T4`lNG%C7pX6nMyf#^!L&^5|qvjO6MLSx^(VM(0i;w<=ovs=Xa7sZ$qVV;b%bf zTSM9fG!tWOsJlhQrxDX_sOW>9Zp)4St;#-0N&7?6{sKe4LE0bk9Qp}to<{aBSN3}h z%eSQc;bI72e^}Z_2+02MB>TH-xqU|`xc$ohuxh`dzX^JeHLUF44Rn4dF!t;`THX6g z*e*8rpv@f3E#*Zs(8alhd2o{|Jsv+%dDI z32v#`a=mI&2_;Hn@2q=OtdIvi|AZc?pxegCALBMiwlNZlcaKj$0n%~1`{*+1g!^_k zeilOW0l1Y9Z=i%5hVB#6S&vk>R<(_80>v016;hkkaCPF5y7P{FHGzC}xO&5`!_}vW z9^8)({h01Mx#lq_!0qi-y5Uv9v@It()6(1I=EooFmL5MXovd5hyQBj@ub4m&LRvNQ zL)J4nE$tzaBQ5D%H%-!QY(}#r$3(wa=fF@;p*Q7C^Fzb?3{iR;-ZO}YfjIHBbOp)FTVbeZ z%9fU98~rOlH)-i8cEj^uRW8H~i&wgELTTx7C{m4o=ZI`72Ch9V-D&81K<}}1TG~qX zA0mm~hEc<^9f*Ewm}04d_W5b)A|o79_7Q4(TAFR>E2Mp$mX4l5?U&|?i^e}(TKdn$ za%7vd^h0=QX3JxTdA9VSgg|^LX(>#)9BJtW<%D~B-H2a%^4~E&{0T~5JT1w7lXp>P zTB<6YmTo$3TKet>rx}%M<{trBO_b+8g2vv44-8@l5GS6Nnn+&G{3t0@w;@|vnq%~@ z2HmWd=ts5uH$LO^(D|9leydS2L)t%~v~-WrzyG(~ zetTNlZ0JvbZdOYnvVR9K|2L~8zt5SLwj%^)wwwXC)NHv>HR(%9ORs#0@Ho=ao`q_% zx~A9OQxTe9fLrAS~GOK*PQN=q%+{@96!xTi1)clxmo_4o8mlh262 zU{xxxbfs;N5&QhyP^RJH77lt_`vX;0yH<{9Y8SGb?z?jbfuUb>^G z9t(Y=K8If)^S2xIF_ikP{WSNa6@9nd-+d{2kByb$Qd3t-Er9oX4!uxRk+lY}xpIXy z+Pxe|G+ge|?1cNd=79@$W%prK+S6Ad_gL@*-8w)vN~s<29Qp@jL&|JA@8qmO(z^iA zv9!fnTyBpXklP~%WEOcNXB{Z6#S^3M(G|JK^npD{ZTRUn9|X7RgR4*~W*ph4=tXO& zy7yD*)-e93@X7v_y$$Mcb?*V7+p2%thsCyausA%6g92>%?u?3bjRJbOU`xsx%I?B_ z(nF!5XLVY*%Y{D75VOr1nDfLu>}19U360&F=7W04ZylsH|45g5P;n|4%3{}P*fZdW zsd`+fv7msS7lu6}XF`_NZ6o8@(VG11u>XD6Y-qG?2KHN{svpNFV0u>x{W#{Zi7h$V zkE605N9A_W8))yKdk~dPJ?>~Tu}Y?LzEm_d2t~8S=!3O7NQ)H zgK~oEt-R}tTJ57@1Ah3=uQOWfQL=APVXow?Bd#X&&Z>}BXu{jJt#;=;^stCssJ*g9 za?{e`;>#m;7b#Pbi)Od+nF~zA?Jhox-|+KKK~y#TK9sPBMZxCX#$!F*$JB14)R$|; zPE~QYae{UmojY6H4HL2(CVDz(7wfKpd5j61j-R8 zN1${BR?sV}(fB5>erdHi(ve7R@yfSWK}f{CZNX?P5}qm6%ePv+>K)gxH~ZWA}poGF^ZVVjd^it;Aq?Z~9k-uB=v5pTLZ;!OoRBHnlj z!?5cO#=`{Qq8G2+M#;;WVx4h9+Jq0%nUt5@IldkF%B@XtFzt=T!;!AZMvkxN3;2uo zfz!!gNV!aI9k2bL=a)m-|1OUD<>ZC*b~3smmX3DDq9Hi|^h>zM(f@gV9#{G5cqq}) z2}^iE+bcDtWY1#sC8J%FeOVEV#}jF{daLYZU|uVs5_5Yrf^(U+Dp#3x(sIj8c9o6t{G**8=ubOHK*T_|h^+yF?s zYq$PKp)de=CElc=C$Q%MYw(P)_BVw>9pDHaY+3-fVMC`A@J{Sm-UUb(cBJLy*}7JE zx@tV<%sRP>Xy*ev7k_@(+(sO;YgWy!{dV;!o2#~oWoIw_<|XIPCp5|9wG*)&eAq@Z z7vah82Y~N1Sfw+K`VROm-|4qz*W6yQ;?$EWpQ(UpY4aY?R}B{m-6Y5KiZqcW{TZMi z23@{`&-BVWr?LD2(9eE%YWfk-@5YnVRPuAs{wi!d$#?&a{^^yh|5DID#%7`XX28(j z@JI)tyL^=+P>w)30_6ylBT$Y&IRgKmBEaA8@%MUqnJK>kq&Q$BUtM~!h@bo9wf7qe z(&jf^{9Pd4HIx^Ye)7V*b@Jlx18I#)7t??DSs_97G5ygUy%J8>#54u@8@f9E@g0BT zcUY0}J-WJTHPmcEv?zpDD<0MKIU4dec=~g5F`*4H{&HPsKl`*Ce|yOG@Z*d0zgF|( zR~cF7IAYl;>!Z41c*9TDjh|ZY9~$s?>~6btx}O_0Y|-$~H0;!Hn}&C3$lrBdwPM9m z?**+JGx2oBTkpHnSGRCcM%EU6Z*iTk?o!|43l-jNd2+oHUxxDc^LA<_zDMQnOye09o#P*W- zQ^ivy@w3f24o>~mg6FLh4<&VWj4D^+r4&9Fz)r2i7uP%WxdBf6872=n@pA;ve<%J- zQED7kin9cNAME7EY;Cb)RJl^r3LbxUYNhy!*ie$6vr>_p=Ib=^QK5)YCCYyl_`aOa z#bEkuji)zq?AKQ{p7SZw@WLxbbI~8p533}8a(p_p9?s9KXN%;Y9OnmsFXjJK?Rr@1 zIbBpOVa-KYO3%|ESBS~u=VwyShfaAKbf#n>v5kq&^`tKI?eAs=ipB@-hIx&rYZFN1@NW(%xOLD za}NGy3Ozs7{OGlQ-NTnOVB1@ERv3deV5l=wx~1{ozD;r z1e3|&mH@Wwl3PSu68pq~aHgYU3q)+S065bQR=uYe>g)vPsrM7gWgBtr5s(Bgym@ z;WMa>nP@D$Fd7z;p*@&t7ryY8IFu=vPAbYRktF?q;iv>aOGaWr640g2SX%gGKl)Jg zZAzd@g;XRYeCbFRs&d%)l8G>W;t;+_yB>b+VQAx;(xJwolHf8-1v{cp2Gyh)UP?V1 zQz_xYWT2x$O3bkZ<#|274metyvpheCFy!YBu*gfR z+@avkB{bxRg?V!^$1=o4K$)kt5tOz?L6+2wgX%y1hq*)G2Q zGrSu#yF9O(88RQM;5x&OVy&vqG(fnt|8&pVAuq{T?80?RS{JfOvHe`92O!eL{_}c#L!;tgJF2E^ahc#W7$^IBxET`5Ut}<A&(tJ zQUCv-QP*^0p?xH}*RVgFziEGleb*8a{9u);EpsVnyUa^#)Kc>Oe#LRBOTbkX{|yht B9Wnp_ delta 3285 zcmZWreQZ-z6u-A`U0=Hn`c}GjtZeOIFhN>A#|EzJ*a{6}qY#Nx(TE>$5Je>;Mj}iy zKY%*2aYQsGWDtTwOazL8BnW;1K{5P+$pZed;RiKCL_`=weV%*oyY;P{8K%An%%PrUG{8NUg!KKt8`x}VuCD9( zw%@8gxAM$4+g4rvdg|FG;uT}tTSZsKb`fEu$ zZ!?8eqR7)G2OLp`u94CKS)f<1;XHXjBHRyN`4rJwjm*Krt^ zSo;CiPR53kQQxMM2o4~4ge5H9l+pOJG}DFufD?APtEdw$x{JqU+|K^3>m9Lcbnrur z_A2XQA0pUZ>x9vx=FDJ*taj6@5KhL@9cfw}yRi=QQ8OkJWC;&)gAVRR*yn^zqju3e zxYsi&Zz5(O$#Qgm9z5=;rxW>17dzmXr=$Q&R*XG;hj&?=2ful0sRJe#mPWA7h4MvI zcErxp!JQa9gm8K(Gf&1&B%@-guVri-MWFQ#*j!jk?Qo{BbR2FZ5$M8g7N}1I>P9$y z1Q86XI3VA9ADS=ln#>>K=0DKE-DsjQ^D1H9VTVpMSK*Ae)Q8nk= zo4{wET=P@irfZmkEh=>R+UR(mgX_8AFWQrfyC;3~CQS8|J8k1-vX14#?V@^G&J3p= z#S+CZ`%=x{Vu|32zoi&%O5G<~xR#;wKivJD8h52|oq=Q2Hj1_hl-$!#`0JS+l z{X)C*2KT*^1LOQmyWG#T8~lY-fp`4twX)&Kx~0H`F>SPrC+lGgq{eJv=QZ2^@s2Ah zKPgf>Z6DEtQ(9>qEmr2RTZA1tN3QDnW5}z1*L8Le#F6(Pf8m<0vnhxE)b&2(`;ZSI zpNlK!g9FniSQ1y^!u0X&y%<2=4X7uh@LPso58}@=z#u>JHvA4?2glfP8XMty-BO>Q zs~oa8Z7N~m{q!160FCG(&11co)EifGS;(Nu6h_*s}M7|-+ zS_K!iWtctOoZC@kATr(a(1r3L!q)mQ=KG8Kd6u3sa5RMOO4ng-dv|YNas2MtoJAZj z!@JVhk0`}a4<8*NS2&)TB^+*8K_7$-F)svSDzwC0MT>CQ1!M$ydlZ@?@qL24g;~*O z;g1F{eVMnr_k$CqOu|-KMu6s4GSCnTB}j#k_X_W#;`+H|3)huY^-$j z8b-J`t~XZN<(6qpOm(JZ_1kn9V?YLud40Avar zTr|sRNQo&jrkV&!uq6Y!>T_UWqFiMF(j=a)HxuP{12|%Zqlrp;v{LBv;C5n`qp?az K`4nb1miz~uDYm=- diff --git a/csrc/tensor.cpp b/csrc/tensor.cpp index b3cfb3e..1a283ec 100644 --- a/csrc/tensor.cpp +++ b/csrc/tensor.cpp @@ -142,7 +142,7 @@ extern "C" { Tensor* sub_tensor(Tensor* tensor1, Tensor* tensor2) { printf("Adding tensor\n"); if (tensor1->ndim != tensor2->ndim) { - fprintf(stderr, "Tensors must have the same number of dimensions %d and %d for addition\n", tensor1->ndim, tensor2->ndim); + fprintf(stderr, "Tensors must have the same number of dimensions %d and %d for subtraction\n", tensor1->ndim, tensor2->ndim); exit(1); } @@ -193,7 +193,7 @@ extern "C" { for (int i = 0; i < ndim; i++) { if (tensor1->shape[i] != tensor2->shape[i]) { - fprintf(stderr, "Tensors must have the same shape %d and %d at index %d for addition\n", tensor1->shape[i], tensor2->shape[i], i); + fprintf(stderr, "Tensors must have the same shape %d and %d at index %d for subtraction\n", tensor1->shape[i], tensor2->shape[i], i); exit(1); } shape[i] = tensor1->shape[i]; @@ -211,4 +211,161 @@ extern "C" { return create_tensor(result_data, shape, ndim); } + + Tensor* elementwise_mul_tensor(Tensor* tensor1, Tensor* tensor2) { + printf("Adding tensor\n"); + if (tensor1->ndim != tensor2->ndim) { + fprintf(stderr, "Tensors must have the same number of dimensions %d and %d for element-wise multiplication\n", tensor1->ndim, tensor2->ndim); + exit(1); + } + + int ndim = tensor1->ndim; + int* shape = (int*)malloc(ndim * sizeof(int)); + if (shape == NULL) { + fprintf(stderr, "Memory allocation failed\n"); + exit(1); + } + + printf("Size: %d\n", tensor1->size); + printf("Data: ["); + for (int i = 0; i < tensor1->size; i++) { + printf("%.2f", tensor1->data[i]); + if (i < tensor1->size - 1) { + printf(", "); + } + } + printf("]\n"); + + printf("Size: %d\n", tensor2->size); + printf("Data: ["); + for (int i = 0; i < tensor2->size; i++) { + printf("%.2f", tensor2->data[i]); + if (i < tensor2->size - 1) { + printf(", "); + } + } + printf("]\n"); + + printf("Shapes : ["); + for (int i = 0; i < tensor1->ndim; i++) { + printf("%d", tensor1->shape[i]); + if (i < tensor1->ndim - 1) { + printf(", "); + } + } + printf("]\n"); + + printf("Shapes : ["); + for (int i = 0; i < tensor2->ndim; i++) { + printf("%d", tensor2->shape[i]); + if (i < tensor2->ndim - 1) { + printf(", "); + } + } + printf("]\n"); + + for (int i = 0; i < ndim; i++) { + if (tensor1->shape[i] != tensor2->shape[i]) { + fprintf(stderr, "Tensors must have the same shape %d and %d at index %d for element-wise multiplication\n", tensor1->shape[i], tensor2->shape[i], i); + exit(1); + } + shape[i] = tensor1->shape[i]; + } + + float* result_data = (float*)malloc(tensor1->size * sizeof(float)); + if (result_data == NULL) { + fprintf(stderr, "Memory allocation failed\n"); + exit(1); + } + + for (int i = 0; i < tensor1->size; i++) { + result_data[i] = tensor1->data[i] * tensor2->data[i]; + } + + return create_tensor(result_data, shape, ndim); + } + + Tensor* matmul_tensor(Tensor* tensor1, Tensor* tensor2) { + // Check if tensors have compatible shapes for matrix multiplication + if (tensor1->shape[1] != tensor2->shape[0]) { + fprintf(stderr, "Incompatible shapes for matrix multiplication\n"); + exit(1); + } + + // Calculate the shape of the result tensor + int ndim = tensor1->ndim + tensor2->ndim - 2; + int* shape = (int*)malloc(ndim * sizeof(int)); + if (shape == NULL) { + fprintf(stderr, "Memory allocation failed\n"); + exit(1); + } + for (int i = 0; i < tensor1->ndim - 1; i++) { + shape[i] = tensor1->shape[i]; + } + for (int i = tensor1->ndim - 1; i < ndim; i++) { + shape[i] = tensor2->shape[i - tensor1->ndim + 2]; + } + + int size = 1; + for (int i = 0; i < ndim; i++) { + size *= shape[i]; + } + + float* result_data = (float*)malloc(size * sizeof(float)); + if (result_data == NULL) { + fprintf(stderr, "Memory allocation failed\n"); + exit(1); + } + + for (int i = 0; i < tensor1->shape[0]; i++) { + for (int j = 0; j < tensor2->shape[1]; j++) { + float sum = 0.0; + for (int k = 0; k < tensor1->shape[1]; k++) { + sum += tensor1->data[i * tensor1->shape[1] + k] * tensor2->data[k * tensor2->shape[1] + j]; + } + result_data[i * tensor2->shape[1] + j] = sum; + } + } + + return create_tensor(result_data, shape, ndim); + } + + void reshape_tensor(Tensor* tensor, int* new_shape, int new_ndim) { + // Calculate the total number of elements in the new shape + int new_size = 1; + for (int i = 0; i < new_ndim; i++) { + new_size *= new_shape[i]; + } + + // Check if the total number of elements matches the current tensor's size + if (new_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); + } + + // Update the shape + tensor->shape = (int*)malloc(new_ndim * sizeof(int)); + if (tensor->shape == NULL) { + fprintf(stderr, "Memory allocation failed\n"); + exit(1); + } + for (int i = 0; i < new_ndim; i++) { + tensor->shape[i] = new_shape[i]; + } + tensor->ndim = new_ndim; + + // Update the strides + tensor->strides = (int*)malloc(new_ndim * sizeof(int)); + if (tensor->strides == NULL) { + fprintf(stderr, "Memory allocation failed\n"); + exit(1); + } + + int stride = 1; + for (int i = new_ndim - 1; i >= 0; i--) { + tensor->strides[i] = stride; + stride *= new_shape[i]; + } + } } + diff --git a/csrc/tensor.h b/csrc/tensor.h index 58e79a2..08052cb 100644 --- a/csrc/tensor.h +++ b/csrc/tensor.h @@ -15,6 +15,8 @@ extern "C" { Tensor* create_tensor(float* data, int* shape, int ndim); Tensor* add_tensor(Tensor* tensor1, Tensor* tensor2); Tensor* sub_tensor(Tensor* tensor1, Tensor* tensor2); + Tensor* elementwise_mul_tensor(Tensor* tensor1, Tensor* tensor2); + void reshape_tensor(Tensor* tensor, int* new_shape, int new_ndim); } #endif /* TENSOR_H */ diff --git a/csrc/tensor.o b/csrc/tensor.o index 363e2a1e5b5e39484d92b63a7ea55e2bbb7411f6..e5df58895e0a2d23d27d7e95e5e2f1412df4fa3e 100644 GIT binary patch delta 3700 zcmaJ@ZEO@p7~Z`W+AH@}XiH0gZlQ*cNb#c<8-g^&kwt~{2Q_ViVkuW!ZSUZE6=Ez$ zJZo}nE-8+QL@^;zBf%(=A~Xt0E6E8tlM6KFjHhXYgqBuqY=0EJc)suMOgkM?Cuw%) zndg0;dETAf-pKY72Yc7riXS_-{)nS5KljOO`c2M8wUPdjQ&Cx`N4op)mW-{JgYi*$p zw}Okt&cwQb$oKjQ*q}`x7`8MLs8O3x$1cRkkP6;0lEfVtB#9yqz!fT@1wCwA6X>iM z^+WEngApD>gP%G?h((TE$f+Jh71NI2AhdJ9K zE0gZV0Vvved(zziH?WN`D2^V^WHR>&q#}bTke~-0OG-QhO7fL30HcM#4wIBR`q#gi z4DEFkAFo;ZT#Src4V}PL2AS9l&k`CsF{hyk*3faYq2pmP0SyJ`q*1fHnvBb5ryb7D zO6pXXI+J?De6~sK2z^+sd}hk>nL?jRgG@<2QHKq6%~M%FHFJ;#4K-p5lPPuVs%cvv zkJJ^TgB=~PnMt}qef;@p3}cWu9Lfkuhr83t;jK(B*QiemdSp6cGe}ypx6aM98FM<^ zkygjP?!ERMZW%`3X-A7nUvSpAI=l7-15M4rHoq@VQRp6LDZS;?93o8jI=8O$wtG9g zzTk#KZJl1NqpLmGcA&kjxk;)E{V#2Yui4*m0BZKNd$rEirUTwit;HYEI+}ukw!`+AFfDzx4HU{ky1>+0C&4QT!rjd^xzZN6&F=RIUeYy15mkM$tm z+^Pjzy;^fuAOIn>&bHUQScw=@rHfp@FIWz(IY}?ON;h9J+oUM2EVv;q&s(-_&V^YE zvK%?N3bM>s53l%?X4Y@2uuczZC4C<-fj(<-^YJh0gD_cC0`wVaA%o7eSz}JAaAYwqNagds$VsdY)ce84HEj zPHeD_pU4j_Bjrto;|B3I4KEmuPZ_=*FqeRTDe0uK4KM};9Sp&B;uaoK&;V2Og*!rl zQ^4;rTsHO~2?h_qXi6c&QQr-JT)$MnJxXYHVWT9JMpYQqLP0*vxloP7U&sk8`N1;GW{EKA36W^JK zzXA`4pl>FI`5K19qX{XZ!4|0?YL+X?_Su3t{M{0tCvFt@V&WGizDqUpE6U44fw-X$ zq=I8?L!ZpUzbx^2fY?0zp^#L-6Z|9;hzW)zzKb2Yy8>SfumX=h4E}cT0BaZ?!q-DA z*di711lt)d^LI;p9-vX+iveyjoV~`OnbUYn654}XhA|-si64-^1>6HXeyX3s2@KpQ zo=c85FnsR!Bb3Sb*#d77?iNQ-5JE35bVR`MmgDyXJWs$M!S=$KYn;jKKpf8>H$Kl2 z?_vd*Ks0U;aM3s<;9}zA43~{hGaT(g_|4~uKNEyRmb zz~F1We>Vgn(KsdGqHzxUcZkMft9=ZYje8i*jc>C~A1h$*zuf9~1Ow6dV*wXieRUq= zpCx17>0z#e{$k?0d;xILST6(shTQ6x7%nG1!f>9L8;1oQ^eyD4H!3W6Xiw=YwWrVz z21?)`|IKm*?s0IO|7ICB(OKMCK<}3RY8x%2*R`s=CD!yos=AlZ<*UkV=SpVFVcK6p-&pm8J;`{}^xCRT q_UT8hs?*fIdXqh|%$haO9WYzVc5}v5Jh{5ueyY-151L3ybNmOnIdUw4cW5I*JOn_nYJ)&ex+u{ z!YDIYQTzDhAI1rj7ns~&%HWw?WopLNAjkj)0(_HqnVK=J;F)|;N^CNRnFfsMVWz>g zf(N1~m2q;FnHdu=-{ehZc1)Ih5VhtSOb&dLjm+&h6ZjzNW=yU$H|Lxn58)+D-U;OW dfXWF>{t4u9K&>sCY-wT6^#Q7~gaPO#IRO2AIU4`~ diff --git a/norch/__pycache__/tensor.cpython-38.pyc b/norch/__pycache__/tensor.cpython-38.pyc index d8a3609688d28acbcd6c68057067a6a5f8b6c6cd..8858bf5c9e3657400b1da54d5f0c9f30feeb3289 100644 GIT binary patch literal 5223 zcmd5=TaO$^74GWW^z>Zz;`MsHagrefi2-{zLPP{)MHaRR0W*Qv$-_Ws%}&+K_PFP= z)jf_qniY_>1iT>e4`v?CAK(@719;%E9^t9~f)Rr6RQ29`kql4iQJ=c@sj74N&Z+sN zRtqgWfBx%EuYJX`{z;9qPk_b^yy@RTxW!pwwXF{HSs zc7)6Q`?mZW=9dz0=(HVXzV+L_@E88lYL|q=177A~-{BQrb#ytC(38r7eqQb9UEiekI4btX7xk*~drY2Hx~XAi`=hZnbUBgtKskGw_zIXLGw~ z`MTwtYtbwpb39Y_y>^NFSg)i?H}8p2I+n`iovfouX(nSXQst%lorzHHh{r?artwU) znt=+UXw(^sC{ke*4aa<%P`w&O2h&bcjQHK`(F7byKmF*#*1g*wt5P@WCF4$}ylxba zGBD}IBH?ND#FDExk0wvQySqOgirpj`?cN%9r$aHy#_6sUT`4-5h$fG+vFz^e_QvDg zk*;NW-#e2>(gqWUq*Db6>^h@gz=D72k>`X-Q*Y4peUM{elDIg7nmnl2;qGJB_PNLX z$1o0F;sHu+ASlaZA}B*%<8_o3G8B|mG8L3Het~bGtdqH*T;VoEy`_93>Dk4do!z}@ zobcTw-YZU#j$2(i#_>aY^zh!qoBjo44wc_`Pg?-2l?TiiV%$c(zwt?P1 zSaU0K(dVHLrieaU_xTxH6n#sEIi7q5@Ev~<8}ZM4p5g*8nBPCF`d3+K06PgtwDu1dY6li}2|f>X&Fh;kA! zQX$!2q-{smskfpn(DbN~?64Rmxk&N|N&1%SG}(@Bs#;+xMx-j}$iB8jZ3U_%MM_QO zoAIUtvldBDlBlKTd18!mjd;S?<58SN(F$7WTOgJlu)4j$HrT4|v(SD^moCHa-$BV| zKAWxFGiG`G?K!QwU_r~#n>IjxYxNjs+D#}A#F zm|Vb{{t~2b9WuOj20>@+$d-5J)-CI^?L!Cb9L>%>^yX|~&)p?Z=3sMfFPscYJAx50 z7zjq7Y-Wu@%z_GaFPq!6yrgxSo(khM9;I1l)CD3WahfS#JcJmz@;j3WN_mkc`l>R5 zkOhzQoHA)TBDHGq0$ZkgdLWEPBO&#Os+hXg^%JHp6xI>J@uxD4ABg0U0SQux{00cU z=>~{Ja%Ukty^3F*Rqc??zFEk>;Inf*8FxEL`u@(#7&l${26)N0iF}R7sVy|AsTZd8 zsFjiiELN>m*{ptkspm9Je6;^gP{|OQGW*n1No9+oXUPs7w4fIF7nPVxwE=_wc>{k? zir`5_txV;j1dj{JN|p#_bMm-~;f`8$4u3PMs~ImNf6^*49z#H*-9}zUP1WHU1oja$ zibx<|qhUl}OTJD-i-V3`NJ^DQ;`7Lahn16R;iq-#IOD6&xD+`dE%iOHOKHrq*WpnO zXmMS8)QM1wiRXk$ZbR0uh00NZ)x%{V3I=H0%n%We*cU7VlxGBBR3h#zS>Mh|gJ4jO zLbNMIy-H(-1Ar3XdB|vFQ3k&nAQp}|Cl|C^vv*t5;hvC&LsqoWIG8MO_gn%{jY3$OP(ApssGwKNXB8%co42>cP zAy2i@bhNdqQn3%^Js`fPanwLpC5pg>m_*SJ!62nIEQ*Tu75K?IaBH?+nD3dX=dgyP zO}eKteGlX;SaGLMaQW^z;?}S!+)+hfd9va|C|wjUM7BX0#NLXvr7QUBZr6POY8$X zqc8P-E&1z{HO}xqK<~oil_ioo9sR1UoN!!O`G4@t?5ho_+bn#_8aWyDQ}*+j{n^#0 z>}u23(M@@>GwKV2oeuIc@JB9NLPN1HS~?VifxVv_AE4+bF~Dj$9`O)|w|fa%&q zBHwxml>Pu~eu7uWAuAP*DPd3MfDBc>$oaK}fihc}HY|fcGdQ-iyi-h8fX1Rg`!QE?qNPDbQ%Y^FU}Y zN|B;BG|o19W7+6%9%lfgX3b#vJJ|8F7<39hA*R1L5-7JjMz&3^`@{rI<&6P;Qmq-Q zD?W=Ru+>V@rNjR-5{I5?k>~)Nib!6C*=+-rC<^$9=7D`i-+Jh3k_M3YWU|x! zAC!PEUaSO3H$ap9{+MddiRf(YWKE;7$eonzvp9fw)g-gbYjL5aEyC z$_XrI&7X7-j~+fF{Du?<(-R0U{$KY3Mtj| z%$ikW@fx`(eRw&1)xhB8T6~eq&_5aGn`jBv=@y8ED_($$-dRxEWWGIHGe)^{`b`tJ z+Eqz0N0y0%L@2{k?#)|w?&zFOzDb>QhmgBOzDl(#J zow8Hfm~MguxPOM*!CFucs=;>P1grEjw_${|VJ99%5kj6ub}e`ct2rfNlG5*-W%4*W}Fu=wA4w8sfKePk;zc>$(cmMzZ delta 848 zcmZ8f&rj4q6rMNzvE6Q4rGOMzQ6dH^5~JbgK_O8SPs9Wz2E_*Ku-T2Q%WS!I!(n4Q zc_8B-u$yquL{EG2=HJmoFCGku2hW~-uYfU5^S<}q*Xhjn&Gg6AyI}gU>)H&*(dNuf z;xYD}ntz9Z%?d{I{o(X2I68bjxn{-1*sf?5-i8fJXz^)(mUF%+EKwX{q8P9ZTa<*| z24KSMm|eEvVDE@Rz=(@sjfJHihh2b&s=+TnpL)ZW434zl`^q)GYvT4K<1Wd#gHT~D zp0E_!Kp4X?u2kVIErok*Gid3lj@fG&DepjPba9uxiM?{lQ*DQq<*O2Km5|FpOqsdW zXoz|onsN+V^~IQ+(W{BrmFKZNfj+7R5JHTmgTO$C0kH7#)JJoycgdOuxT)5Q3kN#B zEnwtwA7K$gi-fWo#GIs}xW(K6G1@6pCiw|Z*@;YIV||D%uomo1ph+XI{1iG`m!S!B zHDkGrWAgGFE@xnedS=M=t$IaBI5Zw^ZVeu&U+&vw8*P?SsFk`L zEgEfxw(#7e67sC$bQ#uk6fS0rH74hAv(U>0OIodbT`nLcM+wRZncdvqjuKICzLN4H zv5N%sSI8xG4b!L29rSy2e9p{qXo*9PgT 0: + result += "\n" + " " * ((depth - 1) * 4) + for i in range(tensor.shape[depth]): + index[depth] = i + result += "[" + result += print_recursively(tensor, depth + 1, index) + "]," + if i < tensor.shape[depth] - 1: + result += "\n" + " " * (depth * 4) + return result.strip(",") + + index = [0] * self.ndim + result = "tensor([" + result += print_recursively(self, 0, index) + result += "])" + return result def __repr__(self): return self.__str__() @@ -87,10 +128,10 @@ class Tensor: def __sub__(self, other): if self.shape != other.shape: - raise ValueError("Tensors must have the same shape for addition") + raise ValueError("Tensors must have the same shape for subtraction") - Tensor._C.add_tensor.argtypes = [ctypes.POINTER(CTensor), ctypes.POINTER(CTensor)] - Tensor._C.add_tensor.restype = ctypes.POINTER(CTensor) + Tensor._C.sub_tensor.argtypes = [ctypes.POINTER(CTensor), ctypes.POINTER(CTensor)] + Tensor._C.sub_tensor.restype = ctypes.POINTER(CTensor) result_tensor_ptr = Tensor._C.sub_tensor(self.tensor, other.tensor) @@ -100,18 +141,56 @@ class Tensor: result_data.ndim = self.ndim return result_data + + def __mul__(self, other): + if self.shape != other.shape: + raise ValueError("Tensors must have the same shape for element-wise multiplication") + + Tensor._C.elementwise_mul_tensor.argtypes = [ctypes.POINTER(CTensor), ctypes.POINTER(CTensor)] + Tensor._C.elementwise_mul_tensor.restype = ctypes.POINTER(CTensor) + + result_tensor_ptr = Tensor._C.elementwise_mul_tensor(self.tensor, other.tensor) + + result_data = Tensor() + result_data.tensor = result_tensor_ptr + result_data.shape = self.shape.copy() + result_data.ndim = self.ndim + + return result_data + + def __matmul__(self, other): + if self.ndim != 2 or other.ndim != 2: + raise ValueError("Matrix multiplication requires 2D tensors") + + if self.shape[1] != other.shape[0]: + raise ValueError("Incompatible shapes for matrix multiplication") + + Tensor._C.matmul_tensor.argtypes = [ctypes.POINTER(CTensor), ctypes.POINTER(CTensor)] + Tensor._C.matmul_tensor.restype = ctypes.POINTER(CTensor) + + result_tensor_ptr = Tensor._C.matmul_tensor(self.tensor, other.tensor) + + result_data = Tensor() + result_data.tensor = result_tensor_ptr + result_data.shape = [self.shape[0], other.shape[1]] + result_data.ndim = 2 + + return result_data if __name__ == "__main__": from tensor import Tensor import time ini = time.time() - a = Tensor([[1, 2, 3], [1, 2, 3]]) - b = Tensor([[1, 2, 3], [1, 2, 3]]) + a = Tensor([[[1, 2, 3], [4, 5, 6], [7, 8, 9], [10, 11, 12]],[[13, 14, 15], [16, 17, 18], [19, 20, 21], [22, 23, 24]]]) + print(a) + b = a.reshape([4, 3, 2]) + #b = Tensor([[1, 2, 3], [1, 2, 3], [1, 2, 3]]) + #a = Tensor([[[[1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]], [[1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]], [[1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]], [[1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]]]]) + #b = Tensor([[[[1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]], [[1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]], [[1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]], [[1, 2, 3], [1, 2, 3], [1, 2, 3], [1, 2, 3]]]]) + #c = a @ b - c = a + b - b - - print(c) + print("\n###########", b) fim = time.time()