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

如何让XGBoost外部内存与AFT生存分析模型协同工作?

解决XGBoost批量迭代器适配AFT模型的删失区间设置问题

要在自定义的BatchedParquetIterator中为AFT模型传递删失区间信息,需要利用XGBoost input_data函数的set_info参数,把label_lower_bound和label_upper_bound随批次数据一起传入。具体修改步骤如下:

1. 更新预处理方法,返回删失区间数据

修改_preprocess方法,从DataFrame中提取出删失时间的上下界,和特征、标签一起返回:

def _preprocess(self, df: pd.DataFrame) -> Tuple[pd.DataFrame, pd.DataFrame, pd.Series, pd.Series]:
    # 原有预处理逻辑:提取特征X和标签y
    # ...
    # 从df中获取删失区间数据,替换为你实际的列名
    y_lower = df['你的下界列名']
    y_upper = df['你的上界列名']
    return X, y, y_lower, y_upper

2. 在迭代器的next方法中传递删失区间信息

在next方法调用input_data时,通过set_info参数传入这两个区间的数组:

def next(self, input_data: Callable):
    if self._it == len(self._file_paths):
      return 0
    
    df = pd.read_parquet(self._file_paths[self._it])
    # 修改接收变量,获取删失区间
    X, y, y_lower, y_upper = self._preprocess(df)
    
    # 传递数据、标签以及删失区间信息
    input_data(
        data=X, 
        label=y,
        set_info={
            'label_lower_bound': y_lower.values,
            'label_upper_bound': y_upper.values
        }
    )
    self._it += 1

    return 1

关键说明

  • set_info参数接受字典格式,键为XGBoost支持的信息名称(固定为label_lower_bound和label_upper_bound),值为对应批次的numpy数组或pandas序列。
  • 必须保证每个批次的y_lower、y_upper长度和当前批次的X、y完全匹配,避免数据错位。

修改后生成的Xy_train(DMatrix对象)就包含了AFT模型所需的删失区间信息,可直接用于训练:

parquet_iterator_train = BatchedParquetIterator(batches)
Xy_train = xgboost.DMatrix(parquet_iterator_train)

# 训练AFT模型示例
model = xgboost.train(
    params={'objective': 'survival:aft', 'eval_metric': 'aft-nloglik'},
    dtrain=Xy_train,
    num_boost_round=100
)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 20:11:25