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

如何在PySpark的PythonModel类中使用@pandas_udf函数?

问题分析与解决方案

核心错误原因

  1. 类型标注不匹配:predict_batch_udf参数标注为pd.Series,但代码中访问data['content']说明输入是DataFrame;同时返回标注为pd.Series,但实际返回两个数组,与定义的返回结构矛盾,触发类型提示错误。
  2. 不必要的pandas_udf使用:MLflow的PythonModel.predict方法本身接收pandas数据(DataFrame/Series),无需用@pandas_udf装饰——这个装饰器是给Spark DataFrame注册UDF用的,放在此处完全不适用。
  3. 方法调用名错误:predict方法里调用self._predict_batch,但类中定义的是predict_batch_udf,命名不匹配会导致运行时错误。

修正后的代码

class RobertaClassifier(PythonModel):

    def load_context(self, context: PythonModelContext):
        import os
        import torch
        from transformers.models.auto import AutoConfig, AutoModelForSequenceClassification
        from transformers.models.auto.tokenization_auto import AutoTokenizer
        
        config_file = os.path.dirname(context.artifacts["config"])
        self.config = AutoConfig.from_pretrained(config_file)
        self.tokenizer = AutoTokenizer.from_pretrained(config_file)
        self.model = AutoModelForSequenceClassification.from_pretrained(config_file, config=self.config)
        
        if torch.cuda.is_available():
            print('[INFO] Model is being sent to CUDA device as GPU is available')
            self.model = self.model.cuda()
        else:
            print('[INFO] Model will use CPU runtime')
        
        self.model.eval()
    
    def _predict_batch(self, data: pd.DataFrame) -> tuple[np.ndarray, np.ndarray]:
        import torch
        import pandas as pd
        
        with torch.no_grad():
            # 假设preprocessing是已定义的文本预处理函数
            inputs_text = preprocessing(data['content'].tolist())
            inputs = self.tokenizer(inputs_text, padding=True, return_tensors='pt', max_length=512, truncation=True)
        
            if self.model.device.type == 'cuda':
                torch.cuda.empty_cache()
                inputs = {k: v.to(self.model.device) for k, v in inputs.items()}

            predictions = self.model(**inputs)
            probs = torch.nn.Softmax(dim=1)(predictions.logits)
            probs = probs.detach().cpu().numpy()

            labels = probs.argmax(axis=1)
            scores = probs.max(axis=1)

            return labels, scores
        
    def predict(self, context: PythonModelContext, data: pd.DataFrame) -> pd.DataFrame:
        import math
        import numpy as np
        
        batch_size = 64
        sample_size = len(data)
        
        labels = np.zeros(sample_size, dtype=int)
        scores = np.zeros(sample_size, dtype=float)

        for batch_idx in range(math.ceil(sample_size / batch_size)):
            bfrom = batch_idx * batch_size
            bto = min(bfrom + batch_size, sample_size)
            
            batch_data = data.iloc[bfrom:bto]
            l, s = self._predict_batch(batch_data)
            labels[bfrom:bto] = l
            scores[bfrom:bto] = s
            
        # 将数字标签映射为文本标签
        text_labels = [self.config.id2label[label] for label in labels]
        return pd.DataFrame({'label': text_labels, 'score': scores})

关键修改点

  • 移除@pandas_udf装饰器,将predict_batch_udf改为普通方法_predict_batch,明确输入为pd.DataFrame,返回为元组(标签数组、分数数组)。
  • 修正predict方法中的调用名,改为self._predict_batch。
  • 调整输入处理逻辑:将data['content']转为列表传入预处理,避免Series操作的潜在问题。
  • 优化设备判断逻辑:用self.model.device.type == 'cuda'替代原有的索引判断,更准确。
  • 修正数组初始化的 dtype,确保与返回值类型匹配。
  • 循环中用min(bfrom + batch_size, sample_size)避免越界。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 11:15:43