| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent ec340ee commit 271dcef
6 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -293,7 +293,8 @@ protected override void DisposeUnmanagedResources(IntPtr handle) | |||
| 293 | 293 | // c_api.TF_CloseSession(handle, tf.Status.Handle); | |
| 294 | 294 | if (tf.Status == null || tf.Status.Handle.IsInvalid) | |
| 295 | 295 | { | |
| 296 | - c_api.TF_DeleteSession(handle, c_api.TF_NewStatus()); | ||
| 296 | + using var status = new Status(); | ||
| 297 | + c_api.TF_DeleteSession(handle, status.Handle); | ||
| 297 | 298 | } | |
| 298 | 299 | else | |
| 299 | 300 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -39,5 +39,25 @@ public void on_epoch_end(int epoch, Dictionary<string, float> epoch_logs) | |||
| 39 | 39 | { | |
| 40 | 40 | callbacks.ForEach(x => x.on_epoch_end(epoch, epoch_logs)); | |
| 41 | 41 | } | |
| 42 | + | ||
| 43 | + public void on_predict_begin() | ||
| 44 | + { | ||
| 45 | + callbacks.ForEach(x => x.on_predict_begin()); | ||
| 46 | + } | ||
| 47 | + | ||
| 48 | + public void on_predict_batch_begin(long step) | ||
| 49 | + { | ||
| 50 | + callbacks.ForEach(x => x.on_predict_batch_begin(step)); | ||
| 51 | + } | ||
| 52 | + | ||
| 53 | + public void on_predict_batch_end(long end_step, Dictionary<string, Tensors> logs) | ||
| 54 | + { | ||
| 55 | + callbacks.ForEach(x => x.on_predict_batch_end(end_step, logs)); | ||
| 56 | + } | ||
| 57 | + | ||
| 58 | + public void on_predict_end() | ||
| 59 | + { | ||
| 60 | + callbacks.ForEach(x => x.on_predict_end()); | ||
| 61 | + } | ||
| 42 | 62 | } | |
| 43 | 63 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -48,5 +48,26 @@ public void on_epoch_end(int epoch, Dictionary<string, float> epoch_logs) | |||
| 48 | 48 | history[log.Key].Add((float)log.Value); | |
| 49 | 49 | } | |
| 50 | 50 | } | |
| 51 | + | ||
| 52 | + public void on_predict_begin() | ||
| 53 | + { | ||
| 54 | + epochs = new List<int>(); | ||
| 55 | + history = new Dictionary<string, List<float>>(); | ||
| 56 | + } | ||
| 57 | + | ||
| 58 | + public void on_predict_batch_begin(long step) | ||
| 59 | + { | ||
| 60 | + | ||
| 61 | + } | ||
| 62 | + | ||
| 63 | + public void on_predict_batch_end(long end_step, Dictionary<string, Tensors> logs) | ||
| 64 | + { | ||
| 65 | + | ||
| 66 | + } | ||
| 67 | + | ||
| 68 | + public void on_predict_end() | ||
| 69 | + { | ||
| 70 | + | ||
| 71 | + } | ||
| 51 | 72 | } | |
| 52 | 73 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -11,5 +11,9 @@ public interface ICallback | |||
| 11 | 11 | void on_train_batch_begin(long step); | |
| 12 | 12 | void on_train_batch_end(long end_step, Dictionary<string, float> logs); | |
| 13 | 13 | void on_epoch_end(int epoch, Dictionary<string, float> epoch_logs); | |
| 14 | + void on_predict_begin(); | ||
| 15 | + void on_predict_batch_begin(long step); | ||
| 16 | + void on_predict_batch_end(long end_step, Dictionary<string, Tensors> logs); | ||
| 17 | + void on_predict_end(); | ||
| 14 | 18 | } | |
| 15 | 19 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,5 +1,4 @@ | |||
| 1 | - using PureHDF; | ||
| 2 | - using System; | ||
| 1 | + using System; | ||
| 3 | 2 | using System.Collections.Generic; | |
| 4 | 3 | using System.Diagnostics; | |
| 5 | 4 | using System.Linq; | |
@@ -77,5 +76,26 @@ void _maybe_init_progbar() | |||
| 77 | 76 | { | |
| 78 | 77 | ||
| 79 | 78 | } | |
| 79 | + | ||
| 80 | + public void on_predict_begin() | ||
| 81 | + { | ||
| 82 | + _reset_progbar(); | ||
| 83 | + _maybe_init_progbar(); | ||
| 84 | + } | ||
| 85 | + | ||
| 86 | + public void on_predict_batch_begin(long step) | ||
| 87 | + { | ||
| 88 | + | ||
| 89 | + } | ||
| 90 | + | ||
| 91 | + public void on_predict_batch_end(long end_step, Dictionary<string, Tensors> logs) | ||
| 92 | + { | ||
| 93 | + | ||
| 94 | + } | ||
| 95 | + | ||
| 96 | + public void on_predict_end() | ||
| 97 | + { | ||
| 98 | + | ||
| 99 | + } | ||
| 80 | 100 | } | |
| 81 | 101 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -5,11 +5,70 @@ | |||
| 5 | 5 | using Tensorflow.Keras.ArgsDefinition; | |
| 6 | 6 | using Tensorflow.Keras.Engine.DataAdapters; | |
| 7 | 7 | using static Tensorflow.Binding; | |
| 8 | + using Tensorflow.Keras.Callbacks; | ||
| 8 | 9 | ||
| 9 | 10 | namespace Tensorflow.Keras.Engine | |
| 10 | 11 | { | |
| 11 | 12 | public partial class Model | |
| 12 | 13 | { | |
| 14 | + public Tensors predict(IDatasetV2 dataset, | ||
| 15 | + int batch_size = -1, | ||
| 16 | + int verbose = 0, | ||
| 17 | + int steps = -1, | ||
| 18 | + int max_queue_size = 10, | ||
| 19 | + int workers = 1, | ||
| 20 | + bool use_multiprocessing = false) | ||
| 21 | + { | ||
| 22 | + var data_handler = new DataHandler(new DataHandlerArgs | ||
| 23 | + { | ||
| 24 | + Dataset = dataset, | ||
| 25 | + BatchSize = batch_size, | ||
| 26 | + StepsPerEpoch = steps, | ||
| 27 | + InitialEpoch = 0, | ||
| 28 | + Epochs = 1, | ||
| 29 | + MaxQueueSize = max_queue_size, | ||
| 30 | + Workers = workers, | ||
| 31 | + UseMultiprocessing = use_multiprocessing, | ||
| 32 | + Model = this, | ||
| 33 | + StepsPerExecution = _steps_per_execution | ||
| 34 | + }); | ||
| 35 | + | ||
| 36 | + var callbacks = new CallbackList(new CallbackParams | ||
| 37 | + { | ||
| 38 | + Model = this, | ||
| 39 | + Verbose = verbose, | ||
| 40 | + Epochs = 1, | ||
| 41 | + Steps = data_handler.Inferredsteps | ||
| 42 | + }); | ||
| 43 | + | ||
| 44 | + Tensor batch_outputs = null; | ||
| 45 | + _predict_counter.assign(0); | ||
| 46 | + callbacks.on_predict_begin(); | ||
| 47 | + foreach (var (epoch, iterator) in data_handler.enumerate_epochs()) | ||
| 48 | + { | ||
| 49 | + foreach (var step in data_handler.steps()) | ||
| 50 | + { | ||
| 51 | + callbacks.on_predict_batch_begin(step); | ||
| 52 | + var tmp_batch_outputs = run_predict_step(iterator); | ||
| 53 | + if (batch_outputs == null) | ||
| 54 | + { | ||
| 55 | + batch_outputs = tmp_batch_outputs[0]; | ||
| 56 | + } | ||
| 57 | + else | ||
| 58 | + { | ||
| 59 | + batch_outputs = tf.concat(new Tensor[] { batch_outputs, tmp_batch_outputs[0] }, axis: 0); | ||
| 60 | + } | ||
| 61 | + | ||
| 62 | + var end_step = step + data_handler.StepIncrement; | ||
| 63 | + callbacks.on_predict_batch_end(end_step, new Dictionary<string, Tensors> { { "outputs", batch_outputs } }); | ||
| 64 | + } | ||
| 65 | + GC.Collect(); | ||
| 66 | + } | ||
| 67 | + | ||
| 68 | + callbacks.on_predict_end(); | ||
| 69 | + return batch_outputs; | ||
| 70 | + } | ||
| 71 | + | ||
| 13 | 72 | /// <summary> | |
| 14 | 73 | /// Generates output predictions for the input samples. | |
| 15 | 74 | /// </summary> | |
| Back | FazBrowse Home | New Git URL |
0 commit comments