| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 5e4f530 commit baf620a
4 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -24,6 +24,7 @@ ICallback fit(NDArray x, NDArray y, | |||
| 24 | 24 | List<ICallback> callbacks = null, | |
| 25 | 25 | float validation_split = 0f, | |
| 26 | 26 | ValidationDataPack validation_data = null, | |
| 27 | + int validation_step = 10, | ||
| 27 | 28 | bool shuffle = true, | |
| 28 | 29 | Dictionary<int, float> class_weight = null, | |
| 29 | 30 | NDArray sample_weight = null, | |
@@ -47,6 +48,20 @@ ICallback fit(IEnumerable<NDArray> x, NDArray y, | |||
| 47 | 48 | int workers = 1, | |
| 48 | 49 | bool use_multiprocessing = false); | |
| 49 | 50 | ||
| 51 | + public ICallback fit(IDatasetV2 dataset, | ||
| 52 | + int batch_size = -1, | ||
| 53 | + int epochs = 1, | ||
| 54 | + int verbose = 1, | ||
| 55 | + List<ICallback> callbacks = null, | ||
| 56 | + IDatasetV2 validation_data = null, | ||
| 57 | + int validation_step = 10, // 间隔多少次会进行一次验证 | ||
| 58 | + bool shuffle = true, | ||
| 59 | + Dictionary<int, float> class_weight = null, | ||
| 60 | + int initial_epoch = 0, | ||
| 61 | + int max_queue_size = 10, | ||
| 62 | + int workers = 1, | ||
| 63 | + bool use_multiprocessing = false); | ||
| 64 | + | ||
| 50 | 65 | void save(string filepath, | |
| 51 | 66 | bool overwrite = true, | |
| 52 | 67 | bool include_optimizer = true, | |
@@ -85,6 +100,14 @@ Tensors predict(Tensors x, | |||
| 85 | 100 | int workers = 1, | |
| 86 | 101 | bool use_multiprocessing = false); | |
| 87 | 102 | ||
| 103 | + public Tensors predict(IDatasetV2 dataset, | ||
| 104 | + int batch_size = -1, | ||
| 105 | + int verbose = 0, | ||
| 106 | + int steps = -1, | ||
| 107 | + int max_queue_size = 10, | ||
| 108 | + int workers = 1, | ||
| 109 | + bool use_multiprocessing = false); | ||
| 110 | + | ||
| 88 | 111 | void summary(int line_length = -1, float[] positions = null); | |
| 89 | 112 | ||
| 90 | 113 | IKerasConfig get_config(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -132,6 +132,7 @@ Dictionary<string, float> evaluate(DataHandler data_handler, CallbackList callba | |||
| 132 | 132 | var end_step = step + data_handler.StepIncrement; | |
| 133 | 133 | if (!is_val) | |
| 134 | 134 | callbacks.on_test_batch_end(end_step, logs); | |
| 135 | + GC.Collect(); | ||
| 135 | 136 | } | |
| 136 | 137 | } | |
| 137 | 138 | callbacks.on_test_end(logs); | |
@@ -167,7 +168,9 @@ Dictionary<string, float> test_step_multi_inputs_function(DataHandler data_handl | |||
| 167 | 168 | Dictionary<string, float> test_step(DataHandler data_handler, Tensors x, Tensors y) | |
| 168 | 169 | { | |
| 169 | 170 | (x,y) = data_handler.DataAdapter.Expand1d(x, y); | |
| 171 | + | ||
| 170 | 172 | var y_pred = Apply(x, training: false); | |
| 173 | + | ||
| 171 | 174 | var loss = compiled_loss.Call(y, y_pred); | |
| 172 | 175 | compiled_metrics.update_state(y, y_pred); | |
| 173 | 176 | return metrics.Select(x => (x.Name, x.result())).ToDictionary(x => x.Item1, x => (float)x.Item2); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -41,6 +41,7 @@ public ICallback fit(NDArray x, NDArray y, | |||
| 41 | 41 | List<ICallback> callbacks = null, | |
| 42 | 42 | float validation_split = 0f, | |
| 43 | 43 | ValidationDataPack validation_data = null, | |
| 44 | + int validation_step = 10, | ||
| 44 | 45 | bool shuffle = true, | |
| 45 | 46 | Dictionary<int, float> class_weight = null, | |
| 46 | 47 | NDArray sample_weight = null, | |
@@ -147,7 +148,7 @@ public ICallback fit(IEnumerable<NDArray> x, NDArray y, | |||
| 147 | 148 | } | |
| 148 | 149 | } | |
| 149 | 150 | ||
| 150 | - public History fit(IDatasetV2 dataset, | ||
| 151 | + public ICallback fit(IDatasetV2 dataset, | ||
| 151 | 152 | int batch_size = -1, | |
| 152 | 153 | int epochs = 1, | |
| 153 | 154 | int verbose = 1, | |
@@ -156,7 +157,6 @@ public History fit(IDatasetV2 dataset, | |||
| 156 | 157 | int validation_step = 10, | |
| 157 | 158 | bool shuffle = true, | |
| 158 | 159 | Dictionary<int, float> class_weight = null, | |
| 159 | - NDArray sample_weight = null, | ||
| 160 | 160 | int initial_epoch = 0, | |
| 161 | 161 | int max_queue_size = 10, | |
| 162 | 162 | int workers = 1, | |
@@ -170,7 +170,7 @@ public History fit(IDatasetV2 dataset, | |||
| 170 | 170 | InitialEpoch = initial_epoch, | |
| 171 | 171 | Epochs = epochs, | |
| 172 | 172 | Shuffle = shuffle, | |
| 173 | - SampleWeight = sample_weight, | ||
| 173 | + ClassWeight = class_weight, | ||
| 174 | 174 | MaxQueueSize = max_queue_size, | |
| 175 | 175 | Workers = workers, | |
| 176 | 176 | UseMultiprocessing = use_multiprocessing, | |
@@ -218,6 +218,7 @@ History FitInternal(DataHandler data_handler, int epochs, int validation_step, i | |||
| 218 | 218 | var end_step = step + data_handler.StepIncrement; | |
| 219 | 219 | End_step = end_step; | |
| 220 | 220 | callbacks.on_train_batch_end(end_step, logs); | |
| 221 | + GC.Collect(); | ||
| 221 | 222 | } | |
| 222 | 223 | ||
| 223 | 224 | if (validation_data != null) | |
@@ -233,11 +234,10 @@ History FitInternal(DataHandler data_handler, int epochs, int validation_step, i | |||
| 233 | 234 | callbacks.on_train_batch_end(End_step, logs); | |
| 234 | 235 | } | |
| 235 | 236 | ||
| 237 | + GC.Collect(); | ||
| 236 | 238 | ||
| 237 | 239 | callbacks.on_epoch_end(epoch, logs); | |
| 238 | 240 | ||
| 239 | - GC.Collect(); | ||
| 240 | - GC.WaitForPendingFinalizers(); | ||
| 241 | 241 | if (stop_training) | |
| 242 | 242 | { | |
| 243 | 243 | break; | |
@@ -282,6 +282,7 @@ History FitInternal(DataHandler data_handler, int epochs, int verbose, List<ICal | |||
| 282 | 282 | var end_step = step + data_handler.StepIncrement; | |
| 283 | 283 | End_step = end_step; | |
| 284 | 284 | callbacks.on_train_batch_end(end_step, logs); | |
| 285 | + GC.Collect(); | ||
| 285 | 286 | } | |
| 286 | 287 | ||
| 287 | 288 | if (validation_data != null) | |
@@ -301,7 +302,6 @@ History FitInternal(DataHandler data_handler, int epochs, int verbose, List<ICal | |||
| 301 | 302 | callbacks.on_epoch_end(epoch, logs); | |
| 302 | 303 | ||
| 303 | 304 | GC.Collect(); | |
| 304 | - GC.WaitForPendingFinalizers(); | ||
| 305 | 305 | if (stop_training) | |
| 306 | 306 | { | |
| 307 | 307 | break; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -102,9 +102,9 @@ Tensors PredictInternal(DataHandler data_handler, int verbose) | |||
| 102 | 102 | for (int i = 0; i < batch_outputs.Length; i++) | |
| 103 | 103 | batch_outputs[i] = tf.concat(new Tensor[] { batch_outputs[i], tmp_batch_outputs[i] }, axis: 0); | |
| 104 | 104 | } | |
| 105 | - | ||
| 106 | 105 | var end_step = step + data_handler.StepIncrement; | |
| 107 | 106 | callbacks.on_predict_batch_end(end_step, new Dictionary<string, Tensors> { { "outputs", batch_outputs } }); | |
| 107 | + GC.Collect(); | ||
| 108 | 108 | } | |
| 109 | 109 | } | |
| 110 | 110 | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments