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

IInitializer for Keras. #355 · MSavameri/TensorFlow.NET@e631c1a · GitHub

Repository navigation

Commit e631c1a

Browse files
committed
IInitializer for Keras. SciSharp#355
1 parent e2f7be6 commit e631c1a

19 files changed

Lines changed: 106 additions & 130 deletions

File tree

‎src/TensorFlowNET.Core/Graphs/Graph.cs‎

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -147,6 +147,7 @@ public ITensorOrOperation as_graph_element(object obj, bool allow_tensor = true,
147147
/// <returns></returns>
148148
public Graph as_default()
149149
{
150+
tf.Context.graph_mode();
150151
return ops.set_default_graph(this);
151152
}
152153

@@ -490,6 +491,7 @@ public void prevent_fetching(Operation op)
490491

491492
protected override void DisposeManagedResources()
492493
{
494+
tf.Context.eager_mode();
493495
ops.default_graph_stack.remove(this);
494496
}
495497

‎src/TensorFlowNET.Core/Keras/Engine/Layer.cs‎

Lines changed: 16 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -18,9 +18,11 @@ limitations under the License.
1818
using System.Collections.Generic;
1919
using System.Linq;
2020
using System.Threading;
21+
using Tensorflow.Contexts;
2122
using Tensorflow.Keras.ArgsDefinition;
2223
using Tensorflow.Keras.Layers;
2324
using Tensorflow.Keras.Utils;
25+
using Tensorflow.Operations.Activation;
2426
using Tensorflow.Train;
2527
using static Tensorflow.Binding;
2628

@@ -46,7 +48,7 @@ public abstract class Layer : AutoTrackable
4648
protected bool built;
4749
public bool Trainable => args.Trainable;
4850
public TF_DataType DType => args.DType;
49-
51+
5052
/// <summary>
5153
/// A stateful layer is a layer whose updates are run during inference too,
5254
/// for instance stateful RNNs.
@@ -110,8 +112,11 @@ public Layer(LayerArgs args)
110112
/// <param name="input"></param>
111113
/// <param name="is_training"></param>
112114
/// <returns></returns>
113-
public Tensor Apply(Tensor[] inputs, bool is_training = false)
115+
public Tensor[] Apply(Tensor[] inputs, bool is_training = false)
114116
{
117+
var input = inputs[0];
118+
Tensor[] outputs = null;
119+
115120
callContext = callContext ?? new ThreadLocal<CallContext>()
116121
{
117122
Value = new CallContext()
@@ -120,7 +125,7 @@ public Tensor Apply(Tensor[] inputs, bool is_training = false)
120125
using var ctxManager = CallContext.enter();
121126

122127
string nameScope = "";
123-
if (tf.Context.executing_eagerly())
128+
if (tf.executing_eagerly())
124129
{
125130
nameScope = name;
126131
}
@@ -129,15 +134,21 @@ public Tensor Apply(Tensor[] inputs, bool is_training = false)
129134
throw new NotImplementedException("");
130135
}
131136

137+
using var graph = tf.keras.backend.get_graph().as_default();
138+
132139
tf_with(ops.name_scope(nameScope), scope =>
133140
{
134141
if (!built)
135142
MaybeBuild(inputs);
136143

137-
call(inputs, is_training: is_training);
144+
outputs = call(inputs, is_training: is_training);
145+
146+
(input, outputs) = _set_connectivity_metadata_(input, outputs);
147+
_handle_activity_regularization(inputs[0], outputs);
148+
_set_mask_metadata(inputs[0], outputs, null);
138149
});
139150

140-
throw new NotImplementedException("");
151+
return outputs;
141152
}
142153

143154
[Obsolete("User Apply()")]

‎src/TensorFlowNET.Core/Keras/Layers/Dense.cs‎

Lines changed: 7 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -30,8 +30,9 @@ namespace Tensorflow.Keras.Layers
3030
public class Dense : Layer
3131
{
3232
DenseArgs args;
33-
protected IVariableV1 kernel;
34-
protected IVariableV1 bias;
33+
IVariableV1 kernel;
34+
IVariableV1 bias;
35+
Activation activation => args.Activation;
3536

3637
public Dense(DenseArgs args) :
3738
base(args)
@@ -74,15 +75,15 @@ protected override Tensor[] call(Tensor[] inputs, bool training = false, Tensor
7475
}
7576
else
7677
{
77-
outputs = gen_math_ops.mat_mul(inputs[0], kernel.Handle);
78+
outputs = gen_math_ops.mat_mul(inputs[0], kernel.AsTensor());
7879
}
7980

8081
if (args.UseBias)
8182
outputs = tf.nn.bias_add(outputs, bias);
82-
//if (args.Activation != null)
83-
//outputs = args.Activation.Activate(outputs);
83+
if (args.Activation != null)
84+
outputs = activation(outputs);
8485

85-
return new[] { outputs, outputs };
86+
return new[] { outputs };
8687
}
8788
}
8889
}

‎src/TensorFlowNET.Core/Keras/Utils/base_layer_utils.cs‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -36,7 +36,11 @@ public static IVariableV1 make_variable(VariableArgs args)
3636

3737
ops.init_scope();
3838

