tensorflow / tensorflow/java

How to update the weight about BlockLSTM ?

Ouverte
#450 1 commentaire 0 réactions 0 personnes assignées Voir sur GitHub

Personne n'a encore pris cette issue.

Langage dominant
Java
Étoiles
928
Forks
227
Métriques de merge des PR
Aucune PR mergée en 30 j

Description

HI :

now I use BlockLSTM for build lstm layer ,but I don't know how to update the lstm weight parameters , if need use blockLSTMGrad or something to do ,the coda is paste here:

object LstmExample {

  def initializeTruncatedNormalTensor(shape: Operand[TInt32], scope: Scope): TFloat32 = {
    TruncatedNormal.seed(1000L)
    //        TruncatedNormal<TFloat32> truncatedNormal = TruncatedNormal.create(scope, shape, TFloat32.DTYPE);
    //        DataType<TFloat32> DTYPE = DataType.create("FLOAT", 1, 4, TFloat32Impl::mapTensor);
    //        DataType DTYPE = DataType.valueOf("FLOAT");
    val truncatedNormal: TruncatedNormal[TFloat32] = TruncatedNormal.create(scope, shape, classOf[TFloat32])
    return truncatedNormal.asTensor
  }
  private def getWeightMatrix(shape: Operand[TInt32], scope: Scope) = { //        Tensor<TFloat32> tensorWeight = TensorValues.initializeTruncatedNormalTensor(shape, scope);
    val tensorWeight = TensorValues.initializeTruncatedNormalTensor(shape, scope)
    Constant.create(scope, tensorWeight)
  }

  def printTensor(tensor: Operand[TFloat32],name:String): Unit ={
    val data: Array[Float] = TensorResources.extractFloats(tensor.asTensor())
    println(s"data:${name},  ${data.mkString(" | ")}")
  }
  def main(args: Array[String]): Unit = {
    val libraryPath = System.getProperty("java.library.path")
    System.out.println(libraryPath)
    implicit val session = TestSession.createTestSession(TestSession.Mode.EAGER) // EagerSession.create()
    implicit val tf = session.getTF // Ops.create(session).withName("test")
    implicit val scope = tf.scope()
    //    val session = EagerSession.create
    //    val tf = Ops.create(session)
    //        Scope scope = new Scope(session);
    //    val scope = session.baseScope()
    val rawInputSequence = Array(Array(Array(0.1f, 0.2f)), Array(Array(0.3f, 0.4f))) //shape (timelen, batch_size, num_inputs).
    val inputSequence = tf.constant(rawInputSequence)
    val inputSize = 2
    val cellSize =  5
    val maximumTimeLength = 2
    val cellShape = Array(1, cellSize)
    val cellDims = Constant.vectorOf(scope, cellShape)
    val seqLenMax = tf.array(maximumTimeLength)
    //        Operand<TFloat32> initialCellState = Zeros.create(scope, cellDims, TFloat32.DTYPE);
    //        Operand<TFloat32> initialHiddenState = Zeros.create(scope, cellDims, TFloat32.DTYPE);
    val initialCellState = Zeros.create(scope, cellDims, classOf[TFloat32])
    val initialHiddenState = Zeros.create(scope, cellDims, classOf[TFloat32])
    val weightShape = Array(inputSize + cellSize, cellSize * 4)
    val weightMatrixDims = Constant.vectorOf(scope, weightShape)
    val weightMatrix = getWeightMatrix(weightMatrixDims, scope)
    //    session.print(weightMatrix)
    val weightGatesShape = Array(cellSize)
    val weightGatesDims = Constant.vectorOf(scope, weightGatesShape)
    val weightInputGate = getWeightMatrix(weightGatesDims, scope)
    printTensor(weightInputGate,"weightInputGate")
    val weightForgetGate = getWeightMatrix(weightGatesDims, scope)
    printTensor(weightForgetGate,"weightForgetGate")
    val weightOutputGate = getWeightMatrix(weightGatesDims, scope)
    printTensor(weightOutputGate ,"weightOutputGate")
    val biasShape = Array(cellSize * 4)
    val biasDim = Constant.vectorOf(scope, biasShape)
    //classOf[TFloat32]
    //        Operand<TFloat32> bias = Zeros.create(scope, biasDim, TFloat32.DTYPE);
    val bias = Zeros.create(scope, biasDim, classOf[TFloat32])
    val blockLSTM = BlockLSTM.create(scope, tf.dtypes.cast(seqLenMax, classOf[TInt64]), inputSequence, initialCellState, initialHiddenState, weightMatrix, weightInputGate, weightForgetGate, weightOutputGate, bias)

//    session.print(blockLSTM.i)
//    println("&&&cs")
//    session.print(blockLSTM.cs)
//    println("&&&f")
//    session.print(blockLSTM.f)
//    println("&&&o")
//    session.print(blockLSTM.o)
//    println("&&&ci")
//    session.print(blockLSTM.ci)
//    println("&&&co")
//    session.print(blockLSTM.co)
//    println("&&&h")
//    session.print(blockLSTM.h)

thanks for your help

Guide de contribution

Ouvrir le guide de contribution

Par où commencer

  1. Lisez l'issue en entier, puis le guide de contribution du projet.
  2. Signalez en commentaire que vous la prenez — cela évite que deux personnes fassent le même travail.
  3. Forkez le dépôt et travaillez sur une branche.
  4. Ouvrez une pull request qui référence le numéro de l'issue.

Piste de recherche

Commencez par l'exemple Scala fourni et l'appel à BlockLSTM.create, puis examinez l'API Java à la recherche de BlockLSTMGrad ou d'autres opérations de gradient. Exécutez l'exemple dans l'environnement Java de TensorFlow et déterminez si les API exposées permettent de mettre à jour les tenseurs de poids affichés. Le travail est terminé lorsque le workflow pris en charge pour la mise à jour des poids est documenté ou que la limitation de l'API est clairement indiquée.

Rédigé par le modèle d'indexation à partir du texte de l'issue.

Évaluation

Stack technique
java, scala, tensorflow
Domaine
machine-learning
Type d'issue
Documentation
Difficulté
5/5
Temps estimé
Plus d'une semaine
Activité
À l'abandon
Clarté
À clarifier
Accessibilité débutants
20/100

Recevez les nouvelles issues par e-mail

Un résumé court des issues GitHub adaptées aux débutants.