FazBrowse GitHub Viewer | Trending |
URL:
| Home
Tools: [Download Repo ZIP]   [Original HTTPS Page]

add tf.image.decode_image #350 · feelsyt/TensorFlow.NET@9d90d74 · GitHub

Commit 9d90d74

Browse files
committed
add tf.image.decode_image SciSharp#350
1 parent f004ab6 commit 9d90d74

9 files changed

Lines changed: 332 additions & 1 deletion

File tree

‎src/TensorFlowNET.Core/APIs/tf.image.cs‎

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,20 @@ public Tensor resize_bilinear(Tensor images, Tensor size, bool align_corners = f
4242

4343
public Tensor convert_image_dtype(Tensor image, TF_DataType dtype, bool saturate = false, string name = null)
4444
=> gen_image_ops.convert_image_dtype(image, dtype, saturate: saturate, name: name);
45+
46+
public Tensor decode_image(Tensor contents, int channels = 0, TF_DataType dtype = TF_DataType.TF_UINT8,
47+
string name = null, bool expand_animations = true)
48+
=> image_ops_impl.decode_image(contents, channels: channels, dtype: dtype,
49+
name: name, expand_animations: expand_animations);
50+
51+
/// <summary>
52+
/// Convenience function to check if the 'contents' encodes a JPEG image.
53+
/// </summary>
54+
/// <param name="contents"></param>
55+
/// <param name="name"></param>
56+
/// <returns></returns>
57+
public static Tensor is_jpeg(Tensor contents, string name = null)
58+
=> image_ops_impl.is_jpeg(contents, name: name);
4559
}
4660
}
4761
}
Lines changed: 32 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,32 @@
1+
/*****************************************************************************
2+
Copyright 2018 The TensorFlow.NET Authors. All Rights Reserved.
3+
4+
Licensed under the Apache License, Version 2.0 (the "License");
5+
you may not use this file except in compliance with the License.
6+
You may obtain a copy of the License at
7+
8+
http://www.apache.org/licenses/LICENSE-2.0
9+
10+
Unless required by applicable law or agreed to in writing, software
11+
distributed under the License is distributed on an "AS IS" BASIS,
12+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13+
See the License for the specific language governing permissions and
14+
limitations under the License.
15+
******************************************************************************/
16+
17+
using System.Collections.Generic;
18+
using Tensorflow.IO;
19+
20+
namespace Tensorflow
21+
{
22+
public partial class tensorflow
23+
{
24+
public strings_internal strings = new strings_internal();
25+
public class strings_internal
26+
{
27+
public Tensor substr(Tensor input, int pos, int len,
28+
string name = null, string @uint = "BYTE")
29+
=> string_ops.substr(input, pos, len, name: name, @uint: @uint);
30+
}
31+
}
32+
}

‎src/TensorFlowNET.Core/Operations/gen_image_ops.py.cs‎

Lines changed: 63 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -88,6 +88,69 @@ public static Tensor decode_jpeg(Tensor contents,
8888
}
8989
}
9090

