Create ML生成的.mlmodel出现额外stateIn输入的疑问求解
关于Create ML生成LSTM模型中
stateIn参数的说明 一、stateIn的由来
Create ML在构建LSTM模型时,会自动暴露LSTM层的状态输入接口。LSTM作为循环神经网络,核心是靠维护内部状态来处理序列数据——每一步的输出不仅依赖当前输入,还和上一步的隐藏状态有关。训练完成后,模型会把这个状态输入保留下来,方便你在预测时灵活控制状态的传递,比如处理连续长序列、重置状态等。
二、预测时该传什么值
分两种常用场景:
- 首次预测或需要重置状态时:传全零张量就行,维度要和模型训练时的LSTM状态维度匹配。比如单层LSTM的状态一般是两个张量(隐藏状态h和细胞状态c),形状通常是
[2, 1, 隐藏单元数](批量大小为1的情况)。用MLMultiArray创建全零数组,数据类型和模型要求一致(大多是float32)。 - 连续序列预测时:直接传入上一次预测输出的
stateOut值。Create ML生成的LSTM模型会同时输出预测结果和stateOut(状态输出),把这个输出作为下一次的stateIn,模型就能延续之前的状态,处理连贯的时序数据。
三、Xcode里的实操示例
自动生成的ReloadInput类中,stateIn是MLMultiArray类型,初始化全零状态可以这么写:
// 假设模型LSTM隐藏单元数为64,批量大小设为1 do { let stateIn = try MLMultiArray(shape: [2, 1, 64], dataType: .float32) // 填充全零值 for index in 0..<stateIn.count { stateIn[index] = 0.0 } // 把stateIn传入模型输入对象 let input = ReloadInput(accX: ..., accY: ..., accZ: ..., gyroX: ..., gyroY: ..., gyroZ: ..., stateIn: stateIn) } catch { // 处理初始化错误 }
如果是连续预测,直接把上一次模型输出的modelOutput.stateOut赋值给下一次输入的stateIn就好。
内容的提问来源于stack exchange,提问作者Qazi Ammar
相关产品推荐
相关产品推荐

