You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

SageMaker本地批量转换报错:无法转字符串为浮点数,自定义代码未执行

问题分析

从错误日志和你的代码来看,核心问题是SageMaker的SKLearn容器在批量转换时,没有调用你自定义的文本处理逻辑(TfidfVectorizer),而是使用了默认的输入处理函数。默认的input_fn会尝试把CSV中的内容直接转成float类型,而你的数据是字符串文本,所以触发了ValueError。

为什么自定义代码没被调用?因为你只在训练阶段(train.py的训练逻辑)里用了TfidfVectorizer,但没有定义SageMaker推理容器要求的标准推理函数(model_fn、input_fn等),容器不知道该怎么用你的自定义逻辑处理输入数据。

解决方案

你需要修改train.py,添加SageMaker SKLearn容器要求的推理函数,同时调整Transformer的初始化方式,确保容器能正确加载你的自定义逻辑。

步骤1:完善train.py,添加推理函数

在train.py中,除了训练逻辑,还要添加以下四个关键函数:

  • model_fn(model_dir):加载训练好的模型和TfidfVectorizer
  • input_fn(input_data, content_type):处理输入的文本数据(用TfidfVectorizer转换)
  • predict_fn(input_data, model):用加载的模型做预测
  • output_fn(prediction, accept):格式化预测结果输出

示例代码如下:

import os
import pandas as pd
from sklearn.feature_extraction.text import TfidfVectorizer
from sklearn.ensemble import RandomForestClassifier  # 替换成你实际使用的模型
import joblib

def train():
    # 你的训练逻辑
    train_data = pd.read_csv(os.environ['SM_CHANNEL_TRAIN'])
    # 假设你的CSV包含text列(待处理文本)和label列(标签)
    X_train = train_data['text']
    y_train = train_data['label']
    
    # 初始化并训练TfidfVectorizer
    vectorizer = TfidfVectorizer()
    X_train_vec = vectorizer.fit_transform(X_train)
    
    # 训练模型
    model = RandomForestClassifier()
    model.fit(X_train_vec, y_train)
    
    # 保存模型和vectorizer到SageMaker指定的model_dir
    model_dir = os.environ['SM_MODEL_DIR']
    joblib.dump(model, os.path.join(model_dir, 'model.joblib'))
    joblib.dump(vectorizer, os.path.join(model_dir, 'vectorizer.joblib'))

def model_fn(model_dir):
    # 加载保存的模型和vectorizer
    model = joblib.load(os.path.join(model_dir, 'model.joblib'))
    vectorizer = joblib.load(os.path.join(model_dir, 'vectorizer.joblib'))
    # 返回包含两者的字典,方便后续函数调用
    return {'model': model, 'vectorizer': vectorizer}

def input_fn(input_data, content_type):
    # 处理输入的CSV文本数据
    if content_type == 'text/csv':
        # 按行读取输入文本,过滤空行
        texts = [line.strip() for line in input_data.split('\n') if line.strip()]
        return texts
    else:
        raise ValueError(f"Unsupported content type: {content_type}")

def predict_fn(input_data, model_dict):
    # 用vectorizer转换文本,再执行预测
    vectorizer = model_dict['vectorizer']
    model = model_dict['model']
    input_vec = vectorizer.transform(input_data)
    predictions = model.predict(input_vec)
    return predictions

def output_fn(prediction, accept):
    # 格式化预测结果为CSV格式输出
    if accept == 'text/csv':
        return '\n'.join(map(str, prediction)), accept
    else:
        raise ValueError(f"Unsupported accept type: {accept}")

if __name__ == '__main__':
    train()

步骤2:调整Transformer的初始化代码

你不需要手动指定entry_point,因为sklearn estimator已经关联了你的source_dir和entry_point,直接调用sklearn.transformer()即可,避免路径错误:

# 移除entry_point参数,estimator已关联你的脚本路径
transformer = sklearn.transformer(
    instance_count=1, 
    instance_type='local_gpu', 
    strategy='MultiRecord', 
    assemble_with='Line'
)
transformer.transform(
    "file://test_messages", 
    content_type='text/csv', 
    split_type='Line'
)
print('Waiting for transform job: ' + transformer.latest_transform_job.job_name)
transformer.wait()
为什么这样能解决问题?
  • SageMaker的SKLearn容器在启动推理服务时,会自动查找model_fn、input_fn等标准函数,如果找到,就会用这些自定义函数处理输入、加载模型和输出结果,而不是用默认的逻辑。
  • 我们把TfidfVectorizer和模型一起保存,在model_fn中加载,这样在批量转换时,input_fn接收文本后,会用vectorizer把字符串转换成模型能处理的特征,避免了默认的float转换错误。

内容的提问来源于stack exchange,提问作者Tony

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.06 12:19:06