modified: Figure_1.png
modified: checkpoints/forward/best_tcn_model.pt modified: checkpoints/forward/training_history.csv new file: checkpoints_mlp/task1_feature_mlp/best_feature_mlp.pt new file: checkpoints_mlp/task1_feature_mlp/training_history.csv new file: checkpoints_rms/forward_rms/best_rms_model.pt new file: checkpoints_rms/forward_rms/training_history.csv modified: evaluation_outputs/forward/evaluation_forward_test_b0_s0.png new file: evaluation_outputs/forward/evaluation_forward_val_b0_s0.png new file: evaluation_outputs/forward/evaluation_forward_val_b3_s0.png new file: evaluation_outputs/forward_rms/evaluation_test_all_samples.csv new file: evaluation_outputs/forward_rms/evaluation_test_s0.png new file: evaluation_outputs/forward_rms/evaluation_train_all_samples.csv new file: evaluation_outputs/forward_rms/evaluation_val_all_samples.csv new file: evaluation_outputs/forward_rms/evaluation_val_s0.png new file: evaluation_outputs/forward_rms/evaluation_val_s0_waveform.png new file: evaluation_outputs/task1_feature_mlp/evaluation_train_all_samples.csv new file: evaluation_outputs/task1_feature_mlp/evaluation_train_curve.png new file: evaluation_outputs/task1_feature_mlp/evaluation_val_all_samples.csv new file: evaluation_outputs/task1_feature_mlp/evaluation_val_curve.png new file: evaluation_outputs/task1_feature_mlp/evaluation_val_s0.png new file: evaluation_outputs/task1_feature_mlp/harmonic_5mm_0.75Hz_prediction.png new file: evaluation_outputs/task1_feature_mlp/harmonic_5mm_1.55Hz_prediction.png new file: scripts/__pycache__/config.cpython-310.pyc new file: scripts/__pycache__/dataset.cpython-310.pyc new file: scripts/__pycache__/model.cpython-310.pyc new file: scripts/config.py new file: scripts/dataset.py new file: scripts/evaluate.py new file: scripts/model.py new file: scripts/predict_single.py new file: scripts/train.py modified: src/__pycache__/config.cpython-310.pyc modified: src/__pycache__/dataset.cpython-310.pyc modified: src/__pycache__/model.cpython-310.pyc modified: src/config.py modified: src/dataset.py modified: src/model.py new file: src_new/__pycache__/config.cpython-310.pyc new file: src_new/__pycache__/dataset.cpython-310.pyc new file: src_new/__pycache__/evaluate.cpython-310.pyc new file: src_new/__pycache__/model.cpython-310.pyc new file: src_new/__pycache__/train.cpython-310.pyc new file: src_new/config.py new file: src_new/dataset.py new file: src_new/evaluate.py new file: src_new/model.py new file: src_new/train.py
BIN
Figure_1.png
|
Before Width: | Height: | Size: 190 KiB After Width: | Height: | Size: 167 KiB |
@@ -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
|
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
|
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,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
|
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,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
|
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,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
|
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.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
|
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,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
|
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,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
|
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.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
|
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.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
|
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.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
|
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
|
||||||
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
|
|
||||||
|
|||||||
|
BIN
checkpoints_mlp/task1_feature_mlp/best_feature_mlp.pt
Normal file
59
checkpoints_mlp/task1_feature_mlp/training_history.csv
Normal file
@@ -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
|
||||||
|
BIN
checkpoints_rms/forward_rms/best_rms_model.pt
Normal file
30
checkpoints_rms/forward_rms/training_history.csv
Normal file
@@ -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
|
||||||
|
|
Before Width: | Height: | Size: 362 KiB After Width: | Height: | Size: 372 KiB |
BIN
evaluation_outputs/forward/evaluation_forward_val_b0_s0.png
Normal file
|
After Width: | Height: | Size: 338 KiB |
BIN
evaluation_outputs/forward/evaluation_forward_val_b3_s0.png
Normal file
|
After Width: | Height: | Size: 323 KiB |
@@ -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
|
||||||
|
BIN
evaluation_outputs/forward_rms/evaluation_test_s0.png
Normal file
|
After Width: | Height: | Size: 172 KiB |
@@ -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
|
||||||
|
@@ -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
|
||||||
|
BIN
evaluation_outputs/forward_rms/evaluation_val_s0.png
Normal file
|
After Width: | Height: | Size: 166 KiB |
BIN
evaluation_outputs/forward_rms/evaluation_val_s0_waveform.png
Normal file
|
After Width: | Height: | Size: 280 KiB |
@@ -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
|
||||||
|
BIN
evaluation_outputs/task1_feature_mlp/evaluation_train_curve.png
Normal file
|
After Width: | Height: | Size: 266 KiB |
@@ -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
|
||||||
|
BIN
evaluation_outputs/task1_feature_mlp/evaluation_val_curve.png
Normal file
|
After Width: | Height: | Size: 128 KiB |
BIN
evaluation_outputs/task1_feature_mlp/evaluation_val_s0.png
Normal file
|
After Width: | Height: | Size: 310 KiB |
|
After Width: | Height: | Size: 247 KiB |
|
After Width: | Height: | Size: 263 KiB |
BIN
scripts/__pycache__/config.cpython-310.pyc
Normal file
BIN
scripts/__pycache__/dataset.cpython-310.pyc
Normal file
BIN
scripts/__pycache__/model.cpython-310.pyc
Normal file
109
scripts/config.py
Normal file
@@ -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()
|
||||||
478
scripts/dataset.py
Normal file
@@ -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),
|
||||||
|
)
|
||||||
224
scripts/evaluate.py
Normal file
@@ -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()
|
||||||
31
scripts/model.py
Normal file
@@ -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)
|
||||||
143
scripts/predict_single.py
Normal file
@@ -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()
|
||||||
274
scripts/train.py
Normal file
@@ -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()
|
||||||
@@ -36,6 +36,11 @@ class DataConfig:
|
|||||||
pin_memory: bool = True
|
pin_memory: bool = True
|
||||||
shuffle_train: bool = True
|
shuffle_train: bool = True
|
||||||
drop_last_train: bool = False
|
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"
|
interpolation_method: str = "linear"
|
||||||
incomplete_file_policy: IncompleteFilePolicy = "skip"
|
incomplete_file_policy: IncompleteFilePolicy = "skip"
|
||||||
@@ -62,9 +67,11 @@ class ModelConfig:
|
|||||||
output_channels: int = 1
|
output_channels: int = 1
|
||||||
tcn_channels: tuple[int, ...] = (32, 32, 64, 64, 128)
|
tcn_channels: tuple[int, ...] = (32, 32, 64, 64, 128)
|
||||||
kernel_size: int = 5
|
kernel_size: int = 5
|
||||||
dropout: float = 0.25
|
dropout: float = 0.15
|
||||||
use_causal_conv: bool = True
|
use_causal_conv: bool = True
|
||||||
dilation_base: int = 2
|
dilation_base: int = 2
|
||||||
|
use_linear_skip_branch: bool = True
|
||||||
|
linear_skip_kernel_size: int = 33
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
@@ -78,7 +85,7 @@ class LossConfig:
|
|||||||
forward_rms_loss_weight: float = 8.0
|
forward_rms_loss_weight: float = 8.0
|
||||||
forward_scale_loss_weight: float = 8.0
|
forward_scale_loss_weight: float = 8.0
|
||||||
forward_underestimate_loss_weight: float = 16.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
|
envelope_kernel_size: int = 33
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -7,7 +7,7 @@ from typing import Any
|
|||||||
import numpy as np
|
import numpy as np
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
import torch
|
import torch
|
||||||
from torch.utils.data import DataLoader, Dataset
|
from torch.utils.data import DataLoader, Dataset, WeightedRandomSampler
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from .config import DataConfig, ExperimentConfig
|
from .config import DataConfig, ExperimentConfig
|
||||||
@@ -109,8 +109,11 @@ class WindowedTimeSeriesDataset(Dataset):
|
|||||||
self.report = report or SplitLoadReport(split=split)
|
self.report = report or SplitLoadReport(split=split)
|
||||||
self.sequence_store: list[dict[str, Any]] = []
|
self.sequence_store: list[dict[str, Any]] = []
|
||||||
self.window_index: list[tuple[int, int]] = []
|
self.window_index: list[tuple[int, int]] = []
|
||||||
|
self.window_weights: list[float] = []
|
||||||
|
|
||||||
min_length = max(config.min_sequence_length, config.window_size)
|
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:
|
for record in records:
|
||||||
sequence_length = int(record.time.shape[0])
|
sequence_length = int(record.time.shape[0])
|
||||||
@@ -135,6 +138,13 @@ class WindowedTimeSeriesDataset(Dataset):
|
|||||||
last_start = sequence_length - config.window_size
|
last_start = sequence_length - config.window_size
|
||||||
for start in range(0, last_start + 1, config.window_stride):
|
for start in range(0, last_start + 1, config.window_stride):
|
||||||
self.window_index.append((record_index, start))
|
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:
|
if not self.window_index:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
@@ -142,6 +152,15 @@ class WindowedTimeSeriesDataset(Dataset):
|
|||||||
"Check file completeness, sequence length, and window configuration."
|
"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:
|
def __len__(self) -> int:
|
||||||
return len(self.window_index)
|
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(
|
def build_datasets(
|
||||||
config: ExperimentConfig | DataConfig,
|
config: ExperimentConfig | DataConfig,
|
||||||
) -> tuple[dict[str, WindowedTimeSeriesDataset], NormalizationStats, dict[str, SplitLoadReport]]:
|
) -> 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]]:
|
) -> tuple[dict[str, DataLoader], dict[str, WindowedTimeSeriesDataset], dict[str, SplitLoadReport]]:
|
||||||
data_config = config.data if isinstance(config, ExperimentConfig) else config
|
data_config = config.data if isinstance(config, ExperimentConfig) else config
|
||||||
datasets, _, reports = build_datasets(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 = {
|
loaders = {
|
||||||
"train": DataLoader(
|
"train": DataLoader(
|
||||||
datasets["train"],
|
datasets["train"],
|
||||||
batch_size=data_config.batch_size,
|
batch_size=data_config.batch_size,
|
||||||
shuffle=data_config.shuffle_train,
|
shuffle=train_shuffle,
|
||||||
|
sampler=train_sampler,
|
||||||
num_workers=data_config.num_workers,
|
num_workers=data_config.num_workers,
|
||||||
pin_memory=data_config.pin_memory,
|
pin_memory=data_config.pin_memory,
|
||||||
drop_last=data_config.drop_last_train,
|
drop_last=data_config.drop_last_train,
|
||||||
|
|||||||
13
src/model.py
@@ -111,11 +111,22 @@ class TemporalConvNet(nn.Module):
|
|||||||
|
|
||||||
self.network = nn.Sequential(*blocks)
|
self.network = nn.Sequential(*blocks)
|
||||||
self.output_head = nn.Conv1d(in_channels, config.output_channels, kernel_size=1)
|
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)
|
self.receptive_field_info = compute_receptive_field(config)
|
||||||
|
|
||||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||||
features = self.network(x)
|
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):
|
class BuildingTCN(nn.Module):
|
||||||
|
|||||||
BIN
src_new/__pycache__/config.cpython-310.pyc
Normal file
BIN
src_new/__pycache__/dataset.cpython-310.pyc
Normal file
BIN
src_new/__pycache__/evaluate.cpython-310.pyc
Normal file
BIN
src_new/__pycache__/model.cpython-310.pyc
Normal file
BIN
src_new/__pycache__/train.cpython-310.pyc
Normal file
103
src_new/config.py
Normal file
@@ -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()
|
||||||
388
src_new/dataset.py
Normal file
@@ -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)
|
||||||
222
src_new/evaluate.py
Normal file
@@ -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()
|
||||||
140
src_new/model.py
Normal file
@@ -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)
|
||||||
327
src_new/train.py
Normal file
@@ -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()
|
||||||