diff --git a/data/surrogate/fea_surrogate_holdout.csv b/data/surrogate/fea_surrogate_holdout.csv new file mode 100644 index 0000000..acf16f9 --- /dev/null +++ b/data/surrogate/fea_surrogate_holdout.csv @@ -0,0 +1,201 @@ +sample,thickness_m,height_m,width_m,youngs_modulus_pa,density_kg_m3,tip_force_n,max_stress_pa,tip_deflection_m,mode_1_frequency_hz +1,0.0015673661079655607,0.04841099674629158,0.07852426675757909,197739429524.0499,8047.462134980211,808.0558923186421,110832000.0,0.008515325,56.28995 +2,0.001039199973926956,0.035477670011132674,0.08876809267456234,218031511424.30328,7968.190184923785,995.9491348513075,261662000.0,0.02469732,45.64068 +3,0.0012323572885739886,0.03672328918190327,0.07038402948587634,217994905341.84094,7620.661449593422,1243.9372998437982,327816000.0,0.03006282,46.89249 +4,0.0019751856295698762,0.03816195914701505,0.097842341697464,209492851972.2291,7691.704572169437,762.6824045406236,94556000.0,0.008619294,48.00584 +5,0.0013311309113572594,0.04752311009445264,0.09996699898696126,215188692211.80255,8026.035041284622,1299.1576003411121,172492000.0,0.01234711,59.55903 +6,0.0014437762569660816,0.044855514316994693,0.0609249744065813,198747500635.71484,8051.029257384443,1192.717245314734,239596000.0,0.01979418,51.26112 +7,0.001188200117606723,0.045047116274764536,0.09472976887840008,217551219689.42175,8020.003309830518,1124.139755269309,185257000.0,0.01384259,56.8891 +8,0.0015856986712866634,0.04378676832443886,0.09724598833583326,196604955408.89792,7982.483799812365,710.8808512232648,90921900.0,0.007723549,52.47048 +9,0.0016782206057048572,0.03927392133320407,0.07954518134759588,224840797559.7173,7822.929309307037,1024.392989993321,169412000.0,0.01406843,50.08283 +10,0.001739238181219579,0.0334400409983221,0.09101531975938702,195980793828.49066,7603.265872271703,803.7618602411698,139677000.0,0.01552385,41.13423 +11,0.0012114953072574599,0.041369402425203664,0.09609392493539576,196277758456.21185,8075.373956459915,1284.416315177556,227452000.0,0.02046817,49.80088 +12,0.0016889600105813264,0.049464936399504694,0.09816731357528284,205574519845.76608,7859.504567247029,904.7656979336423,93391000.0,0.006731061,60.50222 +13,0.0017616333466639916,0.04935738626766732,0.06416444308796024,218998000182.25275,8033.147205953683,1053.2592254806866,150001000.0,0.01022522,58.73076 +14,0.0011158906325899345,0.04819177746988225,0.06751990795967047,223105345139.5546,7626.146318351904,889.6435542631685,190356000.0,0.0130391,60.68956 +15,0.0017539792048929942,0.047920984049290585,0.07621517485660353,210024874093.76862,7677.030032970093,1039.0646319697055,133694000.0,0.009771794,58.42578 +16,0.0017478980794603147,0.03821519000331284,0.06666693233471288,200542339973.5879,7657.801672694839,1013.9662450946366,195002000.0,0.01869998,45.67526 +17,0.0016113559013571166,0.04118723061806716,0.08534868407155294,209655863978.86575,8094.728706374595,1191.8164676735798,181303000.0,0.01538619,50.1423 +18,0.001904378045795319,0.04530012576426517,0.09416551179524907,224009123610.90994,8072.299946375574,700.73977316053,74995600.0,0.005412578,56.91761 +19,0.0011841517017579326,0.043422209849552335,0.0814301379355353,216109335024.6959,7890.419653654939,1118.3160396393848,219975000.0,0.0172086,54.4346 +20,0.0012366465205376053,0.04191624326559128,0.09877328898445728,211478598183.6125,7752.652442256228,1217.9417700235506,203417000.0,0.01675835,53.51325 +21,0.001477650295242563,0.040380761596888956,0.06351935154951731,192871672342.10504,7632.463890251571,1233.908647151695,267714000.0,0.02528744,47.29688 +22,0.0011391530594900873,0.038901055584905836,0.08580041741818421,199904383779.8479,7764.221080588196,972.6192034882511,216524000.0,0.02038321,47.97305 +23,0.0013414587649204384,0.03995590933994123,0.09577433914426714,193966848738.52783,7991.053241140136,1031.8440115950982,173990000.0,0.01639287,48.0107 +24,0.001997740019058419,0.03646652298294407,0.08372509227731485,214305759579.95157,7665.057176008142,1166.1957995801354,173796000.0,0.01626863,45.87965 +25,0.0019504453652020443,0.04560888816184532,0.07319179396066991,221244298995.16376,7870.889131527212,757.0604766942041,97619200.0,0.007115635,56.0919 +26,0.0015847308468650116,0.034796497550405966,0.08430648553844938,198256539475.62366,7961.843150536535,1065.866762455843,204667000.0,0.02168197,41.90765 +27,0.001518114450704905,0.03697332361538613,0.07247611821031821,193223800946.8755,7644.632165625088,1209.1286407867012,255355000.0,0.02623102,44.15172 +28,0.0015510953869590418,0.04310173723064328,0.07616673214705878,216339577782.8144,7674.324813157276,1000.6865913059719,164247000.0,0.0129419,54.01726 +29,0.0013078829044340231,0.043954024247456105,0.08551666682344689,207507947311.71948,7919.368947595397,1062.5481507869151,180168000.0,0.014491,53.96372 +30,0.0016323039296096402,0.042619050299318006,0.07131088774139138,218378896276.89667,8098.492357739378,1074.3483388289787,180670000.0,0.01427215,51.81826 +31,0.0012662323829857342,0.04385812961415239,0.09854084867248417,222105022483.6572,7682.097878035567,1010.85914020214,156758000.0,0.0117647,57.41422 +32,0.001374640585217261,0.043181530306182145,0.09580879207185156,220954722750.09012,7684.425757461321,1265.1243469452004,189879000.0,0.01455575,56.1474 +33,0.0014128711118088686,0.04263144851191791,0.08011786891390649,190963685053.19946,7998.972792548506,930.534945985672,161359000.0,0.01455005,49.62517 +34,0.0018858236858815785,0.04877159536894916,0.08720150114022818,198858451981.69485,7940.143970615229,1272.448790067444,133893000.0,0.01013934,57.4968 +35,0.0013832913378363483,0.034925805266897515,0.08337035174255128,222443673680.8548,7911.087249652723,773.1293081558031,168432000.0,0.01585089,44.89963 +36,0.0016544027090606463,0.04006633413262866,0.060060916218754624,203396588836.3152,7808.437072263594,831.4654067035273,172699000.0,0.01559527,47.16835 +37,0.001795180350446238,0.03377573964731196,0.07936196498458331,201336068539.94766,7829.097868014607,1057.173808044511,199264000.0,0.02143598,40.91513 +38,0.0015327314507462338,0.03321412585666836,0.06964305509662681,211009248463.319,7651.8507684434435,1257.0752121824682,311323000.0,0.0325757,41.50712 +39,0.0012851372219377804,0.049977842602725045,0.08906809082449113,194258965995.7093,7994.663533900413,1085.5063166736395,154567000.0,0.01169386,58.74703 +40,0.0012515389727767862,0.03360113465467106,0.06890033298431736,207171094571.78906,7701.593906230043,849.3436819276468,250670000.0,0.0264168,41.74667 +41,0.00146737933046142,0.03244496881560363,0.07063363466196532,216880640817.15012,7784.928609510936,866.9898568427363,226854000.0,0.02362911,40.94043 +42,0.0019317641146437796,0.04691486795822283,0.0805946602656955,215627287884.17334,7989.7678345730255,748.9988217827794,86605800.0,0.006291821,57.03955 +43,0.0010510299497451068,0.04058826160088657,0.08629979956492509,221031431910.03152,7712.680955471077,921.6938861921277,209075000.0,0.01707135,52.78867 +44,0.0014957825984989325,0.04469300480395855,0.07995508874303417,213257780385.62943,7886.798738401926,834.1548477158365,129804000.0,0.01000414,55.05547 +45,0.001030131533893984,0.040828342194672115,0.061764468126851255,207888117727.47424,7964.369501325811,992.602685753483,302489000.0,0.02622681,48.93406 +46,0.00199469752182398,0.042865168459457134,0.06528878027083283,218729860792.63193,7862.052252307182,1184.697544754751,178596000.0,0.01401675,51.95893 +47,0.001507153280264632,0.040219111283724525,0.0775343224794501,191898625713.53555,8067.006047426973,743.7736729932545,134414000.0,0.01278127,46.66653 +48,0.0015900603540070963,0.03312244033038402,0.06061895643975591,191260121692.74353,7792.664770222287,1277.9252000666083,347056000.0,0.04024837,38.44766 +49,0.0017076061656356163,0.032616518968907916,0.09941516638771641,203883311554.99374,8069.08078964587,862.0177904348878,145178000.0,0.01583295,40.08426 +50,0.001806579334110566,0.035273622519560094,0.07797647280460752,211194660994.50092,7788.316732081006,1223.842636527906,220706000.0,0.02169762,43.70841 +51,0.001579029569964946,0.04398613108884024,0.06795342867111619,196791208367.51953,7974.191898954633,858.7361521025568,148938000.0,0.01266037,50.82732 +52,0.001558451723046742,0.03802386592697958,0.0821478863926264,204618662131.3446,8007.1499358816745,953.9630895256298,170529000.0,0.01605396,46.09497 +53,0.0012799261491605153,0.035058890999861404,0.09043099252087915,206390137981.1957,7889.897804677218,1077.0412211649505,233306000.0,0.02351938,43.9079 +54,0.0015202366326704138,0.038355329618242,0.07237437902214361,196932314451.581,7896.084641247488,753.0812510509969,151963000.0,0.01477142,45.38002 +55,0.0017799788807599734,0.04550854373079786,0.0796164081657828,218585220403.97324,8013.834713066852,785.2507999932986,102512000.0,0.007572096,55.85514 +56,0.0016190241472997676,0.047366695023209796,0.08358216228946458,224193692092.90384,7883.402579366172,737.0700617745523,95664800.0,0.006618217,59.69359 +57,0.0016469864212547156,0.034442699423933486,0.08309069025902543,221469190787.17325,7715.470373130519,1043.7308290238204,198721000.0,0.01904274,44.42043 +58,0.0014370235324109922,0.036202274868402634,0.06546534947359599,191715493809.70163,7832.838294034669,723.6432632452329,179854000.0,0.01903864,42.24999 +59,0.0019730996188538504,0.03659563090903129,0.07368909749537282,222995539275.62463,7770.404683855752,1081.7221158305806,181430000.0,0.01630817,46.09898 +60,0.0019217118843284202,0.04913120934721562,0.0878388544841741,193486654441.72858,7767.312356378798,916.0223419328465,93331600.0,0.00721068,57.73699 +61,0.001259594373757634,0.03582213545210425,0.07384445473260236,202784567053.59656,7777.11742195665,1248.3105117033401,318947000.0,0.0322,43.92393 +62,0.0012991475453806175,0.039142486119581286,0.08165625727179554,209091978811.18076,7747.052142460571,898.9506085660657,183746000.0,0.01645481,48.96949 +63,0.0017137672044462602,0.03626066004377332,0.06332886337594967,192010313550.67264,7727.931079564343,1141.502896828589,249450000.0,0.02633195,42.16212 +64,0.0011333011834980611,0.04982951645483909,0.09968761146998545,206939665433.344,7655.420862052629,1250.6519434908657,182717000.0,0.01298464,62.69824 +65,0.0015475662398414874,0.04707359279255719,0.06437705291739869,210405336797.78888,7709.00192608509,988.9888710189803,167698000.0,0.01247068,56.55851 +66,0.0010474933955920208,0.048522813537894435,0.09620221015576977,206582427168.56354,7727.463533898324,1079.3164680581353,180708000.0,0.01321505,60.73835 +67,0.0011798709190928602,0.048257553287084164,0.07471784790959508,206092816593.9507,7768.00310834795,927.0769009969251,172906000.0,0.01279333,58.46033 +68,0.001156851594157868,0.043430099439163254,0.0618230944464977,214430990851.74573,8061.82215859459,812.5401856300873,205771000.0,0.01626832,52.0954 +69,0.001769035396574269,0.04833523492055179,0.09930878945601779,224587860523.1023,7645.593437699972,783.8534405590175,79026000.0,0.00533134,62.72668 +70,0.0013773824137598475,0.03215374300069462,0.0828228278129575,208549761578.80975,7904.332843901159,1238.2313214303165,301226000.0,0.03279691,40.19211 +71,0.0012220607885406569,0.043648348134545135,0.07000627631562129,208262504446.17737,7830.710306588638,1071.5188828756816,231171000.0,0.01870594,52.99113 +72,0.0014053063997419468,0.04167665282673068,0.0881943479029065,217397929843.15857,7984.405079456771,871.2698126358365,143708000.0,0.01161749,52.40394 +73,0.001560459985306737,0.03399176352148401,0.0764296614188569,190090344269.99475,7847.4729515267,1101.7774455338197,240588000.0,0.02726604,40.08737 +74,0.0017431545005335935,0.041831255265977405,0.09827362038895882,199396351271.67993,7795.059549956641,877.7438444730954,108039000.0,0.009459272,51.08947 +75,0.0011632445157521514,0.038601654958560515,0.09181548183149575,211709981205.5263,7749.754376167714,1158.0596560699196,240669000.0,0.02151633,49.34219 +76,0.0014238026349586122,0.032084431640931534,0.07284631653308135,222730630715.75995,7929.062753435994,1285.773292600422,340944000.0,0.03494258,40.85015 +77,0.001452220127032031,0.046141026630263496,0.08100054687201927,201379077775.92865,7733.082960633141,768.3186022431049,116795000.0,0.009235516,55.76896 +78,0.0010050381508334568,0.039039196870009896,0.09225465000371905,216601934671.5639,7633.1146370658225,708.5098906492628,165711000.0,0.01431988,51.04413 +79,0.00167091847866873,0.04723190232881108,0.09316497133617607,210977220416.02084,7907.138892116954,807.1492130417147,93118800.0,0.006851083,58.25655 +80,0.0017249448132301603,0.0397537920685873,0.06055895817273243,209409954107.4152,7804.036407713038,1199.45692738683,240873000.0,0.02129047,47.48905 +81,0.0018552614678984415,0.049050773814640125,0.08419636524751185,213670125695.8628,8059.228118488878,859.8379219628283,93920100.0,0.006585763,59.2857 +82,0.0010045815546937072,0.04417800350826916,0.06303700055571566,201057047704.15128,7930.034693792971,933.2007950944312,259641000.0,0.02152253,51.94702 +83,0.0015957388207584729,0.037841896723912034,0.08756266775107506,217653798800.5499,7937.26866146517,1028.372539697004,171259000.0,0.01520478,47.77529 +84,0.0013170761011523147,0.037680286766674155,0.08607740941804368,203693561432.2554,8031.645753847682,747.0211108094937,150869000.0,0.01438005,46.02009 +85,0.0012458249668523226,0.04225966234035547,0.07912146538698878,194689486156.2207,7849.429299667002,1106.7723889702525,220033000.0,0.01963419,50.3032 +86,0.0010985291536961942,0.03687269824749808,0.0935577567641524,210541985003.80002,7756.874759832197,1033.280135777542,235266000.0,0.02210673,47.27606 +87,0.0010226119363994104,0.032290222037732706,0.08245157711978324,193797786197.0026,8038.941679769356,801.4795562110792,255056000.0,0.02976665,38.97707 +88,0.001758685706024517,0.04101987047984067,0.09748150103620733,221943527069.08676,7942.550413057737,1228.5185128442276,154763000.0,0.01241199,52.35887 +89,0.0018451945017654008,0.04944734829642077,0.09173963765114318,202447909782.15817,7800.015763122335,923.6970178296677,93303400.0,0.006840789,59.64812 +90,0.0010849620627911612,0.049730735320534375,0.06449949436476743,191462249167.85602,8041.171885594976,1132.2369388357652,274908000.0,0.01921077,56.05293 +91,0.0010421804361146166,0.041106874727668295,0.06395337958369562,219408913622.9559,8043.039854671228,957.6670349708677,278106000.0,0.02268699,50.51724 +92,0.0016205210041619186,0.0336734866208742,0.0751510484112767,199129851657.2115,8036.517723219117,846.4032976835842,183602000.0,0.02005491,40.04865 +93,0.0016015890584605197,0.03552789761231871,0.08963224011892446,199692174977.9141,8053.672930010988,1135.6971819028868,199565000.0,0.02052846,42.8709 +94,0.001867825158171037,0.03264977456976684,0.06375612414193613,201894733035.5186,8086.14286467085,1220.1560538000003,281343000.0,0.03133239,38.12856 +95,0.001833833893776132,0.047997767039807585,0.07838595619617467,217299791106.24197,7757.80436871158,1004.6630921147075,121060000.0,0.008535809,59.29219 +96,0.0015021101149215199,0.03272137577534843,0.06507743708324591,200028754477.94748,7613.913122485475,791.1076712685681,215557000.0,0.02417201,39.71535 +97,0.00110960948597827,0.04637425852318981,0.08068697166775098,192635591013.00272,7880.214899972686,1260.5005314386651,245088000.0,0.02016124,54.65483 +98,0.0010864342163111665,0.04588535358151211,0.07054521754090677,214642055712.81595,7867.212944771455,1162.8207458900522,261432000.0,0.01953364,56.38272 +99,0.0012712314914494344,0.03407915078546634,0.08495219875716017,193035064059.60336,7908.070759851666,1088.3901143260036,259298000.0,0.02879037,41.07504 +100,0.0012077575997844637,0.036071312185907485,0.062457398565266754,219983167837.66458,7641.648439343085,824.5948930351431,250637000.0,0.02321734,45.72593 +101,0.0018142783013398948,0.03754269336086639,0.0733429778915637,204269860813.3154,7669.492686482894,1236.2910289079055,216732000.0,0.02073963,45.67229 +102,0.0018000723311684549,0.044115392038015094,0.08688316999779941,202095964427.29813,7957.127344479432,735.7660205465351,91711800.0,0.007543856,52.81337 +103,0.0011223383950565336,0.04506257244447126,0.07891431092604256,210818100043.21133,7840.231395065621,1020.3783712956795,207118000.0,0.01602046,55.69354 +104,0.0014016621176678084,0.034426809003095894,0.09647405130775297,207786728735.9325,7952.5493647100275,780.6242127969855,150465000.0,0.01529923,43.22142 +105,0.0017704757962709076,0.033421246281669446,0.06472115714929726,223602413896.8454,7843.669182139326,1140.1572235774954,263601000.0,0.02589641,41.84434 +106,0.0013886949749571434,0.03810924648879894,0.06495911775546595,208814004720.71442,7711.499828792607,880.9726732717877,213204000.0,0.01969562,46.63038 +107,0.0013261576239301915,0.04296036438925564,0.09027702569151325,203199911388.25662,7959.615849747899,938.745782440838,154236000.0,0.0129413,52.40961 +108,0.0016363061878319138,0.03425659418774961,0.0683850112635012,213823827522.47827,7791.355904209917,712.5010549793875,162919000.0,0.01632158,42.42912 +109,0.0018181091378455696,0.04770865096845556,0.06652516950301474,203020338114.57333,7721.80895249033,1121.0657229338208,157375000.0,0.01196586,56.09059 +110,0.0013109984066376856,0.040534700617946895,0.07566686478430668,196348698524.18353,7688.311340695252,978.5933232571515,202965000.0,0.01872485,48.81621 +111,0.0011436771085663927,0.035196091833916275,0.060331439712804996,214976831125.78976,7922.745011584477,876.3594389075542,297178000.0,0.02887356,43.31506 +112,0.0016264371980254555,0.03474558856884679,0.08456029857172596,216986285886.1201,7925.728932515355,948.6741420048995,177896000.0,0.01724232,43.84085 +113,0.0015136803230920435,0.0424901365447439,0.07663019763927897,204716580642.82938,7618.907991734484,1252.00327988962,212725000.0,0.01796375,52.12199 +114,0.0017805768073767179,0.0454707071228888,0.09882145058770608,194156254625.91394,8023.235531865769,1115.8495865597542,121064000.0,0.01003115,53.73454 +115,0.001016472325554434,0.04468605733248764,0.0745395563739293,195670729396.858,7628.528337510765,827.2142779425922,195267000.0,0.01642038,53.77557 +116,0.001915427308132725,0.0351236747504744,0.06602565887301658,194450906151.75076,8063.723293756274,719.3564956965558,143477000.0,0.0154269,40.25507 +117,0.0016093089703608106,0.03899192063381924,0.07712151301306412,200745660197.41342,8008.155290062704,1188.7779451012427,211478000.0,0.01982106,46.39475 +118,0.001101714304955322,0.04300330601737026,0.07814752871902858,197012207748.7253,7854.764893128109,754.1545155548057,166230000.0,0.01441112,51.50641 +119,0.0011267730766770275,0.034597779548291345,0.09505068678757259,206782448973.52103,7605.941738359689,1155.6780081692723,273529000.0,0.02783163,44.58752 +120,0.0017004362529490467,0.03681254847138026,0.09382132677645749,215224424890.18192,7704.921338588467,897.3035204828801,137546000.0,0.01266027,47.13577 +121,0.0014346785599403619,0.03961867921442355,0.07262881053022027,190556446077.9698,7780.693109653845,770.3809137117389,156837000.0,0.01525809,46.46342 +122,0.0010645569553994537,0.03461335386673619,0.09156416864312814,212334033769.95935,7615.640569400991,814.3570224012201,209724000.0,0.02080112,45.1069 +123,0.0010698832643761139,0.03720265823742174,0.06151234479735368,205715782197.1977,7653.8112019748005,1292.2463222402646,427943000.0,0.04111955,45.55023 +124,0.001854665400976728,0.0444714803186271,0.0944990999885327,195362792209.40564,7744.512326450841,764.9880095866127,85493600.0,0.007203641,53.41113 +125,0.0014911140369107145,0.045773279079021706,0.08673103282827536,213606809145.14417,8081.499293946279,822.4006908809966,116268000.0,0.008726384,56.12505 +126,0.0015394956552635159,0.04657196105972297,0.062355285149266616,212798779022.04382,7900.620029026496,1275.7912487676494,226186000.0,0.01681181,55.44077 +127,0.0013356540204718737,0.04675096514603997,0.08925402175779344,198102860902.35635,7873.545289035963,793.6455826528144,118055000.0,0.009352423,56.18534 +128,0.0019395142773716277,0.03564282038280661,0.08950462319711852,200248189385.00027,7623.094173822272,1267.3181373603966,187950000.0,0.0192185,43.84175 +129,0.0018947731168686595,0.048970634776547284,0.06280531851964724,197284801209.54675,7718.531888356908,980.6045883096324,134540000.0,0.01026151,56.16346 +130,0.0018801506383905988,0.04778473910878073,0.075836750283087,221797474435.14355,7894.492121077844,1146.9686541833464,139774000.0,0.009701587,58.86648 +131,0.0017291431567619188,0.04439667731413467,0.087624752090802,204446200088.15973,8096.636522840804,883.3771017697487,112417000.0,0.009082237,53.10276 +132,0.0015743556274395235,0.04162955363116832,0.0906345559748095,220351248872.5613,7696.4040722577965,959.5045353450957,139640000.0,0.01114202,53.6082 +133,0.0017856398526993147,0.03746698051689425,0.09778055307018077,195901511747.9943,7611.616109621621,1126.3026422364273,156070000.0,0.01548987,46.07681 +134,0.0010259016107971182,0.044880761291556624,0.06705644934174185,194970505839.49564,8083.656081071819,907.3627067488897,230689000.0,0.01940475,51.73302 +135,0.001860082355902441,0.03854848144068639,0.08882994030522845,190304186730.0698,8027.721594459935,1280.479653229362,179798000.0,0.01792092,44.95892 +136,0.001542604286776804,0.04712174122173453,0.06946435429203932,211621156319.32574,7742.179922842699,1130.2153666844827,180251000.0,0.0133065,57.16057 +137,0.0019447376191786942,0.0422851148591643,0.07542812291255657,205298266274.5523,7857.376779458624,1008.9878235583678,139903000.0,0.0118392,50.57019 +138,0.0012647021238459378,0.035709830364575895,0.07084535437356718,205916444480.02826,7761.976728256548,1047.780579380404,277540000.0,0.02769773,43.98809 +139,0.001346410547180829,0.04866970272097498,0.0741878147172175,190710370487.26245,7664.975924986586,788.8564370438124,129569000.0,0.01027382,56.8145 +140,0.0019109316502197373,0.03302525156876573,0.06574641204496581,206105154795.34467,7975.324757886759,870.9361560519158,188977000.0,0.02037559,39.28643 +141,0.0011148857469977088,0.0464680233751421,0.07496367498340413,201908150953.0058,7879.964310885254,1150.2849743610118,236283000.0,0.01852402,55.60907 +142,0.001464740911934044,0.049711859736627784,0.08570023348301076,214018409316.24356,7897.60834917751,717.947018940474,94184600.0,0.006505656,61.27031 +143,0.0010579979117090685,0.0470124570663052,0.0839531746612341,212486848097.1356,8004.897754203495,1023.4701695805461,197925000.0,0.01455433,57.97167 +144,0.0014800732808251515,0.04012106048632336,0.07527756124140332,201020267934.16144,7638.26756945313,1242.838391783298,234882000.0,0.0213825,48.8582 +145,0.0017912352056033884,0.04325027170806647,0.09421868947751355,202376141880.8521,7678.675861569514,1092.5190360774088,130678000.0,0.01092428,53.24379 +146,0.00165790347241636,0.03951530093860501,0.06929134720330497,197614344801.13553,7826.261439274482,1037.8916718429637,193947000.0,0.01825191,46.57088 +147,0.0012949279288387442,0.04520876209150625,0.08191274123615717,223267172060.27637,7934.115768484856,704.2296231172097,120804000.0,0.008790025,57.15468 +148,0.0014279348485995121,0.03380688728234241,0.08708253689029355,198414717657.72092,7806.5997873960105,940.9015676920483,199523000.0,0.02170553,41.49771 +149,0.0012033565547046893,0.039300725712802624,0.06844237979174157,212631422261.5503,7836.766650910136,739.1036313470431,187598000.0,0.01649932,48.51903 +150,0.0018205626223580078,0.04151748853300725,0.07862703777005536,221655632133.8367,7817.348677921927,729.1531923435158,105772000.0,0.00843653,52.16255 +151,0.0012261832963326308,0.032477073313271404,0.06802175387401485,224803137571.3637,7609.569138392722,1262.054162715256,399987000.0,0.04018244,42.35512 +152,0.0019061186545965573,0.04746940533500796,0.0929129107625721,205199519139.2358,7996.322012441318,982.6512261583089,100317000.0,0.007551378,57.10437 +153,0.0011519650657749714,0.04096400873897511,0.06916738581374614,219798470393.9508,7864.459897792613,1289.1528305942006,320826000.0,0.02619655,51.27512 +154,0.0018794322203857946,0.04864858463337651,0.09666200816022084,207030205384.4084,7698.986723451884,1294.8602703368176,125523000.0,0.00913422,60.08103 +155,0.0016446116014318438,0.04130567961813221,0.09906111012782726,199086099123.769,7774.421096056694,1226.1355442961103,160346000.0,0.0142322,50.66562 +156,0.0019569297695237126,0.040768995044271984,0.09493443194509107,210245696381.0047,7820.20101814875,936.4480199158262,110711000.0,0.009436892,50.69072 +157,0.00147092886469247,0.03968772833832893,0.0884026502650962,191088401872.02737,7817.773839424196,1148.0579915835758,192434000.0,0.0185675,47.36791 +158,0.0016937293980628112,0.037105961117207896,0.06871887695375847,208963657552.7331,8016.631308711516,1174.1953279389306,234688000.0,0.02222964,44.53362 +159,0.001962878250270023,0.03882170500817191,0.09338947801542177,218278172058.9604,7970.9822531522605,855.2666991715206,108720000.0,0.009368498,48.74445 +160,0.0011749996830923153,0.048066486042962375,0.0673468123643724,220627336650.59006,7851.973974489205,1096.4628537414767,224781000.0,0.01561048,59.24859 +161,0.0014165844228562787,0.04759728245773061,0.07201220882125688,224467156763.09308,8079.125049693023,798.9506839471545,131984000.0,0.009092194,58.55143 +162,0.0010135904546021955,0.04239470027244893,0.08647039605624748,207393766727.6782,7787.31021289188,965.2505899874918,214627000.0,0.01789362,53.03119 +163,0.0014574265039940374,0.042763641373258734,0.07761695134610702,219100347378.9571,7916.768848718892,967.3140678898401,166912000.0,0.01308424,53.35344 +164,0.0018272563447773327,0.04207790623804254,0.06681187029713344,208660995580.10422,7723.703316520129,1171.4372107187214,191443000.0,0.01603914,50.65791 +165,0.001241425230609511,0.049605850928105266,0.08462589575517705,223879890303.69653,7985.77424750156,888.9857569858998,137743000.0,0.009116388,62.38694 +166,0.0010938029420502551,0.03416903341659582,0.09680088029212675,197514800249.02713,7705.6472658984185,1179.3099366855215,286527000.0,0.03087322,42.89242 +167,0.0016664810230683086,0.04679942361260607,0.09082859902162574,220478342667.78046,8074.04006396437,836.9563686993101,100047000.0,0.007111387,58.28917 +168,0.0013913526100700338,0.04885617119291578,0.06770376269539231,216020010817.63977,7977.670058847445,903.8389266997189,154605000.0,0.01078968,58.80261 +169,0.0011468392474162526,0.04174615457688004,0.07198845798306255,219283297849.10898,8049.730597192357,1054.3666726831339,248882000.0,0.0199833,51.7342 +170,0.0013237017257300607,0.03594717081009375,0.0808545750173205,193658192786.83762,7671.8566633743785,850.3616812290466,191159000.0,0.02010231,43.66194 +171,0.001715146630955862,0.03289558588495469,0.08024179588399599,213326751314.4753,8000.622575313186,1196.5646053587998,240085000.0,0.02500946,40.75781 +172,0.0010751701596541398,0.03222635227988049,0.09136065104630277,220111630155.42178,7920.162621156419,1104.3435439344362,307461000.0,0.03154529,42.07259 +173,0.0014860476026900197,0.03327125312102871,0.06126001078571163,212996033240.05478,7692.95952653101,910.0497337011886,258359000.0,0.02678262,41.18135 +174,0.0011696735604229606,0.03983066641241566,0.06118457817090618,216437546294.57846,7635.073948270479,963.8923544603529,271466000.0,0.02317023,49.63737 +175,0.0016616455346265724,0.04067764393333946,0.09363816618115932,200355483220.07404,7949.845473224718,975.3968734801692,135244000.0,0.01212933,49.26358 +176,0.0019452658401220404,0.03602861330880891,0.09262781187486727,197890942036.70685,7945.354362549695,1067.1843259899786,151121000.0,0.01545533,43.25951 +177,0.0016829828369754433,0.04666361052301451,0.08984925797336436,194769721270.16266,7779.232119879119,1159.1806083015954,139114000.0,0.01122803,55.58289 +178,0.0013608239529378937,0.03846044616197594,0.09530994736223242,214710511643.2613,7877.172679303085,894.0282726904414,156464000.0,0.01382539,49.03872 +179,0.0019883818841470438,0.04919270910657715,0.0743374524499838,190504984651.28082,7837.810840205586,730.2307881124138,82998300.0,0.006518456,55.98759 +180,0.0012838808027105303,0.033904954904771016,0.07345190558874638,215507655930.34308,7648.553830783976,944.4092689551005,254822000.0,0.02555906,43.31527 +181,0.001398048783856736,0.037626899655335305,0.06994993611148863,202678316135.34772,7735.159326596529,1109.5745446238036,254520000.0,0.02451061,45.6748 +182,0.0019821681539476366,0.03724049568996849,0.09015897857920538,219616201751.56573,7939.995776537255,1204.1711657458736,164877000.0,0.0147233,46.90096 +183,0.0015294408446765045,0.03632673800112669,0.09202514094317588,195572915406.13754,7750.423273345416,1212.7513375573612,211114000.0,0.0216788,44.36539 +184,0.00173248462564456,0.032968934873695184,0.08238643654872985,222360873419.20856,7967.062463769673,1213.7876866588938,235220000.0,0.02343884,41.86569 +185,0.0011949982130658285,0.046073062032126924,0.07739841983767326,192206868421.1705,7814.47814331933,913.9546527966384,173127000.0,0.01437321,54.15843 +186,0.0016957750689699561,0.04627728790025447,0.07104079893183664,195123331817.1899,7867.585511649465,1168.6288802010977,171805000.0,0.01400099,53.51619 +187,0.0019266655315298609,0.03788821823851027,0.08277047972092491,204936764621.52954,7739.356825822281,775.6886648734088,114826000.0,0.01082862,46.34981 +188,0.0012174548641305184,0.04451159303000278,0.08515693502630645,211978413881.79114,7952.080206875343,840.6208568879327,150523000.0,0.01170616,55.15195 +189,0.001195557258039558,0.03868901353122211,0.09707050578181561,201660713347.5619,8090.481374352285,817.0372716725276,157317000.0,0.0147064,47.43178 +190,0.0013532584003990154,0.044314274401153764,0.0768094690842071,222622456296.70026,8011.3512350648025,985.5331678826299,175707000.0,0.01308924,55.31521 +191,0.0014462347749105487,0.03541856060902097,0.09248641605455668,212211853356.54623,8088.607467847734,949.6016343880086,178437000.0,0.01730696,44.28207 +192,0.0018408117604271513,0.0394543148773705,0.066286296087238,192549760882.28403,7811.373122958586,842.6441763458685,149490000.0,0.014467,45.51884 +193,0.0013679864251037004,0.042163633761963304,0.06276441058466486,223490037364.635,7799.695443911549,1093.2096523145299,243296000.0,0.01900282,52.46834 +194,0.0019696935672332333,0.045748633662070264,0.0620769578081231,208043612537.48062,7660.523391733803,1049.9890585313217,153501000.0,0.01187838,54.23168 +195,0.0013560978803749569,0.03736524055361199,0.07173424234388633,199582073076.28568,7602.293351519349,1203.8900777588865,280157000.0,0.02757811,45.5906 +196,0.0010707216216999163,0.036553803162769086,0.08839499701041932,209898537469.9625,8055.227835959226,998.8772201271333,247239000.0,0.02354378,45.74624 +197,0.0018980188368081274,0.04030403788039735,0.07153749572593633,192366718488.7468,7912.677157871229,724.8186642061315,114153000.0,0.01081771,46.42789 +198,0.001873275543240377,0.045320024080245824,0.0813190749723127,204028024128.5546,7732.185338139463,1180.33023927621,145353000.0,0.01154611,54.73734 +199,0.0013043661468336654,0.04352988547097253,0.06595787031753336,203636512916.76678,8018.050736604178,1112.7347490132402,238452000.0,0.01979678,51.21833 +200,0.0018383616248104525,0.0460352514020682,0.09556671687880627,215733806541.34338,7687.098211973478,1015.319046462962,108532000.0,0.008003181,58.27658 diff --git a/docs/figures/ml_accuracy_vs_training_size.png b/docs/figures/ml_accuracy_vs_training_size.png new file mode 100644 index 0000000..bc61f7a Binary files /dev/null and b/docs/figures/ml_accuracy_vs_training_size.png differ diff --git a/docs/figures/ml_break_even_vs_training_size.png b/docs/figures/ml_break_even_vs_training_size.png new file mode 100644 index 0000000..00aa2b2 Binary files /dev/null and b/docs/figures/ml_break_even_vs_training_size.png differ diff --git a/docs/figures/ml_runtime_comparison.png b/docs/figures/ml_runtime_comparison.png new file mode 100644 index 0000000..4edba29 Binary files /dev/null and b/docs/figures/ml_runtime_comparison.png differ diff --git a/docs/validation/ml_break_even.csv b/docs/validation/ml_break_even.csv new file mode 100644 index 0000000..f0a8394 --- /dev/null +++ b/docs/validation/ml_break_even.csv @@ -0,0 +1,5 @@ +training_samples,median_model_training_seconds,median_fea_seconds_per_design,median_surrogate_seconds_per_design,speedup,sequential_equivalent_break_even_queries,incremental_break_even_queries_if_fea_data_is_sunk +50,2.3846008330001496,0.06928459049959201,2.4136844996974106e-05,2870.49075835213,84.44689645187897,34.42947175732142 +100,1.1080441569938555,0.06928459049959201,2.4136844996974106e-05,2870.49075835213,116.03307201872305,15.998222629607955 +250,1.3011211090051802,0.06928459049959201,2.4136844996974106e-05,2870.49075835213,268.87304011569523,18.785916642907495 +500,3.0629117009957554,0.06928459049959201,2.4136844996974106e-05,2870.49075835213,544.3973430903196,44.22309614474418 diff --git a/docs/validation/ml_runtime_benchmark.csv b/docs/validation/ml_runtime_benchmark.csv new file mode 100644 index 0000000..bb6e519 --- /dev/null +++ b/docs/validation/ml_runtime_benchmark.csv @@ -0,0 +1,11 @@ +design,fea_seconds,surrogate_seconds,surrogate_microseconds,speedup +1,0.1308434840029804,2.4605792001239025e-05,24.605792001239024,5317.588801709442 +2,0.05166892499255482,2.463494949915912e-05,24.63494949915912,2097.383028705558 +3,0.06922784099879209,2.4324389996763785e-05,24.324389996763784,2846.025779392718 +4,0.06427983500179835,2.4318238502019085e-05,24.318238502019085,2643.2767733757255 +5,0.0643696709885262,2.4068134996923617e-05,24.068134996923618,2674.4768963924257 +6,0.07119906300795265,2.404276149900397e-05,24.042761499003973,2961.3512994712455 +7,0.06934134000039194,2.420555499702459e-05,24.205554997024592,2864.687052575971 +8,0.07262198999524117,2.4027854000451043e-05,24.027854000451043,3022.40849282162 +9,0.06555216800188646,2.3994212999241428e-05,23.994212999241427,2731.9990867780857 +10,0.07112148200394586,2.39602619985817e-05,23.9602619985817,2968.3098627283716 diff --git a/docs/validation/ml_sample_efficiency.csv b/docs/validation/ml_sample_efficiency.csv new file mode 100644 index 0000000..a24e089 --- /dev/null +++ b/docs/validation/ml_sample_efficiency.csv @@ -0,0 +1,13 @@ +training_samples,training_seed,best_epoch,training_seconds,stress_mae_mpa,stress_mape_percent,deflection_mae_mm,deflection_mape_percent,frequency_mae_hz,frequency_mape_percent +50,11,997,3.4642961810022825,6.866049174085978,3.4474224405541487,0.7498572583446255,4.16635850829526,0.5034619353940796,0.9997557665489964 +50,22,853,2.3846008330001496,6.094562400789273,3.3270179258758508,0.7925686314606283,4.536706605059106,0.5505339110088674,1.105739782763316 +50,33,312,0.9896144550002646,6.125175256295271,3.4069044767291956,0.7332625301228618,4.731202795699453,0.4581106917814379,0.9390807228557378 +100,11,218,1.1080441569938557,4.351308041553677,2.3161539755271727,0.44613879500113873,2.6868529845017135,0.32136209147216155,0.6553676785429958 +100,22,424,1.6328563980059698,3.848124641176021,1.9775999021909814,0.46894848230469455,2.722674610630376,0.3706394474622282,0.7379393394766395 +100,33,196,0.6949700510012917,4.145067825131325,2.2909306535623988,0.49258496094957677,3.0646649134609185,0.26235689897040765,0.5448309486457028 +250,11,280,1.5792550449987175,3.6717981375085804,1.953591187415173,0.3369824003256247,1.9725709129716034,0.18214313826487896,0.3696252886072582 +250,22,186,1.3011211090051802,3.0708207934280036,1.5698838493673906,0.3452453334307398,2.0289620364843013,0.21551091608795475,0.4325561028307659 +250,33,188,1.2517717239970807,3.2632325247831617,1.7428773009875393,0.3359198731630988,2.021027251202072,0.19179065506606421,0.3909628170140283 +500,11,307,3.782049052999355,2.411647806960748,1.2421835441042974,0.20752295927102588,1.1687077888042159,0.11618084285121405,0.24237984274190452 +500,22,206,3.0629117009957554,2.5937687257548916,1.2622236702097585,0.25053278539598955,1.4290929190361925,0.1497163155806846,0.301327590975663 +500,33,170,2.3323324599914486,2.3838750738071783,1.2680807057387091,0.2659226072635731,1.609994265582561,0.14517910115881505,0.29390010263258365 diff --git a/scripts/benchmark_surrogate_runtime.py b/scripts/benchmark_surrogate_runtime.py new file mode 100644 index 0000000..e2e2c24 --- /dev/null +++ b/scripts/benchmark_surrogate_runtime.py @@ -0,0 +1,661 @@ +import csv +from dataclasses import replace +from pathlib import Path +from statistics import median +from time import perf_counter + +import matplotlib.pyplot as plt +import numpy as np +import pandas as pd +import torch + +from bodysimpy.config.loader import load_config +from bodysimpy.ml.dataset import Standardization +from bodysimpy.ml.dataset_generation import evaluate_design_point +from bodysimpy.ml.design_space import ( + SurrogateDesignPoint, + generate_design_points, +) +from bodysimpy.ml.experiment import calculate_runtime_tradeoff +from bodysimpy.ml.model import StructuralSurrogate +from bodysimpy.modeling.crossmember import build_crossmember_model + +BENCHMARK_DESIGNS = 10 +BENCHMARK_SEED = 20260816 + +WARMUP_CALLS = 100 +INFERENCE_REPETITIONS_PER_DESIGN = 2_000 + +CHECKPOINT_PATH = Path("models/checkpoints/structural_surrogate.pt") + +SAMPLE_EFFICIENCY_PATH = Path("docs/validation/ml_sample_efficiency.csv") + +RUNTIME_CSV_PATH = Path("docs/validation/ml_runtime_benchmark.csv") + +BREAK_EVEN_CSV_PATH = Path("docs/validation/ml_break_even.csv") + +RUNTIME_FIGURE_PATH = Path("docs/figures/ml_runtime_comparison.png") + +BREAK_EVEN_FIGURE_PATH = Path("docs/figures/ml_break_even_vs_training_size.png") + + +def _as_float64_array( + value: object, + *, + name: str, +) -> np.ndarray: + """Convert checkpoint preprocessing data to a 1D float array.""" + + array = np.asarray( + value, + dtype=np.float64, + ) + + if array.ndim != 1: + raise ValueError(f"{name} must be a one-dimensional array.") + + return array + + +def load_surrogate_checkpoint() -> tuple[ + StructuralSurrogate, + Standardization, + Standardization, +]: + """Load the trained surrogate and preprocessing statistics.""" + + if not CHECKPOINT_PATH.exists(): + raise FileNotFoundError( + "Surrogate checkpoint was not found at " + f"{CHECKPOINT_PATH}. " + "Run scripts/train_surrogate.py first." + ) + + checkpoint = torch.load( + CHECKPOINT_PATH, + map_location="cpu", + weights_only=False, + ) + + if not isinstance(checkpoint, dict): + raise TypeError("Surrogate checkpoint must contain a dictionary.") + + required_keys = { + "model_state_dict", + "feature_mean", + "feature_standard_deviation", + "target_mean", + "target_standard_deviation", + } + + missing_keys = required_keys - checkpoint.keys() + + if missing_keys: + missing = ", ".join(sorted(missing_keys)) + + raise KeyError(f"Surrogate checkpoint is missing: {missing}") + + model = StructuralSurrogate() + + model.load_state_dict(checkpoint["model_state_dict"]) + + model.eval() + + feature_standardization = Standardization( + mean=_as_float64_array( + checkpoint["feature_mean"], + name="feature_mean", + ), + standard_deviation=_as_float64_array( + checkpoint["feature_standard_deviation"], + name="feature_standard_deviation", + ), + ) + + target_standardization = Standardization( + mean=_as_float64_array( + checkpoint["target_mean"], + name="target_mean", + ), + standard_deviation=_as_float64_array( + checkpoint["target_standard_deviation"], + name="target_standard_deviation", + ), + ) + + return ( + model, + feature_standardization, + target_standardization, + ) + + +def design_to_feature_array( + design: SurrogateDesignPoint, +) -> np.ndarray: + """Convert one structural design into the six ML inputs.""" + + return np.array( + [ + [ + design.thickness_m, + design.height_m, + design.width_m, + design.youngs_modulus_pa, + design.density_kg_m3, + design.tip_force_n, + ] + ], + dtype=np.float64, + ) + + +def predict_single_design( + model: StructuralSurrogate, + design: SurrogateDesignPoint, + *, + feature_standardization: Standardization, + target_standardization: Standardization, +) -> np.ndarray: + """Run complete end-to-end inference for one structural design.""" + + features = design_to_feature_array(design) + + normalized_features = ( + features - feature_standardization.mean + ) / feature_standardization.standard_deviation + + feature_tensor = torch.tensor( + normalized_features, + dtype=torch.float32, + ) + + with torch.inference_mode(): + normalized_prediction = model(feature_tensor).cpu().numpy().astype(np.float64) + + prediction = ( + normalized_prediction * target_standardization.standard_deviation + + target_standardization.mean + ) + + return prediction[0] + + +def warm_up_surrogate( + model: StructuralSurrogate, + designs: tuple[SurrogateDesignPoint, ...], + *, + feature_standardization: Standardization, + target_standardization: Standardization, +) -> None: + """Warm up PyTorch before latency measurements.""" + + for call_index in range(WARMUP_CALLS): + design = designs[call_index % len(designs)] + + predict_single_design( + model, + design, + feature_standardization=(feature_standardization), + target_standardization=(target_standardization), + ) + + +def benchmark_fea( + designs: tuple[SurrogateDesignPoint, ...], +) -> tuple[float, ...]: + """Measure sequential static-plus-modal FEA runtime.""" + + config = load_config("configs/baseline_crossmember.yaml") + + base_model = build_crossmember_model(config) + + benchmark_model = replace( + base_model, + name=(f"{base_model.name}_runtime_benchmark"), + ) + + durations: list[float] = [] + + print() + print("Benchmarking CalculiX FEA") + print("-" * 70) + + for design in designs: + start = perf_counter() + + sample = evaluate_design_point( + benchmark_model, + design, + ) + + elapsed = perf_counter() - start + + durations.append(elapsed) + + print( + f"Design {design.sample_index:>2}: " + f"{elapsed:.6f} s | " + f"stress=" + f"{sample.max_stress_pa / 1e6:.3f} MPa | " + f"deflection=" + f"{sample.tip_deflection_m * 1000.0:.4f} mm | " + f"mode1=" + f"{sample.mode_1_frequency_hz:.4f} Hz" + ) + + return tuple(durations) + + +def benchmark_surrogate( + model: StructuralSurrogate, + designs: tuple[SurrogateDesignPoint, ...], + *, + feature_standardization: Standardization, + target_standardization: Standardization, +) -> tuple[float, ...]: + """Measure end-to-end single-design surrogate inference.""" + + print() + print("Warming up PyTorch surrogate") + print("-" * 70) + + warm_up_surrogate( + model, + designs, + feature_standardization=(feature_standardization), + target_standardization=(target_standardization), + ) + + print(f"Completed {WARMUP_CALLS} warm-up calls.") + + durations: list[float] = [] + + print() + print("Benchmarking PyTorch inference") + print("-" * 70) + + for design in designs: + start = perf_counter() + + for _ in range(INFERENCE_REPETITIONS_PER_DESIGN): + predict_single_design( + model, + design, + feature_standardization=(feature_standardization), + target_standardization=(target_standardization), + ) + + elapsed = perf_counter() - start + + seconds_per_design = elapsed / INFERENCE_REPETITIONS_PER_DESIGN + + durations.append(seconds_per_design) + + print(f"Design {design.sample_index:>2}: {seconds_per_design * 1e6:.3f} microseconds/query") + + return tuple(durations) + + +def write_runtime_csv( + fea_durations: tuple[float, ...], + inference_durations: tuple[float, ...], +) -> None: + """Write individual runtime measurements to CSV.""" + + RUNTIME_CSV_PATH.parent.mkdir( + parents=True, + exist_ok=True, + ) + + with RUNTIME_CSV_PATH.open( + "w", + encoding="utf-8", + newline="", + ) as stream: + writer = csv.writer(stream) + + writer.writerow( + [ + "design", + "fea_seconds", + "surrogate_seconds", + "surrogate_microseconds", + "speedup", + ] + ) + + for index, ( + fea_seconds, + inference_seconds, + ) in enumerate( + zip( + fea_durations, + inference_durations, + strict=True, + ), + start=1, + ): + writer.writerow( + [ + index, + fea_seconds, + inference_seconds, + inference_seconds * 1e6, + fea_seconds / inference_seconds, + ] + ) + + +def write_runtime_plot( + *, + median_fea_seconds: float, + median_inference_seconds: float, +) -> None: + """Plot median FEA and surrogate runtimes.""" + + RUNTIME_FIGURE_PATH.parent.mkdir( + parents=True, + exist_ok=True, + ) + + labels = [ + "CalculiX\nstatic + modal", + "PyTorch\nsurrogate", + ] + + values = [ + median_fea_seconds, + median_inference_seconds, + ] + + plt.figure(figsize=(7, 5)) + + plt.bar( + labels, + values, + ) + + plt.yscale("log") + + plt.ylabel("Median runtime per design [s, log scale]") + + plt.title("FEA vs Surrogate Runtime") + + plt.tight_layout() + + plt.savefig( + RUNTIME_FIGURE_PATH, + dpi=200, + ) + + plt.close() + + +def load_training_runtime_summary() -> dict[int, float]: + """Read median ML training time for each dataset size.""" + + if not SAMPLE_EFFICIENCY_PATH.exists(): + raise FileNotFoundError( + "Sample-efficiency results were not found at " + f"{SAMPLE_EFFICIENCY_PATH}. " + "Run scripts/run_ml_tradeoff_experiment.py first." + ) + + dataframe = pd.read_csv(SAMPLE_EFFICIENCY_PATH) + + required_columns = { + "training_samples", + "training_seconds", + } + + missing_columns = required_columns - set(dataframe.columns) + + if missing_columns: + missing = ", ".join(sorted(missing_columns)) + + raise ValueError(f"Sample-efficiency CSV is missing: {missing}") + + grouped = dataframe.groupby( + "training_samples", + sort=True, + )["training_seconds"].median() + + return { + int(sample_count): float(training_seconds) + for ( + sample_count, + training_seconds, + ) in grouped.items() + } + + +def write_break_even_outputs( + *, + median_fea_seconds: float, + median_inference_seconds: float, +) -> None: + """Calculate full-cost and sunk-data break-even estimates.""" + + training_runtime_by_size = load_training_runtime_summary() + + rows: list[ + tuple[ + int, + float, + float, + float, + float, + float, + float, + ] + ] = [] + + runtime_saving_per_query = median_fea_seconds - median_inference_seconds + + if runtime_saving_per_query <= 0.0: + raise ValueError("Measured surrogate inference is not faster than FEA.") + + for ( + training_samples, + training_seconds, + ) in sorted(training_runtime_by_size.items()): + tradeoff = calculate_runtime_tradeoff( + fea_seconds_per_design=(median_fea_seconds), + inference_seconds_per_design=(median_inference_seconds), + training_samples=(training_samples), + model_training_seconds=(training_seconds), + ) + + incremental_break_even_queries = training_seconds / runtime_saving_per_query + + rows.append( + ( + training_samples, + training_seconds, + tradeoff.fea_seconds_per_design, + tradeoff.inference_seconds_per_design, + tradeoff.speedup, + tradeoff.break_even_queries, + incremental_break_even_queries, + ) + ) + + BREAK_EVEN_CSV_PATH.parent.mkdir( + parents=True, + exist_ok=True, + ) + + with BREAK_EVEN_CSV_PATH.open( + "w", + encoding="utf-8", + newline="", + ) as stream: + writer = csv.writer(stream) + + writer.writerow( + [ + "training_samples", + "median_model_training_seconds", + "median_fea_seconds_per_design", + "median_surrogate_seconds_per_design", + "speedup", + ("sequential_equivalent_break_even_queries"), + ("incremental_break_even_queries_if_fea_data_is_sunk"), + ] + ) + + writer.writerows(rows) + + write_break_even_plot(rows) + + +def write_break_even_plot( + rows: list[ + tuple[ + int, + float, + float, + float, + float, + float, + float, + ] + ], +) -> None: + """Plot break-even query count versus training-set size.""" + + BREAK_EVEN_FIGURE_PATH.parent.mkdir( + parents=True, + exist_ok=True, + ) + + training_sizes = [row[0] for row in rows] + + full_break_even = [row[5] for row in rows] + + incremental_break_even = [row[6] for row in rows] + + plt.figure(figsize=(8, 5)) + + plt.plot( + training_sizes, + full_break_even, + marker="o", + label=("Full FEA-data + training cost"), + ) + + plt.plot( + training_sizes, + incremental_break_even, + marker="o", + label=("Training-only cost (FEA data already available)"), + ) + + plt.xlabel("Training samples") + + plt.ylabel("Break-even future queries") + + plt.title("Surrogate Computational Break-Even") + + plt.grid(True) + + plt.legend() + + plt.tight_layout() + + plt.savefig( + BREAK_EVEN_FIGURE_PATH, + dpi=200, + ) + + plt.close() + + +def main() -> None: + print() + print("BodySimPy Surrogate Runtime Benchmark") + print("=" * 70) + + ( + model, + feature_standardization, + target_standardization, + ) = load_surrogate_checkpoint() + + designs = generate_design_points( + sample_count=BENCHMARK_DESIGNS, + seed=BENCHMARK_SEED, + ) + + print() + print(f"Generated {len(designs)} fresh benchmark designs.") + + fea_durations = benchmark_fea(designs) + + inference_durations = benchmark_surrogate( + model, + designs, + feature_standardization=(feature_standardization), + target_standardization=(target_standardization), + ) + + write_runtime_csv( + fea_durations, + inference_durations, + ) + + median_fea_seconds = float(median(fea_durations)) + + median_inference_seconds = float(median(inference_durations)) + + speedup = median_fea_seconds / median_inference_seconds + + write_runtime_plot( + median_fea_seconds=(median_fea_seconds), + median_inference_seconds=(median_inference_seconds), + ) + + write_break_even_outputs( + median_fea_seconds=(median_fea_seconds), + median_inference_seconds=(median_inference_seconds), + ) + + print() + print("Runtime summary") + print("-" * 70) + + print(f"Median FEA/design: {median_fea_seconds:.6f} s") + + print( + f"Median surrogate/design: " + f"{median_inference_seconds:.9f} s " + f"(" + f"{median_inference_seconds * 1e6:.3f} " + f"microseconds)" + ) + + print(f"Measured speedup: {speedup:,.1f}x") + + print() + print("Generated artifacts") + print("-" * 70) + + print(RUNTIME_CSV_PATH) + + print(BREAK_EVEN_CSV_PATH) + + print(RUNTIME_FIGURE_PATH) + + print(BREAK_EVEN_FIGURE_PATH) + + print() + print("Break-even values are sequential-equivalent estimates.") + + print("Full-cost break-even includes FEA training-data generation.") + + print("Incremental break-even treats the FEA dataset as already available.") + + +if __name__ == "__main__": + main() diff --git a/scripts/generate_surrogate_holdout.py b/scripts/generate_surrogate_holdout.py new file mode 100644 index 0000000..3a38801 --- /dev/null +++ b/scripts/generate_surrogate_holdout.py @@ -0,0 +1,56 @@ +from dataclasses import replace +from pathlib import Path + +from bodysimpy.config.loader import load_config +from bodysimpy.ml.dataset_generation import ( + generate_fea_dataset, + write_fea_dataset, +) +from bodysimpy.ml.design_space import ( + generate_design_points, +) +from bodysimpy.modeling.crossmember import ( + build_crossmember_model, +) + +SAMPLE_COUNT = 200 +RANDOM_SEED = 20260815 +MAX_WORKERS = 4 + + +def main() -> None: + config = load_config("configs/baseline_crossmember.yaml") + + base_model = build_crossmember_model(config) + + holdout_model = replace( + base_model, + name=(f"{base_model.name}_ml_holdout"), + ) + + designs = generate_design_points( + sample_count=SAMPLE_COUNT, + seed=RANDOM_SEED, + ) + + samples = generate_fea_dataset( + holdout_model, + designs, + max_workers=MAX_WORKERS, + ) + + output_path = Path("data/surrogate/fea_surrogate_holdout.csv") + + write_fea_dataset( + samples, + output_path, + ) + + print() + print(f"Generated {len(samples)} independent holdout samples.") + + print(output_path) + + +if __name__ == "__main__": + main() diff --git a/scripts/run_ml_tradeoff_experiment.py b/scripts/run_ml_tradeoff_experiment.py new file mode 100644 index 0000000..9cd4e6f --- /dev/null +++ b/scripts/run_ml_tradeoff_experiment.py @@ -0,0 +1,226 @@ +import csv +from pathlib import Path +from time import perf_counter + +import matplotlib.pyplot as plt +import numpy as np + +from bodysimpy.ml.dataset import ( + load_surrogate_arrays, +) +from bodysimpy.ml.evaluation import ( + calculate_metrics, + predict, +) +from bodysimpy.ml.experiment import ( + build_nested_training_subsets, + split_holdout_indices, +) +from bodysimpy.ml.training import ( + train_surrogate_fixed_validation, +) + +TRAINING_SIZES = ( + 50, + 100, + 250, + 500, +) + +TRAINING_SEEDS = ( + 11, + 22, + 33, +) + + +def main() -> None: + training_features, training_targets = load_surrogate_arrays( + "data/surrogate/fea_surrogate_dataset.csv" + ) + + holdout_features, holdout_targets = load_surrogate_arrays( + "data/surrogate/fea_surrogate_holdout.csv" + ) + + subsets = build_nested_training_subsets( + pool_size=training_features.shape[0], + sample_counts=TRAINING_SIZES, + seed=42, + ) + + holdout_split = split_holdout_indices( + sample_count=holdout_features.shape[0], + validation_count=100, + seed=42, + ) + + validation_features = holdout_features[holdout_split.validation_indices] + + validation_targets = holdout_targets[holdout_split.validation_indices] + + test_features = holdout_features[holdout_split.test_indices] + + test_targets = holdout_targets[holdout_split.test_indices] + + rows: list[dict[str, float | int]] = [] + + for subset in subsets: + for training_seed in TRAINING_SEEDS: + print(f"Training size={subset.sample_count}, seed={training_seed}") + + train_features = training_features[subset.indices] + + train_targets = training_targets[subset.indices] + + start = perf_counter() + + result = train_surrogate_fixed_validation( + train_features, + train_targets, + validation_features, + validation_targets, + seed=training_seed, + ) + + training_seconds = perf_counter() - start + + predictions = predict( + result.model, + test_features, + feature_standardization=(result.feature_standardization), + target_standardization=(result.target_standardization), + ) + + stress_metrics = calculate_metrics( + test_targets[:, 0], + predictions[:, 0], + ) + + displacement_metrics = calculate_metrics( + test_targets[:, 1], + predictions[:, 1], + ) + + frequency_metrics = calculate_metrics( + test_targets[:, 2], + predictions[:, 2], + ) + + rows.append( + { + "training_samples": (subset.sample_count), + "training_seed": (training_seed), + "best_epoch": (result.best_epoch), + "training_seconds": (training_seconds), + "stress_mae_mpa": (stress_metrics.mae / 1e6), + "stress_mape_percent": (stress_metrics.mean_absolute_percentage_error), + "deflection_mae_mm": (displacement_metrics.mae * 1000.0), + "deflection_mape_percent": ( + displacement_metrics.mean_absolute_percentage_error + ), + "frequency_mae_hz": (frequency_metrics.mae), + "frequency_mape_percent": (frequency_metrics.mean_absolute_percentage_error), + } + ) + + output_path = Path("docs/validation/ml_sample_efficiency.csv") + + output_path.parent.mkdir( + parents=True, + exist_ok=True, + ) + + with output_path.open( + "w", + encoding="utf-8", + newline="", + ) as stream: + writer = csv.DictWriter( + stream, + fieldnames=list(rows[0].keys()), + ) + + writer.writeheader() + writer.writerows(rows) + + plot_accuracy(rows) + + +def plot_accuracy( + rows: list[dict[str, float | int]], +) -> None: + figure_path = Path("docs/figures/ml_accuracy_vs_training_size.png") + + figure_path.parent.mkdir( + parents=True, + exist_ok=True, + ) + + metric_fields = ( + ( + "Stress", + "stress_mape_percent", + ), + ( + "Deflection", + "deflection_mape_percent", + ), + ( + "Mode-1 frequency", + "frequency_mape_percent", + ), + ) + + plt.figure(figsize=(8, 5)) + + for label, field in metric_fields: + means: list[float] = [] + deviations: list[float] = [] + + for sample_count in TRAINING_SIZES: + values = np.array( + [float(row[field]) for row in rows if int(row["training_samples"]) == sample_count], + dtype=float, + ) + + means.append(float(np.mean(values))) + + deviations.append( + float( + np.std( + values, + ddof=1, + ) + ) + ) + + plt.errorbar( + TRAINING_SIZES, + means, + yerr=deviations, + marker="o", + capsize=4, + label=label, + ) + + plt.xlabel("Number of FEA training samples") + + plt.ylabel("Held-out test MAPE [%]") + + plt.title("Surrogate Accuracy vs Training-Data Size") + + plt.legend() + plt.grid(True) + plt.tight_layout() + + plt.savefig( + figure_path, + dpi=200, + ) + + plt.close() + + +if __name__ == "__main__": + main() diff --git a/src/bodysimpy/ml/experiment.py b/src/bodysimpy/ml/experiment.py new file mode 100644 index 0000000..3972017 --- /dev/null +++ b/src/bodysimpy/ml/experiment.py @@ -0,0 +1,127 @@ +from dataclasses import dataclass + +import numpy as np +from numpy.typing import NDArray + + +@dataclass(frozen=True, slots=True) +class TrainingSubset: + sample_count: int + indices: NDArray[np.int64] + + +@dataclass(frozen=True, slots=True) +class HoldoutSplit: + validation_indices: NDArray[np.int64] + test_indices: NDArray[np.int64] + + +@dataclass(frozen=True, slots=True) +class RuntimeTradeoff: + fea_seconds_per_design: float + inference_seconds_per_design: float + speedup: float + break_even_queries: float + + +def build_nested_training_subsets( + *, + pool_size: int, + sample_counts: tuple[int, ...], + seed: int, +) -> tuple[TrainingSubset, ...]: + """Create deterministic nested subsets from one training pool.""" + + if pool_size <= 0: + raise ValueError("Training pool size must be positive.") + + if not sample_counts: + raise ValueError("At least one training sample count is required.") + + if any(count <= 0 for count in sample_counts): + raise ValueError("Training sample counts must be positive.") + + if any(count > pool_size for count in sample_counts): + raise ValueError("Training sample count exceeds the available pool.") + + if tuple(sorted(sample_counts)) != sample_counts: + raise ValueError("Training sample counts must be increasing.") + + rng = np.random.default_rng(seed) + + order = rng.permutation(pool_size).astype(np.int64) + + return tuple( + TrainingSubset( + sample_count=count, + indices=order[:count].copy(), + ) + for count in sample_counts + ) + + +def split_holdout_indices( + *, + sample_count: int, + validation_count: int, + seed: int, +) -> HoldoutSplit: + """Split an independent holdout dataset into validation and test sets.""" + + if sample_count <= 1: + raise ValueError("Holdout dataset must contain at least two samples.") + + if validation_count <= 0 or validation_count >= sample_count: + raise ValueError("Validation count must lie inside the holdout sample range.") + + rng = np.random.default_rng(seed) + + order = rng.permutation(sample_count).astype(np.int64) + + return HoldoutSplit( + validation_indices=(order[:validation_count].copy()), + test_indices=(order[validation_count:].copy()), + ) + + +def calculate_runtime_tradeoff( + *, + fea_seconds_per_design: float, + inference_seconds_per_design: float, + training_samples: int, + model_training_seconds: float, +) -> RuntimeTradeoff: + """Calculate surrogate speedup and sequential break-even point.""" + + if fea_seconds_per_design <= 0.0: + raise ValueError("FEA runtime must be positive.") + + if inference_seconds_per_design <= 0.0: + raise ValueError("Inference runtime must be positive.") + + if training_samples <= 0: + raise ValueError("Training sample count must be positive.") + + if model_training_seconds < 0.0: + raise ValueError("Model training runtime cannot be negative.") + + runtime_saving = fea_seconds_per_design - inference_seconds_per_design + + if runtime_saving <= 0.0: + raise ValueError( + "Surrogate inference must be faster than FEA " + "for a positive computational break-even point." + ) + + speedup = fea_seconds_per_design / inference_seconds_per_design + + estimated_upfront_seconds = training_samples * fea_seconds_per_design + model_training_seconds + + break_even_queries = estimated_upfront_seconds / runtime_saving + + return RuntimeTradeoff( + fea_seconds_per_design=(fea_seconds_per_design), + inference_seconds_per_design=(inference_seconds_per_design), + speedup=speedup, + break_even_queries=break_even_queries, + ) diff --git a/src/bodysimpy/ml/model.py b/src/bodysimpy/ml/model.py index 8e8dba6..2b61af9 100644 --- a/src/bodysimpy/ml/model.py +++ b/src/bodysimpy/ml/model.py @@ -26,4 +26,4 @@ def forward( return cast( Tensor, self.network(features), - ) \ No newline at end of file + ) diff --git a/src/bodysimpy/ml/training.py b/src/bodysimpy/ml/training.py index 818322c..905a02a 100644 --- a/src/bodysimpy/ml/training.py +++ b/src/bodysimpy/ml/training.py @@ -36,6 +36,16 @@ class TrainingResult: split: DataSplit +@dataclass(frozen=True, slots=True) +class FixedValidationTrainingResult: + model: StructuralSurrogate + feature_standardization: Standardization + target_standardization: Standardization + train_losses: tuple[float, ...] + validation_losses: tuple[float, ...] + best_epoch: int + + def create_data_split( sample_count: int, *, @@ -259,3 +269,158 @@ def save_training_checkpoint( }, checkpoint_path, ) + + +def train_surrogate_fixed_validation( + train_features: NDArray[np.float64], + train_targets: NDArray[np.float64], + validation_features: NDArray[np.float64], + validation_targets: NDArray[np.float64], + *, + seed: int, + batch_size: int = 32, + learning_rate: float = 1e-3, + maximum_epochs: int = 1000, + patience: int = 50, +) -> FixedValidationTrainingResult: + """Train using an explicitly fixed validation dataset.""" + + torch.manual_seed(seed) + + feature_statistics = fit_standardization(train_features) + + target_statistics = fit_standardization(train_targets) + + normalized_train_features = standardize( + train_features, + feature_statistics, + ) + + normalized_train_targets = standardize( + train_targets, + target_statistics, + ) + + normalized_validation_features = standardize( + validation_features, + feature_statistics, + ) + + normalized_validation_targets = standardize( + validation_targets, + target_statistics, + ) + + train_dataset = StructuralSurrogateDataset( + normalized_train_features, + normalized_train_targets, + ) + + validation_dataset = StructuralSurrogateDataset( + normalized_validation_features, + normalized_validation_targets, + ) + + training_loader: DataLoader[tuple[Tensor, Tensor]] = DataLoader( + train_dataset, + batch_size=batch_size, + shuffle=True, + ) + + validation_loader: DataLoader[tuple[Tensor, Tensor]] = DataLoader( + validation_dataset, + batch_size=batch_size, + shuffle=False, + ) + + device = torch.device("cuda" if torch.cuda.is_available() else "cpu") + + model = StructuralSurrogate().to(device) + + optimizer = Adam( + model.parameters(), + lr=learning_rate, + ) + + loss_function = nn.MSELoss() + + train_losses: list[float] = [] + validation_losses: list[float] = [] + + best_validation_loss = float("inf") + best_model_state = deepcopy(model.state_dict()) + + best_epoch = 0 + epochs_without_improvement = 0 + + for epoch in range( + 1, + maximum_epochs + 1, + ): + model.train() + + batch_losses: list[float] = [] + + for ( + features_batch, + targets_batch, + ) in training_loader: + features_batch = features_batch.to(device) + + targets_batch = targets_batch.to(device) + + optimizer.zero_grad() + + predictions = model(features_batch) + + loss = loss_function( + predictions, + targets_batch, + ) + + loss.backward() + optimizer.step() + + batch_losses.append(float(loss.item())) + + train_loss = float(np.mean(batch_losses)) + + validation_loss = _mean_loss( + model, + validation_loader, + loss_function, + device, + ) + + train_losses.append(train_loss) + + validation_losses.append(validation_loss) + + if validation_loss < best_validation_loss: + best_validation_loss = validation_loss + + best_model_state = deepcopy(model.state_dict()) + + best_epoch = epoch + + epochs_without_improvement = 0 + + else: + epochs_without_improvement += 1 + + if epochs_without_improvement >= patience: + break + + model.load_state_dict(best_model_state) + + model.to("cpu") + model.eval() + + return FixedValidationTrainingResult( + model=model, + feature_standardization=(feature_statistics), + target_standardization=(target_statistics), + train_losses=tuple(train_losses), + validation_losses=tuple(validation_losses), + best_epoch=best_epoch, + ) diff --git a/tests/unit/test_ml_experiment.py b/tests/unit/test_ml_experiment.py new file mode 100644 index 0000000..2ad3c42 --- /dev/null +++ b/tests/unit/test_ml_experiment.py @@ -0,0 +1,70 @@ +import pytest + +from bodysimpy.ml.experiment import ( + build_nested_training_subsets, + calculate_runtime_tradeoff, + split_holdout_indices, +) + + +def test_training_subsets_are_nested() -> None: + subsets = build_nested_training_subsets( + pool_size=500, + sample_counts=(50, 100, 250, 500), + seed=42, + ) + + assert tuple(subset.sample_count for subset in subsets) == ( + 50, + 100, + 250, + 500, + ) + + first = set(subsets[0].indices.tolist()) + second = set(subsets[1].indices.tolist()) + third = set(subsets[2].indices.tolist()) + fourth = set(subsets[3].indices.tolist()) + + assert first < second + assert second < third + assert third < fourth + + +def test_holdout_validation_and_test_are_disjoint() -> None: + split = split_holdout_indices( + sample_count=200, + validation_count=100, + seed=42, + ) + + validation = set(split.validation_indices.tolist()) + + test = set(split.test_indices.tolist()) + + assert len(validation) == 100 + assert len(test) == 100 + assert validation.isdisjoint(test) + + +def test_runtime_tradeoff() -> None: + result = calculate_runtime_tradeoff( + fea_seconds_per_design=10.0, + inference_seconds_per_design=0.01, + training_samples=100, + model_training_seconds=5.0, + ) + + assert result.speedup == pytest.approx(1000.0) + + assert result.break_even_queries > 100.0 + + +def test_runtime_tradeoff_rejects_slower_surrogate() -> None: + with pytest.raises(ValueError): + calculate_runtime_tradeoff( + fea_seconds_per_design=1.0, + inference_seconds_per_design=2.0, + training_samples=100, + model_training_seconds=5.0, + )