DL4J如何加载嵌套数组(数组的数组)结构数据生成DataSet
DL4J加载3D嵌套数组生成DataSet的实现方案
你可以通过调整存储格式+维度重塑的方式快速实现需求,不需要复杂的自定义开发,具体操作如下:
第一步:调整源文件存储格式
推荐将单个样本的嵌套3D特征按固定顺序扁平化存储为单行,后续加载后再还原维度,兼容性最好。
以你给出的样本{{1, 1, 1}, {2, 2, 2}, {3,4,5}, {6, 6, 6}}为例,该样本特征对应shape为[4, 3]的二维数组(4个内部子数组,每个长度为3),你可以按行优先顺序将其展开为12个连续数值,末尾追加对应标签值,单个样本占一行,用逗号分隔即可,示例行格式:
1,1,1,2,2,2,3,4,5,6,6,6,0
其中前12位为展开的特征值,最后1位为标签值。
第二步:加载并重塑维度生成DataSet
读取逻辑和你原有一维数组的读取逻辑基本一致,仅需要在拿到原始DataSet后对特征做维度重塑即可,完整代码示例:
int numLinesToSkip = 0; char delimiter = ','; String filePath = "data.txt"; File file = new File(filePath); // 数据集参数配置 int batchSize = (int) lineCount; // 按你的数据集总大小调整 int featureHeight = 4; // 嵌套数组的子数组数量 int featureWidth = 3; // 每个子数组的元素长度 int featureFlattenLength = featureHeight * featureWidth; int labelIndex = featureFlattenLength; // 标签所在的列索引,紧随特征之后 int numClasses = 2; // 分类数量,可按你的需求调整 // CSV读取逻辑和原有逻辑一致 RecordReader recordReader = new CSVRecordReader(numLinesToSkip, delimiter); recordReader.initialize(new FileSplit(file)); DataSetIterator iterator = new RecordReaderDataSetIterator(recordReader, batchSize, labelIndex, numClasses); DataSet allData = iterator.next(); allData.shuffle(); // 核心逻辑:将扁平化的特征还原为3D结构 // reshape参数说明:-1代表自动匹配样本数量,后两位对应你定义的嵌套数组维度 INDArray reshapedFeatures = allData.getFeatures().reshape(-1, featureHeight, featureWidth); allData.setFeatures(reshapedFeatures); // 至此得到的allData即为包含3D嵌套结构特征的DataSet
可选方案:自定义RecordReader保留嵌套格式
如果你不想扁平化存储,需要在文本中保留嵌套结构的分隔标识(比如用分号分隔子数组,逗号分隔元素,单样本格式为1,1,1;2,2,2;3,4,5;6,6,6,0),可以继承RecordReader重写解析逻辑,按你定义的分隔符逐层解析为double数组后封装为INDArray即可,适合嵌套层级多、结构复杂的场景。
如果你加载的是时序类3D数据,也可以直接使用DL4J自带的
CSVSequenceRecordReader按序列规则读取,不需要手动做维度重塑。
内容的提问来源于stack exchange,提问作者Joehot200
相关产品推荐
相关产品推荐

