| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent a075bba commit 78bd4c7
3 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -1,5 +1,6 @@ | |||
| 1 | 1 | using Tensorflow.Keras.Engine; | |
| 2 | 2 | using Tensorflow.Keras.Saving; | |
| 3 | + using Tensorflow.NumPy; | ||
| 3 | 4 | using Tensorflow.Training; | |
| 4 | 5 | ||
| 5 | 6 | namespace Tensorflow.Keras | |
@@ -18,6 +19,8 @@ public interface ILayer: IWithTrackable, IKerasConfigable | |||
| 18 | 19 | List<IVariableV1> TrainableWeights { get; } | |
| 19 | 20 | List<IVariableV1> NonTrainableWeights { get; } | |
| 20 | 21 | List<IVariableV1> Weights { get; set; } | |
| 22 | + void set_weights(List<NDArray> weights); | ||
| 23 | + List<NDArray> get_weights(); | ||
| 21 | 24 | Shape OutputShape { get; } | |
| 22 | 25 | Shape BatchInputShape { get; } | |
| 23 | 26 | TensorShapeConfig BuildInputShape { get; } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -21,6 +21,7 @@ limitations under the License. | |||
| 21 | 21 | using Tensorflow.Keras.ArgsDefinition.Rnn; | |
| 22 | 22 | using Tensorflow.Keras.Engine; | |
| 23 | 23 | using Tensorflow.Keras.Saving; | |
| 24 | + using Tensorflow.NumPy; | ||
| 24 | 25 | using Tensorflow.Operations; | |
| 25 | 26 | using Tensorflow.Train; | |
| 26 | 27 | using Tensorflow.Util; | |
@@ -71,7 +72,10 @@ public abstract class RnnCell : ILayer, RNNArgs.IRnnArgCell | |||
| 71 | 72 | ||
| 72 | 73 | public List<IVariableV1> TrainableVariables => throw new NotImplementedException(); | |
| 73 | 74 | public List<IVariableV1> TrainableWeights => throw new NotImplementedException(); | |
| 74 | - public List<IVariableV1> Weights => throw new NotImplementedException(); | ||
| 75 | + public List<IVariableV1> Weights { get => throw new NotImplementedException(); set => throw new NotImplementedException(); } | ||
| 76 | + | ||
| 77 | + public List<NDArray> get_weights() => throw new NotImplementedException(); | ||
| 78 | + public void set_weights(List<NDArray> weights) => throw new NotImplementedException(); | ||
| 75 | 79 | public List<IVariableV1> NonTrainableWeights => throw new NotImplementedException(); | |
| 76 | 80 | ||
| 77 | 81 | public Shape OutputShape => throw new NotImplementedException(); | |
@@ -84,8 +88,6 @@ public abstract class RnnCell : ILayer, RNNArgs.IRnnArgCell | |||
| 84 | 88 | protected bool built = false; | |
| 85 | 89 | public bool Built => built; | |
| 86 | 90 | ||
| 87 | - List<IVariableV1> ILayer.Weights { get => throw new NotImplementedException(); set => throw new NotImplementedException(); } | ||
| 88 | - | ||
| 89 | 91 | public RnnCell(bool trainable = true, | |
| 90 | 92 | string name = null, | |
| 91 | 93 | TF_DataType dtype = TF_DataType.DtInvalid, | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -120,6 +120,30 @@ public virtual List<IVariableV1> Weights | |||
| 120 | 120 | } | |
| 121 | 121 | } | |
| 122 | 122 | ||
| 123 | + public virtual void set_weights(List<NDArray> weights) | ||
| 124 | + { | ||
| 125 | + if (Weights.Count() != weights.Count()) throw new ValueError( | ||
| 126 | + $"You called `set_weights` on layer \"{this.name}\"" + | ||
| 127 | + $"with a weight list of length {len(weights)}, but the layer was " + | ||
| 128 | + $"expecting {len(Weights)} weights."); | ||
| 129 | + for (int i = 0; i < weights.Count(); i++) | ||
| 130 | + { | ||
| 131 | + if (weights[i].shape != Weights[i].shape) | ||
| 132 | + { | ||
| 133 | + throw new ValueError($"Layer weight shape {weights[i].shape} not compatible with provided weight shape {Weights[i].shape}"); | ||
| 134 | + } | ||
| 135 | + } | ||
| 136 | + foreach (var (this_w, v_w) in zip(Weights, weights)) | ||
| 137 | + this_w.assign(v_w, read_value: true); | ||
| 138 | + } | ||
| 139 | + | ||
| 140 | + public List<NDArray> get_weights() | ||
| 141 | + { | ||
| 142 | + List<NDArray > weights = new List<NDArray>(); | ||
| 143 | + weights.AddRange(Weights.ConvertAll(x => x.numpy())); | ||
| 144 | + return weights; | ||
| 145 | + } | ||
| 146 | + | ||
| 123 | 147 | protected int id; | |
| 124 | 148 | public int Id => id; | |
| 125 | 149 | protected string name; | |
| Back | FazBrowse Home | New Git URL |
0 commit comments