| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 0114885 commit 675b93a
29 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -71,15 +71,15 @@ public Tensor strided_slice<T>(Tensor input, T[] begin, T[] end, T[] strides = n | |||
| 71 | 71 | public Tensor[] split(Tensor value, int num_split, Tensor axis, string name = null) | |
| 72 | 72 | => array_ops.split( | |
| 73 | 73 | value: value, | |
| 74 | - num_split: num_split, | ||
| 74 | + num_or_size_splits: num_split, | ||
| 75 | 75 | axis: axis, | |
| 76 | 76 | name: name); | |
| 77 | 77 | ||
| 78 | 78 | public Tensor[] split(Tensor value, int num_split, int axis, string name = null) | |
| 79 | 79 | => array_ops.split( | |
| 80 | 80 | value: value, | |
| 81 | - num_split: num_split, | ||
| 82 | - axis: axis, | ||
| 81 | + num_or_size_splits: num_split, | ||
| 82 | + axis: ops.convert_to_tensor(axis), | ||
| 83 | 83 | name: name); | |
| 84 | 84 | ||
| 85 | 85 | public Tensor ensure_shape(Tensor x, Shape shape, string name = null) | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -197,25 +197,11 @@ public bool IsNested() | |||
| 197 | 197 | } | |
| 198 | 198 | else if(NestType is NestType.List) | |
| 199 | 199 | { | |
| 200 | - foreach(var item in ListValue!) | ||
| 201 | - { | ||
| 202 | - if(item.NestType is NestType.List or NestType.Dictionary) | ||
| 203 | - { | ||
| 204 | - return true; | ||
| 205 | - } | ||
| 206 | - } | ||
| 207 | - return false; | ||
| 200 | + return ListValue!.Count > 0; | ||
| 208 | 201 | } | |
| 209 | 202 | else | |
| 210 | 203 | { | |
| 211 | - foreach (var item in DictValue!.Values) | ||
| 212 | - { | ||
| 213 | - if (item.NestType is NestType.List or NestType.Dictionary) | ||
| 214 | - { | ||
| 215 | - return true; | ||
| 216 | - } | ||
| 217 | - } | ||
| 218 | - return false; | ||
| 204 | + return DictValue!.Count > 0; | ||
| 219 | 205 | } | |
| 220 | 206 | } | |
| 221 | 207 | ||
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -352,7 +352,11 @@ bool SetOpAttrScalar(Context ctx, SafeEagerOpHandle op, | |||
| 352 | 352 | c_api.TFE_OpSetAttrFloat(op, key, Convert.ToSingle(value)); | |
| 353 | 353 | break; | |
| 354 | 354 | case TF_AttrType.TF_ATTR_SHAPE: | |
| 355 | - var dims = (value as long[]).ToArray(); | ||
| 355 | + long[] dims; | ||
| 356 | + if (value is Shape shape) dims = shape.dims.ToArray(); | ||
| 357 | + else if (value is long[] longs) dims = longs.ToArray(); | ||
| 358 | + else if (value is int[] ints) dims = ints.Select(x => (long)x).ToArray(); | ||
| 359 | + else dims = ((long[])value).ToArray(); | ||
| 356 | 360 | c_api.TFE_OpSetAttrShape(op, key, dims, dims.Length, status); | |
| 357 | 361 | status.Check(true); | |
| 358 | 362 | break; | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -137,14 +137,14 @@ TapeTensor TapeTensorFromTensor(Tensor tensor) | |||
| 137 | 137 | { | |
| 138 | 138 | dims[i] = c_api.TFE_TensorHandleDim(handle, i, status); | |
| 139 | 139 | } | |
| 140 | - Shape tensor_shape = new(dims); | ||
| 141 | 140 | ||
| 142 | 141 | if(status.Code != TF_Code.TF_OK) | |
| 143 | 142 | { | |
| 144 | 143 | return new TapeTensor(id, TF_DataType.DtInvalid, Shape.Null); | |
| 145 | 144 | } | |
| 146 | 145 | else | |
| 147 | 146 | { | |
| 147 | + Shape tensor_shape = new(dims); | ||
| 148 | 148 | return new TapeTensor(id, dtype, tensor_shape); | |
| 149 | 149 | } | |
| 150 | 150 | } | |
@@ -173,8 +173,12 @@ bool DTypeNeedsHandleData(TF_DataType dtype) | |||
| 173 | 173 | return dtype == dtypes.variant || dtype == dtypes.resource; | |
| 174 | 174 | } | |
| 175 | 175 | ||
| 176 | - bool ListContainNone(long[] list) | ||
| 176 | + bool ListContainNone(long[]? list) | ||
| 177 | 177 | { | |
| 178 | + if(list is null) | ||
| 179 | + { | ||
| 180 | + return true; | ||
| 181 | + } | ||
| 178 | 182 | int len = list.Length; | |
| 179 | 183 | if(len == 0) | |
| 180 | 184 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -90,8 +90,7 @@ private static Tensor[] _ConcatGradHelper(Operation op, Tensor grad, int start_v | |||
| 90 | 90 | ? input_values[0].rank + dim_int | |
| 91 | 91 | : dim_int % input_values[0].rank; | |
| 92 | 92 | var sizes = input_values.Select(x => x.shape[non_neg_concat_dim]).ToArray(); | |
| 93 | - var sizes_tensor = constant_op.constant(sizes); | ||
| 94 | - out_grads = array_ops.split(grad, sizes_tensor, non_neg_concat_dim).ToList(); | ||
| 93 | + out_grads = array_ops.split(grad, sizes.Select(x => (int)x).ToArray(), ops.convert_to_tensor(non_neg_concat_dim)).ToList(); | ||
| 95 | 94 | } | |
| 96 | 95 | else if (constant_op.is_constant(concat_dim)) | |
| 97 | 96 | { | |
@@ -127,7 +126,7 @@ there will be a small number of performance regressions.*/ | |||
| 127 | 126 | new Tensor[] { non_neg_concat_dim, tf.constant(0) }, | |
| 128 | 127 | new Tensor[] { tf.constant(1), tf.constant(-1) }); | |
| 129 | 128 | var squeeze_sizes = array_ops.squeeze(slice); | |
| 130 | - out_grads = array_ops.split(axis: grad, value: squeeze_sizes, num_split: (int)non_neg_concat_dim).ToList(); | ||
| 129 | + out_grads = array_ops.split(axis: grad, value: squeeze_sizes, num_or_size_splits: (int)non_neg_concat_dim).ToList(); | ||
| 131 | 130 | } | |
| 132 | 131 | else | |
| 133 | 132 | { | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -4,8 +4,6 @@ public class LSTMArgs : RNNArgs | |||
| 4 | 4 | { | |
| 5 | 5 | // TODO: maybe change the `RNNArgs` and implement this class. | |
| 6 | 6 | public bool UnitForgetBias { get; set; } | |
| 7 | - public float Dropout { get; set; } | ||
| 8 | - public float RecurrentDropout { get; set; } | ||
| 9 | 7 | public int Implementation { get; set; } | |
| 10 | 8 | } | |
| 11 | 9 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -29,7 +29,7 @@ public class LSTMCellArgs : AutoSerializeLayerArgs | |||
| 29 | 29 | [JsonProperty("unit_forget_bias")] | |
| 30 | 30 | public bool UnitForgetBias { get; set; } = true; | |
| 31 | 31 | [JsonProperty("implementation")] | |
| 32 | - public int Implementation { get; set; } = 2; | ||
| 32 | + public int Implementation { get; set; } = 1; | ||
| 33 | 33 | ||
| 34 | 34 | } | |
| 35 | 35 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -7,12 +7,6 @@ namespace Tensorflow.Keras.ArgsDefinition.Rnn | |||
| 7 | 7 | // TODO(Rinne): add regularizers. | |
| 8 | 8 | public class RNNArgs : AutoSerializeLayerArgs | |
| 9 | 9 | { | |
| 10 | - [JsonProperty("cell")] | ||
| 11 | - // TODO: the cell should be serialized with `serialize_keras_object`. | ||
| 12 | - public IRnnCell Cell { get; set; } = null; | ||
| 13 | - [JsonProperty("cells")] | ||
| 14 | - public IList<IRnnCell> Cells { get; set; } = null; | ||
| 15 | - | ||
| 16 | 10 | [JsonProperty("return_sequences")] | |
| 17 | 11 | public bool ReturnSequences { get; set; } = false; | |
| 18 | 12 | [JsonProperty("return_state")] | |
@@ -25,8 +19,10 @@ public class RNNArgs : AutoSerializeLayerArgs | |||
| 25 | 19 | public bool Unroll { get; set; } = false; | |
| 26 | 20 | [JsonProperty("time_major")] | |
| 27 | 21 | public bool TimeMajor { get; set; } = false; | |
| 22 | + | ||
| 23 | + public int? InputDim { get; set; } | ||
| 24 | + public int? InputLength { get; set; } | ||
| 28 | 25 | // TODO: Add `num_constants` and `zero_output_for_mask`. | |
| 29 | - public Dictionary<string, object> Kwargs { get; set; } = null; | ||
| 30 | 26 | ||
| 31 | 27 | public int Units { get; set; } | |
| 32 | 28 | public Activation Activation { get; set; } | |
@@ -38,21 +34,5 @@ public class RNNArgs : AutoSerializeLayerArgs | |||
| 38 | 34 | public float Dropout { get; set; } = .0f; | |
| 39 | 35 | public bool ZeroOutputForMask { get; set; } = false; | |
| 40 | 36 | public float RecurrentDropout { get; set; } = .0f; | |
| 41 | - | ||
| 42 | - // kernel_regularizer=None, | ||
| 43 | - // recurrent_regularizer=None, | ||
| 44 | - // bias_regularizer=None, | ||
| 45 | - // activity_regularizer=None, | ||
| 46 | - // kernel_constraint=None, | ||
| 47 | - // recurrent_constraint=None, | ||
| 48 | - // bias_constraint=None, | ||
| 49 | - // dropout=0., | ||
| 50 | - // recurrent_dropout=0., | ||
| 51 | - // return_sequences=False, | ||
| 52 | - // return_state=False, | ||
| 53 | - // go_backwards=False, | ||
| 54 | - // stateful=False, | ||
| 55 | - // unroll=False, | ||
| 56 | - // **kwargs): | ||
| 57 | 37 | } | |
| 58 | 38 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -5,7 +5,6 @@ namespace Tensorflow.Keras.ArgsDefinition.Rnn | |||
| 5 | 5 | { | |
| 6 | 6 | public class StackedRNNCellsArgs : LayerArgs | |
| 7 | 7 | { | |
| 8 | - public IList<IRnnCell> Cells { get; set; } | ||
| 9 | - public Dictionary<string, object> Kwargs { get; set; } = null; | ||
| 8 | + public bool ReverseStateOrder = false; | ||
| 10 | 9 | } | |
| 11 | 10 | } | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -182,7 +182,7 @@ public ILayer LSTM(int units, | |||
| 182 | 182 | bool unit_forget_bias = true, | |
| 183 | 183 | float dropout = 0f, | |
| 184 | 184 | float recurrent_dropout = 0f, | |
| 185 | - int implementation = 2, | ||
| 185 | + int implementation = 1, | ||
| 186 | 186 | bool return_sequences = false, | |
| 187 | 187 | bool return_state = false, | |
| 188 | 188 | bool go_backwards = false, | |
| Back | FazBrowse Home | New Git URL |
0 commit comments