| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent bdf229a commit 89fe0bb
5 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -12,4 +12,6 @@ public interface ICallback | |||
| 12 | 12 | void on_predict_batch_begin(long step); | |
| 13 | 13 | void on_predict_batch_end(long end_step, Dictionary<string, Tensors> logs); | |
| 14 | 14 | void on_predict_end(); | |
| 15 | + void on_test_begin(); | ||
| 16 | + void on_test_batch_end(long end_step, IEnumerable<(string, Tensor)> logs); | ||
| 15 | 17 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -20,7 +20,10 @@ public void on_train_begin() | |||
| 20 | 20 | { | |
| 21 | 21 | callbacks.ForEach(x => x.on_train_begin()); | |
| 22 | 22 | } | |
| 23 | - | ||
| 23 | + public void on_test_begin() | ||
| 24 | + { | ||
| 25 | + callbacks.ForEach(x => x.on_test_begin()); | ||
| 26 | + } | ||
| 24 | 27 | public void on_epoch_begin(int epoch) | |
| 25 | 28 | { | |
| 26 | 29 | callbacks.ForEach(x => x.on_epoch_begin(epoch)); | |
@@ -60,4 +63,13 @@ public void on_predict_end() | |||
| 60 | 63 | { | |
| 61 | 64 | callbacks.ForEach(x => x.on_predict_end()); | |
| 62 | 65 | } | |
| 66 | + | ||
| 67 | + public void on_test_batch_begin(long step) | ||
| 68 | + { | ||
| 69 | + callbacks.ForEach(x => x.on_train_batch_begin(step)); | ||
| 70 | + } | ||
| 71 | + public void on_test_batch_end(long end_step, IEnumerable<(string, Tensor)> logs) | ||
| 72 | + { | ||
| 73 | + callbacks.ForEach(x => x.on_test_batch_end(end_step, logs)); | ||
| 74 | + } | ||
| 63 | 75 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -18,15 +18,19 @@ public void on_train_begin() | |||
| 18 | 18 | epochs = new List<int>(); | |
| 19 | 19 | history = new Dictionary<string, List<float>>(); | |
| 20 | 20 | } | |
| 21 | - | ||
| 21 | + public void on_test_begin() | ||
| 22 | + { | ||
| 23 | + epochs = new List<int>(); | ||
| 24 | + history = new Dictionary<string, List<float>>(); | ||
| 25 | + } | ||
| 22 | 26 | public void on_epoch_begin(int epoch) | |
| 23 | 27 | { | |
| 24 | 28 | ||
| 25 | 29 | } | |
| 26 | 30 | ||
| 27 | 31 | public void on_train_batch_begin(long step) | |
| 28 | 32 | { | |
| 29 | - | ||
| 33 | + | ||
| 30 | 34 | } | |
| 31 | 35 | ||
| 32 | 36 | public void on_train_batch_end(long end_step, Dictionary<string, float> logs) | |
@@ -55,16 +59,25 @@ public void on_predict_begin() | |||
| 55 | 59 | ||
| 56 | 60 | public void on_predict_batch_begin(long step) | |
| 57 | 61 | { | |
| 58 | - | ||
| 62 | + | ||
| 59 | 63 | } | |
| 60 | 64 | ||
| 61 | 65 | public void on_predict_batch_end(long end_step, Dictionary<string, Tensors> logs) | |
| 62 | 66 | { | |
| 63 | - | ||
| 67 | + | ||
| 64 | 68 | } | |
| 65 | 69 | ||
| 66 | 70 | public void on_predict_end() | |
| 67 | 71 | { | |
| 68 | - | ||
| 72 | + | ||
| 73 | + } | ||
| 74 | + | ||
| 75 | + public void on_test_batch_begin(long step) | ||
| 76 | + { | ||
| 77 | + | ||
| 78 | + } | ||
| 79 | + | ||
| 80 | + public void on_test_batch_end(long end_step, IEnumerable<(string, Tensor)> logs) | ||
| 81 | + { | ||
| 69 | 82 | } | |
| 70 | 83 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -22,7 +22,10 @@ public void on_train_begin() | |||
| 22 | 22 | _called_in_fit = true; | |
| 23 | 23 | _sw = new Stopwatch(); | |
| 24 | 24 | } | |
| 25 | - | ||
| 25 | + public void on_test_begin() | ||
| 26 | + { | ||
| 27 | + _sw = new Stopwatch(); | ||
| 28 | + } | ||
| 26 | 29 | public void on_epoch_begin(int epoch) | |
| 27 | 30 | { | |
| 28 | 31 | _reset_progbar(); | |
@@ -44,7 +47,7 @@ public void on_train_batch_end(long end_step, Dictionary<string, float> logs) | |||
| 44 | 47 | var progress = ""; | |
| 45 | 48 | var length = 30.0 / _parameters.Steps; | |
| 46 | 49 | for (int i = 0; i < Math.Floor(end_step * length - 1); i++) | |
| 47 | - progress += "="; | ||
| 50 | + progress += "="; | ||
| 48 | 51 | if (progress.Length < 28) | |
| 49 | 52 | progress += ">"; | |
| 50 | 53 | else | |
@@ -84,17 +87,35 @@ public void on_predict_begin() | |||
| 84 | 87 | ||
| 85 | 88 | public void on_predict_batch_begin(long step) | |
| 86 | 89 | { | |
| 87 | - | ||
| 90 | + | ||
| 88 | 91 | } | |
| 89 | 92 | ||
| 90 | 93 | public void on_predict_batch_end(long end_step, Dictionary<string, Tensors> logs) | |
| 91 | 94 | { | |
| 92 | - | ||
| 95 | + | ||
| 93 | 96 | } | |
| 94 | 97 | ||
| 95 | 98 | public void on_predict_end() | |
| 96 | 99 | { | |
| 97 | - | ||
| 100 | + | ||
| 101 | + } | ||
| 102 | + | ||
| 103 | + public void on_test_batch_begin(long step) | ||
| 104 | + { | ||
| 105 | + _sw.Restart(); | ||
| 98 | 106 | } | |
| 107 | + public void on_test_batch_end(long end_step, IEnumerable<(string, Tensor)> logs) | ||
| 108 | + { | ||
| 109 | + _sw.Stop(); | ||
| 110 | + var elapse = _sw.ElapsedMilliseconds; | ||
| 111 | + var results = string.Join(" - ", logs.Select(x => $"{x.Item1}: {(float)x.Item2.numpy():F6}")); | ||
| 112 | + | ||
| 113 | + Binding.tf_output_redirect.Write($"{end_step + 1:D4}/{_parameters.Steps:D4} - {elapse}ms/step - {results}"); | ||
| 114 | + if (!Console.IsOutputRedirected) | ||
| 115 | + { | ||
| 116 | + Console.CursorLeft = 0; | ||
| 117 | + } | ||
| 118 | + } | ||
| 119 | + | ||
| 99 | 120 | } | |
| 100 | 121 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -5,6 +5,10 @@ | |||
| 5 | 5 | using Tensorflow.Keras.ArgsDefinition; | |
| 6 | 6 | using Tensorflow.Keras.Engine.DataAdapters; | |
| 7 | 7 | using static Tensorflow.Binding; | |
| 8 | + using Tensorflow.Keras.Layers; | ||
| 9 | + using Tensorflow.Keras.Utils; | ||
| 10 | + using Tensorflow; | ||
| 11 | + using Tensorflow.Keras.Callbacks; | ||
| 8 | 12 | ||
| 9 | 13 | namespace Tensorflow.Keras.Engine | |
| 10 | 14 | { | |
@@ -31,6 +35,11 @@ public void evaluate(NDArray x, NDArray y, | |||
| 31 | 35 | bool use_multiprocessing = false, | |
| 32 | 36 | bool return_dict = false) | |
| 33 | 37 | { | |
| 38 | + if (x.dims[0] != y.dims[0]) | ||
| 39 | + { | ||
| 40 | + throw new InvalidArgumentError( | ||
| 41 | + $"The array x and y should have same value at dim 0, but got {x.dims[0]} and {y.dims[0]}"); | ||
| 42 | + } | ||
| 34 | 43 | var data_handler = new DataHandler(new DataHandlerArgs | |
| 35 | 44 | { | |
| 36 | 45 | X = x, | |
@@ -46,18 +55,31 @@ public void evaluate(NDArray x, NDArray y, | |||
| 46 | 55 | StepsPerExecution = _steps_per_execution | |
| 47 | 56 | }); | |
| 48 | 57 | ||
| 58 | + var callbacks = new CallbackList(new CallbackParams | ||
| 59 | + { | ||
| 60 | + Model = this, | ||
| 61 | + Verbose = verbose, | ||
| 62 | + Steps = data_handler.Inferredsteps | ||
| 63 | + }); | ||
| 64 | + callbacks.on_test_begin(); | ||
| 65 | + | ||
| 49 | 66 | foreach (var (epoch, iterator) in data_handler.enumerate_epochs()) | |
| 50 | 67 | { | |
| 51 | 68 | reset_metrics(); | |
| 52 | - // callbacks.on_epoch_begin(epoch) | ||
| 69 | + //callbacks.on_epoch_begin(epoch); | ||
| 53 | 70 | // data_handler.catch_stop_iteration(); | |
| 54 | - IEnumerable<(string, Tensor)> results = null; | ||
| 71 | + IEnumerable<(string, Tensor)> logs = null; | ||
| 72 | + | ||
| 55 | 73 | foreach (var step in data_handler.steps()) | |
| 56 | 74 | { | |
| 57 | - // callbacks.on_train_batch_begin(step) | ||
| 58 | - results = test_function(data_handler, iterator); | ||
| 75 | + callbacks.on_train_batch_begin(step); | ||
| 76 | + logs = test_function(data_handler, iterator); | ||
| 77 | + var end_step = step + data_handler.StepIncrement; | ||
| 78 | + callbacks.on_test_batch_end(end_step, logs); | ||
| 59 | 79 | } | |
| 60 | 80 | } | |
| 81 | + GC.Collect(); | ||
| 82 | + GC.WaitForPendingFinalizers(); | ||
| 61 | 83 | } | |
| 62 | 84 | ||
| 63 | 85 | public KeyValuePair<string, float>[] evaluate(IDatasetV2 x) | |
@@ -75,7 +97,8 @@ public KeyValuePair<string, float>[] evaluate(IDatasetV2 x) | |||
| 75 | 97 | reset_metrics(); | |
| 76 | 98 | // callbacks.on_epoch_begin(epoch) | |
| 77 | 99 | // data_handler.catch_stop_iteration(); | |
| 78 | - | ||
| 100 | + | ||
| 101 | + | ||
| 79 | 102 | foreach (var step in data_handler.steps()) | |
| 80 | 103 | { | |
| 81 | 104 | // callbacks.on_train_batch_begin(step) | |
| Back | FazBrowse Home | New Git URL |
0 commit comments