| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 1deaa75 commit 1dd95bd
11 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -91,6 +91,9 @@ BackwardFunction GetGradientFunction(string op_name, | |||
| 91 | 91 | Tensor[] op_outputs) | |
| 92 | 92 | => (output_grads, unneeded_gradients) => | |
| 93 | 93 | { | |
| 94 | + if (ops.gradientFunctions[op_name] == null) | ||
| 95 | + return new Tensor[op_inputs.Length]; | ||
| 96 | + | ||
| 94 | 97 | var gradients = ops.gradientFunctions[op_name](new EagerOperation | |
| 95 | 98 | { | |
| 96 | 99 | Name = op_name, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -15,5 +15,7 @@ public class OpTapeEntry<BackwardFunction, TapeTensor> | |||
| 15 | 15 | public TapeTensor[] output_tensor_info { get; set; } | |
| 16 | 16 | public long[] input_tensor_id { get; set; } | |
| 17 | 17 | public BackwardFunction backward_function { get; set; } | |
| 18 | + public override string ToString() | ||
| 19 | + => $"{op_type}, inputs: {string.Join(",", input_tensor_id)}"; | ||
| 18 | 20 | } | |
| 19 | 21 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -29,12 +29,13 @@ public Tensor[] ComputeGradient(long[] target_tensor_ids, | |||
| 29 | 29 | tensor_tape_, | |
| 30 | 30 | state.op_tape); | |
| 31 | 31 | ||
| 32 | - while (op_stack.Count > 0) | ||
| 32 | + while (!op_stack.empty()) | ||
| 33 | 33 | { | |
| 34 | 34 | var op = op_stack.Dequeue(); | |
| 35 | 35 | if (!state.op_tape.find(op, out var trace)) | |
| 36 | 36 | continue; | |
| 37 | 37 | ||
| 38 | + Console.WriteLine($"ComputeGradient: {state.op_tape[op].op_type}"); | ||
| 38 | 39 | state.op_tape.erase(op); | |
| 39 | 40 | ||
| 40 | 41 | var out_gradients = new List<Tensor>(trace.output_tensor_info.Length); | |
@@ -103,7 +104,7 @@ public Tensor[] ComputeGradient(long[] target_tensor_ids, | |||
| 103 | 104 | } | |
| 104 | 105 | else | |
| 105 | 106 | { | |
| 106 | - throw new NotImplementedException(""); | ||
| 107 | + in_gradients = new Tensor[trace.input_tensor_id.Length]; | ||
| 107 | 108 | } | |
| 108 | 109 | ||
| 109 | 110 | for (int i = 0; i < in_gradients.Length; ++i) | |
@@ -113,17 +114,18 @@ public Tensor[] ComputeGradient(long[] target_tensor_ids, | |||
| 113 | 114 | { | |
| 114 | 115 | var unaggregated_grads = gradients[id]; | |
| 115 | 116 | unaggregated_grads.Add(in_gradients[i]); | |
| 116 | - if(unaggregated_grads.Count > kMinAggregateCount) | ||
| 117 | + if (unaggregated_grads.Count > kMinAggregateCount) | ||
| 117 | 118 | { | |
| 118 | - if(!gradients_size.ContainsKey(id)) | ||
| 119 | + if (!gradients_size.find(id, out var size)) | ||
| 119 | 120 | { | |
| 121 | + size = (long)unaggregated_grads[0].size; | ||
| 122 | + gradients_size.emplace(id, size); | ||
| 120 | 123 | } | |
| 121 | - else | ||
| 122 | - { | ||
| 123 | 124 | ||
| 125 | + if (unaggregated_grads.Count * size * 4 > kMinAggregateBytes) | ||
| 126 | + { | ||
| 127 | + throw new NotImplementedException(""); | ||
| 124 | 128 | } | |
| 125 | - | ||
| 126 | - throw new NotImplementedException(""); | ||
| 127 | 129 | } | |
| 128 | 130 | } | |
| 129 | 131 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -0,0 +1,16 @@ | |||
| 1 | + using System; | ||
| 2 | + using System.Collections.Generic; | ||
| 3 | + using System.Text; | ||
| 4 | + | ||
| 5 | + namespace Tensorflow.Keras.ArgsDefinition | ||
| 6 | + { | ||
| 7 | + public class RMSpropArgs | ||
| 8 | + { | ||
| 9 | + public float LearningRate { get; set; } = 0.001f; | ||
| 10 | + public float RHO { get; set; } = 0.9f; | ||
| 11 | + public float Momentum { get; set; } = 0.0f; | ||
| 12 | + public float Epsilon { get; set; } = 1e-7f; | ||
| 13 | + public bool Centered { get; set; } = false; | ||
| 14 | + public string Name { get; set; } = "RMSprop"; | ||
| 15 | + } | ||
| 16 | + } | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -23,8 +23,6 @@ public class Functional : Model | |||
| 23 | 23 | List<KerasHistory> _input_coordinates; | |
| 24 | 24 | List<KerasHistory> _output_coordinates; | |
| 25 | 25 | public string[] NetworkNodes { get; set; } | |
| 26 | - public Dictionary<int, List<Node>> NodesByDepth { get; set; } | ||
| 27 | - public List<Layer> Layers => _layers; | ||
| 28 | 26 | ||
| 29 | 27 | Dictionary<int, int> tensor_usage_count; | |
| 30 | 28 | public Dictionary<int, int> TensorUsageCount => tensor_usage_count; | |
@@ -43,9 +41,10 @@ public override List<IVariableV1> trainable_variables | |||
| 43 | 41 | } | |
| 44 | 42 | } | |
| 45 | 43 | ||
| 46 | - public Functional(Tensors inputs, Tensors outputs) | ||
| 44 | + public Functional(Tensors inputs, Tensors outputs, string name = null) | ||
| 47 | 45 | : base(new ModelArgs | |
| 48 | 46 | { | |
| 47 | + Name = name, | ||
| 49 | 48 | Inputs = inputs, | |
| 50 | 49 | Outputs = outputs | |
| 51 | 50 | }) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -10,7 +10,7 @@ namespace Tensorflow.Keras.Engine | |||
| 10 | 10 | /// <summary> | |
| 11 | 11 | /// `Model` groups layers into an object with training and inference features. | |
| 12 | 12 | /// </summary> | |
| 13 | - public class Model : Layer | ||
| 13 | + public partial class Model : Layer | ||
| 14 | 14 | { | |
| 15 | 15 | #pragma warning disable CS0169 // The field 'Model._cloning' is never used | |
| 16 | 16 | bool _cloning; | |
@@ -33,12 +33,20 @@ public Model(ModelArgs args) | |||
| 33 | 33 | ||
| 34 | 34 | } | |
| 35 | 35 | ||
| 36 | + public void compile(ILossFunc loss, OptimizerV2 optimizer, string[] metrics) | ||
| 37 | + { | ||
| 38 | + | ||
| 39 | + } | ||
| 40 | + | ||
| 36 | 41 | public void compile(string optimizerName, string lossName) | |
| 37 | 42 | { | |
| 38 | 43 | switch (optimizerName) | |
| 39 | 44 | { | |
| 40 | 45 | case "rmsprop": | |
| 41 | - optimizer = new RMSprop(); | ||
| 46 | + optimizer = new RMSprop(new RMSpropArgs | ||
| 47 | + { | ||
| 48 | + | ||
| 49 | + }); | ||
| 42 | 50 | break; | |
| 43 | 51 | } | |
| 44 | 52 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -30,7 +30,7 @@ namespace Tensorflow.Keras.Engine | |||
| 30 | 30 | /// Each time the output of a layer is used by another layer, | |
| 31 | 31 | /// a node is added to `layer._outbound_nodes`. | |
| 32 | 32 | /// </summary> | |
| 33 | - public class Node | ||
| 33 | + public partial class Node | ||
| 34 | 34 | { | |
| 35 | 35 | NodeArgs args; | |
| 36 | 36 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -39,8 +39,8 @@ public Sequential Sequential(List<Layer> layers = null, | |||
| 39 | 39 | /// <param name="input"></param> | |
| 40 | 40 | /// <param name="output"></param> | |
| 41 | 41 | /// <returns></returns> | |
| 42 | - public Functional Model(Tensors inputs, Tensors outputs) | ||
| 43 | - => new Functional(inputs, outputs); | ||
| 42 | + public Functional Model(Tensors inputs, Tensors outputs, string name = null) | ||
| 43 | + => new Functional(inputs, outputs, name: name); | ||
| 44 | 44 | ||
| 45 | 45 | /// <summary> | |
| 46 | 46 | /// Instantiate a Keras tensor. | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,6 +1,7 @@ | |||
| 1 | 1 | using System; | |
| 2 | 2 | using System.Collections.Generic; | |
| 3 | 3 | using System.Text; | |
| 4 | + using Tensorflow.Keras.ArgsDefinition; | ||
| 4 | 5 | ||
| 5 | 6 | namespace Tensorflow.Keras.Optimizers | |
| 6 | 7 | { | |
@@ -29,5 +30,31 @@ public OptimizerV2 Adam(float learning_rate = 0.001f, | |||
| 29 | 30 | epsilon: epsilon, | |
| 30 | 31 | amsgrad: amsgrad, | |
| 31 | 32 | name: name); | |
| 33 | + | ||
| 34 | + /// <summary> | ||
| 35 | + /// Construct a new RMSprop optimizer. | ||
| 36 | + /// </summary> | ||
| 37 | + /// <param name="learning_rate"></param> | ||
| 38 | + /// <param name="rho"></param> | ||
| 39 | + /// <param name="momentum"></param> | ||
| 40 | + /// <param name="epsilon"></param> | ||
| 41 | + /// <param name="centered"></param> | ||
| 42 | + /// <param name="name"></param> | ||
| 43 | + /// <returns></returns> | ||
| 44 | + public OptimizerV2 RMSprop(float learning_rate = 0.001f, | ||
| 45 | + float rho = 0.9f, | ||
| 46 | + float momentum = 0.0f, | ||
| 47 | + float epsilon = 1e-7f, | ||
| 48 | + bool centered = false, | ||
| 49 | + string name = "RMSprop") | ||
| 50 | + => new RMSprop(new RMSpropArgs | ||
| 51 | + { | ||
| 52 | + LearningRate = learning_rate, | ||
| 53 | + RHO = rho, | ||
| 54 | + Momentum = momentum, | ||
| 55 | + Epsilon = epsilon, | ||
| 56 | + Centered = centered, | ||
| 57 | + Name = name | ||
| 58 | + }); | ||
| 32 | 59 | } | |
| 33 | 60 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,6 +1,7 @@ | |||
| 1 | 1 | using System; | |
| 2 | 2 | using System.Collections.Generic; | |
| 3 | 3 | using System.Text; | |
| 4 | + using Tensorflow.Keras.ArgsDefinition; | ||
| 4 | 5 | ||
| 5 | 6 | namespace Tensorflow.Keras.Optimizers | |
| 6 | 7 | { | |
@@ -9,6 +10,11 @@ namespace Tensorflow.Keras.Optimizers | |||
| 9 | 10 | /// </summary> | |
| 10 | 11 | public class RMSprop : OptimizerV2 | |
| 11 | 12 | { | |
| 13 | + RMSpropArgs args; | ||
| 12 | 14 | ||
| 15 | + public RMSprop(RMSpropArgs args) | ||
| 16 | + { | ||
| 17 | + this.args = args; | ||
| 18 | + } | ||
| 13 | 19 | } | |
| 14 | 20 | } | |
| Back | FazBrowse Home | New Git URL |
0 commit comments