From 90969f2ab36ea316e606cacfc88a74f4314457c9 Mon Sep 17 00:00:00 2001 From: Victor Phan Date: Wed, 12 Nov 2025 20:52:57 +0700 Subject: [PATCH] apply CNN --- __pycache__/new_import_ODC.cpython-310.pyc | Bin 0 -> 24364 bytes cnn_model.py | 274 +++++++++++++++++++++ 2 files changed, 274 insertions(+) create mode 100644 __pycache__/new_import_ODC.cpython-310.pyc create mode 100644 cnn_model.py diff --git a/__pycache__/new_import_ODC.cpython-310.pyc b/__pycache__/new_import_ODC.cpython-310.pyc new file mode 100644 index 0000000000000000000000000000000000000000..1d30f299348c7354f2cd56ccfc831e0ae3acb42f GIT binary patch literal 24364 zcmb7s36vb?U1wKyb$9hKJw1m;qr-ApwrsT{jZR;(e8`e~U}>ya@>RBJdb(<6dQ{y# z`l@@Rnbaf(IV6OLI0ul!hEdpr6>K28ENs?>Wgp4HTb5mRVcCV+103Ey&;oCh2MOSf z?fw1!Ro&AwQh=R#{nc0Bef+=w{Z(sVAeX@3W31EKjusSpMjSL@$*L!a}ti` zB^sKoYf(IDCne3W4M%sAb15fzm^p@*Ze;9CT*|bwZgwT-q#W&Xw`SfhNX-Fz0ME2n zYz*3iQaQ5=e&)LaeEx`yuHcY=P2-=ACRz+LIFRuYaTWMB|iwO42Qqf3N*siQj~FpR}Kpcz++;qdp}4X76<4jD4o@ zl>L;vYu?lLGnm0G-m{J8?B^tHhi5mQx1X1I|IGVWoKtA+WUy(I@3t~VW;WL zI(4Lc)M+_$&RN7RIojz&dHyQ##bl`MwR2CMK3+NX+#|=II9=94^JG<>bz0|}OC_i7 z&wJIyiAME8MYYcR$Q`M*n$z`}j&dvhfl9SjbA2Xf8dd*ny*Yzu_PnaM-Ab)BQ`WnuT^y_{(f(NwQ|1Ro~^W|r%?Dq2N%h>eV_Px~hnZ6? zZ*iv8L|e(DPoBo3c%tsNUFT@aYbgK#Qt}ViFSyQ`dhKk~V8c~2Zbi0~(;j9M%~qpc z^&WPsZ9o}QOD8bGs(R9uO~QnbwxNH4;+&SkV1^rRwOR4cb*k8oimOyhg{Imctqk&W z$FLPA01B=O2hO-nzonQGLBi0J+(P%nlTS2HaYw418YyY_Zyb?4en;{1U)2+dHYn>w zP%J$4K>J+cH9dYy2KwhfPlF`V45q!JEseIdATg&qsX%kmmr!!q2n>H0>Sk8!K9g|F z(}}=vvYQgiDdgpX)Le4T=(arry7_r+IbBaIXOLT9+d=wR;zfJe3`~?wEhN-W@)?+$ z63zfx$u8|_r&lsT_OkXdYdII>R!paOSzFEr`5+URhdKJez{GeVTNBd6lP->9)3>Ttgn_O2)L6u&*1o;>+`adLgcc5M z%L68EHKmWJy0md-wvGnofkA!DYb;@(2v9aM+r)~N3z+zR; znyz~sPje=&_>}MZ@dV3)d$QWD&7!CFtZU83=<+sn3<~MDm0N9Cm=DRIteJbRgC%WXseP%xw4ZCwtiI)! zZr7~%>WC*MCYFj8SJ%Z3G-upH*3yV&ov8v*wk;sF>LDyy(^^XJvQA*=VcOqc0ZjvH z%U~d+K%qeb*4Y6p_4U_2YE8CeJ%QnTEJvs5$f95jaR=6<0-Pi|BM5Te_{mRx(Rxbe z;@8|JdM~FyxLqeq0sFem(3tW%ZkTM=RhaZ^US~UNef?v$$7DJ76I|HiQsLqfErdkSg;fk&f zqky^>0TR=-7LxtWO?b_5)YzM<1wM#fxQ^v_gbSMPc;4boZw+E<$FTfOpNs&eht0H} zt|~yF=klFkaeB5o&yBs6h1nYCf5-y6F~5jC27}{UJL}Ee4Y$!!i{)>T@NnwKNzPxBgnNFS6{U=*Y=KILRs zxV-={?6uDUhAG*XFnjz$&7J3g4)x}In63(3hAFUl4IgD|_4!40ABGqjORZKT)GsJI z!)ajvp#herjY445{4mRbR2HfVQSdP}cdA-D8)jwr7574|JRa&!EzF^qTZwupd9CJ5 zyWZ|N-IO9IbN@S8ws@|->j9m)jy>8xNwRPRp}2bv$j%Z>E6pn8i6~0=1VpPEM$FaH zT-ziE!52FX&I%K-(-)$rxd{H2Kw+B~;+T3G3lO)(Ik!6)xjnXcB_;6t)Nja4#P zj(P^AuOyU%$6I$q4Df}!?%8wn^kYvx@leaf_MV!>r0&@Lblq)Zn0uP;`3e{`@NtzV zjvbwtUko>LPBOj8UtCwL{V|UpDcbPb}y3ndS7T9{L?)x3r zkD58Xo-|vaOgNr&+0Cu ztGf{_-F5$z0+i^(=yNPCduC&ER-Y>22!8cOi-F~5eu{ITvp~+@#-H!V3G@rutYZplaMGqZn3xqVnA$Q@1V zXkDp}zpj;iT(V~C{GP{Z?wYK!ag~z)RRY=7Gy}4&f{uhu*w^-i7T26QjXyjJ=7=mM$GKGvza?gW|4i}p9Kv6 zCEFiQEN$a#Dm4Q0#9^-m=Ec7sk#5nY&6ooL&H@h?akK*Kt^*9V4S7!rm9ges<%{Zv znE!lBLEGa)jxg#FECewLw-lz%&$>`oW#Z!Hlpqv)cm~T?k;8(aK=%dCaj-9qV+>1V zvMSI{x=PqMKTJm?ENC5Mwma6*QzUPeZZ>(Ko( z5#gE(f_Vv%nZz>I8LEbxh$!ZODje;SsByaSInB|XB+^n)!NmDc`E=8CAeIzO9OsO- z0kezYq=0+`^+*|g89Yv}jSM)IZe7HIaMF|5UZGJ8Ae$s1x?BS%M2dtHza(P@oEU_Y z2tHPq2V&0qP$1Omt7(F`Rex+9}p5i z_?wR6I4D6if!E4rXf&OLx}BqZ8fh;>9rQKDcRx(Fp_$J}U*N3sq8kK{&KhY+ZLAw; zXi72$25nBXvm`GCkdRS?Kq0Ni>6Z;VLlLI3Fx!nRKz#8RqHuQoU?YO;yUl&Ka+pMG;O((u z>eGm>;nKM%Jyqwg;RUKTBLGpoOjX;nP_?C&(~bt-Ti=)<2fUnFx+~DYlc(qOxuk@K zlXQ$rFfT$isehN8Y>@meIXO{Unn5eT5$csSmhgKPSbTFJfZYWGcUY>0^Wj$hU z9!9}(J~ScgLiRrID_C0-U;v-~ z_2q*41dA1_&>=weAV98u97*cO5ZD4$vG6HUe6sFbfS50O_T*W2(a!Q%2yYc;7G}GV z1&Kr)nZwp35HvUjN-nt-?tg{SFCzPQ5fHYar~&Z}!l1TU{)@Vy72hcuqdJ7RcZ{U^ zB8t=qx!pgaQ34Rt{NF;*PFzd~TGZxrfJ$;D`Ju$L<^UnSYz79oE5J5{yyY~08TOnhS7Z6WL`sS! zr1tK%NzFFVD>UuK3&wI`c|df?W-t&8LR{>omzIZ?hoLhr1SM#6i_q(iOs4@V=HxZ_ zf=73RljY2X`lgq)4};0vlvs~h0X)R4iokP)CLk(aXabvu5VN``f@hC%1Bi=86HB*< zL8r2cO|D0S1ukbk`x*ZI!2?%sM^#{|Y-A9kIjjK132a9dADo-}C&XjG=v0=a{esL3vbQR{dipyokV#=jBG` zq%bYh^z9Pkw%wWP6*b0|LEm%+#CpFhlrh%`cXd z_V{X(l}fXM4S^;Vnl9h&nIOTQ0rD?Y9B2?{*Cy9!C4p${mTVo`?Ghy47{Sn;1T_yq>P#3pi6(8*OAz~vsS#<4LV z%t>>V85l>yJVRhs+DOBEUuPk7i0nxH8e0>vl|aCDAUg75mR8en%t-hga>3BSrlz&@ z+vyBKJcly{EgdQ3lAFn4{5FcM;V_7gdr(WV=1>NPvTQ{Mv}yhTI1Z6+ArTwEnwsRL z10B4FKw~@OM0O_dkmclLn3w5BhjJ-(tkraH0wbrle02|657X`iaQfik8fT%c3ex~N zuRaAX1iI|dm;=R$8v8TGjW zMuOO7eYwycm@5Va>MdbDq>*!Qr6fmwK*}YTi_xfH_GC#IK0&s!FgNC$JdjEOrpZ!J zl3M!Q;M@@U!#g%?U{ZJ+80YWB5~$B006viT=a_(%q}*UY4d-_=2qjE|Ea9sMSW3!E zvaTLvLdYV+>Pw72j=+W}=Xydxgd>V-ko8eNiGWZ66_yifXVh=-Jp(?ACWVM!qEKLG zDMucmtpoO@ddu$mjX)kqLcvrX0!Z6I3~CcK*6d&n+kmQOKa7jNG*p2Y=(kl?IRtdIXzf zwOJkrlQZ+M^w5k|A+E90u>MeC`9C?HKV)!>!D|f20of*D6%eQd)I+1@)#ud*kQ1h+ zA>Bc2wFla>ooaKo1$cv*NUgB$AG1Vq#+kN@^O56A*CsTw<{bcsVcNM^WT+!Ppxhad zJ^1BNC%+D;;1INaMShGRHirIauOV>h2H2Q)jCaSV0Z=0d3zAdkajlNE&=1AWb+{MP2_*@`uS0+WA@e9(knT% zrlmo)84S-Q;YyhJpxz%&ClA9_WMw!QCWeRPc^@Rtk;Sh-@`Q`)QaYA6N8Uej?mkGI zc#6cCTY)1;Fi;002fA>3kOOwkPA+A|&~eZHz4rh|6C7>*%zGnQ6Z@2~lcy#pm*|m3 zmjY`BN%%7#EN@VMiio1uULWkHfn{s69rYZt$6J!p2eM$KJqJ%U{79iu2_cwv3O??c< zX>EJrB5e<2B%R@M7?@UXv=;iE~2^xv=!)(%>0lphsKPJ<(oV(d=}yMH2Xyw&yKl zqFtYUD@77!UCFEegaHb7rhbR9-(?_hv6C^O(z59U2^ivAu|g?g)Y8$md_(et2NRt4 zWmf&C3@$NPi_^or6pZlMZ*mh#C}fCbnf723MdNMtHWIHPeLXz&0&)XAS~O}`pd)j5>>Vn8x-qZvI1#NJ3!nnkZLS}B zpEqo*8~Y|l&;Ca>l}3CG0C71fpUSv4Bex4cIG@xA#Q^lGv7BBSj`*q|9dSI+B_-6h zHh@d)fW|Tm{LQUE?a9S^f!b-{FCiCxlsKj7NY-|w)N>#aH0w80(Ua&NCBa3l5<$m= zB@(Ux;I!bF0LdS+KLB^&)zoqh%1X-q=^*z|0x^f0P1_}B=u(cYb;rk6S4UVLo)lAi*|;7Qcc`oS`g+GJhoW`^PJ>+!IXD39 ziroWO4ffVfdJNc*46TEAPnP8Xb+)ax7kMpp+XCFvj{pXv;cJ^(8b1l=aN5A&9p2aJ zAxq7oK{=Lk7S%g^eE+D~0$j(#0@x z1xL4v?vP;3;7oAlA9URNI^$U{r{#pJEPE$VMmpajV?6&*{H|Kp*PiUi#GX3WaiLaL ztT4rZbF=eupd+nC#4D`BziUtY{%`Tju!x*^asSb^^mlZVbdPCp^by)f?QYR9lRBMw z@DD8Te|3{F$}7f5@tdTit5~CpJ%~3--!=-U|4l`2mNimC20a`25o7oP0-Q$Z8ED_2 zx)idE4%Zy8ptRG^!AwB$9D1-dDL<4D{V`2srUMUWTIWGarG+4gb7N4CE4KS6mw8}X zp*sUbT`@01(;P&)T(OAS`njse#%qR3#6sa~ghO2waH!dZ?f@L@8tpD*jU<>8^&yl7 z1g|upsEKenh*k&Up@g=;1 za#h)cOg8qIV_i(!ui2$OyKL0r_p^x~U~rLv3@nFUA_Z;N8jJPVg+cnmC<9Xi*sv1( z(uyf~mXQI2C;3{CBBCc)38McZB0&ObeE^cMSR&cL?h32a8pUiXR?NcWa9zc028vm5 z2YrfJIMG8f3(n&*6|p9`2gtzmJAY-6P%B8P7d+e<^-}F)=_I6O% z+tj4`Hu_V4&ERhs{4ImG5CESF6n%#Y-(c`%2H%wr9qi*nX=uV zOGCzO!=%DO!4+U%sO=B61EF>>)b0qiJ45ZRvbmHw(W%0dqY)+%v6Dw1okVbA5|Zm_ z+c*vX+@%6!LikxPT27}qV;fH*6y2lxESohb@xikfj>T}Lubo}m(357)Kpo%J<1gj+ zL=waVv_PWG7oYU6ef-TgK4aBpzxt)-taah5murWtw~`b4r{B_cKqdJqIIhXY-EhghxrM<*xu9qUlm6OXTb$1J2+f@rW8Eo*-EtDlFjXY0+^ zFV`*q%{RVeId8u41+ZAyuGhcNmP(H`A^uyU_&sFpoZ>|a@t@vh?b~Jf?zsxQ5og-7 zWhwlG_*dhFD`d18o7Qq^=7iUW>mc-Bh5v)fWPJi(0&?%lk#KxJ7 zPc@8pmvFKwsdjP&wg7k+Br2c<$QOSbl#24rzWglmc^Nu=iKX+%DL^}x8O3SF`#@aV zgbkrvyC}6y$sd$_6O!v{eyJ~iXie+G>slX)YiIjvk4kM=9nkv*$;Tz8zWlMi{Edi@ zL%YVZo7T0x8Sh(AGkX~(pvu<1S~o-`rPj8%6m)hdb)%GuW@h!3x=BjGurH;yjQH}k*g8QS) zJD7P#6u*=4JMEz(P8Cwgf_r3PChl!;#)H43SE(+{RM>H^l>>GkKm#@4jBaZUa>@;9x?-`vaJ632Lv zI~kd6)ayJLTh$|2Gqs1o1cQAH_A=Pd-~fX=7#w797lS()+|2;rGXU9SaEQTu4Dc<3 zgrXNkXe`v>yJ7GmjCz2nM;JWF;2{Rb7#wAQuNQzCB7n;0QO1QH3+0o*#Zbcx{7uf| zcJaR^G~J4nF_1C3qm3o|pTm-ZR4{Nkhz3go5%k2h%SnYVK=e_Ba)u{gOwE-r2j2r+ ziNg(?N?52=a4Fb}b262R)2acq3DD{Z1mQpo?==?|!x-jexR0|4$1I5D6C5uTgHR4a za7mSRN@|USr~6kzqzA-YKzfc)!lk-WeEg!A0F9(ksGuSfOfzChi9w2`p9#`vP6nxh ze-$KJ`iSrRy?By^P<$wi6PPUiBSd;_r z6C98X+J(3Sqc%sDmoy-LnthPQF%C!?2+?v;(m2{d%(#1f20mSLJ(?Ma)dnPEZ zz6Uw@VijpP31BXorqiq0c^T^|Ji^i0&YQ1aXTkaBp^k7Qrxwq& z;B_Gdcf&g!SI#}aYJa&nkF)@H_voA)SF&J7!Aw9!RR6Lzw@POxDl+oFRf#8PZ zN34IyKzPv*dym6}K<-J(y^y|9f9>O6{V5E6bxwzthAqpAU81^-(^5&&V2$pj3`M2r z5fgJ)@r94b`!FSjlyJl8vvA~A&9(RV)jz>-bhMg6mFg*3Lf+byNF{ ztcab%Wz(4`UM}@XuD{4yzrf&^82k$czl^{x#>2;rdt7~U?1K1jb>9XeH?Hn$uW+>f zb}{z;0!z&O=vMB&iA}zn*KI1!9Ed!^;@tJj_Ba@JT8Ju?!7D%>mH9=P?H}?uQ>|}i zc+0K(FHSai+YC3G&6?y00)n@!PNNd_6bfHKi80Twpn?b^7B2Ab#CUEGMuXdfGRFU& z;5{oNlCxoDRGvGKvf*;#uEg>VtnyfJ``kun>@uxSyOwtcx1rr}XW+7i{M&+CgWW;- ztiG}dhK{Jtt>-=y>}YQe#xH6&pao}RaO-FyxNS>fc>*Kf5|4ZjM!grmeduB9<;2H8 z8-vZ5!`*-99IhSvp59pZTxYEN`o_91*n7=b_n!N|_=Ll_0om?i5}$C$?!y}U5A=V; zA=rbtY?is~!CdxYF8eT-{lN_@Hv~IGKd}w5JrcVSu?dM;i0zly7$yF+#Tna}a5jxWqj6S8*`EvWxdZicfQ!vh9$;gOK*-h@ zHU=Gx!`TL{#f`y)W8ryIFK2tOM{;fs_V?!m+k>0WrqJ$J(e5qXb|Xl{m~IRXaO7AI z?c7KD_Us_`>kjtDW~7lNoy2-DH^XxoHM5YEXp`0SuY z1vN7u(^DDww?y^LFw8#q| zvO+C4gMQ^th`s2pCdGSpDI>M4=S902YA^D80B3ls1RqTB*;sPo;PleKt~e38Fxa}b za0Ry_SSs-20baLkewVySp^k+PbFs5Mq3OEn%0ADztSxaxZ3w_#gu$ z zf`H5F1+i_YNe(ROcN*gE9vUu2j~g{n?Z_QIC+_qWItqu;rJvP=@;!eKW5mN?r)5ig zHN}?!;mqjqMrm7}MVhKIBim)4+BSUSxr%VMiJ2yRwxGu!*>IGPbGwX>@L%eWWbQrR zYC9`SNTJQn`cS`898&OTlbcPkMFKW?d8L^VnRzzc&C2m5&L~h9+RMqyTQ4haf!Foy zxRS)xQQo9ufZ5Flh6@w%pqJ z=(XKya|55P2+rVANhHiba|7jv&yQ~by0nUoZ%>HxH?+BR?!hb@mHBs@%yL}PpqgW5 zY1QHqU9pJoUa&>1BNwP!H79S;U+B(TxuMjWS~@RXz6!z20Tq+*5J=bQca8taPo|e*yA!^w9_LgQB^O zuGI8Vjceu*M;WwTh;#M7+4gQB52x!cYn;qU*&JFO7uJ}}aWoqJZ8IJ9w7#8jt=kSi zK2g%TD(>`mavAkLENTtE-XArCDg1bxed;A>h>8`rxK0K_Oy5&@X25m~X<8fSP>z>cJU}>(CpQ z0!CXVBCG6)G zknyb%Jp>B3o__Q7Um-PoN?qxS%;BIw0B85{m=>pT=yOBfe#q0G?z9ybL=xg>!o#BN|;-BTakw;+9p z5yj6E4Q$7ZM~!h@WEz%#euri4n~5W=K8B9o+CbjQtJGF8$pmiTTx0Tu7#f|lrxAgl zMlv@0LfaGH6+w(}g&QQe)koF7Tm(&B7U9*=ala`gS2f}1NTWI2fg#Rzg+7NDFH;e( zLO49}>tF74;cS(ejM&Hjh!MbdtP~A=mAvAceliu_sR7}x@wpELdgIr6rhoVmL z*a0}3%XNIA0PYXh^MxcfRUIsn5{|^6(YWjU_(#m6A|uR2y;5PQ{%cbw)_S+k8X2@!3!)!J#C1O&D(|au%><-BUz&ar|z(~FE(DVFTFcF zyr+X*qJbGE`L07Zb#}2n5Jq@t^bERCi*N=d0GzIIpAwokeveRY_rRc>7N^Ke=O%O@ z8*m1!VUMG~?&fHu#Pi?=qOVk_PjC%C&VVa+z3sqGpiN>Ie5uOp+YQ+RE`Q_~=`qH8 zKG8m8)TbCHTLh0s+!K<$P(R8yk>15>%^^TymlUE%XtgO8nB#O>ndMAzjLBLhb~j|-bsKh)Tv4&P3m6m9L9Zw#lvS0%r0=M6VZCx>FBkW94)t-dss z4PQ0WYc)Gxt11Xwixa$E@`r5Ss3B_e>t$7_Wx+|S4eWaiL70-wQgks>k23faQ?qhU zil2(JH)F-8ac#D{Df~cL{6Sc?31#h)AE|REpc-`>_?n#Da`qg}r5w7OlHnWWY`QMRYo+b|7=jQIa+G6)i&csC+GVLT2W}$DN;z_0?_6 z8U)x!8#W;yicmX|s>)KUJ4*5&_`p+yJLm$T-ot=ee6^dwe`Y`=WSgQ4p1`F9kwo*+ ztB6jj!2+cyB|384;YjeL2x^b_m7ef$sj(|&MPlmH#0Jys%4hI0gI5@QlEJ4Le38Ld z82mPa-)Ha}4E}<_w-~(50Q5B>UeMxnDQH_){1V^DtHl(jl#V3IA&d3W*XJe>#nd-^dHE%IZbKD#< kZ#IW=W2Rx=h4^-JaQGJUq`pVja_=NQVBVP9THNse09AV$sQ>@~ literal 0 HcmV?d00001 diff --git a/cnn_model.py b/cnn_model.py new file mode 100644 index 0000000..a648b14 --- /dev/null +++ b/cnn_model.py @@ -0,0 +1,274 @@ +""" +PyTorch 1D CNN Model for Land Use Classification +DΓΉng cho 3 channels: NDVI, VH, VV (13 time steps) +""" + +import torch +import torch.nn as nn +import torch.optim as optim +from torch.utils.data import Dataset, DataLoader +import numpy as np +from sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score, confusion_matrix +import matplotlib.pyplot as plt + + +class TimeSeriesDataset(Dataset): + """Custom Dataset for time series data""" + def __init__(self, X, y): + self.X = torch.FloatTensor(X) + self.y = torch.LongTensor(y) + + def __len__(self): + return len(self.X) + + def __getitem__(self, idx): + return self.X[idx], self.y[idx] + + +class CNN1D(nn.Module): + """1D CNN for time series classification + + Input: (batch_size, 3, 13) + - 3 channels: NDVI, VH, VV + - 13 time steps + Output: (batch_size, num_classes) + """ + def __init__(self, num_classes=8, dropout_rate=0.3): + super(CNN1D, self).__init__() + + # 1D Convolutional layers + self.conv1 = nn.Conv1d(in_channels=3, out_channels=32, kernel_size=3, padding=1) + self.bn1 = nn.BatchNorm1d(32) + self.relu1 = nn.ReLU() + self.pool1 = nn.MaxPool1d(kernel_size=2, stride=2) + + self.conv2 = nn.Conv1d(in_channels=32, out_channels=64, kernel_size=3, padding=1) + self.bn2 = nn.BatchNorm1d(64) + self.relu2 = nn.ReLU() + self.pool2 = nn.MaxPool1d(kernel_size=2, stride=2) + + self.conv3 = nn.Conv1d(in_channels=64, out_channels=128, kernel_size=3, padding=1) + self.bn3 = nn.BatchNorm1d(128) + self.relu3 = nn.ReLU() + self.pool3 = nn.MaxPool1d(kernel_size=2, stride=2) + + # Global average pooling + self.global_avg_pool = nn.AdaptiveAvgPool1d(1) + + # Fully connected layers + self.fc1 = nn.Linear(128, 64) + self.dropout = nn.Dropout(dropout_rate) + self.fc2 = nn.Linear(64, num_classes) + + def forward(self, x): + # Conv block 1 + x = self.conv1(x) + x = self.bn1(x) + x = self.relu1(x) + x = self.pool1(x) + + # Conv block 2 + x = self.conv2(x) + x = self.bn2(x) + x = self.relu2(x) + x = self.pool2(x) + + # Conv block 3 + x = self.conv3(x) + x = self.bn3(x) + x = self.relu3(x) + x = self.pool3(x) + + # Global average pooling + x = self.global_avg_pool(x) + x = x.view(x.size(0), -1) + + # FC layers + x = self.fc1(x) + x = self.dropout(x) + x = self.fc2(x) + + return x + + +class CNNTrainer: + """Trainer for PyTorch CNN Model""" + + def __init__(self, num_classes=8, learning_rate=0.001, device=None): + self.device = device or torch.device("cuda" if torch.cuda.is_available() else "cpu") + self.model = CNN1D(num_classes=num_classes).to(self.device) + self.criterion = nn.CrossEntropyLoss() + self.optimizer = optim.Adam(self.model.parameters(), lr=learning_rate) + self.history = {'train_loss': [], 'val_loss': [], 'train_acc': [], 'val_acc': []} + + print(f"πŸš€ Model initialized on device: {self.device}") + print(f" Total parameters: {sum(p.numel() for p in self.model.parameters()):,}") + + def train_epoch(self, train_loader): + """Train one epoch""" + self.model.train() + total_loss = 0 + correct = 0 + total = 0 + + for X_batch, y_batch in train_loader: + X_batch, y_batch = X_batch.to(self.device), y_batch.to(self.device) + + # Forward pass + outputs = self.model(X_batch) + loss = self.criterion(outputs, y_batch) + + # Backward pass + self.optimizer.zero_grad() + loss.backward() + self.optimizer.step() + + # Metrics + total_loss += loss.item() + _, predicted = torch.max(outputs.data, 1) + correct += (predicted == y_batch).sum().item() + total += y_batch.size(0) + + avg_loss = total_loss / len(train_loader) + accuracy = correct / total + return avg_loss, accuracy + + def validate(self, val_loader): + """Validate model""" + self.model.eval() + total_loss = 0 + correct = 0 + total = 0 + + with torch.no_grad(): + for X_batch, y_batch in val_loader: + X_batch, y_batch = X_batch.to(self.device), y_batch.to(self.device) + + outputs = self.model(X_batch) + loss = self.criterion(outputs, y_batch) + + total_loss += loss.item() + _, predicted = torch.max(outputs.data, 1) + correct += (predicted == y_batch).sum().item() + total += y_batch.size(0) + + avg_loss = total_loss / len(val_loader) + accuracy = correct / total + return avg_loss, accuracy + + def fit(self, X_train, y_train, X_val, y_val, epochs=50, batch_size=32, verbose=True): + """Train model with validation""" + train_dataset = TimeSeriesDataset(X_train, y_train) + val_dataset = TimeSeriesDataset(X_val, y_val) + + train_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True) + val_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False) + + print(f"\nπŸ“Š Training start:") + print(f" Train samples: {len(X_train)}") + print(f" Val samples: {len(X_val)}") + print(f" Batch size: {batch_size}") + print(f" Epochs: {epochs}\n") + + for epoch in range(epochs): + train_loss, train_acc = self.train_epoch(train_loader) + val_loss, val_acc = self.validate(val_loader) + + self.history['train_loss'].append(train_loss) + self.history['val_loss'].append(val_loss) + self.history['train_acc'].append(train_acc) + self.history['val_acc'].append(val_acc) + + if verbose and (epoch + 1) % 10 == 0: + print(f"Epoch [{epoch+1}/{epochs}] " + f"Train Loss: {train_loss:.4f}, Acc: {train_acc:.4f} | " + f"Val Loss: {val_loss:.4f}, Acc: {val_acc:.4f}") + + print(f"\nβœ… Training completed!") + print(f" Final Train Acc: {train_acc:.4f}") + print(f" Final Val Acc: {val_acc:.4f}") + + def predict(self, X_test): + """Predict on test data""" + self.model.eval() + X_test = torch.FloatTensor(X_test).to(self.device) + + with torch.no_grad(): + outputs = self.model(X_test) + _, predictions = torch.max(outputs, 1) + + return predictions.cpu().numpy() + + def evaluate(self, X_test, y_test): + """Evaluate on test data""" + y_pred = self.predict(X_test) + + accuracy = accuracy_score(y_test, y_pred) + precision = precision_score(y_test, y_pred, average='weighted', zero_division=0) + recall = recall_score(y_test, y_pred, average='weighted', zero_division=0) + f1 = f1_score(y_test, y_pred, average='weighted', zero_division=0) + + print(f"\nπŸ“ˆ Test Results:") + print(f" Accuracy: {accuracy:.4f}") + print(f" Precision: {precision:.4f}") + print(f" Recall: {recall:.4f}") + print(f" F1-Score: {f1:.4f}") + + return { + 'accuracy': accuracy, + 'precision': precision, + 'recall': recall, + 'f1': f1, + 'predictions': y_pred, + 'confusion_matrix': confusion_matrix(y_test, y_pred) + } + + def plot_history(self): + """Plot training history""" + fig, axes = plt.subplots(1, 2, figsize=(12, 4)) + + # Loss + axes[0].plot(self.history['train_loss'], label='Train Loss') + axes[0].plot(self.history['val_loss'], label='Val Loss') + axes[0].set_xlabel('Epoch') + axes[0].set_ylabel('Loss') + axes[0].set_title('Training and Validation Loss') + axes[0].legend() + axes[0].grid(True) + + # Accuracy + axes[1].plot(self.history['train_acc'], label='Train Acc') + axes[1].plot(self.history['val_acc'], label='Val Acc') + axes[1].set_xlabel('Epoch') + axes[1].set_ylabel('Accuracy') + axes[1].set_title('Training and Validation Accuracy') + axes[1].legend() + axes[1].grid(True) + + plt.tight_layout() + plt.show() + + def save(self, filepath): + """Save model""" + torch.save(self.model.state_dict(), filepath) + print(f"βœ… Model saved to {filepath}") + + def load(self, filepath): + """Load model""" + self.model.load_state_dict(torch.load(filepath, map_location=self.device)) + print(f"βœ… Model loaded from {filepath}") + + +def reshape_for_cnn(X): + """Reshape data for CNN + + Input: (n_samples, n_features) where n_features = 13*3 = 39 (13 timesteps x 3 channels) + Output: (n_samples, 3, 13) - (batch, channels, timesteps) + """ + n_samples = X.shape[0] + n_timesteps = 13 + n_channels = 3 + + # Reshape: (n_samples, 39) -> (n_samples, 3, 13) + X_cnn = X.reshape(n_samples, n_channels, n_timesteps) + return X_cnn