| FazBrowse GitHub Viewer | Trending | | Home |
| Tools: [Download Repo ZIP] [Original HTTPS Page] |
1 parent 0e5f56e commit 318f991
7 files changed
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -26,7 +26,7 @@ public partial class tensorflow | |||
| 26 | 26 | ||
| 27 | 27 | public class nn_internal | |
| 28 | 28 | { | |
| 29 | - public Tensor conv2d(Tensor input, IVariableV1 filter, int[] strides, string padding, bool use_cudnn_on_gpu = true, | ||
| 29 | + public Tensor conv2d(Tensor input, Tensor filter, int[] strides, string padding, bool use_cudnn_on_gpu = true, | ||
| 30 | 30 | string data_format = "NHWC", int[] dilations = null, string name = null) | |
| 31 | 31 | { | |
| 32 | 32 | var parameters = new Conv2dParams | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -133,6 +133,46 @@ public static Tensor[] _SquaredDifferenceGrad(Operation op, Tensor[] grads) | |||
| 133 | 133 | -x_grad | |
| 134 | 134 | }; | |
| 135 | 135 | } | |
| 136 | + | ||
| 137 | + /// <summary> | ||
| 138 | + /// The derivatives for deconvolution. | ||
| 139 | + /// </summary> | ||
| 140 | + /// <param name="op">The Deconvolution op.</param> | ||
| 141 | + /// <param name="grads">The tensor representing the gradient w.r.t. the output</param> | ||
| 142 | + /// <returns>The gradients w.r.t. the input and the filter</returns> | ||
| 143 | + [RegisterGradient("Conv2DBackpropInput")] | ||
| 144 | + public static Tensor[] _Conv2DBackpropInputGrad(Operation op, Tensor[] grads) | ||
| 145 | + { | ||
| 146 | + var grad = grads[0]; | ||
| 147 | + var dilations = op.get_attr_list<int>("dilations"); | ||
| 148 | + var strides = op.get_attr_list<int>("strides"); | ||
| 149 | + var padding = op.get_attr<string>("padding"); | ||
| 150 | + var explicit_paddings = op.get_attr_list<int>("explicit_paddings"); | ||
| 151 | + var use_cudnn_on_gpu = op.get_attr<bool>("use_cudnn_on_gpu"); | ||
| 152 | + var data_format = op.get_attr<string>("data_format"); | ||
| 153 | + | ||
| 154 | + return new Tensor[] | ||
| 155 | + { | ||
| 156 | + gen_nn_ops.conv2d_backprop_filter(grad, array_ops.shape(op.inputs[1]), op.inputs[2], | ||
| 157 | + strides, padding, | ||
| 158 | + use_cudnn_on_gpu: use_cudnn_on_gpu, | ||
| 159 | + explicit_paddings: explicit_paddings, | ||
| 160 | + dilations: dilations, | ||
| 161 | + data_format: data_format), | ||
| 162 | + gen_nn_ops.conv2d(new Conv2dParams | ||
| 163 | + { | ||
| 164 | + Input = grad, | ||
| 165 | + Filter = op.inputs[1], | ||
| 166 | + Strides = strides, | ||
| 167 | + Padding = padding, | ||
| 168 | + DataFormat = data_format, | ||
| 169 | + Dilations = dilations, | ||
| 170 | + ExplicitPaddings = explicit_paddings, | ||
| 171 | + UseCudnnOnGpu = use_cudnn_on_gpu | ||
| 172 | + }) | ||
| 173 | + }; | ||
| 174 | + } | ||
| 175 | + | ||
| 136 | 176 | /// <summary> | |
| 137 | 177 | /// Gradient function for Conv2D. | |
| 138 | 178 | /// </summary> | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -283,7 +283,7 @@ public virtual Operation create_op(string op_type, Tensor[] inputs, TF_DataType[ | |||
| 283 | 283 | // This was causing duplicate graph node name errors, when testing a conv2d autoencoder | |
| 284 | 284 | // https://keras.io/guides/functional_api/#:~:text=keras.,graph%20(DAG)%20of%20layers. | |
| 285 | 285 | // name = name.EndsWith("/") ? ops.name_from_scope_name(name) : unique_name(name); | |
| 286 | - name = name.EndsWith("/") ? unique_name(ops.name_from_scope_name(name)) : unique_name(name); | ||
| 286 | + name = name.EndsWith("/") ? ops.name_from_scope_name(name) : unique_name(name); | ||
| 287 | 287 | var node_def = ops._NodeDef(op_type, name, attrs: attrs); | |
| 288 | 288 | ||
| 289 | 289 | var input_ops = inputs.Select(x => x.op).ToArray(); | |
@@ -386,10 +386,6 @@ public string name_scope(string name) | |||
| 386 | 386 | /// to name the operation being created.</returns> | |
| 387 | 387 | public string unique_name(string name, bool mark_as_used = true) | |
| 388 | 388 | { | |
| 389 | - if (name.EndsWith("basic_r_n_n_cell")) | ||
| 390 | - { | ||
| 391 | - | ||
| 392 | - } | ||
| 393 | 389 | if (!String.IsNullOrEmpty(_name_stack)) | |
| 394 | 390 | name = _name_stack + "/" + name; | |
| 395 | 391 | // For the sake of checking for names in use, we treat names as case | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -42,7 +42,7 @@ public class Conv2dParams | |||
| 42 | 42 | /// <summary> | |
| 43 | 43 | /// A 4-D tensor of shape | |
| 44 | 44 | /// </summary> | |
| 45 | - public IVariableV1 Filter { get; set; } | ||
| 45 | + public Tensor Filter { get; set; } | ||
| 46 | 46 | ||
| 47 | 47 | /// <summary> | |
| 48 | 48 | /// An integer vector representing the tensor shape of `filter` | |
| Original file line number | Diff line number | Diff line change | |
|---|---|---|---|
@@ -36,7 +36,7 @@ public ConvolutionInternal(ConvolutionalArgs args) | |||
| 36 | 36 | name = args.Name; | |
| 37 | 37 | } | |
| 38 | 38 | ||
| 39 | - public Tensor Apply(Tensors input, IVariableV1 filters) | ||
| 39 | + public Tensor Apply(Tensors input, Tensor filters) | ||
| 40 | 40 | { | |
| 41 | 41 | var filters_rank = filters.shape.ndim; | |
| 42 | 42 | var inputs_rank = input.shape.ndim; | |
| Back | FazBrowse Home | New Git URL |
0 commit comments