diff --git a/.gitignore b/.gitignore index 40451d7..e3cc06c 100644 --- a/.gitignore +++ b/.gitignore @@ -1,4 +1,4 @@ -.venv/ .venv310/ .vscode/ -src/__pycache__/ \ No newline at end of file +scripts/__pycache__/ +old/ diff --git a/Figure_1.png b/Figure_1.png deleted file mode 100644 index 3f550d9..0000000 Binary files a/Figure_1.png and /dev/null differ diff --git a/checkpoints/forward/best_tcn_model.pt b/checkpoints/forward/best_tcn_model.pt deleted file mode 100644 index c6235b3..0000000 Binary files a/checkpoints/forward/best_tcn_model.pt and /dev/null differ diff --git a/checkpoints/forward/training_history.csv b/checkpoints/forward/training_history.csv deleted file mode 100644 index c12ced1..0000000 --- a/checkpoints/forward/training_history.csv +++ /dev/null @@ -1,11 +0,0 @@ -epoch,lr,train_loss,train_time_loss,train_weighted_time_loss,train_fft_loss,train_rms_loss,train_scale_loss,train_underestimate_loss,train_envelope_loss,train_rms_error,train_dominant_freq_error,val_loss,val_time_loss,val_weighted_time_loss,val_fft_loss,val_rms_loss,val_scale_loss,val_underestimate_loss,val_envelope_loss,val_rms_error,val_dominant_freq_error -1,0.001,25.599940884788083,1.7562742283884085,0.1756274257347269,2774.0617503040244,0.48780856886,0.4243854204157613,0.1559595560549565,0.8805500887474924,0.48780856886,0.5304876364710495,11.634879671294113,0.5565754904829222,0.05565755005026686,1030.2479747903758,0.19599251089424924,0.30123303316790484,0.07480601036262795,0.6266409343686598,0.19599251089424924,1.0759024783184585 -2,0.001,18.8950198911271,1.807937761522689,0.18079377879511635,1975.3159283332104,0.3459213998801303,0.35067820830165214,0.1092575388907824,0.758364588865694,0.3459213998801303,0.4821035959354592,10.542049144876414,0.5095152341086289,0.05095152441283752,783.4861005585769,0.16758740950247336,0.2732687505154774,0.10411370940634916,0.6904992646184461,0.16758740950247336,0.9969777073497166 -3,0.001,16.69826183678969,1.7823544706938401,0.17823545016207784,1672.4529568654186,0.30989001663226,0.3281119924108937,0.10324830568905147,0.7008864249823228,0.30989001663226,0.5562505104268722,11.536806665617844,0.5393576467859333,0.05393576609163449,807.3707762093379,0.21286767027501402,0.33425827504232014,0.10436723167718999,0.6995668801768072,0.21286767027501402,1.0535930765001758 -4,0.001,14.293646308611024,1.7581631035174963,0.1758163128540201,1394.1521462494472,0.25787185048157313,0.2936001386282579,0.08524404234200153,0.6856944434485346,0.25787185048157313,0.4707189468993712,11.0088351677204,0.512058621850507,0.05120586244196727,716.754319289635,0.18768149871250678,0.30934940792363264,0.12457355485972145,0.7022169437901727,0.18768149871250678,0.9801404230908486 -5,0.001,13.64289491581467,1.7736638503254585,0.1773663880127781,1282.9595152656987,0.2586157521549261,0.2966467724093851,0.0801986690856657,0.662726102291413,0.2586157521549261,0.4816929557024733,12.69674718791041,0.5321119869577473,0.0532112003400408,923.869591548525,0.2018765033832912,0.3501281591838804,0.1300795058286267,0.763439338782738,0.2018765033832912,0.8944807380729035 -6,0.001,11.776648678869572,1.7715154886245728,0.17715155192703572,1060.149999654518,0.21744662122625225,0.2669158032480276,0.07291070573067046,0.6286382666736279,0.21744662122625225,0.49275163173347086,13.566723692006079,0.5599699513665561,0.055996995740409554,1081.8859031940328,0.23347977780062576,0.399347482827203,0.09571884981966738,0.753588840879243,0.23347977780062576,0.802507071705832 -7,0.001,10.598831523139522,1.8021629632643934,0.180216298631902,959.8165948256006,0.1918528508746399,0.24521326632151064,0.06225078732197015,0.5634953882896675,0.1918528508746399,0.43498506576279067,12.335502788938324,0.5373578256574171,0.053735782938270735,897.6921434073613,0.22004679693230267,0.3851976631016567,0.09088767872288309,0.7485740184783936,0.22004679693230267,0.7111648032890617 -8,0.0005,10.096835316352124,1.771101622086651,0.17711016507643573,904.1345197569649,0.18218446649470418,0.2265465495721349,0.06017633231427028,0.5831915904890816,0.18218446649470418,0.4688892364850107,12.481576097422632,0.55741053922423,0.05574105519416003,993.8351429906385,0.20113462980451255,0.37394588877414836,0.08785425563310754,0.7251736295634302,0.20113462980451255,0.7193729794563382 -9,0.0005,9.301918011791301,1.8092092241881028,0.18092092534281173,872.4488234609928,0.15496070823579464,0.2050890047454609,0.05094481654078612,0.5316191723324218,0.15496070823579464,0.45952377956299667,12.33887189010094,0.529184412339638,0.052918442107480146,888.4320715542498,0.2059096264941939,0.35693827361382285,0.11748063429820768,0.7306599390917811,0.2059096264941939,0.7187415812949122 -10,0.0005,9.509747361237148,1.775009681593697,0.17750097143481364,875.1525979671839,0.15211023376235422,0.21900151374767413,0.054768192924488826,0.5556492130711393,0.15211023376235422,0.48492007759302524,11.80303616359316,0.5109482181483301,0.05109482271404102,766.1427895118451,0.20308723608995305,0.33445684663180647,0.13568100256138835,0.7249893361124499,0.20308723608995305,0.7126380656820888 diff --git a/checkpoints_mlp/task1_feature_mlp/best_feature_mlp.pt b/checkpoints_mlp/task1_feature_mlp/best_feature_mlp.pt deleted file mode 100644 index 4e72542..0000000 Binary files a/checkpoints_mlp/task1_feature_mlp/best_feature_mlp.pt and /dev/null differ diff --git a/checkpoints_mlp/task1_feature_mlp/training_history.csv b/checkpoints_mlp/task1_feature_mlp/training_history.csv deleted file mode 100644 index 2ad8868..0000000 --- a/checkpoints_mlp/task1_feature_mlp/training_history.csv +++ /dev/null @@ -1,59 +0,0 @@ -epoch,lr,train_loss,train_relative_loss,train_log_rms_huber_loss,train_mae_loss,train_rms_error,val_loss,val_relative_loss,val_log_rms_huber_loss,val_mae_loss,val_rms_error -1,0.001,1.0684084097544353,0.6919674078623453,0.19923226369751823,1.5134448210398357,0.6919673813713921,0.3821044862270355,0.2511320114135742,0.033154696226119995,0.7073763012886047,0.2511320412158966 -2,0.001,1.0431654651959736,0.6733311547173394,0.19394073883692423,1.49585849708981,0.673331167962816,0.3861435353755951,0.2520703375339508,0.0340907983481884,0.7233673334121704,0.2520703673362732 -3,0.001,1.0058492289649115,0.6443076067500644,0.18731229503949484,1.473716417948405,0.6443075670136346,0.39013171195983887,0.252702534198761,0.03538244962692261,0.7392822504043579,0.2527025640010834 -4,0.001,0.9792734649446275,0.6221125589476691,0.18307040548986858,1.4657206005520291,0.6221125887499915,0.394980788230896,0.253364622592926,0.037405483424663544,0.7570802569389343,0.253364622592926 -5,0.001,0.939600666364034,0.5912809504403008,0.17598583300908408,1.4422023958630033,0.5912809636857774,0.3971262574195862,0.2521932125091553,0.03932402655482292,0.7696000933647156,0.2521932125091553 -6,0.001,0.9127781920962863,0.5688717034127977,0.17321109771728516,1.4266544183095295,0.5688716835445828,0.3954552412033081,0.24867530167102814,0.040780965238809586,0.7746281623840332,0.24867533147335052 -7,0.001,0.8162693844901191,0.4975863430235121,0.15291899773809645,1.3599585427178278,0.4975863430235121,0.39774268865585327,0.24631300568580627,0.044477906078100204,0.7871417999267578,0.24631300568580627 -8,0.001,0.8098740312788222,0.48419485489527386,0.16087035338083902,1.3668427334891424,0.4841948416497972,0.40212470293045044,0.24344877898693085,0.05097588896751404,0.8029600381851196,0.24344877898693085 -9,0.001,0.7296120656861199,0.4106071988741557,0.16139008270369637,1.3197486003239949,0.4106071756945716,0.43612974882125854,0.2614344656467438,0.060854148119688034,0.8603642582893372,0.2614344358444214 -10,0.001,0.7408013741175333,0.38955933849016827,0.2076671322186788,1.3032778369055853,0.3895593351787991,0.4869456887245178,0.29191675782203674,0.07229459285736084,0.9387197494506836,0.29191672801971436 -11,0.001,0.7496903538703918,0.3897787862353855,0.21981331043773228,1.3003438181347318,0.3897787994808621,0.5376651287078857,0.32053330540657043,0.0865083634853363,1.0150035619735718,0.32053327560424805 -12,0.001,0.9797695477803549,0.38233999411265057,0.5365811387697855,1.29995772573683,0.3823399974240197,0.5587358474731445,0.3330889940261841,0.09165894240140915,1.0460175275802612,0.3330889642238617 -13,0.001,0.7886577645937601,0.39104337162441677,0.27514297929075027,1.2750475406646729,0.39104337493578595,0.5636787414550781,0.3352062702178955,0.09356075525283813,1.055345892906189,0.3352062702178955 -14,0.001,0.6367178559303284,0.34643163283665973,0.1475350492530399,1.1975661383734808,0.3464316460821364,0.5635904669761658,0.3346971571445465,0.09383875876665115,1.0567615032196045,0.3346971571445465 -15,0.001,0.6623357137044271,0.36682088838683236,0.15060860332515505,1.2170556386311848,0.3668208916982015,0.5532020926475525,0.32986027002334595,0.08957388252019882,1.0410760641098022,0.32986021041870117 -16,0.001,0.6297085682551066,0.36428216265307534,0.12578680531846154,1.140575302971734,0.36428215106328327,0.5544819235801697,0.33266526460647583,0.08754842728376389,1.041035532951355,0.33266523480415344 -17,0.001,0.6100040011935763,0.36249567733870613,0.10765769663784239,1.111767013867696,0.36249566078186035,0.5716693997383118,0.34491467475891113,0.08935325592756271,1.0649319887161255,0.34491464495658875 -18,0.001,0.5903598533736335,0.3496483862400055,0.10458873874611324,1.0817992157406278,0.34964838955137467,0.601777195930481,0.3637298047542572,0.0955144539475441,1.109410285949707,0.3637298047542572 -19,0.001,0.5701847804917229,0.3414529164632161,0.09985838168197209,1.0255872938368056,0.3414529164632161,0.629698634147644,0.3823082447052002,0.10031794756650925,1.1476796865463257,0.3823082447052002 -20,0.001,0.5357944567998251,0.3132438593440586,0.09814917544523875,0.9929247697194418,0.3132438593440586,0.6656612157821655,0.4028518795967102,0.11059614270925522,1.199081540107727,0.4028518795967102 -21,0.001,0.5140700340270996,0.29829433891508317,0.1006597230831782,0.9352060423956977,0.2982943256696065,0.6950218081474304,0.4190041720867157,0.11983328312635422,1.2409510612487793,0.4190041720867157 -22,0.001,0.5517757270071242,0.31282461020681596,0.12161111583312352,0.9849517875247531,0.31282461020681596,0.6951501965522766,0.42067644000053406,0.11815791577100754,1.2390354871749878,0.4206763803958893 -23,0.001,0.5385147333145142,0.3098441825972663,0.12465377317534553,0.9012014071146647,0.3098441743188434,0.6760208606719971,0.41229119896888733,0.10967614501714706,1.209816813468933,0.41229113936424255 -24,0.001,0.4948686576551861,0.2839176009098689,0.10664872804449664,0.873096740908093,0.283917604221238,0.6479407548904419,0.3981650471687317,0.0997471809387207,1.1664352416992188,0.3981650173664093 -25,0.001,0.4815748400158352,0.2894285586145189,0.08536433428525925,0.8541535006629096,0.2894285437133577,0.628013014793396,0.3875465989112854,0.09387131035327911,1.133752703666687,0.3875465989112854 -26,0.001,0.4902527266078525,0.3052508632342021,0.08093913561768001,0.8286499811543359,0.3052508615785175,0.6259267330169678,0.3861921727657318,0.09331636875867844,1.1316486597061157,0.3861921429634094 -27,0.001,0.4992575960026847,0.3134735193517473,0.0804886631667614,0.8361171748903062,0.31347350279490155,0.6494007110595703,0.399603933095932,0.09916889667510986,1.1694673299789429,0.3996039032936096 -28,0.001,0.45642075273725724,0.2935712155368593,0.07241135980519983,0.7236066361268362,0.29357120229138267,0.6795501708984375,0.4163399338722229,0.10745861381292343,1.2174421548843384,0.4163399338722229 -29,0.0005,0.4504433075586955,0.27966498997476363,0.08081474165535635,0.7344483600722419,0.2796649800406562,0.712536096572876,0.4332484304904938,0.11834639310836792,1.2701857089996338,0.4332483410835266 -30,0.0005,0.4445747633775075,0.2770162257883284,0.0779732180138429,0.7271908124287924,0.27701622247695923,0.7285711765289307,0.44202420115470886,0.12342425435781479,1.293191909790039,0.44202420115470886 -31,0.0005,0.47021543317370945,0.29266008569134605,0.08219371032383707,0.7727337082227071,0.29266009893682265,0.7398308515548706,0.4480714797973633,0.12714092433452606,1.3093578815460205,0.4480714499950409 -32,0.0005,0.4485766556527879,0.27516163720024955,0.07896825671195984,0.7612588869200813,0.27516160408655804,0.7477874755859375,0.45214587450027466,0.1301499605178833,1.3201942443847656,0.45214587450027466 -33,0.0005,0.4806170066197713,0.29431821240319145,0.0897826933198505,0.7930783894326952,0.29431822564866805,0.7456773519515991,0.45039525628089905,0.13018397986888885,1.3176275491714478,0.4503951966762543 -34,0.0005,0.4639296531677246,0.27904501888487077,0.09117302215761608,0.7766991588804457,0.2790450221962399,0.7344075441360474,0.4435312747955322,0.12712228298187256,1.3035639524459839,0.44353124499320984 -35,0.0005,0.4217291673024495,0.25814421640502083,0.07712550130155352,0.704938703113132,0.2581442031595442,0.728592038154602,0.43990880250930786,0.12561193108558655,1.2964953184127808,0.4399087429046631 -36,0.0005,0.4452095528443654,0.2780698604053921,0.07377014433344205,0.7454138000806173,0.2780698537826538,0.719640851020813,0.4346870481967926,0.12287390232086182,1.285322666168213,0.4346870481967926 -37,0.0005,0.4346568849351671,0.27202920781241524,0.0730870481994417,0.7187491787804497,0.27202921443515354,0.7090620994567871,0.42957648634910583,0.11880777031183243,1.2691987752914429,0.42957648634910583 -38,0.0005,0.45190709829330444,0.2777327348788579,0.07873265279663934,0.7674991223547194,0.2777327398459117,0.7021456360816956,0.42580172419548035,0.11658556759357452,1.2593647241592407,0.42580166459083557 -39,0.0005,0.41501253181033665,0.25652576155132717,0.06958397726217906,0.7086585627661811,0.2565257747968038,0.6967095136642456,0.4226696193218231,0.11500495672225952,1.2519077062606812,0.42266955971717834 -40,0.0005,0.426217923561732,0.2659038090043598,0.06590471830632952,0.739237109820048,0.26590381893846726,0.6904255747795105,0.4193277060985565,0.11303865909576416,1.2421258687973022,0.41932764649391174 -41,0.0005,0.474410249127282,0.302842206425137,0.07200946576065487,0.7837395668029785,0.302842206425137,0.6886562705039978,0.41840869188308716,0.11255472153425217,1.2388770580291748,0.41840869188308716 -42,0.0005,0.4341147674454583,0.2751454992426766,0.06902315095067024,0.7146794034375085,0.2751455108324687,0.6964045763015747,0.42251697182655334,0.11516380310058594,1.2500985860824585,0.42251691222190857 -43,0.0005,0.40912215577231514,0.2595863143603007,0.058075524038738675,0.7065279748704698,0.2595863276057773,0.709182620048523,0.430009663105011,0.11880651861429214,1.2671202421188354,0.430009663105011 -44,0.0005,0.44166630175378585,0.27680986291832393,0.07062854783402549,0.7459001806047227,0.2768098645740085,0.7103798389434814,0.4306167662143707,0.11920434236526489,1.2690653800964355,0.4306167662143707 -45,0.0005,0.4364900390307109,0.2806501239538193,0.06813105609681872,0.6982774602042304,0.28065012726518845,0.713029146194458,0.43142566084861755,0.12053372710943222,1.274687647819519,0.43142572045326233 -46,0.0005,0.42031361990504795,0.2599939935737186,0.0675199499560727,0.7311977677875094,0.2599939935737186,0.7204116582870483,0.43440350890159607,0.12379532307386398,1.287744164466858,0.43440350890159607 -47,0.0005,0.41649588611390853,0.2588261928823259,0.06496965015927951,0.7262829906410642,0.2588262077834871,0.7308988571166992,0.43982744216918945,0.12745337188243866,1.3032091856002808,0.43982744216918945 -48,0.0005,0.45140007469389176,0.2842165165477329,0.07611107577880223,0.734001550409529,0.2842165331045787,0.738196074962616,0.4431730806827545,0.13040359318256378,1.314801573753357,0.44317302107810974 -49,0.0005,0.4124910682439804,0.25314582718743217,0.068223740077681,0.7211828695403205,0.25314583049880135,0.7486779689788818,0.4491503834724426,0.1336752474308014,1.3284742832183838,0.4491503834724426 -50,0.00025,0.41983669333987766,0.2572576337390476,0.07067853129572338,0.7304676373799642,0.25725764367315507,0.7499342560768127,0.4504518508911133,0.1333393007516861,1.32985258102417,0.4504518508911133 -51,0.00025,0.39338985085487366,0.24541797240575156,0.061128986792431936,0.6808343132336935,0.2454179906182819,0.7484983205795288,0.4500373899936676,0.1325150579214096,1.3271640539169312,0.4500373899936676 -52,0.00025,0.4347236222691006,0.26843487554126316,0.07120031118392944,0.7525900999704996,0.2684348738855786,0.7462888956069946,0.44955143332481384,0.1310974657535553,1.3227624893188477,0.44955143332481384 -53,0.00025,0.39689934915966457,0.24233278632164001,0.06171891983184549,0.721849156750573,0.24233278466595543,0.7399852871894836,0.44651278853416443,0.1285889446735382,1.3135385513305664,0.44651278853416443 -54,0.00025,0.37512414654095966,0.23281626568900216,0.05609904395209418,0.6682239903344048,0.23281625906626383,0.7348796129226685,0.44354644417762756,0.12702621519565582,1.307090163230896,0.44354644417762756 -55,0.00025,0.42326178153355914,0.271195156706704,0.06771800035817756,0.6751875513129764,0.2711951591902309,0.731861412525177,0.4418794810771942,0.1260313242673874,1.3030563592910767,0.4418794810771942 -56,0.00025,0.40027793248494464,0.2503356287876765,0.0625538213385476,0.6868461138672299,0.25033563044336105,0.7314411997795105,0.44189128279685974,0.12566323578357697,1.3020164966583252,0.44189128279685974 -57,0.00025,0.41923924618297154,0.2631028691927592,0.06492394229604138,0.7162894474135505,0.2631028923723433,0.7334963083267212,0.4432640075683594,0.12621350586414337,1.30381441116333,0.443263977766037 -58,0.00025,0.4435661964946323,0.2827550404601627,0.06622585654258728,0.7409449021021525,0.2827550503942702,0.7361425757408142,0.44472166895866394,0.12714305520057678,1.3070908784866333,0.44472166895866394 diff --git a/checkpoints_rms/forward_rms/best_rms_model.pt b/checkpoints_rms/forward_rms/best_rms_model.pt deleted file mode 100644 index 2230b5c..0000000 Binary files a/checkpoints_rms/forward_rms/best_rms_model.pt and /dev/null differ diff --git a/checkpoints_rms/forward_rms/training_history.csv b/checkpoints_rms/forward_rms/training_history.csv deleted file mode 100644 index bf44b5d..0000000 --- a/checkpoints_rms/forward_rms/training_history.csv +++ /dev/null @@ -1,30 +0,0 @@ -epoch,lr,train_loss,train_relative_loss,train_log_loss,train_mae_loss,train_waveform_l1,train_waveform_huber,train_rms_error,val_loss,val_relative_loss,val_log_loss,val_mae_loss,val_waveform_l1,val_waveform_huber,val_rms_error -1,0.001,0.3693518406814999,0.24889015323585933,0.05584108498361376,0.25861359967125785,0.5032549103101095,0.25580188632011414,0.24889015323585933,0.4390127956867218,0.30896639823913574,0.056898053735494614,0.3188031315803528,0.44043534994125366,0.17367032170295715,0.30896639823913574 -2,0.001,0.3945225642787086,0.2697031597296397,0.06360895559191704,0.27172964480188155,0.4614827699131436,0.224760792321629,0.2697031597296397,0.36929574608802795,0.2602313160896301,0.0452926866710186,0.2584935128688812,0.4390406906604767,0.17246951162815094,0.2602313160896301 -3,0.001,0.5084060496754117,0.3563646740383572,0.09457878602875604,0.32678870028919643,0.4492427276240455,0.19155073165893555,0.3563646740383572,0.30823153257369995,0.2178763449192047,0.03594436123967171,0.20229224860668182,0.4395522177219391,0.1724669635295868,0.2178763449192047 -4,0.001,0.46196970012452865,0.2874237828784519,0.10073509646786584,0.3864523735311296,0.4910557005140517,0.2566720247268677,0.2874237828784519,0.2588622272014618,0.1819354146718979,0.02867879904806614,0.16302458941936493,0.44016343355178833,0.1725275069475174,0.1819354146718979 -5,0.001,0.4498247702916463,0.2731182641453213,0.0863925533162223,0.3938671946525574,0.6069219443533156,0.33671536213821834,0.2731182641453213,0.3132787346839905,0.22292804718017578,0.02889658883213997,0.2158222645521164,0.44235706329345703,0.17352235317230225,0.22292804718017578 -6,0.001,0.42692894405788845,0.28966741760571796,0.05804820607105891,0.3025594221221076,0.5698766128884422,0.3100252277735207,0.28966741760571796,0.4311088025569916,0.3014686703681946,0.060949672013521194,0.30897605419158936,0.4416200518608093,0.17345383763313293,0.3014686703681946 -7,0.001,0.36648913555675083,0.2673240436447991,0.04492858507566982,0.2210414773888058,0.40499014324612087,0.1858146521780226,0.2673240436447991,0.3253621459007263,0.2316536009311676,0.031499434262514114,0.22455067932605743,0.43930211663246155,0.17284171283245087,0.2316536009311676 -8,0.001,0.3705151147312588,0.2639704677793715,0.056356401907073125,0.23181239929464129,0.40243885583347744,0.1668036257227262,0.2639704677793715,0.2693358361721039,0.19200395047664642,0.023184901103377342,0.17564474046230316,0.43946224451065063,0.17288772761821747,0.19200395047664642 -9,0.001,0.3809369206428528,0.238957146803538,0.05438671095503701,0.32686107357343036,0.5658377408981323,0.32192010349697536,0.238957146803538,0.22007109224796295,0.15654480457305908,0.01698182336986065,0.1330064833164215,0.43889322876930237,0.17233915627002716,0.15654480457305908 -10,0.001,0.3033063875304328,0.20820601118935478,0.03028729951216115,0.2125823481215371,0.4863186809751723,0.244431518846088,0.20820601118935478,0.2353043407201767,0.16689710319042206,0.020381871610879898,0.14587727189064026,0.43858274817466736,0.17178986966609955,0.16689710319042206 -11,0.001,0.3418840931521522,0.21859833929273817,0.041131472835938133,0.3211692141162025,0.4186662236849467,0.1973548432191213,0.21859833929273817,0.2819516360759735,0.19871805608272552,0.029886296018958092,0.1862599402666092,0.43857088685035706,0.17136642336845398,0.19871805608272552 -12,0.001,0.327799528837204,0.2257505324151781,0.038094287945164576,0.24139907293849522,0.4269571900367737,0.19686741133530936,0.2257505324151781,0.2265533208847046,0.15818463265895844,0.02364269271492958,0.13944414258003235,0.43839970231056213,0.17068645358085632,0.15818463265895844 -13,0.001,0.5482124156422086,0.2889212518930435,0.05876548960804939,0.7335240874025557,0.7617671754625108,0.4734871983528137,0.2889212518930435,0.13387161493301392,0.09191606938838959,0.01390319224447012,0.05341717228293419,0.4380643963813782,0.1701544225215912,0.09191606938838959 -14,0.001,0.2683628300825755,0.1934465699725681,0.03215170403321584,0.1588393354581462,0.3651085015800264,0.16354632439712682,0.1934465699725681,0.11752602458000183,0.08033295720815659,0.0075136758387088776,0.0473531074821949,0.43740665912628174,0.16951507329940796,0.08033295720815659 -15,0.001,0.23245521552032894,0.15845761530929142,0.021862993792941172,0.1661414752403895,0.4185062315728929,0.17951083762778175,0.15845761530929142,0.13540877401828766,0.09327995032072067,0.009294009767472744,0.06379763782024384,0.4362388551235199,0.16890482604503632,0.09327995032072067 -16,0.001,0.2540251049730513,0.16881393061743843,0.02455571148958471,0.1829415543211831,0.4814679291513231,0.2550777710146374,0.16881393061743843,0.24452388286590576,0.17012111842632294,0.027872953563928604,0.15593497455120087,0.43555423617362976,0.16831809282302856,0.17012111842632294 -17,0.001,0.27790621254179215,0.16291005578305987,0.02768930456497603,0.27147142092386883,0.5733835763401456,0.3216428938839171,0.16291005578305987,0.23126113414764404,0.16358613967895508,0.019583800807595253,0.14588342607021332,0.4346035122871399,0.16748279333114624,0.16358613967895508 -18,0.001,0.21714572608470917,0.16336159076955584,0.02207756083872583,0.11655174030197991,0.29744883709483677,0.09367910772562027,0.16336159076955584,0.21335509419441223,0.15434198081493378,0.013381856493651867,0.12396419048309326,0.43387162685394287,0.1662997454404831,0.15434198081493378 -19,0.001,0.39585938718583846,0.23502571880817413,0.03807514740361108,0.43226980169614154,0.5700339277585348,0.33255258699258167,0.23502571880817413,0.29088473320007324,0.1979036182165146,0.050344835966825485,0.18606995046138763,0.4335789680480957,0.16567720472812653,0.1979036182165146 -20,0.001,0.26503710283173454,0.18777255879508126,0.03789552880658044,0.15700330336888632,0.36711519294314915,0.16105001833703783,0.18777255879508126,0.3211463689804077,0.22426843643188477,0.04961967095732689,0.20329146087169647,0.4336654543876648,0.16470572352409363,0.22426843643188477 -21,0.001,0.28275859852631885,0.18599259356657663,0.043205282702628106,0.20048580318689346,0.4617711802323659,0.22377558714813656,0.18599259356657663,0.26557254791259766,0.19436083734035492,0.022131552919745445,0.15585075318813324,0.43292519450187683,0.1639096587896347,0.19436083734035492 -22,0.001,0.21910542911953396,0.14179434875647226,0.016369602125551965,0.16523817347155678,0.5044849514961243,0.2536436948511336,0.14179434875647226,0.2662939727306366,0.19275937974452972,0.025155413895845413,0.15944524109363556,0.4315638244152069,0.16297321021556854,0.19275937974452972 -23,0.0005,0.3775199121899075,0.2587335771984524,0.04730055698504051,0.2791730182038413,0.45998597972922856,0.23086444040139517,0.2587335771984524,0.24067305028438568,0.1702818125486374,0.027465825900435448,0.1423315852880478,0.43073099851608276,0.16307000815868378,0.1702818125486374 -24,0.0005,0.12755481650431952,0.09668536235888799,0.006956045328277267,0.05454847796095742,0.30247317420111763,0.09360233859883414,0.09668536235888799,0.21739451587200165,0.15576092898845673,0.02215108834207058,0.1178358793258667,0.4309665262699127,0.16340160369873047,0.15576092898845673 -25,0.0005,0.2033023022943073,0.13513341546058655,0.014085318272312483,0.12979511668284735,0.5040803915924497,0.2711007396380107,0.13513341546058655,0.22780413925647736,0.16796176135540009,0.01980750262737274,0.11519914120435715,0.43173280358314514,0.1637372523546219,0.16796176135540009 -26,0.0005,0.26767876081996494,0.19527330001195273,0.029956953910489876,0.14093264937400818,0.4134024911456638,0.19583487345112693,0.19527330001195273,0.22256356477737427,0.16064496338367462,0.0197360347956419,0.12349464744329453,0.43188977241516113,0.16440483927726746,0.16064496338367462 -27,0.0005,0.29266027278370327,0.21245999303128985,0.037366345110866755,0.15680206815401712,0.4156636761294471,0.19693358656432894,0.21245999303128985,0.23951780796051025,0.1713121384382248,0.026933521032333374,0.13419118523597717,0.4317554533481598,0.16476942598819733,0.1713121384382248 -28,0.0005,0.2840205712450875,0.1834611776802275,0.03268412335051431,0.22491087267796198,0.502896871831682,0.25805431852738064,0.1834611776802275,0.20573638379573822,0.1483912318944931,0.018318505957722664,0.10800518840551376,0.43159055709838867,0.16473759710788727,0.1483912318944931 -29,0.0005,0.18021956914001042,0.1047939153181182,0.012183074632452594,0.14656837284564972,0.5704893204900954,0.31154681907759774,0.1047939153181182,0.17089061439037323,0.12431001663208008,0.01182954479008913,0.07795087993144989,0.43144071102142334,0.16469746828079224,0.12431001663208008 diff --git a/checkpoints_tree/task1_tr_tree/best_tr_tree.pkl b/checkpoints_tree/task1_tr_tree/best_tr_tree.pkl deleted file mode 100644 index 1ffd2e5..0000000 Binary files a/checkpoints_tree/task1_tr_tree/best_tr_tree.pkl and /dev/null differ diff --git a/checkpoints_tree/task1_tr_tree/model_selection.csv b/checkpoints_tree/task1_tr_tree/model_selection.csv deleted file mode 100644 index 977afdb..0000000 --- a/checkpoints_tree/task1_tr_tree/model_selection.csv +++ /dev/null @@ -1,4 +0,0 @@ -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/forward/evaluation_forward_test_b0_s0.png b/evaluation_outputs/forward/evaluation_forward_test_b0_s0.png deleted file mode 100644 index 7566f71..0000000 Binary files a/evaluation_outputs/forward/evaluation_forward_test_b0_s0.png and /dev/null differ diff --git a/evaluation_outputs/forward/evaluation_forward_val_b0_s0.png b/evaluation_outputs/forward/evaluation_forward_val_b0_s0.png deleted file mode 100644 index c1904ad..0000000 Binary files a/evaluation_outputs/forward/evaluation_forward_val_b0_s0.png and /dev/null differ diff --git a/evaluation_outputs/forward/evaluation_forward_val_b3_s0.png b/evaluation_outputs/forward/evaluation_forward_val_b3_s0.png deleted file mode 100644 index 7c88303..0000000 Binary files a/evaluation_outputs/forward/evaluation_forward_val_b3_s0.png and /dev/null differ diff --git a/evaluation_outputs/forward_rms/evaluation_test_all_samples.csv b/evaluation_outputs/forward_rms/evaluation_test_all_samples.csv deleted file mode 100644 index 45353cb..0000000 --- a/evaluation_outputs/forward_rms/evaluation_test_all_samples.csv +++ /dev/null @@ -1,3 +0,0 @@ -file_name,true_rms,pred_rms,relative_error_percent -Kobe_seismic_wave.csv,0.9873504638671875,0.8297591805458069,15.961029052734375 -Northridge_seismic_wave.csv,0.36303532123565674,0.6096307635307312,67.92601776123047 diff --git a/evaluation_outputs/forward_rms/evaluation_test_s0.png b/evaluation_outputs/forward_rms/evaluation_test_s0.png deleted file mode 100644 index e3e1d75..0000000 Binary files a/evaluation_outputs/forward_rms/evaluation_test_s0.png and /dev/null differ diff --git a/evaluation_outputs/forward_rms/evaluation_train_all_samples.csv b/evaluation_outputs/forward_rms/evaluation_train_all_samples.csv deleted file mode 100644 index b6d287a..0000000 --- a/evaluation_outputs/forward_rms/evaluation_train_all_samples.csv +++ /dev/null @@ -1,37 +0,0 @@ -file_name,frequency_hz,true_rms,pred_rms,relative_error_percent -harmonic_5mm_1.65Hz.csv,1.649999976158142,0.9560214281082153,0.958398163318634,0.24860689734976008 -harmonic_5mm_1.7Hz.csv,1.7000000476837158,0.7670571804046631,0.7446213364601135,2.9249245711660596 -harmonic_5mm_0.7Hz.csv,0.699999988079071,0.3087632358074188,0.2989788353443146,3.1689007395967814 -harmonic_5mm_1.75Hz.csv,1.75,1.1417795419692993,1.1781681776046753,3.187010652915903 -harmonic_5mm_1Hz.csv,1.0,0.4110860228538513,0.4256628453731537,3.545929977892899 -harmonic_5mm_0.6Hz.csv,0.6000000238418579,0.22619637846946716,0.2362937331199646,4.463977150660003 -harmonic_5mm_2.5Hz.csv,2.5,1.527108907699585,1.4517914056777954,4.932032132223416 -harmonic_5mm_1.1Hz.csv,1.100000023841858,0.4332965314388275,0.4580433964729309,5.711300054013273 -harmonic_5mm_0.8Hz.csv,0.800000011920929,0.3433072865009308,0.31963616609573364,6.895024176870468 -harmonic_5mm_0.9Hz.csv,0.8999999761581421,0.39361461997032166,0.3609310984611511,8.303431796216017 -harmonic_5mm_2Hz.csv,2.0,0.8881404399871826,0.8095211982727051,8.852118220808878 -harmonic_5mm_0.5Hz.csv,0.5,0.185550257563591,0.20308159291744232,9.448294809207178 -harmonic_5mm_2.2Hz.csv,2.200000047683716,1.4068351984024048,1.25899076461792,10.50900872770146 -harmonic_5mm_1.45Hz.csv,1.4500000476837158,0.8364413976669312,0.9315659403800964,11.372529262479622 -harmonic_5mm_2.25Hz.csv,2.25,2.169835090637207,1.8927007913589478,12.77213648512221 -harmonic_5mm_1.95Hz.csv,1.9500000476837158,1.9055689573287964,1.6509041786193848,13.364238419710492 -harmonic_5mm_1.3Hz.csv,1.2999999523162842,0.8282244205474854,0.7093077301979065,14.358027534490079 -harmonic_5mm_2.4Hz.csv,2.4000000953674316,4.0633931159973145,3.3863630294799805,16.66169300361096 -harmonic_5mm_1.8Hz.csv,1.7999999523162842,1.193358063697815,0.9930893182754517,16.781949317189667 -harmonic_5mm_1.4Hz.csv,1.399999976158142,0.5552759170532227,0.6490811109542847,16.893438202555952 -harmonic_5mm_1.15Hz.csv,1.149999976158142,0.8088750243186951,0.6577692627906799,18.680977528671942 -harmonic_5mm_0.85Hz.csv,0.8500000238418579,0.590942919254303,0.47555646300315857,19.5258209366055 -harmonic_5mm_1.6Hz.csv,1.600000023841858,1.2506930828094482,0.928046464920044,25.797425629366955 -harmonic_5mm_1.35Hz.csv,1.350000023841858,1.2584363222122192,0.9336408972740173,25.809444562696697 -harmonic_5mm_1.5Hz.csv,1.5,1.0597823858261108,0.7772730588912964,26.657295942373644 -harmonic_5mm_1.9Hz.csv,1.899999976158142,0.816762387752533,1.040550708770752,27.399440078773996 -harmonic_5mm_1.2Hz.csv,1.2000000476837158,0.8279687762260437,0.5912047028541565,28.595773194622048 -harmonic_5mm_2.15Hz.csv,2.1500000953674316,2.85538649559021,1.9993276596069336,29.98048906182589 -harmonic_5mm_1.05Hz.csv,1.0499999523162842,0.39136433601379395,0.5116328001022339,30.73056306392747 -harmonic_5mm_2.3Hz.csv,2.299999952316284,3.4340105056762695,2.2803103923797607,33.59628956840676 -harmonic_5mm_0.55Hz.csv,0.550000011920929,0.23804350197315216,0.33142292499542236,39.22788156292628 -harmonic_5mm_0.65Hz.csv,0.6499999761581421,0.24615468084812164,0.35735413432121277,45.174624788752894 -harmonic_5mm_2.45Hz.csv,2.450000047683716,4.256887435913086,2.1653497219085693,49.13302842728069 -harmonic_5mm_2.1Hz.csv,2.0999999046325684,1.085241436958313,1.8744914531707764,72.72575385847348 -harmonic_5mm_2.05Hz.csv,2.049999952316284,0.7474663257598877,1.5259369611740112,104.14792059330784 -harmonic_5mm_2.35Hz.csv,2.3499999046325684,1.0703518390655518,2.9562220573425293,176.19161750807072 diff --git a/evaluation_outputs/forward_rms/evaluation_val_all_samples.csv b/evaluation_outputs/forward_rms/evaluation_val_all_samples.csv deleted file mode 100644 index f6408ba..0000000 --- a/evaluation_outputs/forward_rms/evaluation_val_all_samples.csv +++ /dev/null @@ -1,6 +0,0 @@ -file_name,frequency_hz,true_rms,pred_rms,relative_error_percent -harmonic_5mm_1.85Hz.csv,1.850000023841858,1.2355948686599731,1.2751460075378418,3.2009795347210073 -harmonic_5mm_1.55Hz.csv,1.5499999523162842,1.4854224920272827,1.5365926027297974,3.4448186275056623 -harmonic_5mm_0.95Hz.csv,0.949999988079071,0.5067455172538757,0.52765291929245,4.125818843326791 -harmonic_5mm_1.25Hz.csv,1.25,0.49293941259384155,0.5253125429153442,6.56736497314257 -harmonic_5mm_0.75Hz.csv,0.75,0.40764957666397095,0.31443294882774353,22.866852603913454 diff --git a/evaluation_outputs/forward_rms/evaluation_val_s0.png b/evaluation_outputs/forward_rms/evaluation_val_s0.png deleted file mode 100644 index 8d84a2a..0000000 Binary files a/evaluation_outputs/forward_rms/evaluation_val_s0.png and /dev/null differ diff --git a/evaluation_outputs/forward_rms/evaluation_val_s0_waveform.png b/evaluation_outputs/forward_rms/evaluation_val_s0_waveform.png deleted file mode 100644 index 6b4df1b..0000000 Binary files a/evaluation_outputs/forward_rms/evaluation_val_s0_waveform.png and /dev/null differ diff --git a/evaluation_outputs/task1_feature_mlp/evaluation_train_all_samples.csv b/evaluation_outputs/task1_feature_mlp/evaluation_train_all_samples.csv deleted file mode 100644 index ffe85d0..0000000 --- a/evaluation_outputs/task1_feature_mlp/evaluation_train_all_samples.csv +++ /dev/null @@ -1,37 +0,0 @@ -file_name,frequency_hz,true_rms,pred_rms,relative_error_percent,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.9Hz.csv,0.8999999761581421,0.3999128043651581,0.42324697971343994,5.8348157632321165,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_0.7Hz.csv,0.699999988079071,0.2675361931324005,0.24590837955474854,8.08407016801223,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_1.2Hz.csv,1.2000000476837158,0.7224284410476685,0.6560892462730408,9.182804967980266,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_0.6Hz.csv,0.6000000238418579,0.16537058353424072,0.18173347413539886,9.894680330356415,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_0.5Hz.csv,0.5,0.18624502420425415,0.16313482820987701,12.408490424438014,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.8Hz.csv,0.800000011920929,0.3204104006290436,0.25158339738845825,21.480889230019116,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.5Hz.csv,1.5,1.2808952331542969,0.9783489108085632,23.619911645755096,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_0.65Hz.csv,0.6499999761581421,0.2034619301557541,0.2532237470149994,24.45755666485207,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_0.55Hz.csv,0.550000011920929,0.1971234530210495,0.24593766033649445,24.763267164477078,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_2.1Hz.csv,2.0999999046325684,0.9339020252227783,1.1686650514602661,25.137864561487174,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_1.1Hz.csv,1.100000023841858,0.5534171462059021,0.410175621509552,25.883102046689434,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_2.2Hz.csv,2.200000047683716,1.0718663930892944,1.3539294004440308,26.315127442496312,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_1.8Hz.csv,1.7999999523162842,1.6704485416412354,1.2214888334274292,26.876596137029036,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_1.45Hz.csv,1.4500000476837158,1.0070611238479614,0.7258758544921875,27.921370679206674,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_2.35Hz.csv,2.3499999046325684,0.7861254215240479,1.0398896932601929,32.280379795399135,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_1.75Hz.csv,1.75,1.59738028049469,0.9762305021286011,38.8855293852586,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_0.85Hz.csv,0.8500000238418579,0.5794365406036377,0.35035184025764465,39.53577040677142,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.3Hz.csv,1.2999999523162842,0.9550632834434509,0.5524730682373047,42.15325017569723,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_2.5Hz.csv,2.5,0.6538841724395752,0.9343159198760986,42.887067657603815,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_1.7Hz.csv,1.7000000476837158,0.6067864894866943,0.8731490969657898,43.89725415679946,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.15Hz.csv,1.149999976158142,0.9153674244880676,0.5124025344848633,44.022201273829275,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_1.95Hz.csv,1.9500000476837158,2.063512086868286,1.132091999053955,45.13761241049535,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_2.25Hz.csv,2.25,2.1223104000091553,1.1386420726776123,46.34893780510615,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.6Hz.csv,1.600000023841858,1.1856828927993774,0.6328451633453369,46.62610321962223,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_1Hz.csv,1.0,0.3037901520729065,0.4519365727901459,48.76603790687915,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.35Hz.csv,1.350000023841858,1.256479024887085,0.6293237805366516,49.913705834189585,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.15Hz.csv,2.1500000953674316,3.0180647373199463,1.3036904335021973,56.80376178213196,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.65Hz.csv,1.649999976158142,0.5992732048034668,0.9399581551551819,56.8496885261946,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_1.4Hz.csv,1.399999976158142,0.3850800693035126,0.6127961277961731,59.134729799058796,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.45Hz.csv,2.450000047683716,4.167365074157715,1.4023194313049316,66.34997399193875,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_2.3Hz.csv,2.299999952316284,3.7199018001556396,1.1657079458236694,68.66293766747023,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_2Hz.csv,2.0,0.5447156429290771,0.9262616038322449,70.04497958815654,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 -harmonic_5mm_1.05Hz.csv,1.0499999523162842,0.36426904797554016,0.6340927481651306,74.07263990425794,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_2.4Hz.csv,2.4000000953674316,4.107017993927002,1.0405583381652832,74.66389629400348,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_1.9Hz.csv,1.899999976158142,0.5868141055107117,1.1536935567855835,96.60290131940668,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.049999952316284,0.6199118494987488,1.310292363166809,111.36752979093583,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 diff --git a/evaluation_outputs/task1_feature_mlp/evaluation_train_curve.png b/evaluation_outputs/task1_feature_mlp/evaluation_train_curve.png deleted file mode 100644 index ed7cedb..0000000 Binary files a/evaluation_outputs/task1_feature_mlp/evaluation_train_curve.png and /dev/null differ diff --git a/evaluation_outputs/task1_feature_mlp/evaluation_val_all_samples.csv b/evaluation_outputs/task1_feature_mlp/evaluation_val_all_samples.csv deleted file mode 100644 index fbb7f40..0000000 --- a/evaluation_outputs/task1_feature_mlp/evaluation_val_all_samples.csv +++ /dev/null @@ -1,6 +0,0 @@ -file_name,frequency_hz,true_rms,pred_rms,relative_error_percent,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.85Hz.csv,1.850000023841858,1.185911774635315,1.2428719997406006,4.803074421181267,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.25Hz.csv,1.25,0.3802441656589508,0.4504010081291199,18.450471777414265,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 -harmonic_5mm_0.95Hz.csv,0.949999988079071,0.5053804516792297,0.3541945815086365,29.915258824959956,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_0.75Hz.csv,0.75,0.39182430505752563,0.2595634460449219,33.755144156559126,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_1.55Hz.csv,1.5499999523162842,1.4177086353302002,0.9243398308753967,34.80043728025205,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 diff --git a/evaluation_outputs/task1_feature_mlp/evaluation_val_curve.png b/evaluation_outputs/task1_feature_mlp/evaluation_val_curve.png deleted file mode 100644 index e884726..0000000 Binary files a/evaluation_outputs/task1_feature_mlp/evaluation_val_curve.png and /dev/null differ diff --git a/evaluation_outputs/task1_feature_mlp/evaluation_val_s0.png b/evaluation_outputs/task1_feature_mlp/evaluation_val_s0.png deleted file mode 100644 index c96e1bc..0000000 Binary files a/evaluation_outputs/task1_feature_mlp/evaluation_val_s0.png and /dev/null differ diff --git a/evaluation_outputs/task1_feature_mlp/harmonic_5mm_0.75Hz_prediction.png b/evaluation_outputs/task1_feature_mlp/harmonic_5mm_0.75Hz_prediction.png deleted file mode 100644 index 2433d5f..0000000 Binary files a/evaluation_outputs/task1_feature_mlp/harmonic_5mm_0.75Hz_prediction.png and /dev/null differ diff --git a/evaluation_outputs/task1_feature_mlp/harmonic_5mm_1.55Hz_prediction.png b/evaluation_outputs/task1_feature_mlp/harmonic_5mm_1.55Hz_prediction.png deleted file mode 100644 index ba8fe90..0000000 Binary files a/evaluation_outputs/task1_feature_mlp/harmonic_5mm_1.55Hz_prediction.png and /dev/null differ diff --git a/evaluation_outputs/task1_tr_tree/evaluation_train_all_samples.csv b/evaluation_outputs/task1_tr_tree/evaluation_train_all_samples.csv deleted file mode 100644 index 2f1ac82..0000000 --- a/evaluation_outputs/task1_tr_tree/evaluation_train_all_samples.csv +++ /dev/null @@ -1,37 +0,0 @@ -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 deleted file mode 100644 index 59ef4ff..0000000 Binary files a/evaluation_outputs/task1_tr_tree/evaluation_train_curve.png and /dev/null 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 deleted file mode 100644 index ca5a6dc..0000000 --- a/evaluation_outputs/task1_tr_tree/evaluation_val_all_samples.csv +++ /dev/null @@ -1,6 +0,0 @@ -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 deleted file mode 100644 index fdcf242..0000000 Binary files a/evaluation_outputs/task1_tr_tree/evaluation_val_curve.png and /dev/null 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 deleted file mode 100644 index f222294..0000000 Binary files a/evaluation_outputs/task1_tr_tree/harmonic_5mm_1.55Hz_prediction.png and /dev/null differ diff --git a/sanity_check_alignment_forward.png b/sanity_check_alignment_forward.png deleted file mode 100644 index fd65537..0000000 Binary files a/sanity_check_alignment_forward.png and /dev/null differ diff --git a/scripts/README.md b/scripts/README.md new file mode 100644 index 0000000..04e5222 --- /dev/null +++ b/scripts/README.md @@ -0,0 +1,54 @@ +# 建筑结构简谐激励响应预测 (Harmonic Response Prediction) + +## 项目简介 +本项目旨在通过数据驱动的方法,预测建筑结构在简谐波底座激振下的顶层加速度响应。工程利用机器学习算法提取核心物理特征,建立激励与响应之间的传递关系,实现对目标结构响应均方根(RMS)的快速、准确预测。 + +## 技术路线 +本项目的核心技术链路包含数据预处理、特征工程、传递率建模与算法回归四个关键环节,逻辑清晰,具备较强的工程可解释性: + +1. **稳态信号截取** + 由于原始振动数据包含起振与衰减的瞬态过程,系统通过滑动窗口结合变异系数评估,自动截取最平稳的中间段数据,消除非稳态噪声对分析的干扰。 +2. **物理特征工程** + 系统从底座激励信号中提取 4 维极具代表性的物理特征作为模型输入: + - 频率 ($f$) + - 频率的平方 ($f^2$) + - 激励均方根 ($x_{rms}$) + - 理论加速度幅值 ($(2\pi f)^2 \cdot A$) +3. **传递率 (TR) 目标建模** + 模型将预测目标设定为系统的**传递率 (Transmissibility, TR)**,即输出响应与输入激励的均方根比值 ($TR = y_{rms} / x_{rms}$)。 + 在最终推理阶段,通过公式 $\text{预测响应 } y_{rms} = \text{预测 } TR \times \text{实际输入 } x_{rms}$ 还原最终结果,这种无量纲化的处理极大提升了模型的泛化能力。 +4. **距离权重 KNN 回归** + 采用数据标准化 (StandardScaler) 结合距离加权的 K 近邻回归 (KNeighborsRegressor) 算法。通过特征空间中的距离衰减机制,对高维空间中的传递率进行平滑拟合。 + +## 核心模块说明 + +- **`config.py`** + 全局配置中心。集中管理数据路径、传感器编号(基座输入与顶层输出)、信号截取比例以及模型超参数。 +- **`dataset.py`** + 数据处理引擎。负责读取原始 CSV 振动数据、缺失值插值对齐、稳态窗口搜索、傅里叶主频计算以及核心物理特征的打包提取。 +- **`train_final.py`** + 模型训练入口。读取全量简谐波数据,构建特征矩阵与目标向量完成 KNN 模型训练,将最优模型序列化保存,并输出全量数据的拟合评估曲线。 +- **`evaluate.py`** + 批量评估工具。加载已保存的模型权重,对全集数据进行预测评估,生成包含理论拟合曲线、散点对比与误差分布的可视化图表。 +- **`predict_single.py`** + 单样本预测脚本。支持输入特定的单一 CSV 数据文件,提供端到端的预测功能,并绘制包含时域对比、频域谱线及预测 RMS 对比的综合诊断图。 + +## 使用指南 + +**1. 训练模型** +读取配置文件中指定的数据集,执行全量训练并保存模型权重: +```bash +python scripts/train_final.py +``` + +**2. 批量评估** +加载已训练的模型,对样本进行批量评估与统计分析: +```bash +python scripts/evaluate.py +``` + +**3. 单文件预测** +对特定的单一样本进行独立预测,并生成详细的图表报告: +```bash +python scripts/predict_single.py --file path/to/your/harmonic.csv +``` \ No newline at end of file diff --git a/scripts/__pycache__/config.cpython-310.pyc b/scripts/__pycache__/config.cpython-310.pyc index 501a546..e97454d 100644 Binary files a/scripts/__pycache__/config.cpython-310.pyc and b/scripts/__pycache__/config.cpython-310.pyc differ diff --git a/scripts/__pycache__/dataset.cpython-310.pyc b/scripts/__pycache__/dataset.cpython-310.pyc index f65b124..70250b6 100644 Binary files a/scripts/__pycache__/dataset.cpython-310.pyc and b/scripts/__pycache__/dataset.cpython-310.pyc differ diff --git a/scripts/__pycache__/model.cpython-310.pyc b/scripts/__pycache__/model.cpython-310.pyc deleted file mode 100644 index 7690e08..0000000 Binary files a/scripts/__pycache__/model.cpython-310.pyc and /dev/null differ diff --git a/scripts/config.py b/scripts/config.py index 804a77f..46f88a0 100644 --- a/scripts/config.py +++ b/scripts/config.py @@ -5,34 +5,18 @@ from pathlib import Path CORE_FEATURE_NAMES: tuple[str, ...] = ( - "dominant_frequency_hz", + "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", + "x_rms", + "theoretical_accel", ) @dataclass class DataConfig: project_root: Path = field(default_factory=lambda: Path(__file__).resolve().parents[1]) - scenario: str = "Non_TMD" - train_split_name: str = "train" - val_split_name: str = "val" - test_split_name: str = "test" - csv_pattern: str = "*.csv" + data_dir: str = "downloads/Non_TMD" + harmonic_pattern: str = "harmonic*.csv" code_column: str = "code" time_column: str = "time" @@ -49,52 +33,31 @@ class DataConfig: stability_subwindow_count: int = 4 interpolation_method: str = "linear" normalization_eps: float = 1e-6 + harmonic_amplitude_m: float = 0.005 - downloads_dir: Path = field(init=False) - scenario_dir: Path = field(init=False) - train_dir: Path = field(init=False) - val_dir: Path = field(init=False) - test_dir: Path = field(init=False) + data_root: Path = field(init=False) def __post_init__(self) -> None: self.project_root = Path(self.project_root).resolve() - self.downloads_dir = self.project_root / "downloads" - self.scenario_dir = self.downloads_dir / self.scenario - self.train_dir = self.scenario_dir / self.train_split_name - self.val_dir = self.scenario_dir / self.val_split_name - self.test_dir = self.scenario_dir / self.test_split_name + self.data_root = (self.project_root / self.data_dir).resolve() @dataclass class ModelConfig: - input_dim: int = len(CORE_FEATURE_NAMES) - hidden_dims: tuple[int, ...] = (96, 64, 32) - dropout: float = 0.08 + model_name: str = "knn_distance" + n_neighbors: int = 4 + distance_power: int = 2 @dataclass class TrainConfig: - epochs: int = 400 - batch_size: int = 16 - learning_rate: float = 1e-3 - weight_decay: float = 1e-4 - seed: int = 42 - device: str = "cuda" - grad_clip_norm: float = 1.0 - lr_scheduler_patience: int = 20 - lr_scheduler_factor: float = 0.5 - min_learning_rate: float = 1e-6 - early_stop_patience: int = 50 - checkpoint_dir: str = "checkpoints_mlp" - history_name: str = "training_history.csv" - best_model_name: str = "best_feature_mlp.pt" - - -@dataclass -class LossConfig: - relative_rms_weight: float = 1.0 - log_rms_huber_weight: float = 0.75 - mae_weight: float = 0.15 + checkpoint_dir: str = "checkpoints_final" + model_name: str = "task1_final_model.pkl" + fit_csv_name: str = "task1_final_fit_all.csv" + fit_figure_name: str = "task1_final_fit_curve.png" + dense_curve_csv_name: str = "task1_final_dense_curve.csv" + evaluation_dense_csv_name: str = "evaluation_dense_curve.csv" + dense_curve_points: int = 400 @dataclass @@ -102,8 +65,19 @@ class ExperimentConfig: data: DataConfig = field(default_factory=DataConfig) model: ModelConfig = field(default_factory=ModelConfig) train: TrainConfig = field(default_factory=TrainConfig) - loss: LossConfig = field(default_factory=LossConfig) def make_experiment_config() -> ExperimentConfig: return ExperimentConfig() + + +def checkpoint_dir(config: ExperimentConfig) -> Path: + path = config.data.project_root / config.train.checkpoint_dir / "task1_final" + path.mkdir(parents=True, exist_ok=True) + return path + + +def evaluation_dir(config: ExperimentConfig) -> Path: + path = config.data.project_root / "evaluation_outputs" / "task1_final" + path.mkdir(parents=True, exist_ok=True) + return path diff --git a/scripts/dataset.py b/scripts/dataset.py index e0a85c3..31b2567 100644 --- a/scripts/dataset.py +++ b/scripts/dataset.py @@ -3,12 +3,9 @@ from __future__ import annotations from dataclasses import dataclass, field from pathlib import Path import re -from typing import Any import numpy as np import pandas as pd -import torch -from torch.utils.data import DataLoader, Dataset try: from .config import CORE_FEATURE_NAMES, DataConfig, ExperimentConfig @@ -85,7 +82,6 @@ def get_middle_segment( peak_cv = float(split_peaks.std() / max(split_peaks.mean(), config.normalization_eps)) window_rms = calculate_rms(window_y) - # Prefer windows that are stable inside the middle candidate region while keeping enough energy. score = rms_cv + 0.35 * peak_cv - 0.05 * (window_rms / max(candidate_rms, config.normalization_eps)) if score < best_score: best_score = score @@ -134,106 +130,24 @@ def get_dominant_frequency(signal: np.ndarray, sampling_rate: float) -> float: return _parabolic_peak_frequency(freqs, magnitudes, peak_index) -def compute_harmonic_fit_features(signal: np.ndarray, sampling_rate: float, frequency_hz: float) -> tuple[float, float]: - signal = np.asarray(signal, dtype=np.float64).reshape(-1) - if signal.size < 4 or frequency_hz <= 0.0: - return 0.0, 1.0 - time_axis = np.arange(signal.size, dtype=np.float64) / max(sampling_rate, 1e-12) - omega_t = 2.0 * np.pi * frequency_hz * time_axis - design = np.stack([np.sin(omega_t), np.cos(omega_t), np.ones_like(omega_t)], axis=1) - coefficients, _, _, _ = np.linalg.lstsq(design, signal, rcond=None) - fitted = design @ coefficients - harmonic_amplitude = float(np.sqrt(coefficients[0] ** 2 + coefficients[1] ** 2)) - residual = signal - fitted - residual_ratio = calculate_rms(residual) / max(calculate_rms(signal), 1e-6) - return harmonic_amplitude, float(residual_ratio) - - -def compute_spectral_features(signal: np.ndarray, sampling_rate: float) -> tuple[float, float, float, float, float]: - freqs, amplitudes = _compute_windowed_spectrum(signal, sampling_rate) - if amplitudes.size == 0: - return 0.0, 0.0, 0.0, 0.0, 0.0 - - powers = np.square(amplitudes) - amplitudes[0] = 0.0 - powers[0] = 0.0 - band_mask = (freqs >= 0.1) & (freqs <= 5.0) - if not np.any(band_mask): - return 0.0, 0.0, 0.0, 0.0, 0.0 - - band_amplitudes = np.where(band_mask, amplitudes, 0.0) - dominant_index = int(np.argmax(band_amplitudes)) - dominant_amplitude = float(amplitudes[dominant_index]) - dominant_freq = _parabolic_peak_frequency(freqs, amplitudes, dominant_index) - - local_mask = np.abs(freqs - dominant_freq) <= 0.10 - dominant_energy = float(powers[local_mask].sum()) - total_band_energy = float(powers[band_mask].sum()) - dominant_energy_ratio = dominant_energy / max(total_band_energy, 1e-12) - spectral_centroid = float((freqs[band_mask] * powers[band_mask]).sum() / max(total_band_energy, 1e-12)) - - background_mask = band_mask & (~local_mask) - background_level = float(np.median(amplitudes[background_mask])) if np.any(background_mask) else 0.0 - spectral_peak_prominence = dominant_amplitude / max(background_level, 1e-6) - - half_power_level = dominant_amplitude / np.sqrt(2.0) - left_index = dominant_index - right_index = dominant_index - while left_index > 0 and amplitudes[left_index] >= half_power_level: - left_index -= 1 - while right_index < amplitudes.size - 1 and amplitudes[right_index] >= half_power_level: - right_index += 1 - half_power_bandwidth = float(freqs[right_index] - freqs[left_index]) if right_index > left_index else 0.0 - return dominant_amplitude, dominant_energy_ratio, spectral_centroid, spectral_peak_prominence, half_power_bandwidth +def theoretical_acceleration(frequency_hz: float, amplitude_m: float) -> float: + return float((2.0 * np.pi * frequency_hz) ** 2 * amplitude_m) def extract_core_features( signal: np.ndarray, sampling_rate: float, - original_length: int, config: DataConfig, + known_frequency_hz: float | None = None, ) -> np.ndarray: + frequency_hz = known_frequency_hz if known_frequency_hz is not None else get_dominant_frequency(signal, sampling_rate) x_rms = calculate_rms(signal) - x_peak = float(np.max(np.abs(signal))) if signal.size > 0 else 0.0 - x_peak_to_peak = float(np.max(signal) - np.min(signal)) if signal.size > 0 else 0.0 - crest_factor = x_peak / max(x_rms, config.normalization_eps) - dominant_frequency = get_dominant_frequency(signal, sampling_rate) - frequency_squared = dominant_frequency * dominant_frequency - inverse_frequency = 1.0 / max(dominant_frequency, 1e-6) - log_frequency = float(np.log(max(dominant_frequency, 1e-6))) - middle_length_ratio = float(signal.size) / float(max(original_length, 1)) - ( - dominant_amplitude, - dominant_energy_ratio, - spectral_centroid, - spectral_peak_prominence, - half_power_bandwidth, - ) = compute_spectral_features(signal, sampling_rate) - harmonic_fit_amplitude, harmonic_fit_residual_ratio = compute_harmonic_fit_features( - signal, - sampling_rate, - dominant_frequency, - ) - signal_mean = float(np.mean(signal)) if signal.size > 0 else 0.0 return np.asarray( [ - dominant_frequency, - frequency_squared, - inverse_frequency, - log_frequency, + frequency_hz, + frequency_hz * frequency_hz, x_rms, - x_peak, - x_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, - spectral_centroid, - signal_mean, + theoretical_acceleration(frequency_hz, config.harmonic_amplitude_m), ], dtype=np.float32, ) @@ -277,29 +191,26 @@ def load_aligned_signals(file_path: Path, config: DataConfig) -> tuple[np.ndarra @dataclass class SampleRecord: file_path: Path - split: str frequency_hz: float - features: torch.Tensor - x_rms: torch.Tensor - y_rms: torch.Tensor - target_y_rms: torch.Tensor - target_log_y_rms: torch.Tensor - time_middle: torch.Tensor - x_middle: torch.Tensor - y_middle: torch.Tensor + features: np.ndarray + x_rms: float + y_rms: float + target_tr: float + time_middle: np.ndarray + x_middle: np.ndarray + y_middle: np.ndarray sampling_rate: float interpolation_count: int = 0 @dataclass -class SplitLoadReport: - split: str +class LoadReport: loaded_files: list[str] = field(default_factory=list) skipped_files: list[tuple[str, str]] = field(default_factory=list) interpolated_files: dict[str, int] = field(default_factory=dict) def to_lines(self) -> list[str]: - lines = [f"[{self.split}] loaded={len(self.loaded_files)} skipped={len(self.skipped_files)}"] + lines = [f"[all_data] loaded={len(self.loaded_files)} skipped={len(self.skipped_files)}"] for file_name, reason in self.skipped_files: lines.append(f" - skipped {file_name}: {reason}") for file_name, count in self.interpolated_files.items(): @@ -308,45 +219,12 @@ class SplitLoadReport: return lines -@dataclass -class NormalizationStats: - feature_mean: torch.Tensor - feature_std: torch.Tensor - target_mean: torch.Tensor - target_std: torch.Tensor +def list_harmonic_files(config: DataConfig) -> list[Path]: + files = sorted(config.data_root.rglob(config.harmonic_pattern)) + return [file_path for file_path in files if extract_frequency_hz(file_path.name) is not None] -class FeatureDataset(Dataset): - def __init__(self, records: list[SampleRecord], normalization: NormalizationStats) -> None: - self.records = records - self.normalization = normalization - - def __len__(self) -> int: - return len(self.records) - - def __getitem__(self, index: int) -> dict[str, Any]: - record = self.records[index] - feature_norm = (record.features - self.normalization.feature_mean) / self.normalization.feature_std - target_norm = (record.target_y_rms - self.normalization.target_mean) / self.normalization.target_std - return { - "x": feature_norm, - "target": target_norm, - "target_raw": record.target_y_rms, - "target_log_raw": record.target_log_y_rms, - "x_rms_raw": record.x_rms, - "y_rms_raw": record.y_rms, - "frequency_hz": torch.tensor(record.frequency_hz, dtype=torch.float32), - "features_raw": record.features, - "file_name": record.file_path.name, - } - - -def list_split_files(config: DataConfig, split: str) -> list[Path]: - split_dir = getattr(config, f"{split}_dir") - return sorted(split_dir.glob(config.csv_pattern)) - - -def build_record_from_file(file_path: Path, split: str, config: DataConfig) -> SampleRecord: +def build_record_from_file(file_path: Path, config: DataConfig) -> SampleRecord: time_values, x_values, y_values, interpolation_count = load_aligned_signals(file_path, config) time_middle, x_middle, y_middle = get_middle_segment(time_values, x_values, y_values, config) sampling_rate = estimate_sampling_rate(time_middle) @@ -356,34 +234,32 @@ def build_record_from_file(file_path: Path, split: str, config: DataConfig) -> S x_rms = calculate_rms(x_middle) y_rms = calculate_rms(y_middle) - transmission_ratio = float(y_rms / max(x_rms, config.normalization_eps)) - target_y_rms = transmission_ratio - target_log_y_rms = float(np.log(max(transmission_ratio, config.normalization_eps))) - features = extract_core_features(x_middle, sampling_rate, len(x_values), config) + target_tr = float(y_rms / max(x_rms, config.normalization_eps)) + features = extract_core_features(x_middle, sampling_rate, config, known_frequency_hz=frequency_hz) return SampleRecord( file_path=file_path, - split=split, - frequency_hz=frequency_hz, - features=torch.tensor(features, dtype=torch.float32), - x_rms=torch.tensor([x_rms], dtype=torch.float32), - y_rms=torch.tensor([y_rms], dtype=torch.float32), - target_y_rms=torch.tensor([target_y_rms], dtype=torch.float32), - target_log_y_rms=torch.tensor([target_log_y_rms], dtype=torch.float32), - time_middle=torch.tensor(time_middle, dtype=torch.float64), - x_middle=torch.tensor(x_middle[:, None], dtype=torch.float32), - y_middle=torch.tensor(y_middle[:, None], dtype=torch.float32), - sampling_rate=sampling_rate, + frequency_hz=float(frequency_hz), + features=features, + x_rms=float(x_rms), + y_rms=float(y_rms), + target_tr=target_tr, + time_middle=np.asarray(time_middle, dtype=np.float64), + x_middle=np.asarray(x_middle, dtype=np.float32).reshape(-1), + y_middle=np.asarray(y_middle, dtype=np.float32).reshape(-1), + sampling_rate=float(sampling_rate), interpolation_count=interpolation_count, ) -def load_split_records(config: DataConfig, split: str) -> tuple[list[SampleRecord], SplitLoadReport]: +def load_all_records(config: ExperimentConfig | DataConfig) -> tuple[list[SampleRecord], LoadReport]: + data_config = config.data if isinstance(config, ExperimentConfig) else config records: list[SampleRecord] = [] - report = SplitLoadReport(split=split) - for file_path in list_split_files(config, split): + report = LoadReport() + + for file_path in list_harmonic_files(data_config): try: - record = build_record_from_file(file_path, split, config) + record = build_record_from_file(file_path, data_config) except ValueError as error: report.skipped_files.append((file_path.name, str(error))) continue @@ -393,86 +269,26 @@ def load_split_records(config: DataConfig, split: str) -> tuple[list[SampleRecor report.interpolated_files[file_path.name] = record.interpolation_count if not records: - raise RuntimeError(f"No usable records found for split='{split}'.") + raise RuntimeError("No usable harmonic records found in data_root.") return records, report -def fit_normalization(records: list[SampleRecord], config: DataConfig) -> NormalizationStats: - feature_all = torch.stack([record.features for record in records], dim=0) - target_all = torch.cat([record.target_y_rms for record in records], dim=0) - - feature_std = torch.clamp(feature_all.std(dim=0, unbiased=False), min=config.normalization_eps) - target_std = torch.clamp(target_all.std(dim=0, unbiased=False).view(1), min=config.normalization_eps) - - return NormalizationStats( - feature_mean=feature_all.mean(dim=0), - feature_std=feature_std, - target_mean=target_all.mean(dim=0, keepdim=True), - target_std=target_std, - ) +def records_to_frame(records: list[SampleRecord]) -> pd.DataFrame: + rows = [] + for record in records: + row = { + "file_name": record.file_path.name, + "frequency_hz": record.frequency_hz, + "x_rms": record.x_rms, + "y_rms": record.y_rms, + "target_tr": record.target_tr, + "sampling_rate": record.sampling_rate, + } + for feature_name, feature_value in zip(CORE_FEATURE_NAMES, record.features.tolist()): + row[feature_name] = float(feature_value) + rows.append(row) + return pd.DataFrame(rows) -def build_datasets( - config: ExperimentConfig | DataConfig, -) -> tuple[dict[str, FeatureDataset], dict[str, list[SampleRecord]], dict[str, SplitLoadReport]]: - data_config = config.data if isinstance(config, ExperimentConfig) else config - train_records, train_report = load_split_records(data_config, "train") - normalization = fit_normalization(train_records, data_config) - val_records, val_report = load_split_records(data_config, "val") - test_records, test_report = load_split_records(data_config, "test") - - raw_records = { - "train": train_records, - "val": val_records, - "test": test_records, - } - datasets = { - split: FeatureDataset(records, normalization) - for split, records in raw_records.items() - } - reports = { - "train": train_report, - "val": val_report, - "test": test_report, - } - return datasets, raw_records, reports - - -def build_dataloaders( - config: ExperimentConfig | DataConfig, -) -> tuple[dict[str, DataLoader], dict[str, FeatureDataset], dict[str, list[SampleRecord]], dict[str, SplitLoadReport]]: - experiment_config = config if isinstance(config, ExperimentConfig) else ExperimentConfig(data=config) - datasets, raw_records, reports = build_datasets(experiment_config) - batch_size = experiment_config.train.batch_size - loaders = { - "train": DataLoader(datasets["train"], batch_size=batch_size, shuffle=True), - "val": DataLoader(datasets["val"], batch_size=batch_size, shuffle=False), - "test": DataLoader(datasets["test"], batch_size=batch_size, shuffle=False), - } - return loaders, datasets, raw_records, reports - - -def report_to_text(reports: dict[str, SplitLoadReport]) -> str: - lines: list[str] = [] - for split in ("train", "val", "test"): - lines.extend(reports[split].to_lines()) - return "\n".join(lines) - - -def checkpoint_normalization_payload(normalization: NormalizationStats) -> dict[str, list[float]]: - return { - "feature_mean": normalization.feature_mean.tolist(), - "feature_std": normalization.feature_std.tolist(), - "target_mean": normalization.target_mean.tolist(), - "target_std": normalization.target_std.tolist(), - "feature_names": list(CORE_FEATURE_NAMES), - } - - -def normalization_from_payload(payload: dict[str, Any]) -> NormalizationStats: - return NormalizationStats( - feature_mean=torch.tensor(payload["feature_mean"], dtype=torch.float32), - feature_std=torch.tensor(payload["feature_std"], dtype=torch.float32), - target_mean=torch.tensor(payload["target_mean"], dtype=torch.float32), - target_std=torch.tensor(payload["target_std"], dtype=torch.float32), - ) +def report_to_text(report: LoadReport) -> str: + return "\n".join(report.to_lines()) diff --git a/scripts/evaluate.py b/scripts/evaluate.py index 2ecb706..9a6c450 100644 --- a/scripts/evaluate.py +++ b/scripts/evaluate.py @@ -1,105 +1,76 @@ from __future__ import annotations import argparse +import pickle from pathlib import Path import matplotlib.pyplot as plt import numpy as np import pandas as pd -import torch try: - from .config import CORE_FEATURE_NAMES, ExperimentConfig, make_experiment_config - from .dataset import build_dataloaders, normalization_from_payload, report_to_text - from .model import build_model + from .config import CORE_FEATURE_NAMES, evaluation_dir, make_experiment_config + from .dataset import load_all_records, records_to_frame, report_to_text except ImportError: - from config import CORE_FEATURE_NAMES, ExperimentConfig, make_experiment_config - from dataset import build_dataloaders, normalization_from_payload, report_to_text - from model import build_model + from config import CORE_FEATURE_NAMES, evaluation_dir, make_experiment_config + from dataset import load_all_records, records_to_frame, report_to_text def parse_args() -> argparse.Namespace: - parser = argparse.ArgumentParser(description="Evaluate feature MLP 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("--device", type=str, default=None) + parser = argparse.ArgumentParser(description="Evaluate the final all-data task1 model.") parser.add_argument("--checkpoint", type=str, default=None) return parser.parse_args() -def resolve_device(device_name: str | None, config: ExperimentConfig) -> torch.device: - requested = device_name or config.train.device - if requested.startswith("cuda") and not torch.cuda.is_available(): - return torch.device("cpu") - return torch.device(requested) - - -def resolve_checkpoint_path(config: ExperimentConfig, checkpoint_arg: str | None) -> Path: +def resolve_checkpoint_path(project_root: Path, checkpoint_arg: str | None) -> Path: if checkpoint_arg: return Path(checkpoint_arg).resolve() - return config.data.project_root / config.train.checkpoint_dir / "task1_feature_mlp" / config.train.best_model_name + return project_root / "checkpoints_final" / "task1_final" / "task1_final_model.pkl" -def relative_percent_error(true_value: float, pred_value: float) -> float: - denominator = max(abs(true_value), 1e-12) - return abs(pred_value - true_value) / denominator * 100.0 +def relative_error_percent(true_value: float, pred_value: float) -> float: + return abs(pred_value - true_value) / max(abs(true_value), 1e-12) * 100.0 -def predict_rms( - model: torch.nn.Module, - feature_norm: torch.Tensor, - x_rms_raw: torch.Tensor, - target_mean: torch.Tensor, - target_std: torch.Tensor, -) -> torch.Tensor: - pred_norm = model(feature_norm) - pred_tr = torch.clamp(pred_norm * target_std + target_mean, min=1e-6) - return pred_tr * x_rms_raw - - -def evaluate_sample( - model: torch.nn.Module, - dataset, - raw_record, - sample_index: int, - device: torch.device, -) -> dict[str, float | str]: - sample = dataset[sample_index] - normalization = dataset.normalization - feature_norm = sample["x"].unsqueeze(0).to(device) - x_rms_raw = sample["x_rms_raw"].unsqueeze(0).to(device) - target_mean = normalization.target_mean.to(device).view(1, 1) - target_std = normalization.target_std.to(device).view(1, 1) - - with torch.no_grad(): - pred_rms = predict_rms(model, feature_norm, x_rms_raw, target_mean, target_std) - - true_rms = float(sample["y_rms_raw"].item()) - pred_rms_value = float(pred_rms.detach().cpu().numpy().reshape(-1)[0]) - row = { - "file_name": str(sample["file_name"]), - "frequency_hz": float(sample["frequency_hz"].item()), - "true_rms": true_rms, - "pred_rms": pred_rms_value, - "relative_error_percent": relative_percent_error(true_rms, pred_rms_value), - } - features_raw = sample["features_raw"].numpy().reshape(-1) - for name, value in zip(CORE_FEATURE_NAMES, features_raw): - row[name] = float(value) - row["sampling_rate"] = float(raw_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 +def build_dense_curve(model, result_df: pd.DataFrame, config) -> pd.DataFrame: 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 Feature-MLP Evaluation | {split}") + dense_freq = np.linspace( + float(plot_df["frequency_hz"].min()), + float(plot_df["frequency_hz"].max()), + config.train.dense_curve_points, + ) + dense_x_rms = np.interp(dense_freq, plot_df["frequency_hz"], plot_df["x_rms"]) + dense_true_rms = np.interp(dense_freq, plot_df["frequency_hz"], plot_df["true_rms"]) + dense_features = np.column_stack( + [ + dense_freq, + dense_freq**2, + dense_x_rms, + (2.0 * np.pi * dense_freq) ** 2 * config.data.harmonic_amplitude_m, + ] + ) + dense_pred_tr = model.predict(dense_features).clip(min=config.data.normalization_eps) + dense_pred_rms = dense_pred_tr * dense_x_rms + return pd.DataFrame( + { + "frequency_hz": dense_freq, + "x_rms_interp": dense_x_rms, + "true_rms_interp": dense_true_rms, + "pred_tr": dense_pred_tr, + "pred_rms": dense_pred_rms, + } + ) - 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") + +def save_plot(result_df: pd.DataFrame, dense_df: pd.DataFrame, figure_path: Path) -> None: + plot_df = result_df.sort_values("frequency_hz").reset_index(drop=True) + fig, axes = plt.subplots(2, 1, figsize=(12, 9)) + fig.suptitle("Task1 Final Model Evaluation | All Harmonic Data") + + axes[0].plot(dense_df["frequency_hz"], dense_df["true_rms_interp"], color="tab:blue", linewidth=2.0, label="True RMS Guide Curve") + axes[0].plot(dense_df["frequency_hz"], dense_df["pred_rms"], color="tab:orange", linewidth=2.0, label="Pred RMS Dense Curve") + axes[0].scatter(plot_df["frequency_hz"], plot_df["true_rms"], color="tab:blue", s=28, zorder=3, label="True RMS Samples") + axes[0].scatter(plot_df["frequency_hz"], plot_df["pred_rms"], color="tab:orange", s=22, zorder=3, label="Pred RMS Samples") axes[0].set_xlabel("Frequency (Hz)") axes[0].set_ylabel("RMS") axes[0].grid(True, alpha=0.3) @@ -112,112 +83,58 @@ def save_all_samples_plot(result_df: pd.DataFrame, split: str, save_dir: Path) - axes[1].tick_params(axis="x", labelrotation=45) plt.tight_layout() - save_dir.mkdir(parents=True, exist_ok=True) - 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(raw_record, result: dict[str, float | str], split: str, save_dir: Path, sample_index: int) -> Path: - time_middle = raw_record.time_middle.detach().cpu().numpy().reshape(-1) - x_middle = raw_record.x_middle.detach().cpu().numpy().reshape(-1) - y_middle = raw_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 Feature-MLP | {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() - save_dir.mkdir(parents=True, exist_ok=True) - 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() - device = resolve_device(args.device, config) - checkpoint_path = resolve_checkpoint_path(config, args.checkpoint) - if not checkpoint_path.exists(): - raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}") + ckpt_path = resolve_checkpoint_path(config.data.project_root, args.checkpoint) + with ckpt_path.open("rb") as handle: + payload = pickle.load(handle) + model = payload["model"] - checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) - if "normalization" in checkpoint: - normalization = normalization_from_payload(checkpoint["normalization"]) - else: - normalization = None + records, report = load_all_records(config) + print(report_to_text(report)) + data_df = records_to_frame(records) + feature_matrix = data_df.loc[:, CORE_FEATURE_NAMES].to_numpy() + pred_tr = model.predict(feature_matrix).clip(min=config.data.normalization_eps) + pred_rms = pred_tr * data_df["x_rms"].to_numpy() + true_rms = data_df["y_rms"].to_numpy() - loaders, datasets, raw_records, reports = build_dataloaders(config) - del loaders - print(report_to_text(reports)) + result_df = data_df.copy() + result_df["pred_tr"] = pred_tr + result_df["pred_rms"] = pred_rms + result_df["true_rms"] = true_rms + result_df["relative_error_percent"] = [ + relative_error_percent(true, pred) for true, pred in zip(true_rms, pred_rms) + ] + result_df = result_df.sort_values(["frequency_hz", "file_name"]).reset_index(drop=True) - dataset = datasets[args.split] - if normalization is not None: - dataset.normalization = normalization + save_dir = evaluation_dir(config) + csv_path = save_dir / "evaluation_all_samples.csv" + dense_csv_path = save_dir / config.train.evaluation_dense_csv_name + figure_path = save_dir / "evaluation_all_curve.png" + dense_df = build_dense_curve(model, result_df, config) + result_df.to_csv(csv_path, index=False) + dense_df.to_csv(dense_csv_path, index=False) + save_plot(result_df, dense_df, figure_path) - model = build_model(config).to(device) - model.load_state_dict(checkpoint["model_state_dict"]) - model.eval() - - save_dir = config.data.project_root / "evaluation_outputs" / "task1_feature_mlp" - if args.all_samples: - rows = [ - evaluate_sample(model, dataset, raw_records[args.split][index], index, device) - for index in range(len(dataset)) - ] - result_df = pd.DataFrame(rows).sort_values(["relative_error_percent", "file_name"]).reset_index(drop=True) - save_dir.mkdir(parents=True, exist_ok=True) - csv_path = save_dir / f"evaluation_{args.split}_all_samples.csv" - result_df.to_csv(csv_path, index=False) - figure_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: {checkpoint_path}") - print(f"Summary: {summary}") - print(f"CSV saved to: {csv_path}") - if figure_path is not None: - print(f"Figure saved to: {figure_path}") - print(result_df.to_string(index=False)) - return - - result = evaluate_sample(model, dataset, raw_records[args.split][args.sample_index], args.sample_index, device) - figure_path = save_single_sample_plot(raw_records[args.split][args.sample_index], result, args.split, save_dir, args.sample_index) - print(f"Checkpoint: {checkpoint_path}") - 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}") + 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: {payload['model_name']}") + print(f"Summary: {summary}") + print(f"CSV saved to: {csv_path}") + print(f"Dense curve CSV saved to: {dense_csv_path}") print(f"Figure saved to: {figure_path}") + print(result_df[["file_name", "frequency_hz", "true_rms", "pred_rms", "relative_error_percent"]].to_string(index=False)) if __name__ == "__main__": diff --git a/scripts/model.py b/scripts/model.py deleted file mode 100644 index 2b02602..0000000 --- a/scripts/model.py +++ /dev/null @@ -1,31 +0,0 @@ -from __future__ import annotations - -import torch -from torch import nn - -try: - from .config import ExperimentConfig, ModelConfig -except ImportError: - from config import ExperimentConfig, ModelConfig - - -class FeatureMLP(nn.Module): - def __init__(self, config: ModelConfig) -> None: - super().__init__() - layers: list[nn.Module] = [] - input_dim = config.input_dim - for hidden_dim in config.hidden_dims: - layers.append(nn.Linear(input_dim, hidden_dim)) - layers.append(nn.GELU()) - layers.append(nn.Dropout(config.dropout)) - input_dim = hidden_dim - layers.append(nn.Linear(input_dim, 1)) - self.network = nn.Sequential(*layers) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.network(x) - - -def build_model(config: ExperimentConfig | ModelConfig) -> FeatureMLP: - model_config = config.model if isinstance(config, ExperimentConfig) else config - return FeatureMLP(model_config) diff --git a/scripts/predict_single.py b/scripts/predict_single.py index 3446080..028bfb2 100644 --- a/scripts/predict_single.py +++ b/scripts/predict_single.py @@ -1,60 +1,44 @@ from __future__ import annotations import argparse +import pickle from pathlib import Path import matplotlib.pyplot as plt import numpy as np -import torch try: - from .config import CORE_FEATURE_NAMES, make_experiment_config - from .dataset import build_record_from_file, normalization_from_payload - from .model import build_model + from .config import CORE_FEATURE_NAMES, evaluation_dir, make_experiment_config + from .dataset import build_record_from_file except ImportError: - from config import CORE_FEATURE_NAMES, make_experiment_config - from dataset import build_record_from_file, normalization_from_payload - from model import build_model + from config import CORE_FEATURE_NAMES, evaluation_dir, make_experiment_config + from dataset import build_record_from_file def parse_args() -> argparse.Namespace: - parser = argparse.ArgumentParser(description="Predict task1 RMS from a single waveform CSV.") + parser = argparse.ArgumentParser(description="Predict task1 RMS from a single harmonic waveform CSV.") parser.add_argument("--file", type=str, required=True) - parser.add_argument("--device", type=str, default=None) parser.add_argument("--checkpoint", type=str, default=None) return parser.parse_args() -def resolve_device(device_name: str | None, default_device: str) -> torch.device: - requested = device_name or default_device - if requested.startswith("cuda") and not torch.cuda.is_available(): - return torch.device("cpu") - return torch.device(requested) - - -def resolve_checkpoint_path(project_root: Path, checkpoint_dir: str, best_model_name: str, checkpoint_arg: str | None) -> Path: +def resolve_checkpoint_path(project_root: Path, checkpoint_arg: str | None) -> Path: if checkpoint_arg: return Path(checkpoint_arg).resolve() - return project_root / checkpoint_dir / "task1_feature_mlp" / best_model_name + return project_root / "checkpoints_final" / "task1_final" / "task1_final_model.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) - y_middle = record.y_middle.detach().cpu().numpy().reshape(-1) +def save_prediction_figure(file_path: Path, record, pred_rms: float, output_dir: Path) -> Path: + time_middle = record.time_middle.reshape(-1) + x_middle = record.x_middle.reshape(-1) + y_middle = record.y_middle.reshape(-1) sampling_rate = record.sampling_rate - fft_values = np.fft.rfft(x_middle) + fft_values = np.fft.rfft(x_middle - np.mean(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 Single-File Prediction | {file_path.name}") + fig, axes = plt.subplots(4, 1, figsize=(12, 12)) + fig.suptitle(f"Task1 Final 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)") @@ -62,26 +46,28 @@ def save_prediction_figure( 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].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[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) + axes[2].plot(freqs, amplitudes, color="tab:purple") + axes[2].axvline(record.frequency_hz, color="tab:red", linestyle="--", label=f"Frequency = {record.frequency_hz:.4f} Hz") + axes[2].set_xlim(0.0, 5.0) + axes[2].set_title("Input Spectrum") + axes[2].set_xlabel("Frequency (Hz)") + axes[2].set_ylabel("Amplitude") + axes[2].grid(True, alpha=0.3) + axes[2].legend() + + axes[3].bar(["True RMS", "Pred RMS"], [record.y_rms, pred_rms], color=["tab:green", "tab:orange"]) + error_percent = abs(pred_rms - record.y_rms) / max(abs(record.y_rms), 1e-12) * 100.0 + axes[3].set_title( + f"True RMS = {record.y_rms:.6f} | Pred RMS = {pred_rms:.6f} | Error = {error_percent:.2f}%" + ) + axes[3].set_ylabel("RMS") + axes[3].grid(True, axis="y", alpha=0.3) plt.tight_layout() output_dir.mkdir(parents=True, exist_ok=True) @@ -94,49 +80,29 @@ def save_prediction_figure( def main() -> None: args = parse_args() config = make_experiment_config() - device = resolve_device(args.device, config.train.device) - checkpoint_path = resolve_checkpoint_path( - config.data.project_root, - config.train.checkpoint_dir, - config.train.best_model_name, - args.checkpoint, - ) - if not checkpoint_path.exists(): - raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}") - - checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) - normalization = normalization_from_payload(checkpoint["normalization"]) - model = build_model(config).to(device) - model.load_state_dict(checkpoint["model_state_dict"]) - model.eval() + ckpt_path = resolve_checkpoint_path(config.data.project_root, args.checkpoint) + with ckpt_path.open("rb") as handle: + payload = pickle.load(handle) + model = payload["model"] file_path = Path(args.file).resolve() - record = build_record_from_file(file_path, split="predict", config=config.data) - feature_norm = ((record.features - normalization.feature_mean) / normalization.feature_std).unsqueeze(0).to(device) - x_rms_raw = record.x_rms.unsqueeze(0).to(device) - target_mean = normalization.target_mean.to(device).view(1, 1) - target_std = normalization.target_std.to(device).view(1, 1) + record = build_record_from_file(file_path, config.data) + pred_tr = max(float(model.predict([record.features.tolist()])[0]), config.data.normalization_eps) + pred_rms = pred_tr * record.x_rms - with torch.no_grad(): - pred_norm = model(feature_norm) - pred_tr = torch.clamp(pred_norm * target_std + target_mean, min=1e-6) - pred_rms = float((pred_tr * x_rms_raw).detach().cpu().numpy().reshape(-1)[0]) + out_dir = evaluation_dir(config) + fig_path = save_prediction_figure(file_path, record, pred_rms, out_dir) - true_rms = float(record.y_rms.item()) if record.y_rms.numel() > 0 else None - output_dir = config.data.project_root / "evaluation_outputs" / "task1_feature_mlp" - figure_path = save_prediction_figure(file_path, record, pred_rms, true_rms, output_dir) - - print(f"Checkpoint: {checkpoint_path}") + print(f"Checkpoint: {ckpt_path}") + print(f"Model: {payload['model_name']}") print(f"Input file: {file_path}") - print(f"Extracted frequency (Hz): {record.frequency_hz:.6f}") for feature_name, feature_value in zip(CORE_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: {figure_path}") + print(f"True RMS: {record.y_rms:.6f}") + print(f"Relative RMS Error (%): {abs(pred_rms - record.y_rms) / max(abs(record.y_rms), 1e-12) * 100.0:.4f}") + print(f"Figure saved to: {fig_path}") if __name__ == "__main__": diff --git a/scripts/train.py b/scripts/train.py deleted file mode 100644 index ecd1960..0000000 --- a/scripts/train.py +++ /dev/null @@ -1,274 +0,0 @@ -from __future__ import annotations - -import argparse -import random -from dataclasses import asdict -from pathlib import Path -from typing import Any - -import pandas as pd -import torch -import torch.nn.functional as F -from torch import nn -from torch.optim import AdamW -from torch.optim.lr_scheduler import ReduceLROnPlateau - -try: - from .config import ExperimentConfig, make_experiment_config - from .dataset import build_dataloaders, checkpoint_normalization_payload, report_to_text - from .model import build_model -except ImportError: - from config import ExperimentConfig, make_experiment_config - from dataset import build_dataloaders, checkpoint_normalization_payload, report_to_text - from model import build_model - - -def parse_args() -> argparse.Namespace: - parser = argparse.ArgumentParser(description="Train feature-based MLP for task1 RMS prediction.") - parser.add_argument("--epochs", type=int, default=None) - parser.add_argument("--batch-size", type=int, default=None) - parser.add_argument("--device", type=str, default=None) - return parser.parse_args() - - -def set_seed(seed: int) -> None: - random.seed(seed) - torch.manual_seed(seed) - torch.cuda.manual_seed_all(seed) - - -def resolve_device(device_name: str) -> torch.device: - if device_name.startswith("cuda") and not torch.cuda.is_available(): - return torch.device("cpu") - return torch.device(device_name) - - -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 denormalize_target(pred_norm: torch.Tensor, target_mean: torch.Tensor, target_std: torch.Tensor) -> torch.Tensor: - return pred_norm * target_std + target_mean - - -def relative_rms_error(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor: - return torch.abs(pred - target) / torch.clamp(target.abs(), min=1e-6) - - -class RMSLoss(nn.Module): - def __init__(self, relative_rms_weight: float, log_rms_huber_weight: float, mae_weight: float) -> None: - super().__init__() - self.relative_rms_weight = relative_rms_weight - self.log_rms_huber_weight = log_rms_huber_weight - self.mae_weight = mae_weight - - def forward( - self, - pred_tr: torch.Tensor, - target_tr: torch.Tensor, - target_log_tr: torch.Tensor, - ) -> dict[str, torch.Tensor]: - relative_loss = relative_rms_error(pred_tr, target_tr).mean() - log_huber_loss = F.huber_loss(torch.log(torch.clamp(pred_tr, min=1e-6)), target_log_tr) - mae_loss = F.l1_loss(pred_tr, target_tr) - total = ( - self.relative_rms_weight * relative_loss - + self.log_rms_huber_weight * log_huber_loss - + self.mae_weight * mae_loss - ) - return { - "total": total, - "relative": relative_loss, - "log_huber": log_huber_loss, - "mae": mae_loss, - } - - -def run_epoch( - model: nn.Module, - dataloader: torch.utils.data.DataLoader, - optimizer: AdamW | None, - criterion: RMSLoss, - target_mean: torch.Tensor, - target_std: torch.Tensor, - device: torch.device, - grad_clip_norm: float, -) -> dict[str, float]: - is_train = optimizer is not None - model.train(is_train) - - total_loss_sum = 0.0 - relative_loss_sum = 0.0 - log_huber_sum = 0.0 - mae_loss_sum = 0.0 - rms_error_sum = 0.0 - sample_count = 0 - - for batch in dataloader: - x = batch["x"].to(device) - target_tr = batch["target_raw"].to(device) - target_log_tr = batch["target_log_raw"].to(device) - x_rms_raw = batch["x_rms_raw"].to(device) - y_rms_raw = batch["y_rms_raw"].to(device) - - if is_train: - optimizer.zero_grad(set_to_none=True) - - pred_norm = model(x) - pred_tr = torch.clamp(denormalize_target(pred_norm, target_mean, target_std), min=1e-6) - pred_rms = pred_tr * x_rms_raw - losses = criterion(pred_tr, target_tr, target_log_tr) - - if is_train: - losses["total"].backward() - torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip_norm) - optimizer.step() - - batch_size = x.shape[0] - total_loss_sum += losses["total"].detach().item() * batch_size - relative_loss_sum += losses["relative"].detach().item() * batch_size - log_huber_sum += losses["log_huber"].detach().item() * batch_size - mae_loss_sum += losses["mae"].detach().item() * batch_size - rms_error_sum += relative_rms_error(pred_rms.detach(), y_rms_raw.detach()).mean().item() * batch_size - sample_count += batch_size - - return { - "loss": total_loss_sum / sample_count, - "relative_loss": relative_loss_sum / sample_count, - "log_huber_loss": log_huber_sum / sample_count, - "mae_loss": mae_loss_sum / sample_count, - "rms_error": rms_error_sum / sample_count, - } - - -def checkpoint_paths(config: ExperimentConfig) -> tuple[Path, Path]: - root = config.data.project_root / config.train.checkpoint_dir / "task1_feature_mlp" - root.mkdir(parents=True, exist_ok=True) - return root / config.train.best_model_name, root / config.train.history_name - - -def train_model(config: ExperimentConfig) -> None: - set_seed(config.train.seed) - device = resolve_device(config.train.device) - loaders, datasets, raw_records, reports = build_dataloaders(config) - del raw_records - normalization = datasets["train"].normalization - target_mean = normalization.target_mean.to(device).view(1, 1) - target_std = normalization.target_std.to(device).view(1, 1) - - model = build_model(config).to(device) - optimizer = AdamW(model.parameters(), lr=config.train.learning_rate, weight_decay=config.train.weight_decay) - scheduler = ReduceLROnPlateau( - optimizer, - mode="min", - factor=config.train.lr_scheduler_factor, - patience=config.train.lr_scheduler_patience, - min_lr=config.train.min_learning_rate, - ) - criterion = RMSLoss( - relative_rms_weight=config.loss.relative_rms_weight, - log_rms_huber_weight=config.loss.log_rms_huber_weight, - mae_weight=config.loss.mae_weight, - ) - - best_val_error = float("inf") - epochs_without_improvement = 0 - history: list[dict[str, float]] = [] - best_model_path, history_path = checkpoint_paths(config) - - print(f"Device: {device}") - print(report_to_text(reports)) - - for epoch in range(1, config.train.epochs + 1): - train_metrics = run_epoch( - model=model, - dataloader=loaders["train"], - optimizer=optimizer, - criterion=criterion, - target_mean=target_mean, - target_std=target_std, - device=device, - grad_clip_norm=config.train.grad_clip_norm, - ) - val_metrics = run_epoch( - model=model, - dataloader=loaders["val"], - optimizer=None, - criterion=criterion, - target_mean=target_mean, - target_std=target_std, - device=device, - grad_clip_norm=config.train.grad_clip_norm, - ) - scheduler.step(val_metrics["rms_error"]) - current_lr = optimizer.param_groups[0]["lr"] - - history_row = { - "epoch": epoch, - "lr": current_lr, - "train_loss": train_metrics["loss"], - "train_relative_loss": train_metrics["relative_loss"], - "train_log_rms_huber_loss": train_metrics["log_huber_loss"], - "train_mae_loss": train_metrics["mae_loss"], - "train_rms_error": train_metrics["rms_error"], - "val_loss": val_metrics["loss"], - "val_relative_loss": val_metrics["relative_loss"], - "val_log_rms_huber_loss": val_metrics["log_huber_loss"], - "val_mae_loss": val_metrics["mae_loss"], - "val_rms_error": val_metrics["rms_error"], - } - history.append(history_row) - - print( - f"Epoch {epoch:03d} | train_loss={train_metrics['loss']:.6f} | " - f"val_loss={val_metrics['loss']:.6f} | val_rms_error={val_metrics['rms_error']:.6f} | lr={current_lr:.2e}" - ) - - if val_metrics["rms_error"] < best_val_error: - best_val_error = val_metrics["rms_error"] - epochs_without_improvement = 0 - torch.save( - { - "epoch": epoch, - "model_state_dict": model.state_dict(), - "optimizer_state_dict": optimizer.state_dict(), - "best_val_rms_error": best_val_error, - "config": serialize_for_checkpoint(asdict(config)), - "normalization": checkpoint_normalization_payload(normalization), - }, - best_model_path, - ) - else: - epochs_without_improvement += 1 - if epochs_without_improvement >= config.train.early_stop_patience: - print(f"Early stopping triggered after {epoch} epochs.") - break - - history_df = pd.DataFrame(history) - history_df.to_csv(history_path, index=False) - print(f"Best model saved to: {best_model_path}") - print(f"Training history saved to: {history_path}") - - -def main() -> None: - args = parse_args() - config = make_experiment_config() - if args.epochs is not None: - config.train.epochs = args.epochs - if args.batch_size is not None: - config.train.batch_size = args.batch_size - if args.device is not None: - config.train.device = args.device - train_model(config) - - -if __name__ == "__main__": - main() diff --git a/scripts/train_final.py b/scripts/train_final.py new file mode 100644 index 0000000..ce5341d --- /dev/null +++ b/scripts/train_final.py @@ -0,0 +1,178 @@ +from __future__ import annotations + +import pickle +from dataclasses import asdict +from pathlib import Path +from typing import Any + +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +from sklearn.neighbors import KNeighborsRegressor +from sklearn.pipeline import Pipeline +from sklearn.preprocessing import StandardScaler + +try: + from .config import CORE_FEATURE_NAMES, ExperimentConfig, checkpoint_dir, evaluation_dir, make_experiment_config + from .dataset import load_all_records, records_to_frame, report_to_text +except ImportError: + from config import CORE_FEATURE_NAMES, ExperimentConfig, checkpoint_dir, evaluation_dir, make_experiment_config + from dataset import load_all_records, records_to_frame, report_to_text + + +def build_final_model(config: ExperimentConfig) -> Pipeline: + return Pipeline( + [ + ("scaler", StandardScaler()), + ( + "model", + KNeighborsRegressor( + n_neighbors=config.model.n_neighbors, + weights="distance", + p=config.model.distance_power, + ), + ), + ] + ) + + +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 relative_error_percent(true_value: float, pred_value: float) -> float: + return abs(pred_value - true_value) / max(abs(true_value), 1e-12) * 100.0 + + +def build_dense_curve(model: Pipeline, result_df: pd.DataFrame, config: ExperimentConfig) -> pd.DataFrame: + plot_df = result_df.sort_values("frequency_hz").reset_index(drop=True) + dense_freq = np.linspace( + float(plot_df["frequency_hz"].min()), + float(plot_df["frequency_hz"].max()), + config.train.dense_curve_points, + ) + dense_x_rms = np.interp(dense_freq, plot_df["frequency_hz"], plot_df["x_rms"]) + dense_true_rms = np.interp(dense_freq, plot_df["frequency_hz"], plot_df["true_rms"]) + dense_features = np.column_stack( + [ + dense_freq, + dense_freq**2, + dense_x_rms, + (2.0 * np.pi * dense_freq) ** 2 * config.data.harmonic_amplitude_m, + ] + ) + dense_pred_tr = model.predict(dense_features).clip(min=config.data.normalization_eps) + dense_pred_rms = dense_pred_tr * dense_x_rms + return pd.DataFrame( + { + "frequency_hz": dense_freq, + "x_rms_interp": dense_x_rms, + "true_rms_interp": dense_true_rms, + "pred_tr": dense_pred_tr, + "pred_rms": dense_pred_rms, + } + ) + + +def save_fit_plot(result_df: pd.DataFrame, dense_df: pd.DataFrame, figure_path: Path) -> None: + plot_df = result_df.sort_values("frequency_hz").reset_index(drop=True) + fig, axes = plt.subplots(2, 1, figsize=(12, 9)) + fig.suptitle("Task1 Final Model | All Harmonic Data Fit") + + axes[0].plot(dense_df["frequency_hz"], dense_df["true_rms_interp"], color="tab:blue", linewidth=2.0, label="True RMS Guide Curve") + axes[0].plot(dense_df["frequency_hz"], dense_df["pred_rms"], color="tab:orange", linewidth=2.0, label="Pred RMS Dense Curve") + axes[0].scatter(plot_df["frequency_hz"], plot_df["true_rms"], color="tab:blue", s=28, zorder=3, label="True RMS Samples") + axes[0].scatter(plot_df["frequency_hz"], plot_df["pred_rms"], color="tab:orange", s=22, zorder=3, label="Pred RMS Samples") + 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() + plt.savefig(figure_path, dpi=180, bbox_inches="tight") + plt.close(fig) + + +def train_final() -> None: + config = make_experiment_config() + records, report = load_all_records(config) + print(report_to_text(report)) + + data_df = records_to_frame(records) + feature_matrix = data_df.loc[:, CORE_FEATURE_NAMES].to_numpy() + target_tr = data_df["target_tr"].to_numpy() + + model = build_final_model(config) + model.fit(feature_matrix, target_tr) + + pred_tr = model.predict(feature_matrix).clip(min=config.data.normalization_eps) + pred_rms = pred_tr * data_df["x_rms"].to_numpy() + true_rms = data_df["y_rms"].to_numpy() + + result_df = data_df.copy() + result_df["pred_tr"] = pred_tr + result_df["pred_rms"] = pred_rms + result_df["true_rms"] = true_rms + result_df["relative_error_percent"] = [ + relative_error_percent(true, pred) for true, pred in zip(true_rms, pred_rms) + ] + result_df = result_df.sort_values(["frequency_hz", "file_name"]).reset_index(drop=True) + + mean_error = float(result_df["relative_error_percent"].mean()) + median_error = float(result_df["relative_error_percent"].median()) + max_error = float(result_df["relative_error_percent"].max()) + + ckpt_dir = checkpoint_dir(config) + eval_dir = evaluation_dir(config) + csv_path = eval_dir / config.train.fit_csv_name + fig_path = eval_dir / config.train.fit_figure_name + dense_csv_path = eval_dir / config.train.dense_curve_csv_name + model_path = ckpt_dir / config.train.model_name + + dense_df = build_dense_curve(model, result_df, config) + result_df.to_csv(csv_path, index=False) + dense_df.to_csv(dense_csv_path, index=False) + save_fit_plot(result_df, dense_df, fig_path) + + payload = { + "model_name": config.model.model_name, + "model": model, + "config": serialize_for_checkpoint(asdict(config)), + "feature_names": list(CORE_FEATURE_NAMES), + "summary": { + "count": len(result_df), + "mean_error_percent": mean_error, + "median_error_percent": median_error, + "max_error_percent": max_error, + }, + } + with model_path.open("wb") as handle: + pickle.dump(payload, handle) + + print(f"Final model saved to: {model_path}") + print(f"Fit CSV saved to: {csv_path}") + print(f"Dense curve CSV saved to: {dense_csv_path}") + print(f"Fit figure saved to: {fig_path}") + print( + f"All-data self-fit summary | count={len(result_df)} | " + f"mean={mean_error:.4f}% | median={median_error:.4f}% | max={max_error:.4f}%" + ) + print(result_df[["file_name", "frequency_hz", "true_rms", "pred_rms", "relative_error_percent"]].to_string(index=False)) + + +if __name__ == "__main__": + train_final() diff --git a/scripts_tree/__pycache__/config.cpython-310.pyc b/scripts_tree/__pycache__/config.cpython-310.pyc deleted file mode 100644 index 2931abb..0000000 Binary files a/scripts_tree/__pycache__/config.cpython-310.pyc and /dev/null differ diff --git a/scripts_tree/config.py b/scripts_tree/config.py deleted file mode 100644 index f37fc96..0000000 --- a/scripts_tree/config.py +++ /dev/null @@ -1,50 +0,0 @@ -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 deleted file mode 100644 index 35da0c8..0000000 --- a/scripts_tree/evaluate.py +++ /dev/null @@ -1,177 +0,0 @@ -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 deleted file mode 100644 index cab8032..0000000 --- a/scripts_tree/predict_single.py +++ /dev/null @@ -1,115 +0,0 @@ -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 deleted file mode 100644 index aaa0bbf..0000000 --- a/scripts_tree/train.py +++ /dev/null @@ -1,160 +0,0 @@ -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() diff --git a/src/__pycache__/config.cpython-310.pyc b/src/__pycache__/config.cpython-310.pyc deleted file mode 100644 index 21e66cb..0000000 Binary files a/src/__pycache__/config.cpython-310.pyc and /dev/null differ diff --git a/src/__pycache__/config.cpython-314.pyc b/src/__pycache__/config.cpython-314.pyc deleted file mode 100644 index fd6846e..0000000 Binary files a/src/__pycache__/config.cpython-314.pyc and /dev/null differ diff --git a/src/__pycache__/dataset.cpython-310.pyc b/src/__pycache__/dataset.cpython-310.pyc deleted file mode 100644 index 292fa11..0000000 Binary files a/src/__pycache__/dataset.cpython-310.pyc and /dev/null differ diff --git a/src/__pycache__/dataset.cpython-314.pyc b/src/__pycache__/dataset.cpython-314.pyc deleted file mode 100644 index fdcdb6c..0000000 Binary files a/src/__pycache__/dataset.cpython-314.pyc and /dev/null differ diff --git a/src/__pycache__/model.cpython-310.pyc b/src/__pycache__/model.cpython-310.pyc deleted file mode 100644 index 0344524..0000000 Binary files a/src/__pycache__/model.cpython-310.pyc and /dev/null differ diff --git a/src/__pycache__/model.cpython-314.pyc b/src/__pycache__/model.cpython-314.pyc deleted file mode 100644 index 40b7cb6..0000000 Binary files a/src/__pycache__/model.cpython-314.pyc and /dev/null differ diff --git a/src/config.py b/src/config.py deleted file mode 100644 index 20aa77f..0000000 --- a/src/config.py +++ /dev/null @@ -1,133 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass, field -from pathlib import Path -from typing import Literal - - -TaskMode = Literal["forward", "inverse"] -IncompleteFilePolicy = Literal["skip", "raise"] - - -@dataclass -class DataConfig: - project_root: Path = field(default_factory=lambda: Path(__file__).resolve().parents[1]) - scenario: str = "Non_TMD" - task: TaskMode = "forward" - - train_split_name: str = "train" - val_split_name: str = "val" - test_split_name: str = "test" - csv_pattern: str = "*.csv" - - code_column: str = "code" - time_column: str = "time" - base_sensor_code: str = "WSMS00012" - base_axis: str = "value1" - response_sensor_code: str = "WSMS00007" - response_axis: str = "value3" - - window_size: int = 512 - window_stride: int = 64 - min_sequence_length: int = 512 - - batch_size: int = 32 - num_workers: int = 0 - pin_memory: bool = True - shuffle_train: bool = True - drop_last_train: bool = False - use_weighted_train_sampler: bool = True - train_weight_gain_power: float = 1.0 - train_weight_response_rms_power: float = 0.5 - train_weight_min: float = 0.5 - train_weight_max: float = 4.0 - - interpolation_method: str = "linear" - incomplete_file_policy: IncompleteFilePolicy = "skip" - normalization_eps: float = 1e-6 - - downloads_dir: Path = field(init=False) - scenario_dir: Path = field(init=False) - train_dir: Path = field(init=False) - val_dir: Path = field(init=False) - test_dir: Path = field(init=False) - - def __post_init__(self) -> None: - self.project_root = Path(self.project_root).resolve() - self.downloads_dir = self.project_root / "downloads" - self.scenario_dir = self.downloads_dir / self.scenario - self.train_dir = self.scenario_dir / self.train_split_name - self.val_dir = self.scenario_dir / self.val_split_name - self.test_dir = self.scenario_dir / self.test_split_name - - -@dataclass -class ModelConfig: - input_channels: int = 1 - output_channels: int = 1 - tcn_channels: tuple[int, ...] = (32, 32, 64, 64, 128) - kernel_size: int = 5 - dropout: float = 0.15 - use_causal_conv: bool = True - dilation_base: int = 2 - use_linear_skip_branch: bool = True - linear_skip_kernel_size: int = 33 - - -@dataclass -class LossConfig: - time_loss: str = "huber" - huber_delta: float = 1.0 - fft_loss_weight: float = 0.005 - fft_loss_min_hz: float = 0.1 - fft_loss_max_hz: float = 5.0 - forward_time_loss_weight: float = 0.1 - forward_rms_loss_weight: float = 8.0 - forward_scale_loss_weight: float = 8.0 - forward_underestimate_loss_weight: float = 16.0 - forward_envelope_loss_weight: float = 2.0 - envelope_kernel_size: int = 33 - - -@dataclass -class TrainConfig: - epochs: int = 100 - learning_rate: float = 1e-3 - weight_decay: float = 1e-4 - seed: int = 42 - device: str = "cuda" - grad_clip_norm: float = 1.0 - use_amp: bool = True - lr_scheduler_patience: int = 5 - lr_scheduler_factor: float = 0.5 - min_learning_rate: float = 1e-6 - early_stop_patience: int = 8 - checkpoint_dir: str = "checkpoints" - best_model_name: str = "best_tcn_model.pt" - history_name: str = "training_history.csv" - dominant_freq_min_hz: float = 0.1 - dominant_freq_max_hz: float = 5.0 - - -@dataclass -class ExperimentConfig: - data: DataConfig = field(default_factory=DataConfig) - model: ModelConfig = field(default_factory=ModelConfig) - loss: LossConfig = field(default_factory=LossConfig) - train: TrainConfig = field(default_factory=TrainConfig) - - -def make_forward_config() -> ExperimentConfig: - config = ExperimentConfig() - config.data.task = "forward" - config.model.input_channels = 1 - config.model.output_channels = 1 - return config - - -def make_inverse_config() -> ExperimentConfig: - config = ExperimentConfig() - config.data.task = "inverse" - config.model.input_channels = 1 - config.model.output_channels = 1 - return config diff --git a/src/dataset.py b/src/dataset.py deleted file mode 100644 index 88857a0..0000000 --- a/src/dataset.py +++ /dev/null @@ -1,449 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass, field -from pathlib import Path -from typing import Any - -import numpy as np -import pandas as pd -import torch -from torch.utils.data import DataLoader, Dataset, WeightedRandomSampler - -try: - from .config import DataConfig, ExperimentConfig -except ImportError: - from config import DataConfig, ExperimentConfig - - -@dataclass -class NormalizationStats: - x_mean: torch.Tensor - x_std: torch.Tensor - y_mean: torch.Tensor - y_std: torch.Tensor - - def normalize_x(self, tensor: torch.Tensor) -> torch.Tensor: - return (tensor - self.x_mean) / self.x_std - - def denormalize_x(self, tensor: torch.Tensor) -> torch.Tensor: - return tensor * self.x_std + self.x_mean - - def normalize_y(self, tensor: torch.Tensor) -> torch.Tensor: - return (tensor - self.y_mean) / self.y_std - - def denormalize_y(self, tensor: torch.Tensor) -> torch.Tensor: - return tensor * self.y_std + self.y_mean - - -@dataclass -class TensorStandardScaler: - mean: torch.Tensor - std: torch.Tensor - - def transform(self, array: torch.Tensor | np.ndarray) -> torch.Tensor | np.ndarray: - return self._apply(array, inverse=False) - - def inverse_transform(self, array: torch.Tensor | np.ndarray) -> torch.Tensor | np.ndarray: - return self._apply(array, inverse=True) - - def _apply(self, array: torch.Tensor | np.ndarray, inverse: bool) -> torch.Tensor | np.ndarray: - if isinstance(array, torch.Tensor): - mean = self.mean.to(array.device, dtype=array.dtype) - std = self.std.to(array.device, dtype=array.dtype) - return array * std + mean if inverse else (array - mean) / std - - np_array = np.asarray(array, dtype=np.float32) - mean_np = self.mean.detach().cpu().numpy().astype(np.float32) - std_np = self.std.detach().cpu().numpy().astype(np.float32) - return np_array * std_np + mean_np if inverse else (np_array - mean_np) / std_np - - -@dataclass -class SequenceRecord: - file_path: Path - split: str - time: torch.Tensor - x_raw: torch.Tensor - y_raw: torch.Tensor - interpolation_count: int = 0 - - -@dataclass -class SplitLoadReport: - split: str - loaded_files: list[str] = field(default_factory=list) - skipped_files: list[tuple[str, str]] = field(default_factory=list) - interpolated_files: dict[str, int] = field(default_factory=dict) - short_files: list[tuple[str, int]] = field(default_factory=list) - - def to_lines(self) -> list[str]: - lines = [ - f"[{self.split}] loaded={len(self.loaded_files)} " - f"skipped={len(self.skipped_files)} short={len(self.short_files)}" - ] - if self.skipped_files: - for file_name, reason in self.skipped_files: - lines.append(f" - skipped {file_name}: {reason}") - if self.short_files: - for file_name, length in self.short_files: - lines.append(f" - short {file_name}: length={length}") - if self.interpolated_files: - for file_name, count in self.interpolated_files.items(): - if count > 0: - lines.append(f" - interpolated {file_name}: missing_points={count}") - return lines - - -class WindowedTimeSeriesDataset(Dataset): - def __init__( - self, - records: list[SequenceRecord], - normalization: NormalizationStats, - config: DataConfig, - split: str, - report: SplitLoadReport | None = None, - ) -> None: - self.config = config - self.split = split - self.normalization = normalization - self.report = report or SplitLoadReport(split=split) - self.sequence_store: list[dict[str, Any]] = [] - self.window_index: list[tuple[int, int]] = [] - self.window_weights: list[float] = [] - - min_length = max(config.min_sequence_length, config.window_size) - raw_window_gains: list[float] = [] - raw_window_response_rms: list[float] = [] - - for record in records: - sequence_length = int(record.time.shape[0]) - if sequence_length < min_length: - self.report.short_files.append((record.file_path.name, sequence_length)) - continue - - x_norm = normalization.normalize_x(record.x_raw) - y_norm = normalization.normalize_y(record.y_raw) - - record_index = len(self.sequence_store) - self.sequence_store.append( - { - "file_name": record.file_path.name, - "file_path": str(record.file_path), - "time": record.time, - "x": x_norm, - "y": y_norm, - } - ) - - last_start = sequence_length - config.window_size - for start in range(0, last_start + 1, config.window_stride): - self.window_index.append((record_index, start)) - raw_x = record.x_raw[start : start + config.window_size] - raw_y = record.y_raw[start : start + config.window_size] - x_rms = torch.sqrt(torch.mean(raw_x.pow(2))).item() - y_rms = torch.sqrt(torch.mean(raw_y.pow(2))).item() - gain = y_rms / max(x_rms, float(config.normalization_eps)) - raw_window_gains.append(float(gain)) - raw_window_response_rms.append(float(y_rms)) - - if not self.window_index: - raise RuntimeError( - f"No valid windows were created for split='{split}'. " - "Check file completeness, sequence length, and window configuration." - ) - - if split == "train": - self.window_weights = build_window_weights( - gains=raw_window_gains, - response_rms=raw_window_response_rms, - config=config, - ) - else: - self.window_weights = [1.0 for _ in self.window_index] - - def __len__(self) -> int: - return len(self.window_index) - - def __getitem__(self, index: int) -> dict[str, Any]: - record_index, start = self.window_index[index] - record = self.sequence_store[record_index] - end = start + self.config.window_size - - x = record["x"][start:end] - time = record["time"][start:end] - y = record["y"][start:end] - - return { - "x": x, - "y": y, - "time": time, - "file_name": record["file_name"], - "file_path": record["file_path"], - "window_start": torch.tensor(start, dtype=torch.long), - "window_end": torch.tensor(end, dtype=torch.long), - } - - def denormalize_x(self, tensor: torch.Tensor) -> torch.Tensor: - return self.normalization.denormalize_x(tensor) - - def denormalize_y(self, tensor: torch.Tensor) -> torch.Tensor: - return self.normalization.denormalize_y(tensor) - - -def list_split_files(config: DataConfig, split: str) -> list[Path]: - split_dir = getattr(config, f"{split}_dir") - return sorted(split_dir.glob(config.csv_pattern)) - - -def _value_frame(df: pd.DataFrame, config: DataConfig, sensor_code: str, value_column: str) -> pd.DataFrame: - sensor_df = df.loc[df[config.code_column] == sensor_code, [config.time_column, value_column]].copy() - sensor_df = sensor_df.sort_values(config.time_column) - sensor_df = sensor_df.drop_duplicates(subset=config.time_column, keep="first") - sensor_df[config.time_column] = sensor_df[config.time_column].astype("float64") - sensor_df[value_column] = sensor_df[value_column].astype("float32") - return sensor_df - - -def _handle_incomplete_file( - file_path: Path, - split: str, - reason: str, - config: DataConfig, - report: SplitLoadReport, -) -> None: - if config.incomplete_file_policy == "raise": - raise ValueError(f"{file_path.name}: {reason}") - report.skipped_files.append((file_path.name, reason)) - - -def align_file_to_base_timeline( - file_path: Path, - config: DataConfig, - split: str, - report: SplitLoadReport, -) -> SequenceRecord | None: - df = pd.read_csv(file_path) - - base_df = _value_frame(df, config, config.base_sensor_code, config.base_axis) - if base_df.empty: - _handle_incomplete_file( - file_path=file_path, - split=split, - reason=f"missing base sensor {config.base_sensor_code}", - config=config, - report=report, - ) - return None - - response_df = _value_frame(df, config, config.response_sensor_code, config.response_axis) - if response_df.empty: - _handle_incomplete_file( - file_path=file_path, - split=split, - reason=f"missing response sensor {config.response_sensor_code}", - config=config, - report=report, - ) - return None - - aligned = base_df.rename(columns={config.base_axis: "base_signal"}) - response_df = response_df.rename(columns={config.response_axis: "response_signal"}) - aligned = aligned.merge(response_df, on=config.time_column, how="left") - - interpolation_count = int(aligned["response_signal"].isna().sum()) - aligned[["response_signal"]] = aligned[["response_signal"]].interpolate( - method=config.interpolation_method, - limit_direction="both", - ) - aligned[["response_signal"]] = aligned[["response_signal"]].ffill().bfill() - - if aligned["response_signal"].isna().any(): - _handle_incomplete_file( - file_path=file_path, - split=split, - reason="remaining NaN values after alignment/interpolation", - config=config, - report=report, - ) - return None - - if interpolation_count > 0: - report.interpolated_files[file_path.name] = interpolation_count - - time_tensor = torch.tensor(aligned[config.time_column].to_numpy(), dtype=torch.float64) - base_tensor = torch.tensor(aligned[["base_signal"]].to_numpy(), dtype=torch.float32) - response_tensor = torch.tensor(aligned[["response_signal"]].to_numpy(), dtype=torch.float32) - - if config.task == "forward": - x_raw = base_tensor - y_raw = response_tensor - elif config.task == "inverse": - x_raw = response_tensor - y_raw = base_tensor - else: - raise ValueError(f"Unsupported task: {config.task}") - - report.loaded_files.append(file_path.name) - return SequenceRecord( - file_path=file_path, - split=split, - time=time_tensor, - x_raw=x_raw, - y_raw=y_raw, - interpolation_count=interpolation_count, - ) - - -def load_split_records(config: DataConfig, split: str) -> tuple[list[SequenceRecord], SplitLoadReport]: - report = SplitLoadReport(split=split) - records: list[SequenceRecord] = [] - - for file_path in list_split_files(config, split): - record = align_file_to_base_timeline(file_path=file_path, config=config, split=split, report=report) - if record is not None: - records.append(record) - - if not records: - raise RuntimeError( - f"No usable files found in split='{split}' under '{getattr(config, f'{split}_dir')}'." - ) - - return records, report - - -def _safe_feature_std(tensor: torch.Tensor, eps: float) -> torch.Tensor: - std = tensor.std(dim=0, unbiased=False) - eps_tensor = torch.full_like(std, eps) - return torch.maximum(std, eps_tensor) - - -def fit_normalization_stats(records: list[SequenceRecord], config: DataConfig) -> NormalizationStats: - x_all = torch.cat([record.x_raw for record in records], dim=0) - y_all = torch.cat([record.y_raw for record in records], dim=0) - x_mean = x_all.mean(dim=0) - x_std = _safe_feature_std(x_all, config.normalization_eps) - y_mean = y_all.mean(dim=0) - y_std = _safe_feature_std(y_all, config.normalization_eps) - - return NormalizationStats( - x_mean=x_mean, - x_std=x_std, - y_mean=y_mean, - y_std=y_std, - ) - - -def build_window_weights( - gains: list[float], - response_rms: list[float], - config: DataConfig, -) -> list[float]: - if not gains or not response_rms: - return [] - - gain_array = np.asarray(gains, dtype=np.float32) - rms_array = np.asarray(response_rms, dtype=np.float32) - gain_ref = float(np.median(gain_array[gain_array > 0.0])) if np.any(gain_array > 0.0) else 1.0 - rms_ref = float(np.median(rms_array[rms_array > 0.0])) if np.any(rms_array > 0.0) else 1.0 - gain_ref = max(gain_ref, float(config.normalization_eps)) - rms_ref = max(rms_ref, float(config.normalization_eps)) - - normalized_gain = np.power(np.maximum(gain_array / gain_ref, config.normalization_eps), config.train_weight_gain_power) - normalized_rms = np.power( - np.maximum(rms_array / rms_ref, config.normalization_eps), - config.train_weight_response_rms_power, - ) - weights = normalized_gain * normalized_rms - weights = np.clip(weights, config.train_weight_min, config.train_weight_max) - return weights.astype(np.float32).tolist() - - -def build_datasets( - config: ExperimentConfig | DataConfig, -) -> tuple[dict[str, WindowedTimeSeriesDataset], NormalizationStats, dict[str, SplitLoadReport]]: - data_config = config.data if isinstance(config, ExperimentConfig) else config - - train_records, train_report = load_split_records(data_config, "train") - normalization = fit_normalization_stats(train_records, data_config) - - val_records, val_report = load_split_records(data_config, "val") - test_records, test_report = load_split_records(data_config, "test") - - datasets = { - "train": WindowedTimeSeriesDataset(train_records, normalization, data_config, "train", train_report), - "val": WindowedTimeSeriesDataset(val_records, normalization, data_config, "val", val_report), - "test": WindowedTimeSeriesDataset(test_records, normalization, data_config, "test", test_report), - } - reports = {"train": train_report, "val": val_report, "test": test_report} - return datasets, normalization, reports - - -def build_dataloaders( - config: ExperimentConfig | DataConfig, -) -> tuple[dict[str, DataLoader], dict[str, WindowedTimeSeriesDataset], dict[str, SplitLoadReport]]: - data_config = config.data if isinstance(config, ExperimentConfig) else config - datasets, _, reports = build_datasets(config) - train_sampler = None - train_shuffle = data_config.shuffle_train - - if data_config.use_weighted_train_sampler: - train_weights = torch.tensor(datasets["train"].window_weights, dtype=torch.double) - train_sampler = WeightedRandomSampler( - weights=train_weights, - num_samples=len(train_weights), - replacement=True, - ) - train_shuffle = False - - loaders = { - "train": DataLoader( - datasets["train"], - batch_size=data_config.batch_size, - shuffle=train_shuffle, - sampler=train_sampler, - num_workers=data_config.num_workers, - pin_memory=data_config.pin_memory, - drop_last=data_config.drop_last_train, - ), - "val": DataLoader( - datasets["val"], - batch_size=data_config.batch_size, - shuffle=False, - num_workers=data_config.num_workers, - pin_memory=data_config.pin_memory, - drop_last=False, - ), - "test": DataLoader( - datasets["test"], - batch_size=data_config.batch_size, - shuffle=False, - num_workers=data_config.num_workers, - pin_memory=data_config.pin_memory, - drop_last=False, - ), - } - return loaders, datasets, reports - - -def get_dataloaders( - config: ExperimentConfig | DataConfig, -) -> tuple[ - dict[str, DataLoader], - dict[str, WindowedTimeSeriesDataset], - dict[str, SplitLoadReport], - TensorStandardScaler, - TensorStandardScaler, -]: - loaders, datasets, reports = build_dataloaders(config) - normalization = datasets["train"].normalization - scaler_x = TensorStandardScaler(mean=normalization.x_mean, std=normalization.x_std) - scaler_y = TensorStandardScaler(mean=normalization.y_mean, std=normalization.y_std) - return loaders, datasets, reports, scaler_x, scaler_y - - -def report_to_text(reports: dict[str, SplitLoadReport]) -> str: - lines: list[str] = [] - for split in ("train", "val", "test"): - report = reports[split] - lines.extend(report.to_lines()) - return "\n".join(lines) diff --git a/src/evaluate.py b/src/evaluate.py deleted file mode 100644 index 4584eff..0000000 --- a/src/evaluate.py +++ /dev/null @@ -1,274 +0,0 @@ -from __future__ import annotations - -import argparse -from pathlib import Path - -import matplotlib.pyplot as plt -import numpy as np -import torch - -try: - from .config import ExperimentConfig, make_forward_config, make_inverse_config - from .dataset import get_dataloaders, report_to_text - from .model import build_model -except ImportError: - from config import ExperimentConfig, make_forward_config, make_inverse_config - from dataset import get_dataloaders, report_to_text - from model import build_model - - -def calculate_rms(signal: np.ndarray) -> float: - signal = np.asarray(signal, dtype=np.float64).reshape(-1) - if signal.size == 0: - raise ValueError("signal must not be empty") - return float(np.sqrt(np.mean(np.square(signal)))) - - -def get_dominant_frequency(signal: np.ndarray, sampling_rate: float) -> float: - signal = np.asarray(signal, dtype=np.float64).reshape(-1) - if signal.size == 0: - raise ValueError("signal must not be empty") - fft_values = np.fft.rfft(signal) - frequencies = np.fft.rfftfreq(signal.size, d=1.0 / sampling_rate) - magnitudes = np.abs(fft_values) - if magnitudes.size > 0: - magnitudes[0] = 0.0 - dominant_index = int(np.argmax(magnitudes)) - return float(frequencies[dominant_index]) - - -def parse_args() -> argparse.Namespace: - parser = argparse.ArgumentParser(description="Evaluate trained TCN with time-domain and frequency-domain plots.") - parser.add_argument("--task", choices=("forward", "inverse"), default="forward") - parser.add_argument("--split", choices=("val", "test"), default="test") - parser.add_argument("--batch-index", type=int, default=0) - parser.add_argument("--sample-index", type=int, default=0) - parser.add_argument("--sampling-rate", type=float, default=None) - parser.add_argument("--device", type=str, default=None) - parser.add_argument("--checkpoint", type=str, default=None) - parser.add_argument("--save-dir", type=str, default=None) - return parser.parse_args() - - -def make_config(task: str) -> ExperimentConfig: - if task == "forward": - return make_forward_config() - if task == "inverse": - return make_inverse_config() - raise ValueError(f"Unsupported task: {task}") - - -def resolve_device(device_name: str | None, config: ExperimentConfig) -> torch.device: - requested = device_name or config.train.device - if requested.startswith("cuda") and not torch.cuda.is_available(): - return torch.device("cpu") - return torch.device(requested) - - -def default_checkpoint_path(config: ExperimentConfig) -> Path: - return config.data.project_root / config.train.checkpoint_dir / config.data.task / config.train.best_model_name - - -def resolve_checkpoint_path(config: ExperimentConfig, checkpoint_arg: str | None) -> Path: - if checkpoint_arg: - return Path(checkpoint_arg).resolve() - - default_path = default_checkpoint_path(config) - if default_path.exists(): - return default_path - - legacy_path = config.data.project_root / "best_model.pth" - if legacy_path.exists(): - return legacy_path - - return default_path - - -def select_batch(dataloader: torch.utils.data.DataLoader, batch_index: int) -> dict[str, object]: - for current_index, batch in enumerate(dataloader): - if current_index == batch_index: - return batch - raise IndexError(f"batch_index={batch_index} is out of range") - - -def get_sample_1d(batch_tensor: torch.Tensor, sample_index: int) -> torch.Tensor: - sample = batch_tensor[sample_index] - if sample.ndim == 2 and sample.shape[-1] == 1: - sample = sample.squeeze(-1) - return sample - - -def compute_fft_curve(signal: np.ndarray, sampling_rate: float) -> tuple[np.ndarray, np.ndarray]: - signal = np.asarray(signal, dtype=np.float64).reshape(-1) - fft_values = np.fft.rfft(signal) - frequencies = np.fft.rfftfreq(signal.size, d=1.0 / sampling_rate) - magnitudes = np.abs(fft_values) - return frequencies, magnitudes - - -def estimate_sampling_rate_from_time(time_values: np.ndarray) -> float: - time_values = np.asarray(time_values, dtype=np.float64).reshape(-1) - if time_values.size < 2: - return 100.0 - dt = np.diff(time_values) - dt = dt[np.isfinite(dt)] - dt = dt[dt > 0.0] - if dt.size == 0: - return 100.0 - median_dt = float(np.median(dt)) - if median_dt <= 0.0: - return 100.0 - return 1.0 / median_dt - - -def relative_percent_error(true_value: float, pred_value: float) -> float: - denominator = max(abs(true_value), 1e-12) - return abs(pred_value - true_value) / denominator * 100.0 - - -def task_labels(config: ExperimentConfig) -> tuple[str, str]: - if config.data.task == "forward": - return ( - f"{config.data.base_sensor_code}/{config.data.base_axis}", - f"{config.data.response_sensor_code}/{config.data.response_axis}", - ) - return ( - f"{config.data.response_sensor_code}/{config.data.response_axis}", - f"{config.data.base_sensor_code}/{config.data.base_axis}", - ) - - -def evaluate_one_batch( - config: ExperimentConfig, - split: str, - batch_index: int, - sample_index: int, - sampling_rate: float, - checkpoint_path: Path, - device: torch.device, - save_dir: Path, -) -> Path: - loaders, datasets, reports, scaler_x, scaler_y = get_dataloaders(config) - print(report_to_text(reports)) - - dataset = datasets[split] - dataloader = loaders[split] - batch = select_batch(dataloader, batch_index) - - x = batch["x"].to(device) - y = batch["y"].to(device) - - model = build_model(config).to(device) - checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) - model.load_state_dict(checkpoint["model_state_dict"]) - model.eval() - - with torch.no_grad(): - pred_norm = model(x) - - scaler_output = scaler_y if config.data.task == "forward" else scaler_x - pred_physical = scaler_output.inverse_transform(pred_norm).detach().cpu() - true_physical = scaler_output.inverse_transform(y).detach().cpu() - - pred_signal = get_sample_1d(pred_physical, sample_index).numpy() - true_signal = get_sample_1d(true_physical, sample_index).numpy() - time_values = batch["time"][sample_index].detach().cpu().numpy() - time_steps = np.arange(true_signal.shape[0], dtype=np.float64) - effective_sampling_rate = ( - sampling_rate if sampling_rate is not None else estimate_sampling_rate_from_time(time_values) - ) - - true_rms = calculate_rms(true_signal) - pred_rms = calculate_rms(pred_signal) - rms_error_percent = relative_percent_error(true_rms, pred_rms) - - true_freq = get_dominant_frequency(true_signal, sampling_rate=effective_sampling_rate) - pred_freq = get_dominant_frequency(pred_signal, sampling_rate=effective_sampling_rate) - - true_fft_freqs, true_fft_mag = compute_fft_curve(true_signal, sampling_rate=effective_sampling_rate) - pred_fft_freqs, pred_fft_mag = compute_fft_curve(pred_signal, sampling_rate=effective_sampling_rate) - - freq_mask_true = (true_fft_freqs >= 0.0) & (true_fft_freqs <= 5.0) - freq_mask_pred = (pred_fft_freqs >= 0.0) & (pred_fft_freqs <= 5.0) - - input_label, output_label = task_labels(config) - file_name = batch["file_name"][sample_index] - window_start = int(batch["window_start"][sample_index].item()) - - fig, axes = plt.subplots(2, 1, figsize=(14, 9)) - fig.suptitle(f"{config.data.task.upper()} evaluation | {split} | {file_name} | window_start={window_start}") - - axes[0].plot(time_steps, true_signal, linewidth=1.2, label="True") - axes[0].plot(time_steps, pred_signal, linewidth=1.2, label="Pred") - axes[0].set_xlabel("Time Step") - axes[0].set_ylabel("Acceleration") - axes[0].set_title( - f"{output_label} | True RMS: {true_rms:.4f}, Pred RMS: {pred_rms:.4f}, Error: {rms_error_percent:.2f}%" - ) - axes[0].grid(True, alpha=0.3) - axes[0].legend(loc="upper right") - - axes[1].plot(true_fft_freqs[freq_mask_true], true_fft_mag[freq_mask_true], linewidth=1.2, label="True FFT") - axes[1].plot(pred_fft_freqs[freq_mask_pred], pred_fft_mag[freq_mask_pred], linewidth=1.2, label="Pred FFT") - true_peak_mag = np.interp(true_freq, true_fft_freqs, true_fft_mag) - pred_peak_mag = np.interp(pred_freq, pred_fft_freqs, pred_fft_mag) - if 0.0 <= true_freq <= 5.0: - axes[1].scatter([true_freq], [true_peak_mag], s=50, marker="o", label="True Peak") - if 0.0 <= pred_freq <= 5.0: - axes[1].scatter([pred_freq], [pred_peak_mag], s=50, marker="x", label="Pred Peak") - axes[1].set_xlim(0.0, 5.0) - axes[1].set_xlabel("Frequency (Hz)") - axes[1].set_ylabel("Amplitude") - axes[1].set_title(f"True Freq: {true_freq:.4f} Hz, Pred Freq: {pred_freq:.4f} Hz") - axes[1].grid(True, alpha=0.3) - axes[1].legend(loc="upper right") - - plt.tight_layout() - save_dir.mkdir(parents=True, exist_ok=True) - figure_path = save_dir / f"evaluation_{config.data.task}_{split}_b{batch_index}_s{sample_index}.png" - plt.savefig(figure_path, dpi=180, bbox_inches="tight") - plt.show() - - print(f"Input channel: {input_label}") - print(f"Output channel: {output_label}") - print(f"Checkpoint: {checkpoint_path}") - print(f"Sampling Rate Used: {effective_sampling_rate:.6f} Hz") - print(f"Figure saved to: {figure_path}") - print(f"True RMS: {true_rms:.6f}") - print(f"Pred RMS: {pred_rms:.6f}") - print(f"RMS Error (%): {rms_error_percent:.4f}") - print(f"True Dominant Frequency: {true_freq:.6f} Hz") - print(f"Pred Dominant Frequency: {pred_freq:.6f} Hz") - print(f"Dominant Frequency Error: {abs(pred_freq - true_freq):.6f} Hz") - - return figure_path - - -def main() -> None: - args = parse_args() - config = make_config(args.task) - device = resolve_device(args.device, config) - checkpoint_path = resolve_checkpoint_path(config, args.checkpoint) - save_dir = ( - Path(args.save_dir).resolve() - if args.save_dir - else config.data.project_root / "evaluation_outputs" / config.data.task - ) - - if not checkpoint_path.exists(): - raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}") - - evaluate_one_batch( - config=config, - split=args.split, - batch_index=args.batch_index, - sample_index=args.sample_index, - sampling_rate=args.sampling_rate, - checkpoint_path=checkpoint_path, - device=device, - save_dir=save_dir, - ) - - -if __name__ == "__main__": - main() diff --git a/src/model.py b/src/model.py deleted file mode 100644 index 715ba61..0000000 --- a/src/model.py +++ /dev/null @@ -1,162 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass - -import torch -import torch.nn as nn -from torch.nn.utils import weight_norm - -try: - from .config import ExperimentConfig, ModelConfig -except ImportError: - from config import ExperimentConfig, ModelConfig - - -class Chomp1d(nn.Module): - def __init__(self, chomp_size: int) -> None: - super().__init__() - self.chomp_size = chomp_size - - def forward(self, x: torch.Tensor) -> torch.Tensor: - if self.chomp_size == 0: - return x - return x[:, :, :-self.chomp_size].contiguous() - - -class CausalConv1d(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int, - dilation: int = 1, - ) -> None: - super().__init__() - padding = (kernel_size - 1) * dilation - self.net = nn.Sequential( - weight_norm( - nn.Conv1d( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=kernel_size, - padding=padding, - dilation=dilation, - ) - ), - Chomp1d(padding), - ) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.net(x) - - -class TemporalBlock(nn.Module): - def __init__( - self, - in_channels: int, - out_channels: int, - kernel_size: int, - dilation: int, - dropout: float, - ) -> None: - super().__init__() - self.conv1 = CausalConv1d(in_channels, out_channels, kernel_size, dilation=dilation) - self.act1 = nn.GELU() - self.dropout1 = nn.Dropout(dropout) - - self.conv2 = CausalConv1d(out_channels, out_channels, kernel_size, dilation=dilation) - self.act2 = nn.GELU() - self.dropout2 = nn.Dropout(dropout) - - self.residual = nn.Conv1d(in_channels, out_channels, kernel_size=1) if in_channels != out_channels else nn.Identity() - self.final_act = nn.GELU() - - def forward(self, x: torch.Tensor) -> torch.Tensor: - residual = self.residual(x) - out = self.conv1(x) - out = self.act1(out) - out = self.dropout1(out) - out = self.conv2(out) - out = self.act2(out) - out = self.dropout2(out) - return self.final_act(out + residual) - - -@dataclass -class ReceptiveFieldInfo: - receptive_field: int - dilations: tuple[int, ...] - - -class TemporalConvNet(nn.Module): - def __init__(self, config: ModelConfig) -> None: - super().__init__() - blocks: list[nn.Module] = [] - in_channels = config.input_channels - dilations: list[int] = [] - - for level, out_channels in enumerate(config.tcn_channels): - dilation = config.dilation_base ** level - dilations.append(dilation) - blocks.append( - TemporalBlock( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=config.kernel_size, - dilation=dilation, - dropout=config.dropout, - ) - ) - in_channels = out_channels - - self.network = nn.Sequential(*blocks) - self.output_head = nn.Conv1d(in_channels, config.output_channels, kernel_size=1) - self.linear_skip = None - if config.use_linear_skip_branch: - self.linear_skip = CausalConv1d( - in_channels=config.input_channels, - out_channels=config.output_channels, - kernel_size=config.linear_skip_kernel_size, - dilation=1, - ) - self.receptive_field_info = compute_receptive_field(config) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - features = self.network(x) - nonlinear_output = self.output_head(features) - if self.linear_skip is None: - return nonlinear_output - return nonlinear_output + self.linear_skip(x) - - -class BuildingTCN(nn.Module): - def __init__(self, config: ModelConfig) -> None: - super().__init__() - self.config = config - self.tcn = TemporalConvNet(config) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - if x.ndim != 3: - raise ValueError(f"Expected input shape (batch, seq, channels), got {tuple(x.shape)}") - x = x.transpose(1, 2) - y = self.tcn(x) - return y.transpose(1, 2) - - @property - def receptive_field(self) -> int: - return self.tcn.receptive_field_info.receptive_field - - -def compute_receptive_field(config: ModelConfig) -> ReceptiveFieldInfo: - receptive_field = 1 - dilations: list[int] = [] - for level, _ in enumerate(config.tcn_channels): - dilation = config.dilation_base ** level - dilations.append(dilation) - receptive_field += 2 * (config.kernel_size - 1) * dilation - return ReceptiveFieldInfo(receptive_field=receptive_field, dilations=tuple(dilations)) - - -def build_model(config: ExperimentConfig | ModelConfig) -> BuildingTCN: - model_config = config.model if isinstance(config, ExperimentConfig) else config - return BuildingTCN(model_config) diff --git a/src/sanity_check.py b/src/sanity_check.py deleted file mode 100644 index 463c1ef..0000000 --- a/src/sanity_check.py +++ /dev/null @@ -1,75 +0,0 @@ -from __future__ import annotations - -from pathlib import Path - -import matplotlib.pyplot as plt -import torch - -from config import make_forward_config -from dataset import build_dataloaders, report_to_text - - -def _to_cpu(tensor: torch.Tensor) -> torch.Tensor: - return tensor.detach().cpu() - - -def main() -> None: - config = make_forward_config() - config.data.batch_size = 4 - config.data.window_size = 1024 - config.data.window_stride = 256 - - loaders, datasets, reports = build_dataloaders(config) - train_loader = loaders["train"] - train_dataset = datasets["train"] - - print("Dataset loading report:") - print(report_to_text(reports)) - - batch = next(iter(train_loader)) - - time = _to_cpu(batch["time"][0]).numpy() - x_norm = batch["x"][0] - y_norm = batch["y"][0] - - x = _to_cpu(train_dataset.denormalize_x(x_norm)).squeeze(-1).numpy() - y = _to_cpu(train_dataset.denormalize_y(y_norm)).squeeze(-1).numpy() - - if config.data.task == "forward": - input_label = f"Input {config.data.base_sensor_code}/{config.data.base_axis}" - output_label = f"Output {config.data.response_sensor_code}/{config.data.response_axis}" - else: - input_label = f"Input {config.data.response_sensor_code}/{config.data.response_axis}" - output_label = f"Output {config.data.base_sensor_code}/{config.data.base_axis}" - - fig, axes = plt.subplots(2, 1, figsize=(14, 8), sharex=True) - fig.suptitle(f"Alignment sanity check: {batch['file_name'][0]}") - - axes[0].plot(time, x, color="black", linewidth=1.2, label=input_label) - axes[0].set_ylabel("Acceleration") - axes[0].legend(loc="upper right") - axes[0].grid(True, alpha=0.3) - - axes[1].plot(time, y, linewidth=1.1, label=output_label) - axes[1].set_xlabel("Time") - axes[1].set_ylabel("Acceleration") - axes[1].legend(loc="upper right") - axes[1].grid(True, alpha=0.3) - - n_guides = min(12, len(time)) - guide_indices = torch.linspace(0, len(time) - 1, steps=n_guides).round().to(torch.long).tolist() - guide_times = [time[i] for i in guide_indices] - for ax in axes: - for guide_time in guide_times: - ax.axvline(guide_time, color="gray", linestyle="--", linewidth=0.8, alpha=0.35) - - plt.tight_layout() - - output_path = Path(__file__).resolve().parents[1] / f"sanity_check_alignment_{config.data.task}.png" - plt.savefig(output_path, dpi=180, bbox_inches="tight") - print(f"Saved figure to: {output_path}") - plt.show() - - -if __name__ == "__main__": - main() diff --git a/src/train.py b/src/train.py deleted file mode 100644 index 6678878..0000000 --- a/src/train.py +++ /dev/null @@ -1,586 +0,0 @@ -from __future__ import annotations - -import argparse -import math -import random -from dataclasses import asdict -from pathlib import Path -from typing import Any - -import pandas as pd -import torch -import torch.nn.functional as F -from torch import nn -from torch.optim import AdamW -from torch.optim.lr_scheduler import ReduceLROnPlateau - -try: - from .config import ExperimentConfig, make_forward_config, make_inverse_config - from .dataset import build_dataloaders, report_to_text - from .model import build_model, compute_receptive_field -except ImportError: - from config import ExperimentConfig, make_forward_config, make_inverse_config - from dataset import build_dataloaders, report_to_text - from model import build_model, compute_receptive_field - - -class MixedTimeFrequencyLoss(nn.Module): - def __init__( - self, - huber_delta: float, - fft_loss_weight: float, - fft_loss_min_hz: float, - fft_loss_max_hz: float, - forward_time_loss_weight: float, - forward_rms_loss_weight: float, - forward_scale_loss_weight: float, - forward_underestimate_loss_weight: float, - forward_envelope_loss_weight: float, - envelope_kernel_size: int, - task: str, - ) -> None: - super().__init__() - self.huber_delta = huber_delta - self.fft_loss_weight = fft_loss_weight - self.fft_loss_min_hz = fft_loss_min_hz - self.fft_loss_max_hz = fft_loss_max_hz - self.forward_time_loss_weight = forward_time_loss_weight - self.forward_rms_loss_weight = forward_rms_loss_weight - self.forward_scale_loss_weight = forward_scale_loss_weight - self.forward_underestimate_loss_weight = forward_underestimate_loss_weight - self.forward_envelope_loss_weight = forward_envelope_loss_weight - self.envelope_kernel_size = envelope_kernel_size - self.task = task - - def forward( - self, - prediction: torch.Tensor, - target: torch.Tensor, - time_tensor: torch.Tensor, - ) -> dict[str, torch.Tensor]: - raw_time_loss = F.huber_loss(prediction, target, delta=self.huber_delta) - fft_loss = band_limited_fft_loss( - prediction=prediction, - target=target, - time_tensor=time_tensor, - min_hz=self.fft_loss_min_hz, - max_hz=self.fft_loss_max_hz, - ) - rms_pred = sequence_rms(prediction) - rms_target = sequence_rms(target) - rms_loss = F.l1_loss(rms_pred, rms_target) if self.task == "forward" else prediction.new_tensor(0.0) - scale_loss = forward_scale_loss(prediction, target) if self.task == "forward" else prediction.new_tensor(0.0) - underestimate_loss = ( - forward_underestimate_loss(prediction, target) if self.task == "forward" else prediction.new_tensor(0.0) - ) - envelope_loss = ( - forward_envelope_loss(prediction, target, self.envelope_kernel_size) - if self.task == "forward" - else prediction.new_tensor(0.0) - ) - time_loss = ( - raw_time_loss * self.forward_time_loss_weight - if self.task == "forward" - else raw_time_loss - ) - total_loss = ( - time_loss - + self.fft_loss_weight * fft_loss - + self.forward_rms_loss_weight * rms_loss - + self.forward_scale_loss_weight * scale_loss - + self.forward_underestimate_loss_weight * underestimate_loss - + self.forward_envelope_loss_weight * envelope_loss - ) - return { - "total": total_loss, - "time": raw_time_loss, - "weighted_time": time_loss, - "fft": fft_loss, - "rms_loss": rms_loss, - "scale_loss": scale_loss, - "underestimate_loss": underestimate_loss, - "envelope_loss": envelope_loss, - } - - -def parse_args() -> argparse.Namespace: - parser = argparse.ArgumentParser(description="Train single-input single-output TCN for structural response modeling.") - parser.add_argument("--task", choices=("forward", "inverse"), default="forward") - parser.add_argument("--epochs", type=int, default=None) - parser.add_argument("--batch-size", type=int, default=None) - parser.add_argument("--device", type=str, default=None) - return parser.parse_args() - - -def make_config(task: str) -> ExperimentConfig: - if task == "forward": - return make_forward_config() - if task == "inverse": - return make_inverse_config() - raise ValueError(f"Unsupported task: {task}") - - -def set_seed(seed: int) -> None: - random.seed(seed) - torch.manual_seed(seed) - torch.cuda.manual_seed_all(seed) - - -def resolve_device(device_name: str) -> torch.device: - if device_name.startswith("cuda") and not torch.cuda.is_available(): - return torch.device("cpu") - return torch.device(device_name) - - -def move_batch_to_device(batch: dict[str, object], device: torch.device) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - x = batch["x"].to(device) - y = batch["y"].to(device) - time = batch["time"].to(device) - return x, y, time - - -def denormalize_prediction( - prediction_norm: torch.Tensor, - target_norm: torch.Tensor, - y_mean: torch.Tensor, - y_std: torch.Tensor, -) -> tuple[torch.Tensor, torch.Tensor]: - prediction = prediction_norm * y_std + y_mean - target = target_norm * y_std + y_mean - return prediction, target - - -def sequence_rms(signal: torch.Tensor) -> torch.Tensor: - return torch.sqrt(torch.mean(signal ** 2, dim=1)) - - -def rms_error(prediction: torch.Tensor, target: torch.Tensor) -> torch.Tensor: - pred_rms = sequence_rms(prediction) - target_rms = sequence_rms(target) - return torch.abs(pred_rms - target_rms) - - -def forward_scale_loss(prediction: torch.Tensor, target: torch.Tensor) -> torch.Tensor: - epsilon = 1e-6 - pred_std = prediction.std(dim=1, unbiased=False) - target_std = target.std(dim=1, unbiased=False) - pred_peak = prediction.abs().amax(dim=1) - target_peak = target.abs().amax(dim=1) - pred_rms = sequence_rms(prediction) - target_rms = sequence_rms(target) - std_loss = torch.mean(torch.abs(pred_std - target_std) / (target_std.abs() + epsilon)) - peak_loss = torch.mean(torch.abs(pred_peak - target_peak) / (target_peak.abs() + epsilon)) - rms_loss = torch.mean(torch.abs(pred_rms - target_rms) / (target_rms.abs() + epsilon)) - return (std_loss + peak_loss + rms_loss) / 3.0 - - -def forward_underestimate_loss(prediction: torch.Tensor, target: torch.Tensor) -> torch.Tensor: - epsilon = 1e-6 - pred_rms = sequence_rms(prediction) - target_rms = sequence_rms(target) - pred_peak = prediction.abs().amax(dim=1) - target_peak = target.abs().amax(dim=1) - pred_std = prediction.std(dim=1, unbiased=False) - target_std = target.std(dim=1, unbiased=False) - - rms_under = torch.relu((target_rms - pred_rms) / (target_rms.abs() + epsilon)) - peak_under = torch.relu((target_peak - pred_peak) / (target_peak.abs() + epsilon)) - std_under = torch.relu((target_std - pred_std) / (target_std.abs() + epsilon)) - return (rms_under.mean() + peak_under.mean() + std_under.mean()) / 3.0 - - -def local_envelope(signal: torch.Tensor, kernel_size: int) -> torch.Tensor: - signal_abs = signal.squeeze(-1).abs().unsqueeze(1) - effective_kernel = max(3, int(kernel_size)) - if effective_kernel % 2 == 0: - effective_kernel += 1 - padding = effective_kernel // 2 - return F.max_pool1d(signal_abs, kernel_size=effective_kernel, stride=1, padding=padding).squeeze(1) - - -def forward_envelope_loss(prediction: torch.Tensor, target: torch.Tensor, kernel_size: int) -> torch.Tensor: - epsilon = 1e-6 - pred_envelope = local_envelope(prediction, kernel_size) - target_envelope = local_envelope(target, kernel_size) - relative_error = torch.abs(pred_envelope - target_envelope) / (target_envelope.abs() + epsilon) - underestimate_error = torch.relu((target_envelope - pred_envelope) / (target_envelope.abs() + epsilon)) - return relative_error.mean() + 2.0 * underestimate_error.mean() - - -def band_limited_fft_loss( - prediction: torch.Tensor, - target: torch.Tensor, - time_tensor: torch.Tensor, - min_hz: float, - max_hz: float, -) -> torch.Tensor: - prediction_signal = prediction.squeeze(-1) - target_signal = target.squeeze(-1) - if prediction_signal.shape[1] < 4: - pred_fft = torch.abs(torch.fft.rfft(prediction_signal, dim=1)) - target_fft = torch.abs(torch.fft.rfft(target_signal, dim=1)) - return F.mse_loss(pred_fft, target_fft) - - dt = torch.diff(time_tensor, dim=1).mean(dim=1) - dt = torch.clamp(dt, min=1e-6) - pred_fft = torch.abs(torch.fft.rfft(prediction_signal, dim=1)) - target_fft = torch.abs(torch.fft.rfft(target_signal, dim=1)) - fft_bins = pred_fft.shape[1] - n = prediction_signal.shape[1] - bin_indices = torch.arange(fft_bins, device=prediction.device, dtype=prediction.dtype).unsqueeze(0) - frequencies = bin_indices / (dt.unsqueeze(1) * n) - band_mask = (frequencies >= min_hz) & (frequencies <= max_hz) - band_mask[:, 0] = False - - squared_error = (pred_fft - target_fft) ** 2 - masked_error = squared_error * band_mask.to(squared_error.dtype) - valid_counts = band_mask.sum(dim=1).clamp_min(1).to(squared_error.dtype) - per_sample_loss = masked_error.sum(dim=1) / valid_counts - fallback_loss = squared_error.mean(dim=1) - has_valid_band = band_mask.any(dim=1) - per_sample_loss = torch.where(has_valid_band, per_sample_loss, fallback_loss) - return per_sample_loss.mean() - - -def dominant_frequency_error( - prediction: torch.Tensor, - target: torch.Tensor, - time_tensor: torch.Tensor, - min_hz: float, - max_hz: float, -) -> torch.Tensor: - if prediction.shape[1] < 4: - return torch.tensor(0.0, device=prediction.device) - - dt = torch.diff(time_tensor, dim=1).mean(dim=1) - dt = torch.clamp(dt, min=1e-6) - - prediction_signal = prediction.squeeze(-1) - target_signal = target.squeeze(-1) - pred_fft = torch.abs(torch.fft.rfft(prediction_signal, dim=1)) - target_fft = torch.abs(torch.fft.rfft(target_signal, dim=1)) - n = prediction_signal.shape[1] - fft_bins = pred_fft.shape[1] - bin_indices = torch.arange(fft_bins, device=prediction.device, dtype=prediction.dtype).unsqueeze(0) - frequencies = bin_indices / (dt.unsqueeze(1) * n) - band_mask = (frequencies >= min_hz) & (frequencies <= max_hz) - - if not band_mask.any(): - return torch.tensor(0.0, device=prediction.device) - - masked_pred_fft = pred_fft.masked_fill(~band_mask, float("-inf")) - masked_target_fft = target_fft.masked_fill(~band_mask, float("-inf")) - - pred_index = masked_pred_fft.argmax(dim=1) - target_index = masked_target_fft.argmax(dim=1) - pred_freq = pred_index.to(prediction.dtype) / (dt * n) - target_freq = target_index.to(target.dtype) / (dt * n) - return torch.mean(torch.abs(pred_freq - target_freq)) - - -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 monitor_metric_name(config: ExperimentConfig) -> str: - if config.data.task == "forward": - return "rms_error" - return "loss" - - -def run_epoch( - model: nn.Module, - dataloader: torch.utils.data.DataLoader, - optimizer: AdamW | None, - criterion: MixedTimeFrequencyLoss, - y_mean: torch.Tensor, - y_std: torch.Tensor, - device: torch.device, - grad_clip_norm: float, - scaler: torch.cuda.amp.GradScaler, - amp_enabled: bool, - dominant_freq_min_hz: float, - dominant_freq_max_hz: float, -) -> dict[str, float]: - is_train = optimizer is not None - model.train(is_train) - - total_loss_sum = 0.0 - time_loss_sum = 0.0 - weighted_time_loss_sum = 0.0 - fft_loss_sum = 0.0 - rms_loss_sum = 0.0 - scale_loss_sum = 0.0 - underestimate_loss_sum = 0.0 - envelope_loss_sum = 0.0 - rms_error_sum = 0.0 - dominant_freq_error_sum = 0.0 - sample_count = 0 - - for batch in dataloader: - x, y_norm, time = move_batch_to_device(batch, device) - - if is_train: - optimizer.zero_grad(set_to_none=True) - - with torch.amp.autocast(device_type=device.type, enabled=amp_enabled): - prediction_norm = model(x) - prediction, target = denormalize_prediction(prediction_norm, y_norm, y_mean, y_std) - losses = criterion(prediction, target, time) - - if is_train: - scaler.scale(losses["total"]).backward() - scaler.unscale_(optimizer) - torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip_norm) - scaler.step(optimizer) - scaler.update() - - batch_size = x.shape[0] - total_loss_sum += losses["total"].detach().item() * batch_size - time_loss_sum += losses["time"].detach().item() * batch_size - weighted_time_loss_sum += losses["weighted_time"].detach().item() * batch_size - fft_loss_sum += losses["fft"].detach().item() * batch_size - rms_loss_sum += losses["rms_loss"].detach().item() * batch_size - scale_loss_sum += losses["scale_loss"].detach().item() * batch_size - underestimate_loss_sum += losses["underestimate_loss"].detach().item() * batch_size - envelope_loss_sum += losses["envelope_loss"].detach().item() * batch_size - rms_error_sum += rms_error(prediction.detach(), target.detach()).mean().item() * batch_size - dominant_freq_error_sum += dominant_frequency_error( - prediction.detach(), - target.detach(), - time.detach(), - min_hz=dominant_freq_min_hz, - max_hz=dominant_freq_max_hz, - ).item() * batch_size - sample_count += batch_size - - if sample_count == 0: - raise RuntimeError("No samples were processed in the epoch.") - - return { - "loss": total_loss_sum / sample_count, - "time_loss": time_loss_sum / sample_count, - "weighted_time_loss": weighted_time_loss_sum / sample_count, - "fft_loss": fft_loss_sum / sample_count, - "rms_loss": rms_loss_sum / sample_count, - "scale_loss": scale_loss_sum / sample_count, - "underestimate_loss": underestimate_loss_sum / sample_count, - "envelope_loss": envelope_loss_sum / sample_count, - "rms_error": rms_error_sum / sample_count, - "dominant_freq_error": dominant_freq_error_sum / sample_count, - } - - -def checkpoint_paths(config: ExperimentConfig) -> tuple[Path, Path]: - root = config.data.project_root / config.train.checkpoint_dir / config.data.task - root.mkdir(parents=True, exist_ok=True) - return root / config.train.best_model_name, root / config.train.history_name - - -def train_model(config: ExperimentConfig) -> None: - set_seed(config.train.seed) - device = resolve_device(config.train.device) - - dataloaders, datasets, reports = build_dataloaders(config) - train_dataset = datasets["train"] - normalization = train_dataset.normalization - y_mean = normalization.y_mean.to(device).view(1, 1, -1) - y_std = normalization.y_std.to(device).view(1, 1, -1) - - model = build_model(config).to(device) - receptive_field = compute_receptive_field(config.model) - optimizer = AdamW( - model.parameters(), - lr=config.train.learning_rate, - weight_decay=config.train.weight_decay, - ) - scheduler = ReduceLROnPlateau( - optimizer, - mode="min", - factor=config.train.lr_scheduler_factor, - patience=config.train.lr_scheduler_patience, - min_lr=config.train.min_learning_rate, - ) - criterion = MixedTimeFrequencyLoss( - huber_delta=config.loss.huber_delta, - fft_loss_weight=config.loss.fft_loss_weight, - fft_loss_min_hz=config.loss.fft_loss_min_hz, - fft_loss_max_hz=config.loss.fft_loss_max_hz, - forward_time_loss_weight=config.loss.forward_time_loss_weight, - forward_rms_loss_weight=config.loss.forward_rms_loss_weight, - forward_scale_loss_weight=config.loss.forward_scale_loss_weight, - forward_underestimate_loss_weight=config.loss.forward_underestimate_loss_weight, - forward_envelope_loss_weight=config.loss.forward_envelope_loss_weight, - envelope_kernel_size=config.loss.envelope_kernel_size, - task=config.data.task, - ) - - amp_enabled = config.train.use_amp and device.type == "cuda" - scaler = torch.cuda.amp.GradScaler(enabled=amp_enabled) - best_model_path, history_path = checkpoint_paths(config) - - print(f"Task: {config.data.task}") - print(f"Device: {device}") - print( - "TCN receptive field: " - f"{receptive_field.receptive_field} samples, " - f"dilations={receptive_field.dilations}" - ) - print(report_to_text(reports)) - - history_rows: list[dict[str, float | int]] = [] - monitor_name = monitor_metric_name(config) - best_monitor_value = math.inf - stale_epochs = 0 - - for epoch in range(1, config.train.epochs + 1): - train_metrics = run_epoch( - model=model, - dataloader=dataloaders["train"], - optimizer=optimizer, - criterion=criterion, - y_mean=y_mean, - y_std=y_std, - device=device, - grad_clip_norm=config.train.grad_clip_norm, - scaler=scaler, - amp_enabled=amp_enabled, - dominant_freq_min_hz=config.train.dominant_freq_min_hz, - dominant_freq_max_hz=config.train.dominant_freq_max_hz, - ) - - with torch.no_grad(): - val_metrics = run_epoch( - model=model, - dataloader=dataloaders["val"], - optimizer=None, - criterion=criterion, - y_mean=y_mean, - y_std=y_std, - device=device, - grad_clip_norm=config.train.grad_clip_norm, - scaler=scaler, - amp_enabled=amp_enabled, - dominant_freq_min_hz=config.train.dominant_freq_min_hz, - dominant_freq_max_hz=config.train.dominant_freq_max_hz, - ) - - current_monitor_value = val_metrics[monitor_name] - scheduler.step(current_monitor_value) - current_lr = optimizer.param_groups[0]["lr"] - - history_row: dict[str, float | int] = { - "epoch": epoch, - "lr": current_lr, - "train_loss": train_metrics["loss"], - "train_time_loss": train_metrics["time_loss"], - "train_weighted_time_loss": train_metrics["weighted_time_loss"], - "train_fft_loss": train_metrics["fft_loss"], - "train_rms_loss": train_metrics["rms_loss"], - "train_scale_loss": train_metrics["scale_loss"], - "train_underestimate_loss": train_metrics["underestimate_loss"], - "train_envelope_loss": train_metrics["envelope_loss"], - "train_rms_error": train_metrics["rms_error"], - "train_dominant_freq_error": train_metrics["dominant_freq_error"], - "val_loss": val_metrics["loss"], - "val_time_loss": val_metrics["time_loss"], - "val_weighted_time_loss": val_metrics["weighted_time_loss"], - "val_fft_loss": val_metrics["fft_loss"], - "val_rms_loss": val_metrics["rms_loss"], - "val_scale_loss": val_metrics["scale_loss"], - "val_underestimate_loss": val_metrics["underestimate_loss"], - "val_envelope_loss": val_metrics["envelope_loss"], - "val_rms_error": val_metrics["rms_error"], - "val_dominant_freq_error": val_metrics["dominant_freq_error"], - } - history_rows.append(history_row) - - print( - f"Epoch {epoch:03d} | " - f"train_loss={train_metrics['loss']:.6f} | " - f"val_loss={val_metrics['loss']:.6f} | " - f"val_rms={val_metrics['rms_error']:.6f} | " - f"val_dom_freq_err={val_metrics['dominant_freq_error']:.6f} Hz | " - f"monitor=val_{monitor_name}:{current_monitor_value:.6f} | " - f"lr={current_lr:.2e}" - ) - - if current_monitor_value < best_monitor_value: - best_monitor_value = current_monitor_value - stale_epochs = 0 - torch.save( - { - "model_state_dict": model.state_dict(), - "optimizer_state_dict": optimizer.state_dict(), - "config": serialize_for_checkpoint(asdict(config)), - "epoch": epoch, - "best_monitor_name": monitor_name, - "best_monitor_value": best_monitor_value, - "val_metrics": val_metrics, - }, - best_model_path, - ) - else: - stale_epochs += 1 - - if stale_epochs >= config.train.early_stop_patience: - print(f"Early stopping triggered after {epoch} epochs.") - break - - history_df = pd.DataFrame(history_rows) - history_df.to_csv(history_path, index=False) - print(f"Best model saved to: {best_model_path}") - print(f"Training history saved to: {history_path}") - - if best_model_path.exists(): - checkpoint = torch.load(best_model_path, map_location=device, weights_only=False) - model.load_state_dict(checkpoint["model_state_dict"]) - - with torch.no_grad(): - test_metrics = run_epoch( - model=model, - dataloader=dataloaders["test"], - optimizer=None, - criterion=criterion, - y_mean=y_mean, - y_std=y_std, - device=device, - grad_clip_norm=config.train.grad_clip_norm, - scaler=scaler, - amp_enabled=amp_enabled, - dominant_freq_min_hz=config.train.dominant_freq_min_hz, - dominant_freq_max_hz=config.train.dominant_freq_max_hz, - ) - - print( - "Test metrics | " - f"loss={test_metrics['loss']:.6f} | " - f"rms={test_metrics['rms_error']:.6f} | " - f"dominant_freq_error={test_metrics['dominant_freq_error']:.6f} Hz" - ) - - -def main() -> None: - args = parse_args() - config = make_config(args.task) - - if args.epochs is not None: - config.train.epochs = args.epochs - if args.batch_size is not None: - config.data.batch_size = args.batch_size - if args.device is not None: - config.train.device = args.device - - train_model(config) - - -if __name__ == "__main__": - main() diff --git a/src_new/__pycache__/config.cpython-310.pyc b/src_new/__pycache__/config.cpython-310.pyc deleted file mode 100644 index a47c917..0000000 Binary files a/src_new/__pycache__/config.cpython-310.pyc and /dev/null differ diff --git a/src_new/__pycache__/dataset.cpython-310.pyc b/src_new/__pycache__/dataset.cpython-310.pyc deleted file mode 100644 index f2fccec..0000000 Binary files a/src_new/__pycache__/dataset.cpython-310.pyc and /dev/null differ diff --git a/src_new/__pycache__/evaluate.cpython-310.pyc b/src_new/__pycache__/evaluate.cpython-310.pyc deleted file mode 100644 index b17f0d2..0000000 Binary files a/src_new/__pycache__/evaluate.cpython-310.pyc and /dev/null differ diff --git a/src_new/__pycache__/model.cpython-310.pyc b/src_new/__pycache__/model.cpython-310.pyc deleted file mode 100644 index efd447d..0000000 Binary files a/src_new/__pycache__/model.cpython-310.pyc and /dev/null differ diff --git a/src_new/__pycache__/train.cpython-310.pyc b/src_new/__pycache__/train.cpython-310.pyc deleted file mode 100644 index 681b0a7..0000000 Binary files a/src_new/__pycache__/train.cpython-310.pyc and /dev/null differ diff --git a/src_new/config.py b/src_new/config.py deleted file mode 100644 index 467145c..0000000 --- a/src_new/config.py +++ /dev/null @@ -1,103 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass, field -from pathlib import Path - - -@dataclass -class DataConfig: - project_root: Path = field(default_factory=lambda: Path(__file__).resolve().parents[1]) - scenario: str = "Non_TMD" - train_split_name: str = "train" - val_split_name: str = "val" - test_split_name: str = "test" - csv_pattern: str = "*.csv" - - code_column: str = "code" - time_column: str = "time" - base_sensor_code: str = "WSMS00012" - base_axis: str = "value1" - response_sensor_code: str = "WSMS00007" - response_axis: str = "value3" - - max_sequence_length: int = 4096 - min_sequence_length: int = 512 - batch_size: int = 8 - num_workers: int = 0 - pin_memory: bool = True - use_weighted_train_sampler: bool = True - train_weight_power: float = 1.0 - train_weight_min: float = 0.5 - train_weight_max: float = 8.0 - low_frequency_emphasis_power: float = 1.25 - low_frequency_reference_hz: float = 1.0 - - use_steady_state_only: bool = True - steady_state_start_ratio: float = 0.50 - steady_state_min_samples: int = 256 - - interpolation_method: str = "linear" - normalization_eps: float = 1e-6 - - downloads_dir: Path = field(init=False) - scenario_dir: Path = field(init=False) - train_dir: Path = field(init=False) - val_dir: Path = field(init=False) - test_dir: Path = field(init=False) - - def __post_init__(self) -> None: - self.project_root = Path(self.project_root).resolve() - self.downloads_dir = self.project_root / "downloads" - self.scenario_dir = self.downloads_dir / self.scenario - self.train_dir = self.scenario_dir / self.train_split_name - self.val_dir = self.scenario_dir / self.val_split_name - self.test_dir = self.scenario_dir / self.test_split_name - - -@dataclass -class ModelConfig: - input_channels: int = 1 - tcn_channels: tuple[int, ...] = (32, 32, 64, 64) - kernel_size: int = 7 - dropout: float = 0.15 - dilation_base: int = 2 - pooled_feature_dim: int = 128 - - -@dataclass -class TrainConfig: - epochs: int = 120 - learning_rate: float = 1e-3 - weight_decay: float = 1e-4 - seed: int = 42 - device: str = "cuda" - grad_clip_norm: float = 1.0 - use_amp: bool = True - lr_scheduler_patience: int = 8 - lr_scheduler_factor: float = 0.5 - min_learning_rate: float = 1e-6 - early_stop_patience: int = 15 - checkpoint_dir: str = "checkpoints_rms" - best_model_name: str = "best_rms_model.pt" - history_name: str = "training_history.csv" - - -@dataclass -class LossConfig: - relative_rms_weight: float = 1.0 - log_rms_weight: float = 0.5 - mae_weight: float = 0.25 - waveform_l1_weight: float = 0.03 - waveform_huber_weight: float = 0.05 - - -@dataclass -class ExperimentConfig: - data: DataConfig = field(default_factory=DataConfig) - model: ModelConfig = field(default_factory=ModelConfig) - train: TrainConfig = field(default_factory=TrainConfig) - loss: LossConfig = field(default_factory=LossConfig) - - -def make_rms_forward_config() -> ExperimentConfig: - return ExperimentConfig() diff --git a/src_new/dataset.py b/src_new/dataset.py deleted file mode 100644 index 5295819..0000000 --- a/src_new/dataset.py +++ /dev/null @@ -1,388 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass, field -from pathlib import Path -import re -from typing import Any - -import numpy as np -import pandas as pd -import torch -import torch.nn.functional as F -from torch.utils.data import DataLoader, Dataset, WeightedRandomSampler - -try: - from .config import DataConfig, ExperimentConfig -except ImportError: - from config import DataConfig, ExperimentConfig - - -def calculate_rms(signal: np.ndarray) -> float: - signal = np.asarray(signal, dtype=np.float64).reshape(-1) - return float(np.sqrt(np.mean(np.square(signal)))) - - -def extract_frequency_hz(file_name: str) -> float | None: - match = re.search(r"(\d+(?:\.\d+)?)Hz", file_name, flags=re.IGNORECASE) - if match is None: - return None - return float(match.group(1)) - - -@dataclass -class NormalizationStats: - x_mean: torch.Tensor - x_std: torch.Tensor - y_wave_mean: torch.Tensor - y_wave_std: torch.Tensor - aux_mean: torch.Tensor - aux_std: torch.Tensor - target_mean: torch.Tensor - target_std: torch.Tensor - - -@dataclass -class TensorStandardScaler: - mean: torch.Tensor - std: torch.Tensor - - def transform(self, array: torch.Tensor | np.ndarray) -> torch.Tensor | np.ndarray: - if isinstance(array, torch.Tensor): - mean = self.mean.to(array.device, dtype=array.dtype) - std = self.std.to(array.device, dtype=array.dtype) - return (array - mean) / std - np_array = np.asarray(array, dtype=np.float32) - return (np_array - self.mean.cpu().numpy()) / self.std.cpu().numpy() - - def inverse_transform(self, array: torch.Tensor | np.ndarray) -> torch.Tensor | np.ndarray: - if isinstance(array, torch.Tensor): - mean = self.mean.to(array.device, dtype=array.dtype) - std = self.std.to(array.device, dtype=array.dtype) - return array * std + mean - np_array = np.asarray(array, dtype=np.float32) - return np_array * self.std.cpu().numpy() + self.mean.cpu().numpy() - - -@dataclass -class RMSRecord: - file_path: Path - split: str - time: torch.Tensor - x_full: torch.Tensor - x_model_input: torch.Tensor - y_model_target: torch.Tensor - x_rms: torch.Tensor - y_rms: torch.Tensor - target_value: torch.Tensor - aux_features: torch.Tensor - sample_weight: float - interpolation_count: int = 0 - - -@dataclass -class SplitLoadReport: - split: str - loaded_files: list[str] = field(default_factory=list) - skipped_files: list[tuple[str, str]] = field(default_factory=list) - interpolated_files: dict[str, int] = field(default_factory=dict) - - def to_lines(self) -> list[str]: - lines = [f"[{self.split}] loaded={len(self.loaded_files)} skipped={len(self.skipped_files)}"] - for file_name, reason in self.skipped_files: - lines.append(f" - skipped {file_name}: {reason}") - for file_name, count in self.interpolated_files.items(): - if count > 0: - lines.append(f" - interpolated {file_name}: missing_points={count}") - return lines - - -class RMSRegressionDataset(Dataset): - def __init__( - self, - records: list[RMSRecord], - normalization: NormalizationStats, - max_sequence_length: int, - ) -> None: - self.records = records - self.normalization = normalization - self.max_sequence_length = max_sequence_length - self.sample_weights = [record.sample_weight for record in records] - - def __len__(self) -> int: - return len(self.records) - - def __getitem__(self, index: int) -> dict[str, Any]: - record = self.records[index] - x = record.x_model_input - if x.shape[0] > self.max_sequence_length: - x = x[-self.max_sequence_length :] - - pad_length = self.max_sequence_length - x.shape[0] - if pad_length > 0: - x = F.pad(x.transpose(0, 1), (pad_length, 0), value=0.0).transpose(0, 1) - - x_norm = (x - self.normalization.x_mean) / self.normalization.x_std - y = record.y_model_target - if y.shape[0] > self.max_sequence_length: - y = y[-self.max_sequence_length :] - aux_norm = (record.aux_features - self.normalization.aux_mean) / self.normalization.aux_std - target_norm = (record.target_value - self.normalization.target_mean) / self.normalization.target_std - if pad_length > 0: - y = F.pad(y.transpose(0, 1), (pad_length, 0), value=0.0).transpose(0, 1) - y_wave_norm = (y - self.normalization.y_wave_mean) / self.normalization.y_wave_std - - valid_length = min(record.x_model_input.shape[0], self.max_sequence_length) - mask = torch.zeros(self.max_sequence_length, dtype=torch.float32) - mask[-valid_length:] = 1.0 - - return { - "x": x_norm, - "aux": aux_norm, - "y": target_norm, - "y_wave": y_wave_norm, - "y_raw": record.y_rms, - "x_rms_raw": record.x_rms, - "frequency_hz": torch.tensor(record.aux_features[0].item(), dtype=torch.float32), - "mask": mask, - "file_name": record.file_path.name, - "file_path": str(record.file_path), - "valid_length": torch.tensor(valid_length, dtype=torch.long), - "sample_weight": torch.tensor(record.sample_weight, dtype=torch.float32), - } - - -def list_split_files(config: DataConfig, split: str) -> list[Path]: - split_dir = getattr(config, f"{split}_dir") - return sorted(split_dir.glob(config.csv_pattern)) - - -def _value_frame(df: pd.DataFrame, config: DataConfig, sensor_code: str, value_column: str) -> pd.DataFrame: - sensor_df = df.loc[df[config.code_column] == sensor_code, [config.time_column, value_column]].copy() - sensor_df = sensor_df.sort_values(config.time_column) - sensor_df = sensor_df.drop_duplicates(subset=config.time_column, keep="first") - sensor_df[config.time_column] = sensor_df[config.time_column].astype("float64") - sensor_df[value_column] = sensor_df[value_column].astype("float32") - return sensor_df - - -def select_model_segment(signal: np.ndarray, config: DataConfig) -> np.ndarray: - if not config.use_steady_state_only: - return signal - start_index = int(len(signal) * config.steady_state_start_ratio) - max_start = max(0, len(signal) - config.steady_state_min_samples) - start_index = min(start_index, max_start) - return signal[start_index:] - - -def estimate_dominant_frequency(signal: np.ndarray, sampling_rate: float) -> float: - signal = np.asarray(signal, dtype=np.float64).reshape(-1) - if signal.size < 4: - return 0.0 - fft_values = np.fft.rfft(signal) - freqs = np.fft.rfftfreq(signal.size, d=1.0 / sampling_rate) - magnitudes = np.abs(fft_values) - magnitudes[0] = 0.0 - band_mask = (freqs >= 0.1) & (freqs <= 5.0) - if not np.any(band_mask): - return 0.0 - masked_magnitudes = np.where(band_mask, magnitudes, 0.0) - return float(freqs[int(np.argmax(masked_magnitudes))]) - - -def estimate_sampling_rate(time_values: np.ndarray) -> float: - if time_values.size < 2: - return 100.0 - dt = np.diff(time_values) - dt = dt[np.isfinite(dt)] - dt = dt[dt > 0.0] - if dt.size == 0: - return 100.0 - return float(1.0 / np.median(dt)) - - -def compute_sample_weight(y_rms: float, x_rms: float, freq_hz: float, config: DataConfig) -> float: - gain = y_rms / max(x_rms, config.normalization_eps) - low_freq_factor = (config.low_frequency_reference_hz / max(freq_hz, config.normalization_eps)) ** config.low_frequency_emphasis_power - weight = (gain ** config.train_weight_power) * low_freq_factor - return float(np.clip(weight, config.train_weight_min, config.train_weight_max)) - - -def load_split_records(config: DataConfig, split: str) -> tuple[list[RMSRecord], SplitLoadReport]: - records: list[RMSRecord] = [] - report = SplitLoadReport(split=split) - - for file_path in list_split_files(config, split): - df = pd.read_csv(file_path) - base_df = _value_frame(df, config, config.base_sensor_code, config.base_axis) - response_df = _value_frame(df, config, config.response_sensor_code, config.response_axis) - - if base_df.empty or response_df.empty: - report.skipped_files.append((file_path.name, "missing required sensor")) - continue - - aligned = base_df.rename(columns={config.base_axis: "base_signal"}) - aligned = aligned.merge( - response_df.rename(columns={config.response_axis: "response_signal"}), - on=config.time_column, - how="left", - ) - interpolation_count = int(aligned["response_signal"].isna().sum()) - aligned["response_signal"] = aligned["response_signal"].interpolate( - method=config.interpolation_method, - limit_direction="both", - ).ffill().bfill() - - if aligned["response_signal"].isna().any(): - report.skipped_files.append((file_path.name, "remaining NaN after interpolation")) - continue - - time_values = aligned[config.time_column].to_numpy(dtype=np.float64) - x_values = aligned["base_signal"].to_numpy(dtype=np.float32) - y_values = aligned["response_signal"].to_numpy(dtype=np.float32) - if len(x_values) < config.min_sequence_length: - report.skipped_files.append((file_path.name, f"sequence too short: {len(x_values)}")) - continue - - x_model = select_model_segment(x_values, config) - y_model = select_model_segment(y_values, config) - time_model = select_model_segment(time_values, config) - sampling_rate = estimate_sampling_rate(time_model) - file_frequency = extract_frequency_hz(file_path.name) - x_rms = calculate_rms(x_model) - y_rms = calculate_rms(y_model) - target_value = float(np.log(max(y_rms / max(x_rms, config.normalization_eps), config.normalization_eps))) - dominant_freq = estimate_dominant_frequency(x_model, sampling_rate) - feature_frequency = file_frequency if file_frequency is not None else dominant_freq - sample_weight = compute_sample_weight(y_rms=y_rms, x_rms=x_rms, freq_hz=feature_frequency, config=config) if split == "train" else 1.0 - - aux_features = torch.tensor( - [feature_frequency, x_rms, float(len(x_model)) / float(config.max_sequence_length)], - dtype=torch.float32, - ) - records.append( - RMSRecord( - file_path=file_path, - split=split, - time=torch.tensor(time_model, dtype=torch.float64), - x_full=torch.tensor(x_values[:, None], dtype=torch.float32), - x_model_input=torch.tensor(x_model[:, None], dtype=torch.float32), - y_model_target=torch.tensor(y_model[:, None], dtype=torch.float32), - x_rms=torch.tensor([x_rms], dtype=torch.float32), - y_rms=torch.tensor([y_rms], dtype=torch.float32), - target_value=torch.tensor([target_value], dtype=torch.float32), - aux_features=aux_features, - sample_weight=sample_weight, - interpolation_count=interpolation_count, - ) - ) - report.loaded_files.append(file_path.name) - if interpolation_count > 0: - report.interpolated_files[file_path.name] = interpolation_count - - if not records: - raise RuntimeError(f"No usable records found for split='{split}'.") - return records, report - - -def fit_normalization(records: list[RMSRecord], config: DataConfig) -> NormalizationStats: - x_all = torch.cat([record.x_model_input for record in records], dim=0) - y_all = torch.cat([record.y_model_target for record in records], dim=0) - aux_all = torch.stack([record.aux_features for record in records], dim=0) - target_all = torch.cat([record.target_value for record in records], dim=0) - - def safe_std(tensor: torch.Tensor, dim: int) -> torch.Tensor: - std = tensor.std(dim=dim, unbiased=False) - return torch.clamp(std, min=config.normalization_eps) - - return NormalizationStats( - x_mean=x_all.mean(dim=0), - x_std=safe_std(x_all, dim=0), - y_wave_mean=y_all.mean(dim=0), - y_wave_std=safe_std(y_all, dim=0), - aux_mean=aux_all.mean(dim=0), - aux_std=safe_std(aux_all, dim=0), - target_mean=target_all.mean(dim=0, keepdim=True), - target_std=safe_std(target_all, dim=0).view(1), - ) - - -def build_datasets( - config: ExperimentConfig | DataConfig, -) -> tuple[dict[str, RMSRegressionDataset], NormalizationStats, dict[str, SplitLoadReport]]: - data_config = config.data if isinstance(config, ExperimentConfig) else config - train_records, train_report = load_split_records(data_config, "train") - normalization = fit_normalization(train_records, data_config) - val_records, val_report = load_split_records(data_config, "val") - test_records, test_report = load_split_records(data_config, "test") - - datasets = { - "train": RMSRegressionDataset(train_records, normalization, data_config.max_sequence_length), - "val": RMSRegressionDataset(val_records, normalization, data_config.max_sequence_length), - "test": RMSRegressionDataset(test_records, normalization, data_config.max_sequence_length), - } - return datasets, normalization, {"train": train_report, "val": val_report, "test": test_report} - - -def build_dataloaders( - config: ExperimentConfig | DataConfig, -) -> tuple[dict[str, DataLoader], dict[str, RMSRegressionDataset], dict[str, SplitLoadReport]]: - data_config = config.data if isinstance(config, ExperimentConfig) else config - datasets, _, reports = build_datasets(config) - - train_sampler = None - train_shuffle = True - if data_config.use_weighted_train_sampler: - weights = torch.tensor(datasets["train"].sample_weights, dtype=torch.double) - train_sampler = WeightedRandomSampler(weights, num_samples=len(weights), replacement=True) - train_shuffle = False - - loaders = { - "train": DataLoader( - datasets["train"], - batch_size=data_config.batch_size, - shuffle=train_shuffle, - sampler=train_sampler, - num_workers=data_config.num_workers, - pin_memory=data_config.pin_memory, - ), - "val": DataLoader( - datasets["val"], - batch_size=data_config.batch_size, - shuffle=False, - num_workers=data_config.num_workers, - pin_memory=data_config.pin_memory, - ), - "test": DataLoader( - datasets["test"], - batch_size=data_config.batch_size, - shuffle=False, - num_workers=data_config.num_workers, - pin_memory=data_config.pin_memory, - ), - } - return loaders, datasets, reports - - -def get_dataloaders( - config: ExperimentConfig | DataConfig, -) -> tuple[ - dict[str, DataLoader], - dict[str, RMSRegressionDataset], - dict[str, SplitLoadReport], - TensorStandardScaler, - TensorStandardScaler, - TensorStandardScaler, -]: - loaders, datasets, reports = build_dataloaders(config) - normalization = datasets["train"].normalization - x_scaler = TensorStandardScaler(normalization.x_mean, normalization.x_std) - aux_scaler = TensorStandardScaler(normalization.aux_mean, normalization.aux_std) - y_scaler = TensorStandardScaler(normalization.target_mean, normalization.target_std) - return loaders, datasets, reports, x_scaler, aux_scaler, y_scaler - - -def report_to_text(reports: dict[str, SplitLoadReport]) -> str: - lines: list[str] = [] - for split in ("train", "val", "test"): - lines.extend(reports[split].to_lines()) - return "\n".join(lines) diff --git a/src_new/evaluate.py b/src_new/evaluate.py deleted file mode 100644 index be1c595..0000000 --- a/src_new/evaluate.py +++ /dev/null @@ -1,222 +0,0 @@ -from __future__ import annotations - -import argparse -from pathlib import Path - -import matplotlib.pyplot as plt -import numpy as np -import pandas as pd -import torch - -try: - from .config import ExperimentConfig, make_rms_forward_config - from .dataset import get_dataloaders, report_to_text - from .model import build_model -except ImportError: - from config import ExperimentConfig, make_rms_forward_config - from dataset import get_dataloaders, report_to_text - from model import build_model - - -def parse_args() -> argparse.Namespace: - parser = argparse.ArgumentParser(description="Evaluate direct RMS regression model on harmonic data.") - parser.add_argument("--split", choices=("train", "val", "test"), default="val") - parser.add_argument("--sample-index", type=int, default=0) - parser.add_argument("--device", type=str, default=None) - parser.add_argument("--checkpoint", type=str, default=None) - parser.add_argument("--all-samples", action="store_true") - return parser.parse_args() - - -def resolve_device(device_name: str | None, config: ExperimentConfig) -> torch.device: - requested = device_name or config.train.device - if requested.startswith("cuda") and not torch.cuda.is_available(): - return torch.device("cpu") - return torch.device(requested) - - -def resolve_checkpoint_path(config: ExperimentConfig, checkpoint_arg: str | None) -> Path: - if checkpoint_arg: - return Path(checkpoint_arg).resolve() - return config.data.project_root / config.train.checkpoint_dir / "forward_rms" / config.train.best_model_name - - -def relative_percent_error(true_value: float, pred_value: float) -> float: - denominator = max(abs(true_value), 1e-12) - return abs(pred_value - true_value) / denominator * 100.0 - - -def inverse_waveform(array: torch.Tensor, mean: torch.Tensor, std: torch.Tensor) -> torch.Tensor: - mean = mean.to(array.device, dtype=array.dtype) - std = std.to(array.device, dtype=array.dtype) - return array * std + mean - - -def predict_sample( - model: torch.nn.Module, - sample: dict[str, object], - device: torch.device, - y_scaler, - x_scaler, - y_wave_mean: torch.Tensor, - y_wave_std: torch.Tensor, -) -> dict[str, float | str | np.ndarray]: - x = sample["x"].unsqueeze(0).to(device) - aux = sample["aux"].unsqueeze(0).to(device) - mask = sample["mask"].unsqueeze(0).to(device) - x_rms_raw = sample["x_rms_raw"].unsqueeze(0).to(device) - y_true = float(sample["y_raw"].detach().cpu().numpy().reshape(-1)[0]) - frequency_hz = float(sample["frequency_hz"].item()) - file_name = str(sample["file_name"]) - valid_mask = sample["mask"].detach().cpu().numpy() > 0.5 - - with torch.no_grad(): - outputs = model(x, mask, aux) - - pred_target = y_scaler.inverse_transform(outputs["rms"]) - y_pred = float((torch.exp(pred_target) * x_rms_raw).detach().cpu().numpy().reshape(-1)[0]) - pred_wave = inverse_waveform(outputs["waveform"], y_wave_mean, y_wave_std).detach().cpu().numpy().reshape(-1)[valid_mask] - true_wave = inverse_waveform(sample["y_wave"], y_wave_mean.cpu(), y_wave_std.cpu()).detach().cpu().numpy().reshape(-1)[valid_mask] - x_wave = x_scaler.inverse_transform(sample["x"]).reshape(-1)[valid_mask] - return { - "file_name": file_name, - "frequency_hz": frequency_hz, - "true_rms": y_true, - "pred_rms": y_pred, - "relative_error_percent": float(relative_percent_error(y_true, y_pred)), - "x_wave": x_wave.numpy().reshape(-1), - "true_wave": true_wave, - "pred_wave": pred_wave, - } - - -def evaluate_sample( - model: torch.nn.Module, - sample: dict[str, object], - device: torch.device, - y_scaler, - x_scaler, - y_wave_mean: torch.Tensor, - y_wave_std: torch.Tensor, -) -> dict[str, float | str]: - result = predict_sample(model, sample, device, y_scaler, x_scaler, y_wave_mean, y_wave_std) - return { - "file_name": str(result["file_name"]), - "frequency_hz": float(result["frequency_hz"]), - "true_rms": float(result["true_rms"]), - "pred_rms": float(result["pred_rms"]), - "relative_error_percent": float(result["relative_error_percent"]), - } - - -def save_single_sample_figure(result: dict[str, float | str | np.ndarray], split: str, save_dir: Path, sample_index: int) -> Path: - true_rms = float(result["true_rms"]) - pred_rms = float(result["pred_rms"]) - error_percent = relative_percent_error(true_rms, pred_rms) - x_signal = np.asarray(result["x_wave"], dtype=np.float64).reshape(-1) - true_wave = np.asarray(result["true_wave"], dtype=np.float64).reshape(-1) - pred_wave = np.asarray(result["pred_wave"], dtype=np.float64).reshape(-1) - time_steps = np.arange(x_signal.shape[0], dtype=np.float64) - - fig, axes = plt.subplots(3, 1, figsize=(12, 10)) - fig.suptitle(f"Forward RMS regression | {split} | {result['file_name']}") - axes[0].plot(time_steps, x_signal, linewidth=1.2) - axes[0].set_title("Input Base Excitation (Steady-State Segment)") - axes[0].set_xlabel("Time Step") - axes[0].set_ylabel("Acceleration") - axes[0].grid(True, alpha=0.3) - axes[1].plot(time_steps, true_wave, label="True response", linewidth=1.2, color="tab:blue") - axes[1].plot(time_steps, pred_wave, label="Pred response", linewidth=1.2, color="tab:orange", alpha=0.85) - axes[1].set_title("Auxiliary Waveform Head: True vs Predicted Top Response") - axes[1].set_xlabel("Time Step") - axes[1].set_ylabel("Acceleration") - axes[1].grid(True, alpha=0.3) - axes[1].legend() - axes[2].bar(["True RMS", "Pred RMS"], [true_rms, pred_rms], color=["tab:blue", "tab:orange"]) - axes[2].set_title(f"True RMS: {true_rms:.4f}, Pred RMS: {pred_rms:.4f}, Error: {error_percent:.2f}%") - axes[2].set_ylabel("RMS") - axes[2].grid(True, axis="y", alpha=0.3) - plt.tight_layout() - save_dir.mkdir(parents=True, exist_ok=True) - figure_path = save_dir / f"evaluation_{split}_s{sample_index}_waveform.png" - plt.savefig(figure_path, dpi=180, bbox_inches="tight") - plt.show() - plt.close(fig) - return figure_path - - -def main() -> None: - args = parse_args() - config = make_rms_forward_config() - device = resolve_device(args.device, config) - checkpoint_path = resolve_checkpoint_path(config, args.checkpoint) - loaders, datasets, reports, x_scaler, aux_scaler, y_scaler = get_dataloaders(config) - print(report_to_text(reports)) - - if not checkpoint_path.exists(): - raise FileNotFoundError(f"Checkpoint not found: {checkpoint_path}") - - dataset = datasets[args.split] - sample = dataset[args.sample_index] - normalization = datasets["train"].normalization - model = build_model(config).to(device) - checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) - model.load_state_dict(checkpoint["model_state_dict"]) - model.eval() - save_dir = config.data.project_root / "evaluation_outputs" / "forward_rms" - if args.all_samples: - rows = [ - evaluate_sample( - model, - dataset[index], - device, - y_scaler, - x_scaler, - normalization.y_wave_mean, - normalization.y_wave_std, - ) - for index in range(len(dataset)) - ] - result_df = pd.DataFrame(rows) - result_df = result_df.sort_values(["relative_error_percent", "file_name"]).reset_index(drop=True) - csv_path = save_dir / f"evaluation_{args.split}_all_samples.csv" - save_dir.mkdir(parents=True, exist_ok=True) - result_df.to_csv(csv_path, index=False) - 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: {checkpoint_path}") - print(f"Summary: {summary}") - print(f"CSV saved to: {csv_path}") - print(result_df.to_string(index=False)) - return - - result = predict_sample( - model, - sample, - device, - y_scaler, - x_scaler, - normalization.y_wave_mean, - normalization.y_wave_std, - ) - figure_path = save_single_sample_figure( - result=result, - split=args.split, - save_dir=save_dir, - sample_index=args.sample_index, - ) - print(f"Checkpoint: {checkpoint_path}") - print(f"Sample file: {result['file_name']}") - print(f"True RMS: {result['true_rms']:.6f}") - print(f"Pred RMS: {result['pred_rms']:.6f}") - print(f"Relative RMS Error (%): {result['relative_error_percent']:.4f}") - print(f"Figure saved to: {figure_path}") - - -if __name__ == "__main__": - main() diff --git a/src_new/model.py b/src_new/model.py deleted file mode 100644 index 31f1f25..0000000 --- a/src_new/model.py +++ /dev/null @@ -1,140 +0,0 @@ -from __future__ import annotations - -from dataclasses import dataclass - -import torch -import torch.nn as nn -from torch.nn.utils import weight_norm - -try: - from .config import ExperimentConfig, ModelConfig -except ImportError: - from config import ExperimentConfig, ModelConfig - - -class Chomp1d(nn.Module): - def __init__(self, chomp_size: int) -> None: - super().__init__() - self.chomp_size = chomp_size - - def forward(self, x: torch.Tensor) -> torch.Tensor: - if self.chomp_size == 0: - return x - return x[:, :, :-self.chomp_size].contiguous() - - -class CausalConv1d(nn.Module): - def __init__(self, in_channels: int, out_channels: int, kernel_size: int, dilation: int = 1) -> None: - super().__init__() - padding = (kernel_size - 1) * dilation - self.net = nn.Sequential( - weight_norm( - nn.Conv1d( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=kernel_size, - padding=padding, - dilation=dilation, - ) - ), - Chomp1d(padding), - ) - - def forward(self, x: torch.Tensor) -> torch.Tensor: - return self.net(x) - - -class TemporalBlock(nn.Module): - def __init__(self, in_channels: int, out_channels: int, kernel_size: int, dilation: int, dropout: float) -> None: - super().__init__() - self.conv1 = CausalConv1d(in_channels, out_channels, kernel_size, dilation=dilation) - self.act1 = nn.GELU() - self.dropout1 = nn.Dropout(dropout) - self.conv2 = CausalConv1d(out_channels, out_channels, kernel_size, dilation=dilation) - self.act2 = nn.GELU() - self.dropout2 = nn.Dropout(dropout) - self.residual = nn.Conv1d(in_channels, out_channels, kernel_size=1) if in_channels != out_channels else nn.Identity() - self.final_act = nn.GELU() - - def forward(self, x: torch.Tensor) -> torch.Tensor: - residual = self.residual(x) - out = self.dropout1(self.act1(self.conv1(x))) - out = self.dropout2(self.act2(self.conv2(out))) - return self.final_act(out + residual) - - -@dataclass -class ReceptiveFieldInfo: - receptive_field: int - dilations: tuple[int, ...] - - -class RMSRegressionTCN(nn.Module): - def __init__(self, config: ModelConfig) -> None: - super().__init__() - blocks: list[nn.Module] = [] - in_channels = config.input_channels - dilations: list[int] = [] - - for level, out_channels in enumerate(config.tcn_channels): - dilation = config.dilation_base ** level - dilations.append(dilation) - blocks.append( - TemporalBlock( - in_channels=in_channels, - out_channels=out_channels, - kernel_size=config.kernel_size, - dilation=dilation, - dropout=config.dropout, - ) - ) - in_channels = out_channels - - self.encoder = nn.Sequential(*blocks) - self.rms_head = nn.Sequential( - nn.Linear(in_channels * 2 + 3, config.pooled_feature_dim), - nn.GELU(), - nn.Dropout(config.dropout), - nn.Linear(config.pooled_feature_dim, config.pooled_feature_dim), - nn.GELU(), - nn.Dropout(config.dropout), - nn.Linear(config.pooled_feature_dim, 1), - ) - self.waveform_head = nn.Sequential( - nn.Conv1d(in_channels, in_channels, kernel_size=1), - nn.GELU(), - nn.Dropout(config.dropout), - nn.Conv1d(in_channels, 1, kernel_size=1), - ) - self.receptive_field_info = compute_receptive_field(config) - - def forward(self, x: torch.Tensor, mask: torch.Tensor, aux: torch.Tensor) -> dict[str, torch.Tensor]: - if x.ndim != 3: - raise ValueError(f"Expected x shape (batch, seq, channels), got {tuple(x.shape)}") - features = self.encoder(x.transpose(1, 2)) - mask_1d = mask.unsqueeze(1) - masked_features = features * mask_1d - valid_count = mask_1d.sum(dim=2).clamp_min(1.0) - mean_pool = masked_features.sum(dim=2) / valid_count - masked_for_max = features.masked_fill(mask_1d == 0.0, float("-inf")) - max_pool = masked_for_max.max(dim=2).values - max_pool = torch.where(torch.isfinite(max_pool), max_pool, torch.zeros_like(max_pool)) - fused = torch.cat([mean_pool, max_pool, aux], dim=1) - rms_prediction = self.rms_head(fused) - waveform_prediction = self.waveform_head(features).transpose(1, 2) - return {"rms": rms_prediction, "waveform": waveform_prediction} - - -def compute_receptive_field(config: ModelConfig) -> ReceptiveFieldInfo: - receptive_field = 1 - dilations: list[int] = [] - for level, _ in enumerate(config.tcn_channels): - dilation = config.dilation_base ** level - dilations.append(dilation) - receptive_field += 2 * (config.kernel_size - 1) * dilation - return ReceptiveFieldInfo(receptive_field=receptive_field, dilations=tuple(dilations)) - - -def build_model(config: ExperimentConfig | ModelConfig) -> RMSRegressionTCN: - model_config = config.model if isinstance(config, ExperimentConfig) else config - return RMSRegressionTCN(model_config) diff --git a/src_new/train.py b/src_new/train.py deleted file mode 100644 index 179ea1f..0000000 --- a/src_new/train.py +++ /dev/null @@ -1,327 +0,0 @@ -from __future__ import annotations - -import argparse -import random -from dataclasses import asdict -from pathlib import Path -from typing import Any - -import pandas as pd -import torch -import torch.nn.functional as F -from torch import nn -from torch.optim import AdamW -from torch.optim.lr_scheduler import ReduceLROnPlateau - -try: - from .config import ExperimentConfig, make_rms_forward_config - from .dataset import build_dataloaders, report_to_text - from .model import build_model, compute_receptive_field -except ImportError: - from config import ExperimentConfig, make_rms_forward_config - from dataset import build_dataloaders, report_to_text - from model import build_model, compute_receptive_field - - -def parse_args() -> argparse.Namespace: - parser = argparse.ArgumentParser(description="Train TCN for direct RMS regression.") - parser.add_argument("--epochs", type=int, default=None) - parser.add_argument("--batch-size", type=int, default=None) - parser.add_argument("--device", type=str, default=None) - return parser.parse_args() - - -def set_seed(seed: int) -> None: - random.seed(seed) - torch.manual_seed(seed) - torch.cuda.manual_seed_all(seed) - - -def resolve_device(device_name: str) -> torch.device: - if device_name.startswith("cuda") and not torch.cuda.is_available(): - return torch.device("cpu") - return torch.device(device_name) - - -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 denormalize_target( - pred_norm: torch.Tensor, - target_norm: torch.Tensor, - target_mean: torch.Tensor, - target_std: torch.Tensor, -) -> tuple[torch.Tensor, torch.Tensor]: - pred = pred_norm * target_std + target_mean - target = target_norm * target_std + target_mean - return pred, target - - -def relative_rms_error(pred: torch.Tensor, target: torch.Tensor) -> torch.Tensor: - return torch.abs(pred - target) / torch.clamp(target.abs(), min=1e-6) - - -class RMSRegressionLoss(nn.Module): - def __init__( - self, - relative_rms_weight: float, - log_rms_weight: float, - mae_weight: float, - waveform_l1_weight: float, - waveform_huber_weight: float, - ) -> None: - super().__init__() - self.relative_rms_weight = relative_rms_weight - self.log_rms_weight = log_rms_weight - self.mae_weight = mae_weight - self.waveform_l1_weight = waveform_l1_weight - self.waveform_huber_weight = waveform_huber_weight - - def forward( - self, - pred: torch.Tensor, - target: torch.Tensor, - pred_wave: torch.Tensor, - target_wave: torch.Tensor, - mask: torch.Tensor, - ) -> dict[str, torch.Tensor]: - rel_loss = relative_rms_error(pred, target).mean() - log_loss = F.huber_loss(torch.log(torch.clamp(pred, min=1e-6)), torch.log(torch.clamp(target, min=1e-6))) - mae_loss = F.l1_loss(pred, target) - mask_expanded = mask.unsqueeze(-1) - valid_count = mask_expanded.sum().clamp_min(1.0) - wave_residual = (pred_wave - target_wave) * mask_expanded - waveform_l1 = torch.abs(wave_residual).sum() / valid_count - waveform_huber = F.huber_loss(pred_wave * mask_expanded, target_wave * mask_expanded, reduction="sum") / valid_count - total = ( - self.relative_rms_weight * rel_loss - + self.log_rms_weight * log_loss - + self.mae_weight * mae_loss - + self.waveform_l1_weight * waveform_l1 - + self.waveform_huber_weight * waveform_huber - ) - return { - "total": total, - "relative": rel_loss, - "log": log_loss, - "mae": mae_loss, - "waveform_l1": waveform_l1, - "waveform_huber": waveform_huber, - } - - -def run_epoch( - model: nn.Module, - dataloader: torch.utils.data.DataLoader, - optimizer: AdamW | None, - criterion: RMSRegressionLoss, - target_mean: torch.Tensor, - target_std: torch.Tensor, - device: torch.device, - grad_clip_norm: float, - scaler: torch.cuda.amp.GradScaler, - amp_enabled: bool, -) -> dict[str, float]: - is_train = optimizer is not None - model.train(is_train) - - total_loss_sum = 0.0 - relative_loss_sum = 0.0 - log_loss_sum = 0.0 - mae_loss_sum = 0.0 - waveform_l1_sum = 0.0 - waveform_huber_sum = 0.0 - rms_error_sum = 0.0 - sample_count = 0 - - for batch in dataloader: - x = batch["x"].to(device) - aux = batch["aux"].to(device) - mask = batch["mask"].to(device) - y_norm = batch["y"].to(device) - y_wave = batch["y_wave"].to(device) - x_rms_raw = batch["x_rms_raw"].to(device) - y_rms_raw = batch["y_raw"].to(device) - - if is_train: - optimizer.zero_grad(set_to_none=True) - - with torch.amp.autocast(device_type=device.type, enabled=amp_enabled): - outputs = model(x, mask, aux) - pred_target, _ = denormalize_target(outputs["rms"], y_norm, target_mean, target_std) - pred = torch.exp(pred_target) * x_rms_raw - target = y_rms_raw - losses = criterion(pred, target, outputs["waveform"], y_wave, mask) - - if is_train: - scaler.scale(losses["total"]).backward() - scaler.unscale_(optimizer) - torch.nn.utils.clip_grad_norm_(model.parameters(), grad_clip_norm) - scaler.step(optimizer) - scaler.update() - - batch_size = x.shape[0] - total_loss_sum += losses["total"].detach().item() * batch_size - relative_loss_sum += losses["relative"].detach().item() * batch_size - log_loss_sum += losses["log"].detach().item() * batch_size - mae_loss_sum += losses["mae"].detach().item() * batch_size - waveform_l1_sum += losses["waveform_l1"].detach().item() * batch_size - waveform_huber_sum += losses["waveform_huber"].detach().item() * batch_size - rms_error_sum += relative_rms_error(pred.detach(), target.detach()).mean().item() * batch_size - sample_count += batch_size - - return { - "loss": total_loss_sum / sample_count, - "relative_loss": relative_loss_sum / sample_count, - "log_loss": log_loss_sum / sample_count, - "mae_loss": mae_loss_sum / sample_count, - "waveform_l1": waveform_l1_sum / sample_count, - "waveform_huber": waveform_huber_sum / sample_count, - "rms_error": rms_error_sum / sample_count, - } - - -def checkpoint_paths(config: ExperimentConfig) -> tuple[Path, Path]: - root = config.data.project_root / config.train.checkpoint_dir / "forward_rms" - root.mkdir(parents=True, exist_ok=True) - return root / config.train.best_model_name, root / config.train.history_name - -def train_model(config: ExperimentConfig) -> None: - set_seed(config.train.seed) - device = resolve_device(config.train.device) - loaders, datasets, reports = build_dataloaders(config) - train_dataset = datasets["train"] - normalization = train_dataset.normalization - target_mean = normalization.target_mean.to(device).view(1, 1) - target_std = normalization.target_std.to(device).view(1, 1) - - model = build_model(config).to(device) - receptive_field = compute_receptive_field(config.model) - optimizer = AdamW(model.parameters(), lr=config.train.learning_rate, weight_decay=config.train.weight_decay) - scheduler = ReduceLROnPlateau( - optimizer, - mode="min", - factor=config.train.lr_scheduler_factor, - patience=config.train.lr_scheduler_patience, - min_lr=config.train.min_learning_rate, - ) - criterion = RMSRegressionLoss( - relative_rms_weight=config.loss.relative_rms_weight, - log_rms_weight=config.loss.log_rms_weight, - mae_weight=config.loss.mae_weight, - waveform_l1_weight=config.loss.waveform_l1_weight, - waveform_huber_weight=config.loss.waveform_huber_weight, - ) - amp_enabled = config.train.use_amp and device.type == "cuda" - scaler = torch.cuda.amp.GradScaler(enabled=amp_enabled) - - best_val_error = float("inf") - epochs_without_improvement = 0 - history: list[dict[str, float]] = [] - best_model_path, history_path = checkpoint_paths(config) - - print(f"Device: {device}") - print(f"TCN receptive field: {receptive_field.receptive_field} samples, dilations={receptive_field.dilations}") - print(report_to_text(reports)) - - for epoch in range(1, config.train.epochs + 1): - train_metrics = run_epoch( - model=model, - dataloader=loaders["train"], - optimizer=optimizer, - criterion=criterion, - target_mean=target_mean, - target_std=target_std, - device=device, - grad_clip_norm=config.train.grad_clip_norm, - scaler=scaler, - amp_enabled=amp_enabled, - ) - val_metrics = run_epoch( - model=model, - dataloader=loaders["val"], - optimizer=None, - criterion=criterion, - target_mean=target_mean, - target_std=target_std, - device=device, - grad_clip_norm=config.train.grad_clip_norm, - scaler=scaler, - amp_enabled=amp_enabled, - ) - scheduler.step(val_metrics["rms_error"]) - current_lr = optimizer.param_groups[0]["lr"] - - history_row = { - "epoch": epoch, - "lr": current_lr, - "train_loss": train_metrics["loss"], - "train_relative_loss": train_metrics["relative_loss"], - "train_log_loss": train_metrics["log_loss"], - "train_mae_loss": train_metrics["mae_loss"], - "train_waveform_l1": train_metrics["waveform_l1"], - "train_waveform_huber": train_metrics["waveform_huber"], - "train_rms_error": train_metrics["rms_error"], - "val_loss": val_metrics["loss"], - "val_relative_loss": val_metrics["relative_loss"], - "val_log_loss": val_metrics["log_loss"], - "val_mae_loss": val_metrics["mae_loss"], - "val_waveform_l1": val_metrics["waveform_l1"], - "val_waveform_huber": val_metrics["waveform_huber"], - "val_rms_error": val_metrics["rms_error"], - } - history.append(history_row) - - print( - f"Epoch {epoch:03d} | train_loss={train_metrics['loss']:.6f} | " - f"val_loss={val_metrics['loss']:.6f} | val_rms_error={val_metrics['rms_error']:.6f} | lr={current_lr:.2e}" - ) - - if val_metrics["rms_error"] < best_val_error: - best_val_error = val_metrics["rms_error"] - epochs_without_improvement = 0 - torch.save( - { - "epoch": epoch, - "model_state_dict": model.state_dict(), - "optimizer_state_dict": optimizer.state_dict(), - "best_val_rms_error": best_val_error, - "config": serialize_for_checkpoint(asdict(config)), - }, - best_model_path, - ) - else: - epochs_without_improvement += 1 - if epochs_without_improvement >= config.train.early_stop_patience: - print(f"Early stopping triggered after {epoch} epochs.") - break - - history_df = pd.DataFrame(history) - history_df.to_csv(history_path, index=False) - print(f"Best model saved to: {best_model_path}") - print(f"Training history saved to: {history_path}") - - -def main() -> None: - args = parse_args() - config = make_rms_forward_config() - if args.epochs is not None: - config.train.epochs = args.epochs - if args.batch_size is not None: - config.data.batch_size = args.batch_size - if args.device is not None: - config.train.device = args.device - train_model(config) - - -if __name__ == "__main__": - main() diff --git a/src_old/__init__.py b/src_old/__init__.py deleted file mode 100644 index e69de29..0000000 diff --git a/src_old/__pycache__/config.cpython-310.pyc b/src_old/__pycache__/config.cpython-310.pyc deleted file mode 100644 index 2f2bd90..0000000 Binary files a/src_old/__pycache__/config.cpython-310.pyc and /dev/null differ diff --git a/src_old/__pycache__/config.cpython-314.pyc b/src_old/__pycache__/config.cpython-314.pyc deleted file mode 100644 index 26527ba..0000000 Binary files a/src_old/__pycache__/config.cpython-314.pyc and /dev/null differ diff --git a/src_old/__pycache__/dataset.cpython-310.pyc b/src_old/__pycache__/dataset.cpython-310.pyc deleted file mode 100644 index 62ffa0f..0000000 Binary files a/src_old/__pycache__/dataset.cpython-310.pyc and /dev/null differ diff --git a/src_old/__pycache__/dataset.cpython-314.pyc b/src_old/__pycache__/dataset.cpython-314.pyc deleted file mode 100644 index ac9e1d2..0000000 Binary files a/src_old/__pycache__/dataset.cpython-314.pyc and /dev/null differ diff --git a/src_old/__pycache__/evaluate.cpython-314.pyc b/src_old/__pycache__/evaluate.cpython-314.pyc deleted file mode 100644 index 9602c93..0000000 Binary files a/src_old/__pycache__/evaluate.cpython-314.pyc and /dev/null differ diff --git a/src_old/__pycache__/model.cpython-310.pyc b/src_old/__pycache__/model.cpython-310.pyc deleted file mode 100644 index 7c7a526..0000000 Binary files a/src_old/__pycache__/model.cpython-310.pyc and /dev/null differ diff --git a/src_old/__pycache__/model.cpython-314.pyc b/src_old/__pycache__/model.cpython-314.pyc deleted file mode 100644 index db0b38a..0000000 Binary files a/src_old/__pycache__/model.cpython-314.pyc and /dev/null differ diff --git a/src_old/__pycache__/train.cpython-310.pyc b/src_old/__pycache__/train.cpython-310.pyc deleted file mode 100644 index b3f47ed..0000000 Binary files a/src_old/__pycache__/train.cpython-310.pyc and /dev/null differ diff --git a/src_old/__pycache__/train.cpython-314.pyc b/src_old/__pycache__/train.cpython-314.pyc deleted file mode 100644 index 17071e3..0000000 Binary files a/src_old/__pycache__/train.cpython-314.pyc and /dev/null differ diff --git a/src_old/config.py b/src_old/config.py deleted file mode 100644 index f79f9be..0000000 --- a/src_old/config.py +++ /dev/null @@ -1,36 +0,0 @@ -import os - -# Data Configuration -# 使用基于当前文件的绝对路径拼接,以防止你在不同目录下运行报错 -DATA_DIR = os.path.join(os.path.dirname(os.path.dirname(__file__)), 'downloads') -SEQ_LEN = 512 # 滑动窗口的长度 (时间序列步数) -STEP_SIZE = 20 # 滑动窗口的步长 -BATCH_SIZE = 128 - -# Features Configuration -INPUT_SENSOR = 'WSMS00012' -OUTPUT_SENSORS = ['WSMS00007', 'WSMS00008', 'WSMS00009', 'WSMS00010', 'WSMS00011'] -INPUT_AXIS = 'value1' # 底部传感器输入轴(课程要求:X轴) -OUTPUT_AXIS = 'value3' # 目标传感器输出轴 - -# Model Configuration -CHANNELS = [64, 64, 128, 128, 256, 256] # TCN 各层通道数 -KERNEL_SIZE = 5 -DROPOUT = 0.1 - -# Training Configuration -LEARNING_RATE = 3e-4 -EPOCHS = 50 -WEIGHT_DECAY = 1e-4 -ENABLE_EARLY_STOP = False -EARLY_STOP_PATIENCE = 15 -SPECTRAL_LOSS_WEIGHT = 0.25 -CORR_LOSS_WEIGHT = 0.05 -MSE_LOSS_WEIGHT = 0.5 -STD_LOSS_WEIGHT = 0.3 -PEAK_LOSS_WEIGHT = 0.2 -CORR_LOSS_EPS = 1e-8 - -# Device -import torch -DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu' diff --git a/src_old/dataset.py b/src_old/dataset.py deleted file mode 100644 index 1c87f65..0000000 --- a/src_old/dataset.py +++ /dev/null @@ -1,188 +0,0 @@ -import os -import glob -import pandas as pd -import numpy as np -import torch -from torch.utils.data import Dataset, DataLoader -from sklearn.preprocessing import StandardScaler -from config import * - -class MultiOutputStandardizer: - """Per-channel standardization that ignores missing labels via masks.""" - def __init__(self, n_outputs): - self.n_outputs = n_outputs - self.mean_ = np.zeros(n_outputs, dtype=np.float32) - self.scale_ = np.ones(n_outputs, dtype=np.float32) - self.fitted = False - - def fit(self, y_sequences, mask_sequences): - means = [] - scales = [] - for c in range(self.n_outputs): - valid_values = [] - for y_seq, m_seq in zip(y_sequences, mask_sequences): - valid = m_seq[:, c] > 0.5 - if np.any(valid): - valid_values.append(y_seq[valid, c]) - if len(valid_values) == 0: - means.append(0.0) - scales.append(1.0) - continue - vals = np.concatenate(valid_values, axis=0) - mean = float(np.mean(vals)) - std = float(np.std(vals)) - if std < 1e-6: - std = 1.0 - means.append(mean) - scales.append(std) - - self.mean_ = np.asarray(means, dtype=np.float32) - self.scale_ = np.asarray(scales, dtype=np.float32) - self.fitted = True - - def transform(self, y): - if not self.fitted: - raise RuntimeError("MultiOutputStandardizer must be fitted before transform.") - return (y - self.mean_) / self.scale_ - - def inverse_transform(self, y): - if not self.fitted: - raise RuntimeError("MultiOutputStandardizer must be fitted before inverse_transform.") - return y * self.scale_ + self.mean_ - - -class BuildingDataset(Dataset): - def __init__(self, file_paths, seq_len, step_size, scaler_X=None, scaler_Y=None, fit_scaler=False): - self.seq_len = seq_len - self.X_data = [] - self.Y_data = [] - self.M_data = [] - self.scaler_X = scaler_X if scaler_X is not None else StandardScaler() - self.scaler_Y = scaler_Y if scaler_Y is not None else MultiOutputStandardizer(len(OUTPUT_SENSORS)) - - raw_X = [] - raw_Y = [] - raw_M = [] - - for f in file_paths: - # 读取数据 - df = pd.read_csv(f) - # 使用长表格式: code, type, time, value1, value2, value3 - - # 提取 012 的输入轴作为基准 - df_in = df[df['code'] == INPUT_SENSOR][['time', INPUT_AXIS]].rename(columns={INPUT_AXIS: 'input_signal'}) - if len(df_in) == 0: - # 若无输入传感器(如自由衰减数据),则补零 - df_in = pd.DataFrame({'time': df['time'].unique()}) - df_in['input_signal'] = 0.0 - - # 提取 007~011 的 Z 轴并按时间戳逐步合并 (使用 left join 确保以 df_in 的时间为基准) - df_merged = df_in - for sens in OUTPUT_SENSORS: - df_out_sens = df[df['code'] == sens][['time', OUTPUT_AXIS]].rename(columns={OUTPUT_AXIS: f'out_{sens}'}) - df_merged = pd.merge(df_merged, df_out_sens, on='time', how='left') - - df_merged = df_merged.sort_values('time').reset_index(drop=True) - - if len(df_merged) == 0: - print(f"Warning: Skipping file {f} due to no overlapping timestamps across required sensors.") - continue - - x_seq = df_merged['input_signal'].values.reshape(-1, 1).astype(np.float32) - - # 提取所有 target 传感器列与可用性掩码 - out_cols = [f'out_{sens}' for sens in OUTPUT_SENSORS] - y_seq = np.zeros((len(df_merged), len(OUTPUT_SENSORS)), dtype=np.float32) - m_seq = np.zeros((len(df_merged), len(OUTPUT_SENSORS)), dtype=np.float32) - for c, col in enumerate(out_cols): - series = df_merged[col] - observed = ~series.isna() - m_seq[:, c] = observed.astype(np.float32) - if observed.any(): - filled = series.interpolate(method='linear').bfill().ffill() - y_seq[:, c] = filled.fillna(0.0).values.astype(np.float32) - else: - y_seq[:, c] = 0.0 - - raw_X.append(x_seq) - raw_Y.append(y_seq) - raw_M.append(m_seq) - - if len(raw_X) == 0: - raise ValueError("未能从文件中构造出有效序列,请检查数据路径与传感器编码配置。") - - # 拼接所有文件数据进行 fit - X_all = np.vstack(raw_X) - - if fit_scaler: - self.scaler_X.fit(X_all) - self.scaler_Y.fit(raw_Y, raw_M) - - # 切分窗口 - for x_seq, y_seq, m_seq in zip(raw_X, raw_Y, raw_M): - x_seq_scaled = self.scaler_X.transform(x_seq) - y_seq_scaled = self.scaler_Y.transform(y_seq) - y_seq_scaled = np.where(m_seq > 0.5, y_seq_scaled, 0.0).astype(np.float32) - - for i in range(0, len(x_seq_scaled) - seq_len + 1, step_size): - x_win = x_seq_scaled[i:i+seq_len] - y_win = y_seq_scaled[i:i+seq_len] - m_win = m_seq[i:i+seq_len] - if np.sum(m_win) <= 0: - continue - self.X_data.append(x_win) - self.Y_data.append(y_win) - self.M_data.append(m_win) - - self.X_data = np.array(self.X_data) - self.Y_data = np.array(self.Y_data) - self.M_data = np.array(self.M_data) - - def __len__(self): - return len(self.X_data) - - def __getitem__(self, idx): - return ( - torch.tensor(self.X_data[idx], dtype=torch.float32), - torch.tensor(self.Y_data[idx], dtype=torch.float32), - torch.tensor(self.M_data[idx], dtype=torch.float32), - ) - -def get_dataloaders(condition='Non_TMD', include_free_vib=False): - """ - condition: 'Non_TMD' 或者是 'TMD' - include_free_vib: 是否在训练集中加入自由振动与自由衰减数据 - """ - base_dir = os.path.join(DATA_DIR, condition) - - # 手动区分的子目录 - train_files = glob.glob(os.path.join(base_dir, 'train', '*.csv')) - if include_free_vib: - train_files += glob.glob(os.path.join(base_dir, 'free_vib', '*.csv')) - - val_files = glob.glob(os.path.join(base_dir, 'val', '*.csv')) - test_files = glob.glob(os.path.join(base_dir, 'test', '*.csv')) - - print(f"[{condition}] Train files: {len(train_files)}, Val files: {len(val_files)}, Test files: {len(test_files)}") - - if len(train_files) == 0: - raise ValueError(f"错误: 在 {base_dir}/train 目录下未找到训练文件!请检查路径是否正确。") - if len(val_files) == 0: - raise ValueError(f"错误: 在 {base_dir}/val 目录下未找到验证文件!请检查路径是否正确。") - - train_dataset = BuildingDataset(train_files, SEQ_LEN, STEP_SIZE, fit_scaler=True) - val_dataset = BuildingDataset(val_files, SEQ_LEN, STEP_SIZE, - scaler_X=train_dataset.scaler_X, - scaler_Y=train_dataset.scaler_Y, fit_scaler=False) - # 若某条件(如 TMD)下没有 test 数据,可以处理一下防止报错 - test_loader = None - if len(test_files) > 0: - test_dataset = BuildingDataset(test_files, SEQ_LEN, STEP_SIZE, - scaler_X=train_dataset.scaler_X, - scaler_Y=train_dataset.scaler_Y, fit_scaler=False) - test_loader = DataLoader(test_dataset, batch_size=BATCH_SIZE, shuffle=False) - - train_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True) - val_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False) - - return train_loader, val_loader, test_loader, train_dataset.scaler_X, train_dataset.scaler_Y diff --git a/src_old/evaluate.py b/src_old/evaluate.py deleted file mode 100644 index af4d348..0000000 --- a/src_old/evaluate.py +++ /dev/null @@ -1,73 +0,0 @@ -import torch -import numpy as np -import matplotlib.pyplot as plt -from config import * -from dataset import get_dataloaders -from model import BuildingTCN - -def evaluate_model(): - _, _, test_loader, scaler_X, scaler_Y = get_dataloaders() - if test_loader is None: - raise ValueError("当前数据配置下没有可用的测试集。") - - model = BuildingTCN(input_size=1, output_size=5, num_channels=CHANNELS, - kernel_size=KERNEL_SIZE, dropout=DROPOUT).to(DEVICE) - model.load_state_dict(torch.load('best_model.pth', map_location=DEVICE)) - model.eval() - - all_preds = [] - all_targets = [] - - with torch.no_grad(): - for inputs, targets, masks in test_loader: - inputs = inputs.to(DEVICE) - outputs = model(inputs) - - all_preds.append(outputs.cpu().numpy()) - all_targets.append(targets.numpy()) - - # Concatenate results - all_preds = np.concatenate(all_preds, axis=0) - all_targets = np.concatenate(all_targets, axis=0) - - # Inverse transform to physical scale for metrics - all_preds_inv = scaler_Y.inverse_transform(all_preds.reshape(-1, len(OUTPUT_SENSORS))).reshape(all_preds.shape) - all_targets_inv = scaler_Y.inverse_transform(all_targets.reshape(-1, len(OUTPUT_SENSORS))).reshape(all_targets.shape) - - print("\n=== Test Metrics (All Windows) ===") - for i, sens in enumerate(OUTPUT_SENSORS): - pred_i = all_preds_inv[:, :, i].reshape(-1) - target_i = all_targets_inv[:, :, i].reshape(-1) - mae = float(np.mean(np.abs(pred_i - target_i))) - pred_std = float(np.std(pred_i)) - target_std = float(np.std(target_i)) - amp_ratio = pred_std / (target_std + 1e-12) - pred_peak = float(np.max(np.abs(pred_i))) - target_peak = float(np.max(np.abs(target_i))) - peak_ratio = pred_peak / (target_peak + 1e-12) - corr = float(np.corrcoef(pred_i, target_i)[0, 1]) if len(pred_i) > 1 else float("nan") - - print( - f"{sens}: MAE={mae:.4f}, Corr={corr:.4f}, " - f"AmpRatio(std_pred/std_true)={amp_ratio:.4f}, " - f"PeakRatio(max|pred|/max|true|)={peak_ratio:.4f}" - ) - - # 取一个 batch 的第一条序列进行可视化 - sample_pred = all_preds_inv[0] - sample_target = all_targets_inv[0] - - plt.figure(figsize=(15, 10)) - for i, sens in enumerate(OUTPUT_SENSORS): - plt.subplot(5, 1, i+1) - plt.plot(sample_target[:, i], label='True', alpha=0.7) - plt.plot(sample_pred[:, i], label='Pred', alpha=0.7, linestyle='--') - plt.title(f'Sensor {sens} Z-axis Response (Test: Earthquake)') - plt.legend() - - plt.tight_layout() - plt.savefig('test_results.png') - print("Evaluation done. Result saved to test_results.png") - -if __name__ == '__main__': - evaluate_model() diff --git a/src_old/model.py b/src_old/model.py deleted file mode 100644 index f45f44b..0000000 --- a/src_old/model.py +++ /dev/null @@ -1,76 +0,0 @@ -import torch -import torch.nn as nn -from torch.nn.utils import weight_norm - -class Chomp1d(nn.Module): - def __init__(self, chomp_size): - super(Chomp1d, self).__init__() - self.chomp_size = chomp_size - - def forward(self, x): - return x[:, :, :-self.chomp_size].contiguous() - -class TemporalBlock(nn.Module): - def __init__(self, n_inputs, n_outputs, kernel_size, stride, dilation, padding, dropout=0.2): - super(TemporalBlock, self).__init__() - self.conv1 = weight_norm(nn.Conv1d(n_inputs, n_outputs, kernel_size, - stride=stride, padding=padding, dilation=dilation)) - self.chomp1 = Chomp1d(padding) - self.relu1 = nn.ReLU() - self.dropout1 = nn.Dropout(dropout) - - self.conv2 = weight_norm(nn.Conv1d(n_outputs, n_outputs, kernel_size, - stride=stride, padding=padding, dilation=dilation)) - self.chomp2 = Chomp1d(padding) - self.relu2 = nn.ReLU() - self.dropout2 = nn.Dropout(dropout) - - self.net = nn.Sequential(self.conv1, self.chomp1, self.relu1, self.dropout1, - self.conv2, self.chomp2, self.relu2, self.dropout2) - self.downsample = nn.Conv1d(n_inputs, n_outputs, 1) if n_inputs != n_outputs else None - self.relu = nn.ReLU() - self.init_weights() - - def init_weights(self): - self.conv1.weight.data.normal_(0, 0.01) - self.conv2.weight.data.normal_(0, 0.01) - if self.downsample is not None: - self.downsample.weight.data.normal_(0, 0.01) - - def forward(self, x): - out = self.net(x) - res = x if self.downsample is None else self.downsample(x) - return self.relu(out + res) - -class TemporalConvNet(nn.Module): - def __init__(self, num_inputs, num_channels, kernel_size=2, dropout=0.2): - super(TemporalConvNet, self).__init__() - layers = [] - num_levels = len(num_channels) - for i in range(num_levels): - dilation_size = 2 ** i - in_channels = num_inputs if i == 0 else num_channels[i-1] - out_channels = num_channels[i] - layers += [TemporalBlock(in_channels, out_channels, kernel_size, stride=1, dilation=dilation_size, - padding=(kernel_size-1) * dilation_size, dropout=dropout)] - - self.network = nn.Sequential(*layers) - - def forward(self, x): - return self.network(x) - -class BuildingTCN(nn.Module): - def __init__(self, input_size, output_size, num_channels, kernel_size=3, dropout=0.2): - super(BuildingTCN, self).__init__() - self.tcn = TemporalConvNet(input_size, num_channels, kernel_size, dropout=dropout) - self.linear = nn.Linear(num_channels[-1], output_size) - - def forward(self, x): - # x shape: (batch, seq_len, input_size) - # TCN needs shape: (batch, input_size, seq_len) - x = x.transpose(1, 2) - y = self.tcn(x) - # y shape: (batch, num_channels, seq_len) - # linear needs shape: (batch, seq_len, num_channels) - y = y.transpose(1, 2) - return self.linear(y) diff --git a/src_old/train.py b/src_old/train.py deleted file mode 100644 index f201ff2..0000000 --- a/src_old/train.py +++ /dev/null @@ -1,241 +0,0 @@ -import torch -import torch.optim as optim -from config import * -from dataset import get_dataloaders -from model import BuildingTCN - -def masked_l1_loss(pred, target, mask): - diff = torch.abs(pred - target) * mask - denom = mask.sum().clamp(min=1.0) - return diff.sum() / denom - -def masked_mse_loss(pred, target, mask): - sq = ((pred - target) ** 2) * mask - denom = mask.sum().clamp(min=1.0) - return sq.sum() / denom - -def masked_spectral_mag_loss(pred, target, mask): - # Apply mask in time domain first so missing labels do not pollute spectrum. - pred_masked = pred * mask - target_masked = target * mask - - pred_fft = torch.fft.rfft(pred_masked, dim=1) - target_fft = torch.fft.rfft(target_masked, dim=1) - pred_mag = torch.abs(pred_fft) - target_mag = torch.abs(target_fft) - pred_mag = torch.nan_to_num(pred_mag, nan=0.0, posinf=1e6, neginf=0.0) - target_mag = torch.nan_to_num(target_mag, nan=0.0, posinf=1e6, neginf=0.0) - - # Weight each sample/channel by its valid-label ratio. - valid_ratio = (mask.sum(dim=1) / mask.shape[1]).clamp(min=0.0, max=1.0) - freq_weight = valid_ratio.unsqueeze(1).expand_as(pred_mag) - - diff = torch.abs(pred_mag - target_mag) * freq_weight - denom = freq_weight.sum().clamp(min=1.0) - return diff.sum() / denom - -def masked_corr_loss(pred, target, mask, eps=1e-8): - # Compute per-sample/per-channel Pearson correlation with masked timesteps. - count = mask.sum(dim=1) # (batch, channels) - valid = count > 1.0 - count_safe = count.clamp(min=1.0) - - pred_mean = (pred * mask).sum(dim=1) / count_safe - target_mean = (target * mask).sum(dim=1) / count_safe - - pred_centered = (pred - pred_mean.unsqueeze(1)) * mask - target_centered = (target - target_mean.unsqueeze(1)) * mask - - cov = (pred_centered * target_centered).sum(dim=1) - pred_var = (pred_centered ** 2).sum(dim=1) - target_var = (target_centered ** 2).sum(dim=1) - - # Put eps inside sqrt to avoid infinite gradients around zero variance. - denom = torch.sqrt((pred_var * target_var).clamp(min=0.0) + eps) - corr = cov / denom - corr = torch.nan_to_num(corr, nan=0.0, posinf=0.0, neginf=0.0) - corr = torch.where(valid, corr, torch.zeros_like(corr)) - - valid_float = valid.float() - denom = valid_float.sum().clamp(min=1.0) - return ((1.0 - corr) * valid_float).sum() / denom - -def masked_std_loss(pred, target, mask): - count = mask.sum(dim=1).clamp(min=1.0) - pred_mean = (pred * mask).sum(dim=1) / count - target_mean = (target * mask).sum(dim=1) / count - - pred_centered = (pred - pred_mean.unsqueeze(1)) * mask - target_centered = (target - target_mean.unsqueeze(1)) * mask - - pred_std = torch.sqrt((pred_centered ** 2).sum(dim=1) / count + 1e-8) - target_std = torch.sqrt((target_centered ** 2).sum(dim=1) / count + 1e-8) - - valid = mask.sum(dim=1) > 1.0 - valid_float = valid.float() - denom = valid_float.sum().clamp(min=1.0) - return (torch.abs(pred_std - target_std) * valid_float).sum() / denom - -def masked_peak_loss(pred, target, mask): - abs_pred = torch.abs(pred) * mask - abs_target = torch.abs(target) * mask - pred_peak = abs_pred.max(dim=1).values - target_peak = abs_target.max(dim=1).values - - valid = mask.sum(dim=1) > 0.0 - valid_float = valid.float() - denom = valid_float.sum().clamp(min=1.0) - return (torch.abs(pred_peak - target_peak) * valid_float).sum() / denom - -def train_model(): - train_loader, val_loader, _, _, _ = get_dataloaders() - - # 初始化前向模型 (1 -> 5) - model = BuildingTCN(input_size=1, output_size=5, num_channels=CHANNELS, - kernel_size=KERNEL_SIZE, dropout=DROPOUT).to(DEVICE) - - optimizer = optim.AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY) - scheduler = optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=8, factor=0.5) - - best_val_loss = float('inf') - no_improve_epochs = 0 - early_stop_patience = EARLY_STOP_PATIENCE - - for epoch in range(EPOCHS): - model.train() - train_loss = 0.0 - train_time_loss = 0.0 - train_spec_loss = 0.0 - train_corr_loss = 0.0 - train_mse_loss = 0.0 - train_std_loss = 0.0 - train_peak_loss = 0.0 - train_batches = 0 - - for inputs, targets, masks in train_loader: - inputs = inputs.to(DEVICE) - targets = targets.to(DEVICE) - masks = masks.to(DEVICE) - - optimizer.zero_grad() - outputs = model(inputs) - time_loss = masked_l1_loss(outputs, targets, masks) - mse_loss = masked_mse_loss(outputs, targets, masks) - spec_loss = masked_spectral_mag_loss(outputs, targets, masks) - corr_loss = masked_corr_loss(outputs, targets, masks, eps=CORR_LOSS_EPS) - std_loss = masked_std_loss(outputs, targets, masks) - peak_loss = masked_peak_loss(outputs, targets, masks) - loss = ( - time_loss - + MSE_LOSS_WEIGHT * mse_loss - + SPECTRAL_LOSS_WEIGHT * spec_loss - + CORR_LOSS_WEIGHT * corr_loss - + STD_LOSS_WEIGHT * std_loss - + PEAK_LOSS_WEIGHT * peak_loss - ) - - if not torch.isfinite(loss): - print(" [Warn] Non-finite train loss encountered. Skip this batch.") - continue - - loss.backward() - grad_norm = torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) - if not torch.isfinite(grad_norm): - print(" [Warn] Non-finite gradient norm encountered. Skip optimizer step for this batch.") - optimizer.zero_grad(set_to_none=True) - continue - - optimizer.step() - - train_loss += loss.item() - train_time_loss += time_loss.item() - train_mse_loss += mse_loss.item() - train_spec_loss += spec_loss.item() - train_corr_loss += corr_loss.item() - train_std_loss += std_loss.item() - train_peak_loss += peak_loss.item() - train_batches += 1 - - # 验证 - model.eval() - val_loss = 0.0 - val_time_loss = 0.0 - val_spec_loss = 0.0 - val_corr_loss = 0.0 - val_mse_loss = 0.0 - val_std_loss = 0.0 - val_peak_loss = 0.0 - val_batches = 0 - with torch.no_grad(): - for inputs, targets, masks in val_loader: - inputs = inputs.to(DEVICE) - targets = targets.to(DEVICE) - masks = masks.to(DEVICE) - outputs = model(inputs) - time_loss = masked_l1_loss(outputs, targets, masks) - mse_loss = masked_mse_loss(outputs, targets, masks) - spec_loss = masked_spectral_mag_loss(outputs, targets, masks) - corr_loss = masked_corr_loss(outputs, targets, masks, eps=CORR_LOSS_EPS) - std_loss = masked_std_loss(outputs, targets, masks) - peak_loss = masked_peak_loss(outputs, targets, masks) - loss = ( - time_loss - + MSE_LOSS_WEIGHT * mse_loss - + SPECTRAL_LOSS_WEIGHT * spec_loss - + CORR_LOSS_WEIGHT * corr_loss - + STD_LOSS_WEIGHT * std_loss - + PEAK_LOSS_WEIGHT * peak_loss - ) - if not torch.isfinite(loss): - continue - val_loss += loss.item() - val_time_loss += time_loss.item() - val_mse_loss += mse_loss.item() - val_spec_loss += spec_loss.item() - val_corr_loss += corr_loss.item() - val_std_loss += std_loss.item() - val_peak_loss += peak_loss.item() - val_batches += 1 - - train_den = max(train_batches, 1) - val_den = max(val_batches, 1) - train_loss /= train_den - train_time_loss /= train_den - train_mse_loss /= train_den - train_spec_loss /= train_den - train_corr_loss /= train_den - train_std_loss /= train_den - train_peak_loss /= train_den - val_loss /= val_den - val_time_loss /= val_den - val_mse_loss /= val_den - val_spec_loss /= val_den - val_corr_loss /= val_den - val_std_loss /= val_den - val_peak_loss /= val_den - - scheduler.step(val_loss) - - print( - f"Epoch {epoch+1}/{EPOCHS} | " - f"Train Loss: {train_loss:.4f} " - f"(L1={train_time_loss:.4f}, MSE={train_mse_loss:.4f}, Spec={train_spec_loss:.4f}, " - f"Corr={train_corr_loss:.4f}, Std={train_std_loss:.4f}, Peak={train_peak_loss:.4f}) | " - f"Val Loss: {val_loss:.4f} " - f"(L1={val_time_loss:.4f}, MSE={val_mse_loss:.4f}, Spec={val_spec_loss:.4f}, " - f"Corr={val_corr_loss:.4f}, Std={val_std_loss:.4f}, Peak={val_peak_loss:.4f})" - ) - - if val_loss < best_val_loss: - best_val_loss = val_loss - torch.save(model.state_dict(), 'best_model.pth') - print(" --> Saved Best Model") - no_improve_epochs = 0 - else: - no_improve_epochs += 1 - if ENABLE_EARLY_STOP and no_improve_epochs >= early_stop_patience: - print(f"Early stopping at epoch {epoch+1}") - break - -if __name__ == '__main__': - train_model()