
本文详解如何将 XGBoost 的外部内存迭代器(External Memory Iterator)与加速失效时间(AFT)生存分析模型结合使用,关键在于通过 input_data() 回调函数直接传入 label_lower_bound 和 label_upper_bound,而非依赖构建完成后的 set_float_info()。
本文详解如何将 xgboost 的外部内存迭代器(external memory iterator)与加速失效时间(aft)生存分析模型结合使用,关键在于通过 `input_data()` 回调函数直接传入 `label_lower_bound` 和 `label_upper_bound`,而非依赖构建完成后的 `set_float_info()`。
XGBoost 的 AFT 模型要求每个样本提供一个生存时间区间(即 [t_low, t_high]),用于处理右删失、左删失或区间删失数据。标准用法中,常在 DMatrix 构建完成后调用 set_float_info() 设置上下界:
dtrain = xgboost.DMatrix(X, label=y)
dtrain.set_float_info('label_lower_bound', y_lower)
dtrain.set_float_info('label_upper_bound', y_upper)
但在外部内存迭代器(如 BatchedParquetIterator)场景下,DMatrix 是由 XGBoost 内部动态构建的,用户无法在构造后手动调用 set_float_info() —— 因为 DMatrix 实例尚未暴露给用户。
幸运的是,XGBoost 的设计已为此预留接口:input_data 回调函数签名与 DMatrix 构造器完全一致。这意味着你可以在 next() 方法中,直接将区间信息作为关键字参数传入:
class BatchedParquetIterator(xgboost.DataIter):
def __init__(self, file_paths: List[str], preprocess_fn=None):
self._file_paths = file_paths
self._it = 0
self._preprocess = preprocess_fn or self._default_preprocess
super().__init__(cache_prefix="cache_aft")
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_lower, y_upper = self._preprocess(df) # 注意:返回三个数组
# ✅ 正确方式:直接通过 input_data 传入区间标签
input_data(
data=X,
label=y_lower, # AFT 中 label 解释为下界(兼容旧版行为)
label_lower_bound=y_lower,
label_upper_bound=y_upper
)
self._it += 1
return 1
def reset(self):
self._it = 0
def _default_preprocess(self, df: pd.DataFrame):
# 假设 df 包含 't_lower', 't_upper', 'event' 等列
X = df.drop(columns=['t_lower', 't_upper', 'event'])
y_lower = df['t_lower'].values.astype(np.float32)
y_upper = df['t_upper'].values.astype(np.float32)
return X, y_lower, y_upper
⚠️ 重要注意事项:
- label_lower_bound 和 label_upper_bound 必须为一维 numpy.ndarray(dtype=float32 或 float64),长度与 data 行数严格一致;
- label 参数仍需提供(通常设为 label_lower_bound 的副本),否则 XGBoost 可能报错或行为未定义(尽管 AFT 实际使用区间而非单点标签);
- 所有 float info(如 label_lower_bound)必须在 input_data() 调用时一次性传入,不可在 next() 返回后补设;
- 若使用 xgboost.train(),需显式指定 objective='survival:aft' 和 aft_loss_distribution 等参数:
params = {
'objective': 'survival:aft',
'aft_loss_distribution': 'normal',
'aft_loss_distribution_scale': 1.0,
'tree_method': 'hist',
'max_depth': 6,
'learning_rate': 0.1
}
dtrain = xgboost.DMatrix(BatchedParquetIterator(train_files))
model = xgboost.train(params, dtrain, num_boost_round=100)
✅ 总结:外部内存 + AFT 的核心在于利用 input_data 的完整构造签名,将区间信息“随数据流同步注入”,而非事后修补。该方案完全兼容大规模 Parquet/CSV 分块读取场景,兼顾内存效率与生存分析建模需求。