39-
Func<Tensor> init_val = () => args.Initializer.call(args.Shape, dtype: args.DType);
39+
Func<Tensor> init_val = () => args.Initializer.Apply(new InitializerArgs
40+
{
41+
Shape = args.Shape,
42+
DType = args.DType
43+
});
4044

4145
var variable_dtype = args.DType.as_base_dtype();
4246
var v = tf.Variable(init_val,

‎src/TensorFlowNET.Core/Operations/Initializers/Constant.cs‎

Lines changed: 7 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -29,27 +29,18 @@ public Constant(T value, TF_DataType dtype = TF_DataType.TF_FLOAT, bool verify_s
2929
_verify_shape = verify_shape;
3030
}
3131

32-
public Tensor call(TensorShape shape, TF_DataType dtype = TF_DataType.DtInvalid, bool? verify_shape = null)
32+
public Tensor Apply(InitializerArgs args)
3333
{
34-
if (dtype == TF_DataType.DtInvalid)
35-
dtype = this.dtype;
34+
if (args.DType == TF_DataType.DtInvalid)
35+
args.DType = this.dtype;
3636

37-
if (!verify_shape.HasValue)
38-
verify_shape = _verify_shape;
37+
if (!args.VerifyShape.HasValue)
38+
args.VerifyShape = _verify_shape;
3939

40-
return constant_op._constant_impl(value, dtype, shape,
40+
return constant_op._constant_impl(value, args.DType, args.Shape,
4141
name: "Const",
42-
verify_shape: verify_shape.Value,
42+
verify_shape: args.VerifyShape.Value,
4343
allow_broadcast: false);
4444
}
45-
46-
public object get_config()
47-
{
48-
return new
49-
{
50-
value,
51-
dtype = dtype.name()
52-
};
53-
}
5445
}
5546
}

‎src/TensorFlowNET.Core/Operations/Initializers/GlorotUniform.cs‎

Lines changed: 0 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -30,18 +30,5 @@ public GlorotUniform(float scale = 1.0f,
3030
{
3131

3232
}
33-
34-
#pragma warning disable CS0114 // Member hides inherited member; missing override keyword
35-
public object get_config()
36-
#pragma warning restore CS0114 // Member hides inherited member; missing override keyword
37-
{
38-
return new
39-
{
40-
scale = _scale,
41-
mode = _mode,
42-
seed = _seed,
43-
dtype = _dtype
44-
};
45-
}
4633
}
4734
}

‎src/TensorFlowNET.Core/Operations/Initializers/IInitializer.cs‎

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -18,7 +18,6 @@ namespace Tensorflow
1818
{
1919
public interface IInitializer
2020
{
21-
Tensor call(TensorShape shape, TF_DataType dtype = TF_DataType.DtInvalid, bool? verify_shape = null);
22-
object get_config();
21+
Tensor Apply(InitializerArgs args);
2322
}
2423
}
Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,13 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Text;
4+
5+
namespace Tensorflow
6+
{
7+
public class InitializerArgs
8+
{
9+
public TensorShape Shape { get; set; }
10+
public TF_DataType DType { get; set; }
11+
public bool? VerifyShape { get; set; } = null;
12+
}
13+
}

‎src/TensorFlowNET.Core/Operations/Initializers/Ones.cs‎

Lines changed: 4 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -25,17 +25,12 @@ public Ones(TF_DataType dtype = TF_DataType.TF_FLOAT)
2525
this.dtype = dtype;
2626
}
2727

28-
public Tensor call(TensorShape shape, TF_DataType dtype = TF_DataType.DtInvalid, bool? verify_shape = null)
28+
public Tensor Apply(InitializerArgs args)
2929
{
30-
if (dtype == TF_DataType.DtInvalid)
31-
dtype = this.dtype;
30+
if (args.DType == TF_DataType.DtInvalid)
31+
args.DType = this.dtype;
3232

33-
return array_ops.ones(shape.dims, dtype);
34-
}
35-
36-
public object get_config()
37-
{
38-
return new { dtype = dtype.name() };
33+
return array_ops.ones(args.Shape, dtype);
3934
}
4035
}
4136
}

‎src/TensorFlowNET.Core/Operations/Initializers/RandomNormal.cs‎

Lines changed: 4 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -38,22 +38,11 @@ public RandomNormal(float mean = 0.0f,
3838
this.dtype = dtype;
3939
}
4040

41-
public Tensor call(TensorShape shape, TF_DataType dtype = TF_DataType.DtInvalid, bool? verify_shape = null)
41+
public Tensor Apply(InitializerArgs args)
4242
{
43-
if (dtype == TF_DataType.DtInvalid)
44-
dtype = this.dtype;
45-
return random_ops.random_normal(shape, mean, stddev, dtype, seed: seed);
46-
}
47-
48-
public object get_config()
49-
{
50-
return new
51-
{
52-
mean,
53-
stddev,
54-
seed,
55-
dtype
56-
};
43+
if (args.DType == TF_DataType.DtInvalid)
44+
args.DType = this.dtype;
45+
return random_ops.random_normal(args.Shape, mean, stddev, dtype, seed: seed);
5746
}
5847
}
5948
}

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL