diff --git a/Figure_1.png b/Figure_1.png index 18fd9a5..3f550d9 100644 Binary files a/Figure_1.png and b/Figure_1.png differ diff --git a/checkpoints/forward/best_tcn_model.pt b/checkpoints/forward/best_tcn_model.pt index 5b2f19a..c6235b3 100644 Binary files a/checkpoints/forward/best_tcn_model.pt and b/checkpoints/forward/best_tcn_model.pt differ diff --git a/checkpoints/forward/training_history.csv b/checkpoints/forward/training_history.csv index 2ae6565..c12ced1 100644 --- a/checkpoints/forward/training_history.csv +++ b/checkpoints/forward/training_history.csv @@ -1,14 +1,11 @@ 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,32.23490004269582,0.9063257234838774,0.09063257336757093,2462.537229142099,0.5202015545570625,0.5072746667659508,0.23183272828189833,0.9878060196368199,0.5202015545570625,0.6667674588373308,16.848691019518622,0.5030307245665583,0.05030307315033058,712.1304384428879,0.21638367042459292,0.3054182426682834,0.16712227438030572,0.7986705179872184,0.21638367042459292,0.9571996228951049 -2,0.001,21.43670712776904,0.9317702313639084,0.09317702508338217,1694.0047123747052,0.35049294160222105,0.3616391454102858,0.1121402242625097,0.672775794312639,0.35049294160222105,0.6002240721274537,18.15417513354071,0.5811605972462687,0.058116060392609956,948.0657793242356,0.23726162314414978,0.3922087324076685,0.15394838554142365,0.7320991713425209,0.23726162314414978,0.8361816405868994 -3,0.001,18.081004331696708,0.949055330933265,0.09490553458344261,1372.9708672289578,0.28145609332143134,0.3131425976753235,0.09374795892750318,0.6080609613432074,0.28145609332143134,0.616063318727439,16.73266759412042,0.5420598767954727,0.05420598895128431,673.0913622625943,0.238979727286717,0.33492567107595245,0.1804515516449665,0.7293170896069757,0.238979727286717,0.7597824622933206 -4,0.001,15.158426779621053,0.9718681652590914,0.09718681784030402,1005.0683164776497,0.2189214069325969,0.2836238658934269,0.08367646379695046,0.5845890998278024,0.2189214069325969,0.6179393338571753,17.278400026518725,0.5685602380283947,0.056856024560743366,892.716255977236,0.22618299260221678,0.36168746264844104,0.1538948353444194,0.6990852653980255,0.22618299260221678,0.7597824622933206 -5,0.001,13.61913053044733,0.9698394334541177,0.09698394513776842,859.7200766509434,0.19106561557020782,0.2557676170232161,0.07612784807833861,0.5538543634257227,0.19106561557020782,0.6155450957215706,17.31632186626566,0.5785790409507423,0.057857905097048856,895.00206572434,0.21634554451909557,0.35666010605877846,0.1514351657302729,0.7220557270378902,0.21634554451909557,0.7402091190436203 -6,0.001,12.678725962368947,0.9694307513956754,0.0969430767351164,788.4346546676924,0.1713324544845887,0.23463592369039105,0.07080332469195127,0.5323761788741598,0.1713324544845887,0.6164817462027068,18.423015594482422,0.5721222347226637,0.057212224808232535,888.629751271215,0.25913861959144985,0.4007708651238474,0.16395779438959113,0.7525067822686557,0.25913861959144985,0.7218985721265917 -7,0.001,12.268796533908484,0.967802310889622,0.09678023287429,744.56112426182,0.1675024907684551,0.2299804154713199,0.06777841198029665,0.5231116253812358,0.1675024907684551,0.5961900787542253,17.48041952067408,0.5510032737049563,0.055100328578003524,704.0817605380354,0.2379121222886546,0.37086750232967836,0.19612120528673305,0.7370918125941835,0.2379121222886546,0.7088496766840734 -8,0.001,11.889431161700555,0.967772865070487,0.09677728824317455,712.115890215028,0.1614244244289848,0.2226250537161557,0.06655820054968573,0.5118434342010966,0.1614244244289848,0.6022199021486582,18.410391971982758,0.5461424337378864,0.05461424415738418,738.1983037488214,0.2533088120920905,0.3795644687167529,0.21159803308546543,0.7770289305982918,0.2533088120920905,0.707797346401937 -9,0.001,11.234980547203207,0.9654024828155086,0.09654024967326308,670.3140682004532,0.14827435323089924,0.20859187953876998,0.062428659900038874,0.4916352005499714,0.14827435323089924,0.5623616312479263,16.51186601046858,0.5469307437025267,0.05469307598882708,658.5074381335028,0.23887184570575581,0.3477302913008065,0.18191729235494958,0.6951427665250055,0.23887184570575581,0.7065345500594454 -10,0.001,10.992490152143082,0.9626774315564137,0.09626774419591112,649.039663350807,0.14625250653557056,0.20260991113928128,0.061842711396374796,0.48383018318212256,0.14625250653557056,0.568747647623218,17.501161114922887,0.5506949491541961,0.05506949608439002,671.2075416301859,0.2541997463538729,0.373840541675173,0.19259302314884705,0.7480303932880533,0.2541997463538729,0.708428744563363 -11,0.0005,10.806822551871246,0.9680146614335617,0.09680146787245318,637.6420590382702,0.14451587228280194,0.2035975716305229,0.0572100906404403,0.4776927213061531,0.14451587228280194,0.5574953460125556,17.13407661174906,0.5288547395632185,0.05288547495829648,618.1973874322299,0.2417878899081,0.34800343832065317,0.21005815672206468,0.7388678706925491,0.2417878899081,0.7122171335564683 -12,0.0005,10.255595112746617,0.9616962041494981,0.09616962194723903,599.1373590433373,0.12593858316540718,0.18720266560338578,0.05598305191247249,0.4703599948365733,0.12593858316540718,0.5670544286673502,16.41022307297279,0.5580467443014013,0.055804675894564594,691.4893851444639,0.23402822017669678,0.3565609496215294,0.17255781141334567,0.6764166149599798,0.23402822017669678,0.6966426454535165 -13,0.0005,9.979247138185322,0.9624017550135558,0.09624017732885648,580.7163422782467,0.12185739905063836,0.18553815514973873,0.05258095577657926,0.4598706980358879,0.12185739905063836,0.5524671217591512,16.273903156148975,0.5295915886245924,0.05295915985158805,541.1639246447332,0.23793373128463483,0.3361044392503541,0.20454522344315873,0.7062619534032099,0.23793373128463483,0.7056926858376642 +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 new file mode 100644 index 0000000..4e72542 Binary files /dev/null and b/checkpoints_mlp/task1_feature_mlp/best_feature_mlp.pt differ diff --git a/checkpoints_mlp/task1_feature_mlp/training_history.csv b/checkpoints_mlp/task1_feature_mlp/training_history.csv new file mode 100644 index 0000000..2ad8868 --- /dev/null +++ b/checkpoints_mlp/task1_feature_mlp/training_history.csv @@ -0,0 +1,59 @@ +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 new file mode 100644 index 0000000..2230b5c Binary files /dev/null and b/checkpoints_rms/forward_rms/best_rms_model.pt differ diff --git a/checkpoints_rms/forward_rms/training_history.csv b/checkpoints_rms/forward_rms/training_history.csv new file mode 100644 index 0000000..bf44b5d --- /dev/null +++ b/checkpoints_rms/forward_rms/training_history.csv @@ -0,0 +1,30 @@ +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/evaluation_outputs/forward/evaluation_forward_test_b0_s0.png b/evaluation_outputs/forward/evaluation_forward_test_b0_s0.png index 6b1a762..7566f71 100644 Binary files a/evaluation_outputs/forward/evaluation_forward_test_b0_s0.png and b/evaluation_outputs/forward/evaluation_forward_test_b0_s0.png differ diff --git a/evaluation_outputs/forward/evaluation_forward_val_b0_s0.png b/evaluation_outputs/forward/evaluation_forward_val_b0_s0.png new file mode 100644 index 0000000..c1904ad Binary files /dev/null and b/evaluation_outputs/forward/evaluation_forward_val_b0_s0.png differ diff --git a/evaluation_outputs/forward/evaluation_forward_val_b3_s0.png b/evaluation_outputs/forward/evaluation_forward_val_b3_s0.png new file mode 100644 index 0000000..7c88303 Binary files /dev/null and b/evaluation_outputs/forward/evaluation_forward_val_b3_s0.png differ diff --git a/evaluation_outputs/forward_rms/evaluation_test_all_samples.csv b/evaluation_outputs/forward_rms/evaluation_test_all_samples.csv new file mode 100644 index 0000000..45353cb --- /dev/null +++ b/evaluation_outputs/forward_rms/evaluation_test_all_samples.csv @@ -0,0 +1,3 @@ +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 new file mode 100644 index 0000000..e3e1d75 Binary files /dev/null and b/evaluation_outputs/forward_rms/evaluation_test_s0.png differ diff --git a/evaluation_outputs/forward_rms/evaluation_train_all_samples.csv b/evaluation_outputs/forward_rms/evaluation_train_all_samples.csv new file mode 100644 index 0000000..b6d287a --- /dev/null +++ b/evaluation_outputs/forward_rms/evaluation_train_all_samples.csv @@ -0,0 +1,37 @@ +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 new file mode 100644 index 0000000..f6408ba --- /dev/null +++ b/evaluation_outputs/forward_rms/evaluation_val_all_samples.csv @@ -0,0 +1,6 @@ +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 new file mode 100644 index 0000000..8d84a2a Binary files /dev/null and b/evaluation_outputs/forward_rms/evaluation_val_s0.png differ diff --git a/evaluation_outputs/forward_rms/evaluation_val_s0_waveform.png b/evaluation_outputs/forward_rms/evaluation_val_s0_waveform.png new file mode 100644 index 0000000..6b4df1b Binary files /dev/null and b/evaluation_outputs/forward_rms/evaluation_val_s0_waveform.png 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 new file mode 100644 index 0000000..ffe85d0 --- /dev/null +++ b/evaluation_outputs/task1_feature_mlp/evaluation_train_all_samples.csv @@ -0,0 +1,37 @@ +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 new file mode 100644 index 0000000..ed7cedb Binary files /dev/null and b/evaluation_outputs/task1_feature_mlp/evaluation_train_curve.png 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 new file mode 100644 index 0000000..fbb7f40 --- /dev/null +++ b/evaluation_outputs/task1_feature_mlp/evaluation_val_all_samples.csv @@ -0,0 +1,6 @@ +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 new file mode 100644 index 0000000..e884726 Binary files /dev/null and b/evaluation_outputs/task1_feature_mlp/evaluation_val_curve.png differ diff --git a/evaluation_outputs/task1_feature_mlp/evaluation_val_s0.png b/evaluation_outputs/task1_feature_mlp/evaluation_val_s0.png new file mode 100644 index 0000000..c96e1bc Binary files /dev/null and b/evaluation_outputs/task1_feature_mlp/evaluation_val_s0.png 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 new file mode 100644 index 0000000..2433d5f Binary files /dev/null and b/evaluation_outputs/task1_feature_mlp/harmonic_5mm_0.75Hz_prediction.png 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 new file mode 100644 index 0000000..ba8fe90 Binary files /dev/null and b/evaluation_outputs/task1_feature_mlp/harmonic_5mm_1.55Hz_prediction.png differ diff --git a/scripts/__pycache__/config.cpython-310.pyc b/scripts/__pycache__/config.cpython-310.pyc new file mode 100644 index 0000000..501a546 Binary files /dev/null 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 new file mode 100644 index 0000000..f65b124 Binary files /dev/null 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 new file mode 100644 index 0000000..7690e08 Binary files /dev/null and b/scripts/__pycache__/model.cpython-310.pyc differ diff --git a/scripts/config.py b/scripts/config.py new file mode 100644 index 0000000..804a77f --- /dev/null +++ b/scripts/config.py @@ -0,0 +1,109 @@ +from __future__ import annotations + +from dataclasses import dataclass, field +from pathlib import Path + + +CORE_FEATURE_NAMES: tuple[str, ...] = ( + "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", +) + + +@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" + + middle_segment_start_ratio: float = 0.20 + middle_segment_end_ratio: float = 0.80 + min_segment_length: int = 512 + steady_window_ratio: float = 0.25 + steady_window_stride_ratio: float = 0.05 + stability_subwindow_count: int = 4 + 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_dim: int = len(CORE_FEATURE_NAMES) + hidden_dims: tuple[int, ...] = (96, 64, 32) + dropout: float = 0.08 + + +@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 + + +@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_experiment_config() -> ExperimentConfig: + return ExperimentConfig() diff --git a/scripts/dataset.py b/scripts/dataset.py new file mode 100644 index 0000000..e0a85c3 --- /dev/null +++ b/scripts/dataset.py @@ -0,0 +1,478 @@ +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 +except ImportError: + from config import CORE_FEATURE_NAMES, 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)) + + +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 get_middle_segment( + time_values: np.ndarray, + x_values: np.ndarray, + y_values: np.ndarray, + config: DataConfig, +) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + length = len(x_values) + search_start = int(length * config.middle_segment_start_ratio) + search_end = int(length * config.middle_segment_end_ratio) + search_start = max(0, min(search_start, length - 1)) + search_end = max(search_start + 1, min(search_end, length)) + + if search_end - search_start < config.min_segment_length: + center = length // 2 + half = config.min_segment_length // 2 + search_start = max(0, center - half) + search_end = min(length, search_start + config.min_segment_length) + search_start = max(0, search_end - config.min_segment_length) + + search_length = search_end - search_start + window_length = max(config.min_segment_length, int(length * config.steady_window_ratio)) + window_length = min(window_length, search_length) + if search_length <= window_length: + return ( + time_values[search_start:search_end], + x_values[search_start:search_end], + y_values[search_start:search_end], + ) + + stride = max(1, int(length * config.steady_window_stride_ratio)) + candidate_y = y_values[search_start:search_end] + candidate_rms = calculate_rms(candidate_y) + best_score = float("inf") + best_slice = slice(search_start, search_start + window_length) + + for window_start in range(search_start, search_end - window_length + 1, stride): + window_end = window_start + window_length + window_y = y_values[window_start:window_end] + split_windows = np.array_split(window_y, config.stability_subwindow_count) + split_rms = np.asarray([calculate_rms(chunk) for chunk in split_windows], dtype=np.float64) + split_peaks = np.asarray([float(np.max(np.abs(chunk))) for chunk in split_windows], dtype=np.float64) + + rms_cv = float(split_rms.std() / max(split_rms.mean(), config.normalization_eps)) + 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 + best_slice = slice(window_start, window_end) + + return time_values[best_slice], x_values[best_slice], y_values[best_slice] + + +def _compute_windowed_spectrum(signal: np.ndarray, sampling_rate: float) -> tuple[np.ndarray, np.ndarray]: + signal = np.asarray(signal, dtype=np.float64).reshape(-1) + if signal.size < 4: + return np.asarray([], dtype=np.float64), np.asarray([], dtype=np.float64) + centered = signal - np.mean(signal) + window = np.hanning(signal.size) + scale = max(np.sum(window), 1e-12) + fft_values = np.fft.rfft(centered * window) + freqs = np.fft.rfftfreq(signal.size, d=1.0 / sampling_rate) + magnitudes = (2.0 / scale) * np.abs(fft_values) + return freqs, magnitudes + + +def _parabolic_peak_frequency(freqs: np.ndarray, magnitudes: np.ndarray, peak_index: int) -> float: + if peak_index <= 0 or peak_index >= magnitudes.size - 1: + return float(freqs[peak_index]) + alpha = magnitudes[peak_index - 1] + beta = magnitudes[peak_index] + gamma = magnitudes[peak_index + 1] + denominator = alpha - 2.0 * beta + gamma + if abs(denominator) < 1e-12: + return float(freqs[peak_index]) + offset = 0.5 * (alpha - gamma) / denominator + bin_width = float(freqs[1] - freqs[0]) + return float(freqs[peak_index] + offset * bin_width) + + +def get_dominant_frequency(signal: np.ndarray, sampling_rate: float) -> float: + freqs, magnitudes = _compute_windowed_spectrum(signal, sampling_rate) + if magnitudes.size == 0: + return 0.0 + magnitudes[0] = 0.0 + band_mask = (freqs >= 0.1) & (freqs <= 5.0) + if not np.any(band_mask): + return 0.0 + band_magnitudes = np.where(band_mask, magnitudes, 0.0) + peak_index = int(np.argmax(band_magnitudes)) + 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 extract_core_features( + signal: np.ndarray, + sampling_rate: float, + original_length: int, + config: DataConfig, +) -> np.ndarray: + 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, + 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, + ], + dtype=np.float32, + ) + + +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 load_aligned_signals(file_path: Path, config: DataConfig) -> tuple[np.ndarray, np.ndarray, np.ndarray, int]: + 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: + raise ValueError(f"Missing required sensor in {file_path.name}") + + aligned = base_df.rename(columns={config.base_axis: "base_signal"}).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(): + raise ValueError(f"Remaining NaN after interpolation in {file_path.name}") + + 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) + return time_values, x_values, y_values, interpolation_count + + +@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 + sampling_rate: 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 + + +@dataclass +class NormalizationStats: + feature_mean: torch.Tensor + feature_std: torch.Tensor + target_mean: torch.Tensor + target_std: torch.Tensor + + +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: + 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) + frequency_hz = extract_frequency_hz(file_path.name) + if frequency_hz is None: + frequency_hz = get_dominant_frequency(x_middle, sampling_rate) + + 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) + + 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, + interpolation_count=interpolation_count, + ) + + +def load_split_records(config: DataConfig, split: str) -> tuple[list[SampleRecord], SplitLoadReport]: + records: list[SampleRecord] = [] + report = SplitLoadReport(split=split) + for file_path in list_split_files(config, split): + try: + record = build_record_from_file(file_path, split, config) + except ValueError as error: + report.skipped_files.append((file_path.name, str(error))) + continue + records.append(record) + report.loaded_files.append(file_path.name) + if record.interpolation_count > 0: + report.interpolated_files[file_path.name] = record.interpolation_count + + if not records: + raise RuntimeError(f"No usable records found for split='{split}'.") + 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 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), + ) diff --git a/scripts/evaluate.py b/scripts/evaluate.py new file mode 100644 index 0000000..2ecb706 --- /dev/null +++ b/scripts/evaluate.py @@ -0,0 +1,224 @@ +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 CORE_FEATURE_NAMES, ExperimentConfig, make_experiment_config + from .dataset import build_dataloaders, normalization_from_payload, report_to_text + from .model import build_model +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 + + +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.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: + 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 + + +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 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 + 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}") + + 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() + 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}") + + checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False) + if "normalization" in checkpoint: + normalization = normalization_from_payload(checkpoint["normalization"]) + else: + normalization = None + + loaders, datasets, raw_records, reports = build_dataloaders(config) + del loaders + print(report_to_text(reports)) + + dataset = datasets[args.split] + if normalization is not None: + dataset.normalization = normalization + + 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}") + print(f"Figure saved to: {figure_path}") + + +if __name__ == "__main__": + main() diff --git a/scripts/model.py b/scripts/model.py new file mode 100644 index 0000000..2b02602 --- /dev/null +++ b/scripts/model.py @@ -0,0 +1,31 @@ +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 new file mode 100644 index 0000000..3446080 --- /dev/null +++ b/scripts/predict_single.py @@ -0,0 +1,143 @@ +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 CORE_FEATURE_NAMES, make_experiment_config + from .dataset import build_record_from_file, normalization_from_payload + from .model import build_model +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 + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description="Predict task1 RMS from a single 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: + if checkpoint_arg: + return Path(checkpoint_arg).resolve() + return project_root / checkpoint_dir / "task1_feature_mlp" / best_model_name + + +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) + 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 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() + 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() + + 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) + + 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]) + + 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"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 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}") + + +if __name__ == "__main__": + main() diff --git a/scripts/train.py b/scripts/train.py new file mode 100644 index 0000000..ecd1960 --- /dev/null +++ b/scripts/train.py @@ -0,0 +1,274 @@ +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/src/__pycache__/config.cpython-310.pyc b/src/__pycache__/config.cpython-310.pyc index c1b0c40..21e66cb 100644 Binary files a/src/__pycache__/config.cpython-310.pyc and b/src/__pycache__/config.cpython-310.pyc differ diff --git a/src/__pycache__/dataset.cpython-310.pyc b/src/__pycache__/dataset.cpython-310.pyc index 4e576b7..292fa11 100644 Binary files a/src/__pycache__/dataset.cpython-310.pyc and b/src/__pycache__/dataset.cpython-310.pyc differ diff --git a/src/__pycache__/model.cpython-310.pyc b/src/__pycache__/model.cpython-310.pyc index 5fd9585..0344524 100644 Binary files a/src/__pycache__/model.cpython-310.pyc and b/src/__pycache__/model.cpython-310.pyc differ diff --git a/src/config.py b/src/config.py index b9e9896..20aa77f 100644 --- a/src/config.py +++ b/src/config.py @@ -36,6 +36,11 @@ class DataConfig: 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" @@ -62,9 +67,11 @@ class ModelConfig: output_channels: int = 1 tcn_channels: tuple[int, ...] = (32, 32, 64, 64, 128) kernel_size: int = 5 - dropout: float = 0.25 + 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 @@ -78,7 +85,7 @@ class LossConfig: 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 = 8.0 + forward_envelope_loss_weight: float = 2.0 envelope_kernel_size: int = 33 diff --git a/src/dataset.py b/src/dataset.py index f738663..88857a0 100644 --- a/src/dataset.py +++ b/src/dataset.py @@ -7,7 +7,7 @@ from typing import Any import numpy as np import pandas as pd import torch -from torch.utils.data import DataLoader, Dataset +from torch.utils.data import DataLoader, Dataset, WeightedRandomSampler try: from .config import DataConfig, ExperimentConfig @@ -109,8 +109,11 @@ class WindowedTimeSeriesDataset(Dataset): 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]) @@ -135,6 +138,13 @@ class WindowedTimeSeriesDataset(Dataset): 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( @@ -142,6 +152,15 @@ class WindowedTimeSeriesDataset(Dataset): "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) @@ -314,6 +333,31 @@ def fit_normalization_stats(records: list[SequenceRecord], config: DataConfig) - ) +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]]: @@ -339,12 +383,24 @@ def build_dataloaders( ) -> 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=data_config.shuffle_train, + 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, diff --git a/src/model.py b/src/model.py index 1a3f662..715ba61 100644 --- a/src/model.py +++ b/src/model.py @@ -111,11 +111,22 @@ class TemporalConvNet(nn.Module): 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) - return self.output_head(features) + 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): diff --git a/src_new/__pycache__/config.cpython-310.pyc b/src_new/__pycache__/config.cpython-310.pyc new file mode 100644 index 0000000..a47c917 Binary files /dev/null and b/src_new/__pycache__/config.cpython-310.pyc differ diff --git a/src_new/__pycache__/dataset.cpython-310.pyc b/src_new/__pycache__/dataset.cpython-310.pyc new file mode 100644 index 0000000..f2fccec Binary files /dev/null and b/src_new/__pycache__/dataset.cpython-310.pyc differ diff --git a/src_new/__pycache__/evaluate.cpython-310.pyc b/src_new/__pycache__/evaluate.cpython-310.pyc new file mode 100644 index 0000000..b17f0d2 Binary files /dev/null and b/src_new/__pycache__/evaluate.cpython-310.pyc differ diff --git a/src_new/__pycache__/model.cpython-310.pyc b/src_new/__pycache__/model.cpython-310.pyc new file mode 100644 index 0000000..efd447d Binary files /dev/null and b/src_new/__pycache__/model.cpython-310.pyc differ diff --git a/src_new/__pycache__/train.cpython-310.pyc b/src_new/__pycache__/train.cpython-310.pyc new file mode 100644 index 0000000..681b0a7 Binary files /dev/null and b/src_new/__pycache__/train.cpython-310.pyc differ diff --git a/src_new/config.py b/src_new/config.py new file mode 100644 index 0000000..467145c --- /dev/null +++ b/src_new/config.py @@ -0,0 +1,103 @@ +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 new file mode 100644 index 0000000..5295819 --- /dev/null +++ b/src_new/dataset.py @@ -0,0 +1,388 @@ +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 new file mode 100644 index 0000000..be1c595 --- /dev/null +++ b/src_new/evaluate.py @@ -0,0 +1,222 @@ +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 new file mode 100644 index 0000000..31f1f25 --- /dev/null +++ b/src_new/model.py @@ -0,0 +1,140 @@ +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 new file mode 100644 index 0000000..179ea1f --- /dev/null +++ b/src_new/train.py @@ -0,0 +1,327 @@ +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()