How to update the weight about BlockLSTM ?
還沒有人認領這個 Issue。
評估
- 難度
- 5/5
- 預估耗時
- 一週以上
- 新手友好度
- 20/100
- Issue 類型
- 文件
- 描述清晰度
- 需要釐清
- 活躍度
- 停滯
研究方向
從提供的 Scala 範例和 BlockLSTM.create 呼叫開始,接著檢查 Java API 是否有 BlockLSTMGrad 或其他梯度運算。在 TensorFlow Java 環境中執行範例,並確認公開的 API 是否支援更新所顯示的權重張量。支援的權重更新工作流程已完成文件記錄,或 API 限制已清楚說明,即表示完成。
由索引模型根據 Issue 內容生成。
描述
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
- 主要語言
- Java
- 星號
- 928
- 分支
- 227
- PR 合併指標
- 30 天內沒有已合併 PR
貢獻指南
從這裡開始
- 先讀完整個 Issue,再讀專案的貢獻指南。
- 在 Issue 下留言說明你要接手 —— 這能避免兩個人做同樣的事。
- Fork 儲存庫,在一個分支上完成修改。
- 送出 Pull Request,並在描述裡引用這個 Issue 編號。
tensorflow/java 的其他 Issue
-
難度 2/5 1-3 小時 新手友好度 65/100
tensorflow/java#653 · 1 則留言 · 4 個 reaction ·
-
難度 5/5 一週以上 新手友好度 25/100
tensorflow/java#621 · 4 則留言 ·
-
難度 5/5 一週以上 新手友好度 25/100
tensorflow/java#617 · 3 則留言 ·
-
難度 2/5 1-3 小時 新手友好度 55/100
tensorflow/java#615 · 1 則留言 ·
-
難度 5/5 一週以上 新手友好度 25/100
tensorflow/java#614 · 1 則留言 ·
相似的 Issue
-
bug
難度 1/5 1 小時以內 新手友好度 90/100
apache/cloudstack#14222 ·
-
難度 2/5 1-3 小時 新手友好度 88/100
-
1.0.0-alpha2 Type/Improvement
難度 2/5 1-3 小時 新手友好度 68/100
wso2/dpdp-accelerator#272 ·
-
難度 2/5 1-3 小時 新手友好度 82/100
infinispan/infinispan#18150 ·
-
area/frontend
難度 2/5 1-3 小時 新手友好度 65/100