91+
public static Tensor decode_gif(Tensor contents,
92+
string name = null)
93+
{
94+
// Add nodes to the TensorFlow graph.
95+
if (tf.context.executing_eagerly())
96+
{
97+
throw new NotImplementedException("decode_gif");
98+
}
99+
else
100+
{
101+
var _op = _op_def_lib._apply_op_helper("DecodeGif", name: name, args: new
102+
{
103+
contents
104+
});
105+
106+
return _op.output;
107+
}
108+
}
109+
110+
public static Tensor decode_png(Tensor contents,
111+
int channels = 0,
112+
TF_DataType dtype = TF_DataType.TF_UINT8,
113+
string name = null)
114+
{
115+
// Add nodes to the TensorFlow graph.
116+
if (tf.context.executing_eagerly())
117+
{
118+
throw new NotImplementedException("decode_png");
119+
}
120+
else
121+
{
122+
var _op = _op_def_lib._apply_op_helper("DecodePng", name: name, args: new
123+
{
124+
contents,
125+
channels,
126+
dtype
127+
});
128+
129+
return _op.output;
130+
}
131+
}
132+
133+
public static Tensor decode_bmp(Tensor contents,
134+
int channels = 0,
135+
string name = null)
136+
{
137+
// Add nodes to the TensorFlow graph.
138+
if (tf.context.executing_eagerly())
139+
{
140+
throw new NotImplementedException("decode_bmp");
141+
}
142+
else
143+
{
144+
var _op = _op_def_lib._apply_op_helper("DecodeBmp", name: name, args: new
145+
{
146+
contents,
147+
channels
148+
});
149+
150+
return _op.output;
151+
}
152+
}
153+
91154
public static Tensor resize_bilinear(Tensor images, Tensor size, bool align_corners = false, string name = null)
92155
{
93156
if (tf.context.executing_eagerly())
Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,42 @@
1+
/*****************************************************************************
2+
Copyright 2018 The TensorFlow.NET Authors. All Rights Reserved.
3+
4+
Licensed under the Apache License, Version 2.0 (the "License");
5+
you may not use this file except in compliance with the License.
6+
You may obtain a copy of the License at
7+
8+
http://www.apache.org/licenses/LICENSE-2.0
9+
10+
Unless required by applicable law or agreed to in writing, software
11+
distributed under the License is distributed on an "AS IS" BASIS,
12+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13+
See the License for the specific language governing permissions and
14+
limitations under the License.
15+
******************************************************************************/
16+
17+
using System;
18+
using System.Collections.Generic;
19+
using System.Text;
20+
21+
namespace Tensorflow
22+
{
23+
public class gen_string_ops
24+
{
25+
static readonly OpDefLibrary _op_def_lib;
26+
static gen_string_ops() { _op_def_lib = new OpDefLibrary(); }
27+
28+
public static Tensor substr(Tensor input, int pos, int len,
29+
string name = null, string @uint = "BYTE")
30+
{
31+
var _op = _op_def_lib._apply_op_helper("Substr", name: name, args: new
32+
{
33+
input,
34+
pos,
35+
len,
36+
unit = @uint
37+
});
38+
39+
return _op.output;
40+
}
41+
}
42+
}

‎src/TensorFlowNET.Core/Operations/image_ops_impl.cs‎

Lines changed: 105 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,11 +17,116 @@ limitations under the License.
1717
using System;
1818
using System.Collections.Generic;
1919
using System.Text;
20+
using static Tensorflow.Binding;
2021

2122
namespace Tensorflow
2223
{
2324
public class image_ops_impl
2425
{
26+
public static Tensor decode_image(Tensor contents, int channels = 0, TF_DataType dtype = TF_DataType.TF_UINT8,
27+
string name = null, bool expand_animations = true)
28+
{
29+
Tensor substr = null;
2530

31+
Func<ITensorOrOperation> _jpeg = () =>
32+
{
33+
int jpeg_channels = channels;
34+
var good_channels = math_ops.not_equal(jpeg_channels, 4, name: "check_jpeg_channels");
35+
string channels_msg = "Channels must be in (None, 0, 1, 3) when decoding JPEG 'images'";
36+
var assert_channels = control_flow_ops.Assert(good_channels, new string[] { channels_msg });
37+
return tf_with(ops.control_dependencies(new[] { assert_channels }), delegate
38+
{
39+
return convert_image_dtype(gen_image_ops.decode_jpeg(contents, channels), dtype);
40+
});
41+
};
42+
43+
Func<ITensorOrOperation> _gif = () =>
44+
{
45+
int gif_channels = channels;
46+
var good_channels = math_ops.logical_and(
47+
math_ops.not_equal(gif_channels, 1, name: "check_gif_channels"),
48+
math_ops.not_equal(gif_channels, 4, name: "check_gif_channels"));
49+
50+
string channels_msg = "Channels must be in (None, 0, 3) when decoding GIF images";
51+
var assert_channels = control_flow_ops.Assert(good_channels, new string[] { channels_msg });
52+
return tf_with(ops.control_dependencies(new[] { assert_channels }), delegate
53+
{
54+
var result = convert_image_dtype(gen_image_ops.decode_gif(contents), dtype);
55+
if (!expand_animations)
56+
// result = array_ops.gather(result, 0);
57+
throw new NotImplementedException("");
58+
return result;
59+
});
60+
};
61+
62+
Func<ITensorOrOperation> _bmp = () =>
63+
{
64+
int bmp_channels = channels;
65+
var signature = string_ops.substr(contents, 0, 2);
66+
var is_bmp = math_ops.equal(signature, "BM", name: "is_bmp");
67+
string decode_msg = "Unable to decode bytes as JPEG, PNG, GIF, or BMP";
68+
var assert_decode = control_flow_ops.Assert(is_bmp, new string[] { decode_msg });
69+
var good_channels = math_ops.not_equal(bmp_channels, 1, name: "check_channels");
70+
string channels_msg = "Channels must be in (None, 0, 3) when decoding BMP images";
71+
var assert_channels = control_flow_ops.Assert(good_channels, new string[] { channels_msg });
72+
return tf_with(ops.control_dependencies(new[] { assert_decode, assert_channels }), delegate
73+
{
74+
return convert_image_dtype(gen_image_ops.decode_bmp(contents), dtype);
75+
});
76+
};
77+
78+
Func<ITensorOrOperation> _png = () =>
79+
{
80+
return convert_image_dtype(gen_image_ops.decode_png(
81+
contents,
82+
channels,
83+
dtype: dtype),
84+
dtype);
85+
};
86+
87+
Func<ITensorOrOperation> check_gif = () =>
88+
{
89+
var is_gif = math_ops.equal(substr, "\x47\x49\x46", name: "is_gif");
90+
return control_flow_ops.cond(is_gif, _gif, _bmp, name: "cond_gif");
91+
};
92+
93+
Func<ITensorOrOperation> check_png = () =>
94+
{
95+
return control_flow_ops.cond(_is_png(contents), _png, check_gif, name: "cond_png");
96+
};
97+
98+
return tf_with(ops.name_scope(name, "decode_image"), scope =>
99+
{
100+
substr = string_ops.substr(contents, 0, 3);
101+
return control_flow_ops.cond(is_jpeg(contents), _jpeg, check_png, name: "cond_jpeg");
102+
});
103+
}
104+
105+
public static Tensor is_jpeg(Tensor contents, string name = null)
106+
{
107+
return tf_with(ops.name_scope(name, "is_jpeg"), scope =>
108+
{
109+
var substr = string_ops.substr(contents, 0, 3);
110+
return math_ops.equal(substr, "\xff\xd8\xff", name: name);
111+
});
112+
}
113+
114+
public static Tensor _is_png(Tensor contents, string name = null)
115+
{
116+
return tf_with(ops.name_scope(name, "is_png"), scope =>
117+
{
118+
var substr = string_ops.substr(contents, 0, 3);
119+
return math_ops.equal(substr, @"\211PN", name: name);
120+
});
121+
}
122+
123+
public static Tensor convert_image_dtype(Tensor image, TF_DataType dtype, bool saturate = false,
124+
string name = null)
125+
{
126+
if (dtype == image.dtype)
127+
return array_ops.identity(image, name: name);
128+
129+
throw new NotImplementedException("");
130+
}
26131
}
27132
}

‎src/TensorFlowNET.Core/Operations/math_ops.cs‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -168,6 +168,9 @@ public static Tensor sqrt(Tensor x, string name = null)
168168
public static Tensor multiply<Tx, Ty>(Tx x, Ty y, string name = null)
169169
=> gen_math_ops.mul(x, y, name: name);
170170

171+
public static Tensor not_equal<Tx, Ty>(Tx x, Ty y, string name = null)
172+
=> gen_math_ops.not_equal(x, y, name: name);
173+
171174
public static Tensor mul_no_nan<Tx, Ty>(Tx x, Ty y, string name = null)
172175
=> gen_math_ops.mul_no_nan(x, y, name: name);
173176

@@ -264,6 +267,9 @@ public static Tensor log(Tensor x, string name = null)
264267
return gen_math_ops.log(x, name);
265268
}
266269

