Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
Show all changes
22 commits
Select commit Hold shift + click to select a range
9e92687
Initial commit of gradient descent optimizers.
Craigacp Nov 5, 2019
68a353c
Adding Apache 2.0 license header to all optimizer files.
Craigacp Nov 5, 2019
b3f4be8
Bug fix for the MNISTTest.
Craigacp Dec 6, 2019
d1868ea
Refactor to uptake latest tensorflow-core changes.
Craigacp Jan 31, 2020
6d189cc
Added type safety and updates for new api.
Craigacp Jan 31, 2020
53e438a
Small changes, plus a fix for DataTypes to include references to the …
Craigacp Jan 31, 2020
83140b4
Repackaging the optimizers into tensorflow-training, org.tensorflow.t…
Craigacp Feb 7, 2020
b2ac923
Initial commit of gradient descent optimizers.
Craigacp Nov 5, 2019
e7eb2e8
Adding Apache 2.0 license header to all optimizer files.
Craigacp Nov 5, 2019
3d63564
Bug fix for the MNISTTest.
Craigacp Dec 6, 2019
ed71dc5
Refactor to uptake latest tensorflow-core changes.
Craigacp Jan 31, 2020
b054449
Added type safety and updates for new api.
Craigacp Jan 31, 2020
b29be50
Repackaging the optimizers into tensorflow-training, org.tensorflow.t…
Craigacp Feb 7, 2020
6cdb55c
Delete pom.xml
Craigacp Feb 8, 2020
b9d64c5
Googlify with IntelliJ's Google Java Style Guide formatter.
Craigacp Feb 11, 2020
6ae5ace
Bumping the copyright year, and switching to try-with-resources in th…
Craigacp Feb 12, 2020
56429e8
Updating variableWithInit to use @Endpoint.
Craigacp Feb 25, 2020
5d8cb69
Refactorings after code review.
Craigacp Feb 25, 2020
51f5d47
Adding a couple of lines to the gitignore.
Craigacp Feb 25, 2020
66876ed
Adding a bit of documentation, threading the named operations through…
Craigacp Feb 25, 2020
7a2fd25
Adding a guard to prevent variableWithInit being called on an EagerSe…
Craigacp Feb 25, 2020
1b98f52
Update Ops.java
karllessard Mar 2, 2020
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Prev Previous commit
Next Next commit
Added type safety and updates for new api.
  • Loading branch information
Craigacp committed Feb 25, 2020
commit b0544494775ede7f856f454cfc2f3f8c20303f24
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,6 @@
import org.tensorflow.Graph;
import org.tensorflow.Operand;
import org.tensorflow.Output;
import org.tensorflow.op.core.Constant;
import org.tensorflow.op.core.Variable;
import org.tensorflow.types.TFloat32;
import org.tensorflow.types.family.TType;
Expand Down Expand Up @@ -60,18 +59,16 @@ protected void createSlots(List<Output<? extends TType>> variables) {
}

private <T extends TType> void createAdaDeltaSlot(Output<T> v) {
Operand<T> accumulatorInitializer = tf.fill(tf.shape(v), (Constant<T>) tf.constant(0.0f, TFloat32.DTYPE));//v.dataType()));
Operand<T> accumulatorInitializer = tf.fill(tf.shape(v), tf.dtypes.cast(tf.constant(0.0f, TFloat32.DTYPE),v.dataType()));
createSlot(v.asOutput(), ACCUMULATOR, accumulatorInitializer);
Operand<T> updateInitializer = tf.fill(tf.shape(v), (Constant<T>) tf.constant(0.0f, TFloat32.DTYPE));//v.dataType()));
Operand<T> updateInitializer = tf.fill(tf.shape(v), tf.dtypes.cast(tf.constant(0.0f, TFloat32.DTYPE),v.dataType()));
createSlot(v.asOutput(), ACCUMULATOR_UPDATE, updateInitializer);
}

