@@ -111,18 +111,29 @@ def compareModels(num_model: int = 10, num_sample: int = 1000):
111111 num_dependent_feature = NUM_DEPENDENT_FEATURE , num_value = 10 ,
112112 is_new_multiplier = False )
113113 calculators = []
114+ max_err_arrs = []
114115 for idx in range (num_model ):
115116 model = makeModel ()
116117 runner = ModelRunnerNN .deserialize (model , MULTI_SERIALIZE_PATH % idx )
117118 _ , max_err_ser = runner .makeRelativeError (test_dl )
119+ max_err_arrs .append (np .reshape (max_err_ser .values , (- 1 , 1 )))
118120 calculator = AccuracyCalculator (cast (np .ndarray , max_err_ser .values ))
119121 calculators .append (calculator )
120122 # Plot the CDF of the errors from all calculators
121123 AccuracyCalculator .plotCDFComparison (calculators , is_plot = True )
124+ # Plot the "oracle" ensemble
125+ max_err_arr = np .hstack (max_err_arrs )
126+ abs_oracle_err_arr = np .min (np .abs (max_err_arr ), axis = 1 )
127+ nonabs_oracle_err_arr = np .min (max_err_arr , axis = 1 )
128+ pos_sel = abs_oracle_err_arr != nonabs_oracle_err_arr
129+ oracle_err_arr = abs_oracle_err_arr .copy ()
130+ oracle_err_arr [pos_sel ] *= - 1
131+ oracle_calculator = AccuracyCalculator (oracle_err_arr )
132+ AccuracyCalculator .plotCDFComparison ([oracle_calculator ], is_plot = True )
122133
123134if __name__ == "__main__" :
124135 #train(num_epoch=50)
125136 #evaluate()
126137 num_model = 10
127- makeModels (num_model = num_model , num_sample = 1000 , num_epoch = 7000 )
128- compareModels (num_model = num_model )
138+ # makeModels(num_model=num_model, num_sample=1000, num_epoch=7000)
139+ compareModels (num_model = 7 , num_sample = 100000 )
0 commit comments