270+
public static Tensor logical_and(Tensor x, Tensor y, string name = null)
271+
=> gen_math_ops.logical_and(x, y, name: name);
272+
267273
public static Tensor lgamma(Tensor x, string name = null)
268274
=> gen_math_ops.lgamma(x, name: name);
269275

Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,38 @@
1+
/*****************************************************************************
2+
Copyright 2018 The TensorFlow.NET Authors. All Rights Reserved.
3+
4+
Licensed under the Apache License, Version 2.0 (the "License");
5+
you may not use this file except in compliance with the License.
6+
You may obtain a copy of the License at
7+
8+
http://www.apache.org/licenses/LICENSE-2.0
9+
10+
Unless required by applicable law or agreed to in writing, software
11+
distributed under the License is distributed on an "AS IS" BASIS,
12+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
13+
See the License for the specific language governing permissions and
14+
limitations under the License.
15+
******************************************************************************/
16+
17+
using System;
18+
using System.Collections.Generic;
19+
using System.Text;
20+
21+
namespace Tensorflow
22+
{
23+
public class string_ops
24+
{
25+
/// <summary>
26+
/// Return substrings from `Tensor` of strings.
27+
/// </summary>
28+
/// <param name="input"></param>
29+
/// <param name="pos"></param>
30+
/// <param name="len"></param>
31+
/// <param name="name"></param>
32+
/// <param name="uint"></param>
33+
/// <returns></returns>
34+
public static Tensor substr(Tensor input, int pos, int len,
35+
string name = null, string @uint = "BYTE")
36+
=> gen_string_ops.substr(input, pos, len, name: name, @uint: @uint);
37+
}
38+
}

‎test/TensorFlowNET.UnitTest/GraphTest.cs‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -416,12 +416,13 @@ public void ImportGraphDef_MissingUnusedInputMappings()
416416

417417
}
418418

419+
[TestMethod]
419420
public void ImportGraphMeta()
420421
{
421422
var dir = "my-save-dir/";
422423
using (var sess = tf.Session())
423424
{
424-
var new_saver = tf.train.import_meta_graph(dir + "my-model-10000.meta");
425+
var new_saver = tf.train.import_meta_graph(@"D:\tmp\resnet_v2_101_2017_04_14\eval.graph");
425426
new_saver.restore(sess, dir + "my-model-10000");
426427
var labels = tf.constant(0, dtype: tf.int32, shape: new int[] { 100 }, name: "labels");
427428
var batch_size = tf.size(labels);
Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,30 @@
1+
using Microsoft.VisualStudio.TestTools.UnitTesting;
2+
using System;
3+
using System.Collections.Generic;
4+
using System.IO;
5+
using System.Text;
6+
using Tensorflow;
7+
using static Tensorflow.Binding;
8+
9+
namespace TensorFlowNET.UnitTest
10+
{
11+
[TestClass]
12+
public class ImageTest
13+
{
14+
string imgPath = "../../../../../data/shasta-daisy.jpg";
15+
Tensor contents;
16+
17+
public ImageTest()
18+
{
19+
imgPath = Path.GetFullPath(imgPath);
20+
contents = tf.read_file(imgPath);
21+
}
22+
23+
[TestMethod]
24+
public void decode_image()
25+
{
26+
var img = tf.image.decode_image(contents);
27+
Assert.AreEqual(img.name, "decode_image/cond_jpeg/Merge:0");
28+
}
29+
}
30+
}

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL