tensorflow / tensorflow/java

How to update the weight about BlockLSTM ?

オープン
#450 コメント 1 件 リアクション 0 件 担当者 0 名 GitHub で見る

まだ誰も着手していません。

主要言語
Java
スター
928
フォーク
227
PR マージ指標
30日以内にマージされた PR はありません

説明

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

コントリビューションガイド

コントリビューションガイドを開く

はじめの一歩

  1. issue を最後まで読み、次にプロジェクトのコントリビューションガイドを読みます。
  2. 着手することを issue にコメントします — 二人が同じ作業をするのを防げます。
  3. リポジトリをフォークし、ブランチを切って変更します。
  4. issue 番号を参照したプルリクエストを送ります。

調査の方向性

提供された Scala の例と BlockLSTM.create 呼び出しから始め、次に Java API で BlockLSTMGrad またはその他の勾配演算を調査します。TensorFlow Java 環境で例を実行し、公開されている API が表示された重みテンソルの更新をサポートしているかどうかを確認します。サポートされている重み更新のワークフローが文書化されるか、API の制限が明確に記載されれば完了です。

索引モデルが issue の本文から書いたものです。

評価

技術スタック
java, scala, tensorflow
領域
machine-learning
issue の種類
ドキュメント
難易度
5/5
見積もり時間
1週間以上
活発さ
停滞
明瞭さ
説明が足りない
初心者へのやさしさ
20/100

新しい issue をメールで受け取る

初心者向けの GitHub issue を短くまとめたダイジェスト。