FazBrowse GitHub Viewer
|
Trending
|
URL:
|
Home
Tools:
[Download Repo ZIP]
[View Raw Code]
[Original HTTPS Page]
TensorFlow.NET/src/TensorFlowNET.Hub/MnistModelLoader.cs at master · baradgur/TensorFlow.NET · GitHub
baradgur
/
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.Hub
/
MnistModelLoader.cs
Copy path
More file actions
More file actions
Latest commit
History
History
History
184 lines (136 loc) · 7.96 KB
Breadcrumbs
TensorFlow.NET
/
src
/
TensorFlowNET.Hub
/
MnistModelLoader.cs
Copy path
File metadata and controls
184 lines (136 loc) · 7.96 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
using
System
;
using
System
.
Threading
.
Tasks
;
using
System
.
Collections
.
Generic
;
using
System
.
Text
;
using
System
.
IO
;
using
NumSharp
;
namespace
Tensorflow
.
Hub
{
public
class
MnistModelLoader
:
IModelLoader
<
MnistDataSet
>
{
private
const
string
DEFAULT_SOURCE_URL
=
"https://storage.googleapis.com/cvdf-datasets/mnist/"
;
private
const
string
TRAIN_IMAGES
=
"train-images-idx3-ubyte.gz"
;
private
const
string
TRAIN_LABELS
=
"train-labels-idx1-ubyte.gz"
;
private
const
string
TEST_IMAGES
=
"t10k-images-idx3-ubyte.gz"
;
private
const
string
TEST_LABELS
=
"t10k-labels-idx1-ubyte.gz"
;
public
static
async
Task
<
Datasets
<
MnistDataSet
>
>
LoadAsync
(
string
trainDir
,
bool
oneHot
=
false
,
int
?
trainSize
=
null
,
int
?
validationSize
=
null
,
int
?
testSize
=
null
,
bool
showProgressInConsole
=
false
)
{
var
loader
=
new
MnistModelLoader
(
)
;
var
setting
=
new
ModelLoadSetting
{
TrainDir
=
trainDir
,
OneHot
=
oneHot
,
ShowProgressInConsole
=
showProgressInConsole
}
;
if
(
trainSize
.
HasValue
)
setting
.
TrainSize
=
trainSize
.
Value
;
if
(
validationSize
.
HasValue
)
setting
.
ValidationSize
=
validationSize
.
Value
;
if
(
testSize
.
HasValue
)
setting
.
TestSize
=
testSize
.
Value
;
return
await
loader
.
LoadAsync
(
setting
)
;
}
public
async
Task
<
Datasets
<
MnistDataSet
>
>
LoadAsync
(
ModelLoadSetting
setting
)
{
if
(
setting
.
TrainSize
.
HasValue
&&
setting
.
ValidationSize
>=
setting
.
TrainSize
.
Value
)
throw
new
ArgumentException
(
"Validation set should be smaller than training set"
)
;
var
sourceUrl
=
setting
.
SourceUrl
;
if
(
string
.
IsNullOrEmpty
(
sourceUrl
)
)
sourceUrl
=
DEFAULT_SOURCE_URL
;
// load train images
await
this
.
DownloadAsync
(
sourceUrl
+
TRAIN_IMAGES
,
setting
.
TrainDir
,
TRAIN_IMAGES
,
showProgressInConsole
:
setting
.
ShowProgressInConsole
)
.
ShowProgressInConsole
(
setting
.
ShowProgressInConsole
)
;
await
this
.
UnzipAsync
(
Path
.
Combine
(
setting
.
TrainDir
,
TRAIN_IMAGES
)
,
setting
.
TrainDir
,
showProgressInConsole
:
setting
.
ShowProgressInConsole
)
.
ShowProgressInConsole
(
setting
.
ShowProgressInConsole
)
;
var
trainImages
=
ExtractImages
(
Path
.
Combine
(
setting
.
TrainDir
,
Path
.
GetFileNameWithoutExtension
(
TRAIN_IMAGES
)
)
,
limit
:
setting
.
TrainSize
)
;
// load train labels
await
this
.
DownloadAsync
(
sourceUrl
+
TRAIN_LABELS
,
setting
.
TrainDir
,
TRAIN_LABELS
,
showProgressInConsole
:
setting
.
ShowProgressInConsole
)
.
ShowProgressInConsole
(
setting
.
ShowProgressInConsole
)
;
await
this
.
UnzipAsync
(
Path
.
Combine
(
setting
.
TrainDir
,
TRAIN_LABELS
)
,
setting
.
TrainDir
,
showProgressInConsole
:
setting
.
ShowProgressInConsole
)
.
ShowProgressInConsole
(
setting
.
ShowProgressInConsole
)
;
var
trainLabels
=
ExtractLabels
(
Path
.
Combine
(
setting
.
TrainDir
,
Path
.
GetFileNameWithoutExtension
(
TRAIN_LABELS
)
)
,
one_hot
:
setting
.
OneHot
,
limit
:
setting
.
TrainSize
)
;
// load test images
await
this
.
DownloadAsync
(
sourceUrl
+
TEST_IMAGES
,
setting
.
TrainDir
,
TEST_IMAGES
,
showProgressInConsole
:
setting
.
ShowProgressInConsole
)
.
ShowProgressInConsole
(
setting
.
ShowProgressInConsole
)
;
await
this
.
UnzipAsync
(
Path
.
Combine
(
setting
.
TrainDir
,
TEST_IMAGES
)
,
setting
.
TrainDir
,
showProgressInConsole
:
setting
.
ShowProgressInConsole
)
.
ShowProgressInConsole
(
setting
.
ShowProgressInConsole
)
;
var
testImages
=
ExtractImages
(
Path
.
Combine
(
setting
.
TrainDir
,
Path
.
GetFileNameWithoutExtension
(
TEST_IMAGES
)
)
,
limit
:
setting
.
TestSize
)
;
// load test labels
await
this
.
DownloadAsync
(
sourceUrl
+
TEST_LABELS
,
setting
.
TrainDir
,
TEST_LABELS
,
showProgressInConsole
:
setting
.
ShowProgressInConsole
)
.
ShowProgressInConsole
(
setting
.
ShowProgressInConsole
)
;
await
this
.
UnzipAsync
(
Path
.
Combine
(
setting
.
TrainDir
,
TEST_LABELS
)
,
setting
.
TrainDir
,
showProgressInConsole
:
setting
.
ShowProgressInConsole
)
.
ShowProgressInConsole
(
setting
.
ShowProgressInConsole
)
;
var
testLabels
=
ExtractLabels
(
Path
.
Combine
(
setting
.
TrainDir
,
Path
.
GetFileNameWithoutExtension
(
TEST_LABELS
)
)
,
one_hot
:
setting
.
OneHot
,
limit
:
setting
.
TestSize
)
;
var
end
=
trainImages
.
shape
[
0
]
;
var
validationSize
=
setting
.
ValidationSize
;
var
validationImages
=
trainImages
[
np
.
arange
(
validationSize
)
]
;
var
validationLabels
=
trainLabels
[
np
.
arange
(
validationSize
)
]
;
trainImages
=
trainImages
[
np
.
arange
(
validationSize
,
end
)
]
;
trainLabels
=
trainLabels
[
np
.
arange
(
validationSize
,
end
)
]
;
var
dtype
=
setting
.
DataType
;
var
reshape
=
setting
.
ReShape
;
var
train
=
new
MnistDataSet
(
trainImages
,
trainLabels
,
dtype
,
reshape
)
;
var
validation
=
new
MnistDataSet
(
validationImages
,
validationLabels
,
dtype
,
reshape
)
;
var
test
=
new
MnistDataSet
(
testImages
,
testLabels
,
dtype
,
reshape
)
;
return
new
Datasets
<
MnistDataSet
>
(
train
,
validation
,
test
)
;
}
private
NDArray
ExtractImages
(
string
file
,
int
?
limit
=
null
)
{
if
(
!
Path
.
IsPathRooted
(
file
)
)
file
=
Path
.
Combine
(
AppContext
.
BaseDirectory
,
file
)
;
using
(
var
bytestream
=
new
FileStream
(
file
,
FileMode
.
Open
)
)
{
var
magic
=
Read32
(
bytestream
)
;
if
(
magic
!=
2051
)
throw
new
Exception
(
$
"Invalid magic number
{
magic
}
in MNIST image file:
{
file
}
"
)
;
var
num_images
=
Read32
(
bytestream
)
;
num_images
=
limit
==
null
?
num_images
:
Math
.
Min
(
num_images
,
(
int
)
limit
)
;
var
rows
=
Read32
(
bytestream
)
;
var
cols
=
Read32
(
bytestream
)
;
var
buf
=
new
byte
[
rows
*
cols
*
num_images
]
;
bytestream
.
Read
(
buf
,
0
,
buf
.
Length
)
;
var
data
=
np
.
frombuffer
(
buf
,
np
.
@byte
)
;
data
=
data
.
reshape
(
num_images
,
rows
,
cols
,
1
)
;
return
data
;
}
}
private
NDArray
ExtractLabels
(
string
file
,
bool
one_hot
=
false
,
int
num_classes
=
10
,
int
?
limit
=
null
)
{
if
(
!
Path
.
IsPathRooted
(
file
)
)
file
=
Path
.
Combine
(
AppContext
.
BaseDirectory
,
file
)
;
using
(
var
bytestream
=
new
FileStream
(
file
,
FileMode
.
Open
)
)
{
var
magic
=
Read32
(
bytestream
)
;
if
(
magic
!=
2049
)
throw
new
Exception
(
$
"Invalid magic number
{
magic
}
in MNIST label file:
{
file
}
"
)
;
var
num_items
=
Read32
(
bytestream
)
;
num_items
=
limit
==
null
?
num_items
:
Math
.
Min
(
num_items
,
(
int
)
limit
)
;
var
buf
=
new
byte
[
num_items
]
;
bytestream
.
Read
(
buf
,
0
,
buf
.
Length
)
;
var
labels
=
np
.
frombuffer
(
buf
,
np
.
uint8
)
;
if
(
one_hot
)
return
DenseToOneHot
(
labels
,
num_classes
)
;
return
labels
;
}
}
private
NDArray
DenseToOneHot
(
NDArray
labels_dense
,
int
num_classes
)
{
var
num_labels
=
labels_dense
.
shape
[
0
]
;
var
index_offset
=
np
.
arange
(
num_labels
)
*
num_classes
;
var
labels_one_hot
=
np
.
zeros
(
num_labels
,
num_classes
)
;
var
labels
=
labels_dense
.
Data
<
byte
>
(
)
;
for
(
int
row
=
0
;
row
<
num_labels
;
row
++
)
{
var
col
=
labels
[
row
]
;
labels_one_hot
.
SetData
(
1.0
,
row
,
col
)
;
}
return
labels_one_hot
;
}
private
int
Read32
(
FileStream
bytestream
)
{
var
buffer
=
new
byte
[
sizeof
(
uint
)
]
;
var
count
=
bytestream
.
Read
(
buffer
,
0
,
4
)
;
return
np
.
frombuffer
(
buffer
,
">u4"
)
.
Data
<
int
>
(
)
[
0
]
;
}
}
}
Back
|
FazBrowse Home
|
New Git URL