| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent f566505 commit 432ae20
7 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,13 +1,15 @@ | |||
| 1 | 1 | using MethodBoundaryAspect.Fody.Attributes; | |
| 2 | 2 | using System; | |
| 3 | 3 | using System.Collections.Generic; | |
| 4 | + using System.Diagnostics; | ||
| 4 | 5 | using System.Linq; | |
| 5 | 6 | using Tensorflow.Eager; | |
| 6 | 7 | using Tensorflow.Functions; | |
| 7 | 8 | using static Tensorflow.Binding; | |
| 8 | 9 | ||
| 9 | 10 | namespace Tensorflow.NumPy | |
| 10 | 11 | { | |
| 12 | + [DebuggerStepThrough] | ||
| 11 | 13 | public sealed class AutoNumPyAttribute : OnMethodBoundaryAspect | |
| 12 | 14 | { | |
| 13 | 15 | bool _changedMode = false; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -10,7 +10,7 @@ public partial class np | |||
| 10 | 10 | { | |
| 11 | 11 | [AutoNumPy] | |
| 12 | 12 | public static NDArray argmax(NDArray a, Axis axis = null) | |
| 13 | - => new NDArray(math_ops.argmax(a, axis)); | ||
| 13 | + => new NDArray(math_ops.argmax(a, axis ?? 0)); | ||
| 14 | 14 | ||
| 15 | 15 | [AutoNumPy] | |
| 16 | 16 | public static NDArray argsort(NDArray a, Axis axis = null) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -31,7 +31,7 @@ public NDArray(long[] value, Shape? shape = null) | |||
| 31 | 31 | public NDArray(IntPtr address, Shape shape, TF_DataType dtype) | |
| 32 | 32 | : base(address, shape, dtype) { NewEagerTensorHandle(); } | |
| 33 | 33 | ||
| 34 | - public NDArray(Tensor tensor, bool eval = true) : base(tensor.Handle) | ||
| 34 | + public NDArray(Tensor tensor, bool clone = false) : base(tensor.Handle, clone: clone) | ||
| 35 | 35 | { | |
| 36 | 36 | if (_handle is null) | |
| 37 | 37 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -35,7 +35,7 @@ public int InputListLength(string name) | |||
| 35 | 35 | tf.Status.Check(true); | |
| 36 | 36 | return num; | |
| 37 | 37 | } | |
| 38 | - public int NumInputs => c_api.TF_OperationNumInputs(_handle); | ||
| 38 | + public int NumInputs => _handle == IntPtr.Zero ? -1 : c_api.TF_OperationNumInputs(_handle); | ||
| 39 | 39 | private TF_DataType[] _input_types => _inputs_val._inputs.Select(x => x.dtype).ToArray(); | |
| 40 | 40 | ||
| 41 | 41 | protected InputList _inputs_val; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -23,7 +23,7 @@ namespace Tensorflow | |||
| 23 | 23 | { | |
| 24 | 24 | public partial class Operation | |
| 25 | 25 | { | |
| 26 | - public int NumOutputs => c_api.TF_OperationNumOutputs(_handle); | ||
| 26 | + public int NumOutputs => _handle == IntPtr.Zero ? -1 : c_api.TF_OperationNumOutputs(_handle); | ||
| 27 | 27 | public TF_DataType OutputType(int index) => c_api.TF_OperationOutputType(_tf_output(index)); | |
| 28 | 28 | ||
| 29 | 29 | public int OutputListLength(string name) | |
@@ -38,7 +38,7 @@ public int OutputListLength(string name) | |||
| 38 | 38 | public virtual Tensor[] outputs => _outputs; | |
| 39 | 39 | public Tensor output => _outputs.FirstOrDefault(); | |
| 40 | 40 | ||
| 41 | - public int NumControlOutputs => c_api.TF_OperationNumControlOutputs(_handle); | ||
| 41 | + public int NumControlOutputs => _handle == IntPtr.Zero ? -1 : c_api.TF_OperationNumControlOutputs(_handle); | ||
| 42 | 42 | ||
| 43 | 43 | public int OutputNumConsumers(int index) => c_api.TF_OperationOutputNumConsumers(new TF_Output(_handle, index)); | |
| 44 | 44 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -39,9 +39,12 @@ public Tensor() | |||
| 39 | 39 | /// Create a Tensor object from an existing TF handle | |
| 40 | 40 | /// </summary> | |
| 41 | 41 | /// <param name="handle">Handle to a <see cref="Tensor"/> object.</param> | |
| 42 | - public Tensor(SafeTensorHandle handle) | ||
| 42 | + public unsafe Tensor(SafeTensorHandle handle, bool clone = false) | ||
| 43 | 43 | { | |
| 44 | 44 | _handle = handle; | |
| 45 | + if (clone) | ||
| 46 | + _handle = TF_NewTensor(shape, dtype, data: TensorDataPointer.ToPointer()); | ||
| 47 | + | ||
| 45 | 48 | isCreatedInGraphMode = !tf.executing_eagerly(); | |
| 46 | 49 | } | |
| 47 | 50 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -55,7 +55,7 @@ protected NDArray GetNDArray(TF_DataType dtype) | |||
| 55 | 55 | return new NDArray(str, shape); | |
| 56 | 56 | } | |
| 57 | 57 | ||
| 58 | - return new NDArray(this); | ||
| 58 | + return new NDArray(this, clone: true); | ||
| 59 | 59 | } | |
| 60 | 60 | ||
| 61 | 61 | /// <summary> | |
| Back | FazBrowse Home | New Git URL |
0 commit comments