如何通过SageMaker Java客户端生成RecordIO?训练报错排查
嘿,我来帮你搞定这个问题!首先明确说:你完全可以用Amazon SageMaker Java客户端生成符合要求的RecordIO封装Protobuf训练数据,而且这应该能解决你现在遇到的「No training data processed」报错——毕竟你之前用Protostuff生成的Protobuf格式大概率不符合Linear Learner的预期规范。下面我给你详细讲怎么做,顺便帮你排查之前的问题根源。
为什么你之前的Protostuff生成数据会报错?
先给你理清楚:SageMaker的Linear Learner算法要求的RecordIO-Protobuf不是随便的POJO序列化结果,它必须遵循官方定义的FeatureVector结构(包含标签和特征的特定Protobuf schema)。你用Protostuff从自定义POJO生成的二进制,算法根本没法解析,所以哪怕文件非空,也会被判定为没有有效训练数据。而LibSVM格式是算法明确支持的,所以能正常跑起来。
用SageMaker Java客户端生成正确格式的步骤
1. 先引入必要的依赖
如果是Maven项目,把这些依赖加到你的pom.xml里(用最新稳定版就行):
<dependency> <groupId>software.amazon.awssdk</groupId> <artifactId>sagemaker</artifactId> <version>2.20.0</version> </dependency> <dependency> <groupId>software.amazon.awssdk</groupId> <artifactId>sagemaker-runtime</artifactId> <version>2.20.0</version> </dependency> <dependency> <groupId>software.amazon.awssdk</groupId> <artifactId>s3</artifactId> <version>2.20.0</version> </dependency>
2. 构建符合要求的FeatureVector
Linear Learner只认官方的FeatureVector结构,你可以用SDK里的类来构建单条训练记录:
import software.amazon.awssdk.services.sagemaker.model.FeatureVector; import software.amazon.awssdk.services.sagemaker.model.Feature; import java.util.ArrayList; import java.util.List; // 把你的Spark DataFrame里的一行数据转换成FeatureVector FeatureVector createTrainingRecord(double label, List<Double> featureValues) { List<Feature> featuresList = new ArrayList<>(); // 遍历特征值,给每个特征指定名称(这里用索引当名称,也可以自定义) for (int idx = 0; idx < featureValues.size(); idx++) { featuresList.add(Feature.builder() .featureName(String.valueOf(idx)) .value(featureValues.get(idx)) .build()); } return FeatureVector.builder() .label(label) .features(featuresList) .build(); }
3. 序列化并封装成RecordIO文件
接下来用SDK提供的RecordWriter把FeatureVector序列化成RecordIO格式,然后上传到S3(因为SageMaker训练任务需要从S3读取数据):
import software.amazon.awssdk.core.sync.RequestBody; import software.amazon.awssdk.services.s3.S3Client; import software.amazon.awssdk.services.s3.model.PutObjectRequest; import software.amazon.awssdk.services.sagemaker.util.RecordWriter; import java.io.FileOutputStream; import java.io.IOException; import java.nio.file.Path; import java.nio.file.Paths; import java.util.List; void generateRecordIoTrainingData(List<FeatureVector> allTrainingRecords, String localOutputPath, String s3Bucket, String s3Key) throws IOException { // 先写入本地文件 try (FileOutputStream fos = new FileOutputStream(localOutputPath); RecordWriter recordWriter = RecordWriter.create(fos)) { for (FeatureVector record : allTrainingRecords) { recordWriter.write(record); } } // 上传到S3 S3Client s3Client = S3Client.create(); Path filePath = Paths.get(localOutputPath); s3Client.putObject(PutObjectRequest.builder() .bucket(s3Bucket) .key(s3Key) .build(), RequestBody.fromFile(filePath)); }
4. 配置训练任务使用这个数据
最后在创建Linear Learner训练任务时,指定训练通道的内容类型为application/x-recordio-protobuf:
import software.amazon.awssdk.services.sagemaker.model.Channel; import software.amazon.awssdk.services.sagemaker.model.DataSource; import software.amazon.awssdk.services.sagemaker.model.S3DataSource; import software.amazon.awssdk.services.sagemaker.model.InputMode; // 构建训练通道 Channel trainingChannel = Channel.builder() .channelName("train") .dataSource(DataSource.builder() .s3DataSource(S3DataSource.builder() .s3Uri(String.format("s3://%s/%s", s3Bucket, s3Key)) .s3DataType("S3Prefix") .build()) .build()) .contentType("application/x-recordio-protobuf") .inputMode(InputMode.FILE) .build();
额外小建议
- 如果你的数据是从Spark DataFrame来的,其实可以用Spark的SageMaker连接器直接生成RecordIO-Protobuf格式,这样不用手动转POJO,效率更高。
- 生成RecordIO文件后,你可以用本地的SageMaker工具简单验证下文件是否能被解析,比如用
aws sagemaker-runtime的本地测试功能,避免上传到S3后才发现问题。
内容的提问来源于stack exchange,提问作者ljl97114

