| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent af9b64c commit e93f112
55 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -14,10 +14,15 @@ You may obtain a copy of the License at | |||
| 14 | 14 | limitations under the License. | |
| 15 | 15 | ******************************************************************************/ | |
| 16 | 16 | ||
| 17 | + using Tensorflow.Gradients; | ||
| 18 | + | ||
| 17 | 19 | namespace Tensorflow | |
| 18 | 20 | { | |
| 19 | 21 | public partial class tensorflow | |
| 20 | 22 | { | |
| 23 | + public GradientActor GradientTape() | ||
| 24 | + => new GradientActor(); | ||
| 25 | + | ||
| 21 | 26 | public Tensor[] gradients(Tensor[] ys, | |
| 22 | 27 | Tensor[] xs, | |
| 23 | 28 | Tensor[] grad_ys = null, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,32 @@ | |||
| 1 | + /***************************************************************************** | ||
| 2 | + Copyright 2018 The TensorFlow.NET Authors. All Rights Reserved. | ||
| 3 | + | ||
| 4 | + Licensed under the Apache License, Version 2.0 (the "License"); | ||
| 5 | + you may not use this file except in compliance with the License. | ||
| 6 | + You may obtain a copy of the License at | ||
| 7 | + | ||
| 8 | + http://www.apache.org/licenses/LICENSE-2.0 | ||
| 9 | + | ||
| 10 | + Unless required by applicable law or agreed to in writing, software | ||
| 11 | + distributed under the License is distributed on an "AS IS" BASIS, | ||
| 12 | + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. | ||
| 13 | + See the License for the specific language governing permissions and | ||
| 14 | + limitations under the License. | ||
| 15 | + ******************************************************************************/ | ||
| 16 | + | ||
| 17 | + using Tensorflow.Keras; | ||
| 18 | + using Tensorflow.Keras.Engine; | ||
| 19 | + using Tensorflow.Keras.Optimizers; | ||
| 20 | + | ||
| 21 | + namespace Tensorflow | ||
| 22 | + { | ||
| 23 | + public partial class tensorflow | ||
| 24 | + { | ||
| 25 | + public KerasOptimizers optimizers => new KerasOptimizers(); | ||
| 26 | + | ||
| 27 | + public class KerasOptimizers | ||
| 28 | + { | ||
| 29 | + public SGD SGD(float learning_rate) => new SGD(learning_rate); | ||
| 30 | + } | ||
| 31 | + } | ||
| 32 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -41,6 +41,14 @@ public static void add<T>(this IList<T> list, T element) | |||
| 41 | 41 | public static void append<T>(this IList<T> list, T element) | |
| 42 | 42 | => list.Add(element); | |
| 43 | 43 | ||
| 44 | + public static T[] concat<T>(this IList<T> list1, IList<T> list2) | ||
| 45 | + { | ||
| 46 | + var list = new List<T>(); | ||
| 47 | + list.AddRange(list1); | ||
| 48 | + list.AddRange(list2); | ||
| 49 | + return list.ToArray(); | ||
| 50 | + } | ||
| 51 | + | ||
| 44 | 52 | public static void extend<T>(this List<T> list, IEnumerable<T> elements) | |
| 45 | 53 | => list.AddRange(elements); | |
| 46 | 54 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,14 @@ | |||
| 1 | + using NumSharp; | ||
| 2 | + using System; | ||
| 3 | + using System.Collections.Generic; | ||
| 4 | + using System.Text; | ||
| 5 | + using Tensorflow.Eager; | ||
| 6 | + | ||
| 7 | + namespace Tensorflow.Eager | ||
| 8 | + { | ||
| 9 | + public partial class EagerTensor | ||
| 10 | + { | ||
| 11 | + public static explicit operator TFE_TensorHandle(EagerTensor tensor) | ||
| 12 | + => tensor.tfe_tensor_handle; | ||
| 13 | + } | ||
| 14 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -5,30 +5,39 @@ | |||
| 5 | 5 | ||
| 6 | 6 | namespace Tensorflow.Eager | |
| 7 | 7 | { | |
| 8 | - public class EagerTensor : Tensor | ||
| 8 | + public partial class EagerTensor : Tensor | ||
| 9 | 9 | { | |
| 10 | + Status status = new Status(); | ||
| 11 | + TFE_TensorHandle tfe_tensor_handle; | ||
| 10 | 12 | public EagerTensor(IntPtr handle) : base(handle) | |
| 11 | 13 | { | |
| 14 | + tfe_tensor_handle = handle; | ||
| 15 | + _handle = c_api.TFE_TensorHandleResolve(handle, status); | ||
| 12 | 16 | } | |
| 13 | 17 | ||
| 14 | 18 | public EagerTensor(string value, string device_name) : base(value) | |
| 15 | 19 | { | |
| 20 | + tfe_tensor_handle = c_api.TFE_NewTensorHandle(_handle, status); | ||
| 16 | 21 | } | |
| 17 | 22 | ||
| 18 | 23 | public EagerTensor(int value, string device_name) : base(value) | |
| 19 | 24 | { | |
| 25 | + tfe_tensor_handle = c_api.TFE_NewTensorHandle(_handle, status); | ||
| 20 | 26 | } | |
| 21 | 27 | ||
| 22 | 28 | public EagerTensor(float[] value, string device_name) : base(value) | |
| 23 | 29 | { | |
| 30 | + tfe_tensor_handle = c_api.TFE_NewTensorHandle(_handle, status); | ||
| 24 | 31 | } | |
| 25 | 32 | ||
| 26 | 33 | public EagerTensor(double[] value, string device_name) : base(value) | |
| 27 | 34 | { | |
| 35 | + tfe_tensor_handle = c_api.TFE_NewTensorHandle(_handle, status); | ||
| 28 | 36 | } | |
| 29 | 37 | ||
| 30 | 38 | public EagerTensor(NDArray value, string device_name) : base(value) | |
| 31 | 39 | { | |
| 40 | + tfe_tensor_handle = c_api.TFE_NewTensorHandle(_handle, status); | ||
| 32 | 41 | } | |
| 33 | 42 | ||
| 34 | 43 | public override string ToString() | |
@@ -51,6 +60,8 @@ private string GetFormattedString() | |||
| 51 | 60 | { | |
| 52 | 61 | case TF_DataType.TF_STRING: | |
| 53 | 62 | return $"b'{(string)nd}'"; | |
| 63 | + case TF_DataType.TF_BOOL: | ||
| 64 | + return (nd.GetByte(0) > 0).ToString(); | ||
| 54 | 65 | default: | |
| 55 | 66 | return nd.ToString(); | |
| 56 | 67 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -32,31 +32,41 @@ public Tensor execute(Context ctx, string op_name, Tensor[] inputs, object[] att | |||
| 32 | 32 | ctx.ensure_initialized(); | |
| 33 | 33 | using (var status = new Status()) | |
| 34 | 34 | { | |
| 35 | - var retVals = wrap_tfe_src.TFE_Py_Execute(ctx, ctx.device_name, op_name, inputs, attrs, 1, status); | ||
| 35 | + var retVals = wrap_tfe_src.TFE_Execute(ctx, ctx.device_name, op_name, inputs, attrs, 1, status); | ||
| 36 | 36 | ||
| 37 | - var t = c_api.TFE_TensorHandleResolve(retVals[0], status); | ||
| 38 | - status.Check(true); | ||
| 39 | - | ||
| 40 | - return new EagerTensor(t); | ||
| 37 | + return new EagerTensor(retVals[0]); | ||
| 41 | 38 | } | |
| 42 | 39 | } | |
| 43 | 40 | ||
| 44 | - public (TF_DataType, Tensor) args_to_matching_eager(Tensor[] l, Context ctx, TF_DataType default_dtype = TF_DataType.DtInvalid) | ||
| 41 | + public (TF_DataType, Tensor[]) args_to_matching_eager(Context ctx, TF_DataType default_dtype = TF_DataType.DtInvalid, object[] args = null) | ||
| 45 | 42 | { | |
| 46 | - var dtype = default_dtype; | ||
| 47 | - if(dtype == TF_DataType.DtInvalid) | ||
| 48 | - { | ||
| 49 | - var tensor = ops.convert_to_tensor(l, dtype, preferred_dtype: default_dtype, ctx: ctx); | ||
| 43 | + if (args.Length == 0 && default_dtype != TF_DataType.DtInvalid) | ||
| 44 | + return (default_dtype, null); | ||
| 50 | 45 | ||
| 51 | - if (dtype == TF_DataType.DtInvalid) | ||
| 52 | - dtype = tensor.dtype; | ||
| 46 | + if (args.Count(x => x is EagerTensor) == args.Length) | ||
| 47 | + return ((args[0] as EagerTensor).dtype, args.Select(x => x as EagerTensor).ToArray()); | ||
| 53 | 48 | ||
| 54 | - return (dtype, tensor); | ||
| 49 | + var dtype = TF_DataType.DtInvalid; | ||
| 50 | + foreach (var x in args) | ||
| 51 | + { | ||
| 52 | + if (x is EagerTensor et) | ||
| 53 | + dtype = et.dtype; | ||
| 55 | 54 | } | |
| 56 | - else | ||
| 55 | + | ||
| 56 | + if (dtype == TF_DataType.DtInvalid) | ||
| 57 | 57 | { | |
| 58 | - return (dtype, l[0]); | ||
| 58 | + var ret = new List<Tensor>(); | ||
| 59 | + foreach (var t in args) | ||
| 60 | + { | ||
| 61 | + ret.Add(ops.convert_to_tensor(t, dtype, preferred_dtype: default_dtype, ctx: ctx)); | ||
| 62 | + if (dtype == TF_DataType.DtInvalid) | ||
| 63 | + dtype = ret.Last().dtype; | ||
| 64 | + } | ||
| 65 | + | ||
| 66 | + return (dtype, ret.ToArray()); | ||
| 59 | 67 | } | |
| 68 | + else | ||
| 69 | + throw new NotImplementedException(""); | ||
| 60 | 70 | } | |
| 61 | 71 | ||
| 62 | 72 | public void record_gradient(string op_name, InputList inputs, Dictionary<string, object> attrs, Tensor[] results, string name = null) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -101,6 +101,16 @@ public partial class c_api | |||
| 101 | 101 | [DllImport(TensorFlowLibName)] | |
| 102 | 102 | public static extern TFE_Op TFE_NewOp(IntPtr ctx, string op_or_function_name, IntPtr status); | |
| 103 | 103 | ||
| 104 | + /// <summary> | ||
| 105 | + /// | ||
| 106 | + /// </summary> | ||
| 107 | + /// <param name="ctx">TFE_Context*</param> | ||
| 108 | + /// <param name="op_or_function_name">const char*</param> | ||
| 109 | + /// <param name="status">TF_Status*</param> | ||
| 110 | + /// <param name="op_to_reset">TFE_Op*</param> | ||
| 111 | + [DllImport(TensorFlowLibName)] | ||
| 112 | + public static extern void TFE_OpReset(IntPtr ctx, string op_or_function_name, IntPtr status, IntPtr op_to_reset); | ||
| 113 | + | ||
| 104 | 114 | /// <summary> | |
| 105 | 115 | /// | |
| 106 | 116 | /// </summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,7 +1,7 @@ | |||
| 1 | 1 | using System.Collections.Generic; | |
| 2 | 2 | using System.Linq; | |
| 3 | 3 | using System; | |
| 4 | - using static Tensorflow.OpDef.Types; | ||
| 4 | + using Tensorflow.Gradients; | ||
| 5 | 5 | ||
| 6 | 6 | namespace Tensorflow.Eager | |
| 7 | 7 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -10,31 +10,40 @@ namespace Tensorflow.Eager | |||
| 10 | 10 | /// </summary> | |
| 11 | 11 | public partial class wrap_tfe_src | |
| 12 | 12 | { | |
| 13 | - public static IntPtr[] TFE_Py_Execute(Context ctx, | ||
| 13 | + public static IntPtr[] TFE_Execute(Context ctx, | ||
| 14 | 14 | string device_name, | |
| 15 | 15 | string op_name, | |
| 16 | 16 | Tensor[] inputs, | |
| 17 | 17 | object[] attrs, | |
| 18 | 18 | int num_outputs, | |
| 19 | 19 | Status status) | |
| 20 | - => TFE_Py_ExecuteCancelable(ctx, device_name, op_name, inputs, attrs, num_outputs, status); | ||
| 20 | + => TFE_ExecuteCancelable(ctx, device_name, op_name, inputs, attrs, num_outputs, status); | ||
| 21 | 21 | ||
| 22 | - public static IntPtr[] TFE_Py_ExecuteCancelable(Context ctx, | ||
| 22 | + public static IntPtr[] TFE_ExecuteCancelable(Context ctx, | ||
| 23 | 23 | string device_name, | |
| 24 | 24 | string op_name, | |
| 25 | 25 | Tensor[] inputs, | |
| 26 | 26 | object[] attrs, | |
| 27 | 27 | int num_outputs, | |
| 28 | 28 | Status status) | |
| 29 | 29 | { | |
| 30 | - var op = c_api.TFE_NewOp(ctx, op_name, status); | ||
| 30 | + var op = GetOp(ctx, op_name, status); | ||
| 31 | 31 | status.Check(true); | |
| 32 | 32 | c_api.TFE_OpSetDevice(op, device_name, status); | |
| 33 | 33 | if(status.ok()) | |
| 34 | 34 | { | |
| 35 | 35 | for (int i = 0; i < inputs.Length; ++i) | |
| 36 | 36 | { | |
| 37 | - var tensor_handle = c_api.TFE_NewTensorHandle(inputs[i], status); | ||
| 37 | + TFE_TensorHandle tensor_handle; | ||
| 38 | + switch (inputs[i]) | ||
| 39 | + { | ||
| 40 | + case EagerTensor et: | ||
| 41 | + tensor_handle = (TFE_TensorHandle)et; | ||
| 42 | + break; | ||
| 43 | + default: | ||
| 44 | + tensor_handle = c_api.TFE_NewTensorHandle(inputs[i], status); | ||
| 45 | + break; | ||
| 46 | + } | ||
| 38 | 47 | c_api.TFE_OpAddInput(op, tensor_handle, status); | |
| 39 | 48 | } | |
| 40 | 49 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -22,7 +22,7 @@ public static EagerTensor TFE_FastPathExecute(Context ctx, | |||
| 22 | 22 | var attr_list_sizes = new Dictionary<string, long>(); | |
| 23 | 23 | using (var status = new Status()) | |
| 24 | 24 | { | |
| 25 | - var op = c_api.TFE_NewOp(ctx, opName, status); | ||
| 25 | + var op = GetOp(ctx, opName, status); | ||
| 26 | 26 | ||
| 27 | 27 | var op_def = Graph.TFE_GetOpDef(opName); | |
| 28 | 28 | ||
@@ -101,11 +101,31 @@ public static EagerTensor TFE_FastPathExecute(Context ctx, | |||
| 101 | 101 | c_api.TFE_Execute(op, retVals, ref num_retvals, status); | |
| 102 | 102 | status.Check(true); | |
| 103 | 103 | ||
| 104 | - var t = c_api.TFE_TensorHandleResolve(retVals[0], status); | ||
| 105 | - status.Check(true); | ||
| 104 | + return num_retvals == 0 ? null : new EagerTensor(retVals[0]); | ||
| 105 | + } | ||
| 106 | + } | ||
| 106 | 107 | ||
| 107 | - return new EagerTensor(t); | ||
| 108 | + private static TFE_Op GetOp(Context ctx, string op_or_function_name, Status status) | ||
| 109 | + { | ||
| 110 | + var maybe_op = ReleaseThreadLocalOp(); | ||
| 111 | + if (maybe_op != IntPtr.Zero) | ||
| 112 | + { | ||
| 113 | + c_api.TFE_OpReset(ctx, op_or_function_name, status, maybe_op); | ||
| 114 | + } | ||
| 115 | + else | ||
| 116 | + { | ||
| 117 | + maybe_op = c_api.TFE_NewOp(ctx, op_or_function_name, status); | ||
| 118 | + op = maybe_op; | ||
| 108 | 119 | } | |
| 120 | + | ||
| 121 | + status.Check(true); | ||
| 122 | + return maybe_op; | ||
| 123 | + } | ||
| 124 | + | ||
| 125 | + static TFE_Op op; | ||
| 126 | + private static TFE_Op ReleaseThreadLocalOp() | ||
| 127 | + { | ||
| 128 | + return op; | ||
| 109 | 129 | } | |
| 110 | 130 | ||
| 111 | 131 | /// <summary> | |
@@ -126,19 +146,19 @@ private static bool AddInputToOp(object inputs, | |||
| 126 | 146 | { | |
| 127 | 147 | TFE_TensorHandle input_handle; | |
| 128 | 148 | ||
| 149 | + // ConvertToTensor(); | ||
| 129 | 150 | switch (inputs) | |
| 130 | 151 | { | |
| 131 | - case Tensor input: | ||
| 132 | - input_handle = c_api.TFE_NewTensorHandle(input, status); | ||
| 152 | + case EagerTensor input: | ||
| 153 | + input_handle = (TFE_TensorHandle)input; | ||
| 133 | 154 | break; | |
| 134 | - case Tensor[] input_list: | ||
| 135 | - input_handle = c_api.TFE_NewTensorHandle(input_list[0], status); | ||
| 155 | + case EagerTensor[] input_list: | ||
| 156 | + input_handle = (TFE_TensorHandle)input_list[0]; | ||
| 136 | 157 | break; | |
| 137 | 158 | default: | |
| 138 | 159 | throw new NotImplementedException(""); | |
| 139 | 160 | } | |
| 140 | 161 | ||
| 141 | - | ||
| 142 | 162 | if(add_type_attr && !string.IsNullOrEmpty(input_arg.TypeAttr)) | |
| 143 | 163 | { | |
| 144 | 164 | var dtype = c_api.TFE_TensorHandleDataType(input_handle); | |
| Back | FazBrowse Home | New Git URL |
0 commit comments