| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -15,7 +15,7 @@ public interface ILayer: IWithTrackable, IKerasConfigable | |||
| 15 | 15 | List<ILayer> Layers { get; } | |
| 16 | 16 | List<INode> InboundNodes { get; } | |
| 17 | 17 | List<INode> OutboundNodes { get; } | |
| 18 | - Tensors Apply(Tensors inputs, Tensors states = null, bool training = false, IOptionalArgs? optional_args = null); | ||
| 18 | + Tensors Apply(Tensors inputs, Tensors states = null, bool? training = false, IOptionalArgs? optional_args = null); | ||
| 19 | 19 | List<IVariableV1> TrainableVariables { get; } | |
| 20 | 20 | List<IVariableV1> TrainableWeights { get; } | |
| 21 | 21 | List<IVariableV1> NonTrainableWeights { get; } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -145,7 +145,7 @@ private Tensor _zero_state_tensors(object state_size, Tensor batch_size, TF_Data | |||
| 145 | 145 | throw new NotImplementedException("_zero_state_tensors"); | |
| 146 | 146 | } | |
| 147 | 147 | ||
| 148 | - public Tensors Apply(Tensors inputs, Tensors state = null, bool is_training = false, IOptionalArgs? optional_args = null) | ||
| 148 | + public Tensors Apply(Tensors inputs, Tensors state = null, bool? is_training = false, IOptionalArgs? optional_args = null) | ||
| 149 | 149 | { | |
| 150 | 150 | throw new NotImplementedException(); | |
| 151 | 151 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -13,7 +13,7 @@ public partial class Layer | |||
| 13 | 13 | /// <param name="state"></param> | |
| 14 | 14 | /// <param name="training"></param> | |
| 15 | 15 | /// <returns></returns> | |
| 16 | - public virtual Tensors Apply(Tensors inputs, Tensors states = null, bool training = false, IOptionalArgs? optional_args = null) | ||
| 16 | + public virtual Tensors Apply(Tensors inputs, Tensors states = null, bool? training = false, IOptionalArgs? optional_args = null) | ||
| 17 | 17 | { | |
| 18 | 18 | if (callContext.Value == null) | |
| 19 | 19 | callContext.Value = new CallContext(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -142,6 +142,7 @@ public History fit(IDatasetV2 dataset, | |||
| 142 | 142 | int verbose = 1, | |
| 143 | 143 | List<ICallback> callbacks = null, | |
| 144 | 144 | IDatasetV2 validation_data = null, | |
| 145 | + int validation_step = 10, // 间隔多少次会进行一次验证 | ||
| 145 | 146 | bool shuffle = true, | |
| 146 | 147 | int initial_epoch = 0, | |
| 147 | 148 | int max_queue_size = 10, | |
@@ -164,11 +165,11 @@ public History fit(IDatasetV2 dataset, | |||
| 164 | 165 | }); | |
| 165 | 166 | ||
| 166 | 167 | ||
| 167 | - return FitInternal(data_handler, epochs, verbose, callbacks, validation_data: validation_data, | ||
| 168 | + return FitInternal(data_handler, epochs, validation_step, verbose, callbacks, validation_data: validation_data, | ||
| 168 | 169 | train_step_func: train_step_function); | |
| 169 | 170 | } | |
| 170 | 171 | ||
| 171 | - History FitInternal(DataHandler data_handler, int epochs, int verbose, List<ICallback> callbackList, IDatasetV2 validation_data, | ||
| 172 | + History FitInternal(DataHandler data_handler, int epochs, int validation_step, int verbose, List<ICallback> callbackList, IDatasetV2 validation_data, | ||
| 172 | 173 | Func<DataHandler, OwnedIterator, Dictionary<string, float>> train_step_func) | |
| 173 | 174 | { | |
| 174 | 175 | stop_training = false; | |
@@ -207,6 +208,9 @@ History FitInternal(DataHandler data_handler, int epochs, int verbose, List<ICal | |||
| 207 | 208 | ||
| 208 | 209 | if (validation_data != null) | |
| 209 | 210 | { | |
| 211 | + if (validation_step > 0 && epoch ==0 || (epoch) % validation_step != 0) | ||
| 212 | + continue; | ||
| 213 | + | ||
| 210 | 214 | var val_logs = evaluate(validation_data); | |
| 211 | 215 | foreach(var log in val_logs) | |
| 212 | 216 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -393,7 +393,7 @@ protected override Tensors Call(Tensors inputs, Tensors initial_state = null, bo | |||
| 393 | 393 | } | |
| 394 | 394 | } | |
| 395 | 395 | ||
| 396 | - public override Tensors Apply(Tensors inputs, Tensors initial_states = null, bool training = false, IOptionalArgs? optional_args = null) | ||
| 396 | + public override Tensors Apply(Tensors inputs, Tensors initial_states = null, bool? training = false, IOptionalArgs? optional_args = null) | ||
| 397 | 397 | { | |
| 398 | 398 | RnnOptionalArgs? rnn_optional_args = optional_args as RnnOptionalArgs; | |
| 399 | 399 | if (optional_args is not null && rnn_optional_args is null) | |
| Back | FazBrowse Home | New Git URL |
0 commit comments