| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent e2f7be6 commit e631c1a
19 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -147,6 +147,7 @@ public ITensorOrOperation as_graph_element(object obj, bool allow_tensor = true, | |||
| 147 | 147 | /// <returns></returns> | |
| 148 | 148 | public Graph as_default() | |
| 149 | 149 | { | |
| 150 | + tf.Context.graph_mode(); | ||
| 150 | 151 | return ops.set_default_graph(this); | |
| 151 | 152 | } | |
| 152 | 153 | ||
@@ -490,6 +491,7 @@ public void prevent_fetching(Operation op) | |||
| 490 | 491 | ||
| 491 | 492 | protected override void DisposeManagedResources() | |
| 492 | 493 | { | |
| 494 | + tf.Context.eager_mode(); | ||
| 493 | 495 | ops.default_graph_stack.remove(this); | |
| 494 | 496 | } | |
| 495 | 497 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -18,9 +18,11 @@ limitations under the License. | |||
| 18 | 18 | using System.Collections.Generic; | |
| 19 | 19 | using System.Linq; | |
| 20 | 20 | using System.Threading; | |
| 21 | + using Tensorflow.Contexts; | ||
| 21 | 22 | using Tensorflow.Keras.ArgsDefinition; | |
| 22 | 23 | using Tensorflow.Keras.Layers; | |
| 23 | 24 | using Tensorflow.Keras.Utils; | |
| 25 | + using Tensorflow.Operations.Activation; | ||
| 24 | 26 | using Tensorflow.Train; | |
| 25 | 27 | using static Tensorflow.Binding; | |
| 26 | 28 | ||
@@ -46,7 +48,7 @@ public abstract class Layer : AutoTrackable | |||
| 46 | 48 | protected bool built; | |
| 47 | 49 | public bool Trainable => args.Trainable; | |
| 48 | 50 | public TF_DataType DType => args.DType; | |
| 49 | - | ||
| 51 | + | ||
| 50 | 52 | /// <summary> | |
| 51 | 53 | /// A stateful layer is a layer whose updates are run during inference too, | |
| 52 | 54 | /// for instance stateful RNNs. | |
@@ -110,8 +112,11 @@ public Layer(LayerArgs args) | |||
| 110 | 112 | /// <param name="input"></param> | |
| 111 | 113 | /// <param name="is_training"></param> | |
| 112 | 114 | /// <returns></returns> | |
| 113 | - public Tensor Apply(Tensor[] inputs, bool is_training = false) | ||
| 115 | + public Tensor[] Apply(Tensor[] inputs, bool is_training = false) | ||
| 114 | 116 | { | |
| 117 | + var input = inputs[0]; | ||
| 118 | + Tensor[] outputs = null; | ||
| 119 | + | ||
| 115 | 120 | callContext = callContext ?? new ThreadLocal<CallContext>() | |
| 116 | 121 | { | |
| 117 | 122 | Value = new CallContext() | |
@@ -120,7 +125,7 @@ public Tensor Apply(Tensor[] inputs, bool is_training = false) | |||
| 120 | 125 | using var ctxManager = CallContext.enter(); | |
| 121 | 126 | ||
| 122 | 127 | string nameScope = ""; | |
| 123 | - if (tf.Context.executing_eagerly()) | ||
| 128 | + if (tf.executing_eagerly()) | ||
| 124 | 129 | { | |
| 125 | 130 | nameScope = name; | |
| 126 | 131 | } | |
@@ -129,15 +134,21 @@ public Tensor Apply(Tensor[] inputs, bool is_training = false) | |||
| 129 | 134 | throw new NotImplementedException(""); | |
| 130 | 135 | } | |
| 131 | 136 | ||
| 137 | + using var graph = tf.keras.backend.get_graph().as_default(); | ||
| 138 | + | ||
| 132 | 139 | tf_with(ops.name_scope(nameScope), scope => | |
| 133 | 140 | { | |
| 134 | 141 | if (!built) | |
| 135 | 142 | MaybeBuild(inputs); | |
| 136 | 143 | ||
| 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); | ||
| 138 | 149 | }); | |
| 139 | 150 | ||
| 140 | - throw new NotImplementedException(""); | ||
| 151 | + return outputs; | ||
| 141 | 152 | } | |
| 142 | 153 | ||
| 143 | 154 | [Obsolete("User Apply()")] | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -30,8 +30,9 @@ namespace Tensorflow.Keras.Layers | |||
| 30 | 30 | public class Dense : Layer | |
| 31 | 31 | { | |
| 32 | 32 | DenseArgs args; | |
| 33 | - protected IVariableV1 kernel; | ||
| 34 | - protected IVariableV1 bias; | ||
| 33 | + IVariableV1 kernel; | ||
| 34 | + IVariableV1 bias; | ||
| 35 | + Activation activation => args.Activation; | ||
| 35 | 36 | ||
| 36 | 37 | public Dense(DenseArgs args) : | |
| 37 | 38 | base(args) | |
@@ -74,15 +75,15 @@ protected override Tensor[] call(Tensor[] inputs, bool training = false, Tensor | |||
| 74 | 75 | } | |
| 75 | 76 | else | |
| 76 | 77 | { | |
| 77 | - outputs = gen_math_ops.mat_mul(inputs[0], kernel.Handle); | ||
| 78 | + outputs = gen_math_ops.mat_mul(inputs[0], kernel.AsTensor()); | ||
| 78 | 79 | } | |
| 79 | 80 | ||
| 80 | 81 | if (args.UseBias) | |
| 81 | 82 | 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); | ||
| 84 | 85 | ||
| 85 | - return new[] { outputs, outputs }; | ||
| 86 | + return new[] { outputs }; | ||
| 86 | 87 | } | |
| 87 | 88 | } | |
| 88 | 89 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -36,7 +36,11 @@ public static IVariableV1 make_variable(VariableArgs args) | |||
| 36 | 36 | ||
| 37 | 37 | ops.init_scope(); | |
| 38 | 38 | ||
| 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 | + }); | ||
| 40 | 44 | ||
| 41 | 45 | var variable_dtype = args.DType.as_base_dtype(); | |
| 42 | 46 | var v = tf.Variable(init_val, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -29,27 +29,18 @@ public Constant(T value, TF_DataType dtype = TF_DataType.TF_FLOAT, bool verify_s | |||
| 29 | 29 | _verify_shape = verify_shape; | |
| 30 | 30 | } | |
| 31 | 31 | ||
| 32 | - public Tensor call(TensorShape shape, TF_DataType dtype = TF_DataType.DtInvalid, bool? verify_shape = null) | ||
| 32 | + public Tensor Apply(InitializerArgs args) | ||
| 33 | 33 | { | |
| 34 | - if (dtype == TF_DataType.DtInvalid) | ||
| 35 | - dtype = this.dtype; | ||
| 34 | + if (args.DType == TF_DataType.DtInvalid) | ||
| 35 | + args.DType = this.dtype; | ||
| 36 | 36 | ||
| 37 | - if (!verify_shape.HasValue) | ||
| 38 | - verify_shape = _verify_shape; | ||
| 37 | + if (!args.VerifyShape.HasValue) | ||
| 38 | + args.VerifyShape = _verify_shape; | ||
| 39 | 39 | ||
| 40 | - return constant_op._constant_impl(value, dtype, shape, | ||
| 40 | + return constant_op._constant_impl(value, args.DType, args.Shape, | ||
| 41 | 41 | name: "Const", | |
| 42 | - verify_shape: verify_shape.Value, | ||
| 42 | + verify_shape: args.VerifyShape.Value, | ||
| 43 | 43 | allow_broadcast: false); | |
| 44 | 44 | } | |
| 45 | - | ||
| 46 | - public object get_config() | ||
| 47 | - { | ||
| 48 | - return new | ||
| 49 | - { | ||
| 50 | - value, | ||
| 51 | - dtype = dtype.name() | ||
| 52 | - }; | ||
| 53 | - } | ||
| 54 | 45 | } | |
| 55 | 46 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -30,18 +30,5 @@ public GlorotUniform(float scale = 1.0f, | |||
| 30 | 30 | { | |
| 31 | 31 | ||
| 32 | 32 | } | |
| 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 | - } | ||
| 46 | 33 | } | |
| 47 | 34 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -18,7 +18,6 @@ namespace Tensorflow | |||
| 18 | 18 | { | |
| 19 | 19 | public interface IInitializer | |
| 20 | 20 | { | |
| 21 | - Tensor call(TensorShape shape, TF_DataType dtype = TF_DataType.DtInvalid, bool? verify_shape = null); | ||
| 22 | - object get_config(); | ||
| 21 | + Tensor Apply(InitializerArgs args); | ||
| 23 | 22 | } | |
| 24 | 23 | } | |
| Original file line number | Diff line number | Diff 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 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -25,17 +25,12 @@ public Ones(TF_DataType dtype = TF_DataType.TF_FLOAT) | |||
| 25 | 25 | this.dtype = dtype; | |
| 26 | 26 | } | |
| 27 | 27 | ||
| 28 | - public Tensor call(TensorShape shape, TF_DataType dtype = TF_DataType.DtInvalid, bool? verify_shape = null) | ||
| 28 | + public Tensor Apply(InitializerArgs args) | ||
| 29 | 29 | { | |
| 30 | - if (dtype == TF_DataType.DtInvalid) | ||
| 31 | - dtype = this.dtype; | ||
| 30 | + if (args.DType == TF_DataType.DtInvalid) | ||
| 31 | + args.DType = this.dtype; | ||
| 32 | 32 | ||
| 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); | ||
| 39 | 34 | } | |
| 40 | 35 | } | |
| 41 | 36 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -38,22 +38,11 @@ public RandomNormal(float mean = 0.0f, | |||
| 38 | 38 | this.dtype = dtype; | |
| 39 | 39 | } | |
| 40 | 40 | ||
| 41 | - public Tensor call(TensorShape shape, TF_DataType dtype = TF_DataType.DtInvalid, bool? verify_shape = null) | ||
| 41 | + public Tensor Apply(InitializerArgs args) | ||
| 42 | 42 | { | |
| 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); | ||
| 57 | 46 | } | |
| 58 | 47 | } | |
| 59 | 48 | } | |
| Back | FazBrowse Home | New Git URL |
0 commit comments