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

compute_gradients · attackgithub/TensorFlow.NET@4c2f76e · GitHub

Repository navigation

Commit 4c2f76e

Browse files
committed
compute_gradients
1 parent 62a2634 commit 4c2f76e

7 files changed

Lines changed: 130 additions & 1 deletion

File tree

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,14 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Text;
4+
5+
namespace Tensorflow
6+
{
7+
public static class Distribute
8+
{
9+
public static VariableAggregationType get_loss_reduction()
10+
{
11+
return VariableAggregationType.MEAN;
12+
}
13+
}
14+
}
Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,13 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Text;
4+
5+
namespace Tensorflow
6+
{
7+
public enum GateGradientType
8+
{
9+
GATE_NONE = 0,
10+
GATE_OP = 1,
11+
GATE_GRAPH = 2
12+
}
13+
}
Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,16 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Text;
4+
5+
namespace Tensorflow
6+
{
7+
public class GradientDescentOptimizer : Optimizer
8+
{
9+
public GradientDescentOptimizer(double learning_rate, bool use_locking = false, string name = "GradientDescent")
10+
: base(learning_rate, use_locking, name)
11+
{
12+
LearningRate = learning_rate;
13+
LearningRateTensor = null;
14+
}
15+
}
16+
}
Lines changed: 55 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,55 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Text;
4+
using distribute_lib = Tensorflow.Distribute;
5+
6+
namespace Tensorflow
7+
{
8+
/// <summary>
9+
/// Base class for optimizers.
10+
/// This class defines the API to add Ops to train a model. You never use this
11+
/// class directly, but instead instantiate one of its subclasses such as
12+
/// `GradientDescentOptimizer`, `AdagradOptimizer`, or `MomentumOptimizer`.
13+
/// </summary>
14+
public abstract class Optimizer
15+
{
16+
public string Name { get; set; }
17+
public double LearningRate { get; set; }
18+
public Tensor LearningRateTensor { get; set; }
19+
20+
public Optimizer(double learning_rate, bool use_locking, string name = "")
21+
{
22+
if (String.IsNullOrEmpty(name))
23+
throw new NotImplementedException("Must specify the optimizer name");
24+
25+
Name = name;
26+
}
27+
28+
/// <summary>
29+
/// Add operations to minimize `loss` by updating `var_list`
30+
/// </summary>
31+
/// <param name="loss"></param>
32+
/// <returns></returns>
33+
public Optimizer minimize(Tensor loss, GateGradientType gate_gradients = GateGradientType.GATE_OP)
34+
{
35+
compute_gradients(loss, gate_gradients);
36+
return this;
37+
}
38+
39+
/// <summary>
40+
/// Compute gradients of `loss` for the variables in `var_list`.
41+
/// </summary>
42+
/// <param name="loss"></param>
43+
/// <param name="gate_gradients"></param>
44+
public List<KeyValuePair<object, object>> compute_gradients(Tensor loss, GateGradientType gate_gradients = GateGradientType.GATE_OP)
45+
{
46+
int num_towers = 1;
47+
if(distribute_lib.get_loss_reduction() == VariableAggregationType.MEAN)
48+
{
49+
50+
}
51+
52+
return null;
53+
}
54+
}
55+
}
Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,14 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Text;
4+
5+
namespace Tensorflow
6+
{
7+
public enum VariableAggregationType
8+
{
9+
NONE = 0,
10+
SUM = 1,
11+
MEAN = 2,
12+
ONLY_FIRST_TOWER = 3
13+
}
14+
}
Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,17 @@
1+
using System;
2+
using System.Collections.Generic;
3+
using System.Text;
4+
5+
namespace Tensorflow
6+
{
7+
public static partial class tf
8+
{
9+
public static class train
10+
{
11+
public static Optimizer GradientDescentOptimizer(double learning_rate)
12+
{
13+
return new GradientDescentOptimizer(learning_rate);
14+
}
15+
}
16+
}
17+
}

‎test/TensorFlowNET.Examples/LinearRegression.cs‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,7 @@ public void Run()
4747

4848
// radient descent
4949
// Note, minimize() knows to modify W and b because Variable objects are trainable=True by default
50-
// var optimizer = tf.train.GradientDescentOptimizer(learning_rate).minimize(cost);
50+
var optimizer = tf.train.GradientDescentOptimizer(learning_rate).minimize(cost);
5151
}
5252
}
5353
}

0 commit comments

Comments
 (0)

Back | FazBrowse Home | New Git URL