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

fix all unit test. · feelsyt/TensorFlow.NET@254ba33 · GitHub

fix all unit test. · feelsyt/TensorFlow.NET@254ba33 · GitHub
Skip to content

Navigation Menu

Commit 254ba33

Browse files
committed
fix all unit test.
1 parent bcb28f3 commit 254ba33

24 files changed

Lines changed: 297 additions & 199 deletions

File tree

‎src/TensorFlowNET.Console/MemoryMonitor.cs‎

Lines changed: 16 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -12,13 +12,26 @@ public class MemoryMonitor
1212
{
1313
public void WarmUp()
1414
{
15+
var x1 = tf.Variable(10, name: "x");
16+
17+
tf.compat.v1.disable_eager_execution();
18+
var input = np.array(4);
19+
var nd = tf.reshape(input, new int[] { 1, 1});
20+
var z = nd[0, 0];
1521
while (true)
1622
{
17-
var ones = np.ones((128, 128));
18-
Thread.Sleep(1);
23+
var x = tf.placeholder(tf.float64, shape: (1024, 1024));
24+
var log = tf.log(x);
25+
26+
using (var sess = tf.Session())
27+
{
28+
var ones = np.ones((1024, 1024), dtype: np.float64);
29+
var o = sess.run(log, new FeedItem(x, ones));
30+
}
31+
// Thread.Sleep(1);
1932
}
2033

21-
TensorShape shape = (1, 32, 32, 3);
34+
Shape shape = (1, 32, 32, 3);
2235
np.arange(shape.size).astype(np.float32).reshape(shape.dims);
2336

2437
print($"tensorflow native version: v{tf.VERSION}");

‎src/TensorFlowNET.Core/APIs/tf.math.cs‎

Lines changed: 7 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,9 @@ public Tensor log(Tensor x, string name = null)
3333
public Tensor erf(Tensor x, string name = null)
3434
=> math_ops.erf(x, name);
3535

36+
public Tensor sum(Tensor x, Axis? axis = null, string name = null)
37+
=> math_ops.reduce_sum(x, axis: axis, name: name);
38+
3639
/// <summary>
3740
///
3841
/// </summary>
@@ -492,40 +495,21 @@ public Tensor reduce_all(Tensor input_tensor, Axis? axis = null, bool keepdims =
492495
public Tensor reduce_prod(Tensor input_tensor, Axis? axis = null, bool keepdims = false, string name = null)
493496
=> math_ops.reduce_prod(input_tensor, axis: axis, keepdims: keepdims, name: name);
494497

495-
/// <summary>
496-
/// Computes the sum of elements across dimensions of a tensor.
497-
/// </summary>
498-
/// <param name="input_tensors"></param>
499-
/// <param name="axis"></param>
500-
/// <param name="keepdims"></param>
501-
/// <param name="name"></param>
502-
/// <returns></returns>
503-
public Tensor reduce_sum(Tensor[] input_tensors, int? axis = null, bool keepdims = false, string name = null)
504-
=> math_ops.reduce_sum(input_tensors, axis: axis, keepdims: keepdims, name: name);
505-
506498
/// <summary>
507499
/// Computes the sum of elements across dimensions of a tensor.
508500
/// </summary>
509501
/// <param name="input"></param>
510502
/// <param name="axis"></param>
511503
/// <returns></returns>
512-
public Tensor reduce_sum(Tensor input, int? axis = null, int? reduction_indices = null,
504+
public Tensor reduce_sum(Tensor input, Axis? axis = null, Axis? reduction_indices = null,
513505
bool keepdims = false, string name = null)
514506
{
515-
if (!axis.HasValue && reduction_indices.HasValue && !keepdims)
516-
return math_ops.reduce_sum(input, reduction_indices.Value);
517-
else if (axis.HasValue && !reduction_indices.HasValue && !keepdims)
518-
return math_ops.reduce_sum(input, axis.Value);
519-
else if (axis.HasValue && !reduction_indices.HasValue && keepdims)
520-
return math_ops.reduce_sum(input, keepdims: keepdims, axis: axis.Value, name: name);
507+
if(keepdims)
508+
return math_ops.reduce_sum(input, axis: constant_op.constant(axis ?? reduction_indices), keepdims: keepdims, name: name);
521509
else
522-
return math_ops.reduce_sum(input, keepdims: keepdims, name: name);
510+
return math_ops.reduce_sum(input, axis: constant_op.constant(axis ?? reduction_indices));
523511
}
524512

525-
public Tensor reduce_sum(Tensor input, Shape axis, int? reduction_indices = null,
526-
bool keepdims = false, string name = null)
527-
=> math_ops.reduce_sum(input, axis, keepdims: keepdims, name: name);
528-
529513
/// <summary>
530514
/// Computes the maximum of elements across dimensions of a tensor.
531515
/// </summary>

‎src/TensorFlowNET.Core/Gradients/nn_grad.cs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -70,7 +70,7 @@ public static Tensor[] _SoftmaxGrad(Operation op, Tensor[] grads)
7070

7171
var softmax = op.outputs[0];
7272
var mul = grad_softmax * softmax;
73-
var sum_channels = math_ops.reduce_sum(mul, -1, keepdims: true);
73+
var sum_channels = math_ops.reduce_sum(mul, axis: constant_op.constant(-1), keepdims: true);
7474
var sub = grad_softmax - sum_channels;
7575
return new Tensor[] { sub * softmax };
7676
}

‎src/TensorFlowNET.Core/NumPy/Axis.cs‎

Lines changed: 28 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,20 @@
1-
using System;
1+
/*****************************************************************************
2+
Copyright 2021 Haiping Chen. 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 System;
218
using System.Collections.Generic;
319
using System.Linq;
420
using System.Text;
@@ -7,6 +23,8 @@ namespace Tensorflow
723
{
824
public record Axis(params int[] axis)
925
{
26+
public int size => axis == null ? -1 : axis.Length;
27+
1028
public int this[int index] => axis[index];
1129

1230
public static implicit operator int[]?(Axis axis)
@@ -16,19 +34,22 @@ public static implicit operator Axis(int axis)
1634
=> new Axis(axis);
1735

1836
public static implicit operator Axis((int, int) axis)
19-
=> new Axis(axis);
37+
=> new Axis(axis.Item1, axis.Item2);
2038

2139
public static implicit operator Axis((int, int, int) axis)
22-
=> new Axis(axis);
40+
=> new Axis(axis.Item1, axis.Item2, axis.Item3);
2341

2442
public static implicit operator Axis(int[] axis)
2543
=> new Axis(axis);
2644

27-
public static implicit operator Axis(long[] shape)
28-
=> new Axis(shape.Select(x => (int)x).ToArray());
45+
public static implicit operator Axis(long[] axis)
46+
=> new Axis(axis.Select(x => (int)x).ToArray());
47+
48+
public static implicit operator Axis(Shape axis)
49+
=> new Axis(axis.dims.Select(x => (int)x).ToArray());
2950

30-
public static implicit operator Axis(Shape shape)
31-
=> new Axis(shape.dims.Select(x => (int)x).ToArray());
51+
public static implicit operator Tensor(Axis axis)
52+
=> constant_op.constant(axis);
3253
}
3354
}
3455

‎src/TensorFlowNET.Core/NumPy/NDArray.Implicit.cs‎

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,12 +6,22 @@ namespace Tensorflow.NumPy
66
{
77
public partial class NDArray
88
{
9+
public void Deconstruct(out byte blue, out byte green, out byte red)
10+
{
11+
blue = (byte)dims[0];
12+
green = (byte)dims[1];
13+
red = (byte)dims[2];
14+
}
15+
916
public static implicit operator NDArray(Array array)
1017
=> new NDArray(array);
1118

1219
public static implicit operator bool(NDArray nd)
1320
=> nd._tensor.ToArray<bool>()[0];
1421

22+
public static implicit operator byte(NDArray nd)
23+
=> nd._tensor.ToArray<byte>()[0];
24+
1525
public static implicit operator byte[](NDArray nd)
1626
=> nd.ToByteArray();
1727

‎src/TensorFlowNET.Core/NumPy/NDArray.Index.cs‎

Lines changed: 23 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,22 @@ public NDArray this[params int[] index]
3030

3131
set
3232
{
33-
33+
var offset = ShapeHelper.GetOffset(shape, index);
34+
unsafe
35+
{
36+
if (dtype == TF_DataType.TF_BOOL)
37+
*((bool*)data + offset) = value;
38+
else if (dtype == TF_DataType.TF_UINT8)
39+
*((byte*)data + offset) = value;
40+
else if (dtype == TF_DataType.TF_INT32)
41+
*((int*)data + offset) = value;
42+
else if (dtype == TF_DataType.TF_INT64)
43+
*((long*)data + offset) = value;
44+
else if (dtype == TF_DataType.TF_FLOAT)
45+
*((float*)data + offset) = value;
46+
else if (dtype == TF_DataType.TF_DOUBLE)
47+
*((double*)data + offset) = value;
48+
}
3449
}
3550
}
3651

@@ -43,7 +58,13 @@ public NDArray this[params Slice[] slices]
4358

4459
set
4560
{
46-
61+
var pos = _tensor[slices];
62+
var len = value.bytesize;
63+
unsafe
64+
{
65+
System.Buffer.MemoryCopy(value.data.ToPointer(), pos.TensorDataPointer.ToPointer(), len, len);
66+
}
67+
// _tensor[slices].assign(constant_op.constant(value));
4768
}
4869
}
4970

‎src/TensorFlowNET.Core/NumPy/Numpy.Math.cs‎

Lines changed: 6 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -10,18 +10,18 @@ namespace Tensorflow.NumPy
1010
public partial class np
1111
{
1212
public static NDArray log(NDArray x)
13-
=> throw new NotImplementedException("");
13+
=> tf.log(x);
1414

1515
public static NDArray prod(NDArray array, Axis? axis = null, Type? dtype = null, bool keepdims = false)
16-
=> tf.reduce_prod(ops.convert_to_tensor(array), axis: axis);
16+
=> tf.reduce_prod(array, axis: axis);
1717

1818
public static NDArray prod<T>(params T[] array) where T : unmanaged
1919
=> tf.reduce_prod(ops.convert_to_tensor(array));
2020

21-
public static NDArray multiply(in NDArray x1, in NDArray x2)
22-
=> throw new NotImplementedException("");
21+
public static NDArray multiply(NDArray x1, NDArray x2)
22+
=> tf.multiply(x1, x2);
2323

24-
public static NDArray sum(NDArray x1)
25-
=> throw new NotImplementedException("");
24+
public static NDArray sum(NDArray x1, Axis? axis = null)
25+
=> tf.math.sum(x1, axis);
2626
}
2727
}
Lines changed: 87 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,87 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Linq;
4+
using System.Text;
5+
6+
namespace Tensorflow.NumPy
7+
{
8+
internal class ShapeHelper
9+
{
10+
public static long GetSize(Shape shape)
11+
{
12+
// scalar
13+
if (shape.ndim == 0)
14+
return 1;
15+
16+
var computed = 1L;
17+
for (int i = 0; i < shape.ndim; i++)
18+
{
19+
var val = shape.dims[i];
20+
if (val == 0)
21+
return 0;
22+
else if (val < 0)
23+
continue;
24+
computed *= val;
25+
}
26+
27+
return computed;
28+
}
29+
30+
public static long[] GetStrides(Shape shape)
31+
{
32+
var strides = new long[shape.ndim];
33+
34+
if (shape.ndim == 0)
35+
return strides;
36+
37+
strides[strides.Length - 1] = 1;
38+
for (int idx = strides.Length - 1; idx >= 1; idx--)
39+
strides[idx - 1] = strides[idx] * shape.dims[idx];
40+
41+
return strides;
42+
}
43+
44+
public static bool Equals(Shape shape, object target)
45+
{
46+
switch (target)
47+
{
48+
case Shape shape1:
49+
if (shape.ndim == -1 && shape1.ndim == -1)
50+
return false;
51+
else if (shape.ndim != shape1.ndim)
52+
return false;
53+
return Enumerable.SequenceEqual(shape1.dims, shape.dims);
54+
case long[] shape2:
55+
if (shape.ndim != shape2.Length)
56+
return false;
57+
return Enumerable.SequenceEqual(shape.dims, shape2);
58+
default:
59+
return false;
60+
}
61+
}
62+
63+
public static string ToString(Shape shape)
64+
{
65+
return shape.ndim switch
66+
{
67+
-1 => "<unknown>",
68+
0 => "()",
69+
1 => $"({shape.dims[0]},)",
70+
_ => $"({string.Join(", ", shape.dims).Replace("-1", "None")})"
71+
};
72+
}
73+
74+
public static long GetOffset(Shape shape, params int[] indices)
75+
{
76+
if (shape.ndim == 0 && indices.Length == 1)
77+
return indices[0];
78+
79+
long offset = 0;
80+
var strides = shape.strides;
81+
for (int i = 0; i < indices.Length; i++)
82+
offset += strides[i] * indices[i];
83+
84+
return offset;
85+
}
86+
}
87+
}

‎src/TensorFlowNET.Core/Numpy/NDArray.cs‎

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,20 @@
1-
using System;
1+
/*****************************************************************************
2+
Copyright 2021 Haiping Chen. 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 System;
218
using System.Collections.Generic;
319
using System.Linq;
420
using System.Text;

‎src/TensorFlowNET.Core/Numpy/Numpy.cs‎

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,20 @@
1-
using System;
1+
/*****************************************************************************
2+
Copyright 2021 Haiping Chen. 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 System;
218
using System.Collections;
319
using System.Collections.Generic;
420
using System.Numerics;

0 commit comments

Comments
 (0)

Footer

© 2026 GitHub, Inc.

Back | FazBrowse Home | New Git URL