| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 1dd95bd commit 9f2adcf
5 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,20 @@ | |||
| 1 | + using System; | ||
| 2 | + using System.Collections.Generic; | ||
| 3 | + using System.Text; | ||
| 4 | + using Tensorflow.Keras.Utils; | ||
| 5 | + | ||
| 6 | + namespace Tensorflow.Keras.Engine | ||
| 7 | + { | ||
| 8 | + public partial class Model | ||
| 9 | + { | ||
| 10 | + /// <summary> | ||
| 11 | + /// Prints a string summary of the network. | ||
| 12 | + /// </summary> | ||
| 13 | + public void summary(int line_length = -1, float[] positions = null) | ||
| 14 | + { | ||
| 15 | + layer_utils.print_summary(this, | ||
| 16 | + line_length: line_length, | ||
| 17 | + positions: positions); | ||
| 18 | + } | ||
| 19 | + } | ||
| 20 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,18 @@ | |||
| 1 | + using System; | ||
| 2 | + using System.Collections.Generic; | ||
| 3 | + using System.Text; | ||
| 4 | + | ||
| 5 | + namespace Tensorflow.Keras.Engine | ||
| 6 | + { | ||
| 7 | + public partial class Node | ||
| 8 | + { | ||
| 9 | + public IEnumerable<(Layer, int, int, Tensor)> iterate_inbound() | ||
| 10 | + { | ||
| 11 | + foreach(var kt in KerasInputs) | ||
| 12 | + { | ||
| 13 | + var (layer, node_index, tensor_index) = kt.KerasHistory; | ||
| 14 | + yield return (layer, node_index, tensor_index, kt); | ||
| 15 | + } | ||
| 16 | + } | ||
| 17 | + } | ||
| 18 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -2,12 +2,18 @@ | |||
| 2 | 2 | using System; | |
| 3 | 3 | using System.Collections.Generic; | |
| 4 | 4 | using System.Text; | |
| 5 | + using Tensorflow; | ||
| 6 | + using Tensorflow.Keras.Engine; | ||
| 7 | + using Tensorflow.Keras.Layers; | ||
| 5 | 8 | using static Tensorflow.Binding; | |
| 6 | 9 | ||
| 7 | 10 | namespace TensorFlowNET.UnitTest | |
| 8 | 11 | { | |
| 9 | 12 | public class EagerModeTestBase : PythonTest | |
| 10 | 13 | { | |
| 14 | + protected KerasApi keras = tf.keras; | ||
| 15 | + protected LayersApi layers = tf.keras.layers; | ||
| 16 | + | ||
| 11 | 17 | [TestInitialize] | |
| 12 | 18 | public void TestInit() | |
| 13 | 19 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -16,13 +16,30 @@ namespace TensorFlowNET.UnitTest.Keras | |||
| 16 | 16 | [TestClass] | |
| 17 | 17 | public class LayersTest : EagerModeTestBase | |
| 18 | 18 | { | |
| 19 | + | ||
| 19 | 20 | [TestMethod] | |
| 20 | 21 | public void Sequential() | |
| 21 | 22 | { | |
| 22 | 23 | var model = tf.keras.Sequential(); | |
| 23 | 24 | model.add(tf.keras.Input(shape: 16)); | |
| 24 | 25 | } | |
| 25 | 26 | ||
| 27 | + [TestMethod] | ||
| 28 | + public void Functional() | ||
| 29 | + { | ||
| 30 | + var inputs = keras.Input(shape: 784); | ||
| 31 | + Assert.AreEqual((None, 784), inputs.TensorShape); | ||
| 32 | + | ||
| 33 | + var dense = layers.Dense(64, activation: "relu"); | ||
| 34 | + var x = dense.Apply(inputs); | ||
| 35 | + | ||
| 36 | + x = layers.Dense(64, activation: "relu").Apply(x); | ||
| 37 | + var outputs = layers.Dense(10).Apply(x); | ||
| 38 | + | ||
| 39 | + var model = keras.Model(inputs, outputs, name: "mnist_model"); | ||
| 40 | + model.summary(); | ||
| 41 | + } | ||
| 42 | + | ||
| 26 | 43 | /// <summary> | |
| 27 | 44 | /// https://www.tensorflow.org/api_docs/python/tf/keras/layers/Embedding | |
| 28 | 45 | /// </summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -16,10 +16,7 @@ public class PythonTest | |||
| 16 | 16 | { | |
| 17 | 17 | #region python compatibility layer | |
| 18 | 18 | protected PythonTest self { get => this; } | |
| 19 | - protected object None | ||
| 20 | - { | ||
| 21 | - get { return null; } | ||
| 22 | - } | ||
| 19 | + protected int None => -1; | ||
| 23 | 20 | #endregion | |
| 24 | 21 | ||
| 25 | 22 | #region pytest assertions | |
@@ -150,7 +147,7 @@ public void assertProtoEquals(object toProto, object o) | |||
| 150 | 147 | ||
| 151 | 148 | protected object _eval_tensor(object tensor) | |
| 152 | 149 | { | |
| 153 | - if (tensor == None) | ||
| 150 | + if (tensor == null) | ||
| 154 | 151 | return None; | |
| 155 | 152 | //else if (callable(tensor)) | |
| 156 | 153 | // return self._eval_helper(tensor()) | |
| Back | FazBrowse Home | New Git URL |
0 commit comments