| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 6dd15f7 commit 70f873e
6 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -3,4 +3,6 @@ | |||
| 3 | 3 | global using System.Text; | |
| 4 | 4 | global using System.Collections; | |
| 5 | 5 | global using System.Data; | |
| 6 | - global using System.Linq; | ||
| 6 | + global using System.Linq; | ||
| 7 | + global using Tensorflow.Keras.Engine; | ||
| 8 | + global using Tensorflow.Framework.Models; | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,5 +1,6 @@ | |||
| 1 | 1 | using System; | |
| 2 | 2 | using Tensorflow.Framework.Models; | |
| 3 | + using Tensorflow.Keras.Engine; | ||
| 3 | 4 | using Tensorflow.Keras.Layers.Rnn; | |
| 4 | 5 | using Tensorflow.NumPy; | |
| 5 | 6 | using static Google.Protobuf.Reflection.FieldDescriptorProto.Types; | |
@@ -135,7 +136,7 @@ public ILayer EinsumDense(string equation, | |||
| 135 | 136 | public ILayer GlobalMaxPooling1D(string data_format = "channels_last"); | |
| 136 | 137 | public ILayer GlobalMaxPooling2D(string data_format = "channels_last"); | |
| 137 | 138 | ||
| 138 | - public Tensors Input(Shape shape = null, | ||
| 139 | + public KerasTensor Input(Shape shape = null, | ||
| 139 | 140 | int batch_size = -1, | |
| 140 | 141 | string name = null, | |
| 141 | 142 | TF_DataType dtype = TF_DataType.DtInvalid, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,40 @@ | |||
| 1 | + namespace Tensorflow.Keras.Engine; | ||
| 2 | + | ||
| 3 | + /// <summary> | ||
| 4 | + /// A representation of a Keras in/output during Functional API construction. | ||
| 5 | + /// </summary> | ||
| 6 | + public class KerasTensor | ||
| 7 | + { | ||
| 8 | + private Tensor _tensor; | ||
| 9 | + public void SetTensor(Tensors tensor) | ||
| 10 | + => _tensor = tensor; | ||
| 11 | + | ||
| 12 | + private TensorSpec _type_spec; | ||
| 13 | + private string _name; | ||
| 14 | + | ||
| 15 | + public KerasTensor(TensorSpec type_spec, string name = null) | ||
| 16 | + { | ||
| 17 | + _type_spec = type_spec; | ||
| 18 | + _name = name; | ||
| 19 | + } | ||
| 20 | + | ||
| 21 | + public static KerasTensor from_tensor(Tensor tensor) | ||
| 22 | + { | ||
| 23 | + var type_spec = tensor.ToTensorSpec(); | ||
| 24 | + var kt = new KerasTensor(type_spec, name: tensor.name); | ||
| 25 | + kt.SetTensor(tensor); | ||
| 26 | + return kt; | ||
| 27 | + } | ||
| 28 | + | ||
| 29 | + public static implicit operator Tensors(KerasTensor kt) | ||
| 30 | + => kt._tensor; | ||
| 31 | + | ||
| 32 | + public static implicit operator Tensor(KerasTensor kt) | ||
| 33 | + => kt._tensor; | ||
| 34 | + | ||
| 35 | + public static implicit operator KerasTensor(Tensor tensor) | ||
| 36 | + => from_tensor(tensor); | ||
| 37 | + | ||
| 38 | + public static implicit operator KerasTensor(Tensors tensors) | ||
| 39 | + => from_tensor(tensors.First()); | ||
| 40 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -76,7 +76,7 @@ public void track_variable(IVariableV1 v) | |||
| 76 | 76 | _GRAPH_VARIABLES[graph.graph_key] = v; | |
| 77 | 77 | } | |
| 78 | 78 | ||
| 79 | - public Tensor placeholder(Shape shape = null, | ||
| 79 | + public KerasTensor placeholder(Shape shape = null, | ||
| 80 | 80 | int ndim = -1, | |
| 81 | 81 | TF_DataType dtype = TF_DataType.DtInvalid, | |
| 82 | 82 | bool sparse = false, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -4,4 +4,5 @@ | |||
| 4 | 4 | global using System.Linq; | |
| 5 | 5 | global using static Tensorflow.Binding; | |
| 6 | 6 | global using static Tensorflow.KerasApi; | |
| 7 | - global using Tensorflow.NumPy; | ||
| 7 | + global using Tensorflow.NumPy; | ||
| 8 | + global using Tensorflow.Keras.Engine; | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -466,7 +466,7 @@ public ILayer Flatten(string data_format = null) | |||
| 466 | 466 | /// In this case, values of 'None' in the 'shape' argument represent ragged dimensions. For more information about RaggedTensors, see this guide. | |
| 467 | 467 | /// </param> | |
| 468 | 468 | /// <returns>A tensor.</returns> | |
| 469 | - public Tensors Input(Shape shape = null, | ||
| 469 | + public KerasTensor Input(Shape shape = null, | ||
| 470 | 470 | int batch_size = -1, | |
| 471 | 471 | string name = null, | |
| 472 | 472 | TF_DataType dtype = TF_DataType.DtInvalid, | |
| Back | FazBrowse Home | New Git URL |
0 commit comments