| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent d3f19f4 commit 400cde2
6 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -3,6 +3,7 @@ | |||
| 3 | 3 | using System.Collections.Generic; | |
| 4 | 4 | using System.Linq; | |
| 5 | 5 | using Tensorflow.Framework.Models; | |
| 6 | + using static Tensorflow.Binding; | ||
| 6 | 7 | ||
| 7 | 8 | namespace Tensorflow | |
| 8 | 9 | { | |
@@ -98,6 +99,20 @@ public IDatasetV2 apply_options() | |||
| 98 | 99 | return dataset; | |
| 99 | 100 | } | |
| 100 | 101 | ||
| 102 | + public Tensor dataset_cardinality(string name = null) | ||
| 103 | + { | ||
| 104 | + if (tf.Context.executing_eagerly()) | ||
| 105 | + { | ||
| 106 | + var results = tf.Runner.TFE_FastPathExecute(tf.Context, tf.Context.DeviceName, | ||
| 107 | + "DatasetCardinality", name, | ||
| 108 | + null, | ||
| 109 | + variant_tensor); | ||
| 110 | + return results[0]; | ||
| 111 | + } | ||
| 112 | + | ||
| 113 | + throw new NotImplementedException(""); | ||
| 114 | + } | ||
| 115 | + | ||
| 101 | 116 | public override string ToString() | |
| 102 | 117 | => $"{GetType().Name} shapes: {string.Join(", ", structure.Select(x => x.shape))}, types: {string.Join(", ", structure.Select(x => "tf." + x.dtype.as_numpy_name()))}"; | |
| 103 | 118 | ||
@@ -117,7 +132,9 @@ public override string ToString() | |||
| 117 | 132 | break; | |
| 118 | 133 | } | |
| 119 | 134 | ||
| 120 | - yield return (results[0], results.Length == 1 ? null : results[1]); | ||
| 135 | + yield return results.Length == 2 | ||
| 136 | + ? (results[0], results[1]) | ||
| 137 | + : (null, results[0]); | ||
| 121 | 138 | } | |
| 122 | 139 | } | |
| 123 | 140 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -74,5 +74,7 @@ IDatasetV2 map(Func<Tensor, (Tensor, Tensor), (Tensor, Tensor)> map_func, | |||
| 74 | 74 | /// </summary> | |
| 75 | 75 | /// <returns></returns> | |
| 76 | 76 | IDatasetV2 apply_options(); | |
| 77 | + | ||
| 78 | + Tensor dataset_cardinality(string name = null); | ||
| 77 | 79 | } | |
| 78 | 80 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -15,11 +15,11 @@ public MapDataset(IDatasetV2 input_dataset, | |||
| 15 | 15 | bool preserve_cardinality = false, | |
| 16 | 16 | bool use_legacy_function = false) : base(input_dataset) | |
| 17 | 17 | { | |
| 18 | - var func = new ConcreteFunction($"autograph_{map_func.Method.Name}"); | ||
| 19 | - var input = tf.placeholder(input_dataset.element_spec[0].dtype, name: "input"); | ||
| 18 | + using var func = new ConcreteFunction($"autograph_{map_func.Method.Name}"); | ||
| 19 | + var input = tf.placeholder(input_dataset.element_spec[0].dtype); | ||
| 20 | 20 | var output = map_func(input); | |
| 21 | 21 | func.ToGraph(input, output); | |
| 22 | - | ||
| 22 | + | ||
| 23 | 23 | structure = func.OutputStructure; | |
| 24 | 24 | ||
| 25 | 25 | variant_tensor = ops.map_dataset(input_dataset.variant_tensor, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -130,6 +130,9 @@ ForwardBackwardCall SelectForwardAndBackwardFunctions(Tensors args, int possible | |||
| 130 | 130 | return new ForwardBackwardCall(functions, args, tape_watching: true); | |
| 131 | 131 | } | |
| 132 | 132 | ||
| 133 | + public override string ToString() | ||
| 134 | + => Name; | ||
| 135 | + | ||
| 133 | 136 | public void Dispose() | |
| 134 | 137 | { | |
| 135 | 138 | c_api.TFE_ContextRemoveFunction(tf.Context.Handle, Name, tf.Status.Handle); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -2,10 +2,11 @@ | |||
| 2 | 2 | ||
| 3 | 3 | namespace Tensorflow.Keras.ArgsDefinition | |
| 4 | 4 | { | |
| 5 | - public class TensorLikeDataAdapterArgs | ||
| 5 | + public class DataAdapterArgs | ||
| 6 | 6 | { | |
| 7 | 7 | public Tensor X { get; set; } | |
| 8 | 8 | public Tensor Y { get; set; } | |
| 9 | + public IDatasetV2 Dataset { get; set; } | ||
| 9 | 10 | public int BatchSize { get; set; } = 32; | |
| 10 | 11 | public int Steps { get; set; } | |
| 11 | 12 | public int Epochs { get; set; } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -6,6 +6,7 @@ public class DataHandlerArgs | |||
| 6 | 6 | { | |
| 7 | 7 | public Tensor X { get; set; } | |
| 8 | 8 | public Tensor Y { get; set; } | |
| 9 | + public IDatasetV2 Dataset { get; set; } | ||
| 9 | 10 | public int BatchSize { get; set; } = 32; | |
| 10 | 11 | public int StepsPerEpoch { get; set; } = -1; | |
| 11 | 12 | public int InitialEpoch { get; set; } = 0; | |
| Back | FazBrowse Home | New Git URL |
0 commit comments