@Override
protected <T extends TType> Operand<T> applyDense(Output<T> gradient, Output<T> variable) {
@SuppressWarnings("unchecked") // suppressed as the slots are created to have the dtype of the variable.
Variable<T> accumSlot = (Variable<T>) getSlot(variable,ACCUMULATOR).get();
@SuppressWarnings("unchecked")
Variable<T> accumUpdateSlot = (Variable<T>) getSlot(variable,ACCUMULATOR_UPDATE).get();
Variable<T> accumSlot = getSlot(variable,ACCUMULATOR).get();
Variable<T> accumUpdateSlot = getSlot(variable,ACCUMULATOR_UPDATE).get();
return tf.train.applyAdadelta(variable, accumSlot, accumUpdateSlot,
tf.constant(learningRate, gradient.dataType()),
tf.constant(rho, gradient.dataType()),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,6 @@
import org.tensorflow.Graph;
import org.tensorflow.Operand;
import org.tensorflow.Output;
import org.tensorflow.op.core.Constant;
import org.tensorflow.op.core.Variable;
import org.tensorflow.types.TFloat32;
import org.tensorflow.types.family.TType;
Expand Down Expand Up @@ -57,14 +56,13 @@ protected void createSlots(List<Output<? extends TType>> variables) {
}

private <T extends TType> void createAdaGradSlot(Output<T> v) {
Operand<T> initializer = tf.fill(tf.shape(v), (Constant<T>) tf.constant(initialAccumulatorValue, TFloat32.DTYPE));//v.dataType()));
Operand<T> initializer = tf.fill(tf.shape(v), tf.dtypes.cast(tf.constant(initialAccumulatorValue, TFloat32.DTYPE),v.dataType()));
createSlot(v.asOutput(), ACCUMULATOR, initializer);
}

@Override
protected <T extends TType> Operand<T> applyDense(Output<T> gradient, Output<T> variable) {
@SuppressWarnings("unchecked") // suppressed as the slots are created to have the dtype of the variable.
Variable<T> slot = (Variable<T>) getSlot(variable,ACCUMULATOR).get();
Variable<T> slot = getSlot(variable,ACCUMULATOR).get();
return tf.train.applyAdagrad(variable, slot, tf.constant(learningRate, gradient.dataType()), gradient);
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,6 @@
import org.tensorflow.Output;
import org.tensorflow.op.Op;
import org.tensorflow.op.core.Assign;
import org.tensorflow.op.core.Constant;
import org.tensorflow.op.core.Variable;
import org.tensorflow.tools.Shape;
import org.tensorflow.types.TFloat32;
Expand Down Expand Up @@ -63,7 +62,7 @@ public AdaGradDA(Graph graph, float learningRate, float initialAccumulatorValue,
}

@Override
protected Optional<Operand> prepare(String name) {
protected Optional<Operand<?>> prepare(String name) {
return Optional.of(tf.assignAdd(globalStep,tf.constant(1L)));
}

Expand All @@ -78,18 +77,16 @@ protected void createSlots(List<Output<? extends TType>> variables) {
}

private <T extends TType> void createAdaGradDASlot(Output<T> v) {
Operand<T> initializer = tf.fill(tf.shape(v), (Constant<T>) tf.constant(0.0f, TFloat32.DTYPE));//v.dataType()));
Operand<T> initializer = tf.fill(tf.shape(v), tf.dtypes.cast(tf.constant(0.0f, TFloat32.DTYPE),v.dataType()));
createSlot(v.asOutput(), ACCUMULATOR, initializer);
Operand<T> sqInitializer = tf.fill(tf.shape(v), (Constant<T>) tf.constant(initialAccumulatorValue, TFloat32.DTYPE));//v.dataType()));
Operand<T> sqInitializer = tf.fill(tf.shape(v), tf.dtypes.cast(tf.constant(initialAccumulatorValue, TFloat32.DTYPE),v.dataType()));
createSlot(v.asOutput(), SQUARED_ACCUMULATOR, sqInitializer);
}

@Override
protected <T extends TType> Operand<T> applyDense(Output<T> gradient, Output<T> variable) {
@SuppressWarnings("unchecked") // suppressed as the slots are created to have the dtype of the variable.
Variable<T> gradSlot = (Variable<T>) getSlot(variable,ACCUMULATOR).get();
@SuppressWarnings("unchecked")
Variable<T> gradSquaredSlot = (Variable<T>) getSlot(variable,SQUARED_ACCUMULATOR).get();
Variable<T> gradSlot = getSlot(variable,ACCUMULATOR).get();
Variable<T> gradSquaredSlot = getSlot(variable,SQUARED_ACCUMULATOR).get();
return tf.train.applyAdagradDa(variable, gradSlot, gradSquaredSlot, gradient,
tf.constant(learningRate, gradient.dataType()),
tf.constant(l1Strength, gradient.dataType()),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -69,7 +69,7 @@ public Adam(Graph graph, float learningRate, float betaOne, float betaTwo, float
@Override
protected void createSlots(List<Output<? extends TType>> variables) {
for (Output<? extends TType> v : variables) {
createAdamSlot(v);
createAdamSlot(v.asOutput());
}
betaOnePower = tf.withName("beta1_power").variable(Shape.scalar(),TFloat32.DTYPE);
Assign<TFloat32> betaOnePowerInit = tf.assign(betaOnePower, tf.constant(betaOne, TFloat32.DTYPE));
Expand All @@ -80,7 +80,7 @@ protected void createSlots(List<Output<? extends TType>> variables) {
}

@Override
protected Optional<Operand> prepare(String scopeName) {
protected Optional<Operand<?>> prepare(String scopeName) {
betaOneConst = tf.constant(betaOne);
betaTwoConst = tf.constant(betaTwo);
learningRateConst = tf.constant(learningRate);
Expand All @@ -89,18 +89,16 @@ protected Optional<Operand> prepare(String scopeName) {
}

private <T extends TType> void createAdamSlot(Output<T> v) {
Operand<T> firstMomentInitializer = tf.fill(tf.shape(v), (Constant<T>) tf.constant(0.0f, TFloat32.DTYPE));//v.dataType()));
Operand<T> firstMomentInitializer = tf.fill(tf.shape(v), tf.dtypes.cast(tf.constant(0.0f, TFloat32.DTYPE),v.dataType()));
createSlot(v.asOutput(), FIRST_MOMENT, firstMomentInitializer);
Operand<T> secondMomentInitializer = tf.fill(tf.shape(v), (Constant<T>) tf.constant(0.0f, TFloat32.DTYPE));//v.dataType()));
Operand<T> secondMomentInitializer = tf.fill(tf.shape(v), tf.dtypes.cast(tf.constant(0.0f, TFloat32.DTYPE),v.dataType()));
createSlot(v.asOutput(), SECOND_MOMENT, secondMomentInitializer);
}

@Override
protected <T extends TType> Operand<T> applyDense(Output<T> gradient, Output<T> variable) {
@SuppressWarnings("unchecked") // suppressed as the slots are created to have the dtype of the variable.
Variable<T> firstMomentSlot = (Variable<T>) getSlot(variable,FIRST_MOMENT).get();
@SuppressWarnings("unchecked") // suppressed as the slots are created to have the dtype of the variable.
Variable<T> secondMomentSlot = (Variable<T>) getSlot(variable,SECOND_MOMENT).get();
Variable<T> firstMomentSlot = getSlot(variable,FIRST_MOMENT).get();
Variable<T> secondMomentSlot = getSlot(variable,SECOND_MOMENT).get();
return tf.train.applyAdam(variable, firstMomentSlot, secondMomentSlot,
tf.dtypes.cast(betaOnePower,gradient.dataType()),
tf.dtypes.cast(betaTwoPower,gradient.dataType()),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
import org.tensorflow.Graph;
import org.tensorflow.Operand;
import org.tensorflow.Output;
import org.tensorflow.types.TFloat32;
import org.tensorflow.types.family.TType;


Expand All @@ -35,7 +36,7 @@ public GradientDescent(Graph graph, float learningRate) {

@Override
protected <T extends TType> Operand<T> applyDense(Output<T> gradient, Output<T> variable) {
return tf.train.applyGradientDescent(variable, tf.constant(learningRate, gradient.dataType()), gradient);
return tf.train.applyGradientDescent(variable, tf.dtypes.cast(tf.constant(learningRate, TFloat32.DTYPE), gradient.dataType()), gradient);
}

@Override
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,6 @@
import org.tensorflow.Graph;
import org.tensorflow.Operand;
import org.tensorflow.Output;
import org.tensorflow.op.core.Constant;
import org.tensorflow.op.core.Variable;
import org.tensorflow.op.train.ApplyMomentum;
import org.tensorflow.types.TFloat32;
Expand Down Expand Up @@ -57,14 +56,13 @@ protected void createSlots(List<Output<? extends TType>> variables) {
}

private <T extends TType> void createMomentumSlot(Output<T> v) {
Operand<T> initializer = tf.fill(tf.shape(v), (Constant<T>) tf.constant(0.0f, TFloat32.DTYPE));//v.dataType()));
Operand<T> initializer = tf.fill(tf.shape(v), tf.dtypes.cast(tf.constant(0.0f, TFloat32.DTYPE),v.dataType()));
createSlot(v.asOutput(), MOMENTUM, initializer);
}

@Override
protected <T extends TType> Operand<T> applyDense(Output<T> gradient, Output<T> variable) {
@SuppressWarnings("unchecked") // suppressed as the slots are created to have the dtype of the variable.
Variable<T> slot = (Variable<T>) getSlot(variable,MOMENTUM).get();
Variable<T> slot = getSlot(variable,MOMENTUM).get();
return tf.train.applyMomentum(variable, slot, tf.constant(learningRate, gradient.dataType()), gradient, tf.constant(momentum, gradient.dataType()), ApplyMomentum.useNesterov(useNesterov));
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -25,9 +25,7 @@
import org.tensorflow.op.core.Assign;
import org.tensorflow.op.core.NoOp;
import org.tensorflow.op.core.Variable;
import org.tensorflow.types.TFloat32;
import org.tensorflow.types.family.TType;
import org.tensorflow.sandbox.util.Pair;

import java.util.ArrayList;
import java.util.HashMap;
Expand All @@ -52,7 +50,7 @@ public abstract class Optimizer {
* Global state variables
*/
//TODO make this be used.
protected final List<Variable> globals;
protected final List<Variable<?>> globals;

/**
* The Graph this optimizer is operating on.
Expand All @@ -76,12 +74,12 @@ public Op minimize(Operand<?> loss) {
}

public Op minimize(Operand<?> loss, String name) {
List<Pair<Output<?>, Output<? extends TType>>> gradsAndVars = computeGradients(loss);
List<GradAndVar<?>> gradsAndVars = computeGradients(loss);

return applyGradients(gradsAndVars, name);
}

public List<Pair<Output<?>, Output<? extends TType>>> computeGradients(Operand<?> loss) {
public <T extends TType> List<GradAndVar<?>> computeGradients(Operand<?> loss) {
List<Operation> variables = new ArrayList<>();
Iterator<Operation> opItr = graph.operations();
while (opItr.hasNext()) {
Expand All @@ -98,26 +96,30 @@ public List<Pair<Output<?>, Output<? extends TType>>> computeGradients(Operand<?
}

Output<?>[] gradients = graph.addGradients(loss.asOutput(), variableOutputArray);
List<Pair<Output<?>, Output<? extends TType>>> gradVarPairs = new ArrayList<>();
List<GradAndVar<? extends TType>> gradVarPairs = new ArrayList<>();

for (int i = 0; i < variableOutputArray.length; i++) {
gradVarPairs.add(new Pair<>(gradients[i], variableOutputArray[i]));
@SuppressWarnings("unchecked")
Output<T> typedGrad = (Output<T>) gradients[i];
@SuppressWarnings("unchecked")
Output<T> typedVar = (Output<T>) variableOutputArray[i];
gradVarPairs.add(new GradAndVar<>(typedGrad, typedVar));
}

return gradVarPairs;
}

public Op applyGradients(List<Pair<Output<?>, Output<? extends TType>>> gradsAndVars, String name) {
List<Output<? extends TType>> variables = gradsAndVars.stream().map(Pair::getB).collect(Collectors.toList());
public Op applyGradients(List<GradAndVar<? extends TType>> gradsAndVars, String name) {
List<Output<? extends TType>> variables = gradsAndVars.stream().map(GradAndVar::getVariable).collect(Collectors.toList());

createSlots(variables);

Optional<Operand> prepOp = prepare(name+"/prepare");
Optional<Operand<? extends TType>> prepOp = prepare(name+"/prepare");

List<Operand<?>> updateOps = new ArrayList<>();
List<Operand<? extends TType>> updateOps = new ArrayList<>();
prepOp.ifPresent(updateOps::add);
for (Pair pair : gradsAndVars) {
updateOps.add(applyDense((Output)pair.getA(),(Output)pair.getB()));
for (GradAndVar<? extends TType> pair : gradsAndVars) {
updateOps.add(applyDense(pair));
}

return finish(updateOps,name);
Expand All @@ -129,7 +131,7 @@ public Op applyGradients(List<Pair<Output<?>, Output<? extends TType>>> gradsAnd
* @param slotName The slot name.
* @return The slot or {@link Optional#empty}.
*/
public Optional<Variable<?>> getSlot(Output<?> var, String slotName) {
public <T extends TType> Optional<Variable<T>> getSlot(Output<T> var, String slotName) {
return getSlot(var.op().name(),slotName);
}

Expand All @@ -139,12 +141,14 @@ public Optional<Variable<?>> getSlot(Output<?> var, String slotName) {
* @param slotName The slot name.
* @return The slot or {@link Optional#empty}.
*/
public Optional<Variable<?>> getSlot(String varName, String slotName) {
Map<String,Variable<?>> variables = slots.get(slotName);
private <T extends TType> Optional<Variable<T>> getSlot(String varName, String slotName) {
Map<String,Variable<? extends TType>> variables = slots.get(slotName);
if (variables != null) {
Variable<?> slot = variables.get(varName);
Variable<? extends TType> slot = variables.get(varName);
if (slot != null) {
return Optional.of(slot);
@SuppressWarnings("unchecked") // This method should only be called when the type is known.
Optional<Variable<T>> opt = Optional.of((Variable<T>)slot);
return opt;
} else {
return Optional.empty();
}
Expand All @@ -162,11 +166,11 @@ public Optional<Variable<?>> getSlot(String varName, String slotName) {
* @param <T> The type of the variable.
*/
protected <T extends TType> void createSlot(Output<T> variable, String slotName, Operand<T> initializer) {
Variable<T> slot = (Variable<T>) tf.withName(createName(variable, slotName)).variable(variable.shape(), TFloat32.DTYPE);
Variable<T> slot = tf.withName(createName(variable, slotName)).variable(variable.shape(), variable.dataType());
Assign<T> slotInit = tf.assign(slot, initializer);
graph.addInitializer(slotInit);
String varName = variable.op().name();
Map<String,Variable<?>> variables = slots.computeIfAbsent(slotName,(k) -> new HashMap<>());
Map<String,Variable<? extends TType>> variables = slots.computeIfAbsent(slotName,(k) -> new HashMap<>());
variables.put(varName,slot);
}

Expand All @@ -175,7 +179,7 @@ protected <T extends TType> void createSlot(Output<T> variable, String slotName,
*
* @param scopeName The scope name to use for any variable creations.
*/
protected Optional<Operand> prepare(String scopeName) {
protected Optional<Operand<? extends TType>> prepare(String scopeName) {
return Optional.empty();
}

Expand All @@ -185,6 +189,10 @@ protected Optional<Operand> prepare(String scopeName) {
*/
protected void createSlots(List<Output<? extends TType>> variables) { }

private <T extends TType> Operand<T> applyDense(GradAndVar<T> gradVarPair) {
return applyDense(gradVarPair.getGradient(),gradVarPair.getVariable());
}

/**
* Generates the gradient update operations for the specific variable and gradient.
* @param gradient The gradient to use.
Expand Down Expand Up @@ -213,7 +221,25 @@ protected Op finish(List<Operand<?>> updateOperations, String name) {
*/
public abstract String getOptimizerName();

public static String createName(Output<?> variable, String slotName) {
public static String createName(Output<? extends TType> variable, String slotName) {
return variable.op().name() + "-" + slotName;
}

public static class GradAndVar<T extends TType> {
private final Output<T> gradient;
private final Output<T> variable;

public GradAndVar(Output<T> gradient, Output<T> variable) {
this.gradient = gradient;
this.variable = variable;
}

public Output<T> getGradient() {
return gradient;
}

public Output<T> getVariable() {
return variable;
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,6 @@
import org.tensorflow.Graph;
import org.tensorflow.Operand;
import org.tensorflow.Output;
import org.tensorflow.op.core.Constant;
import org.tensorflow.op.core.Variable;
import org.tensorflow.types.TFloat32;
import org.tensorflow.types.family.TType;
Expand Down Expand Up @@ -64,25 +63,22 @@ protected void createSlots(List<Output<? extends TType>> variables) {
}

private <T extends TType> void createRMSPropSlot(Output<T> v) {
Operand<T> rmsInitializer = tf.fill(tf.shape(v), (Constant<T>) tf.constant(1.0f, TFloat32.DTYPE));//v.dataType()));
Operand<T> rmsInitializer = tf.fill(tf.shape(v), tf.dtypes.cast(tf.constant(1.0f, TFloat32.DTYPE),v.dataType()));
createSlot(v.asOutput(), RMS, rmsInitializer);
Operand<T> momentumInitializer = tf.fill(tf.shape(v), (Constant<T>) tf.constant(0.0f, TFloat32.DTYPE));//v.dataType()));
Operand<T> momentumInitializer = tf.fill(tf.shape(v), tf.dtypes.cast(tf.constant(0.0f, TFloat32.DTYPE),v.dataType()));
createSlot(v.asOutput(), MOMENTUM, momentumInitializer);
if (centered) {
Operand<T> mgInitializer = tf.fill(tf.shape(v), (Constant<T>) tf.constant(0.0f, TFloat32.DTYPE));//v.dataType()));
Operand<T> mgInitializer = tf.fill(tf.shape(v), tf.dtypes.cast(tf.constant(0.0f, TFloat32.DTYPE),v.dataType()));
createSlot(v.asOutput(), MG, mgInitializer);
}
}

@Override
protected <T extends TType> Operand<T> applyDense(Output<T> gradient, Output<T> variable) {
@SuppressWarnings("unchecked") // suppressed as the slots are created to have the dtype of the variable.
Variable<T> rmsSlot = (Variable<T>) getSlot(variable,RMS).get();
@SuppressWarnings("unchecked") // suppressed as the slots are created to have the dtype of the variable.
Variable<T> momentumSlot = (Variable<T>) getSlot(variable,MOMENTUM).get();
Variable<T> rmsSlot = getSlot(variable,RMS).get();
Variable<T> momentumSlot = getSlot(variable,MOMENTUM).get();
if (centered) {
@SuppressWarnings("unchecked") // suppressed as the slots are created to have the dtype of the variable.
Variable<T> mgSlot = (Variable<T>) getSlot(variable, MG).get();
Variable<T> mgSlot = getSlot(variable, MG).get();
return tf.train.applyCenteredRmsProp(variable, mgSlot, rmsSlot, momentumSlot,
tf.constant(learningRate, gradient.dataType()),
tf.constant(decay, gradient.dataType()),
Expand Down
Loading