| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
2 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -85,5 +85,11 @@ public static NDArray dot(NDArray x1, NDArray x2, NDArray? axes = null, string? | |||
| 85 | 85 | ||
| 86 | 86 | [AutoNumPy] | |
| 87 | 87 | public static NDArray add(NDArray x, NDArray y) => new NDArray(math_ops.add(x, y)); | |
| 88 | + | ||
| 89 | + [AutoNumPy] | ||
| 90 | + public static NDArray greater(NDArray x, NDArray y) => new NDArray(tf.greater(x, y)); | ||
| 91 | + | ||
| 92 | + [AutoNumPy] | ||
| 93 | + public static NDArray less(NDArray x, NDArray y) => new NDArray(tf.less(x, y)); | ||
| 88 | 94 | } | |
| 89 | 95 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -19,8 +19,10 @@ public class EarlyStopping: ICallback | |||
| 19 | 19 | string _monitor; | |
| 20 | 20 | string _mode; | |
| 21 | 21 | bool _restore_best_weights; | |
| 22 | - List<IVariableV1>? _best_weights; | ||
| 22 | + List<NDArray>? _best_weights; | ||
| 23 | 23 | CallbackParams _parameters; | |
| 24 | + Func<NDArray, NDArray, NDArray> _monitor_op; | ||
| 25 | + | ||
| 24 | 26 | public Dictionary<string, List<float>>? history { get; set; } | |
| 25 | 27 | // user need to pass a CallbackParams to EarlyStopping, CallbackParams at least need the model | |
| 26 | 28 | public EarlyStopping(CallbackParams parameters,string monitor = "val_loss", float min_delta = 0f, int patience = 0, | |
@@ -38,17 +40,49 @@ public EarlyStopping(CallbackParams parameters,string monitor = "val_loss", floa | |||
| 38 | 40 | _min_delta = Math.Abs(min_delta); | |
| 39 | 41 | _restore_best_weights = restore_best_weights; | |
| 40 | 42 | _mode = mode; | |
| 41 | - if (mode != "auto" && mode != "min" && mode != "max") | ||
| 43 | + | ||
| 44 | + if (_mode != "auto" && _mode != "min" && _mode != "max") | ||
| 45 | + { | ||
| 46 | + Console.WriteLine($"EarlyStopping mode {_mode} is unknown, fallback to auto mode."); | ||
| 47 | + _mode = "auto"; | ||
| 48 | + } | ||
| 49 | + | ||
| 50 | + if (_mode == "min") | ||
| 51 | + { | ||
| 52 | + _monitor_op = np.less; | ||
| 53 | + } | ||
| 54 | + else if (_mode == "max") | ||
| 55 | + { | ||
| 56 | + _monitor_op = np.greater; | ||
| 57 | + } | ||
| 58 | + else | ||
| 59 | + { | ||
| 60 | + if (_monitor.EndsWith("acc") || _monitor.EndsWith("accuracy") || _monitor.EndsWith("auc")) | ||
| 61 | + { | ||
| 62 | + _monitor_op = np.greater; | ||
| 63 | + } | ||
| 64 | + else | ||
| 65 | + { | ||
| 66 | + _monitor_op = np.less; | ||
| 67 | + } | ||
| 68 | + } | ||
| 69 | + | ||
| 70 | + if (_monitor_op == np.greater) | ||
| 42 | 71 | { | |
| 43 | - Console.WriteLine("EarlyStopping mode %s is unknown, fallback to auto mode.", mode); | ||
| 72 | + _min_delta *= 1; | ||
| 73 | + } | ||
| 74 | + else | ||
| 75 | + { | ||
| 76 | + _min_delta *= -1; | ||
| 44 | 77 | } | |
| 45 | 78 | } | |
| 46 | 79 | public void on_train_begin() | |
| 47 | 80 | { | |
| 48 | 81 | _wait = 0; | |
| 49 | 82 | _stopped_epoch = 0; | |
| 83 | + _best = _monitor_op == np.less ? (float)np.Inf : (float)-np.Inf; | ||
| 84 | + _best_weights = null; | ||
| 50 | 85 | _best_epoch = 0; | |
| 51 | - _best = (float)np.Inf; | ||
| 52 | 86 | } | |
| 53 | 87 | ||
| 54 | 88 | public void on_epoch_begin(int epoch) | |
@@ -74,7 +108,7 @@ public void on_epoch_end(int epoch, Dictionary<string, float> epoch_logs) | |||
| 74 | 108 | // Restore the weights after first epoch if no progress is ever made. | |
| 75 | 109 | if (_restore_best_weights && _best_weights == null) | |
| 76 | 110 | { | |
| 77 | - _best_weights = _parameters.Model.Weights; | ||
| 111 | + _best_weights = _parameters.Model.get_weights(); | ||
| 78 | 112 | } | |
| 79 | 113 | _wait += 1; | |
| 80 | 114 | ||
@@ -83,7 +117,7 @@ public void on_epoch_end(int epoch, Dictionary<string, float> epoch_logs) | |||
| 83 | 117 | _best = current; | |
| 84 | 118 | _best_epoch = epoch; | |
| 85 | 119 | if (_restore_best_weights) | |
| 86 | - _best_weights = _parameters.Model.TrainableWeights; | ||
| 120 | + _best_weights = _parameters.Model.get_weights(); | ||
| 87 | 121 | // Only restart wait if we beat both the baseline and our previous best. | |
| 88 | 122 | if (_baseline == 0f || _is_improvement(current, _baseline)) | |
| 89 | 123 | _wait = 0; | |
@@ -99,7 +133,7 @@ public void on_epoch_end(int epoch, Dictionary<string, float> epoch_logs) | |||
| 99 | 133 | { | |
| 100 | 134 | Console.WriteLine($"Restoring model weights from the end of the best epoch: {_best_epoch + 1}"); | |
| 101 | 135 | } | |
| 102 | - _parameters.Model.Weights = _best_weights; | ||
| 136 | + _parameters.Model.set_weights(_best_weights); | ||
| 103 | 137 | } | |
| 104 | 138 | } | |
| 105 | 139 | } | |
@@ -131,21 +165,7 @@ float get_monitor_value(Dictionary<string, float> logs) | |||
| 131 | 165 | } | |
| 132 | 166 | public bool _is_improvement(float monitor_value, float reference_value) | |
| 133 | 167 | { | |
| 134 | - bool less_op = (monitor_value - _min_delta) < reference_value; | ||
| 135 | - bool greater_op = (monitor_value - _min_delta) >= reference_value; | ||
| 136 | - if (_mode == "min") | ||
| 137 | - return less_op; | ||
| 138 | - else if (_mode == "max") | ||
| 139 | - return greater_op; | ||
| 140 | - else | ||
| 141 | - { | ||
| 142 | - if (_monitor.EndsWith("acc") || _monitor.EndsWith("accuracy") || _monitor.EndsWith("auc")) | ||
| 143 | - { | ||
| 144 | - return greater_op; | ||
| 145 | - } | ||
| 146 | - else | ||
| 147 | - return less_op; | ||
| 148 | - } | ||
| 168 | + return _monitor_op(monitor_value - _min_delta, reference_value); | ||
| 149 | 169 | } | |
| 150 | 170 | ||
| 151 | 171 | public void on_test_end(Dictionary<string, float> logs) | |
| Back | FazBrowse Home | New Git URL |
0 commit comments