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

change shape to int[], comply with NumSharp v0.9 · Oceania2018/TensorFlow.NET@b71b4a2 · GitHub

Commit b71b4a2

Browse files
committed
change shape to int[], comply with NumSharp v0.9
1 parent 00003a6 commit b71b4a2

9 files changed

Lines changed: 109 additions & 81 deletions

File tree

‎src/TensorFlowNET.Core/Keras/Sequence.cs‎

Lines changed: 0 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -62,10 +62,6 @@ public NDArray pad_sequences(NDArray sequences,
6262
{
6363
switch(sequences[i])
6464
{
65-
case int[] data:
66-
for (int j = 0; j < nd.shape[1]; j++)
67-
nd[i, j] = j < data.Length ? data[j] : value;
68-
break;
6965
default:
7066
throw new NotImplementedException("pad_sequences");
7167
}

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

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -291,7 +291,7 @@ private static Tensor _may_reduce_to_scalar(bool keepdims, Tensor axis, Tensor o
291291
!keepdims &&
292292
axis == null)
293293
// We want set_shape to be reflected in the C API graph for when we run it.
294-
output.shape = new long[0];
294+
output.shape = new int[0];
295295
return output;
296296
}
297297

@@ -300,7 +300,7 @@ private static Tensor _may_reduce_to_scalar(bool keepdims, int[] axis, Tensor ou
300300
if (!common_shapes.has_fully_defined_shape(output) &&
301301
!keepdims &&
302302
axis == null)
303-
output.shape = new long[0];
303+
output.shape = new int[0];
304304
return output;
305305
}
306306

‎src/TensorFlowNET.Core/Sessions/_FetchHandler.cs‎

Lines changed: 18 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -52,7 +52,24 @@ public NDArray build_results(BaseSession session, NDArray[] tensor_values)
5252
{
5353
if (is_op)
5454
{
55-
full_values.Add(null);
55+
if(tensor_values.Length > 0)
56+
{
57+
switch (tensor_values[0].dtype.Name)
58+
{
59+
case "Int32":
60+
full_values.Add(float.NaN);
61+
break;
62+
case "Single":
63+
full_values.Add(float.NaN);
64+
break;
65+
default:
66+
throw new NotImplementedException($"build_results tensor_values[0] {tensor_values[0].dtype.Name}");
67+
}
68+
}
69+
else
70+
{
71+
full_values.Add(null);
72+
}
5673
}
5774
else
5875
{

‎src/TensorFlowNET.Core/Sessions/_FetchMapper.cs‎

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
using NumSharp;
22
using System;
33
using System.Collections.Generic;
4+
using System.Linq;
45
using System.Text;
56

67
namespace Tensorflow
@@ -21,7 +22,17 @@ public static _FetchMapper for_fetch(object fetch)
2122

2223
public virtual NDArray build_results(List<object> values)
2324
{
24-
return values.ToArray();
25+
var type = values[0].GetType();
26+
var nd = new NDArray(type, values.Count);
27+
28+
switch (type.Name)
29+
{
30+
case "Single":
31+
nd.SetData(values.Select(x => (float)x).ToArray());
32+
break;
33+
}
34+
35+
return nd;
2536
}
2637

2738
public virtual List<object> unique_fetches()

‎src/TensorFlowNET.Core/Tensors/Tensor.cs‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -46,7 +46,7 @@ public partial class Tensor : Python, IDisposable, ITensorOrOperation
4646

4747
private TF_Output? _tf_output;
4848

49-
public long[] shape
49+
public int[] shape
5050
{
5151
get
5252
{
@@ -63,15 +63,15 @@ public long[] shape
6363
dims[i] = c_api.TF_Dim(_handle, i);
6464
}
6565

66-
return dims;
66+
return dims.Select(x => Convert.ToInt32(x)).ToArray();
6767
}
6868

6969
set
7070
{
7171
if (value == null)
7272
c_api.TF_GraphSetTensorShape(this.graph, this._as_tf_output(), null, -1, status);
7373
else
74-
c_api.TF_GraphSetTensorShape(this.graph, this._as_tf_output(), value, value.Length, status);
74+
c_api.TF_GraphSetTensorShape(this.graph, this._as_tf_output(), value.Select(x => Convert.ToInt64(x)).ToArray(), value.Length, status);
7575
}
7676
}
7777

‎src/TensorFlowNET.Core/Tensors/tensor_util.cs‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -321,6 +321,11 @@ public static TensorShape to_shape(long[] dims)
321321
return new TensorShape(dims.Select(x => (int)x).ToArray());
322322
}
323323

324+
public static TensorShape to_shape(int[] dims)
325+
{
326+
return new TensorShape(dims);
327+
}
328+
324329
public static TensorShape as_shape(this Shape shape)
325330
{
326331
return new TensorShape(shape.Dimensions);

‎src/TensorFlowNET.Core/Variables/gen_state_ops.py.cs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -21,7 +21,7 @@ public class gen_state_ops
2121
/// <param name="container"></param>
2222
/// <param name="shared_name"></param>
2323
/// <returns></returns>
24-
public static Tensor variable_v2(long[] shape, TF_DataType dtype, string name = null, string container = "", string shared_name = "")
24+
public static Tensor variable_v2(int[] shape, TF_DataType dtype, string name = null, string container = "", string shared_name = "")
2525
{
2626
var _op = _op_def_lib._apply_op_helper("VariableV2", name: name, args: new { dtype, shape, container, shared_name });
2727

‎src/TensorFlowNET.Core/Variables/state_ops.cs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,7 +15,7 @@ public class state_ops
1515
/// <param name="container"></param>
1616
/// <param name="shared_name"></param>
1717
/// <returns></returns>
18-
public static Tensor variable_op_v2(long[] shape,
18+
public static Tensor variable_op_v2(int[] shape,
1919
TF_DataType dtype,
2020
string name = "Variable",
2121
string container = "",

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL