FazBrowse GitHub Viewer
|
Trending
|
URL:
|
Home
Tools:
[Download Repo ZIP]
[View Raw Code]
[Original HTTPS Page]
TensorFlow.NET/src/TensorFlowNET.Core/Operations/map_fn.cs at master · ehtick/TensorFlow.NET · GitHub
ehtick
/
TensorFlow.NET
Public
forked from
SciSharp/TensorFlow.NET
Notifications
You must be signed in to change notification settings
Fork
0
Star
0
Code
Pull requests
0
Actions
Projects
Security and quality
0
Insights
Additional navigation options
Code
Pull requests
Actions
Projects
Security and quality
Insights
Expand file tree
Breadcrumbs
TensorFlow.NET
/
src
/
TensorFlowNET.Core
/
Operations
/
map_fn.cs
Copy path
More file actions
More file actions
Latest commit
History
History
History
185 lines (156 loc) · 7.27 KB
Breadcrumbs
TensorFlow.NET
/
src
/
TensorFlowNET.Core
/
Operations
/
map_fn.cs
Copy path
File metadata and controls
185 lines (156 loc) · 7.27 KB
Raw
Copy raw file
Download raw file
Open symbols panel
Edit and raw actions
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
using
System
;
using
System
.
Collections
.
Generic
;
using
System
.
Linq
;
using
Tensorflow
.
Framework
;
using
Tensorflow
.
Util
;
using
static
Tensorflow
.
Binding
;
namespace
Tensorflow
{
#pragma warning disable
CS0659
// 'Operation' overrides Object.Equals(object o) but does not override Object.GetHashCode()
public
partial
class
Operation
#pragma warning restore
CS0659
// 'Operation' overrides Object.Equals(object o) but does not override Object.GetHashCode()
{
/// <summary>
/// map on the list of tensors unpacked from `elems` on dimension 0.
/// </summary>
/// <param name="fn"></param>
/// <param name="elems"></param>
/// <param name="dtype"></param>
/// <param name="parallel_iterations"></param>
/// <param name="back_prop"></param>
/// <param name="swap_memory"></param>
/// <param name="infer_shape"></param>
/// <param name="name"></param>
/// <returns>A tensor or (possibly nested) sequence of tensors.</returns>
public
static
Tensor
map_fn
(
Func
<
Tensor
,
Tensor
>
fn
,
Tensor
elems
,
TF_DataType
dtype
=
TF_DataType
.
DtInvalid
,
int
parallel_iterations
=
10
,
bool
back_prop
=
true
,
bool
swap_memory
=
false
,
bool
infer_shape
=
true
,
string
name
=
null
)
{
bool
input_is_sequence
=
nest
.
is_sequence
(
elems
)
;
Tensor
[
]
input_flatten
(
Tensor
x
)
=>
input_is_sequence
?
nest
.
flatten
(
x
)
.
ToArray
(
)
:
new
[
]
{
x
}
;
Tensor
input_pack
(
Tensor
[
]
x
)
=>
input_is_sequence
?
(
Tensor
)
nest
.
pack_sequence_as
(
elems
,
x
)
:
x
[
0
]
;
bool
output_is_sequence
;
Func
<
Tensor
,
Tensor
[
]
>
output_flatten
;
Func
<
Tensor
[
]
,
Tensor
>
output_pack
;
if
(
dtype
==
TF_DataType
.
DtInvalid
)
{
output_is_sequence
=
input_is_sequence
;
output_flatten
=
input_flatten
;
output_pack
=
input_pack
;
}
else
{
output_is_sequence
=
nest
.
is_sequence
(
dtype
)
;
output_flatten
=
(
x
)
=>
output_is_sequence
?
nest
.
flatten
(
x
)
.
ToArray
(
)
:
new
[
]
{
x
}
;
output_pack
=
(
x
)
=>
output_is_sequence
?
(
Tensor
)
nest
.
pack_sequence_as
(
dtype
,
x
)
:
x
[
0
]
;
}
var
elems_flat
=
input_flatten
(
elems
)
;
return
tf_with
(
ops
.
name_scope
(
name
,
"map"
,
elems_flat
)
,
delegate
{
//if in_graph_mode:
//# Any get_variable calls in fn will cache the first call locally
//# and not issue repeated network I/O requests for each iteration.
//varscope = vs.get_variable_scope()
//varscope_caching_device_was_none = False
//if varscope.caching_device is None:
// # TODO(ebrevdo): Change to using colocate_with here and in other
// # methods.
// varscope.set_caching_device(lambda op: op.device)
// varscope_caching_device_was_none = True
elems_flat
=
elems_flat
.
Select
(
elem
=>
ops
.
convert_to_tensor
(
elem
,
name
:
"elem"
)
)
.
ToArray
(
)
;
dtype
=
elems_flat
.
Select
(
elem
=>
elem
.
dtype
)
.
First
(
)
;
var
dtype_flat
=
new
[
]
{
dtype
}
;
// Convert elems to tensor array. n may be known statically.
var
static_shape
=
elems_flat
[
0
]
.
shape
;
var
n
=
static_shape
[
0
]
;
// TensorArrays are always flat
var
elems_ta
=
elems_flat
.
Select
(
elem
=>
tf
.
TensorArray
(
dtype
:
elem
.
dtype
,
size
:
Convert
.
ToInt32
(
n
)
,
dynamic_size
:
false
,
infer_shape
:
true
)
)
.
ToArray
(
)
;
// Unpack elements
var
elems_ta_1
=
new
List
<
TensorArray
>
(
)
;
foreach
(
var
(
elem_ta
,
elem
)
in
zip
(
elems_ta
,
elems_flat
)
)
elems_ta_1
.
Add
(
elem_ta
.
unstack
(
elem
)
)
;
elems_ta
=
elems_ta_1
.
ToArray
(
)
;
var
i
=
constant_op
.
constant
(
0
)
;
var
accs_ta
=
dtype_flat
.
Select
(
dt
=>
tf
.
TensorArray
(
dtype
:
dt
,
size
:
Convert
.
ToInt32
(
n
)
,
dynamic_size
:
false
,
infer_shape
:
infer_shape
)
)
.
ToArray
(
)
;
BodyItem
compute
(
BodyItem
item
)
{
var
packed_values
=
input_pack
(
elems_ta
.
Select
(
elem_ta
=>
elem_ta
.
read
(
item
.
I
)
)
.
ToArray
(
)
)
;
var
packed_fn_values
=
fn
(
packed_values
)
;
//nest.assert_same_structure(dtype or elems, packed_fn_values)
var
flat_fn_values
=
output_flatten
(
packed_fn_values
)
;
for
(
int
j
=
0
;
j
<
item
.
Accs_ta
.
Length
;
j
++
)
{
item
.
Accs_ta
[
j
]
.
write
(
item
.
I
,
flat_fn_values
[
j
]
)
;
}
return
new
BodyItem
(
item
.
I
+
1
,
item
.
Accs_ta
)
;
}
var
r_a
=
control_flow_ops
.
while_loop
(
(
x
)
=>
x
.
I
<
n
,
compute
,
new
BodyItem
(
i
,
accs_ta
)
,
parallel_iterations
:
parallel_iterations
,
back_prop
:
back_prop
,
swap_memory
:
swap_memory
,
maximum_iterations
:
tf
.
constant
(
n
)
)
;
var
results_flat
=
r_a
.
Accs_ta
.
Select
(
r
=>
r
.
stack
(
)
)
.
ToArray
(
)
;
var
n_static
=
new
Dimension
(
tensor_shape
.
dimension_value
(
elems_flat
[
0
]
.
shape
.
with_rank_at_least
(
1
)
.
dims
[
0
]
)
)
;
foreach
(
var
elem
in
elems_flat
.
Skip
(
1
)
)
{
n_static
.
merge_with
(
new
Dimension
(
tensor_shape
.
dimension_value
(
elem
.
shape
.
with_rank_at_least
(
1
)
.
dims
[
0
]
)
)
)
;
}
foreach
(
Tensor
r
in
results_flat
)
{
r
.
shape
=
new
Shape
(
n_static
)
.
concatenate
(
r
.
dims
.
Skip
(
1
)
.
ToArray
(
)
)
;
}
// todo get working when the above caching_device is fixed
//if (in_graph_mode && varscope_caching_device_was_none) {
// varscope.set_caching_device(None);
//}
return
output_pack
(
results_flat
)
;
}
)
;
}
internal
class
BodyItem
:
ICanBeFlattened
,
IPackable
<
BodyItem
>
,
IFromMergeVars
<
BodyItem
>
{
public
Tensor
I
{
get
;
set
;
}
public
TensorArray
[
]
Accs_ta
{
get
;
set
;
}
public
BodyItem
(
)
{
}
public
BodyItem
(
Tensor
i
,
TensorArray
[
]
accs_ta
)
{
I
=
i
;
Accs_ta
=
accs_ta
;
}
public
object
[
]
Flatten
(
)
{
var
elements
=
new
List
<
object
>
{
I
}
;
elements
.
AddRange
(
Accs_ta
)
;
return
elements
.
ToArray
(
)
;
}
public
BodyItem
Pack
(
object
[
]
sequences
)
{
I
=
sequences
[
0
]
as
Tensor
;
Accs_ta
=
new
[
]
{
sequences
[
1
]
as
TensorArray
}
;
return
new
BodyItem
(
I
,
Accs_ta
)
;
}
public
BodyItem
FromMergeVars
(
ITensorOrTensorArray
[
]
merge_vars
)
{
I
=
(
Tensor
)
merge_vars
[
1
]
;
Accs_ta
=
new
[
]
{
(
TensorArray
)
merge_vars
[
2
]
}
;
return
this
;
}
}
}
}
Back
|
FazBrowse Home
|
New Git URL