Skip to content

Commit 4076cd1

Browse files
committed
Eager.Variables.
1 parent 095261a commit 4076cd1

7 files changed

Lines changed: 139 additions & 18 deletions

File tree

src/TensorFlowNET.Core/Eager/c_api.eager.cs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -85,7 +85,7 @@ public partial class c_api
8585
/// <param name="num_dims">const int</param>
8686
/// <param name="out_status">TF_Status*</param>
8787
[DllImport(TensorFlowLibName)]
88-
public static extern void TFE_OpSetAttrShape(IntPtr op, string attr_name, long[] dims, int num_dims, Status out_status);
88+
public static extern void TFE_OpSetAttrShape(IntPtr op, string attr_name, long[] dims, int num_dims, IntPtr out_status);
8989

9090
/// <summary>
9191
///

test/TensorFlowNET.UnitTest/CApiTest.cs

Lines changed: 36 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -52,6 +52,9 @@ protected void TF_SetAttrBool(OperationDescription desc, string attrName, bool v
5252
protected TF_DataType TFE_TensorHandleDataType(IntPtr h)
5353
=> c_api.TFE_TensorHandleDataType(h);
5454

55+
protected int TFE_TensorHandleNumDims(IntPtr h, IntPtr status)
56+
=> c_api.TFE_TensorHandleNumDims(h, status);
57+
5558
protected TF_Code TF_GetCode(Status s)
5659
=> s.Code;
5760

@@ -79,9 +82,18 @@ protected void TFE_OpAddInput(IntPtr op, IntPtr h, IntPtr status)
7982
protected void TFE_OpSetAttrType(IntPtr op, string attr_name, TF_DataType value)
8083
=> c_api.TFE_OpSetAttrType(op, attr_name, value);
8184

85+
protected void TFE_OpSetAttrShape(IntPtr op, string attr_name, long[] dims, int num_dims, IntPtr out_status)
86+
=> c_api.TFE_OpSetAttrShape(op, attr_name, dims, num_dims, out_status);
87+
88+
protected void TFE_OpSetAttrString(IntPtr op, string attr_name, string value, uint length)
89+
=> c_api.TFE_OpSetAttrString(op, attr_name, value, length);
90+
8291
protected IntPtr TFE_NewOp(IntPtr ctx, string op_or_function_name, IntPtr status)
8392
=> c_api.TFE_NewOp(ctx, op_or_function_name, status);
8493

94+
protected void TFE_Execute(IntPtr op, IntPtr[] retvals, ref int num_retvals, IntPtr status)
95+
=> c_api.TFE_Execute(op, retvals, ref num_retvals, status);
96+
8597
protected IntPtr TFE_NewContextOptions()
8698
=> c_api.TFE_NewContextOptions();
8799

@@ -139,37 +151,49 @@ protected IntPtr TFE_TensorHandleCopyToDevice(IntPtr h, IntPtr ctx, string devic
139151
protected void TFE_OpSetDevice(IntPtr op, string device_name, IntPtr status)
140152
=> c_api.TFE_OpSetDevice(op, device_name, status);
141153

142-
protected unsafe void memcpy(void * src, IntPtr dst, ulong size)
154+
protected unsafe void memcpy<T>(T* dst, void* src, ulong size)
155+
where T : unmanaged
143156
{
144-
Buffer.MemoryCopy(src, dst.ToPointer(), size, size);
157+
Buffer.MemoryCopy(src, dst, size, size);
145158
}
146159

147-
protected unsafe void memcpy<T>(T[] src, IntPtr dst, ulong size)
160+
protected unsafe void memcpy<T>(void* dst, T* src, ulong size)
148161
where T : unmanaged
149162
{
150-
fixed (void* p = &src[0])
151-
Buffer.MemoryCopy(p, dst.ToPointer(), size, size);
163+
Buffer.MemoryCopy(src, dst, size, size);
152164
}
153165

154-
protected unsafe void memcpy<T>(T[] src, IntPtr dst, long size)
155-
where T : unmanaged
166+
protected unsafe void memcpy(void * dst, IntPtr src, ulong size)
156167
{
157-
fixed (void* p = &src[0])
158-
Buffer.MemoryCopy(p, dst.ToPointer(), size, size);
168+
Buffer.MemoryCopy(src.ToPointer(), dst, size, size);
159169
}
160170

161-
protected unsafe void memcpy<T>(IntPtr src, T[] dst, ulong size)
171+
protected unsafe void memcpy<T>(T[] dst, IntPtr src, ulong size)
162172
where T : unmanaged
163173
{
164174
fixed (void* p = &dst[0])
165175
Buffer.MemoryCopy(src.ToPointer(), p, size, size);
166176
}
167177

168-
protected unsafe void memcpy<T>(IntPtr src, T[] dst, long size)
169-
where T: unmanaged
178+
protected unsafe void memcpy<T>(T[] dst, IntPtr src, long size)
179+
where T : unmanaged
170180
{
171181
fixed (void* p = &dst[0])
172182
Buffer.MemoryCopy(src.ToPointer(), p, size, size);
173183
}
184+
185+
protected unsafe void memcpy<T>(IntPtr dst, T[] src, ulong size)
186+
where T : unmanaged
187+
{
188+
fixed (void* p = &src[0])
189+
Buffer.MemoryCopy(p, dst.ToPointer(), size, size);
190+
}
191+
192+
protected unsafe void memcpy<T>(IntPtr dst, T[] src, long size)
193+
where T: unmanaged
194+
{
195+
fixed (void* p = &src[0])
196+
Buffer.MemoryCopy(p, dst.ToPointer(), size, size);
197+
}
174198
}
175199
}

test/TensorFlowNET.UnitTest/Eager/CApi.Eager.Execute_MatMul_CPU.cs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -43,7 +43,7 @@ unsafe void Execute_MatMul_CPU(bool async)
4343
ASSERT_EQ(TF_OK, TF_GetCode(status), TF_Message(status));
4444
var product = new float[4];
4545
EXPECT_EQ(product.Length * sizeof(float), (int)TF_TensorByteSize(t));
46-
memcpy(TF_TensorData(t), product, TF_TensorByteSize(t));
46+
memcpy(product, TF_TensorData(t), TF_TensorByteSize(t));
4747

4848
c_api.TF_DeleteTensor(t);
4949
EXPECT_EQ(7f, product[0]);

test/TensorFlowNET.UnitTest/Eager/CApi.Eager.TensorHandle.cs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -22,7 +22,7 @@ public unsafe void TensorHandle()
2222
ASSERT_EQ(16ul, c_api.TF_TensorByteSize(t));
2323

2424
var data = new float[] { 0f, 0f, 0f, 0f };
25-
memcpy(c_api.TF_TensorData(t), data, data.Length * sizeof(float));
25+
memcpy(data, c_api.TF_TensorData(t), data.Length * sizeof(float));
2626

2727
EXPECT_EQ(1.0f, data[0]);
2828
EXPECT_EQ(2.0f, data[1]);

test/TensorFlowNET.UnitTest/Eager/CApi.Eager.TensorHandleDevices.cs

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -61,10 +61,10 @@ public unsafe void TensorHandleDevices()
6161

6262
TFE_DeleteTensorHandle(hcpu);
6363
// not export api
64-
/*var executor = TFE_ContextGetExecutorForThread(ctx);
64+
var executor = TFE_ContextGetExecutorForThread(ctx);
6565
TFE_ExecutorWaitForAllPendingNodes(executor, status);
6666
ASSERT_EQ(TF_OK, TF_GetCode(status), TF_Message(status));
67-
TFE_DeleteExecutor(executor);*/
67+
TFE_DeleteExecutor(executor);
6868
TFE_DeleteContext(ctx);
6969
}
7070
}
Lines changed: 56 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,56 @@
1+
using Microsoft.VisualStudio.TestTools.UnitTesting;
2+
using System;
3+
using Tensorflow;
4+
using Tensorflow.Eager;
5+
using Buffer = System.Buffer;
6+
7+
namespace TensorFlowNET.UnitTest.Eager
8+
{
9+
public partial class CApiEagerTest
10+
{
11+
/// <summary>
12+
/// TEST(CAPI, Variables)
13+
/// </summary>
14+
[TestMethod]
15+
public unsafe void Variables()
16+
{
17+
var status = c_api.TF_NewStatus();
18+
var opts = TFE_NewContextOptions();
19+
var ctx = TFE_NewContext(opts, status);
20+
ASSERT_EQ(TF_OK, TF_GetCode(status), TF_Message(status));
21+
TFE_DeleteContextOptions(opts);
22+
23+
var var_handle = CreateVariable(ctx, 12.0f, status);
24+
ASSERT_EQ(TF_OK, TF_GetCode(status), TF_Message(status));
25+
26+
var op = TFE_NewOp(ctx, "ReadVariableOp", status);
27+
ASSERT_EQ(TF_OK, TF_GetCode(status), TF_Message(status));
28+
TFE_OpSetAttrType(op, "dtype", TF_FLOAT);
29+
TFE_OpAddInput(op, var_handle, status);
30+
ASSERT_EQ(TF_OK, TF_GetCode(status), TF_Message(status));
31+
int num_retvals = 1;
32+
var value_handle = new[] { IntPtr.Zero };
33+
TFE_Execute(op, value_handle, ref num_retvals, status);
34+
TFE_DeleteOp(op);
35+
36+
ASSERT_EQ(TF_OK, TF_GetCode(status), TF_Message(status));
37+
ASSERT_EQ(1, num_retvals);
38+
EXPECT_EQ(TF_FLOAT, TFE_TensorHandleDataType(value_handle[0]));
39+
EXPECT_EQ(0, TFE_TensorHandleNumDims(value_handle[0], status));
40+
ASSERT_EQ(TF_OK, TF_GetCode(status), TF_Message(status));
41+
var value = 0f; // new float[1];
42+
var t = TFE_TensorHandleResolve(value_handle[0], status);
43+
ASSERT_EQ(TF_OK, TF_GetCode(status), TF_Message(status));
44+
ASSERT_EQ(sizeof(float), (int)TF_TensorByteSize(t));
45+
memcpy(&value, TF_TensorData(t).ToPointer(), sizeof(float));
46+
c_api.TF_DeleteTensor(t);
47+
EXPECT_EQ(12.0f, value);
48+
49+
TFE_DeleteTensorHandle(var_handle);
50+
TFE_DeleteTensorHandle(value_handle[0]);
51+
TFE_DeleteContext(ctx);
52+
CHECK_EQ(TF_OK, TF_GetCode(status), TF_Message(status));
53+
TF_DeleteStatus(status);
54+
}
55+
}
56+
}

test/TensorFlowNET.UnitTest/Eager/CApi.Eager.cs

Lines changed: 42 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@ IntPtr TestMatrixTensorHandle()
1515
var dims = new long[] { 2, 2 };
1616
var data = new float[] { 1.0f, 2.0f, 3.0f, 4.0f };
1717
var t = c_api.TF_AllocateTensor(TF_FLOAT, dims, dims.Length, (ulong)data.Length * sizeof(float));
18-
memcpy(data, c_api.TF_TensorData(t), data.Length * sizeof(float));
18+
memcpy(c_api.TF_TensorData(t), data, data.Length * sizeof(float));
1919

2020
var status = c_api.TF_NewStatus();
2121
var th = c_api.TFE_NewTensorHandle(t, status);
@@ -79,5 +79,46 @@ IntPtr ShapeOp(IntPtr ctx, IntPtr a)
7979

8080
return op;
8181
}
82+
83+
unsafe IntPtr CreateVariable(IntPtr ctx, float value, IntPtr status)
84+
{
85+
var op = TFE_NewOp(ctx, "VarHandleOp", status);
86+
if (TF_GetCode(status) != TF_OK) return IntPtr.Zero;
87+
TFE_OpSetAttrType(op, "dtype", TF_FLOAT);
88+
TFE_OpSetAttrShape(op, "shape", new long[0], 0, status);
89+
TFE_OpSetAttrString(op, "container", "", 0);
90+
TFE_OpSetAttrString(op, "shared_name", "", 0);
91+
if (TF_GetCode(status) != TF_OK) return IntPtr.Zero;
92+
var var_handle = new IntPtr[1];
93+
int num_retvals = 1;
94+
TFE_Execute(op, var_handle, ref num_retvals, status);
95+
TFE_DeleteOp(op);
96+
if (TF_GetCode(status) != TF_OK) return IntPtr.Zero;
97+
CHECK_EQ(1, num_retvals);
98+
99+
// Assign 'value' to it.
100+
op = TFE_NewOp(ctx, "AssignVariableOp", status);
101+
if (TF_GetCode(status) != TF_OK) return IntPtr.Zero;
102+
TFE_OpSetAttrType(op, "dtype", TF_FLOAT);
103+
TFE_OpAddInput(op, var_handle[0], status);
104+
105+
// Convert 'value' to a TF_Tensor then a TFE_TensorHandle.
106+
var t = c_api.TF_AllocateTensor(TF_DataType.TF_FLOAT, new long[0], 0, sizeof(float));
107+
memcpy(TF_TensorData(t).ToPointer(), &value, TF_TensorByteSize(t));
108+
109+
var value_handle = c_api.TFE_NewTensorHandle(t, status);
110+
if (TF_GetCode(status) != TF_OK) return IntPtr.Zero;
111+
112+
TFE_OpAddInput(op, value_handle, status);
113+
if (TF_GetCode(status) != TF_OK) return IntPtr.Zero;
114+
115+
num_retvals = 0;
116+
c_api.TFE_Execute(op, null, ref num_retvals, status);
117+
TFE_DeleteOp(op);
118+
if (TF_GetCode(status) != TF_OK) return IntPtr.Zero;
119+
CHECK_EQ(0, num_retvals);
120+
121+
return var_handle[0];
122+
}
82123
}
83124
}

0 commit comments

Comments
 (0)