| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 740ca28 commit 4fa14fe
10 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,5 +1,6 @@ | |||
| 1 | 1 | using System; | |
| 2 | 2 | using System.Collections.Generic; | |
| 3 | + using System.Runtime.InteropServices; | ||
| 3 | 4 | using System.Text; | |
| 4 | 5 | ||
| 5 | 6 | namespace Tensorflow | |
@@ -14,16 +15,16 @@ public class Operation | |||
| 14 | 15 | ||
| 15 | 16 | private Status status = new Status(); | |
| 16 | 17 | ||
| 17 | - public string name { get; } | ||
| 18 | - public string optype { get; } | ||
| 19 | - public string device { get; } | ||
| 20 | - public int NumOutputs { get; } | ||
| 21 | - public TF_DataType OutputType { get; } | ||
| 22 | - public int OutputListLength { get; } | ||
| 23 | - public int NumInputs { get; } | ||
| 24 | - public int NumConsumers { get; } | ||
| 25 | - public int NumControlInputs { get; } | ||
| 26 | - public int NumControlOutputs { get; } | ||
| 18 | + public string name => c_api.StringPiece(c_api.TF_OperationName(_handle)); | ||
| 19 | + public string optype => c_api.StringPiece(c_api.TF_OperationOpType(_handle)); | ||
| 20 | + public string device => c_api.StringPiece(c_api.TF_OperationDevice(_handle)); | ||
| 21 | + public int NumOutputs => c_api.TF_OperationNumOutputs(_handle); | ||
| 22 | + public TF_DataType OutputType => c_api.TF_OperationOutputType(new TF_Output(_handle, 0)); | ||
| 23 | + public int OutputListLength => c_api.TF_OperationOutputListLength(_handle, "output", status); | ||
| 24 | + public int NumInputs => c_api.TF_OperationNumInputs(_handle); | ||
| 25 | + public int NumConsumers => c_api.TF_OperationOutputNumConsumers(new TF_Output(_handle, 0)); | ||
| 26 | + public int NumControlInputs => c_api.TF_OperationNumControlInputs(_handle); | ||
| 27 | + public int NumControlOutputs => c_api.TF_OperationNumControlOutputs(_handle); | ||
| 27 | 28 | ||
| 28 | 29 | private Tensor[] _outputs; | |
| 29 | 30 | public Tensor[] outputs => _outputs; | |
@@ -35,17 +36,6 @@ public Operation(IntPtr handle) | |||
| 35 | 36 | return; | |
| 36 | 37 | ||
| 37 | 38 | _handle = handle; | |
| 38 | - | ||
| 39 | - name = c_api.TF_OperationName(_handle); | ||
| 40 | - optype = c_api.TF_OperationOpType(_handle); | ||
| 41 | - device = "";// c_api.TF_OperationDevice(_handle); | ||
| 42 | - NumOutputs = c_api.TF_OperationNumOutputs(_handle); | ||
| 43 | - OutputType = c_api.TF_OperationOutputType(new TF_Output(_handle, 0)); | ||
| 44 | - OutputListLength = c_api.TF_OperationOutputListLength(_handle, "output", status); | ||
| 45 | - NumInputs = c_api.TF_OperationNumInputs(_handle); | ||
| 46 | - NumConsumers = c_api.TF_OperationOutputNumConsumers(new TF_Output(_handle, 0)); | ||
| 47 | - NumControlInputs = c_api.TF_OperationNumControlInputs(_handle); | ||
| 48 | - NumControlOutputs = c_api.TF_OperationNumControlOutputs(_handle); | ||
| 49 | 39 | } | |
| 50 | 40 | ||
| 51 | 41 | public Operation(Graph g, string opType, string oper_name) | |
@@ -62,8 +52,8 @@ public Operation(NodeDef node_def, Graph g, List<Tensor> inputs = null, TF_DataT | |||
| 62 | 52 | Graph = g; | |
| 63 | 53 | ||
| 64 | 54 | _id_value = Graph._next_id(); | |
| 55 | + | ||
| 65 | 56 | _handle = ops._create_c_op(g, node_def, inputs); | |
| 66 | - NumOutputs = c_api.TF_OperationNumOutputs(_handle); | ||
| 67 | 57 | ||
| 68 | 58 | _outputs = new Tensor[NumOutputs]; | |
| 69 | 59 | for (int i = 0; i < NumOutputs; i++) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -38,7 +38,7 @@ public static partial class c_api | |||
| 38 | 38 | public static extern IntPtr TF_NewOperation(IntPtr graph, string opType, string oper_name); | |
| 39 | 39 | ||
| 40 | 40 | [DllImport(TensorFlowLibName)] | |
| 41 | - public static extern string TF_OperationDevice(IntPtr oper); | ||
| 41 | + public static extern IntPtr TF_OperationDevice(IntPtr oper); | ||
| 42 | 42 | ||
| 43 | 43 | /// <summary> | |
| 44 | 44 | /// Sets `output_attr_value` to the binary-serialized AttrValue proto | |
@@ -50,13 +50,13 @@ public static partial class c_api | |||
| 50 | 50 | public static extern int TF_OperationGetAttrValueProto(IntPtr oper, string attr_name, IntPtr output_attr_value, IntPtr status); | |
| 51 | 51 | ||
| 52 | 52 | [DllImport(TensorFlowLibName)] | |
| 53 | - public static extern string TF_OperationName(IntPtr oper); | ||
| 53 | + public static extern IntPtr TF_OperationName(IntPtr oper); | ||
| 54 | 54 | ||
| 55 | 55 | [DllImport(TensorFlowLibName)] | |
| 56 | 56 | public static extern int TF_OperationNumInputs(IntPtr oper); | |
| 57 | 57 | ||
| 58 | 58 | [DllImport(TensorFlowLibName)] | |
| 59 | - public static extern string TF_OperationOpType(IntPtr oper); | ||
| 59 | + public static extern IntPtr TF_OperationOpType(IntPtr oper); | ||
| 60 | 60 | ||
| 61 | 61 | /// <summary> | |
| 62 | 62 | /// Get the number of control inputs to an operation. | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -30,12 +30,12 @@ public static unsafe IntPtr _create_c_op(Graph graph, NodeDef node_def, List<Ten | |||
| 30 | 30 | // Add inputs | |
| 31 | 31 | if(inputs != null && inputs.Count > 0) | |
| 32 | 32 | { | |
| 33 | - /*foreach (var op_input in inputs) | ||
| 33 | + foreach (var op_input in inputs) | ||
| 34 | 34 | { | |
| 35 | 35 | c_api.TF_AddInput(op_desc, op_input._as_tf_output()); | |
| 36 | - }*/ | ||
| 36 | + } | ||
| 37 | 37 | ||
| 38 | - c_api.TF_AddInputList(op_desc, inputs.Select(x => x._as_tf_output()).ToArray(), inputs.Count); | ||
| 38 | + //c_api.TF_AddInputList(op_desc, inputs.Select(x => x._as_tf_output()).ToArray(), inputs.Count); | ||
| 39 | 39 | } | |
| 40 | 40 | ||
| 41 | 41 | var status = new Status(); | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -15,7 +15,7 @@ public class Status | |||
| 15 | 15 | /// <summary> | |
| 16 | 16 | /// Error message | |
| 17 | 17 | /// </summary> | |
| 18 | - public string Message => c_api.TF_Message(_handle); | ||
| 18 | + public string Message => c_api.StringPiece(c_api.TF_Message(_handle)); | ||
| 19 | 19 | ||
| 20 | 20 | /// <summary> | |
| 21 | 21 | /// Error code | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -12,15 +12,15 @@ public static partial class c_api | |||
| 12 | 12 | /// </summary> | |
| 13 | 13 | /// <param name="s"></param> | |
| 14 | 14 | [DllImport(TensorFlowLibName)] | |
| 15 | - public static unsafe extern void TF_DeleteStatus(IntPtr s); | ||
| 15 | + public static extern void TF_DeleteStatus(IntPtr s); | ||
| 16 | 16 | ||
| 17 | 17 | /// <summary> | |
| 18 | 18 | /// Return the code record in *s. | |
| 19 | 19 | /// </summary> | |
| 20 | 20 | /// <param name="s"></param> | |
| 21 | 21 | /// <returns></returns> | |
| 22 | 22 | [DllImport(TensorFlowLibName)] | |
| 23 | - public static extern unsafe TF_Code TF_GetCode(IntPtr s); | ||
| 23 | + public static extern TF_Code TF_GetCode(IntPtr s); | ||
| 24 | 24 | ||
| 25 | 25 | /// <summary> | |
| 26 | 26 | /// Return a pointer to the (null-terminated) error message in *s. | |
@@ -30,7 +30,7 @@ public static partial class c_api | |||
| 30 | 30 | /// <param name="s"></param> | |
| 31 | 31 | /// <returns></returns> | |
| 32 | 32 | [DllImport(TensorFlowLibName)] | |
| 33 | - public static extern unsafe string TF_Message(IntPtr s); | ||
| 33 | + public static extern IntPtr TF_Message(IntPtr s); | ||
| 34 | 34 | ||
| 35 | 35 | /// <summary> | |
| 36 | 36 | /// Return a new status object. | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -12,20 +12,27 @@ namespace Tensorflow | |||
| 12 | 12 | /// The API leans towards simplicity and uniformity instead of convenience | |
| 13 | 13 | /// since most usage will be by language specific wrappers. | |
| 14 | 14 | /// | |
| 15 | - /// The params type mapping between .net and c_api | ||
| 15 | + /// The params type mapping between c_api and .NET | ||
| 16 | 16 | /// TF_XX** => ref IntPtr (TF_Operation** op) => (ref IntPtr op) | |
| 17 | 17 | /// TF_XX* => IntPtr (TF_Graph* graph) => (IntPtr graph) | |
| 18 | 18 | /// struct => struct (TF_Output output) => (TF_Output output) | |
| 19 | + /// struct* => struct (TF_Output* output) => (TF_Output[] output) | ||
| 19 | 20 | /// const char* => string | |
| 20 | 21 | /// int32_t => int | |
| 21 | 22 | /// int64_t* => long[] | |
| 22 | 23 | /// size_t* => unlong[] | |
| 23 | 24 | /// void* => IntPtr | |
| 25 | + /// string => IntPtr c_api.StringPiece(IntPtr) | ||
| 24 | 26 | /// </summary> | |
| 25 | 27 | public static partial class c_api | |
| 26 | 28 | { | |
| 27 | 29 | public const string TensorFlowLibName = "tensorflow"; | |
| 28 | 30 | ||
| 31 | + public static string StringPiece(IntPtr handle) | ||
| 32 | + { | ||
| 33 | + return handle == IntPtr.Zero ? String.Empty : Marshal.PtrToStringAnsi(handle); | ||
| 34 | + } | ||
| 35 | + | ||
| 29 | 36 | public delegate void Deallocator(IntPtr data, IntPtr size, ref bool deallocator); | |
| 30 | 37 | ||
| 31 | 38 | [DllImport(TensorFlowLibName)] | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -37,7 +37,7 @@ public static void enable_eager_execution() | |||
| 37 | 37 | context.default_execution_mode = Context.EAGER_MODE; | |
| 38 | 38 | } | |
| 39 | 39 | ||
| 40 | - public static string VERSION => Marshal.PtrToStringAnsi(c_api.TF_Version()); | ||
| 40 | + public static string VERSION => c_api.StringPiece(c_api.TF_Version()); | ||
| 41 | 41 | ||
| 42 | 42 | public static Graph get_default_graph() | |
| 43 | 43 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -39,17 +39,14 @@ public void c_api_Graph() | |||
| 39 | 39 | // Test not found errors in TF_Operation*() query functions. | |
| 40 | 40 | Assert.AreEqual(-1, c_api.TF_OperationOutputListLength(feed, "bogus", s)); | |
| 41 | 41 | Assert.AreEqual(TF_Code.TF_INVALID_ARGUMENT, s.Code); | |
| 42 | - //Assert.IsFalse(c_test_util.GetAttrValue(feed, "missing", ref attr_value, s)); | ||
| 43 | - //Assert.AreEqual("Operation '' has no attr named 'missing'.", s.Message); | ||
| 42 | + Assert.IsFalse(c_test_util.GetAttrValue(feed, "missing", ref attr_value, s)); | ||
| 43 | + Assert.AreEqual("Operation 'feed' has no attr named 'missing'.", s.Message); | ||
| 44 | 44 | ||
| 45 | 45 | // Make a constant oper with the scalar "3". | |
| 46 | 46 | var three = c_test_util.ScalarConst(3, graph, s); | |
| 47 | 47 | ||
| 48 | 48 | // Add oper. | |
| 49 | 49 | var add = c_test_util.Add(feed, three, graph, s); | |
| 50 | - | ||
| 51 | - NodeDef node_def = null; | ||
| 52 | - c_test_util.GetNodeDef(feed, ref node_def); | ||
| 53 | 50 | } | |
| 54 | 51 | } | |
| 55 | 52 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -43,7 +43,7 @@ public void addInPlaceholder() | |||
| 43 | 43 | public void addInConstant() | |
| 44 | 44 | { | |
| 45 | 45 | var a = tf.constant(4.0f); | |
| 46 | - var b = tf.placeholder(tf.float32); | ||
| 46 | + var b = tf.constant(5.0f); | ||
| 47 | 47 | var c = tf.add(a, b); | |
| 48 | 48 | ||
| 49 | 49 | using (var sess = tf.Session()) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -23,11 +23,13 @@ public static void AddOpHelper(Operation l, Operation r, Graph graph, Status s, | |||
| 23 | 23 | { | |
| 24 | 24 | var desc = c_api.TF_NewOperation(graph, "AddN", name); | |
| 25 | 25 | ||
| 26 | - c_api.TF_AddInputList(desc, new TF_Output[] | ||
| 26 | + var inputs = new TF_Output[] | ||
| 27 | 27 | { | |
| 28 | 28 | new TF_Output(l, 0), | |
| 29 | 29 | new TF_Output(r, 0), | |
| 30 | - }, 2); | ||
| 30 | + }; | ||
| 31 | + | ||
| 32 | + c_api.TF_AddInputList(desc, inputs, inputs.Length); | ||
| 31 | 33 | ||
| 32 | 34 | op = c_api.TF_FinishOperation(desc, s); | |
| 33 | 35 | s.Check(); | |
| Back | FazBrowse Home | New Git URL |
0 commit comments