| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,4 +1,5 @@ | |||
| 1 | - using Tensorflow.NumPy; | ||
| 1 | + using OneOf; | ||
| 2 | + using Tensorflow.NumPy; | ||
| 2 | 3 | ||
| 3 | 4 | namespace Tensorflow.Util | |
| 4 | 5 | { | |
@@ -8,10 +9,10 @@ namespace Tensorflow.Util | |||
| 8 | 9 | /// </summary> | |
| 9 | 10 | public class ValidationDataPack | |
| 10 | 11 | { | |
| 11 | - public NDArray val_x; | ||
| 12 | - public NDArray val_y; | ||
| 13 | - public NDArray val_sample_weight = null; | ||
| 14 | - | ||
| 12 | + internal OneOf<NDArray, NDArray[]> val_x; | ||
| 13 | + internal NDArray val_y; | ||
| 14 | + internal NDArray val_sample_weight = null; | ||
| 15 | + public bool val_x_is_array = false; | ||
| 15 | 16 | public ValidationDataPack((NDArray, NDArray) validation_data) | |
| 16 | 17 | { | |
| 17 | 18 | this.val_x = validation_data.Item1; | |
@@ -27,15 +28,17 @@ public ValidationDataPack((NDArray, NDArray, NDArray) validation_data) | |||
| 27 | 28 | ||
| 28 | 29 | public ValidationDataPack((IEnumerable<NDArray>, NDArray) validation_data) | |
| 29 | 30 | { | |
| 30 | - this.val_x = validation_data.Item1.ToArray()[0]; | ||
| 31 | + this.val_x = validation_data.Item1.ToArray(); | ||
| 31 | 32 | this.val_y = validation_data.Item2; | |
| 33 | + val_x_is_array = true; | ||
| 32 | 34 | } | |
| 33 | 35 | ||
| 34 | 36 | public ValidationDataPack((IEnumerable<NDArray>, NDArray, NDArray) validation_data) | |
| 35 | 37 | { | |
| 36 | - this.val_x = validation_data.Item1.ToArray()[0]; | ||
| 38 | + this.val_x = validation_data.Item1.ToArray(); | ||
| 37 | 39 | this.val_y = validation_data.Item2; | |
| 38 | 40 | this.val_sample_weight = validation_data.Item3; | |
| 41 | + val_x_is_array = true; | ||
| 39 | 42 | } | |
| 40 | 43 | ||
| 41 | 44 | public static implicit operator ValidationDataPack((NDArray, NDArray) validation_data) | |
@@ -52,15 +55,24 @@ public static implicit operator ValidationDataPack((IEnumerable<NDArray>, NDArra | |||
| 52 | 55 | ||
| 53 | 56 | public void Deconstruct(out NDArray val_x, out NDArray val_y) | |
| 54 | 57 | { | |
| 55 | - val_x = this.val_x; | ||
| 58 | + val_x = this.val_x.AsT0; | ||
| 56 | 59 | val_y = this.val_y; | |
| 57 | 60 | } | |
| 58 | 61 | ||
| 59 | 62 | public void Deconstruct(out NDArray val_x, out NDArray val_y, out NDArray val_sample_weight) | |
| 60 | 63 | { | |
| 61 | - val_x = this.val_x; | ||
| 64 | + val_x = this.val_x.AsT0; | ||
| 65 | + val_y = this.val_y; | ||
| 66 | + val_sample_weight = this.val_sample_weight; | ||
| 67 | + } | ||
| 68 | + | ||
| 69 | + // add a unuse parameter to make it different from Deconstruct(out NDArray val_x, out NDArray val_y, out NDArray val_sample_weight) | ||
| 70 | + public void Deconstruct(out NDArray[] val_x_array, out NDArray val_y, out NDArray val_sample_weight, out NDArray unuse) | ||
| 71 | + { | ||
| 72 | + val_x_array = this.val_x.AsT1; | ||
| 62 | 73 | val_y = this.val_y; | |
| 63 | 74 | val_sample_weight = this.val_sample_weight; | |
| 75 | + unuse = null; | ||
| 64 | 76 | } | |
| 65 | 77 | } | |
| 66 | 78 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -92,9 +92,17 @@ public static ((IEnumerable<NDArray>, NDArray, NDArray), ValidationDataPack) tra | |||
| 92 | 92 | var train_y = y[new Slice(0, train_count)]; | |
| 93 | 93 | var val_x = x.Select(x => x[new Slice(train_count)] as NDArray); | |
| 94 | 94 | var val_y = y[new Slice(train_count)]; | |
| 95 | - NDArray tmp_sample_weight = sample_weight; | ||
| 96 | - sample_weight = sample_weight[new Slice(0, train_count)]; | ||
| 97 | - ValidationDataPack validation_data = (val_x, val_y, tmp_sample_weight[new Slice(train_count)]); | ||
| 95 | + | ||
| 96 | + ValidationDataPack validation_data; | ||
| 97 | + if (sample_weight != null) | ||
| 98 | + { | ||
| 99 | + validation_data = (val_x, val_y, sample_weight[new Slice(train_count)]); | ||
| 100 | + sample_weight = sample_weight[new Slice(0, train_count)]; | ||
| 101 | + } | ||
| 102 | + else | ||
| 103 | + { | ||
| 104 | + validation_data = (val_x, val_y); | ||
| 105 | + } | ||
| 98 | 106 | return ((train_x, train_y, sample_weight), validation_data); | |
| 99 | 107 | } | |
| 100 | 108 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -70,13 +70,19 @@ public Dictionary<string, float> evaluate(NDArray x, NDArray y, | |||
| 70 | 70 | return evaluate(data_handler, callbacks, is_val, test_function); | |
| 71 | 71 | } | |
| 72 | 72 | ||
| 73 | - public Dictionary<string, float> evaluate(IEnumerable<Tensor> x, Tensor y, int verbose = 1, bool is_val = false) | ||
| 73 | + public Dictionary<string, float> evaluate( | ||
| 74 | + IEnumerable<Tensor> x, | ||
| 75 | + Tensor y, | ||
| 76 | + int verbose = 1, | ||
| 77 | + NDArray sample_weight = null, | ||
| 78 | + bool is_val = false) | ||
| 74 | 79 | { | |
| 75 | 80 | var data_handler = new DataHandler(new DataHandlerArgs | |
| 76 | 81 | { | |
| 77 | 82 | X = new Tensors(x.ToArray()), | |
| 78 | 83 | Y = y, | |
| 79 | 84 | Model = this, | |
| 85 | + SampleWeight = sample_weight, | ||
| 80 | 86 | StepsPerExecution = _steps_per_execution | |
| 81 | 87 | }); | |
| 82 | 88 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -7,6 +7,7 @@ | |||
| 7 | 7 | using System.Diagnostics; | |
| 8 | 8 | using Tensorflow.Keras.Callbacks; | |
| 9 | 9 | using Tensorflow.Util; | |
| 10 | + using OneOf; | ||
| 10 | 11 | ||
| 11 | 12 | namespace Tensorflow.Keras.Engine | |
| 12 | 13 | { | |
@@ -287,10 +288,24 @@ History FitInternal(DataHandler data_handler, int epochs, int verbose, List<ICal | |||
| 287 | 288 | ||
| 288 | 289 | if (validation_data != null) | |
| 289 | 290 | { | |
| 290 | - // Because evaluate calls call_test_batch_end, this interferes with our output on the screen | ||
| 291 | - // so we need to pass a is_val parameter to stop on_test_batch_end | ||
| 292 | - var (val_x, val_y, val_sample_weight) = validation_data; | ||
| 293 | - var val_logs = evaluate(val_x, val_y, sample_weight:val_sample_weight, is_val:true); | ||
| 291 | + NDArray val_x; | ||
| 292 | + NDArray[] val_x_array; | ||
| 293 | + NDArray val_y; | ||
| 294 | + NDArray val_sample_weight; | ||
| 295 | + Dictionary<string, float> val_logs; | ||
| 296 | + if (!validation_data.val_x_is_array) | ||
| 297 | + { | ||
| 298 | + (val_x, val_y, val_sample_weight) = validation_data; | ||
| 299 | + // Because evaluate calls call_test_batch_end, this interferes with our output on the screen | ||
| 300 | + // so we need to pass a is_val parameter to stop on_test_batch_end | ||
| 301 | + val_logs = evaluate(val_x, val_y, sample_weight: val_sample_weight, is_val: true); | ||
| 302 | + | ||
| 303 | + } | ||
| 304 | + else | ||
| 305 | + { | ||
| 306 | + (val_x_array, val_y, val_sample_weight, _) = validation_data; | ||
| 307 | + val_logs = evaluate(val_x_array, val_y, sample_weight: val_sample_weight, is_val: true); | ||
| 308 | + } | ||
| 294 | 309 | foreach (var log in val_logs) | |
| 295 | 310 | { | |
| 296 | 311 | logs["val_" + log.Key] = log.Value; | |
| Back | FazBrowse Home | New Git URL |
0 commit comments