| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 7465337 commit 404c803
21 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -19,6 +19,8 @@ public class DatasetV2 : IDatasetV2 | |||
| 19 | 19 | ||
| 20 | 20 | public TensorSpec[] structure { get; set; } | |
| 21 | 21 | ||
| 22 | + public int FirstInputTensorCount { get; set; } = 1; | ||
| 23 | + | ||
| 22 | 24 | public Shape[] output_shapes => structure.Select(x => x.shape).ToArray(); | |
| 23 | 25 | ||
| 24 | 26 | public TF_DataType[] output_types => structure.Select(x => x.dtype).ToArray(); | |
@@ -131,6 +133,7 @@ public IDatasetV2 apply_options() | |||
| 131 | 133 | ||
| 132 | 134 | // (4) Apply stats aggregator options | |
| 133 | 135 | ||
| 136 | + dataset.FirstInputTensorCount = this.FirstInputTensorCount; | ||
| 134 | 137 | return dataset; | |
| 135 | 138 | } | |
| 136 | 139 | ||
@@ -142,7 +145,7 @@ public override string ToString() | |||
| 142 | 145 | $"types: {string.Join(", ", structure.Select(x => "tf." + x.dtype.as_numpy_name()))}, " + | |
| 143 | 146 | $"len: {length}"; | |
| 144 | 147 | ||
| 145 | - public IEnumerator<(Tensor, Tensor)> GetEnumerator() | ||
| 148 | + public IEnumerator<(Tensors, Tensors)> GetEnumerator() | ||
| 146 | 149 | { | |
| 147 | 150 | using var ownedIterator = new OwnedIterator(this); | |
| 148 | 151 | ||
@@ -158,7 +161,8 @@ public override string ToString() | |||
| 158 | 161 | break; | |
| 159 | 162 | } | |
| 160 | 163 | ||
| 161 | - yield return (results[0], results.Length == 1 ? null : results[1]); | ||
| 164 | + yield return (new Tensors(results.Take(FirstInputTensorCount)), results.Length == FirstInputTensorCount ? | ||
| 165 | + null : new Tensors(results.Skip(FirstInputTensorCount))); | ||
| 162 | 166 | } | |
| 163 | 167 | } | |
| 164 | 168 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -4,7 +4,7 @@ | |||
| 4 | 4 | ||
| 5 | 5 | namespace Tensorflow | |
| 6 | 6 | { | |
| 7 | - public interface IDatasetV2 : IEnumerable<(Tensor, Tensor)> | ||
| 7 | + public interface IDatasetV2 : IEnumerable<(Tensors, Tensors)> | ||
| 8 | 8 | { | |
| 9 | 9 | string[] class_names { get; set; } | |
| 10 | 10 | ||
@@ -18,6 +18,8 @@ public interface IDatasetV2 : IEnumerable<(Tensor, Tensor)> | |||
| 18 | 18 | ||
| 19 | 19 | TensorSpec[] structure { get; set; } | |
| 20 | 20 | ||
| 21 | + int FirstInputTensorCount { get; set; } | ||
| 22 | + | ||
| 21 | 23 | /// <summary> | |
| 22 | 24 | /// Caches the elements in this dataset. | |
| 23 | 25 | /// </summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -27,7 +27,8 @@ void _create_iterator(IDatasetV2 dataset) | |||
| 27 | 27 | _dataset = dataset; | |
| 28 | 28 | _element_spec = dataset.element_spec; | |
| 29 | 29 | // _flat_output_types = | |
| 30 | - (_iterator_resource, _deleter) = ops.anonymous_iterator_v2(_dataset.output_types, _dataset.output_shapes); | ||
| 30 | + _iterator_resource = ops.anonymous_iterator_v3(_dataset.output_types, _dataset.output_shapes); | ||
| 31 | + // TODO(Rinne): deal with graph mode. | ||
| 31 | 32 | ops.make_iterator(dataset.variant_tensor, _iterator_resource); | |
| 32 | 33 | } | |
| 33 | 34 | ||
@@ -48,7 +49,7 @@ public Tensor[] next() | |||
| 48 | 49 | ||
| 49 | 50 | public void Dispose() | |
| 50 | 51 | { | |
| 51 | - tf.Runner.Execute(tf.Context, "DeleteIterator", 0, new[] { _iterator_resource, _deleter }, null); | ||
| 52 | + //tf.Runner.Execute(tf.Context, "DeleteIterator", 0, new[] { _iterator_resource, _deleter }, null); | ||
| 52 | 53 | } | |
| 53 | 54 | } | |
| 54 | 55 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -5,8 +5,8 @@ namespace Tensorflow.Keras.ArgsDefinition | |||
| 5 | 5 | { | |
| 6 | 6 | public class DataAdapterArgs: IKerasConfig | |
| 7 | 7 | { | |
| 8 | - public Tensor X { get; set; } | ||
| 9 | - public Tensor Y { get; set; } | ||
| 8 | + public Tensors X { get; set; } | ||
| 9 | + public Tensors Y { get; set; } | ||
| 10 | 10 | public IDatasetV2 Dataset { get; set; } | |
| 11 | 11 | public int BatchSize { get; set; } = 32; | |
| 12 | 12 | public int Steps { get; set; } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -5,8 +5,8 @@ namespace Tensorflow.Keras.ArgsDefinition | |||
| 5 | 5 | { | |
| 6 | 6 | public class DataHandlerArgs: IKerasConfig | |
| 7 | 7 | { | |
| 8 | - public Tensor X { get; set; } | ||
| 9 | - public Tensor Y { get; set; } | ||
| 8 | + public Tensors X { get; set; } | ||
| 9 | + public Tensors Y { get; set; } | ||
| 10 | 10 | public IDatasetV2 Dataset { get; set; } | |
| 11 | 11 | public int BatchSize { get; set; } = 32; | |
| 12 | 12 | public int StepsPerEpoch { get; set; } = -1; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -24,6 +24,17 @@ ICallback fit(NDArray x, NDArray y, | |||
| 24 | 24 | int workers = 1, | |
| 25 | 25 | bool use_multiprocessing = false); | |
| 26 | 26 | ||
| 27 | + ICallback fit(IEnumerable<NDArray> x, NDArray y, | ||
| 28 | + int batch_size = -1, | ||
| 29 | + int epochs = 1, | ||
| 30 | + int verbose = 1, | ||
| 31 | + float validation_split = 0f, | ||
| 32 | + bool shuffle = true, | ||
| 33 | + int initial_epoch = 0, | ||
| 34 | + int max_queue_size = 10, | ||
| 35 | + int workers = 1, | ||
| 36 | + bool use_multiprocessing = false); | ||
| 37 | + | ||
| 27 | 38 | void save(string filepath, | |
| 28 | 39 | bool overwrite = true, | |
| 29 | 40 | bool include_optimizer = true, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -14,7 +14,76 @@ public void Deconstruct(out byte blue, out byte green, out byte red) | |||
| 14 | 14 | red = data[2]; | |
| 15 | 15 | } | |
| 16 | 16 | ||
| 17 | - public static implicit operator NDArray(Array array) | ||
| 17 | + public static implicit operator NDArray(int[] array) | ||
| 18 | + => new NDArray(array); | ||
| 19 | + | ||
| 20 | + public static implicit operator NDArray(byte[] array) | ||
| 21 | + => new NDArray(array); | ||
| 22 | + | ||
| 23 | + public static implicit operator NDArray(float[] array) | ||
| 24 | + => new NDArray(array); | ||
| 25 | + | ||
| 26 | + public static implicit operator NDArray(double[] array) | ||
| 27 | + => new NDArray(array); | ||
| 28 | + | ||
| 29 | + public static implicit operator NDArray(long[] array) | ||
| 30 | + => new NDArray(array); | ||
| 31 | + | ||
| 32 | + public static implicit operator NDArray(bool[] array) | ||
| 33 | + => new NDArray(array); | ||
| 34 | + | ||
| 35 | + public static implicit operator NDArray(uint[] array) | ||
| 36 | + => new NDArray(array); | ||
| 37 | + | ||
| 38 | + public static implicit operator NDArray(ulong[] array) | ||
| 39 | + => new NDArray(array); | ||
| 40 | + | ||
| 41 | + public static implicit operator NDArray(int[,] array) | ||
| 42 | + => new NDArray(array); | ||
| 43 | + | ||
| 44 | + public static implicit operator NDArray(byte[,] array) | ||
| 45 | + => new NDArray(array); | ||
| 46 | + | ||
| 47 | + public static implicit operator NDArray(float[,] array) | ||
| 48 | + => new NDArray(array); | ||
| 49 | + | ||
| 50 | + public static implicit operator NDArray(double[,] array) | ||
| 51 | + => new NDArray(array); | ||
| 52 | + | ||
| 53 | + public static implicit operator NDArray(long[,] array) | ||
| 54 | + => new NDArray(array); | ||
| 55 | + | ||
| 56 | + public static implicit operator NDArray(bool[,] array) | ||
| 57 | + => new NDArray(array); | ||
| 58 | + | ||
| 59 | + public static implicit operator NDArray(uint[,] array) | ||
| 60 | + => new NDArray(array); | ||
| 61 | + | ||
| 62 | + public static implicit operator NDArray(ulong[,] array) | ||
| 63 | + => new NDArray(array); | ||
| 64 | + | ||
| 65 | + public static implicit operator NDArray(int[,,] array) | ||
| 66 | + => new NDArray(array); | ||
| 67 | + | ||
| 68 | + public static implicit operator NDArray(byte[,,] array) | ||
| 69 | + => new NDArray(array); | ||
| 70 | + | ||
| 71 | + public static implicit operator NDArray(float[,,] array) | ||
| 72 | + => new NDArray(array); | ||
| 73 | + | ||
| 74 | + public static implicit operator NDArray(double[,,] array) | ||
| 75 | + => new NDArray(array); | ||
| 76 | + | ||
| 77 | + public static implicit operator NDArray(long[,,] array) | ||
| 78 | + => new NDArray(array); | ||
| 79 | + | ||
| 80 | + public static implicit operator NDArray(bool[,,] array) | ||
| 81 | + => new NDArray(array); | ||
| 82 | + | ||
| 83 | + public static implicit operator NDArray(uint[,,] array) | ||
| 84 | + => new NDArray(array); | ||
| 85 | + | ||
| 86 | + public static implicit operator NDArray(ulong[,,] array) | ||
| 18 | 87 | => new NDArray(array); | |
| 19 | 88 | ||
| 20 | 89 | public unsafe static implicit operator bool(NDArray nd) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -25,7 +25,7 @@ private NDArray OpenEntry(ZipArchiveEntry entry) | |||
| 25 | 25 | return array; | |
| 26 | 26 | ||
| 27 | 27 | using var s = entry.Open(); | |
| 28 | - return LoadMatrix(s); | ||
| 28 | + return (NDArray)LoadMatrix(s); | ||
| 29 | 29 | } | |
| 30 | 30 | ||
| 31 | 31 | public Array LoadMatrix(Stream stream) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -49,5 +49,8 @@ public IEnumerator<NDArray> GetEnumerator() | |||
| 49 | 49 | ||
| 50 | 50 | IEnumerator IEnumerable.GetEnumerator() | |
| 51 | 51 | => GetEnumerator(); | |
| 52 | + | ||
| 53 | + public static explicit operator NDArray(Array array) | ||
| 54 | + => new NDArray(array); | ||
| 52 | 55 | } | |
| 53 | 56 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,6 +1,9 @@ | |||
| 1 | 1 | using System; | |
| 2 | + using Tensorflow.Contexts; | ||
| 3 | + using Tensorflow.Eager; | ||
| 2 | 4 | using Tensorflow.Framework.Models; | |
| 3 | 5 | using Tensorflow.Functions; | |
| 6 | + using Tensorflow.Operations; | ||
| 4 | 7 | using static Tensorflow.Binding; | |
| 5 | 8 | ||
| 6 | 9 | namespace Tensorflow | |
@@ -220,6 +223,37 @@ public Tensor model_dataset(Tensor input_dataset, | |||
| 220 | 223 | return (results[0], results[1]); | |
| 221 | 224 | } | |
| 222 | 225 | ||
| 226 | + public Tensor anonymous_iterator_v3(TF_DataType[] output_types, Shape[] output_shapes, string name = null) | ||
| 227 | + { | ||
| 228 | + var ctx = tf.Context; | ||
| 229 | + Dictionary<string, object> attrs = new(); | ||
| 230 | + attrs["output_types"] = output_types; | ||
| 231 | + attrs["output_shapes"] = output_shapes; | ||
| 232 | + if (ctx.executing_eagerly()) | ||
| 233 | + { | ||
| 234 | + try | ||
| 235 | + { | ||
| 236 | + var result = tf.Runner.TFE_FastPathExecute(new FastPathOpExecInfo("AnonymousIteratorV3", name) | ||
| 237 | + { | ||
| 238 | + attrs = attrs | ||
| 239 | + }); | ||
| 240 | + return result[0]; | ||
| 241 | + } | ||
| 242 | + catch (Exception) | ||
| 243 | + { | ||
| 244 | + return anonymous_iterator_v3_eager_fallback(output_types, output_shapes, name, ctx); | ||
| 245 | + } | ||
| 246 | + } | ||
| 247 | + return tf.OpDefLib._apply_op_helper("AnonymousIteratorV3", name, attrs).outputs[0]; | ||
| 248 | + } | ||
| 249 | + | ||
| 250 | + public Tensor anonymous_iterator_v3_eager_fallback(TF_DataType[] output_types, Shape[] output_shapes, string name, Context ctx) | ||
| 251 | + { | ||
| 252 | + object[] attrs = new object[] { output_types, output_shapes }; | ||
| 253 | + var result = execute.quick_execute("AnonymousIteratorV3", 1, new Tensor[] { }, attrs, ctx, name); | ||
| 254 | + return result[0]; | ||
| 255 | + } | ||
| 256 | + | ||
| 223 | 257 | /// <summary> | |
| 224 | 258 | /// Makes a new iterator from the given `dataset` and stores it in `iterator`. | |
| 225 | 259 | /// </summary> | |
| Back | FazBrowse Home | New Git URL |
0 commit comments