diff --git a/checkpoints_tree/task1_tr_tree/best_tr_tree.pkl b/checkpoints_tree/task1_tr_tree/best_tr_tree.pkl new file mode 100644 index 0000000..1ffd2e5 Binary files /dev/null and b/checkpoints_tree/task1_tr_tree/best_tr_tree.pkl differ diff --git a/checkpoints_tree/task1_tr_tree/model_selection.csv b/checkpoints_tree/task1_tr_tree/model_selection.csv new file mode 100644 index 0000000..977afdb --- /dev/null +++ b/checkpoints_tree/task1_tr_tree/model_selection.csv @@ -0,0 +1,4 @@ +model_name,train_mean_rms_error,train_median_rms_error,train_max_rms_error,val_mean_rms_error,val_median_rms_error,val_max_rms_error +random_forest,0.23011051198123578,0.1289880713751052,1.3614534820442807,0.2941758399359114,0.1350525270089895,0.8030963210135198 +extra_trees,3.6456671909608616e-08,2.9450879227826572e-08,1.1542847206743363e-07,0.30134282357310554,0.1285832699374185,0.989878917287992 +gradient_boosting,0.012023494394761644,0.005998110300540536,0.11812203696846045,0.3579803351621235,0.3397381373027084,0.9497835717961135 diff --git a/evaluation_outputs/task1_tr_tree/evaluation_train_all_samples.csv b/evaluation_outputs/task1_tr_tree/evaluation_train_all_samples.csv new file mode 100644 index 0000000..2f1ac82 --- /dev/null +++ b/evaluation_outputs/task1_tr_tree/evaluation_train_all_samples.csv @@ -0,0 +1,37 @@ +file_name,frequency_hz,true_rms,pred_rms,relative_error_percent,x_rms,pred_tr,true_tr,dominant_frequency_hz,frequency_squared,inverse_frequency_hz,log_frequency_hz,input_rms,input_peak_abs,input_peak_to_peak,crest_factor,middle_length_ratio,dominant_amplitude,dominant_energy_ratio,harmonic_fit_amplitude,harmonic_fit_residual_ratio,spectral_peak_prominence,half_power_bandwidth_hz,spectral_centroid_hz,signal_mean,sampling_rate +harmonic_5mm_1.95Hz.csv,1.95,2.063512086868286,2.042942868521436,0.996806293394051,0.4040626883506775,5.056004742383957,5.106910705566406,1.9470664262771606,3.7910678386688232,0.5135931372642517,0.6663238406181335,0.4040626883506775,1.629040002822876,2.9249401092529297,4.031651496887207,0.25,0.407725989818573,0.9061827659606934,0.44422510266304016,0.629830002784729,43.91968536376953,0.09677428752183914,2.039226531982422,-0.004885347560048103,50.000047683761295 +harmonic_5mm_0.7Hz.csv,0.7,0.2675361931324005,0.2647816168893098,1.0296088207129102,0.11656410992145538,2.2715535430908202,2.295184850692749,0.6991758346557617,0.4888468384742737,1.43025541305542,-0.3578530251979828,0.11656410992145538,0.33765000104904175,0.6533100008964539,2.8966891765594482,0.25,0.1183394342660904,0.8892508745193481,0.11721993237733841,0.7034800052642822,37.61368942260742,0.06349212676286697,0.9012212753295898,-0.00016539028729312122,50.000047683761295 +harmonic_5mm_0.8Hz.csv,0.8,0.3204104006290436,0.3104756411225298,3.1006357743098976,0.09064009040594101,3.4253677344322204,3.5349743366241455,0.7963964343070984,0.6342473030090332,1.2556560039520264,-0.22765816748142242,0.09064009040594101,0.43685001134872437,0.7310900092124939,4.819611549377441,0.2499212622642517,0.0584954135119915,0.7228869199752808,0.06087201461195946,0.8802145719528198,26.92901039123535,0.09451805055141449,1.1917314529418945,0.002280592219904065,50.000047683761295 +harmonic_5mm_1.45Hz.csv,1.45,1.0070611238479614,1.0508663280991015,4.349805906890859,0.2516467273235321,4.175958651542664,4.0018839836120605,1.451728343963623,2.107515335083008,0.6888340711593628,0.372754842042923,0.2516467273235321,0.9840899705886841,1.9117000102996826,3.9106009006500244,0.25,0.25519290566444397,0.8691993951797485,0.26430585980415344,0.6701168417930603,45.977691650390625,0.09523818641901016,1.7056236267089844,-0.00019088915723841637,50.000047683761295 +harmonic_5mm_1.1Hz.csv,1.1,0.5534171462059021,0.5779597429583696,4.434737326215095,0.1373506784439087,4.207913273572922,4.0292277336120605,1.0999548435211182,1.2099006175994873,0.9091282486915588,0.0952691063284874,0.1373506784439087,0.587689995765686,1.1190800666809082,4.2787556648254395,0.24994152784347534,0.1195499375462532,0.8312234282493591,0.11857906728982925,0.7920843362808228,37.334434509277344,0.04679461568593979,1.324162483215332,-7.055706373648718e-05,50.000047683761295 +harmonic_5mm_1.65Hz.csv,1.65,0.5992732048034668,0.6282565440955618,4.836415020691626,0.38267070055007935,1.6417680872678757,1.566028356552124,1.6503766775131226,2.723743200302124,0.605922281742096,0.5010035634040833,0.38267070055007935,1.171970009803772,2.2650599479675293,3.0626070499420166,0.24991999566555023,0.4040696918964386,0.9350709319114685,0.4694216847419739,0.4976775050163269,43.57965850830078,0.0960308238863945,1.6984570026397705,-0.006530857179313898,50.000047683761295 +harmonic_5mm_0.5Hz.csv,0.5,0.18624502420425415,0.19870795051296228,6.6916828312363865,0.06567567586898804,3.025594299316406,2.835829496383667,0.5023733973503113,0.25237902998924255,1.990551233291626,-0.6884115934371948,0.06567567586898804,0.22301000356674194,0.4394000172615051,3.395625352859497,0.25,0.024631651118397713,0.31016308069229126,0.026469262316823006,0.9584437608718872,16.04189109802246,0.050632961094379425,2.159400463104248,0.0017537976382300258,50.000047683761295 +harmonic_5mm_0.85Hz.csv,0.85,0.5794365406036377,0.5367616244860894,7.364899022952703,0.12664538621902466,4.238303822278977,4.575267791748047,0.845729649066925,0.7152586579322815,1.1824109554290771,-0.1675555258989334,0.12664538621902466,0.5095099806785583,0.9128699898719788,4.023123264312744,0.24991869926452637,0.0956917256116867,0.6594006419181824,0.09168906509876251,0.8589633107185364,18.147825241088867,0.06506186723709106,1.4032906293869019,-0.0011623100144788623,50.000047683761295 +harmonic_5mm_1.15Hz.csv,1.15,0.9153674244880676,0.8472856717514132,7.437642078505262,0.18818436563014984,4.502423295974731,4.864205360412598,1.1485493183135986,1.319165587425232,0.8706635236740112,0.13849970698356628,0.18818436563014984,0.7081000208854675,1.3734800815582275,3.7627995014190674,0.24991869926452637,0.17172867059707642,0.8639681935310364,0.183344766497612,0.7253570556640625,30.369056701660156,0.0975928083062172,1.386907696723938,-0.0019338843412697315,50.000047683761295 +harmonic_5mm_2.1Hz.csv,2.1,0.9339020252227783,1.003448067990676,7.446824280235163,0.6165724992752075,1.6274616029262543,1.5146669149398804,2.10196590423584,4.4182610511779785,0.4757450819015503,0.7428730726242065,0.6165724992752075,1.728060007095337,3.2244300842285156,2.802687406539917,0.25,0.7569394707679749,0.9819692969322205,0.7959626913070679,0.40733572840690613,133.85641479492188,0.11111121624708176,2.100722312927246,-0.005007810425013304,50.000047683761295 +harmonic_5mm_0.65Hz.csv,0.65,0.2034619301557541,0.22107336420719614,8.655886650618205,0.1112922951579094,1.986421107530594,1.8281761407852173,0.6577304005622864,0.43260928988456726,1.5203797817230225,-0.4189601540565491,0.1112922951579094,0.35517001152038574,0.6907100081443787,3.191326141357422,0.25,0.06641923636198044,0.5974180698394775,0.06589993834495544,0.9084223508834839,14.282920837402344,0.09230778366327286,1.486210823059082,-0.00035510817542672157,50.000047683761295 +harmonic_5mm_1Hz.csv,1.0,0.3037901520729065,0.33260700774091306,9.485776767737626,0.21532277762889862,1.544690308213234,1.4108593463897705,0.9989057183265686,0.9978126883506775,1.0010954141616821,-0.0010948515264317393,0.21532277762889862,0.635129988193512,1.2681899070739746,2.94966459274292,0.2499224841594696,0.2571507394313812,0.9758802056312561,0.2638331651687622,0.49583709239959717,152.86651611328125,0.0930522009730339,1.0530376434326172,-0.0009750677854754031,50.000047683761295 +harmonic_5mm_1.2Hz.csv,1.2,0.7224284410476685,0.6515492682358728,9.811237872778989,0.3020855486392975,2.156836933016777,2.391469717025757,1.1994376182556152,1.4386504888534546,0.8337240815162659,0.1818527728319168,0.3020855486392975,0.8440300226211548,1.5216000080108643,2.7940099239349365,0.24991999566555023,0.3333832025527954,0.9823232889175415,0.38784340023994446,0.4198954999446869,124.32585906982422,0.0960308238863945,1.2258714437484741,0.0011010364396497607,50.000047683761295 +harmonic_5mm_1.35Hz.csv,1.35,1.256479024887085,1.1329108880206473,9.83447669391395,0.20795053243637085,5.447982627153396,6.042201519012451,1.3494648933410645,1.8210554122924805,0.7410345077514648,0.29970812797546387,0.20795053243637085,0.8708500266075134,1.7264100313186646,4.187775135040283,0.24991999566555023,0.1956998109817505,0.8739701509475708,0.19486407935619354,0.7484625577926636,37.17776870727539,0.0640205442905426,1.5017311573028564,-0.005424206610769033,50.000047683761295 +harmonic_5mm_2.3Hz.csv,2.3,3.7199018001556396,3.330646787548538,10.464120654766079,0.9912338852882385,3.360101825594902,3.7527992725372314,2.301363945007324,5.296276569366455,0.43452489376068115,0.8335019946098328,0.9912338852882385,1.9464600086212158,3.861459970474243,1.963673710823059,0.24992592632770538,1.217545509338379,0.9911264181137085,1.3459256887435913,0.2792012691497803,227.761474609375,0.08891531825065613,2.3008110523223877,0.008378428407013416,50.000047683761295 +harmonic_5mm_0.9Hz.csv,0.9,0.3999128043651581,0.4468199892673058,11.729353096510772,0.1796521544456482,2.4871396095752716,2.2260396480560303,0.8994534015655518,0.8090164065361023,1.1117863655090332,-0.10596802830696106,0.1796521544456482,0.5928999781608582,1.0781500339508057,3.3002665042877197,0.2499224841594696,0.2150648981332779,0.9518226981163025,0.21449175477027893,0.5359740853309631,58.302303314208984,0.06203479692339897,0.9868842959403992,-0.00024062092415988445,50.000047683761295 +harmonic_5mm_1.75Hz.csv,1.75,1.59738028049469,1.4050853159227565,12.038145638831969,0.35962721705436707,3.907060559630394,4.441766738891602,1.7508330345153809,3.065416097640991,0.5711566805839539,0.5600916743278503,0.35962721705436707,1.2678200006484985,2.4300899505615234,3.5253727436065674,0.25,0.39048516750335693,0.9655163288116455,0.4300072193145752,0.5336334109306335,81.74029541015625,0.08955232053995132,1.7760101556777954,-0.000693939218763262,50.000047683761295 +harmonic_5mm_2.15Hz.csv,2.15,3.0180647373199463,2.6414293752430624,12.479366576190065,0.46591538190841675,5.669332839846611,6.477710247039795,2.1465415954589844,4.607641220092773,0.4658656418323517,0.7638580203056335,0.46591538190841675,1.5632699728012085,3.0576400756835938,3.3552658557891846,0.2499212622642517,0.48979803919792175,0.9034159779548645,0.49743175506591797,0.6561670303344727,72.98995971679688,0.063012033700943,2.2847402095794678,0.002827510703355074,50.000047683761295 +harmonic_5mm_1.5Hz.csv,1.5,1.2808952331542969,1.110302433240289,13.318247698830973,0.45307496190071106,2.450593227624893,2.827115535736084,1.5013059377670288,2.2539196014404297,0.666086733341217,0.40633538365364075,0.45307496190071106,1.1585400104522705,2.150049924850464,2.5570602416992188,0.25,0.5677548050880432,0.9871437549591064,0.5882158875465393,0.3955877423286438,194.53369140625,0.08695660531520844,1.5181952714920044,0.0008904285496100783,50.000047683761295 +harmonic_5mm_1.3Hz.csv,1.3,0.9550632834434509,0.8240282575900374,13.72003595206496,0.2030215710401535,4.058821204900742,4.704245567321777,1.3018183708190918,1.6947312355041504,0.7681562900543213,0.26376205682754517,0.2030215710401535,0.7108299732208252,1.34552001953125,3.501253366470337,0.24993710219860077,0.20925211906433105,0.9342537522315979,0.22753745317459106,0.6098353862762451,55.399173736572266,0.07549075782299042,1.397734522819519,-0.00013690412743017077,50.000047683761295 +harmonic_5mm_1.6Hz.csv,1.6,1.1856828927993774,1.011715705384904,14.672319932333666,0.22232241928577423,4.550668837785721,5.333168029785156,1.60205078125,2.5665667057037354,0.6241999268531799,0.47128453850746155,0.22232241928577423,0.8100799918174744,1.6100399494171143,3.643717050552368,0.2499159723520279,0.20403607189655304,0.9129946231842041,0.22558331489562988,0.6967292428016663,54.344844818115234,0.10087434202432632,1.6988458633422852,0.0011404975084587932,50.000047683761295 +harmonic_5mm_2.45Hz.csv,2.45,4.167365074157715,3.5243981540885345,15.428619970357008,0.5156936049461365,6.83428710436821,8.081088066101074,2.4505887031555176,6.00538444519043,0.40806522965431213,0.8963282704353333,0.5156936049461365,1.7411500215530396,3.2864298820495605,3.376326560974121,0.25,0.5547626614570618,0.9294771552085876,0.5594695210456848,0.6415255665779114,73.60872650146484,0.06451619416475296,2.469957113265991,0.0013848240487277508,50.000047683761295 +harmonic_5mm_1.05Hz.csv,1.05,0.36426904797554016,0.42230971018971875,15.933459770119113,0.2716720998287201,1.5544831819534302,1.3408409357070923,1.0530656576156616,1.1089472770690918,0.9496083855628967,0.051705583930015564,0.2716720998287201,0.7901600003242493,1.5045499801635742,2.908506393432617,0.24991999566555023,0.3438712954521179,0.9576663374900818,0.3379869759082794,0.4775521755218506,72.97195434570312,0.0640205442905426,1.126065731048584,-0.0008716708398424089,50.000047683761295 +harmonic_5mm_1.8Hz.csv,1.8,1.6704485416412354,1.3841861638257655,17.136856998552837,0.5062151551246643,2.734383097410202,3.2998785972595215,1.8012458086013794,3.2444865703582764,0.5551713109016418,0.5884785652160645,0.5062151551246643,1.4020500183105469,2.7015299797058105,2.769672155380249,0.25,0.6510956883430481,0.9815194606781006,0.6610633730888367,0.3853244483470917,126.02283477783203,0.05555560812354088,1.819625735282898,0.0018060111906379461,50.000047683761295 +harmonic_5mm_2.4Hz.csv,2.4,4.107017993927002,3.396024555134596,17.311670897077725,0.9670910835266113,3.5115870810747145,4.246774673461914,2.4017505645751953,5.768405437469482,0.4163629710674286,0.876197874546051,0.9670910835266113,2.3246400356292725,4.529560089111328,2.403744697570801,0.25,1.1933718919754028,0.9901480078697205,1.3087358474731445,0.2874786853790283,216.48052978515625,0.10344837605953217,2.4013094902038574,-8.548868208890781e-05,50.000047683761295 +harmonic_5mm_2.2Hz.csv,2.2,1.0718663930892944,1.2945454730228163,20.774891476140446,0.6085917353630066,2.1271164194345475,1.7612240314483643,2.201188325881958,4.845230579376221,0.45430004596710205,0.7889974117279053,0.6085917353630066,1.7638200521469116,3.3627500534057617,2.8981993198394775,0.25,0.7086807489395142,0.9761964082717896,0.7861699461936951,0.4083978533744812,121.00875854492188,0.07894744724035263,2.1985433101654053,0.0010015374282374978,50.000047683761295 +harmonic_5mm_0.6Hz.csv,0.6,0.16537058353424072,0.2013263053238447,21.742513705382912,0.09397729486227036,2.142286662101746,1.7596864700317383,0.5964841246604919,0.3557933568954468,1.6764904260635376,-0.5167025923728943,0.09397729486227036,0.3058300018310547,0.5446599721908569,3.2542967796325684,0.2499224841594696,0.08657485246658325,0.8864659667015076,0.09075998514890671,0.7302541732788086,35.94514465332031,0.0930522009730339,0.8173193335533142,0.001446693786419928,50.000047683761295 +harmonic_5mm_2.25Hz.csv,2.25,2.1223104000091553,1.6541852357127729,22.057337338325393,0.45399144291877747,3.643648490548134,4.67478084564209,2.2518274784088135,5.070727348327637,0.4440837502479553,0.8117421269416809,0.45399144291877747,1.468690037727356,2.881589889526367,3.235060930252075,0.24992366135120392,0.4949429929256439,0.961538553237915,0.5433851480484009,0.5339846014976501,91.52095031738281,0.09163112193346024,2.2496755123138428,-0.001197385834529996,50.000047683761295 +harmonic_5mm_1.9Hz.csv,1.9,0.5868141055107117,0.7590829336435948,29.356626999099728,0.5143030285835266,1.4759449030160905,1.1409889459609985,1.9012991189956665,3.614938497543335,0.5259561538696289,0.6425374150276184,0.5143030285835266,1.3888100385665894,2.6350998878479004,2.7003729343414307,0.25,0.6525437235832214,0.9834353923797607,0.670233428478241,0.38811808824539185,157.31686401367188,0.06250005960464478,1.9118845462799072,-0.0027283257804811,50.000047683761295 +harmonic_5mm_2.05Hz.csv,2.05,0.6199118494987488,0.82180861364958,32.568624767228094,0.7315996885299683,1.1233036680221558,0.8473374843597412,2.0498123168945312,4.201730251312256,0.48784953355789185,0.71774822473526,0.7315996885299683,1.6353000402450562,3.2010200023651123,2.235239028930664,0.24991999566555023,0.9387410283088684,0.9784685373306274,0.9522533416748047,0.3911336362361908,110.50621795654297,0.0640205442905426,2.0580592155456543,-0.004247209522873163,50.000047683761295 +harmonic_5mm_0.55Hz.csv,0.55,0.1971234530210495,0.2787700623838687,41.41902351624325,0.109112448990345,2.5548877782821657,1.8066083192825317,0.5490954518318176,0.30150580406188965,1.8211770057678223,-0.5994830131530762,0.109112448990345,0.40608999133110046,0.751579999923706,3.72175669670105,0.25,0.06793846935033798,0.5611140727996826,0.07492320239543915,0.874541163444519,13.549410820007324,0.08955232053995132,1.263006329536438,-0.0011556772515177727,50.000047683761295 +harmonic_5mm_1.7Hz.csv,1.7,0.6067864894866943,0.8862511380642422,46.05650478703942,0.32135581970214844,2.7578499710559843,1.8882076740264893,1.700693964958191,2.892360210418701,0.5879952311515808,0.531036376953125,0.32135581970214844,1.1075899600982666,2.123849868774414,3.446615695953369,0.25,0.3294788897037506,0.9566228985786438,0.3822683095932007,0.5411480069160461,79.01654815673828,0.09523818641901016,1.7275936603546143,0.0009157023159787059,50.000047683761295 +harmonic_5mm_1.4Hz.csv,1.4,0.3850800693035126,0.6183495861111026,60.57688657569331,0.2212747484445572,2.7944878051280977,1.740280270576477,1.3995026350021362,1.9586076736450195,0.7145395278930664,0.3361169397830963,0.2212747484445572,0.8126099705696106,1.4906599521636963,3.6724026203155518,0.2499212622642517,0.2121341973543167,0.9263311624526978,0.2458721548318863,0.6176945567131042,50.935482025146484,0.09451805055141449,1.501712679862976,0.0024810773320496082,50.000047683761295 +harmonic_5mm_2.5Hz.csv,2.5,0.6538841724395752,1.1258108364036792,72.17282262749896,0.5744112730026245,1.9599386177062987,1.1383553743362427,2.5015766620635986,6.257885932922363,0.39974790811538696,0.9169211983680725,0.5744112730026245,1.5267499685287476,2.5278899669647217,2.6579387187957764,0.2499280571937561,0.7362810969352722,0.9838377833366394,0.7437620162963867,0.4004853665828705,234.6715087890625,0.05757058039307594,2.530334234237671,-0.01246642041951418,50.000047683761295 +harmonic_5mm_2.35Hz.csv,2.35,0.7861254215240479,1.6179925774542805,105.81863060954142,0.991222620010376,1.6323200709819794,0.7930866479873657,2.351473808288574,5.529428958892822,0.42526522278785706,0.8550422787666321,0.991222620010376,2.2557199001312256,4.422339916229248,2.2756946086883545,0.2499212622642517,1.1917885541915894,0.9816745519638062,1.3294483423233032,0.31509023904800415,127.34346771240234,0.09451805055141449,2.3649213314056396,-0.004312594421207905,50.000047683761295 +harmonic_5mm_2Hz.csv,2.0,0.5447156429290771,1.2863206517188583,136.14534820442807,0.7751342058181763,1.6594812124967575,0.7027371525764465,2.0002939701080322,4.001175880432129,0.49992653727531433,0.693294107913971,0.7751342058181763,1.559309959411621,3.019509792327881,2.011664390563965,0.24990099668502808,0.8933139443397522,0.9935725927352905,1.0516541004180908,0.28168511390686035,319.8121337890625,0.11885906755924225,2.0122017860412598,0.008680449798703194,50.000047683761295 diff --git a/evaluation_outputs/task1_tr_tree/evaluation_train_curve.png b/evaluation_outputs/task1_tr_tree/evaluation_train_curve.png new file mode 100644 index 0000000..59ef4ff Binary files /dev/null and b/evaluation_outputs/task1_tr_tree/evaluation_train_curve.png differ diff --git a/evaluation_outputs/task1_tr_tree/evaluation_val_all_samples.csv b/evaluation_outputs/task1_tr_tree/evaluation_val_all_samples.csv new file mode 100644 index 0000000..ca5a6dc --- /dev/null +++ b/evaluation_outputs/task1_tr_tree/evaluation_val_all_samples.csv @@ -0,0 +1,6 @@ +file_name,frequency_hz,true_rms,pred_rms,relative_error_percent,x_rms,pred_tr,true_tr,dominant_frequency_hz,frequency_squared,inverse_frequency_hz,log_frequency_hz,input_rms,input_peak_abs,input_peak_to_peak,crest_factor,middle_length_ratio,dominant_amplitude,dominant_energy_ratio,harmonic_fit_amplitude,harmonic_fit_residual_ratio,spectral_peak_prominence,half_power_bandwidth_hz,spectral_centroid_hz,signal_mean,sampling_rate +harmonic_5mm_0.75Hz.csv,0.75,0.39182430505752563,0.4018011530100954,2.546255508857473,0.11687792837619781,3.4377846920490267,3.3524234294891357,0.7515671849250793,0.5648532509803772,1.3305530548095703,-0.285594642162323,0.11687792837619781,0.4074699878692627,0.8070399761199951,3.4862868785858154,0.24992701411247253,0.10824807733297348,0.8746803998947144,0.11673416197299957,0.7099294662475586,36.5,0.08761690557003021,0.8937567472457886,-0.001955104758962989,50.000047683761295 +harmonic_5mm_0.95Hz.csv,0.95,0.5053804516792297,0.48249583166035853,4.528196518648947,0.1381060928106308,3.4936607201099394,3.6593642234802246,0.9513587951660156,0.9050835967063904,1.0511281490325928,-0.04986399784684181,0.1381060928106308,0.5125100016593933,1.0134000778198242,3.7109878063201904,0.25,0.09424196928739548,0.7232213020324707,0.11362786591053009,0.8133426308631897,20.55999183654785,0.09677428752183914,1.3476505279541016,0.00015772903861943632,50.000047683761295 +harmonic_5mm_1.85Hz.csv,1.85,1.185911774635315,1.0257513926611004,13.50525270089895,0.6145902872085571,1.6690003308057786,1.9295974969863892,1.8481981754302979,3.4158363342285156,0.5410675406455994,0.6142112016677856,0.6145902872085571,1.3770899772644043,2.6619200706481934,2.240663528442383,0.25,0.7274389863014221,0.9514397978782654,0.789516270160675,0.42080339789390564,85.657958984375,0.09677428752183914,1.8637527227401733,-0.002481284085661173,50.000047683761295 +harmonic_5mm_1.55Hz.csv,1.55,1.4177086353302002,0.7627473327797605,46.198583138198344,0.392645925283432,1.942583084821701,3.610654354095459,1.5492842197418213,2.4002816677093506,0.6454593539237976,0.4377930164337158,0.392645925283432,1.105049967765808,2.1589999198913574,2.8143677711486816,0.25,0.4746793210506439,0.9668325185775757,0.5014731884002686,0.43158969283103943,69.77510833740234,0.09836074709892273,1.584490418434143,-0.0011147483019158244,50.000047683761295 +harmonic_5mm_1.25Hz.csv,1.25,0.3802441656589508,0.6856168561865096,80.30963210135198,0.15971125662326813,4.292852430582046,2.3808226585388184,1.2483434677124023,1.5583614110946655,0.8010615706443787,0.22181743383407593,0.15971125662326813,0.5707299709320068,1.0413799285888672,3.5735113620758057,0.24992592632770538,0.1439778357744217,0.8661479949951172,0.14893445372581482,0.7513272762298584,56.41539764404297,0.05927687883377075,1.4980274438858032,-0.0021986484061926603,50.000047683761295 diff --git a/evaluation_outputs/task1_tr_tree/evaluation_val_curve.png b/evaluation_outputs/task1_tr_tree/evaluation_val_curve.png new file mode 100644 index 0000000..fdcf242 Binary files /dev/null and b/evaluation_outputs/task1_tr_tree/evaluation_val_curve.png differ diff --git a/evaluation_outputs/task1_tr_tree/harmonic_5mm_1.55Hz_prediction.png b/evaluation_outputs/task1_tr_tree/harmonic_5mm_1.55Hz_prediction.png new file mode 100644 index 0000000..f222294 Binary files /dev/null and b/evaluation_outputs/task1_tr_tree/harmonic_5mm_1.55Hz_prediction.png differ diff --git a/scripts_tree/__pycache__/config.cpython-310.pyc b/scripts_tree/__pycache__/config.cpython-310.pyc new file mode 100644 index 0000000..2931abb Binary files /dev/null and b/scripts_tree/__pycache__/config.cpython-310.pyc differ diff --git a/scripts_tree/config.py b/scripts_tree/config.py new file mode 100644 index 0000000..f37fc96 --- /dev/null +++ b/scripts_tree/config.py @@ -0,0 +1,50 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path +import sys + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + +from scripts.config import CORE_FEATURE_NAMES, DataConfig + + +@dataclass +class TreeModelConfig: + candidate_models: tuple[str, ...] = ("extra_trees", "random_forest", "gradient_boosting") + random_state: int = 42 + + +@dataclass +class TrainConfig: + checkpoint_dir: str = "checkpoints_tree" + summary_name: str = "model_selection.csv" + best_model_name: str = "best_tr_tree.pkl" + + +@dataclass +class ExperimentConfig: + data: DataConfig = field(default_factory=DataConfig) + model: TreeModelConfig = field(default_factory=TreeModelConfig) + train: TrainConfig = field(default_factory=TrainConfig) + + +def make_experiment_config() -> ExperimentConfig: + return ExperimentConfig() + + +def checkpoint_dir(config: ExperimentConfig) -> Path: + path = config.data.project_root / config.train.checkpoint_dir / "task1_tr_tree" + path.mkdir(parents=True, exist_ok=True) + return path + + +def evaluation_dir(config: ExperimentConfig) -> Path: + path = config.data.project_root / "evaluation_outputs" / "task1_tr_tree" + path.mkdir(parents=True, exist_ok=True) + return path + + +FEATURE_NAMES = CORE_FEATURE_NAMES diff --git a/scripts_tree/evaluate.py b/scripts_tree/evaluate.py new file mode 100644 index 0000000..35da0c8 --- /dev/null +++ b/scripts_tree/evaluate.py @@ -0,0 +1,177 @@ +from __future__ import annotations + +import argparse +import pickle +from pathlib import Path +import sys + +import matplotlib.pyplot as plt +import pandas as pd + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + +from scripts.dataset import build_datasets, report_to_text + +try: + from .config import FEATURE_NAMES, evaluation_dir, make_experiment_config +except ImportError: + from config import FEATURE_NAMES, evaluation_dir, make_experiment_config + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Evaluate tree model on train/val/test splits.") + parser.add_argument("--split", choices=("train", "val", "test"), default="val") + parser.add_argument("--sample-index", type=int, default=0) + parser.add_argument("--all-samples", action="store_true") + parser.add_argument("--checkpoint", type=str, default=None) + return parser.parse_args() + + +def resolve_checkpoint_path(project_root: Path, checkpoint_arg: str | None) -> Path: + if checkpoint_arg: + return Path(checkpoint_arg).resolve() + return project_root / "checkpoints_tree" / "task1_tr_tree" / "best_tr_tree.pkl" + + +def relative_percent_error(true_value: float, pred_value: float) -> float: + return abs(pred_value - true_value) / max(abs(true_value), 1e-12) * 100.0 + + +def evaluate_record(model, record) -> dict[str, float | str]: + feature_matrix = [record.features.tolist()] + pred_tr = max(float(model.predict(feature_matrix)[0]), 1e-6) + x_rms = float(record.x_rms.item()) + y_rms = float(record.y_rms.item()) + pred_rms = pred_tr * x_rms + row = { + "file_name": record.file_path.name, + "frequency_hz": float(record.frequency_hz), + "true_rms": y_rms, + "pred_rms": pred_rms, + "relative_error_percent": relative_percent_error(y_rms, pred_rms), + "x_rms": x_rms, + "pred_tr": pred_tr, + "true_tr": float(record.target_y_rms.item()), + } + for name, value in zip(FEATURE_NAMES, record.features.tolist()): + row[name] = float(value) + row["sampling_rate"] = float(record.sampling_rate) + return row + + +def save_all_samples_plot(result_df: pd.DataFrame, split: str, save_dir: Path) -> Path | None: + if result_df["frequency_hz"].isna().any(): + return None + plot_df = result_df.sort_values("frequency_hz").reset_index(drop=True) + fig, axes = plt.subplots(2, 1, figsize=(12, 8)) + fig.suptitle(f"Task1 TR Tree Evaluation | {split}") + + axes[0].plot(plot_df["frequency_hz"], plot_df["true_rms"], marker="o", label="True RMS") + axes[0].plot(plot_df["frequency_hz"], plot_df["pred_rms"], marker="o", label="Pred RMS") + axes[0].set_xlabel("Frequency (Hz)") + axes[0].set_ylabel("RMS") + axes[0].grid(True, alpha=0.3) + axes[0].legend() + + axes[1].bar(plot_df["frequency_hz"].astype(str), plot_df["relative_error_percent"], color="tab:orange") + axes[1].set_xlabel("Frequency (Hz)") + axes[1].set_ylabel("Relative Error (%)") + axes[1].grid(True, axis="y", alpha=0.3) + axes[1].tick_params(axis="x", labelrotation=45) + + plt.tight_layout() + figure_path = save_dir / f"evaluation_{split}_curve.png" + plt.savefig(figure_path, dpi=180, bbox_inches="tight") + plt.close(fig) + return figure_path + + +def save_single_sample_plot(record, result: dict[str, float | str], split: str, save_dir: Path, sample_index: int) -> Path: + time_middle = record.time_middle.detach().cpu().numpy().reshape(-1) + x_middle = record.x_middle.detach().cpu().numpy().reshape(-1) + y_middle = record.y_middle.detach().cpu().numpy().reshape(-1) + frequency_hz = float(result["frequency_hz"]) + pred_rms = float(result["pred_rms"]) + true_rms = float(result["true_rms"]) + + fig, axes = plt.subplots(3, 1, figsize=(12, 10)) + fig.suptitle(f"Task1 TR Tree | {split} | {result['file_name']}") + + axes[0].plot(time_middle, x_middle, color="tab:blue") + axes[0].set_title("Input Base Excitation (Middle Segment)") + axes[0].set_xlabel("Time") + axes[0].set_ylabel("Acceleration") + axes[0].grid(True, alpha=0.3) + + axes[1].plot(time_middle, y_middle, color="tab:green") + axes[1].set_title("True Top Response (Middle Segment)") + axes[1].set_xlabel("Time") + axes[1].set_ylabel("Acceleration") + axes[1].grid(True, alpha=0.3) + + axes[2].bar(["True RMS", "Pred RMS"], [true_rms, pred_rms], color=["tab:green", "tab:orange"]) + axes[2].set_title( + f"Freq: {frequency_hz:.4f} Hz | True RMS: {true_rms:.6f} | Pred RMS: {pred_rms:.6f} | " + f"Error: {float(result['relative_error_percent']):.2f}%" + ) + axes[2].set_ylabel("RMS") + axes[2].grid(True, axis="y", alpha=0.3) + + plt.tight_layout() + figure_path = save_dir / f"evaluation_{split}_s{sample_index}.png" + plt.savefig(figure_path, dpi=180, bbox_inches="tight") + plt.close(fig) + return figure_path + + +def main() -> None: + args = parse_args() + config = make_experiment_config() + ckpt_path = resolve_checkpoint_path(config.data.project_root, args.checkpoint) + with ckpt_path.open("rb") as handle: + bundle = pickle.load(handle) + model = bundle["model"] + + _, raw_records, reports = build_datasets(config.data) + print(report_to_text(reports)) + records = raw_records[args.split] + save_dir = evaluation_dir(config) + + if args.all_samples: + rows = [evaluate_record(model, record) for record in records] + result_df = pd.DataFrame(rows).sort_values(["relative_error_percent", "file_name"]).reset_index(drop=True) + csv_path = save_dir / f"evaluation_{args.split}_all_samples.csv" + result_df.to_csv(csv_path, index=False) + fig_path = save_all_samples_plot(result_df, args.split, save_dir) + summary = { + "count": len(result_df), + "mean_error_percent": float(result_df["relative_error_percent"].mean()), + "median_error_percent": float(result_df["relative_error_percent"].median()), + "max_error_percent": float(result_df["relative_error_percent"].max()), + "min_error_percent": float(result_df["relative_error_percent"].min()), + } + print(f"Checkpoint: {ckpt_path}") + print(f"Model: {bundle['model_name']}") + print(f"Summary: {summary}") + print(f"CSV saved to: {csv_path}") + if fig_path is not None: + print(f"Figure saved to: {fig_path}") + print(result_df.to_string(index=False)) + return + + record = records[args.sample_index] + result = evaluate_record(model, record) + fig_path = save_single_sample_plot(record, result, args.split, save_dir, args.sample_index) + print(f"Checkpoint: {ckpt_path}") + print(f"Model: {bundle['model_name']}") + print(f"Sample file: {result['file_name']}") + print(f"True RMS: {float(result['true_rms']):.6f}") + print(f"Pred RMS: {float(result['pred_rms']):.6f}") + print(f"Relative RMS Error (%): {float(result['relative_error_percent']):.4f}") + print(f"Figure saved to: {fig_path}") + + +if __name__ == "__main__": + main() diff --git a/scripts_tree/predict_single.py b/scripts_tree/predict_single.py new file mode 100644 index 0000000..cab8032 --- /dev/null +++ b/scripts_tree/predict_single.py @@ -0,0 +1,115 @@ +from __future__ import annotations + +import argparse +import pickle +from pathlib import Path +import sys + +import matplotlib.pyplot as plt +import numpy as np + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + +from scripts.dataset import build_record_from_file + +try: + from .config import FEATURE_NAMES, evaluation_dir, make_experiment_config +except ImportError: + from config import FEATURE_NAMES, evaluation_dir, make_experiment_config + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Predict task1 RMS from a single waveform CSV with tree model.") + parser.add_argument("--file", type=str, required=True) + parser.add_argument("--checkpoint", type=str, default=None) + return parser.parse_args() + + +def resolve_checkpoint_path(project_root: Path, checkpoint_arg: str | None) -> Path: + if checkpoint_arg: + return Path(checkpoint_arg).resolve() + return project_root / "checkpoints_tree" / "task1_tr_tree" / "best_tr_tree.pkl" + + +def save_prediction_figure(file_path: Path, record, pred_rms: float, true_rms: float | None, output_dir: Path) -> Path: + time_middle = record.time_middle.detach().cpu().numpy().reshape(-1) + x_middle = record.x_middle.detach().cpu().numpy().reshape(-1) + sampling_rate = record.sampling_rate + fft_values = np.fft.rfft(x_middle) + freqs = np.fft.rfftfreq(x_middle.size, d=1.0 / sampling_rate) + amplitudes = np.abs(fft_values) + + fig, axes = plt.subplots(3, 1, figsize=(12, 10)) + fig.suptitle(f"Task1 TR Tree Single-File Prediction | {file_path.name}") + + axes[0].plot(time_middle, x_middle, color="tab:blue") + axes[0].set_title("Input Base Excitation (Middle Segment)") + axes[0].set_xlabel("Time") + axes[0].set_ylabel("Acceleration") + axes[0].grid(True, alpha=0.3) + + axes[1].plot(freqs, amplitudes, color="tab:purple") + axes[1].axvline(record.frequency_hz, color="tab:red", linestyle="--", label=f"Dominant freq = {record.frequency_hz:.4f} Hz") + axes[1].set_xlim(0.0, 5.0) + axes[1].set_title("Input Spectrum") + axes[1].set_xlabel("Frequency (Hz)") + axes[1].set_ylabel("Amplitude") + axes[1].grid(True, alpha=0.3) + axes[1].legend() + + labels = ["Pred RMS"] if true_rms is None else ["True RMS", "Pred RMS"] + values = [pred_rms] if true_rms is None else [true_rms, pred_rms] + colors = ["tab:orange"] if true_rms is None else ["tab:green", "tab:orange"] + axes[2].bar(labels, values, color=colors) + title = f"Predicted RMS = {pred_rms:.6f}" + if true_rms is not None: + error_percent = abs(pred_rms - true_rms) / max(abs(true_rms), 1e-12) * 100.0 + title = f"True RMS = {true_rms:.6f} | Pred RMS = {pred_rms:.6f} | Error = {error_percent:.2f}%" + axes[2].set_title(title) + axes[2].set_ylabel("RMS") + axes[2].grid(True, axis="y", alpha=0.3) + + plt.tight_layout() + output_dir.mkdir(parents=True, exist_ok=True) + figure_path = output_dir / f"{file_path.stem}_prediction.png" + plt.savefig(figure_path, dpi=180, bbox_inches="tight") + plt.close(fig) + return figure_path + + +def main() -> None: + args = parse_args() + config = make_experiment_config() + ckpt_path = resolve_checkpoint_path(config.data.project_root, args.checkpoint) + with ckpt_path.open("rb") as handle: + bundle = pickle.load(handle) + model = bundle["model"] + + file_path = Path(args.file).resolve() + record = build_record_from_file(file_path, split="predict", config=config.data) + pred_tr = max(float(model.predict([record.features.tolist()])[0]), 1e-6) + pred_rms = pred_tr * float(record.x_rms.item()) + true_rms = float(record.y_rms.item()) if record.y_rms.numel() > 0 else None + + out_dir = evaluation_dir(config) + fig_path = save_prediction_figure(file_path, record, pred_rms, true_rms, out_dir) + + print(f"Checkpoint: {ckpt_path}") + print(f"Model: {bundle['model_name']}") + print(f"Input file: {file_path}") + print(f"Extracted frequency (Hz): {record.frequency_hz:.6f}") + for feature_name, feature_value in zip(FEATURE_NAMES, record.features.tolist()): + print(f"{feature_name}: {feature_value:.6f}") + print(f"Predicted TR: {pred_tr:.6f}") + print(f"Predicted RMS: {pred_rms:.6f}") + if true_rms is not None: + error_percent = abs(pred_rms - true_rms) / max(abs(true_rms), 1e-12) * 100.0 + print(f"True RMS: {true_rms:.6f}") + print(f"Relative RMS Error (%): {error_percent:.4f}") + print(f"Figure saved to: {fig_path}") + + +if __name__ == "__main__": + main() diff --git a/scripts_tree/train.py b/scripts_tree/train.py new file mode 100644 index 0000000..aaa0bbf --- /dev/null +++ b/scripts_tree/train.py @@ -0,0 +1,160 @@ +from __future__ import annotations + +import pickle +from dataclasses import asdict +from pathlib import Path +from typing import Any +import sys + +import pandas as pd +from sklearn.ensemble import ExtraTreesRegressor, GradientBoostingRegressor, RandomForestRegressor + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +if str(PROJECT_ROOT) not in sys.path: + sys.path.insert(0, str(PROJECT_ROOT)) + +from scripts.dataset import build_datasets, report_to_text + +try: + from .config import FEATURE_NAMES, ExperimentConfig, checkpoint_dir, make_experiment_config +except ImportError: + from config import FEATURE_NAMES, ExperimentConfig, checkpoint_dir, make_experiment_config + + +def build_regressor(model_name: str, random_state: int): + if model_name == "extra_trees": + return ExtraTreesRegressor( + n_estimators=600, + max_depth=None, + min_samples_leaf=1, + min_samples_split=2, + random_state=random_state, + ) + if model_name == "random_forest": + return RandomForestRegressor( + n_estimators=500, + max_depth=None, + min_samples_leaf=1, + min_samples_split=2, + random_state=random_state, + ) + if model_name == "gradient_boosting": + return GradientBoostingRegressor( + n_estimators=300, + learning_rate=0.03, + max_depth=3, + random_state=random_state, + loss="huber", + ) + raise ValueError(f"Unsupported model: {model_name}") + + +def records_to_frame(records: list[Any]) -> pd.DataFrame: + rows = [] + for record in records: + row = { + "file_name": record.file_path.name, + "frequency_hz": float(record.frequency_hz), + "x_rms": float(record.x_rms.item()), + "y_rms": float(record.y_rms.item()), + "target_tr": float(record.target_y_rms.item()), + } + for name, value in zip(FEATURE_NAMES, record.features.tolist()): + row[name] = float(value) + rows.append(row) + return pd.DataFrame(rows) + + +def evaluate_frame(model, frame: pd.DataFrame) -> dict[str, float]: + feature_matrix = frame.loc[:, FEATURE_NAMES].to_numpy() + pred_tr = model.predict(feature_matrix) + pred_tr = pred_tr.clip(min=1e-6) + pred_rms = pred_tr * frame["x_rms"].to_numpy() + true_rms = frame["y_rms"].to_numpy() + relative_error = abs(pred_rms - true_rms) / true_rms.clip(min=1e-6) + return { + "mean_rms_error": float(relative_error.mean()), + "median_rms_error": float(pd.Series(relative_error).median()), + "max_rms_error": float(relative_error.max()), + } + + +def serialize_for_checkpoint(value: Any) -> Any: + if isinstance(value, Path): + return str(value) + if isinstance(value, dict): + return {key: serialize_for_checkpoint(sub_value) for key, sub_value in value.items()} + if isinstance(value, tuple): + return [serialize_for_checkpoint(item) for item in value] + if isinstance(value, list): + return [serialize_for_checkpoint(item) for item in value] + return value + + +def train() -> None: + config = make_experiment_config() + datasets, raw_records, reports = build_datasets(config.data) + del datasets + print(report_to_text(reports)) + + train_df = records_to_frame(raw_records["train"]) + val_df = records_to_frame(raw_records["val"]) + + feature_matrix = train_df.loc[:, FEATURE_NAMES].to_numpy() + target = train_df["target_tr"].to_numpy() + + results: list[dict[str, Any]] = [] + best_bundle: dict[str, Any] | None = None + best_val_error = float("inf") + + for model_name in config.model.candidate_models: + model = build_regressor(model_name, config.model.random_state) + model.fit(feature_matrix, target) + + train_metrics = evaluate_frame(model, train_df) + val_metrics = evaluate_frame(model, val_df) + row = { + "model_name": model_name, + "train_mean_rms_error": train_metrics["mean_rms_error"], + "train_median_rms_error": train_metrics["median_rms_error"], + "train_max_rms_error": train_metrics["max_rms_error"], + "val_mean_rms_error": val_metrics["mean_rms_error"], + "val_median_rms_error": val_metrics["median_rms_error"], + "val_max_rms_error": val_metrics["max_rms_error"], + } + results.append(row) + print( + f"{model_name}: train_mean={train_metrics['mean_rms_error']:.6f} | " + f"val_mean={val_metrics['mean_rms_error']:.6f}" + ) + + if val_metrics["mean_rms_error"] < best_val_error: + best_val_error = val_metrics["mean_rms_error"] + best_bundle = { + "model_name": model_name, + "model": model, + "config": serialize_for_checkpoint(asdict(config)), + "feature_names": list(FEATURE_NAMES), + "val_metrics": val_metrics, + "train_metrics": train_metrics, + } + + summary_df = pd.DataFrame(results).sort_values("val_mean_rms_error").reset_index(drop=True) + ckpt_dir = checkpoint_dir(config) + summary_path = ckpt_dir / config.train.summary_name + summary_df.to_csv(summary_path, index=False) + print(f"Model selection saved to: {summary_path}") + print(summary_df.to_string(index=False)) + + if best_bundle is None: + raise RuntimeError("No valid model trained.") + + best_path = ckpt_dir / config.train.best_model_name + with best_path.open("wb") as handle: + pickle.dump(best_bundle, handle) + print(f"Best tree model saved to: {best_path}") + print(f"Best model: {best_bundle['model_name']} | val_mean_rms_error={best_val_error:.6f}") + + +if __name__ == "__main__": + train()