FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

Completed `TEST(CAPI, Graph)` porting · attackgithub/TensorFlow.NET@64cc96b · GitHub

Repository navigation

Commit 64cc96b

Browse files
committed
Completed TEST(CAPI, Graph) porting
1 parent d5b81e9 commit 64cc96b

12 files changed

Lines changed: 241 additions & 24 deletions

File tree

‎TensorFlow.NET.sln‎

Lines changed: 0 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -9,8 +9,6 @@ Project("{9A19103F-16F7-4668-BE54-9A1E7A4F7556}") = "TensorFlowNET.Core", "src\T
99
EndProject
1010
Project("{9A19103F-16F7-4668-BE54-9A1E7A4F7556}") = "TensorFlowNET.Examples", "test\TensorFlowNET.Examples\TensorFlowNET.Examples.csproj", "{1FE60088-157C-4140-91AB-E96B915E4BAE}"
1111
EndProject
12-
Project("{9A19103F-16F7-4668-BE54-9A1E7A4F7556}") = "NumSharp.Core", "..\NumSharp\src\NumSharp.Core\NumSharp.Core.csproj", "{DA680126-DA60-4CE3-9094-72C355C081D3}"
13-
EndProject
1412
Global
1513
GlobalSection(SolutionConfigurationPlatforms) = preSolution
1614
Debug|Any CPU = Debug|Any CPU
@@ -29,10 +27,6 @@ Global
2927
{1FE60088-157C-4140-91AB-E96B915E4BAE}.Debug|Any CPU.Build.0 = Debug|Any CPU
3028
{1FE60088-157C-4140-91AB-E96B915E4BAE}.Release|Any CPU.ActiveCfg = Release|Any CPU
3129
{1FE60088-157C-4140-91AB-E96B915E4BAE}.Release|Any CPU.Build.0 = Release|Any CPU
32-
{DA680126-DA60-4CE3-9094-72C355C081D3}.Debug|Any CPU.ActiveCfg = Debug|Any CPU
33-
{DA680126-DA60-4CE3-9094-72C355C081D3}.Debug|Any CPU.Build.0 = Debug|Any CPU
34-
{DA680126-DA60-4CE3-9094-72C355C081D3}.Release|Any CPU.ActiveCfg = Release|Any CPU
35-
{DA680126-DA60-4CE3-9094-72C355C081D3}.Release|Any CPU.Build.0 = Release|Any CPU
3630
EndGlobalSection
3731
GlobalSection(SolutionProperties) = preSolution
3832
HideSolutionNode = FALSE

‎src/TensorFlowNET.Core/Graphs/c_api.graph.cs‎

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,6 +31,25 @@ public static partial class c_api
3131
[DllImport(TensorFlowLibName)]
3232
public static extern void TF_GraphGetTensorShape(IntPtr graph, TF_Output output, long[] dims, int num_dims, IntPtr status);
3333

34+
/// <summary>
35+
/// Iterate through the operations of a graph.
36+
/// </summary>
37+
/// <param name="graph"></param>
38+
/// <param name="pos"></param>
39+
/// <returns></returns>
40+
[DllImport(TensorFlowLibName)]
41+
public static extern IntPtr TF_GraphNextOperation(IntPtr graph, ref uint pos);
42+
43+
/// <summary>
44+
/// Returns the operation in the graph with `oper_name`. Returns nullptr if
45+
/// no operation found.
46+
/// </summary>
47+
/// <param name="graph"></param>
48+
/// <param name="oper_name"></param>
49+
/// <returns></returns>
50+
[DllImport(TensorFlowLibName)]
51+
public static extern IntPtr TF_GraphOperationByName(IntPtr graph, string oper_name);
52+
3453
/// <summary>
3554
/// Sets the shape of the Tensor referenced by `output` in `graph` to
3655
/// the shape described by `dims` and `num_dims`.

‎src/TensorFlowNET.Core/Operations/Operation.cs‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -116,6 +116,8 @@ public override bool Equals(object obj)
116116
{
117117
case IntPtr val:
118118
return val == _handle;
119+
case Operation val:
120+
return val._handle == _handle;
119121
}
120122

121123
return base.Equals(obj);
Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,26 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Text;
4+
5+
namespace Tensorflow
6+
{
7+
public class OperationDescription
8+
{
9+
private IntPtr _handle;
10+
11+
public OperationDescription(IntPtr handle)
12+
{
13+
_handle = handle;
14+
}
15+
16+
public static implicit operator OperationDescription(IntPtr handle)
17+
{
18+
return new OperationDescription(handle);
19+
}
20+
21+
public static implicit operator IntPtr(OperationDescription desc)
22+
{
23+
return desc._handle;
24+
}
25+
}
26+
}

‎src/TensorFlowNET.Core/Operations/TF_OperationDescription.cs‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ namespace Tensorflow
99
public struct TF_OperationDescription
1010
{
1111
public IntPtr node_builder;
12-
//public TF_Graph graph;
12+
public IntPtr graph;
13+
public IntPtr colocation_constraints;
1314
}
1415
}

‎src/TensorFlowNET.Core/TensorFlowNET.Core.csproj‎

Lines changed: 1 addition & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -33,11 +33,7 @@
3333

3434
<ItemGroup>
3535
<PackageReference Include="Google.Protobuf" Version="3.6.1" />
36-
<PackageReference Include="NumSharp" Version="0.6.2" />
37-
</ItemGroup>
38-
39-
<ItemGroup>
40-
<ProjectReference Include="..\..\..\NumSharp\src\NumSharp.Core\NumSharp.Core.csproj" />
36+
<PackageReference Include="NumSharp" Version="0.6.3" />
4137
</ItemGroup>
4238

