Custom BaseActivationFunction gives Array workspace validation error
Any function that extends BaseActivationFunction as in the docs (https://deeplearning4j.konduit.ai/nd4j/activations) will throw various errors when used. When used in ActivationLayer:
Exception in thread "main" java.lang.IllegalStateException: Feed forward (inference): array (ACTIVATIONS) workspace validation failed (vertex binary_step1 - class: ActivationLayer) - array is defined in incorrect workspace [...] Caused by: org.nd4j.linalg.workspace.ND4JWorkspaceException: Array workspace validation failed: Array of type ACTIVATIONS should be in workspace "WS_LAYER_ACT_3" but is in workspace "WS_LAYER_WORKING_MEM"
When used as activation in a standard layer:
Caused by: org.nd4j.linalg.workspace.ND4JWorkspaceException: Array workspace validation failed: Array of type ACTIVATIONS should be detached (no workspace) but is in workspace: WS_LAYER_ACT_2
Example:
BinaryStepActivation.java
package dl;
import org.nd4j.linalg.activations.BaseActivationFunction;
import org.nd4j.linalg.api.ndarray.INDArray;
import org.nd4j.common.primitives.Pair;
import org.nd4j.linalg.factory.Nd4j;
/**
* 0 if x<0; 1 if x>0
*/
public class BinaryStepActivation extends BaseActivationFunction {
@Override
public INDArray getActivation(INDArray in, boolean training) {
INDArray bstep = Nd4j.ones(in.shape()).mul(in.gt(0));
// must: return workspaceMgr.leverageTo(ArrayType.ACTIVATIONS, ret);
return bstep;
}
@Override
public Pair<INDArray, INDArray> backprop(INDArray in, INDArray epsilon) {
// Compute activation gradient and multiply by upstream gradient
// Return: Pair(gradient, null)
// The second element is null for activations without learnable parameters
INDArray gradient = Nd4j.zeros(in.shape()); // your derivative
// derivative is 0 https://en.wikipedia.org/wiki/Activation_function
// gradient.muli(epsilon);
return new Pair<>(gradient, null);
}
}Main.java
int n_batch = 4;
int nIn = 8;
int nOut = 8;
MultiLayerConfiguration conf = new NeuralNetConfiguration.Builder()
.seed(123)
.dataType(DataType.FLOAT)
.updater(new Adam())
.gradientNormalization(GradientNormalization.RenormalizeL2PerLayer)
.gradientNormalizationThreshold(1.0)
.list()
.layer(new DenseLayer.Builder()
.nIn(nIn)
.nOut(nOut)
.activation(Activation.SIGMOID) // supposedly pairs well with mse
.weightInit(WeightInit.XAVIER_UNIFORM)
.build()
)
.layer(new DenseLayer.Builder()
.nIn(nIn)
.nOut(nOut)
.activation(new BinaryStepActivation())
.weightInit(WeightInit.XAVIER_UNIFORM)
.build()
).build();
MultiLayerNetwork model = new MultiLayerNetwork(conf);
model.init();
// test inference
INDArray input = Nd4j.rand(n_batch, nIn);
INDArray output = model.output(input);Source: deeplearning4j/deeplearning4j