| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 4b4f7f8 commit a22e92d
6 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,31 @@ | |||
| 1 | + using System; | ||
| 2 | + using System.Collections.Generic; | ||
| 3 | + | ||
| 4 | + namespace Tensorflow.NumPy | ||
| 5 | + { | ||
| 6 | + public partial class NumPyImpl | ||
| 7 | + { | ||
| 8 | + public NDArray average(NDArray a, int axis = -1, NDArray? weights = null, bool returned = false) | ||
| 9 | + { | ||
| 10 | + var dtype = NumPyUtils.GetResultType(a.dtype, np.float64); | ||
| 11 | + if(weights is null) | ||
| 12 | + { | ||
| 13 | + var tensorA = math_ops.cast(a, dtype); | ||
| 14 | + var nd = math_ops.reduce_mean(tensorA, axis); | ||
| 15 | + return new NDArray(nd); | ||
| 16 | + } | ||
| 17 | + else | ||
| 18 | + { | ||
| 19 | + var tensorW = math_ops.cast(weights, dtype); | ||
| 20 | + if(a.rank != weights.rank) | ||
| 21 | + { | ||
| 22 | + var weights_sum = math_ops.reduce_sum(tensorW); | ||
| 23 | + var axes = ops.convert_to_tensor(new[,] { { axis }, { 0 } }); | ||
| 24 | + var avg = math_ops.tensordot(a, weights, axes) / weights_sum; | ||
| 25 | + } | ||
| 26 | + | ||
| 27 | + throw new NotImplementedException(""); | ||
| 28 | + } | ||
| 29 | + } | ||
| 30 | + } | ||
| 31 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -14,5 +14,9 @@ public partial class np | |||
| 14 | 14 | ||
| 15 | 15 | [AutoNumPy] | |
| 16 | 16 | public static NDArray amax(NDArray x, int axis = 0) => new NDArray(tf.math.argmax(x, axis)); | |
| 17 | + | ||
| 18 | + [AutoNumPy] | ||
| 19 | + public static NDArray average(NDArray a, int axis = -1, NDArray? weights = null, bool returned = false) | ||
| 20 | + => tf.numpy.average(a, axis: axis, weights: weights, returned: returned); | ||
| 17 | 21 | } | |
| 18 | 22 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,19 @@ | |||
| 1 | + using System; | ||
| 2 | + using System.Text; | ||
| 3 | + | ||
| 4 | + namespace Tensorflow.NumPy | ||
| 5 | + { | ||
| 6 | + internal class NumPyUtils | ||
| 7 | + { | ||
| 8 | + public static TF_DataType GetResultType(params TF_DataType[] dtypes) | ||
| 9 | + { | ||
| 10 | + var resultDType = dtypes[0]; | ||
| 11 | + for(int i = 1; i < dtypes.Length; i++) | ||
| 12 | + { | ||
| 13 | + if (dtypes[i].get_datatype_size() > resultDType.get_datatype_size()) | ||
| 14 | + resultDType = dtypes[i]; | ||
| 15 | + } | ||
| 16 | + return resultDType; | ||
| 17 | + } | ||
| 18 | + } | ||
| 19 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -929,6 +929,72 @@ Tensor _tensordot_reshape(Tensor a, int[] axes, bool flipped = false) | |||
| 929 | 929 | throw new NotImplementedException("tensordot"); | |
| 930 | 930 | } | |
| 931 | 931 | ||
| 932 | + public static Tensor tensordot(Tensor x, Tensor y, Tensor axes, string name = null) | ||
| 933 | + { | ||
| 934 | + Tensor _tensordot_reshape(Tensor a, int[] axes, bool flipped = false) | ||
| 935 | + { | ||
| 936 | + if (a.shape.IsFullyDefined && isinstance(axes, (typeof(List<object>), typeof(Tuple)))) | ||
| 937 | + { | ||
| 938 | + var shape_a = a.shape.dims; | ||
| 939 | + | ||
| 940 | + // axes | ||
| 941 | + int iter = 0; | ||
| 942 | + foreach (int i in axes) | ||
| 943 | + { | ||
| 944 | + if (i >= 0) | ||
| 945 | + axes[0 + iter] = i; | ||
| 946 | + else | ||
| 947 | + axes[0 + iter] = i + len(shape_a); | ||
| 948 | + iter++; | ||
| 949 | + } | ||
| 950 | + | ||
| 951 | + // free | ||
| 952 | + int[] free = { }; | ||
| 953 | + iter = 0; | ||
| 954 | + foreach (int i in Enumerable.Range(0, len(axes))) | ||
| 955 | + if (!Array.Exists(axes, i => i == i)) | ||
| 956 | + free[free.Length] = i; | ||
| 957 | + | ||
| 958 | + // free_dims | ||
| 959 | + int[] free_dims = { }; | ||
| 960 | + foreach (int i in free) | ||
| 961 | + free_dims[free_dims.Length] = (int)shape_a[i]; | ||
| 962 | + | ||
| 963 | + int prod_free = (int)np.prod(free_dims); | ||
| 964 | + | ||
| 965 | + // prod_axes | ||
| 966 | + int[] prod_axes_pre = { }; | ||
| 967 | + foreach (int i in axes) | ||
| 968 | + prod_axes_pre[prod_axes_pre.Length] = (int)shape_a[i]; | ||
| 969 | + int prod_axes = (int)np.prod(prod_axes_pre); | ||
| 970 | + | ||
| 971 | + // perm | ||
| 972 | + Tensor perm; | ||
| 973 | + if (flipped) | ||
| 974 | + perm = ops.convert_to_tensor(list(free)) + ops.convert_to_tensor(free); | ||
| 975 | + else | ||
| 976 | + perm = ops.convert_to_tensor(list(free)) + ops.convert_to_tensor(free) | ||
| 977 | + + ops.convert_to_tensor(list(axes)); | ||
| 978 | + | ||
| 979 | + // new_shape | ||
| 980 | + Shape new_shape; | ||
| 981 | + if (flipped) | ||
| 982 | + new_shape = new Shape(new int[] { prod_axes, prod_free }); | ||
| 983 | + else | ||
| 984 | + new_shape = new Shape(new int[] { prod_free, prod_axes }); | ||
| 985 | + } | ||
| 986 | + | ||
| 987 | + throw new NotImplementedException("_tensordot_reshape"); | ||
| 988 | + } | ||
| 989 | + | ||
| 990 | + return tf_with(ops.name_scope(name, "Tensordot", new { x, y, axes }), scope => | ||
| 991 | + { | ||
| 992 | + name = scope; | ||
| 993 | + var (a_axes, b_axes) = (axes[0], axes[1]); | ||
| 994 | + return x; | ||
| 995 | + }); | ||
| 996 | + } | ||
| 997 | + | ||
| 932 | 998 | public static Tensor truediv(Tensor x, Tensor y, string name = null) | |
| 933 | 999 | => _truediv_python3(x, y, name); | |
| 934 | 1000 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -78,7 +78,7 @@ public partial class Tensor : DisposableObject, | |||
| 78 | 78 | /// <summary> | |
| 79 | 79 | /// The name of the device on which this tensor will be produced, or null. | |
| 80 | 80 | /// </summary> | |
| 81 | - public virtual string Device => op.Device; | ||
| 81 | + public virtual string Device => op?.Device; | ||
| 82 | 82 | public long[] dims => shape.dims; | |
| 83 | 83 | ||
| 84 | 84 | /// <summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,32 @@ | |||
| 1 | + using Microsoft.VisualStudio.TestTools.UnitTesting; | ||
| 2 | + using System; | ||
| 3 | + using System.Collections.Generic; | ||
| 4 | + using System.Linq; | ||
| 5 | + using System.Text; | ||
| 6 | + using Tensorflow; | ||
| 7 | + using Tensorflow.NumPy; | ||
| 8 | + | ||
| 9 | + namespace TensorFlowNET.UnitTest.NumPy | ||
| 10 | + { | ||
| 11 | + /// <summary> | ||
| 12 | + /// https://numpy.org/doc/stable/reference/routines.statistics.html | ||
| 13 | + /// </summary> | ||
| 14 | + [TestClass] | ||
| 15 | + public class StatisticsTest : EagerModeTestBase | ||
| 16 | + { | ||
| 17 | + [TestMethod] | ||
| 18 | + public void average() | ||
| 19 | + { | ||
| 20 | + var data = np.arange(1, 5); | ||
| 21 | + var avg = np.average(data); | ||
| 22 | + Assert.AreEqual(avg, 2.5); | ||
| 23 | + | ||
| 24 | + data = np.arange(6).reshape((3, 2)); | ||
| 25 | + avg = np.average(data, axis: 1); | ||
| 26 | + assertAllEqual(avg.ToArray<double>(), new[] { 0.5, 2.5, 4.5 }); | ||
| 27 | + | ||
| 28 | + // avg = np.average(data, axis: 1, weights: new[] { 1.0 / 4, 3.0 / 4 }); | ||
| 29 | + // assertAllEqual(avg.ToArray<double>(), new[] { 0.75, 2.75, 4.75 }); | ||
| 30 | + } | ||
| 31 | + } | ||
| 32 | + } | ||
| Back | FazBrowse Home | New Git URL |
0 commit comments