4339
<ItemGroup>

‎src/TensorFlowNET.Core/c_api.cs‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -22,6 +22,7 @@ namespace Tensorflow
2222
/// int32_t => int
2323
/// int64_t* => long[]
2424
/// size_t* => unlong[]
25+
/// size_t* => ref uint
2526
/// void* => IntPtr
2627
/// string => IntPtr c_api.StringPiece(IntPtr)
2728
/// </summary>

‎test/TensorFlowNET.Examples/TensorFlowNET.Examples.csproj‎

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -6,11 +6,10 @@
66
</PropertyGroup>
77

88
<ItemGroup>
9-
<PackageReference Include="NumSharp" Version="0.6.2" />
9+
<PackageReference Include="NumSharp" Version="0.6.3" />
1010
</ItemGroup>
1111

1212
<ItemGroup>
13-
<ProjectReference Include="..\..\..\NumSharp\src\NumSharp.Core\NumSharp.Core.csproj" />
1413
<ProjectReference Include="..\..\src\TensorFlowNET.Core\TensorFlowNET.Core.csproj" />
1514
</ItemGroup>
1615

‎test/TensorFlowNET.UnitTest/GraphTest.cs‎

Lines changed: 81 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -104,21 +104,98 @@ public void c_api_Graph()
104104
Assert.IsFalse(found_placeholder);
105105
found_placeholder = true;
106106
}
107-
/*else if (IsScalarConst(n, 3))
107+
else if (c_test_util.IsScalarConst(n, 3))
108108
{
109109
Assert.IsFalse(found_scalar_const);
110110
found_scalar_const = true;
111111
}
112-
else if (IsAddN(n, 2))
112+
else if (c_test_util.IsAddN(n, 2))
113113
{
114114
Assert.IsFalse(found_add);
115115
found_add = true;
116116
}
117117
else
118118
{
119-
ADD_FAILURE() << "Unexpected NodeDef: " << ProtoDebugString(n);
120-
}*/
119+
Assert.Fail($"Unexpected NodeDef: {n}");
120+
}
121121
}
122+
Assert.IsTrue(found_placeholder);
123+
Assert.IsTrue(found_scalar_const);
124+
Assert.IsTrue(found_add);
125+
126+
// Add another oper to the graph.
127+
var neg = c_test_util.Neg(add, graph, s);
128+
Assert.AreEqual(TF_Code.TF_OK, s.Code);
129+
130+
// Serialize to NodeDef.
131+
var node_def = c_test_util.GetNodeDef(neg);
132+
133+
// Validate NodeDef is what we expect.
134+
Assert.IsTrue(c_test_util.IsNeg(node_def, "add"));
135+
136+
// Serialize to GraphDef.
137+
var graph_def2 = c_test_util.GetGraphDef(graph);
138+
139+
// Compare with first GraphDef + added NodeDef.
140+
graph_def.Node.Add(node_def);
141+
Assert.AreEqual(graph_def.ToString(), graph_def2.ToString());
142+
143+
// Look up some nodes by name.
144+
Operation neg2 = c_api.TF_GraphOperationByName(graph, "neg");
145+
Assert.AreEqual(neg, neg2);
146+
var node_def2 = c_test_util.GetNodeDef(neg2);
147+
Assert.AreEqual(node_def.ToString(), node_def2.ToString());
148+
149+
Operation feed2 = c_api.TF_GraphOperationByName(graph, "feed");
150+
Assert.AreEqual(feed, feed2);
151+
node_def = c_test_util.GetNodeDef(feed);
152+
node_def2 = c_test_util.GetNodeDef(feed2);
153+
Assert.AreEqual(node_def.ToString(), node_def2.ToString());
154+
155+
// Test iterating through the nodes of a graph.
156+
found_placeholder = false;
157+
found_scalar_const = false;
158+
found_add = false;
159+
bool found_neg = false;
160+
uint pos = 0;
161+
Operation oper;
162+
163+
while((oper = c_api.TF_GraphNextOperation(graph, ref pos)) != IntPtr.Zero)
164+
{
165+
if (oper.Equals(feed))
166+
{
167+
Assert.IsFalse(found_placeholder);
168+
found_placeholder = true;
169+
}
170+
else if (oper.Equals(three))
171+
{
172+
Assert.IsFalse(found_scalar_const);
173+
found_scalar_const = true;
174+
}
175+
else if (oper.Equals(add))
176+
{
177+
Assert.IsFalse(found_add);
178+
found_add = true;
179+
}
180+
else if (oper.Equals(neg))
181+
{
182+
Assert.IsFalse(found_neg);
183+
found_neg = true;
184+
}
185+
else
186+
{
187+
node_def = c_test_util.GetNodeDef(oper);
188+
Assert.Fail($"Unexpected Node: {node_def.ToString()}");
189+
}
190+
}
191+
192+
Assert.IsTrue(found_placeholder);
193+
Assert.IsTrue(found_scalar_const);
194+
Assert.IsTrue(found_add);
195+
Assert.IsTrue(found_neg);
196+
197+
graph.Dispose();
198+
s.Dispose();
122199
}
123200
}
124201
}

‎test/TensorFlowNET.UnitTest/StatusTest.cs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,7 @@ public void SetStatus()
2323
var s = new Status();
2424
s.SetStatus(TF_Code.TF_CANCELLED, "cancel");
2525
Assert.AreEqual(s.Code, TF_Code.TF_CANCELLED);
26-
// Assert.AreEqual(s.Message, "cancel");
26+
Assert.AreEqual(s.Message, "cancel");
2727
}
2828

2929
[TestMethod]

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL