FazBrowse GitHub Viewer
|
Trending
|
URL:
|
Home
Tools:
[Download Repo ZIP]
[View Raw Code]
[Original HTTPS Page]
TensorFlow.NET/src/TensorFlowNET.Keras/Optimizers/RMSprop.cs at master · feelsyt/TensorFlow.NET · GitHub
feelsyt
/
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.Keras
/
Optimizers
/
RMSprop.cs
Copy path
More file actions
More file actions
Latest commit
History
History
History
78 lines (72 loc) · 2.99 KB
Breadcrumbs
TensorFlow.NET
/
src
/
TensorFlowNET.Keras
/
Optimizers
/
RMSprop.cs
Copy path
File metadata and controls
78 lines (72 loc) · 2.99 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
using
System
;
using
System
.
Collections
.
Generic
;
using
Tensorflow
.
Keras
.
ArgsDefinition
;
namespace
Tensorflow
.
Keras
.
Optimizers
{
/// <summary>
/// Optimizer that implements the RMSprop algorithm.
/// </summary>
public
class
RMSprop
:
OptimizerV2
{
RMSpropArgs
args
;
bool
centered
=>
args
.
Centered
;
protected
override
string
_name
=>
"RMSprop"
;
public
RMSprop
(
RMSpropArgs
args
)
:
base
(
args
)
{
this
.
args
=
args
;
_set_hyper
(
"rho"
,
args
.
RHO
)
;
_set_hyper
(
"momentum"
,
args
.
Momentum
)
;
}
protected
override
void
_create_slots
(
IVariableV1
[
]
var_list
)
{
foreach
(
var
var
in
var_list
)
add_slot
(
var
,
"rms"
)
;
if
(
_momentum
)
foreach
(
var
var
in
var_list
)
add_slot
(
var
,
"momentum"
)
;
if
(
centered
)
foreach
(
var
var
in
var_list
)
add_slot
(
var
,
"mg"
)
;
}
protected
override
void
_prepare_local
(
DeviceDType
device_dtype
,
Dictionary
<
DeviceDType
,
Dictionary
<
string
,
Tensor
>
>
_apply_state
)
{
base
.
_prepare_local
(
device_dtype
,
_apply_state
)
;
var
rho
=
array_ops
.
identity
(
_get_hyper
(
"rho"
,
device_dtype
.
DType
)
)
;
_apply_state
[
device_dtype
]
[
"neg_lr_t"
]
=
-
_apply_state
[
device_dtype
]
[
"lr_t"
]
;
_apply_state
[
device_dtype
]
[
"epsilon"
]
=
ops
.
convert_to_tensor
(
args
.
Epsilon
,
dtype
:
device_dtype
.
DType
)
;
_apply_state
[
device_dtype
]
[
"rho"
]
=
rho
;
_apply_state
[
device_dtype
]
[
"momentum"
]
=
array_ops
.
identity
(
_get_hyper
(
"momentum"
,
device_dtype
.
DType
)
)
;
_apply_state
[
device_dtype
]
[
"one_minus_rho"
]
=
1.0f
-
rho
;
}
protected
override
Operation
_resource_apply_dense
(
IVariableV1
var
,
Tensor
grad
,
Dictionary
<
DeviceDType
,
Dictionary
<
string
,
Tensor
>
>
_apply_state
)
{
Dictionary
<
string
,
Tensor
>
coefficients
=
null
;
foreach
(
var
state
in
_apply_state
)
{
if
(
state
.
Key
.
DType
==
var
.
dtype
.
as_base_dtype
(
)
&&
state
.
Key
.
Device
==
var
.
Device
)
{
coefficients
=
state
.
Value
;
break
;
}
}
var
rms
=
get_slot
(
var
,
"rms"
)
;
if
(
_momentum
)
{
throw
new
NotImplementedException
(
""
)
;
}
else
{
var
rms_t
=
coefficients
[
"rho"
]
*
rms
.
AsTensor
(
)
+
coefficients
[
"one_minus_rho"
]
*
math_ops
.
square
(
grad
)
;
rms_t
=
state_ops
.
assign
(
rms
,
rms_t
,
use_locking
:
_use_locking
)
;
var
denom_t
=
rms_t
;
if
(
centered
)
{
throw
new
NotImplementedException
(
""
)
;
}
var
var_t
=
var
.
AsTensor
(
)
-
coefficients
[
"lr_t"
]
*
grad
/
(
math_ops
.
sqrt
(
denom_t
)
+
coefficients
[
"epsilon"
]
)
;
return
state_ops
.
assign
(
var
,
var_t
,
use_locking
:
_use_locking
)
.
op
;
}
}
}
}
Back
|
FazBrowse Home
|
New Git URL