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

add Keras model.summary(). · MSavameri/TensorFlow.NET@9f2adcf · GitHub

Repository navigation

Commit 9f2adcf

Browse files
committed
add Keras model.summary().
1 parent 1dd95bd commit 9f2adcf

5 files changed

Lines changed: 63 additions & 5 deletions

File tree

Lines changed: 20 additions & 0 deletions
Original file line numberDiff line numberDiff 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+
}
Lines changed: 18 additions & 0 deletions
Original file line numberDiff line numberDiff 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+
}

‎test/TensorFlowNET.UnitTest/EagerModeTestBase.cs‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2,12 +2,18 @@
22
using System;
33
using System.Collections.Generic;
44
using System.Text;
5+
using Tensorflow;
6+
using Tensorflow.Keras.Engine;
7+
using Tensorflow.Keras.Layers;
58
using static Tensorflow.Binding;
69

710
namespace TensorFlowNET.UnitTest
811
{
912
public class EagerModeTestBase : PythonTest
1013
{
14+
protected KerasApi keras = tf.keras;
15+
protected LayersApi layers = tf.keras.layers;
16+
1117
[TestInitialize]
1218
public void TestInit()
1319
{

‎test/TensorFlowNET.UnitTest/Keras/LayersTest.cs‎

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,13 +16,30 @@ namespace TensorFlowNET.UnitTest.Keras
1616
[TestClass]
1717
public class LayersTest : EagerModeTestBase
1818
{
19+
1920
[TestMethod]
2021
public void Sequential()
2122
{
2223
var model = tf.keras.Sequential();
2324
model.add(tf.keras.Input(shape: 16));
2425
}
2526

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+
2643
/// <summary>
2744
/// https://www.tensorflow.org/api_docs/python/tf/keras/layers/Embedding
2845
/// </summary>

‎test/TensorFlowNET.UnitTest/PythonTest.cs‎

Lines changed: 2 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -16,10 +16,7 @@ public class PythonTest
1616
{
1717
#region python compatibility layer
1818
protected PythonTest self { get => this; }
19-
protected object None
20-
{
21-
get { return null; }
22-
}
19+
protected int None => -1;
2320
#endregion
2421

2522
#region pytest assertions
@@ -150,7 +147,7 @@ public void assertProtoEquals(object toProto, object o)
150147

151148
protected object _eval_tensor(object tensor)
152149
{
153-
if (tensor == None)
150+
if (tensor == null)
154151
return None;
155152
//else if (callable(tensor))
156153
// return self._eval_helper(tensor())

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL