Skip to content

Commit 504576e

Browse files
committed
Refined comparisons of models to evaluate ensembles.
1 parent 0463bdb commit 504576e

2 files changed

Lines changed: 14 additions & 2 deletions

File tree

git_reset.sh

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,2 +1,3 @@
11
#! /bin/bash
22
git reset --hard HEAD
3+
git pull

scripts/run_polynomial_evaluation.py

Lines changed: 13 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -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

123134
if __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

Comments
 (0)