如何在Deeplearning4j UCISequenceClassification LSTM示例中添加多特征
Deeplearning4j UCISequenceClassification多特征适配方案
你核心问题是用错了SequenceRecordReader的类型,CSVMultiSequenceRecordReader用于同一样本对应多个独立序列文件的场景,并不适配你的单文件多列特征需求,按以下步骤调整即可:
1. 确认输入文件格式正确性
你调整的文本替换逻辑是对的,最终生成的每个样本csv文件需要满足:
- 每行对应一个时间步
- 每列对应一个特征,列之间用你指定的分隔符(你用的是
|)分隔
比如2个特征、3个时间步的样本文件内容如下:
0.23|1 0.56|1 0.71|1
2. 替换为正确的RecordReader
不需要更换为CSVMultiSequenceRecordReader,直接使用原示例的CSVSequenceRecordReader,构造时传入分隔符参数即可:
// 第一个参数是跳过的表头行数,无表头填0;第二个参数是你使用的列分隔符 SequenceRecordReader trainFeatures = new CSVSequenceRecordReader(0, "|"); trainFeatures.initialize(new NumberedFileInputSplit(featuresDirTrain.getAbsolutePath() + "/%d.csv", 0, 449));
CSVSequenceRecordReader原生支持单文件内多列特征的读取,会自动将每行的多列映射为多个特征。
3. 同步调整网络输入维度
原示例网络第一层的输入维度nIn=1对应单特征,现在有N个特征就把nIn修改为N:
// 2个特征的示例,多特征直接改成对应数量 new LSTM.Builder().nIn(2).nOut(100).activation(Activation.TANH).build()
多特征扩展方式
如果需要新增更多特征,只需要:
- 每个时间步的行内新增对应特征列,用分隔符隔开
- 将
nIn参数修改为特征总数量即可,无需调整其他读取逻辑
内容的提问来源于stack exchange,提问作者Corius
相关产品推荐
相关产品推荐

